mirror of
https://gitflic.ru/project/openide/openide.git
synced 2026-09-27 10:03:11 +07:00
Revert "MLP-33 IJ Imports Ranking"
This reverts commit 47140fd0301283a10966e14c65df9a08d128ec39. GitOrigin-RevId: 81fa4946bd76885bcc599dfae9bfb9feaa11fbb3
This commit is contained in:
committed by
intellij-monorepo-bot
parent
313c1619d2
commit
44b9d24ae8
@@ -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
@@ -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>
|
||||
@@ -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
@@ -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
@@ -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>
|
||||
Generated
-1
@@ -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"),
|
||||
|
||||
-1
@@ -30,7 +30,6 @@ object PythonCommunityPluginModules {
|
||||
"intellij.python.pydev",
|
||||
"intellij.python.sdk",
|
||||
"intellij.python.terminal",
|
||||
"intellij.python.ml.features",
|
||||
)
|
||||
|
||||
/**
|
||||
|
||||
@@ -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() }
|
||||
}
|
||||
+28
-27
@@ -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)
|
||||
}
|
||||
+2
-3
@@ -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) }
|
||||
}
|
||||
}
|
||||
}
|
||||
+11
@@ -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 }
|
||||
}
|
||||
}
|
||||
|
||||
+26
@@ -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>
|
||||
-47
@@ -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
|
||||
}
|
||||
}
|
||||
-26
@@ -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
|
||||
-24
@@ -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" }
|
||||
}
|
||||
-202
@@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
-44
@@ -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
|
||||
}
|
||||
-93
@@ -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,
|
||||
)
|
||||
}
|
||||
-29
@@ -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>
|
||||
@@ -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"/>
|
||||
|
||||
+6
-7
@@ -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
|
||||
|
||||
+8
-11
@@ -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 })
|
||||
}
|
||||
}
|
||||
|
||||
+14
-10
@@ -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))
|
||||
}
|
||||
}
|
||||
|
||||
+10
-10
@@ -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
|
||||
|
||||
|
||||
+6
-7
@@ -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)
|
||||
|
||||
+8
-9
@@ -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)
|
||||
|
||||
+10
-11
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
+9
-11
@@ -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
|
||||
}
|
||||
|
||||
+9
-13
@@ -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
|
||||
|
||||
+3
-19
@@ -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
-1
@@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user