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 +}