From 5b74f366d6d2f31cd5d6b097f0e7e8b1914a6dd9 Mon Sep 17 00:00:00 2001
From: "Kirill.Krylov"
Date: Mon, 6 Jan 2025 06:40:52 +0100
Subject: [PATCH] [ml-completion] connect ml completion ranking to new AB
mechanism
GitOrigin-RevId: 7a70bc2a04506ff148e9f3096fbc505e54f57af1
---
.../intellij.completionMlRanking.iml | 1 +
.../completion-ml-ranking/plugin-content.yaml | 4 +-
.../resources/META-INF/plugin.xml | 7 +-
.../resources/experiment.json | 161 --------------
.../MLRankingExperimentConnector.kt | 205 ++++++++++++++++++
.../experiments/MLRankingExperimentFetcher.kt | 16 ++
.../ml/experiments/MLRankingExperiments.kt | 104 +++++++++
7 files changed, 335 insertions(+), 163 deletions(-)
create mode 100644 plugins/completion-ml-ranking/src/com/intellij/completion/ml/experiments/MLRankingExperimentConnector.kt
create mode 100644 plugins/completion-ml-ranking/src/com/intellij/completion/ml/experiments/MLRankingExperimentFetcher.kt
create mode 100644 plugins/completion-ml-ranking/src/com/intellij/completion/ml/experiments/MLRankingExperiments.kt
diff --git a/plugins/completion-ml-ranking/intellij.completionMlRanking.iml b/plugins/completion-ml-ranking/intellij.completionMlRanking.iml
index 96ff2411b256..2cae3c451192 100644
--- a/plugins/completion-ml-ranking/intellij.completionMlRanking.iml
+++ b/plugins/completion-ml-ranking/intellij.completionMlRanking.iml
@@ -44,6 +44,7 @@
+
diff --git a/plugins/completion-ml-ranking/plugin-content.yaml b/plugins/completion-ml-ranking/plugin-content.yaml
index dbb9352ff9a5..1395971b061b 100644
--- a/plugins/completion-ml-ranking/plugin-content.yaml
+++ b/plugins/completion-ml-ranking/plugin-content.yaml
@@ -1,3 +1,5 @@
- name: lib/completionMlRanking.jar
modules:
- - name: intellij.completionMlRanking
\ No newline at end of file
+ - name: intellij.completionMlRanking
+ - name: intellij.completionMlRanking.experiments
+ - name: intellij.completionMlRanking.ultimate
\ No newline at end of file
diff --git a/plugins/completion-ml-ranking/resources/META-INF/plugin.xml b/plugins/completion-ml-ranking/resources/META-INF/plugin.xml
index e86cf46cabb5..94ea6570f48a 100644
--- a/plugins/completion-ml-ranking/resources/META-INF/plugin.xml
+++ b/plugins/completion-ml-ranking/resources/META-INF/plugin.xml
@@ -1,4 +1,4 @@
-
+
com.intellij.completion.ml.ranking
Machine Learning Code Completion
JetBrains
@@ -13,6 +13,10 @@
Editor | General | Code Completion | "Machine Learning Assistant Code Completion" section.
]]>
+
+
+
+
@@ -42,6 +46,7 @@
+
messages.MlCompletionBundle
diff --git a/plugins/completion-ml-ranking/resources/experiment.json b/plugins/completion-ml-ranking/resources/experiment.json
index 1b3675fa42b5..0d857cf360d3 100644
--- a/plugins/completion-ml-ranking/resources/experiment.json
+++ b/plugins/completion-ml-ranking/resources/experiment.json
@@ -44,20 +44,6 @@
"showArrows": false,
"calculateFeatures": true
},
- {
- "number": 15,
- "description": "Full Line - ONNX",
- "useMLRanking": true,
- "showArrows": false,
- "calculateFeatures": true
- },
- {
- "number": 16,
- "description": "Full Line - ONNX + trigger/filter model",
- "useMLRanking": true,
- "showArrows": false,
- "calculateFeatures": true
- },
{
"number": 17,
"description": "Completion performance (early lookup & MLRanking)",
@@ -85,153 +71,6 @@
"useMLRanking": true,
"showArrows": false,
"calculateFeatures": true
- },
- {
- "number": 23,
- "description": "Full Line - Llama 1536ctx drop initialization on type",
- "useMLRanking": true,
- "showArrows": false,
- "calculateFeatures": true
- },
- {
- "number": 24,
- "description": "Full Line - RAG",
- "useMLRanking": true,
- "showArrows": false,
- "calculateFeatures": true
- },
- {
- "number": 25,
- "description": "Full Line - disable stash by newline",
- "useMLRanking": true,
- "showArrows": false,
- "calculateFeatures": true
- },
- {
- "number": 26,
- "description": "Full Line - FIM",
- "useMLRanking": true,
- "showArrows": false,
- "calculateFeatures": true
- },
- {
- "number": 27,
- "description": "Full Line - FIM and RAG",
- "useMLRanking": true,
- "showArrows": false,
- "calculateFeatures": true
- },
- {
- "number": 28,
- "description": "Full Line - Llama + filter model",
- "useMLRanking": true,
- "showArrows": false,
- "calculateFeatures": true
- },
- {
- "number": 29,
- "description": "Full Line - 150ms debounce",
- "useMLRanking": true,
- "showArrows": false,
- "calculateFeatures": true
- },
- {
- "number": 30,
- "description": "Full Line - 50 ms debounce",
- "useMLRanking": true,
- "showArrows": false,
- "calculateFeatures": true
- },
- {
- "number": 31,
- "description": "Full Line - RAG + Filter Model",
- "useMLRanking": true,
- "showArrows": false,
- "calculateFeatures": true
- },
- {
- "number": 32,
- "description": "Full Line - FIM + Filter Model",
- "useMLRanking": true,
- "showArrows": false,
- "calculateFeatures": true
- },
- {
- "number": 33,
- "description": "Full Line - Llama + filter model + debounce 100",
- "useMLRanking": true,
- "showArrows": false,
- "calculateFeatures": true
- },
- {
- "number": 34,
- "description": "Full Line - PHP WA5 + Filter model",
- "useMLRanking": true,
- "showArrows": false,
- "calculateFeatures": true
- },
- {
- "number": 35,
- "description": "Full Line - PHP WA5",
- "useMLRanking": true,
- "showArrows": false,
- "calculateFeatures": true
- },
- {
- "number": 36,
- "description": "Full Line - Filter simplified",
- "useMLRanking": true,
- "showArrows": false,
- "calculateFeatures": true
- },
- {
- "number": 37,
- "description": "Full Line - FIM + 50 ms debounce",
- "useMLRanking": true,
- "showArrows": false,
- "calculateFeatures": true
- },
- {
- "number": 38,
- "description": "Full Line - Llama + filter model V2",
- "useMLRanking": true,
- "showArrows": false,
- "calculateFeatures": true
- },
- {
- "number": 39,
- "description": "Full Line - PHP WA5 + Filter model V2",
- "useMLRanking": true,
- "showArrows": false,
- "calculateFeatures": true
- },
- {
- "number": 40,
- "description": "Full Line - RAG + filter model V2",
- "useMLRanking": true,
- "showArrows": false,
- "calculateFeatures": true
- },
- {
- "number": 41,
- "description": "Full Line - Rust filtered data",
- "useMLRanking": true,
- "showArrows": false,
- "calculateFeatures": true
- },
- {
- "number": 42,
- "description": "Full Line - Rust filtered data + old filter model",
- "useMLRanking": true,
- "showArrows": false,
- "calculateFeatures": true
- },
- {
- "number": 43,
- "description": "Full Line - Rust filtered data + new filter model",
- "useMLRanking": true,
- "showArrows": false,
- "calculateFeatures": true
}
],
"languages": [
diff --git a/plugins/completion-ml-ranking/src/com/intellij/completion/ml/experiments/MLRankingExperimentConnector.kt b/plugins/completion-ml-ranking/src/com/intellij/completion/ml/experiments/MLRankingExperimentConnector.kt
new file mode 100644
index 000000000000..ad69fbd73b3a
--- /dev/null
+++ b/plugins/completion-ml-ranking/src/com/intellij/completion/ml/experiments/MLRankingExperimentConnector.kt
@@ -0,0 +1,205 @@
+@file:Suppress("PublicApiImplicitType")
+
+package com.intellij.completion.ml.experiments
+
+import com.intellij.lang.Language
+
+@Suppress("PublicApiImplicitType")
+interface MLRankingExperimentConnector {
+ fun connect(languageId: String): List
+
+ companion object {
+ fun getExperiments(language: Language): List? {
+ return HardcodedMLRankingExperimentConnector::class.sealedSubclasses
+ .firstOrNull { checkMatchingLanguage(language, it.objectInstance?.language) }
+ ?.objectInstance?.connect(language.id)
+ }
+
+ private fun checkMatchingLanguage(language: Language, languageId: String?): Boolean {
+ if (languageId == null) return false
+
+ val baseLanguages = Language.getRegisteredLanguages().filter { language.isKindOf(it) }
+ return baseLanguages.any { languageId.equals(it.id, ignoreCase = true) }
+ }
+ }
+}
+
+sealed class HardcodedMLRankingExperimentConnector(val language: String) : MLRankingExperimentConnector {
+ companion object {
+ object Java : HardcodedMLRankingExperimentConnector("java") {
+ override fun connect(languageId: String) = experiments(languageId) {
+ eap {
+ groups = listOf(
+ CommonExperiments.ControlA,
+ CommonExperiments.ControlB,
+ )
+ }
+ }
+ }
+
+ object Python : HardcodedMLRankingExperimentConnector("python") {
+ override fun connect(languageId: String) = experiments(languageId) {
+ eap {
+ groups = listOf(
+ CommonExperiments.ControlA,
+ CommonExperiments.ControlB,
+ )
+ }
+ }
+ }
+
+ object Kotlin : HardcodedMLRankingExperimentConnector("kotlin") {
+ override fun connect(languageId: String) = experiments(languageId) {
+ eap {
+ groups = listOf(
+ CommonExperiments.ControlA,
+ CommonExperiments.ControlB,
+ )
+ }
+ }
+ }
+
+ object Scala : HardcodedMLRankingExperimentConnector("scala") {
+ override fun connect(languageId: String) = experiments(languageId) {
+ eap {
+ groups = listOf(
+ CommonExperiments.ControlA,
+ CommonExperiments.ControlB,
+ )
+ }
+ }
+ }
+
+ object Php : HardcodedMLRankingExperimentConnector("php") {
+ override fun connect(languageId: String) = experiments(languageId) {
+ eap {
+ groups = listOf(
+ CommonExperiments.ControlA,
+ CommonExperiments.ControlB,
+ )
+ }
+ }
+ }
+
+ object Javascript : HardcodedMLRankingExperimentConnector("javascript") {
+ override fun connect(languageId: String) = experiments(languageId) {
+ eap {
+ groups = listOf(
+ CommonExperiments.ControlA,
+ CommonExperiments.ControlB,
+ )
+ }
+ }
+ }
+
+ object Ruby : HardcodedMLRankingExperimentConnector("ruby") {
+ override fun connect(languageId: String) = experiments(languageId) {
+ eap {
+ groups = listOf(
+ CommonExperiments.ControlA,
+ CommonExperiments.ControlB,
+ )
+ }
+ }
+ }
+
+ object Go : HardcodedMLRankingExperimentConnector("go") {
+ override fun connect(languageId: String) = experiments(languageId) {
+ eap {
+ groups = listOf(
+ CommonExperiments.ControlA,
+ CommonExperiments.ControlB,
+ )
+ }
+ }
+ }
+
+ object Rust : HardcodedMLRankingExperimentConnector("rust") {
+ override fun connect(languageId: String) = experiments(languageId) {
+ eap {
+ groups = listOf(
+ CommonExperiments.ControlA,
+ CommonExperiments.ControlB,
+ )
+ }
+ }
+ }
+
+ object Swift : HardcodedMLRankingExperimentConnector("swift") {
+ override fun connect(languageId: String) = experiments(languageId) {
+ eap {
+ groups = listOf(
+ CommonExperiments.ControlA,
+ CommonExperiments.ControlB,
+ )
+ }
+ }
+ }
+
+ object Mysql : HardcodedMLRankingExperimentConnector("mysql") {
+ override fun connect(languageId: String) = experiments(languageId) {
+ eap {
+ groups = listOf(
+ CommonExperiments.ControlA,
+ CommonExperiments.ControlB,
+ )
+ }
+ }
+ }
+
+ object Csharp : HardcodedMLRankingExperimentConnector("C#") {
+ override fun connect(languageId: String) = experiments(languageId) {
+ eap {
+ groups = listOf(
+ CommonExperiments.ControlA,
+ CommonExperiments.ControlB,
+ )
+ }
+ }
+ }
+
+ object ObjectiveC : HardcodedMLRankingExperimentConnector("objectivec") {
+ override fun connect(languageId: String) = experiments(languageId) {
+ eap {
+ groups = listOf(
+ CommonExperiments.ControlA,
+ CommonExperiments.ControlB,
+ )
+ }
+ }
+ }
+
+ object Html : HardcodedMLRankingExperimentConnector("html") {
+ override fun connect(languageId: String) = experiments(languageId) {
+ eap {
+ groups = listOf(
+ CommonExperiments.ControlA,
+ CommonExperiments.ControlB,
+ )
+ }
+ }
+ }
+
+ object Css : HardcodedMLRankingExperimentConnector("css") {
+ override fun connect(languageId: String) = experiments(languageId) {
+ eap {
+ groups = listOf(
+ CommonExperiments.ControlA,
+ CommonExperiments.ControlB,
+ )
+ }
+ }
+ }
+
+ object ShellScript : HardcodedMLRankingExperimentConnector("shell script") {
+ override fun connect(languageId: String) = experiments(languageId) {
+ eap {
+ groups = listOf(
+ CommonExperiments.ControlA,
+ CommonExperiments.ControlB,
+ )
+ }
+ }
+ }
+ }
+}
diff --git a/plugins/completion-ml-ranking/src/com/intellij/completion/ml/experiments/MLRankingExperimentFetcher.kt b/plugins/completion-ml-ranking/src/com/intellij/completion/ml/experiments/MLRankingExperimentFetcher.kt
new file mode 100644
index 000000000000..a56cb2253a8c
--- /dev/null
+++ b/plugins/completion-ml-ranking/src/com/intellij/completion/ml/experiments/MLRankingExperimentFetcher.kt
@@ -0,0 +1,16 @@
+package com.intellij.completion.ml.experiments
+
+import com.intellij.lang.Language
+import com.intellij.openapi.extensions.ExtensionPointName
+
+@Suppress("PublicApiImplicitType")
+interface MLRankingExperimentFetcher {
+ fun getExperimentGroup(language: Language): MLRankingExperiment?
+ fun getExperimentNumber(language: Language): Int? = getExperimentGroup(language)?.id
+
+ companion object {
+ val EP_NAME = ExtensionPointName("com.intellij.completion.ml.experimentFetcher")
+
+ fun getInstance() = EP_NAME.extensionList.firstOrNull()
+ }
+}
diff --git a/plugins/completion-ml-ranking/src/com/intellij/completion/ml/experiments/MLRankingExperiments.kt b/plugins/completion-ml-ranking/src/com/intellij/completion/ml/experiments/MLRankingExperiments.kt
new file mode 100644
index 000000000000..3e4497c4485d
--- /dev/null
+++ b/plugins/completion-ml-ranking/src/com/intellij/completion/ml/experiments/MLRankingExperiments.kt
@@ -0,0 +1,104 @@
+package com.intellij.completion.ml.experiments
+
+interface MLRankingExperiment {
+ val id: Int
+
+ val useMLRanking: Boolean?
+ val showArrows: Boolean?
+ val calculateFeatures: Boolean?
+
+ val targetLanguageId: String?
+
+ companion object {
+ const val SEED: Long = 111L
+ const val NO_EXP: Int = -1
+ }
+}
+
+data class RawMLRankingExperimentData(
+ val isEAP: Boolean,
+ val userFraction: Double,
+ val releaseVersion: String,
+ val groups: List,
+ /** Must not be filled manually! */
+ val languageId: String = "",
+)
+
+class ExperimentDataHolder() {
+ private val _experiments: MutableList = mutableListOf()
+ val experiments: List
+ get() = _experiments.toList()
+
+ fun updateLanguageIds(newLanguageId: String) {
+ _experiments.replaceAll { it.copy(languageId = newLanguageId) }
+ }
+
+ fun addExperiment(data: RawMLRankingExperimentData) {
+ _experiments += data
+ }
+}
+
+class ExperimentDataForRelease(var releaseVersion: String) {
+ var userFraction: Double = 0.0
+ lateinit var groups: List
+
+ fun toExperimentData(): RawMLRankingExperimentData = RawMLRankingExperimentData(
+ isEAP = false,
+ userFraction = userFraction,
+ groups = groups,
+ releaseVersion = releaseVersion,
+ )
+}
+
+class ExperimentDataForEAP {
+ lateinit var groups: List
+
+ fun toExperimentData(): RawMLRankingExperimentData = RawMLRankingExperimentData(
+ isEAP = true,
+ userFraction = 1.0,
+ groups = groups,
+ releaseVersion = "EAP",
+ )
+}
+
+fun MLRankingExperimentConnector.experiments(languageId: String, block: ExperimentDataHolder.() -> Unit): List {
+ return ExperimentDataHolder().apply(block).also { it.updateLanguageIds(languageId) }.experiments
+}
+
+fun ExperimentDataHolder.release(version: String, block: ExperimentDataForRelease.() -> Unit): ExperimentDataForRelease {
+ return ExperimentDataForRelease(releaseVersion = version).apply(block).also {
+ addExperiment(it.toExperimentData())
+ }
+}
+
+fun ExperimentDataHolder.eap(block: ExperimentDataForEAP.() -> Unit): ExperimentDataForEAP {
+ return ExperimentDataForEAP().apply(block).also {
+ addExperiment(it.toExperimentData())
+ }
+}
+
+enum class CommonExperiments(
+ override val id: Int,
+ override val useMLRanking: Boolean? = null,
+ override val showArrows: Boolean? = null,
+ override val calculateFeatures: Boolean? = null,
+) : MLRankingExperiment {
+ NoExperiment(id = -1),
+
+ StandardCompletion(id = 1, useMLRanking = false, showArrows = false, calculateFeatures = false),
+ DefaultMLModelWithoutArrows(id = 8, useMLRanking = true, showArrows = false, calculateFeatures = true),
+ DefaultMLModelWithArrows(id = 11, useMLRanking = true, showArrows = true, calculateFeatures = true),
+ StandardCompletionNoFeatures(id = 12, useMLRanking = false, showArrows = false, calculateFeatures = false),
+ FirstExperimentalModel(id = 13, useMLRanking = true, showArrows = false, calculateFeatures = true),
+ SecondExperimentalModel(id = 14, useMLRanking = true, showArrows = false, calculateFeatures = true),
+ CompletionPerformanceEarlyML(id = 17, useMLRanking = true, showArrows = false, calculateFeatures = true),
+ CompletionPerformanceEarlyLookup(id = 18, useMLRanking = false, showArrows = false, calculateFeatures = true),
+ CompletionPerformance(id = 19, useMLRanking = false, showArrows = false, calculateFeatures = true),
+ FullLineLlama1536ctx(id = 22, useMLRanking = false, showArrows = false, calculateFeatures = true),
+
+ ControlA(id = 50),
+ ControlB(id = 51),
+ ;
+
+ override val targetLanguageId: String? = null
+}