[ml-completion] connect ml completion ranking to new AB mechanism

GitOrigin-RevId: 7a70bc2a04506ff148e9f3096fbc505e54f57af1
This commit is contained in:
Kirill.Krylov
2025-01-06 07:40:18 +00:00
committed by intellij-monorepo-bot
parent 51985dbbd7
commit 5b74f366d6
7 changed files with 335 additions and 163 deletions
@@ -44,6 +44,7 @@
<orderEntry type="library" name="jackson" level="project" />
<orderEntry type="library" name="kotlinx-serialization-json" level="project" />
<orderEntry type="library" name="kotlinx-serialization-core" level="project" />
<orderEntry type="library" name="kotlin-reflect" level="project" />
<orderEntry type="module" module-name="intellij.platform.util.jdom" />
<orderEntry type="module" module-name="intellij.platform.ml" />
<orderEntry type="library" name="opentelemetry" level="project" />
@@ -1,3 +1,5 @@
- name: lib/completionMlRanking.jar
modules:
- name: intellij.completionMlRanking
- name: intellij.completionMlRanking
- name: intellij.completionMlRanking.experiments
- name: intellij.completionMlRanking.ultimate
@@ -1,4 +1,4 @@
<idea-plugin package="com.intellij.completion.ml">
<idea-plugin xmlns:xi="http://www.w3.org/2001/XInclude" package="com.intellij.completion.ml">
<id>com.intellij.completion.ml.ranking</id>
<name>Machine Learning Code Completion</name>
<vendor>JetBrains</vendor>
@@ -13,6 +13,10 @@
Editor | General | Code Completion | "Machine Learning Assistant Code Completion" section.</p>
]]></description>
<xi:include href="/META-INF/ml-ranking-ultimate.xml">
<xi:fallback/>
</xi:include>
<dependencies>
<plugin id="com.intellij.modules.lang"/>
<module name="intellij.platform.vcs.impl"/>
@@ -42,6 +46,7 @@
<extensionPoint qualifiedName="com.intellij.completion.ml.featuresOverride" beanClass="com.intellij.lang.LanguageExtensionPoint" dynamic="true">
<with attribute="implementationClass" implements="com.intellij.completion.ml.features.RankingFeaturesOverrides"/>
</extensionPoint>
<extensionPoint qualifiedName="com.intellij.completion.ml.experimentFetcher" interface="com.intellij.completion.ml.experiments.MLRankingExperimentFetcher" dynamic="true"/>
</extensionPoints>
<resource-bundle>messages.MlCompletionBundle</resource-bundle>
@@ -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": [
@@ -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<RawMLRankingExperimentData>
companion object {
fun getExperiments(language: Language): List<RawMLRankingExperimentData>? {
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,
)
}
}
}
}
}
@@ -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<MLRankingExperimentFetcher>("com.intellij.completion.ml.experimentFetcher")
fun getInstance() = EP_NAME.extensionList.firstOrNull()
}
}
@@ -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<MLRankingExperiment>,
/** Must not be filled manually! */
val languageId: String = "",
)
class ExperimentDataHolder() {
private val _experiments: MutableList<RawMLRankingExperimentData> = mutableListOf()
val experiments: List<RawMLRankingExperimentData>
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<MLRankingExperiment>
fun toExperimentData(): RawMLRankingExperimentData = RawMLRankingExperimentData(
isEAP = false,
userFraction = userFraction,
groups = groups,
releaseVersion = releaseVersion,
)
}
class ExperimentDataForEAP {
lateinit var groups: List<MLRankingExperiment>
fun toExperimentData(): RawMLRankingExperimentData = RawMLRankingExperimentData(
isEAP = true,
userFraction = 1.0,
groups = groups,
releaseVersion = "EAP",
)
}
fun MLRankingExperimentConnector.experiments(languageId: String, block: ExperimentDataHolder.() -> Unit): List<RawMLRankingExperimentData> {
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
}