Revert "MLP-33 IJ Imports Ranking"

This reverts commit 47140fd0301283a10966e14c65df9a08d128ec39.

GitOrigin-RevId: 81fa4946bd76885bcc599dfae9bfb9feaa11fbb3
This commit is contained in:
Andrey Vokin
2024-12-09 21:19:47 +00:00
committed by intellij-monorepo-bot
parent 313c1619d2
commit 44b9d24ae8
56 changed files with 888 additions and 883 deletions
@@ -1,18 +0,0 @@
<component name="libraryTable">
<library name="jetbrains.ml.models.python.imports.ranking.model" type="repository">
<properties include-transitive-deps="false" maven-id="com.jetbrains.ml.models:python-imports-ranking-model:lilac-coua">
<verification>
<artifact url="file://$MAVEN_REPOSITORY$/com/jetbrains/ml/models/python-imports-ranking-model/lilac-coua/python-imports-ranking-model-lilac-coua.jar">
<sha256sum>0516e6383101eb995b8f4d9c7dc217578bbc7467a806895d5c4e2390c27bba7d</sha256sum>
</artifact>
</verification>
</properties>
<CLASSES>
<root url="jar://$MAVEN_REPOSITORY$/com/jetbrains/ml/models/python-imports-ranking-model/lilac-coua/python-imports-ranking-model-lilac-coua.jar!/" />
</CLASSES>
<JAVADOC />
<SOURCES>
<root url="jar://$MAVEN_REPOSITORY$/com/jetbrains/ml/models/python-imports-ranking-model/lilac-coua/python-imports-ranking-model-lilac-coua-sources.jar!/" />
</SOURCES>
</library>
</component>
@@ -1,16 +0,0 @@
<component name="libraryTable">
<library name="jetbrains.mlapi.catboost.shadow.need.slf4j" type="repository">
<properties include-transitive-deps="false" maven-id="com.jetbrains.mlapi:catboost-shadow-need-slf4j:1.2.5">
<verification>
<artifact url="file://$MAVEN_REPOSITORY$/com/jetbrains/mlapi/catboost-shadow-need-slf4j/1.2.5/catboost-shadow-need-slf4j-1.2.5.jar">
<sha256sum>94a06c3fdd4aa638df59c8c783a379ca43a5b933187760e219e09e84f1633d76</sha256sum>
</artifact>
</verification>
</properties>
<CLASSES>
<root url="jar://$MAVEN_REPOSITORY$/com/jetbrains/mlapi/catboost-shadow-need-slf4j/1.2.5/catboost-shadow-need-slf4j-1.2.5.jar!/" />
</CLASSES>
<JAVADOC />
<SOURCES />
</library>
</component>
+20
View File
@@ -0,0 +1,20 @@
<component name="libraryTable">
<library name="jetbrains.mlapi.extension" type="repository">
<properties include-transitive-deps="false" maven-id="com.jetbrains.mlapi:extension:34">
<verification>
<artifact url="file://$MAVEN_REPOSITORY$/com/jetbrains/mlapi/extension/34/extension-34.jar">
<sha256sum>1d18aa7a0d891cab230ddb67e3d1174830e1472c7373c37296c23a3905040a7d</sha256sum>
</artifact>
</verification>
</properties>
<CLASSES>
<root url="jar://$MAVEN_REPOSITORY$/com/jetbrains/mlapi/extension/34/extension-34.jar!/" />
</CLASSES>
<JAVADOC>
<root url="jar://$MAVEN_REPOSITORY$/com/jetbrains/mlapi/extension/34/extension-34-javadoc.jar!/" />
</JAVADOC>
<SOURCES>
<root url="jar://$MAVEN_REPOSITORY$/com/jetbrains/mlapi/extension/34/extension-34-sources.jar!/" />
</SOURCES>
</library>
</component>
-20
View File
@@ -1,20 +0,0 @@
<component name="libraryTable">
<library name="jetbrains.mlapi.ml.building.blocks" type="repository">
<properties include-transitive-deps="false" maven-id="com.jetbrains.mlapi:ml-building-blocks:58">
<verification>
<artifact url="file://$MAVEN_REPOSITORY$/com/jetbrains/mlapi/ml-building-blocks/58/ml-building-blocks-58.jar">
<sha256sum>65b9f465f00474bcbfff5949458631a92c6a4cf0587f05fac4d0acd65cfae4cf</sha256sum>
</artifact>
</verification>
</properties>
<CLASSES>
<root url="jar://$MAVEN_REPOSITORY$/com/jetbrains/mlapi/ml-building-blocks/58/ml-building-blocks-58.jar!/" />
</CLASSES>
<JAVADOC>
<root url="jar://$MAVEN_REPOSITORY$/com/jetbrains/mlapi/ml-building-blocks/58/ml-building-blocks-58-javadoc.jar!/" />
</JAVADOC>
<SOURCES>
<root url="jar://$MAVEN_REPOSITORY$/com/jetbrains/mlapi/ml-building-blocks/58/ml-building-blocks-58-sources.jar!/" />
</SOURCES>
</library>
</component>
-18
View File
@@ -1,18 +0,0 @@
<component name="libraryTable">
<library name="jetbrains.mlapi.ml.feature.api" type="repository">
<properties include-transitive-deps="false" maven-id="com.jetbrains.mlapi:ml-feature-api:58">
<verification>
<artifact url="file://$MAVEN_REPOSITORY$/com/jetbrains/mlapi/ml-feature-api/58/ml-feature-api-58.jar">
<sha256sum>8f0300408884fe7ea8b3665f4bc2e7f725cda763806ae92677cf141f4d781193</sha256sum>
</artifact>
</verification>
</properties>
<CLASSES>
<root url="jar://$MAVEN_REPOSITORY$/com/jetbrains/mlapi/ml-feature-api/58/ml-feature-api-58.jar!/" />
</CLASSES>
<JAVADOC />
<SOURCES>
<root url="jar://$MAVEN_REPOSITORY$/com/jetbrains/mlapi/ml-feature-api/58/ml-feature-api-58-sources.jar!/" />
</SOURCES>
</library>
</component>
+20
View File
@@ -0,0 +1,20 @@
<component name="libraryTable">
<library name="jetbrains.mlapi.usage" type="repository">
<properties include-transitive-deps="false" maven-id="com.jetbrains.mlapi:usage:34">
<verification>
<artifact url="file://$MAVEN_REPOSITORY$/com/jetbrains/mlapi/usage/34/usage-34.jar">
<sha256sum>b28271a408b89f4a9e00e4a849543b5e9206333472ee379027d969e5d1b42e6f</sha256sum>
</artifact>
</verification>
</properties>
<CLASSES>
<root url="jar://$MAVEN_REPOSITORY$/com/jetbrains/mlapi/usage/34/usage-34.jar!/" />
</CLASSES>
<JAVADOC>
<root url="jar://$MAVEN_REPOSITORY$/com/jetbrains/mlapi/usage/34/usage-34-javadoc.jar!/" />
</JAVADOC>
<SOURCES>
<root url="jar://$MAVEN_REPOSITORY$/com/jetbrains/mlapi/usage/34/usage-34-sources.jar!/" />
</SOURCES>
</library>
</component>
-1
View File
@@ -898,7 +898,6 @@
<module fileurl="file://$PROJECT_DIR$/python/helpers/tests/intellij.python.helpers.tests.iml" filepath="$PROJECT_DIR$/python/helpers/tests/intellij.python.helpers.tests.iml" />
<module fileurl="file://$PROJECT_DIR$/python/IntelliLang-python/intellij.python.langInjection.iml" filepath="$PROJECT_DIR$/python/IntelliLang-python/intellij.python.langInjection.iml" />
<module fileurl="file://$PROJECT_DIR$/python/python-markdown/intellij.python.markdown.iml" filepath="$PROJECT_DIR$/python/python-markdown/intellij.python.markdown.iml" />
<module fileurl="file://$PROJECT_DIR$/python/intellij.python.ml.features/intellij.python.ml.features.iml" filepath="$PROJECT_DIR$/python/intellij.python.ml.features/intellij.python.ml.features.iml" />
<module fileurl="file://$PROJECT_DIR$/python/python-parser/intellij.python.parser.iml" filepath="$PROJECT_DIR$/python/python-parser/intellij.python.parser.iml" />
<module fileurl="file://$PROJECT_DIR$/python/python-psi-api/intellij.python.psi.iml" filepath="$PROJECT_DIR$/python/python-psi-api/intellij.python.psi.iml" />
<module fileurl="file://$PROJECT_DIR$/python/python-psi-impl/intellij.python.psi.impl.iml" filepath="$PROJECT_DIR$/python/python-psi-impl/intellij.python.psi.impl.iml" />
@@ -1323,9 +1323,8 @@ object CommunityLibraryLicenses {
jetbrainsLibrary("intellij.remoterobot.remote.fixtures"),
jetbrainsLibrary("intellij.remoterobot.robot.server.core"),
jetbrainsLibrary("jetbrains.intellij.deps.rwmutex.idea"),
jetbrainsLibrary("jetbrains.ml.models.python.imports.ranking.model"),
jetbrainsLibrary("jetbrains.mlapi.ml.building.blocks"),
jetbrainsLibrary("jetbrains.mlapi.ml.feature.api"),
jetbrainsLibrary("jetbrains.mlapi.extension"),
jetbrainsLibrary("jetbrains.mlapi.usage"),
jetbrainsLibrary("jshell-frontend"),
jetbrainsLibrary("jvm-native-trusted-roots"),
jetbrainsLibrary("kotlin-gradle-plugin-idea"),
@@ -30,7 +30,6 @@ object PythonCommunityPluginModules {
"intellij.python.pydev",
"intellij.python.sdk",
"intellij.python.terminal",
"intellij.python.ml.features",
)
/**
+1 -1
View File
@@ -11,6 +11,6 @@
<orderEntry type="library" name="kotlin-stdlib" level="project" />
<orderEntry type="library" name="kotlinx-coroutines-core" level="project" />
<orderEntry type="module" module-name="intellij.platform.core" />
<orderEntry type="library" exported="" name="jetbrains.mlapi.ml.feature.api" level="project" />
<orderEntry type="library" exported="" name="jetbrains.mlapi.extension" level="project" />
</component>
</module>
@@ -50,7 +50,6 @@
<orderEntry type="module" module-name="intellij.platform.ml" />
<orderEntry type="module" module-name="intellij.platform.statistics" />
<orderEntry type="library" name="kotlin-reflect" level="project" />
<orderEntry type="library" exported="" name="jetbrains.mlapi.ml.building.blocks" level="project" />
<orderEntry type="library" exported="" name="jetbrains.mlapi.catboost.shadow.need.slf4j" level="project" />
<orderEntry type="library" exported="" name="jetbrains.mlapi.usage" level="project" />
</component>
</module>
@@ -1,134 +0,0 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.intellij.platform.ml.impl
import com.intellij.openapi.diagnostic.thisLogger
import com.intellij.openapi.util.registry.Registry
import com.jetbrains.ml.building.blocks.model.MLModel
import com.jetbrains.ml.building.blocks.model.MLModelLoader
import com.jetbrains.ml.building.blocks.task.MLTask
import com.jetbrains.ml.features.api.MLUnit
import com.jetbrains.ml.features.api.feature.FeatureDeclaration
import kotlinx.coroutines.CoroutineScope
import kotlinx.coroutines.Dispatchers
import kotlinx.coroutines.async
import org.jetbrains.annotations.ApiStatus
import java.nio.file.Path
import java.util.concurrent.atomic.AtomicReference
import kotlin.io.path.exists
/**
* A utility class that orchestrates ML models for a machine learning task.
*
* It is aware of a default [defaultLoader] and when [modelPathRegistryKey]
* is not empty, it loads a model from the corresponding path
*/
@ApiStatus.Internal
open class MLModelSelector<M : MLModel<P>, P : Any>(
private val defaultLoader: MLModelLoader<M, P>,
registryModelLoaderFactory: (Path) -> MLModelLoader<M, P>,
private val modelPathRegistryKey: String,
private val coroutineScope: CoroutineScope,
) {
private val isUsingRegistry by lazy { Registry.`is`(modelPathRegistryKey) }
private val registryLoader = RegistryModelLoader(modelPathRegistryKey, defaultLoader, registryModelLoaderFactory)
private val modelLoader by lazy {
if (isUsingRegistry) registryLoader else defaultLoader
}
private var state: AtomicReference<State<M, P>> = AtomicReference(State.NotLoaded)
/**
* Starts loading the ML model if it has not been loaded yet.
* Call this function before calling [getModel] to make it
* return not null on the first call, solving the problem of
* cold start.
*/
fun prepareModel(
features: Map<MLUnit<*>, List<FeatureDeclaration<*>>>,
parameters: Map<String, Any>? = null,
) {
when (state.get()) {
is State.Loaded -> return
is State.Loading -> return
is State.NotLoaded -> startLoading(features, parameters)
}
}
fun prepareModel(task: MLTask, parameters: Map<String, Any>? = null) {
prepareModel(task.treeFeaturesFlattened, parameters)
}
/**
* Returns an instance of an ML model if it has been loaded. Null otherwise.
* The model starts to load during the first
*
* If registry value of [modelPathRegistryKey] is not empty, the model is loaded from the specified path
*/
fun getModel(
features: Map<MLUnit<*>, List<FeatureDeclaration<*>>>,
parameters: Map<String, Any>? = null,
): M? {
return when (val s = state.get()) {
is State.Loaded -> s.model
is State.Loading -> null
is State.NotLoaded -> {
startLoading(features, parameters)
null
}
}
}
fun getModel(
task: MLTask,
parameters: Map<String, Any>? = null,
): M? {
return getModel(task.treeFeaturesFlattened, parameters)
}
private fun startLoading(
features: Map<MLUnit<*>, List<FeatureDeclaration<*>>>,
parameters: Map<String, Any>?,
) {
if (!state.compareAndSet(State.NotLoaded, State.Loading)) return
coroutineScope.async(Dispatchers.IO) {
try {
val model = modelLoader.loadModel(features, parameters)
assert(state.compareAndSet(State.Loading, State.Loaded(model)))
} catch (e: Throwable) {
thisLogger().error("Unable to load ML model", e)
state.compareAndSet(State.Loading, State.NotLoaded)
}
}
}
private sealed interface State<in M : MLModel<@UnsafeVariance P>, in P : Any> {
data object NotLoaded : State<MLModel<Any>, Any>
data object Loading : State<MLModel<Any>, Any>
class Loaded<M : MLModel<P>, P : Any>(val model: M) : State<M, P>
}
}
private class RegistryModelLoader<M : MLModel<P>, P : Any>(
private val registryKey: String,
defaultLoader: MLModelLoader<M, P>,
private val registryModelLoaderFactory: (Path) -> MLModelLoader<M, P>,
) : MLModelLoader<M, P> {
override val predictionClass: Class<P> = defaultLoader.predictionClass
override suspend fun loadModel(features: Map<MLUnit<*>, List<FeatureDeclaration<*>>>, parameters: Map<String, Any>?): M {
val key = Registry.get(registryKey)
require(key.isRestartRequired()) {
"Registry key `$registryKey` must require restart"
}
val path = key.asString().takeUnless { it.isEmpty() }?.let { Path.of(it) }
requireNotNull(path) { "Registry key `$registryKey` must contain non-empty string" }
require(path.isAbsolute) { "Registry key `$registryKey` must be an absolute path" }
require(path.exists()) { "Registry key `$registryKey` be an existing path" }
val loader = registryModelLoaderFactory(path)
return loader.loadModel(features, parameters)
}
}
private val MLTask.treeFeaturesFlattened: Map<MLUnit<*>, List<FeatureDeclaration<*>>>
get() = treeFeatures.flatMap { it.entries }.associateBy({ it.key }, { it.value })
@@ -0,0 +1,90 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.intellij.platform.ml.impl.tools
import com.intellij.openapi.components.Service
import com.intellij.openapi.diagnostic.Logger
import com.intellij.openapi.diagnostic.debug
import com.intellij.util.application
import com.intellij.util.messages.Topic
import com.jetbrains.ml.platform.MLApiPlatform.ExtensionController
import org.jetbrains.annotations.ApiStatus
@ApiStatus.Internal
@Service(Service.Level.APP)
class IJPlatform : com.jetbrains.ml.platform.MLApiPlatform(
featureProviders = EP_NAME_FEATURE_PROVIDER.extensionList,
mlUnitProviders = emptyList(),
) {
override val taskListeners: Map<String, List<com.jetbrains.ml.monitoring.MLTaskListenerTyped<*, *>>>
get() = KeyedMessagingProvider.collect(MLTaskListenerTyped.TOPIC)
override fun addTaskListener(taskId: String, taskListener: com.jetbrains.ml.monitoring.MLTaskListenerTyped<*, *>): com.jetbrains.ml.platform.MLApiPlatform.ExtensionController {
val connection = application.messageBus.connect()
fun <M : com.jetbrains.ml.model.MLModel<P>, P : Any> capturingType(taskListenerTyped: com.jetbrains.ml.monitoring.MLTaskListenerTyped<M, P>) {
connection.subscribe(MLTaskListenerTyped.TOPIC, object : MessageBusMLTaskListenerProvider<M, P> {
override fun provide(collector: (com.jetbrains.ml.monitoring.MLTaskListenerTyped<M, P>, String) -> Unit) = collector(taskListenerTyped, taskId)
})
}
capturingType(taskListener)
return ExtensionController { connection.disconnect() }
}
override val loggingListeners: Map<String, List<com.jetbrains.ml.monitoring.MLTaskLoggingListener>>
get() = KeyedMessagingProvider.collect(MLTaskLoggingListener.TOPIC)
override fun addLoggingListener(taskId: String, loggingListener: com.jetbrains.ml.monitoring.MLTaskLoggingListener): ExtensionController {
val connection = application.messageBus.connect()
connection.subscribe(MLTaskLoggingListener.TOPIC, object : MessageBusMLTaskLoggingListenerProvider {
override fun provide(collector: (com.jetbrains.ml.monitoring.MLTaskLoggingListener, String) -> Unit) = collector(loggingListener, taskId)
})
return ExtensionController { connection.disconnect() }
}
override val systemLoggerBuilder: com.jetbrains.ml.platform.SystemLoggerBuilder = object : com.jetbrains.ml.platform.SystemLoggerBuilder {
override fun build(clazz: Class<*>): com.jetbrains.ml.platform.SystemLogger {
return IJSystemLogger(Logger.getInstance(clazz))
}
override fun build(name: String): com.jetbrains.ml.platform.SystemLogger {
return IJSystemLogger(Logger.getInstance(name))
}
private inner class IJSystemLogger(private val baseLogger: Logger) : com.jetbrains.ml.platform.SystemLogger {
override fun info(data: () -> String) = baseLogger.info(data())
override fun warn(data: () -> String) = baseLogger.warn(data())
override fun debug(data: () -> String) = baseLogger.debug { data() }
override fun error(e: Throwable) = baseLogger.error(e)
}
}
}
@ApiStatus.Internal
sealed interface KeyedMessagingProvider<T> {
fun provide(collector: (T, String) -> Unit)
companion object {
fun <P : KeyedMessagingProvider<*>, T> collect(topic: Topic<P>): Map<String, List<T>> {
val collected = mutableMapOf<String, MutableList<T>>()
application.messageBus.syncPublisher(topic).provide { it, keyOfIt ->
@Suppress("UNCHECKED_CAST")
collected.getOrPut(keyOfIt) { mutableListOf() }.add(it as T)
}
return collected
}
}
}
@ApiStatus.Internal
interface MessageBusMLTaskListenerProvider<M : com.jetbrains.ml.model.MLModel<P>, P : Any> : KeyedMessagingProvider<com.jetbrains.ml.monitoring.MLTaskListenerTyped<M, P>>
@ApiStatus.Internal
interface MessageBusMLTaskLoggingListenerProvider : KeyedMessagingProvider<com.jetbrains.ml.monitoring.MLTaskLoggingListener>
@@ -0,0 +1,7 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.intellij.platform.ml.impl.tools
import com.intellij.openapi.extensions.ExtensionPointName
import com.jetbrains.ml.FeatureProvider
internal val EP_NAME_FEATURE_PROVIDER: ExtensionPointName<FeatureProvider> = ExtensionPointName.create("com.intellij.platform.ml.featureProvider")
@@ -0,0 +1,31 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.intellij.platform.ml.impl.tools
import com.intellij.util.messages.Topic
import org.jetbrains.annotations.ApiStatus
@ApiStatus.Internal
interface MLTaskListenerTyped<M : com.jetbrains.ml.model.MLModel<P>, P : Any> : com.jetbrains.ml.monitoring.MLTaskListenerTyped<M, P>, MessageBusMLTaskListenerProvider<M, P> {
val task: com.jetbrains.ml.MLTask<M, P>
override fun provide(collector: (com.jetbrains.ml.monitoring.MLTaskListenerTyped<M, P>, String) -> Unit) {
collector(this, task.id)
}
companion object {
val TOPIC: Topic<MessageBusMLTaskListenerProvider<*, *>> = Topic.create<MessageBusMLTaskListenerProvider<*, *>>("ml.task", MessageBusMLTaskListenerProvider::class.java)
}
}
@ApiStatus.Internal
interface MLTaskLoggingListener : com.jetbrains.ml.monitoring.MLTaskLoggingListener, MessageBusMLTaskLoggingListenerProvider {
val task: com.jetbrains.ml.MLTask<*, *>
override fun provide(collector: (com.jetbrains.ml.monitoring.MLTaskLoggingListener, String) -> Unit) {
collector(this, task.id)
}
companion object {
val TOPIC = Topic.create<MessageBusMLTaskLoggingListenerProvider>("ml.logging", MessageBusMLTaskLoggingListenerProvider::class.java)
}
}
@@ -0,0 +1,21 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.intellij.platform.ml.impl.tools
import com.intellij.internal.statistic.eventLog.EventLogGroup
import com.intellij.openapi.Disposable
import com.intellij.openapi.observable.util.whenDisposed
import com.intellij.platform.ml.impl.tools.logs.IntelliJFusEventRegister
import com.jetbrains.ml.MLTask
import com.jetbrains.ml.model.MLModel
import org.jetbrains.annotations.ApiStatus
@ApiStatus.Internal
fun <M : MLModel<P>, P : Any> EventLogGroup.registerMLTaskLogging(
task: MLTask<M, P>,
parentDisposable: Disposable,
eventPrefix: String = task.id,
) {
val componentRegister = IntelliJFusEventRegister(this)
val listenerController = componentRegister.registerLogging(task, eventPrefix)
parentDisposable.whenDisposed { listenerController.remove() }
}
@@ -1,16 +1,16 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.intellij.platform.ml.impl.logs.fus
package com.intellij.platform.ml.impl.tools.logs
import com.intellij.internal.statistic.eventLog.FeatureUsageData
import com.intellij.internal.statistic.eventLog.events.PrimitiveEventField
import com.intellij.internal.statistic.eventLog.events.VarargEventId
import com.intellij.lang.Language
import com.intellij.openapi.util.Version
import com.intellij.platform.ml.impl.logs.fus.ConverterObjectDescription.Companion.asIJObjectDescription
import com.intellij.platform.ml.impl.logs.fus.ConverterOfEnum.Companion.toIJConverter
import com.intellij.platform.ml.impl.logs.fus.IJEventPairConverter.Companion.typedBuild
import com.jetbrains.ml.building.blocks.logs.FusEventLogger
import com.jetbrains.ml.building.blocks.logs.FusEventRegister
import com.intellij.platform.ml.impl.tools.logs.ConverterObjectDescription.Companion.asIJObjectDescription
import com.intellij.platform.ml.impl.tools.logs.ConverterOfEnum.Companion.toIJConverter
import com.intellij.platform.ml.impl.tools.logs.IJEventPairConverter.Companion.typedBuild
import com.jetbrains.ml.logs.schema.EventField
import com.jetbrains.ml.logs.schema.EventPair
import org.jetbrains.annotations.ApiStatus
import com.intellij.internal.statistic.eventLog.EventLogGroup as IJEventLogGroup
import com.intellij.internal.statistic.eventLog.events.BooleanEventField as IJBooleanEventField
@@ -28,35 +28,35 @@ import com.intellij.internal.statistic.eventLog.events.ObjectEventData as IJObje
import com.intellij.internal.statistic.eventLog.events.ObjectEventField as IJObjectEventField
import com.intellij.internal.statistic.eventLog.events.ObjectListEventField as IJObjectListEventField
import com.intellij.internal.statistic.eventLog.events.StringEventField as IJStringEventField
import com.jetbrains.ml.features.api.logs.BooleanEventField as MLBooleanEventField
import com.jetbrains.ml.features.api.logs.ClassEventField as MLClassEventField
import com.jetbrains.ml.features.api.logs.DoubleEventField as MLDoubleEventField
import com.jetbrains.ml.features.api.logs.EnumEventField as MLEnumEventField
import com.jetbrains.ml.features.api.logs.EventField as MLEventField
import com.jetbrains.ml.features.api.logs.EventPair as MLEventPair
import com.jetbrains.ml.features.api.logs.FloatEventField as MLFloatEventField
import com.jetbrains.ml.features.api.logs.IntEventField as MLIntEventField
import com.jetbrains.ml.features.api.logs.LongEventField as MLLongEventField
import com.jetbrains.ml.features.api.logs.ObjectDescription as MLObjectDescription
import com.jetbrains.ml.features.api.logs.ObjectEventData as MLObjectEventData
import com.jetbrains.ml.features.api.logs.ObjectEventField as MLObjectEventField
import com.jetbrains.ml.features.api.logs.ObjectListEventField as MLObjectListEventField
import com.jetbrains.ml.features.api.logs.StringEventField as MLStringEventField
import com.jetbrains.ml.logs.schema.BooleanEventField as MLBooleanEventField
import com.jetbrains.ml.logs.schema.ClassEventField as MLClassEventField
import com.jetbrains.ml.logs.schema.DoubleEventField as MLDoubleEventField
import com.jetbrains.ml.logs.schema.EnumEventField as MLEnumEventField
import com.jetbrains.ml.logs.schema.EventField as MLEventField
import com.jetbrains.ml.logs.schema.EventPair as MLEventPair
import com.jetbrains.ml.logs.schema.FloatEventField as MLFloatEventField
import com.jetbrains.ml.logs.schema.IntEventField as MLIntEventField
import com.jetbrains.ml.logs.schema.LongEventField as MLLongEventField
import com.jetbrains.ml.logs.schema.ObjectDescription as MLObjectDescription
import com.jetbrains.ml.logs.schema.ObjectEventData as MLObjectEventData
import com.jetbrains.ml.logs.schema.ObjectEventField as MLObjectEventField
import com.jetbrains.ml.logs.schema.ObjectListEventField as MLObjectListEventField
import com.jetbrains.ml.logs.schema.StringEventField as MLStringEventField
@ApiStatus.Internal
class IntelliJFusEventRegister(private val baseEventGroup: IJEventLogGroup) : FusEventRegister {
class IntelliJFusEventRegister(private val baseEventGroup: IJEventLogGroup) : com.jetbrains.ml.logs.FusEventRegister {
private class Logger(
private val varargEventId: VarargEventId,
private val objectDescription: ConverterObjectDescription
) : FusEventLogger {
override suspend fun log(eventPairs: List<MLEventPair<*>>) {
) : com.jetbrains.ml.logs.FusEventLogger {
override fun log(eventPairs: List<MLEventPair<*>>) {
val ijEventPairs = objectDescription.buildEventPairs(eventPairs)
varargEventId.log(*ijEventPairs.toTypedArray())
}
}
override fun registerEvent(name: String, eventFields: List<MLEventField<*>>): FusEventLogger {
override fun registerEvent(name: String, eventFields: List<EventField<*>>): com.jetbrains.ml.logs.FusEventLogger {
val objectDescription = ConverterObjectDescription(MLObjectDescription(eventFields))
val varargEventId = baseEventGroup.registerVarargEvent(name, null, *objectDescription.getFields())
return Logger(varargEventId, objectDescription)
@@ -89,6 +89,7 @@ private fun <L> createConverter(mlEventField: MLEventField<L>): IJEventPairConve
is VersionEventField -> ConverterOfVersion(mlEventField) as IJEventPairConverter<L, *>
}
}
else -> throw IllegalArgumentException(
"""
Conversion of ${mlEventField.javaClass.simpleName} is not possible.
@@ -101,7 +102,7 @@ private fun <L> createConverter(mlEventField: MLEventField<L>): IJEventPairConve
private class ConverterOfCustom<T>(mlEventField: IJCustomEventField<T>) : IJEventPairConverter<T, T> {
override val ijEventField: IJEventField<T> = mlEventField.baseIJEventField
override fun buildEventPair(mlEventPair: MLEventPair<T>): IJEventPair<T> {
override fun buildEventPair(mlEventPair: EventPair<T>): IJEventPair<T> {
return ijEventField with mlEventPair.data
}
}
@@ -113,7 +114,7 @@ private class ConverterOfString(mlEventField: MLStringEventField) : IJEventPairC
description = mlEventField.lazyDescription()
)
override fun buildEventPair(mlEventPair: MLEventPair<String>): IJEventPair<String?> {
override fun buildEventPair(mlEventPair: EventPair<String>): IJEventPair<String?> {
return ijEventField with mlEventPair.data
}
}
@@ -126,7 +127,7 @@ private class ConverterOfLanguage(mlEventField: LanguageEventField) : IJEventPai
}
}
private class ConverterOfVersion(mlEventField: com.intellij.platform.ml.impl.logs.fus.VersionEventField) : IJEventPairConverter<Version, Version?> {
private class ConverterOfVersion(mlEventField: com.intellij.platform.ml.impl.tools.logs.VersionEventField) : IJEventPairConverter<Version, Version?> {
private class VersionEventField(override val name: String, override val description: String?) : PrimitiveEventField<Version?>() {
override val validationRule: List<String>
get() = listOf("{regexp#version}")
@@ -0,0 +1,32 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.intellij.platform.ml.impl.tools.logs
import com.intellij.lang.Language
import com.jetbrains.ml.logs.schema.EventField
import org.jetbrains.annotations.ApiStatus
/**
* Something, that is dedicated for one language only.
*/
@ApiStatus.Internal
interface LanguageSpecific {
val language: Language
}
/**
* The analyzer, that adds information about ML model's language to logs.
*/
@ApiStatus.Internal
class ModelLanguageAnalyser<M, P : Any> : com.jetbrains.ml.analysis.MLTaskAnalyserTyped<M, P>
where M : com.jetbrains.ml.model.MLModel<P>,
M : LanguageSpecific {
private val LANGUAGE = LanguageEventField("model_language") { "The programming language the ML model is trained for" }
override fun startMLSessionAnalysis(sessionInfo: com.jetbrains.ml.session.MLSessionInfo<M, P>) = object : com.jetbrains.ml.analysis.MLSessionAnalyserTyped<M, P> {
override suspend fun analyseSession(tree: com.jetbrains.ml.tree.MLTree.ATopNode<M, P>?) = buildList {
tree?.mlModel?.let { add(LANGUAGE with it.language) }
}
}
override val declaration: List<EventField<*>> = listOf(LANGUAGE)
}
@@ -0,0 +1,34 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.intellij.platform.ml.impl.tools.logs
import com.intellij.openapi.util.Version
import com.jetbrains.ml.logs.schema.EventField
import org.jetbrains.annotations.ApiStatus
/**
* Something, that has versions.
*/
@ApiStatus.Internal
interface Versioned {
val version: Version?
}
/**
* Adds model's version to the ML logs.
*/
@ApiStatus.Internal
class ModelVersionAnalyser<M, P : Any> : com.jetbrains.ml.analysis.MLTaskAnalyserTyped<M, P>
where M : com.jetbrains.ml.model.MLModel<P>,
M : Versioned {
companion object {
private val VERSION = VersionEventField("model_version") { "Version of the ML model" }
}
override fun startMLSessionAnalysis(sessionInfo: com.jetbrains.ml.session.MLSessionInfo<M, P>) = object : com.jetbrains.ml.analysis.MLSessionAnalyserTyped<M, P> {
override suspend fun analyseSession(tree: com.jetbrains.ml.tree.MLTree.ATopNode<M, P>?) = buildList {
tree?.mlModel?.version?.let { add(VERSION with it) }
}
}
override val declaration: List<EventField<*>> = listOf(VERSION)
}
@@ -1,13 +1,12 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.intellij.platform.ml.impl.logs.fus
package com.intellij.platform.ml.impl.tools.logs
import com.intellij.lang.Language
import com.intellij.openapi.util.Version
import com.jetbrains.ml.features.api.logs.CustomRuleEventField
import com.jetbrains.ml.logs.schema.CustomRuleEventField
import org.jetbrains.annotations.ApiStatus
import com.intellij.internal.statistic.eventLog.events.EventField as IJEventField
@ApiStatus.Internal
sealed class IJSpecificEventField<T>(name: String, lazyDescription: () -> String) : CustomRuleEventField<T>(name, lazyDescription)
@@ -0,0 +1,142 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.intellij.platform.ml.impl.tools.model
import com.intellij.internal.ml.DecisionFunction
import com.jetbrains.ml.*
import com.jetbrains.ml.model.MLModel
import org.jetbrains.annotations.ApiStatus
/**
* A wrapper for using legacy [DecisionFunction] ML API`s [MLModel].
*/
@ApiStatus.Internal
open class RegressionModel private constructor(
private val decisionFunction: DecisionFunction,
private val modelUnits: Set<MLUnit<*>>,
globalUnits: Set<MLUnit<*>>,
private val featureSerialization: FeatureNameSerialization,
) : MLModel<Double> {
constructor(
decisionFunction: DecisionFunction,
featureSerialization: FeatureNameSerialization,
globalUnits: Set<MLUnit<*>>,
) : this(
decisionFunction = decisionFunction,
modelUnits = decisionFunction.featuresOrder.map { featureMapper ->
featureSerialization.deserialize(featureMapper.featureName, globalUnits.associateBy { it.name }).first
}.toSet(),
globalUnits = globalUnits,
featureSerialization = featureSerialization
)
override val knownFeatures: Map<MLUnit<*>, FeatureFilter> = createKnownFeatures(
DecisionFunctionWrapper(decisionFunction, globalUnits, featureSerialization),
modelUnits
)
override fun predict(taskUnits: List<MLUnitsMap>, contexts: List<Any?>, features: Map<MLUnit<*>, List<Feature>>): Double {
val array = DoubleArray(decisionFunction.featuresOrder.size)
val featurePerSerializedName: Map<String, Feature> = features
.flatMap { (unit, unitFeatures) -> unitFeatures.map { unit to it } }
.associate { (unit, feature) -> featureSerialization.serialize(unit, feature.declaration.name) to feature }
require(features.keys == modelUnits) {
"Given features units are ${features.keys}, but this model needs ${modelUnits}"
}
for (featureI in decisionFunction.featuresOrder.indices) {
val featureMapper = decisionFunction.featuresOrder[featureI]
val featureSerializedName = featureMapper.featureName
val featureValue = featureMapper.asArrayValue(featurePerSerializedName[featureSerializedName]?.value)
array[featureI] = featureValue
}
return decisionFunction.predict(array)
}
interface FeatureNameSerialization {
fun serialize(mlUnit: MLUnit<*>, featureName: String): String
fun deserialize(serializedFeatureName: String, unitsPerName: Map<String, MLUnit<*>>): Pair<MLUnit<*>, String>
}
object DefaultFeatureSerialization : FeatureNameSerialization {
private const val SERIALIZED_FEATURE_SEPARATOR = '/'
override fun serialize(mlUnit: MLUnit<*>, featureName: String): String {
return mlUnit.name + SERIALIZED_FEATURE_SEPARATOR + featureName
}
override fun deserialize(serializedFeatureName: String, unitsPerName: Map<String, MLUnit<*>>): Pair<MLUnit<*>, String> {
val indexOfLastSeparator = serializedFeatureName.indexOfLast { it == SERIALIZED_FEATURE_SEPARATOR }
require(indexOfLastSeparator >= 0) { "Feature name '$serializedFeatureName' does not contain ML unit's name" }
val unitName = serializedFeatureName.slice(0 until indexOfLastSeparator)
val featureName = serializedFeatureName.slice(indexOfLastSeparator until serializedFeatureName.length)
val featureUnit = requireNotNull(unitsPerName[unitName]) {
"""
Serialized feature '$serializedFeatureName' has tier $unitName,
but all available tiers are ${unitsPerName.keys}
""".trimIndent()
}
return featureUnit to featureName
}
}
private class DecisionFunctionWrapper(
private val decisionFunction: DecisionFunction,
globalUnits: Set<MLUnit<*>>,
private val featureNameSerialization: FeatureNameSerialization,
) {
private val availableUnitsPerName: Map<String, MLUnit<*>> = globalUnits.associateBy { it.name }
val knownFeatures: Map<MLUnit<*>, Set<String>> = run {
val knownFeaturesSerializedNames = decisionFunction.featuresOrder.map { it.featureName }.toSet()
knownFeaturesSerializedNames
.map { featureNameSerialization.deserialize(it, availableUnitsPerName) }
.groupBy({ it.first }, { it.second })
.mapValues { it.value.toSet() }
}
fun getUnknownFeatures(unit: MLUnit<*>, featuresNames: Set<String>): Set<String> {
val knownFeatures = knownFeatures[unit] ?: return featuresNames
return featuresNames.filterNot { it in knownFeatures }.toSet()
}
}
companion object {
private fun createKnownFeatures(
decisionFunction: DecisionFunctionWrapper,
modelUnits: Set<MLUnit<*>>,
): Map<MLUnit<*>, FeatureFilter> {
fun createFeatureSelector(unit: MLUnit<*>) = object : FeatureFilter {
init {
val knownFeatures = decisionFunction.knownFeatures
knownFeatures.forEach { (unit, features) ->
val nonConsistentlyKnownFeatures = decisionFunction.getUnknownFeatures(unit, features)
require(nonConsistentlyKnownFeatures.isEmpty()) {
"These features are known and unknown at the same time: $nonConsistentlyKnownFeatures"
}
}
}
override fun accept(featureDeclarations: Set<FeatureDeclaration<*>>): Set<FeatureDeclaration<*>> {
val availableFeaturesPerName = featureDeclarations.associateBy { it.name }
val availableFeaturesNames = featureDeclarations.map { it.name }.toSet()
val unknownFeaturesNames = decisionFunction.getUnknownFeatures(unit, availableFeaturesNames)
val knownAvailableFeaturesNames = availableFeaturesNames - unknownFeaturesNames
val knownAvailableFeatures = knownAvailableFeaturesNames.map { availableFeaturesPerName.getValue(it) }.toSet()
return knownAvailableFeatures
}
override fun accept(featureDeclaration: FeatureDeclaration<*>): Boolean {
val unknown = decisionFunction.getUnknownFeatures(unit, setOf(featureDeclaration.name))
return unknown.isEmpty()
}
}
return modelUnits.associateWith { createFeatureSelector(it) }
}
}
}
@@ -18,6 +18,8 @@ import kotlin.time.Duration
import kotlin.time.Duration.Companion.seconds
import kotlin.time.DurationUnit
import com.intellij.platform.ml.feature.FeatureDeclaration as OldFeatureDeclaration
import com.jetbrains.ml.Feature as NewFeature
import com.jetbrains.ml.FeatureDeclaration as NewFeatureDeclaration
@ApiStatus.Internal
@Service
@@ -49,6 +51,12 @@ class TypingSpeedTracker {
}
}
fun getTypingSpeedNewEventPairs(): Collection<Pair<EventPair<*>, NewFeature>> = DECAY_DURATIONS_NEW.mapNotNull { (decayDuration, eventFieldAndFeature) ->
typingSpeeds[decayDuration]?.let {
(eventFieldAndFeature.first with it) to (eventFieldAndFeature.second with it)
}
}
@TestOnly
fun getTypingSpeed(decayDuration: Duration): Float? = typingSpeeds[decayDuration]
@@ -88,9 +96,12 @@ class TypingSpeedTracker {
),
OldFeatureDeclaration.float("typing_speed_${it}s").nullable())
}
private val DECAY_DURATIONS_NEW = listOf(1, 2, 5, 30)
.associate { it.seconds to Pair(EventFields.Float("typing_speed_${it}s"), NewFeatureDeclaration.float("typing_speed_${it}s").nullable()) }
fun getInstance(): TypingSpeedTracker = service()
fun getEventFields(): Array<EventField<*>> = DECAY_DURATIONS.values.map { it.first }.toTypedArray()
fun getFeatures(): Set<OldFeatureDeclaration<*>> = DECAY_DURATIONS.values.map { it.second }.toSet()
fun getFeaturesNew(): List<NewFeatureDeclaration<*>> = DECAY_DURATIONS_NEW.values.map { it.second }
}
}
@@ -0,0 +1,26 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.intellij.codeInsight.inline.completion.ml
import com.intellij.codeInsight.inline.completion.logs.TypingSpeedTracker
import com.jetbrains.ml.*
internal class TypingSpeedFeatureProvider : FeatureProvider(MLUnitTyping) {
object Features {
val TIME_SINCE_LAST_TYPING = FeatureDeclaration.long("time_since_last_typing").nullable()
}
override val featureComputationPolicy = FeatureComputationPolicy(true, true)
override suspend fun computeFeatures(units: MLUnitsMap, usefulFeaturesFilter: FeatureFilter): List<Feature> = buildList {
with(units[MLUnitTyping]) {
val timeSinceLastTyping = getTimeSinceLastTyping()
if (timeSinceLastTyping != null) {
add(Features.TIME_SINCE_LAST_TYPING.with(timeSinceLastTyping))
addAll(getTypingSpeedNewEventPairs().map { it.second })
}
}
}
override val featureDeclarations = TypingSpeedTracker.getFeaturesNew() + Features.TIME_SINCE_LAST_TYPING
}
@@ -6,7 +6,7 @@ import com.intellij.lang.Language
import com.intellij.openapi.editor.Editor
import com.intellij.platform.ml.Tier
import com.intellij.psi.PsiFile
import com.jetbrains.ml.features.api.MLUnit
import com.jetbrains.ml.MLUnit
import org.jetbrains.annotations.ApiStatus
@ApiStatus.Internal
@@ -547,6 +547,9 @@
<extensionPoint qualifiedName="com.intellij.platform.ml.taskListener"
interface="com.intellij.platform.ml.monitoring.MLTaskGroupListener"
dynamic="true"/>
<extensionPoint qualifiedName="com.intellij.platform.ml.featureProvider"
interface="com.jetbrains.ml.FeatureProvider"
dynamic="true"/>
<extensionPoint name="defender.config" interface="com.intellij.diagnostic.WindowsDefenderChecker$Extension" dynamic="true" />
<extensionPoint name="authorizationProvider" interface="com.intellij.ide.impl.AuthorizationProvider" dynamic="true" />
@@ -1697,6 +1697,7 @@
<!-- ML API -->
<platform.ml.descriptor implementation="com.intellij.codeInsight.inline.completion.ml.TypingFeatures"/>
<platform.ml.environmentExtender implementation="com.intellij.codeInsight.inline.completion.ml.TypingSpeedProvider"/>
<platform.ml.featureProvider implementation="com.intellij.codeInsight.inline.completion.ml.TypingSpeedFeatureProvider"/>
<!-- Presentation Assistant -->
<applicationService serviceImplementation="com.intellij.platform.ide.impl.presentationAssistant.PresentationAssistant"/>
@@ -11,6 +11,5 @@
<orderEntry type="module" module-name="intellij.python.community.plugin.java" />
<orderEntry type="module" module-name="intellij.python.community.plugin.minor" />
<orderEntry type="module" module-name="intellij.python.community.plugin.modules" />
<orderEntry type="library" name="jetbrains.ml.models.python.imports.ranking.model" level="project" />
</component>
</module>
@@ -7,7 +7,6 @@
<orderEntry type="module" module-name="intellij.python.langInjection" scope="RUNTIME" />
<orderEntry type="module" module-name="intellij.python.copyright" scope="RUNTIME" />
<orderEntry type="module" module-name="intellij.python.terminal" scope="RUNTIME" />
<orderEntry type="module" module-name="intellij.python.ml.features" scope="RUNTIME" />
<orderEntry type="module" module-name="intellij.restructuredtext" scope="RUNTIME" />
<orderEntry type="module" module-name="intellij.python.grazie" scope="RUNTIME" />
<orderEntry type="module" module-name="intellij.python.markdown" scope="RUNTIME" />
@@ -1,23 +0,0 @@
<?xml version="1.0" encoding="UTF-8"?>
<module type="JAVA_MODULE" version="4">
<component name="NewModuleRootManager" inherit-compiler-output="true">
<exclude-output />
<content url="file://$MODULE_DIR$">
<sourceFolder url="file://$MODULE_DIR$/src" isTestSource="false" />
</content>
<orderEntry type="inheritedJdk" />
<orderEntry type="sourceFolder" forTests="false" />
<orderEntry type="library" name="kotlin-stdlib" level="project" />
<orderEntry type="module" module-name="intellij.platform.core" />
<orderEntry type="module" module-name="intellij.platform.statistics" />
<orderEntry type="module" module-name="intellij.platform.ml.impl" />
<orderEntry type="module" module-name="intellij.python.psi.impl" />
<orderEntry type="module" module-name="intellij.platform.ide" />
<orderEntry type="module" module-name="intellij.platform.ide.impl" />
<orderEntry type="module" module-name="intellij.python.community.impl" />
<orderEntry type="module" module-name="intellij.platform.concurrency" />
<orderEntry type="module" module-name="intellij.platform.core.ui" />
<orderEntry type="module" module-name="intellij.platform.ml" />
<orderEntry type="library" name="jetbrains.ml.models.python.imports.ranking.model" level="project" />
</component>
</module>
@@ -1,47 +0,0 @@
package com.intellij.python.ml.features.imports
import com.intellij.ide.DataManager
import com.intellij.openapi.actionSystem.DataContext
import com.intellij.openapi.application.ApplicationManager
import com.intellij.openapi.components.Service
import com.intellij.openapi.ui.popup.JBPopupFactory
import com.intellij.openapi.util.Computable
import com.intellij.util.Consumer
import com.jetbrains.python.codeInsight.imports.ImportCandidateHolder
import com.jetbrains.python.codeInsight.imports.PyImportChooser
import org.jetbrains.concurrency.AsyncPromise
import org.jetbrains.concurrency.Promise
import org.jetbrains.concurrency.resolvedPromise
class PyMLImportChooser : PyImportChooser() {
override fun selectImport(
sources: MutableList<out ImportCandidateHolder>,
useQualifiedImport: Boolean
): Promise<ImportCandidateHolder?> {
if (ApplicationManager.getApplication().isUnitTestMode()) {
return resolvedPromise<ImportCandidateHolder?>(sources.get(0))
}
val result = AsyncPromise<ImportCandidateHolder?>()
launchMLRanking(sources) { mlRanking: RateableRankingResult? ->
// GUI part
DataManager.getInstance().getDataContextFromFocus().doWhenDone(Consumer { dataContext: DataContext? ->
val popup = JBPopupFactory.getInstance()
.createPopupChooserBuilder<ImportCandidateHolder>(mlRanking!!.order)
.setItemChosenCallback(Consumer { item: ImportCandidateHolder? ->
result.setResult(item)
mlRanking.submitSelectedItem(item!!)
})
.setCancelCallback(Computable {
mlRanking.submitPopUpClosed()
true
})
processPopup(popup, useQualifiedImport)
popup.createPopup()
.showInBestPositionFor(dataContext!!)
})
Unit
}
return result
}
}
@@ -1,26 +0,0 @@
# Ranking in fixing missing import quick fix
## Registry
* `quickfix.ranking.ml` - Status of the system
* `IN_EXPERIMENT` - ML is enabled depending on machine id.
* `ENABLED` - ML is enabled
* `DISABLED` - ML is disabled. ML features aren't computed
* `quickfix.ranking.ml.model.path` - Path to a local model
## ML Model
ML model loaded from resources of library `jetbrains.ml.models.python.imports.ranking.model`
(the dependency is located in `community/python/python-psi-impl/intellij.python.psi.impl.iml`).
When you want, you can set `quickfix.ranking.ml.model.path` to an absolute path on your machine
and try out a freshly trained ML model.
Note, that right now a regression model is used.
You may change that by tweaking `com.intellij.python.ml.features.imports.MLModelService`.
## Contact
Initial implementation: Gleb Marin
ML Model, A/B Experiment Pipelines: Nikita Ermolenko
PyCharm Support: Andrey Matveev
@@ -1,24 +0,0 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.intellij.python.ml.features.imports
import com.jetbrains.ml.features.api.logs.BooleanEventField
import com.jetbrains.ml.features.api.logs.EnumEventField
import com.jetbrains.ml.features.api.logs.IntEventField
import com.jetbrains.ml.features.api.logs.LongEventField
internal object ContextAnalysis {
val CANCELLED = BooleanEventField("selection_cancelled") { "No item has been selected" }
val SELECTED_POSITION = IntEventField("selected_position") { "The position of the selected import statement" }
val MODEL_UNAVAILABLE = BooleanEventField("model_unavailable") { "ML model was unavailable" }
val SELECTED_POSITION_INITIAL = IntEventField("selected_position_initial") { "The position of the selected import statement, if the final ML ranking would not have happened" }
val TIME_MS_TO_DISPLAY = LongEventField("time_ms_before_displayed") { "Duration from the quickfix start until when the imports were displayed" }
val TIME_MS_BEFORE_CLOSED = LongEventField("time_ms_before_closed") { "Duration from the quickfix start until the pop-up was closed" }
val ML_LOGGING_STATE = EnumEventField.of<LoggingOption>("ml_logging_state", { "State of the ML session logging" })
val ML_ENABLED = BooleanEventField("ml_enabled") { "Machine Learning ranking is enabled" }
val ML_STATUS_CORRESPONDS_TO_BUCKET = BooleanEventField("ml_status_corresponds_to_bucket") { "Field 'ml_enabled' corresponds to the ML bucket % 2" }
}
internal object CandidateAnalysis {
val MLFieldCorrectElement = BooleanEventField("is_correct") { "The candidate was chosen" }
}
@@ -1,202 +0,0 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.intellij.python.ml.features.imports
import com.intellij.openapi.application.writeAction
import com.intellij.openapi.components.Service
import com.intellij.openapi.components.service
import com.intellij.openapi.diagnostic.debug
import com.intellij.openapi.diagnostic.logger
import com.intellij.util.application
import com.jetbrains.ml.building.blocks.model.MLModel
import com.jetbrains.ml.building.blocks.task.MLFeaturesTree
import com.jetbrains.ml.features.api.logs.EventPair
import com.jetbrains.python.codeInsight.imports.ImportCandidateHolder
import com.jetbrains.python.codeInsight.imports.mlapi.ImportCandidateContext
import com.jetbrains.python.codeInsight.imports.mlapi.ImportRankingContext
import com.jetbrains.python.codeInsight.imports.mlapi.MLUnitImportCandidate
import com.jetbrains.python.codeInsight.imports.mlapi.MLUnitImportRankingContext
import kotlinx.coroutines.CoroutineScope
import kotlinx.coroutines.Dispatchers
import kotlinx.coroutines.delay
import kotlinx.coroutines.launch
import java.util.concurrent.atomic.AtomicBoolean
import kotlin.time.Duration.Companion.seconds
@Service(Service.Level.APP)
internal class MLApiComputations(
val coroutineScope: CoroutineScope,
)
internal sealed class FinalCandidatesRanker(
protected val contextFeatures: MLFeaturesTree,
protected val contextLogs: MutableList<EventPair<*>>,
protected val timestampStarted: Long,
) {
abstract fun launchMLRanking(initialCandidatesOrder: MutableList<out ImportCandidateHolder>, displayResult: (RateableRankingResult) -> Unit)
abstract val mlEnabled: Boolean
}
private class ExperimentalMLRanker(
mlContextFeatures: MLFeaturesTree,
mlContextLogs: MutableList<EventPair<*>>,
timestampStarted: Long,
private val mlModel: MLModel<Double>,
) : FinalCandidatesRanker(mlContextFeatures, mlContextLogs, timestampStarted) {
override val mlEnabled = true
override fun launchMLRanking(initialCandidatesOrder: MutableList<out ImportCandidateHolder>, displayResult: (RateableRankingResult) -> Unit) {
service<MLApiComputations>().coroutineScope.launch(Dispatchers.Default) {
contextFeatures.computeFeaturesFromProviders(
MLUnitImportRankingContext with ImportRankingContext(initialCandidatesOrder),
)
val importCandidatesFeatures = initialCandidatesOrder.associateWith { candidate ->
val candidateFeatures = contextFeatures.addChild()
candidateFeatures.computeFeaturesFromProviders(
MLUnitImportCandidate with ImportCandidateContext(initialCandidatesOrder, candidate)
)
candidateFeatures
}
val predictionsByCandidateTree = contextFeatures
.predictForChildren(mlModel, emptyMap())
.toMap()
val scoreByCandidate = initialCandidatesOrder.associateWith { candidate ->
val candidateTree = importCandidatesFeatures.getValue(candidate)
val score = predictionsByCandidateTree.getValue(candidateTree)
score
}
val relevanceCandidateOrder = scoreByCandidate.toList()
.sortedByDescending { it.second }
.map { it.first }
logger<ExperimentalMLRanker>().debug {
"Sorted import candidates:\n" +
relevanceCandidateOrder.mapIndexed { index, candidate ->
val score = scoreByCandidate.getValue(candidate)
val initialIndex = initialCandidatesOrder.indexOf(candidate)
"\t#${index + 1}: '${candidate.presentableText}' (score: ${score}, initial index: $initialIndex)"
}.joinToString("\n")
}
writeAction {
displayResult(RateableRankingResult(contextFeatures, contextLogs, importCandidatesFeatures, initialCandidatesOrder, relevanceCandidateOrder, timestampStarted))
}
}
}
}
private class InitialOrderKeepingRanker(
contextFeatures: MLFeaturesTree,
contextLogs: MutableList<EventPair<*>>,
timestampStarted: Long,
) : FinalCandidatesRanker(contextFeatures, contextLogs, timestampStarted) {
override val mlEnabled = false
override fun launchMLRanking(initialCandidatesOrder: MutableList<out ImportCandidateHolder>, displayResult: (RateableRankingResult) -> Unit) {
application.runWriteAction {
val importCandidatesFeatures = initialCandidatesOrder.associateWith { candidate ->
contextFeatures.addChild()
}
displayResult(RateableRankingResult(contextFeatures, contextLogs, importCandidatesFeatures, initialCandidatesOrder, initialCandidatesOrder, timestampStarted))
}
}
}
internal fun launchMLRanking(initialCandidatesOrder: MutableList<out ImportCandidateHolder>, displayResult: (RateableRankingResult) -> Unit) {
// In EDT
val timestampStarted = System.currentTimeMillis()
return MLTaskPyCharmImportStatementsRanking.lifetime {
val mlContextFeatures = newMLFeaturesTree()
val mlContextLogs = mutableListOf<EventPair<*>>()
val rankingStatus = service<FinalImportRankingStatusService>().status
var modelWasReady = true
val ranker = if (rankingStatus.mlEnabled) {
val mlModel = service<MLModelService>().getModel(MLTaskPyCharmImportStatementsRanking)
if (mlModel != null) {
ExperimentalMLRanker(mlContextFeatures, mlContextLogs, timestampStarted, mlModel)
} else {
modelWasReady = false
mlContextLogs.add(ContextAnalysis.MODEL_UNAVAILABLE with true)
InitialOrderKeepingRanker(mlContextFeatures, mlContextLogs, timestampStarted)
}
}
else {
InitialOrderKeepingRanker(mlContextFeatures, mlContextLogs, timestampStarted)
}
mlContextLogs.add(ContextAnalysis.ML_STATUS_CORRESPONDS_TO_BUCKET with (rankingStatus.mlStatusCorrespondsToBucket && modelWasReady))
mlContextLogs.add(ContextAnalysis.ML_ENABLED with (rankingStatus.mlEnabled && modelWasReady))
ranker.launchMLRanking(initialCandidatesOrder) { result ->
val timestampDisplayed = System.currentTimeMillis()
displayResult(result)
mlContextLogs.add(ContextAnalysis.TIME_MS_TO_DISPLAY with (timestampDisplayed - timestampStarted))
}
}
}
internal class RateableRankingResult(
private val mlContextFeatures: MLFeaturesTree,
private val contextAnalysis: MutableList<EventPair<*>>,
private val importCandidates: Map<ImportCandidateHolder, MLFeaturesTree>,
val orderInitial: List<ImportCandidateHolder>,
val order: List<ImportCandidateHolder>,
private val timestampStarted: Long,
) {
private var submitted = AtomicBoolean(false)
private val analysis: MutableMap<MLFeaturesTree, MutableList<EventPair<*>>> = mutableMapOf(
mlContextFeatures to contextAnalysis
)
fun submitPopUpClosed() {
val timestampClosed = System.currentTimeMillis()
contextAnalysis.add(ContextAnalysis.TIME_MS_BEFORE_CLOSED with timestampClosed - timestampStarted)
service<MLApiComputations>().coroutineScope.launch(Dispatchers.IO) {
delay(1.seconds)
synchronized(this@RateableRankingResult) {
if (submitted.get()) return@launch
submitted.set(true)
contextAnalysis.add(ContextAnalysis.CANCELLED with true)
}
submitLogs()
}
}
fun submitSelectedItem(selected: ImportCandidateHolder) = synchronized(this) {
// In EDT
require(submitted.compareAndSet(false, true)) { "Some feedback to the ml ranker was already submitted" }
analysis[requireNotNull(importCandidates[selected])] = mutableListOf(CandidateAnalysis.MLFieldCorrectElement with true)
val selectedIndex = order.indexOf(selected)
val selectedIndexOriginal = orderInitial.indexOf(selected)
require(selectedIndex >= 0 && selectedIndexOriginal >= 0)
contextAnalysis.add(ContextAnalysis.SELECTED_POSITION with selectedIndex)
contextAnalysis.add(ContextAnalysis.SELECTED_POSITION_INITIAL with selectedIndexOriginal)
service<MLApiComputations>().coroutineScope.launch(Dispatchers.IO) {
submitLogs()
}
}
private suspend fun submitLogs() {
val logger = PyCharmImportsRankingLogs.mlLogger
val loggingOption = getLoggingOption(mlContextFeatures)
contextAnalysis.add(ContextAnalysis.ML_LOGGING_STATE with loggingOption)
when (loggingOption) {
LoggingOption.FULL -> logger.log(mlContextFeatures, analysis)
LoggingOption.NO_TREE -> {
val treeStem = logger.newMLLogsTree()
treeStem.addAnalysis(contextAnalysis)
logger.log(mlContextFeatures.sessionId, treeStem)
}
LoggingOption.SKIP -> Unit
}
}
}
@@ -1,44 +0,0 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.intellij.python.ml.features.imports
import com.intellij.internal.statistic.eventLog.EventLogGroup
import com.intellij.internal.statistic.service.fus.collectors.CounterUsagesCollector
import com.intellij.platform.ml.impl.logs.MLEventLoggerProvider.Companion.ML_RECORDER_ID
import com.intellij.platform.ml.impl.logs.fus.IntelliJFusEventRegister
import com.intellij.util.application
import com.jetbrains.ml.building.blocks.logs.FusLoggerFactories
import com.jetbrains.ml.building.blocks.logs.extractEventFields
import com.jetbrains.ml.building.blocks.task.MLFeaturesTree
internal object PyCharmImportsRankingLogs : CounterUsagesCollector() {
private val GROUP = EventLogGroup("pycharm.quickfix.imports", 4, ML_RECORDER_ID)
val mlLogger = FusLoggerFactories
.withOneEvent(
fusEventName = "pycharm_import_statements_ranking",
fusEventRegister = IntelliJFusEventRegister(GROUP),
treeFeatures = MLTaskPyCharmImportStatementsRanking.treeFeatures,
treeAnalysisFields = listOf(
extractEventFields(ContextAnalysis),
extractEventFields(CandidateAnalysis)
)
)
override fun getGroup() = GROUP
}
internal enum class LoggingOption {
FULL,
NO_TREE,
SKIP
}
internal const val RELEASE_FULL_LOGGING_PERCENT = 10
internal fun getLoggingOption(tree: MLFeaturesTree): LoggingOption {
if (application.isEAP || application.isUnitTestMode) {
return LoggingOption.FULL
}
if ((tree.sessionId * 37) % 100 > RELEASE_FULL_LOGGING_PERCENT) return LoggingOption.NO_TREE
return LoggingOption.FULL
}
@@ -1,93 +0,0 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.intellij.python.ml.features.imports
import com.intellij.codeInsight.inline.completion.ml.MLUnitTyping
import com.intellij.openapi.components.Service
import com.intellij.openapi.components.service
import com.intellij.openapi.project.Project
import com.intellij.openapi.startup.ProjectActivity
import com.intellij.platform.ml.impl.MLModelSelector
import com.intellij.python.ml.features.imports.MissingTypingFeaturesLoader.Companion.fillingTypingFeatures
import com.jetbrains.ml.building.blocks.model.MLModel
import com.jetbrains.ml.building.blocks.model.MLModelLoader
import com.jetbrains.ml.building.blocks.model.catboost.DistributionFormat
import com.jetbrains.ml.building.blocks.model.catboost.ModelDistributionReaders
import com.jetbrains.ml.building.blocks.model.classical.MLModelLoaders
import com.jetbrains.ml.features.api.MLUnit
import com.jetbrains.ml.features.api.feature.Feature
import com.jetbrains.ml.features.api.feature.FeatureDeclaration
import java.nio.file.Path
@Service
class MLModelService : MLModelSelector<MLModel<Double>, Double>(
defaultLoader = MLModelLoaders.Regressor(
reader = ModelDistributionReaders.FromJavaResources(MLModelService::class.java, Path.of("python-imports-ranking-model")),
format = DistributionFormat(),
).fillingTypingFeatures(),
registryModelLoaderFactory = { path ->
MLModelLoaders.Regressor(
reader = ModelDistributionReaders.FromDirectory(path),
format = DistributionFormat(),
).fillingTypingFeatures()
},
modelPathRegistryKey = "quickfix.ranking.ml.model.path",
coroutineScope = service<MLApiComputations>().coroutineScope,
)
private class MissingTypingFeaturesLoader(private val baseLoader: MLModelLoader<MLModel<Double>, Double>) : MLModelLoader<MLModel<Double>, Double> by baseLoader {
override suspend fun loadModel(features: Map<MLUnit<*>, List<FeatureDeclaration<*>>>, parameters: Map<String, Any>?): MLModel<Double> {
val model = baseLoader.loadModel(features + mapOf(
MLUnitTyping to listOf(
TypingFeatures.SINCE_LAST_TYPING,
TypingFeatures.TYPING_SPEED_1S,
TypingFeatures.TYPING_SPEED_2S,
TypingFeatures.TYPING_SPEED_30S,
TypingFeatures.TYPING_SPEED_5S
)
), parameters)
return object : MLModel<Double> by model {
override fun predict(features: Map<MLUnit<*>, List<Feature>>, parameters: Map<String, Any>): Double {
return model.predict(features + mapOf(MLUnitTyping to fillTypingFeatures()), parameters)
}
override fun predictBatch(contextFeatures: Map<MLUnit<*>, List<Feature>>, itemFeatures: List<Map<MLUnit<*>, List<Feature>>>, parameters: Map<String, Any>): List<Double> {
return model.predictBatch(contextFeatures + mapOf(MLUnitTyping to fillTypingFeatures()), itemFeatures, parameters)
}
}
}
companion object {
fun MLModelLoader<MLModel<Double>, Double>.fillingTypingFeatures(): MLModelLoader<MLModel<Double>, Double> {
return MissingTypingFeaturesLoader(this)
}
}
}
/**
* TODO: Should be removed, because they don't exist anymore
*/
internal object TypingFeatures {
val SINCE_LAST_TYPING = FeatureDeclaration.int("time_since_last_typing")
val TYPING_SPEED_1S = FeatureDeclaration.double("typing_speed_1s")
val TYPING_SPEED_2S = FeatureDeclaration.double("typing_speed_2s")
val TYPING_SPEED_30S = FeatureDeclaration.double("typing_speed_30s")
val TYPING_SPEED_5S = FeatureDeclaration.double("typing_speed_5s")
}
private class QuickfixRankingModelLoading : ProjectActivity {
override suspend fun execute(project: Project) {
service<MLModelService>().prepareModel(MLTaskPyCharmImportStatementsRanking)
}
}
private fun fillTypingFeatures(): List<Feature> {
return listOf(
TypingFeatures.SINCE_LAST_TYPING with 1,
TypingFeatures.TYPING_SPEED_2S with 0.0,
TypingFeatures.TYPING_SPEED_30S with 0.0,
TypingFeatures.TYPING_SPEED_5S with 0.0,
TypingFeatures.TYPING_SPEED_1S with 0.0,
)
}
@@ -1,29 +0,0 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.intellij.python.ml.features.imports
import com.jetbrains.ml.building.blocks.task.MLTask
import com.jetbrains.python.codeInsight.imports.mlapi.features.*
val MLTaskPyCharmImportStatementsRanking = MLTask(
id = "pycharm_import_statements_ranking",
treeFeaturesPerLevel = listOf(
{
featuresFromProviders(
CandidatesListFeatures,
BaseProjectFeatures
)
},
{
featuresFromProviders(
ImportsFeatures,
RelevanceEvaluationFeatures,
ImportCandidateRelativeFeatures,
NeighborFilesImportsFeatures,
PrimitiveImportFeatures,
PsiStructureFeatures,
OpenFilesImportsFeatures,
)
}
)
)
@@ -1,14 +0,0 @@
<idea-plugin package="com.intellij.python.ml.features">
<extensions defaultExtensionNs="com.intellij">
<statistics.counterUsagesCollector implementationClass="com.intellij.python.ml.features.imports.PyCharmImportsRankingLogs"/>
<registryKey
key="quickfix.ranking.ml.model.path"
defaultValue=""
restartRequired="true"
description="Absolute path to a custom ML model used for ranking missing imports"/>
<postStartupActivity implementation="com.intellij.python.ml.features.imports.QuickfixRankingModelLoading"/>
<applicationService serviceInterface="com.jetbrains.python.codeInsight.imports.ImportChooser"
serviceImplementation="com.intellij.python.ml.features.imports.PyMLImportChooser"
overrides="true"/>
</extensions>
</idea-plugin>
-5
View File
@@ -4,10 +4,6 @@
files:
- name: $MAVEN_REPOSITORY$/org/apache/ws/xmlrpc/xmlrpc/2/xmlrpc-2.jar
reason: withProjectLibrary
- name: jetbrains.ml.models.python.imports.ranking.model
files:
- name: $MAVEN_REPOSITORY$/com/jetbrains/ml/models/python-imports-ranking-model/lilac-coua/python-imports-ranking-model-lilac-coua.jar
reason: <- intellij.python.ml.features
- name: jsr305
files:
- name: $MAVEN_REPOSITORY$/com/google/code/findbugs/jsr305/3/jsr305-3.jar
@@ -46,7 +42,6 @@
- name: intellij.python.grazie
- name: intellij.python.langInjection
- name: intellij.python.markdown
- name: intellij.python.ml.features
- name: intellij.python.terminal
- name: lib/python-common.jar
modules:
@@ -52,7 +52,6 @@ The Python plug-in provides smart editing for Python scripts. The feature set of
<module name="intellij.python.grazie"/>
<module name="intellij.python.langInjection"/>
<module name="intellij.python.markdown"/>
<module name="intellij.python.ml.features"/>
<module name="intellij.python.terminal"/>
</content>
@@ -522,6 +521,7 @@ The Python plug-in provides smart editing for Python scripts. The feature set of
<statistics.counterUsagesCollector implementationClass="com.jetbrains.python.run.runAnything.PyRunAnythingCollector"/>
<statistics.counterUsagesCollector implementationClass="com.jetbrains.python.debugger.statistics.PyDataViewerCollector"/>
<statistics.counterUsagesCollector implementationClass="com.jetbrains.python.sdk.installer.BinaryInstallerUsagesCollector"/>
<statistics.counterUsagesCollector implementationClass="com.jetbrains.python.codeInsight.imports.mlapi.PyCharmImportsRankingLogs"/>
<!-- Code-insight IDE bridge -->
<applicationService serviceInterface="com.jetbrains.python.PythonRuntimeService"
@@ -30,6 +30,17 @@
<projectService serviceInterface="com.jetbrains.python.debugger.PySignatureCacheManager"
serviceImplementation="com.jetbrains.python.debugger.PySignatureCacheManagerImpl"/>
<!-- ML API -->
<platform.ml.featureProvider implementation="com.jetbrains.python.codeInsight.imports.mlapi.features.CandidatesListFeatures"/>
<platform.ml.featureProvider implementation="com.jetbrains.python.codeInsight.imports.mlapi.features.ImportCandidateRelativeFeatures"/>
<platform.ml.featureProvider implementation="com.jetbrains.python.codeInsight.imports.mlapi.features.PrimitiveImportFeatures"/>
<platform.ml.featureProvider implementation="com.jetbrains.python.codeInsight.imports.mlapi.features.PsiStructureFeatures"/>
<platform.ml.featureProvider implementation="com.jetbrains.python.codeInsight.imports.mlapi.features.RelevanceEvaluationFeatures"/>
<platform.ml.featureProvider implementation="com.jetbrains.python.codeInsight.imports.mlapi.features.ImportsFeatures"/>
<platform.ml.featureProvider implementation="com.jetbrains.python.codeInsight.imports.mlapi.features.OpenFilesImportsFeatures"/>
<platform.ml.featureProvider implementation="com.jetbrains.python.codeInsight.imports.mlapi.features.NeighborFilesImportsFeatures"/>
<platform.ml.featureProvider implementation="com.jetbrains.python.codeInsight.imports.mlapi.features.BaseProjectFeatures"/>
<stubIndex implementation="com.jetbrains.python.psi.stubs.PyClassNameIndex"/>
<stubIndex implementation="com.jetbrains.python.psi.stubs.PyClassNameIndexInsensitive"/>
<stubIndex implementation="com.jetbrains.python.psi.stubs.PyFunctionNameIndex"/>
@@ -10,10 +10,9 @@ import com.intellij.psi.PsiManager
import com.intellij.psi.search.FileTypeIndex
import com.intellij.psi.search.GlobalSearchScope
import com.intellij.psi.util.PsiTreeUtil
import com.jetbrains.ml.features.api.feature.*
import com.jetbrains.ml.*
import com.jetbrains.python.PythonFileType
import com.jetbrains.python.codeInsight.imports.mlapi.ImportRankingContext
import com.jetbrains.python.codeInsight.imports.mlapi.ImportRankingContextFeatures
import com.jetbrains.python.codeInsight.imports.mlapi.MLUnitImportCandidatesList
import com.jetbrains.python.psi.*
private val interestingClasses = arrayOf(
@@ -32,7 +31,7 @@ enum class FileExtensionType {
PXI // .pxi files
}
object BaseProjectFeatures : ImportRankingContextFeatures() {
class BaseProjectFeatures : FeatureProvider(MLUnitImportCandidatesList) {
object Features {
val NUM_PYTHON_FILES_IN_PROJECT = FeatureDeclaration.int("num_python_files_in_project") {
"The estimated amount of files in the project (by a power of 2)"
@@ -42,10 +41,10 @@ object BaseProjectFeatures : ImportRankingContextFeatures() {
}
override val featureComputationPolicy = FeatureComputationPolicy(false, true)
override val featureDeclarations = extractFeatureDeclarations(Features)
override val featureDeclarations = extractFieldsAsFeatureDeclarations(Features)
override suspend fun computeFeatures(instance: ImportRankingContext, filter: FeatureFilter): List<Feature>? = buildList {
val candidates = instance.candidates
override suspend fun computeFeatures(units: MLUnitsMap, usefulFeaturesFilter: FeatureFilter): List<Feature> = buildList {
val candidates = units[MLUnitImportCandidatesList]
if (candidates.isEmpty()) return@buildList
val project = candidates[0].importable?.project ?: return@buildList
@@ -1,14 +1,10 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.jetbrains.python.codeInsight.imports.mlapi.features
import com.jetbrains.ml.features.api.feature.Feature
import com.jetbrains.ml.features.api.feature.FeatureDeclaration
import com.jetbrains.ml.features.api.feature.FeatureFilter
import com.jetbrains.ml.features.api.feature.extractFeatureDeclarations
import com.jetbrains.python.codeInsight.imports.mlapi.ImportRankingContext
import com.jetbrains.python.codeInsight.imports.mlapi.ImportRankingContextFeatures
import com.jetbrains.ml.*
import com.jetbrains.python.codeInsight.imports.mlapi.MLUnitImportCandidatesList
object CandidatesListFeatures : ImportRankingContextFeatures() {
class CandidatesListFeatures : FeatureProvider(MLUnitImportCandidatesList) {
object Features {
val LENGTH = FeatureDeclaration.int("n_candidates") {
"The amount of import candidates"
@@ -18,10 +14,11 @@ object CandidatesListFeatures : ImportRankingContextFeatures() {
}.nullable()
}
override val featureDeclarations = extractFeatureDeclarations(Features)
override val featureDeclarations = extractFieldsAsFeatureDeclarations(Features)
override suspend fun computeFeatures(instance: ImportRankingContext, filter: FeatureFilter): List<Feature> = buildList {
add(Features.LENGTH with instance.candidates.size)
add(Features.HIGHEST_OLD_RELEVANCE with instance.candidates.maxOfOrNull { it.relevance })
override suspend fun computeFeatures(units: MLUnitsMap, usefulFeaturesFilter: FeatureFilter) = buildList<Feature> {
val candidates = units[MLUnitImportCandidatesList]
add(Features.LENGTH with candidates.size)
add(Features.HIGHEST_OLD_RELEVANCE with candidates.maxOfOrNull { it.relevance })
}
}
@@ -1,23 +1,27 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.jetbrains.python.codeInsight.imports.mlapi.features
import com.jetbrains.ml.features.api.feature.Feature
import com.jetbrains.ml.features.api.feature.FeatureDeclaration
import com.jetbrains.ml.features.api.feature.FeatureFilter
import com.jetbrains.ml.features.api.feature.extractFeatureDeclarations
import com.jetbrains.python.codeInsight.imports.mlapi.ImportCandidateContext
import com.jetbrains.python.codeInsight.imports.mlapi.ImportCandidateFeatures
import com.jetbrains.ml.*
import com.jetbrains.python.codeInsight.imports.mlapi.MLUnitImportCandidate
import com.jetbrains.python.codeInsight.imports.mlapi.MLUnitImportCandidatesList
object ImportCandidateRelativeFeatures : ImportCandidateFeatures() {
class ImportCandidateRelativeFeatures : FeatureProvider(MLUnitImportCandidate) {
object Features {
val RELATIVE_POSITION = FeatureDeclaration.int("original_position") {
"The import's position without ML ranking"
}
}
override val featureDeclarations = extractFeatureDeclarations(Features)
override val unitSight = setOf(
MLUnitImportCandidatesList
)
override suspend fun computeFeatures(instance: ImportCandidateContext, filter: FeatureFilter): List<Feature> = buildList<Feature> {
add(Features.RELATIVE_POSITION with instance.candidates.indexOf(instance.candidate))
override val featureDeclarations = extractFieldsAsFeatureDeclarations(Features)
override suspend fun computeFeatures(units: MLUnitsMap, usefulFeaturesFilter: FeatureFilter) = buildList<Feature> {
val candidate = units[MLUnitImportCandidate]
val allCandidates = units[MLUnitImportCandidatesList]
add(Features.RELATIVE_POSITION with allCandidates.indexOf(candidate))
}
}
@@ -5,16 +5,15 @@ import com.intellij.openapi.application.readAction
import com.intellij.openapi.fileEditor.FileEditorManager
import com.intellij.psi.PsiFile
import com.intellij.psi.PsiManager
import com.intellij.psi.util.QualifiedName
import com.jetbrains.ml.features.api.feature.*
import com.jetbrains.python.codeInsight.imports.mlapi.ImportCandidateContext
import com.jetbrains.python.codeInsight.imports.mlapi.ImportCandidateFeatures
import com.jetbrains.ml.*
import com.jetbrains.python.codeInsight.imports.mlapi.MLUnitImportCandidate
import com.jetbrains.python.psi.PyFile
import one.util.streamex.StreamEx
import com.intellij.psi.util.QualifiedName
import com.jetbrains.python.psi.PyFromImportStatement
import com.jetbrains.python.psi.PyImportElement
import one.util.streamex.StreamEx
object ImportsFeatures : ImportCandidateFeatures() {
class ImportsFeatures : FeatureProvider(MLUnitImportCandidate) {
object Features {
val EXISTING_IMPORT_FROM_PREFIX = FeatureDeclaration.int("existing_import_from_prefix") { "A maximal prefix for which there exists some import from" }.nullable()
val EXISTING_IMPORT_PREFIX = FeatureDeclaration.int("existing_import_prefix") { "A maximal prefix for which there exists some import" }.nullable()
@@ -24,11 +23,12 @@ object ImportsFeatures : ImportCandidateFeatures() {
override val featureComputationPolicy = FeatureComputationPolicy(false, true)
override val featureDeclarations = extractFeatureDeclarations(Features)
override val featureDeclarations = extractFieldsAsFeatureDeclarations(Features)
override suspend fun computeFeatures(units: MLUnitsMap, usefulFeaturesFilter: FeatureFilter) = buildList {
val importCandidate = units[MLUnitImportCandidate]
val project = importCandidate.importable?.project ?: return@buildList
override suspend fun computeFeatures(instance: ImportCandidateContext, filter: FeatureFilter): List<Feature> = buildList {
val importCandidate = instance.candidate
val project = readAction { importCandidate.importable?.project } ?: return@buildList
val editor = readAction { FileEditorManager.getInstance(project).selectedEditor } ?: return@buildList
val psiFile = readAction { PsiManager.getInstance(project).findFile(editor.file) } ?: return@buildList
@@ -5,16 +5,15 @@ import com.intellij.openapi.application.readAction
import com.intellij.openapi.fileEditor.FileEditorManager
import com.intellij.openapi.vfs.VirtualFile
import com.intellij.psi.PsiManager
import com.jetbrains.ml.features.api.feature.*
import com.jetbrains.python.codeInsight.imports.mlapi.ImportCandidateContext
import com.jetbrains.python.codeInsight.imports.mlapi.ImportCandidateFeatures
import com.jetbrains.ml.*
import com.jetbrains.python.codeInsight.imports.mlapi.MLUnitImportCandidate
import com.jetbrains.python.psi.PyFile
import kotlinx.coroutines.async
import kotlinx.coroutines.coroutineScope
private const val FILES_TO_WATCH = 8
object NeighborFilesImportsFeatures : ImportCandidateFeatures() {
class NeighborFilesImportsFeatures : FeatureProvider(MLUnitImportCandidate) {
object Features {
val NEIGHBOR_FILES_EXISTING_IMPORT_FROM_PREFIX = FeatureDeclaration.int("neighbor_files_existing_import_from_prefix") { "A maximal prefix for which there exists some import from" }.nullable()
val NEIGHBOR_FILES_EXISTING_IMPORT_PREFIX = FeatureDeclaration.int("neighbor_files_existing_import_prefix") { "A maximal prefix for which there exists some import" }.nullable()
@@ -24,10 +23,10 @@ object NeighborFilesImportsFeatures : ImportCandidateFeatures() {
override val featureComputationPolicy = FeatureComputationPolicy(false, true)
override val featureDeclarations = extractFeatureDeclarations(Features)
override val featureDeclarations = extractFieldsAsFeatureDeclarations(Features)
override suspend fun computeFeatures(instance: ImportCandidateContext, filter: FeatureFilter): List<Feature> = coroutineScope {
val importCandidate = instance.candidate
override suspend fun computeFeatures(units: MLUnitsMap, usefulFeaturesFilter: FeatureFilter): List<Feature> = coroutineScope {
val importCandidate = units[MLUnitImportCandidate]
if (importCandidate.path == null) {
return@coroutineScope buildList<Feature> {
add(Features.NEIGHBOR_FILES_EXISTING_IMPORT_FROM_PREFIX with 0)
@@ -4,16 +4,15 @@ package com.jetbrains.python.codeInsight.imports.mlapi.features
import com.intellij.openapi.application.readAction
import com.intellij.openapi.fileEditor.FileEditorManager
import com.intellij.psi.PsiManager
import com.jetbrains.ml.features.api.feature.*
import com.jetbrains.python.codeInsight.imports.mlapi.ImportCandidateContext
import com.jetbrains.python.codeInsight.imports.mlapi.ImportCandidateFeatures
import com.jetbrains.ml.*
import com.jetbrains.python.codeInsight.imports.mlapi.MLUnitImportCandidate
import com.jetbrains.python.psi.PyFile
import kotlinx.coroutines.async
import kotlinx.coroutines.coroutineScope
import kotlinx.coroutines.*
import kotlin.collections.copyOfRange
private const val FILES_TO_WATCH = 8
object OpenFilesImportsFeatures : ImportCandidateFeatures() {
class OpenFilesImportsFeatures : FeatureProvider(MLUnitImportCandidate) {
object Features {
val OPEN_FILES_EXISTING_IMPORT_FROM_PREFIX = FeatureDeclaration.int("open_files_existing_import_from_prefix") { "A maximal prefix for which there exists some import from" }.nullable()
val OPEN_FILES_EXISTING_IMPORT_PREFIX = FeatureDeclaration.int("open_files_existing_import_prefix") { "A maximal prefix for which there exists some import" }.nullable()
@@ -23,10 +22,10 @@ object OpenFilesImportsFeatures : ImportCandidateFeatures() {
override val featureComputationPolicy = FeatureComputationPolicy(false, true)
override val featureDeclarations = extractFeatureDeclarations(Features)
override val featureDeclarations = extractFieldsAsFeatureDeclarations(Features)
override suspend fun computeFeatures(instance: ImportCandidateContext, filter: FeatureFilter): List<Feature> = coroutineScope {
val importCandidate = instance.candidate
override suspend fun computeFeatures(units: MLUnitsMap, usefulFeaturesFilter: FeatureFilter) : List<Feature> = coroutineScope {
val importCandidate = units[MLUnitImportCandidate]
if (importCandidate.path == null) {
return@coroutineScope buildList<Feature> {
add(Features.OPEN_FILES_EXISTING_IMPORT_FROM_PREFIX with 0)
@@ -1,15 +1,11 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.jetbrains.python.codeInsight.imports.mlapi.features
import com.jetbrains.ml.features.api.feature.Feature
import com.jetbrains.ml.features.api.feature.FeatureDeclaration
import com.jetbrains.ml.features.api.feature.FeatureFilter
import com.jetbrains.ml.features.api.feature.extractFeatureDeclarations
import com.jetbrains.python.codeInsight.imports.mlapi.ImportCandidateContext
import com.jetbrains.python.codeInsight.imports.mlapi.ImportCandidateFeatures
import com.jetbrains.ml.*
import com.jetbrains.python.codeInsight.imports.mlapi.MLUnitImportCandidate
object PrimitiveImportFeatures : ImportCandidateFeatures() {
class PrimitiveImportFeatures : FeatureProvider(MLUnitImportCandidate) {
object Features {
val RELEVANCE = FeatureDeclaration.int("old_relevance") {
"""
@@ -21,10 +17,13 @@ object PrimitiveImportFeatures : ImportCandidateFeatures() {
}.nullable()
}
override val featureDeclarations = extractFeatureDeclarations(Features)
override val featureDeclarations = extractFieldsAsFeatureDeclarations(Features)
override suspend fun computeFeatures(instance: ImportCandidateContext, filter: FeatureFilter): List<Feature>? = buildList {
add(Features.RELEVANCE with instance.candidate.relevance)
add(Features.COMPONENT_COUNT with instance.candidate.path?.componentCount)
override suspend fun computeFeatures(units: MLUnitsMap, usefulFeaturesFilter: FeatureFilter) = buildList {
val importCandidate = units[MLUnitImportCandidate]
add(Features.RELEVANCE with importCandidate.relevance)
add(Features.COMPONENT_COUNT with importCandidate.path?.componentCount)
}
}
@@ -3,28 +3,26 @@ package com.jetbrains.python.codeInsight.imports.mlapi.features
import com.intellij.openapi.application.readAction
import com.intellij.psi.PsiElement
import com.jetbrains.ml.features.api.feature.Feature
import com.jetbrains.ml.features.api.feature.FeatureDeclaration
import com.jetbrains.ml.features.api.feature.FeatureFilter
import com.jetbrains.ml.features.api.feature.extractFeatureDeclarations
import com.jetbrains.python.codeInsight.imports.mlapi.ImportCandidateContext
import com.jetbrains.python.codeInsight.imports.mlapi.ImportCandidateFeatures
import com.jetbrains.ml.*
import com.jetbrains.python.codeInsight.imports.ImportCandidateHolder
import com.jetbrains.python.codeInsight.imports.mlapi.MLUnitImportCandidate
object PsiStructureFeatures : ImportCandidateFeatures() {
class PsiStructureFeatures : FeatureProvider(MLUnitImportCandidate) {
object Features {
val PSI_CLASS = FeatureDeclaration.aClass("importable_class") { "PSI class of the imported element" }.nullable()
val PSI_PARENT = (1..4).map { i -> FeatureDeclaration.aClass("psi_parent_$i") { "PSI parent #$i" }.nullable() }
}
override val featureDeclarations = extractFeatureDeclarations(Features)
override val featureDeclarations = extractFieldsAsFeatureDeclarations(Features)
override suspend fun computeFeatures(units: MLUnitsMap, usefulFeaturesFilter: FeatureFilter) = buildList {
val importCandidate: ImportCandidateHolder = units[MLUnitImportCandidate]
override suspend fun computeFeatures(instance: ImportCandidateContext, filter: FeatureFilter): List<Feature> = buildList {
readAction {
add(Features.PSI_CLASS with (instance.candidate.importable?.javaClass))
add(Features.PSI_CLASS with (importCandidate.importable?.javaClass))
Features.PSI_PARENT.withIndex().forEach { (i, featureDeclaration) ->
var parent: PsiElement? = instance.candidate.importable
var parent: PsiElement? = importCandidate.importable
repeat(i) {
parent = parent?.parent
}
@@ -6,9 +6,8 @@ import com.intellij.openapi.module.ModuleUtilCore
import com.intellij.openapi.projectRoots.Sdk
import com.intellij.openapi.vfs.VirtualFile
import com.intellij.psi.util.QualifiedName
import com.jetbrains.ml.features.api.feature.*
import com.jetbrains.python.codeInsight.imports.mlapi.ImportCandidateContext
import com.jetbrains.python.codeInsight.imports.mlapi.ImportCandidateFeatures
import com.jetbrains.ml.*
import com.jetbrains.python.codeInsight.imports.mlapi.MLUnitImportCandidate
import com.jetbrains.python.sdk.PythonSdkUtil
enum class UnderscoresType {
@@ -21,7 +20,6 @@ enum class UnderscoresType {
DOUBLE_LEADING_TRAILING, // __something__
IRREGULAR // not from the above
}
enum class ModuleSourceType {
STD_LIB,
EXTERNAL_LIB,
@@ -29,7 +27,9 @@ enum class ModuleSourceType {
}
object RelevanceEvaluationFeatures : ImportCandidateFeatures() {
class RelevanceEvaluationFeatures : FeatureProvider(MLUnitImportCandidate) {
object Features {
val UNDERSCORES_IN_PATH = FeatureDeclaration.int("underscores_in_path") { "number of prefix and suffix underscores in path" }.nullable()
val MODULE_SOURCE_TYPE = FeatureDeclaration.enum<ModuleSourceType>("module_source_type") { "info about lib being std, local, or external" }.nullable()
@@ -38,10 +38,10 @@ object RelevanceEvaluationFeatures : ImportCandidateFeatures() {
override val featureComputationPolicy = FeatureComputationPolicy(false, true)
override val featureDeclarations = extractFeatureDeclarations(Features)
override val featureDeclarations = extractFieldsAsFeatureDeclarations(Features)
override suspend fun computeFeatures(instance: ImportCandidateContext, filter: FeatureFilter): List<Feature> = buildList {
val importCandidate = instance.candidate
override suspend fun computeFeatures(units: MLUnitsMap, usefulFeaturesFilter: FeatureFilter) = buildList {
val importCandidate = units[MLUnitImportCandidate]
add(Features.UNDERSCORES_IN_PATH with countBoundaryUnderscores(importCandidate.path))
readAction {
Features.UNDERSCORES_TYPES_OF_PACKAGES.withIndex().forEach { (i, featureDeclaration) ->
@@ -64,10 +64,7 @@ object RelevanceEvaluationFeatures : ImportCandidateFeatures() {
ModuleUtilCore.findModuleForFile(vFile, baseElement.project) == null -> ModuleSourceType.EXTERNAL_LIB
else -> ModuleSourceType.LOCAL_LIB
}
}
)
} else {
add(Features.MODULE_SOURCE_TYPE with null)
})
}
}
@@ -75,7 +72,6 @@ object RelevanceEvaluationFeatures : ImportCandidateFeatures() {
if (qName == null) return 0
return qName.components.sumOf { it.takeWhile { it == '_' }.length + it.takeLastWhile { it == '_' }.length }
}
private fun getUnderscoresType(s: String): UnderscoresType {
val leadingUnderscores = s.takeWhile { it == '_' }.length
val trailingUnderscores = s.reversed().takeWhile { it == '_' }.length
@@ -1,25 +1,9 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.jetbrains.python.codeInsight.imports.mlapi
import com.jetbrains.ml.features.api.MLUnit
import com.jetbrains.ml.features.api.feature.FeatureProvider
import com.jetbrains.ml.MLUnit
import com.jetbrains.python.codeInsight.imports.ImportCandidateHolder
val MLUnitImportCandidate = MLUnit<ImportCandidateHolder>("import_candidate")
data class ImportRankingContext(
val candidates: List<ImportCandidateHolder>
)
data class ImportCandidateContext(
val candidates: List<ImportCandidateHolder>,
val candidate: ImportCandidateHolder,
)
val MLUnitImportRankingContext = MLUnit<ImportRankingContext>("import_candidates_list")
val MLUnitImportCandidate = MLUnit<ImportCandidateContext>("import_candidate")
abstract class ImportRankingContextFeatures : FeatureProvider<ImportRankingContext>(MLUnitImportRankingContext)
abstract class ImportCandidateFeatures : FeatureProvider<ImportCandidateContext>(MLUnitImportCandidate)
val MLUnitImportCandidatesList = MLUnit<List<ImportCandidateHolder>>("import_candidates_list")
@@ -7,13 +7,13 @@ import com.intellij.openapi.application.ApplicationManager;
import com.intellij.openapi.editor.colors.EditorColorsManager;
import com.intellij.openapi.editor.colors.EditorColorsScheme;
import com.intellij.openapi.editor.colors.EditorFontType;
import com.intellij.openapi.ui.popup.IPopupChooserBuilder;
import com.intellij.openapi.ui.popup.JBPopupFactory;
import com.intellij.psi.PsiElement;
import com.intellij.ui.SimpleColoredComponent;
import com.intellij.ui.SimpleTextAttributes;
import com.intellij.util.Consumer;
import com.jetbrains.python.PyPsiBundle;
import kotlin.Unit;
import org.jetbrains.concurrency.AsyncPromise;
import org.jetbrains.concurrency.Promise;
import org.jetbrains.concurrency.Promises;
@@ -22,34 +22,39 @@ import javax.swing.*;
import java.awt.*;
import java.util.List;
public class PyImportChooser implements ImportChooser {
import static com.jetbrains.python.codeInsight.imports.mlapi.MlImplementationKt.launchMLRanking;
public final class PyImportChooser implements ImportChooser {
@Override
public Promise<ImportCandidateHolder> selectImport(List<? extends ImportCandidateHolder> sources, boolean useQualifiedImport) {
if (ApplicationManager.getApplication().isUnitTestMode()) {
return Promises.resolvedPromise(sources.get(0));
}
AsyncPromise<ImportCandidateHolder> result = new AsyncPromise<>();
// GUI part
DataManager.getInstance().getDataContextFromFocus().doWhenDone((Consumer<DataContext>)dataContext -> {
var popup = JBPopupFactory.getInstance()
.createPopupChooserBuilder(sources)
.setItemChosenCallback(item -> {
result.setResult(item);
});
processPopup(popup, useQualifiedImport);
popup.createPopup()
.showInBestPositionFor(dataContext);
launchMLRanking(sources, (mlRanking) -> {
// GUI part
DataManager.getInstance().getDataContextFromFocus().doWhenDone((Consumer<DataContext>)dataContext -> JBPopupFactory.getInstance()
.createPopupChooserBuilder(mlRanking.getOrder())
.setRenderer(new CellRenderer())
.setTitle(useQualifiedImport ? PyPsiBundle.message("ACT.qualify.with.module") : PyPsiBundle.message("ACT.from.some.module.import"))
.setItemChosenCallback(item -> {
result.setResult(item);
mlRanking.submitSelectedItem(item);
})
.setCancelCallback(() -> {
mlRanking.submitPopUpClosed();
return true;
})
.setNamerForFiltering(o -> o.getPresentableText())
.createPopup()
.showInBestPositionFor(dataContext));
return Unit.INSTANCE;
});
return result;
}
protected void processPopup(IPopupChooserBuilder<? extends ImportCandidateHolder> popup, boolean useQualifiedImport) {
popup.setRenderer(new CellRenderer())
.setTitle(useQualifiedImport ? PyPsiBundle.message("ACT.qualify.with.module") : PyPsiBundle.message("ACT.from.some.module.import"))
.setNamerForFiltering(o -> o.getPresentableText());
return result;
}
// Stolen from FQNameCellRenderer
@@ -1,5 +1,5 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.intellij.python.ml.features.imports
package com.jetbrains.python.codeInsight.imports.mlapi
import com.intellij.internal.statistic.eventLog.EventLogConfiguration
import com.intellij.openapi.components.Service
@@ -0,0 +1,42 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.jetbrains.python.codeInsight.imports.mlapi
import com.jetbrains.ml.analysis.MLTaskFeedback
import com.jetbrains.ml.analysis.MLTaskFeedbackField
import com.jetbrains.ml.analysis.MLUnitFeedbackField
import com.jetbrains.ml.logs.schema.BooleanEventField
import com.jetbrains.ml.logs.schema.EnumEventField
import com.jetbrains.ml.logs.schema.IntEventField
import com.jetbrains.ml.logs.schema.LongEventField
// suspendable analysis -- written after the imports were ranked
internal val MLFieldCorrectElement = MLUnitFeedbackField(
unit = MLUnitImportCandidate,
field = BooleanEventField("is_correct") { "The candidate was chosen" },
throwOnTimeout = false
)
internal object MLFeedbackCorrectElementPosition : MLTaskFeedback() {
val CANCELLED = BooleanEventField("selection_cancelled") { "No item has been selected" }
val SELECTED_POSITION = IntEventField("selected_position") { "The position of the selected import statement" }
val INVALID_SELECTED_ITEM = BooleanEventField("invalid_selected_item") { "The selected item hasn't been found in the initial list" }
override val feedbackDeclaration = listOf(CANCELLED, SELECTED_POSITION, INVALID_SELECTED_ITEM)
}
internal val ML_FEEDBACK_TIME_MS_TO_DISPLAY = MLTaskFeedbackField(
field = LongEventField("time_ms_before_displayed") { "Duration from the quickfix start until when the imports were displayed" },
throwOnTimeout = true,
)
internal val ML_FEEDBACK_TIME_MS_BEFORE_CLOSED = MLTaskFeedbackField(
field = LongEventField("time_ms_before_closed") { "Duration from the quickfix start until the pop-up was closed" },
)
// runtime analysis -- written during the imports are ranked
internal val ML_LOGGING_STATE: EnumEventField<LoggingOption> = EnumEventField.of<LoggingOption>("ml_logging_state", { "State of the ML session logging" })
internal val ML_ENABLED = BooleanEventField("ml_enabled") { "Machine Learning ranking is enabled" }
internal val ML_STATUS_CORRESPONDS_TO_BUCKET = BooleanEventField("ml_status_corresponds_to_bucket") { "Field 'ml_enabled' corresponds to the ML bucket % 2" }
@@ -0,0 +1,136 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.jetbrains.python.codeInsight.imports.mlapi
import com.intellij.codeInsight.inline.completion.logs.TypingSpeedTracker
import com.intellij.codeInsight.inline.completion.ml.MLUnitTyping
import com.intellij.openapi.application.writeAction
import com.intellij.openapi.components.Service
import com.intellij.openapi.components.service
import com.intellij.util.application
import com.jetbrains.ml.MockMLModel
import com.jetbrains.ml.platform.MLApiTaskExecutor
import com.jetbrains.ml.session.MLSession
import com.jetbrains.python.codeInsight.imports.ImportCandidateHolder
import kotlinx.coroutines.CoroutineScope
import kotlinx.coroutines.Dispatchers
import kotlinx.coroutines.delay
import kotlinx.coroutines.launch
import java.util.concurrent.atomic.AtomicBoolean
import kotlin.time.Duration.Companion.seconds
@Service(Service.Level.APP)
private class MLApiComputations(
val coroutineScope: CoroutineScope,
)
internal sealed class FinalCandidatesRanker(
protected val mlSession: MLSession<MockMLModel, Unit>,
protected val timestampStarted: Long,
) {
abstract fun launchMLRanking(initialCandidatesOrder: MutableList<out ImportCandidateHolder>, displayResult: (RateableRankingResult) -> Unit)
abstract val mlEnabled: Boolean
}
private class ExperimentalMLRanker(
private val coroutineScope: CoroutineScope,
mlSession: MLSession<MockMLModel, Unit>, timestampStarted: Long,
) : FinalCandidatesRanker(mlSession, timestampStarted) {
override val mlEnabled = true
override fun launchMLRanking(initialCandidatesOrder: MutableList<out ImportCandidateHolder>, displayResult: (RateableRankingResult) -> Unit) {
coroutineScope.launch(Dispatchers.Default) {
val mlTreeRoot = mlSession.buildRoot(
MLUnitImportCandidatesList with initialCandidatesOrder,
MLUnitTyping with TypingSpeedTracker.getInstance(),
).computingFeatures { computeAllFeaturesNow() }
for (candidate in initialCandidatesOrder) {
mlTreeRoot
.onlyLog(MLUnitImportCandidate with candidate)
.computingFeatures { computeAllFeaturesNow() }
}
mlSession.finish()
writeAction {
displayResult(RateableRankingResult(mlSession, initialCandidatesOrder, timestampStarted))
}
}
}
}
private class InitialOrderKeepingRanker(
mlSession: MLSession<MockMLModel, Unit>, timestampStarted: Long,
) : FinalCandidatesRanker(mlSession, timestampStarted) {
override val mlEnabled = false
override fun launchMLRanking(initialCandidatesOrder: MutableList<out ImportCandidateHolder>, displayResult: (RateableRankingResult) -> Unit) {
mlSession.finish()
application.runWriteAction {
displayResult(RateableRankingResult(mlSession, initialCandidatesOrder, timestampStarted))
}
}
}
internal fun launchMLRanking(initialCandidatesOrder: MutableList<out ImportCandidateHolder>, displayResult: (RateableRankingResult) -> Unit) {
// In EDT
val timestampStarted = System.currentTimeMillis()
val mlCoroutineScope = service<MLApiComputations>().coroutineScope
val mlSession = service<MLTaskPyCharmImportStatementsRanking>().task
.startMLSession(taskExecutor = MLApiTaskExecutor.configure(mlCoroutineScope))
val rankingStatus = service<FinalImportRankingStatusService>().status
val ranker = if (rankingStatus.mlEnabled)
ExperimentalMLRanker(mlCoroutineScope, mlSession, timestampStarted)
else
InitialOrderKeepingRanker(mlSession, timestampStarted)
mlSession.writeRuntimeSessionAnalysis(ML_LOGGING_STATE with getLoggingOption(mlSession))
mlSession.writeRuntimeSessionAnalysis(ML_ENABLED with rankingStatus.mlEnabled)
mlSession.writeRuntimeSessionAnalysis(ML_STATUS_CORRESPONDS_TO_BUCKET with rankingStatus.mlStatusCorrespondsToBucket)
ranker.launchMLRanking(initialCandidatesOrder) { result ->
val timestampDisplayed = System.currentTimeMillis()
displayResult(result)
ML_FEEDBACK_TIME_MS_TO_DISPLAY.feedback(mlSession, timestampDisplayed - timestampStarted)
}
}
internal class RateableRankingResult(
private val mlSession: MLSession<*, *>,
val order: List<ImportCandidateHolder>,
private val timestampStarted: Long,
) {
private var submitted = AtomicBoolean(false)
fun submitPopUpClosed() {
val timestampClosed = System.currentTimeMillis()
ML_FEEDBACK_TIME_MS_BEFORE_CLOSED.feedback(mlSession, timestampClosed - timestampStarted)
service<MLApiComputations>().coroutineScope.launch {
delay(1.seconds)
synchronized(this@RateableRankingResult) {
if (submitted.get()) return@launch
submitted.set(true)
MLFeedbackCorrectElementPosition.feedbackEventPairs(mlSession, MLFeedbackCorrectElementPosition.CANCELLED with true)
MLFieldCorrectElement.cancelFeedback(mlSession)
}
}
}
fun submitSelectedItem(selected: ImportCandidateHolder) = synchronized(this) {
// In EDT
require(submitted.compareAndSet(false, true)) { "Some feedback to the ml ranker was already submitted" }
MLFieldCorrectElement.feedback(mlSession, selected, true)
val selectedIndex = order.indexOf(selected)
if (selectedIndex == -1)
MLFeedbackCorrectElementPosition.feedbackEventPairs(mlSession, MLFeedbackCorrectElementPosition.INVALID_SELECTED_ITEM with true)
else
MLFeedbackCorrectElementPosition.feedbackEventPairs(mlSession, MLFeedbackCorrectElementPosition.SELECTED_POSITION with selectedIndex)
}
}
@@ -0,0 +1,80 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.jetbrains.python.codeInsight.imports.mlapi
import com.intellij.internal.statistic.eventLog.EventLogGroup
import com.intellij.internal.statistic.service.fus.collectors.CounterUsagesCollector
import com.intellij.openapi.components.Service
import com.intellij.openapi.components.service
import com.intellij.platform.ml.impl.logs.MLEventLoggerProvider.Companion.ML_RECORDER_ID
import com.intellij.platform.ml.impl.tools.registerMLTaskLogging
import com.intellij.util.application
import com.jetbrains.ml.MLUnit
import com.jetbrains.ml.logs.*
import com.jetbrains.ml.logs.schema.BooleanEventField
import com.jetbrains.ml.logs.schema.EventField
import com.jetbrains.ml.logs.schema.EventPair
import com.jetbrains.ml.model.MLModel
import com.jetbrains.ml.session.MLSession
import com.jetbrains.ml.session.MLSessionInfo
import com.jetbrains.ml.tree.LevelFeaturesSchema
import com.jetbrains.ml.tree.MLTree
import com.jetbrains.python.PythonPluginDisposable
@Service
private class EventLogGroupHolder {
val group by lazy {
EventLogGroup("pycharm.quickfix.imports", 3, ML_RECORDER_ID).also {
it.registerMLTaskLogging(service<MLTaskPyCharmImportStatementsRanking>().task,
parentDisposable = service<PythonPluginDisposable>())
}
}
}
class PyCharmImportsRankingLogs : CounterUsagesCollector() {
override fun getGroup() = service<EventLogGroupHolder>().group
}
internal enum class LoggingOption {
FULL,
NO_TREE,
SKIP
}
internal const val RELEASE_LOGGING_PERCENT = 5
internal fun getLoggingOption(mlSession: MLSession<*, *>): LoggingOption {
if (application.isEAP || application.isUnitTestMode) {
return LoggingOption.FULL
}
if (mlSession.id % 100 > RELEASE_LOGGING_PERCENT) return LoggingOption.SKIP
return LoggingOption.NO_TREE
}
object PyCharmImportsRankingLogger : MLSessionLoggerProvider<Any> {
private val baseLoggerProvider = EntireSessionLoggerProvider<Any, Boolean>(BooleanEventField("prediction") { "ML model prediction" }) { null }
override fun createMLSessionLogger(
eventPrefix: String,
taskId: String,
treeAnalysisDeclaration: Map<MLUnit<*>, List<EventField<*>>>,
sessionAnalysisDeclaration: List<EventField<*>>,
featuresDeclaration: List<LevelFeaturesSchema>,
fusEventRegister: FusEventRegister,
): MLSessionLogger<Any> {
val baseSessionLogger = baseLoggerProvider.createMLSessionLogger(eventPrefix, taskId, treeAnalysisDeclaration, sessionAnalysisDeclaration, featuresDeclaration, fusEventRegister)
return object : MLSessionLogger<Any> {
override suspend fun logMLSession(sessionInfo: MLSessionInfo<out MLModel<Any>, Any>, analysisException: Throwable?, sessionAnalysis: List<EventPair<*>>, structure: MLTree.ATopNode<out MLModel<Any>, Any>?): MLSessionLoggingOutcome {
val loggingOption = checkNotNull(sessionAnalysis.find { it.field == ML_LOGGING_STATE }).data as LoggingOption
return when (loggingOption) {
LoggingOption.FULL -> baseSessionLogger.logMLSession(sessionInfo, analysisException, sessionAnalysis, structure)
LoggingOption.NO_TREE -> {
baseSessionLogger.logMLSession(sessionInfo, analysisException, sessionAnalysis, null)
MLSessionLoggingOutcome.Custom("no tree logged")
}
LoggingOption.SKIP -> MLSessionLoggingOutcome.FilteredOut
}
}
}
}
}
@@ -0,0 +1,37 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.jetbrains.python.codeInsight.imports.mlapi
import com.intellij.codeInsight.inline.completion.ml.MLUnitTyping
import com.intellij.openapi.components.Service
import com.intellij.openapi.components.service
import com.intellij.platform.ml.impl.tools.IJPlatform
import com.jetbrains.ml.MLTaskBuilder
import kotlin.time.Duration.Companion.days
internal val LONGEST_POPUP_LIFE_DURATION = 365.days
@Service
class MLTaskPyCharmImportStatementsRanking {
val task by lazy {
MLTaskBuilder(
taskId = "pycharm_import_statements_ranking",
taskLevels = listOf(
setOf(MLUnitImportCandidatesList, MLUnitTyping),
setOf(MLUnitImportCandidate)
),
platform = service<IJPlatform>()
).buildWithoutMLModel {
suspendableTreeAnalysis(MLFieldCorrectElement)
suspendableSessionAnalysis(MLFeedbackCorrectElementPosition)
suspendableSessionAnalysis(ML_FEEDBACK_TIME_MS_TO_DISPLAY)
suspendableSessionAnalysis(ML_FEEDBACK_TIME_MS_BEFORE_CLOSED)
runtimeSessionAnalysis(ML_LOGGING_STATE, ML_ENABLED, ML_STATUS_CORRESPONDS_TO_BUCKET)
logger = PyCharmImportsRankingLogger
maxAnalysisDuration = LONGEST_POPUP_LIFE_DURATION
}
}
}