diff --git a/platform/ml-api/src/com/intellij/platform/ml/Feature.kt b/platform/ml-api/src/com/intellij/platform/ml/Feature.kt index 2e9d721b22a0..31474a48de18 100644 --- a/platform/ml-api/src/com/intellij/platform/ml/Feature.kt +++ b/platform/ml-api/src/com/intellij/platform/ml/Feature.kt @@ -87,4 +87,10 @@ sealed class Feature { : TypedFeature(name, value) { override val valueType = FeatureValueType.Version } + + companion object { + fun Feature.toCompactString(): String { + return "${this.declaration.name}=${this.value}" + } + } } diff --git a/platform/ml-api/src/com/intellij/platform/ml/TierDescriptor.kt b/platform/ml-api/src/com/intellij/platform/ml/TierDescriptor.kt index fb1c350bcce5..06702c50f9d4 100644 --- a/platform/ml-api/src/com/intellij/platform/ml/TierDescriptor.kt +++ b/platform/ml-api/src/com/intellij/platform/ml/TierDescriptor.kt @@ -44,7 +44,7 @@ interface TierDescriptor : TierRequester { * @param usefulFeaturesFilter Accepts features, that could make any difference to compute this time. * A feature is considered to be useful if an ML model is aware of the feature, * or it is explicitly said that "ML model is not aware of this feature, but it must be logged" - * (@see [com.intellij.platform.ml.impl.approach.LogDrivenModelInference]). + * (@see [com.intellij.platform.ml.impl.LogDrivenModelInference]). * * @throws IllegalArgumentException If a feature is missing: it is accepted by [usefulFeaturesFilter] * and it is declared, but not present in the result. @@ -73,13 +73,23 @@ interface TierDescriptor : TierRequester { * @param usefulFeaturesFilter Accepts features, that make sense to calculate at the time of this invocation. * If the filter does not accept a feature, it means that its computation will not make any difference. */ - fun couldBeUseful(usefulFeaturesFilter: FeatureFilter): Boolean { - return descriptionDeclaration.any { usefulFeaturesFilter.accept(it) } - } + fun couldBeUseful(usefulFeaturesFilter: FeatureFilter): Boolean companion object { val EP_NAME: ExtensionPointName = ExtensionPointName.create("com.intellij.platform.ml.descriptor") } + + abstract class Default(final override val tier: Tier<*>) : TierDescriptor { + override val additionallyRequiredTiers: Set> + get() = emptySet() + + override val descriptionPolicy: DescriptionPolicy + get() = DescriptionPolicy(false, false) + + override fun couldBeUseful(usefulFeaturesFilter: FeatureFilter): Boolean { + return descriptionDeclaration.any { usefulFeaturesFilter.accept(it) } + } + } } @ApiStatus.Internal diff --git a/platform/ml-impl/intellij.platform.ml.impl.iml b/platform/ml-impl/intellij.platform.ml.impl.iml index 9de0fd35792a..9ede63710bce 100644 --- a/platform/ml-impl/intellij.platform.ml.impl.iml +++ b/platform/ml-impl/intellij.platform.ml.impl.iml @@ -47,5 +47,6 @@ + \ No newline at end of file diff --git a/platform/ml-impl/resources/META-INF/ml.xml b/platform/ml-impl/resources/META-INF/ml.xml index 48e562829999..36cd166cdf71 100644 --- a/platform/ml-impl/resources/META-INF/ml.xml +++ b/platform/ml-impl/resources/META-INF/ml.xml @@ -6,9 +6,6 @@ - - @@ -23,7 +20,7 @@ interface="com.intellij.platform.ml.impl.turboComplete.SmartPipelineRunner" dynamic="true"/> , P : Any>( + final override val task: MLTask

, + final override val apiPlatform: MLApiPlatform, + private val sessionDetails: SessionDetails, + private val thisBuilder: MLTaskApproachBuilder

, +) : MLTaskApproach

{ + private val sessionLevels = sessionDetails.getLevels(task) + + interface SessionDetails, P : Any> { + fun interface Builder, P : Any> { + fun build(apiPlatform: MLApiPlatform): SessionDetails + } + + /** + * The analyzers, that are analyzing session's tree-like structure, giving additional features, + * not used by the ML model, to the nodes. + */ + val structureAnalysers: Collection> + + /** + * The analyzers, that are giving features to the model, that was used during the session. + */ + val mlModelAnalysers: Collection> + + /** + * Provides an ML model to use during session's lifetime. + */ + val mlModelProvider: MLModel.Provider + + /** + * Performs description's computation. + * Could perform caching mechanisms to avoid recomputing features every time. + */ + val descriptionComputer: DescriptionComputer + + /** + * Tiers that do not make a part of te [task], but they could be described and passed to the ML model. + * + * The size of this list must correspond to the number of levels in the solved [task]. + */ + val additionallyDescribedTiers: List>> + + /** + * Declares features, that are not used by the ML model, but must be computed anyway, + * so they make it to logs. + * + * A feature cannot be simultaneously declared as "not used description" and as used by the [mlModelProvider]'s + * provided model. + * If a feature is not declared as "not used but still computed" or as "used by the model", then it will be computed. + * + * It must contain explicitly declared selectors for each tier used in [task], as well as in [additionallyDescribedTiers]. + * + * @param callParameters Contains any additional parameters that you passed in [com.intellij.platform.ml.impl.MLTaskApproach.startMLSession] + * The tier set is equal to the one that was declared in [com.intellij.platform.ml.impl.MLTask.callParameters]'s first level. + * @param mlModel The model that was acquired by the [mlModelProvider] + */ + fun getNotUsedDescription(callParameters: Environment, mlModel: M): PerTier + + companion object { + fun , P : Any> SessionDetails.getLevels(task: MLTask

): List>, Set>>> { + return (task.levels zip additionallyDescribedTiers).map { LevelSignature(it.first, it.second) } + } + + fun , P : Any> SessionDetails.validate(task: MLTask

) { + require(task.levels.size == additionallyDescribedTiers.size) { + "Task $task has ${task.levels.size} levels, when 'additionallyDescribedTiers' has ${additionallyDescribedTiers.size}" + } + + val maybeDuplicatedTaskTiers: List> = getLevels(task).flatMap { it.main.toList() + it.additional } + val taskTiers: Set> = maybeDuplicatedTaskTiers.toSet() + + require(maybeDuplicatedTaskTiers.size == taskTiers.size) { + "There are duplicated tiers in the declaration: ${maybeDuplicatedTaskTiers.groupBy { it }.filter { it.value.size > 1 }.keys}" + } + } + } + } + + open class Builder, P : Any>( + override val task: MLTask

, + private val details: SessionDetails.Builder, + ) : MLTaskApproachBuilder

{ + final override fun buildApproach(apiPlatform: MLApiPlatform): MLTaskApproach

{ + val sessionDetails = details.build(apiPlatform) + sessionDetails.validate(task) + val modelInference = buildModelInference(apiPlatform, sessionDetails) + return modelInference + } + + open fun buildModelInference(apiPlatform: MLApiPlatform, sessionDetails: SessionDetails): LogDrivenModelInference { + return LogDrivenModelInference(task, apiPlatform, sessionDetails, this) + } + + override fun buildApproachSessionDeclaration(apiPlatform: MLApiPlatform): MLTaskApproach.SessionDeclaration { + val sessionDetails = details.build(apiPlatform) + sessionDetails.validate(task) + val sessionAnalyser = MLSessionAnalyser(sessionDetails.structureAnalysers, sessionDetails.mlModelAnalysers) + val levels = (task.levels zip sessionDetails.additionallyDescribedTiers).map { LevelSignature(it.first, it.second) } + return MLTaskApproach.SessionDeclaration( + sessionFeatures = sessionAnalyser.sessionAnalysisDeclaration, + levelsScheme = levels.map { levelTiers -> + LevelSignature( + buildMainTiersScheme(sessionAnalyser, levelTiers.main, apiPlatform), + buildAdditionalTiersScheme(levelTiers.additional, apiPlatform), + ) + } + ) + } + + private fun buildTierDescriptionDeclaration(tierDescriptors: Collection): Set> { + return tierDescriptors.flatMap { + if (it is ObsoleteTierDescriptor) it.partialDescriptionDeclaration else it.descriptionDeclaration + }.toSet() + } + + private fun buildMainTiersScheme(sessionAnalyser: MLSessionAnalyser, tiers: Set>, apiEnvironment: MLApiPlatform): PerTier { + val tiersDescriptors = apiEnvironment.getDescriptorsOfTiers(tiers) + + return tiers.associateWith { tier -> + val tierDescriptors = tiersDescriptors.getValue(tier) + val descriptionDeclaration = buildTierDescriptionDeclaration(tierDescriptors) + val analysisDeclaration = sessionAnalyser.structureAnalysisDeclaration[tier] ?: emptySet() + MainTierScheme(descriptionDeclaration, analysisDeclaration) + } + } + + private fun buildAdditionalTiersScheme(tiers: Set>, apiEnvironment: MLApiPlatform): PerTier { + val tiersDescriptors = apiEnvironment.getDescriptorsOfTiers(tiers) + + return tiers.associateWith { tier -> + val tierDescriptors = tiersDescriptors.getValue(tier) + val descriptionDeclaration = buildTierDescriptionDeclaration(tierDescriptors) + AdditionalTierScheme(descriptionDeclaration) + } + } + } + + override suspend fun startSession(callParameters: Environment, permanentSessionEnvironment: Environment): Session.StartOutcome

{ + return startSessionMonitoring(callParameters, permanentSessionEnvironment) + } + + private suspend fun startSessionMonitoring(callParameters: Environment, permanentSessionEnvironment: Environment): Session.StartOutcome

{ + val approachListener = apiPlatform.getJoinedListenerForTask(thisBuilder, permanentSessionEnvironment) + try { + return acquireModelAndStartSession(callParameters, permanentSessionEnvironment, approachListener) + } + catch (e: Throwable) { + approachListener.onFailedToStartSessionWithException(e) + return Session.StartOutcome.UncaughtException(e) + } + } + + private suspend fun acquireModelAndStartSession(unsafeCallParameters: Environment, + permanentSessionEnvironment: Environment, + approachListener: MLApproachListener): Session.StartOutcome

{ + val callParameters = unsafeCallParameters.narrowedTo(task.callParameters.first()) + val extendedPermanentSessionEnvironment = Environment.joined(callParameters, permanentSessionEnvironment) + + val mlModel: M = run { + val nullableMlModel = sessionDetails.mlModelProvider.provideModel(callParameters, extendedPermanentSessionEnvironment, sessionLevels) + if (nullableMlModel == null) { + val failure = ModelNotAcquiredOutcome

() + approachListener.onFailedToStartSession(failure) + return failure + } + nullableMlModel + } + + var sessionListener: MLSessionListener? = null + + val sessionAnalyser = MLSessionAnalyser(sessionDetails.structureAnalysers, sessionDetails.mlModelAnalysers) + + val analyseThenLogStructure = SessionTreeHandler, M, P> { treeRoot -> + sessionListener?.onSessionDescriptionFinished(treeRoot) + sessionAnalyser.analyseSessionStructure(treeRoot).thenApplyAsync { analysedSession -> + sessionListener?.onSessionAnalysisFinished(analysedSession) + }.exceptionally { + thisLogger().error(it) + } + } + + val notUsedDescription = sessionDetails.getNotUsedDescription(callParameters, mlModel) + validateNotUsedDescription(notUsedDescription) + + val session = if (sessionLevels.size == 1) { + val collector = SolitaryLeafCollector.build( + callParameters, task, sessionDetails.descriptionComputer, notUsedDescription, + permanentSessionEnvironment, sessionLevels.first().additional, mlModel, apiPlatform, sessionLevels.first(), mlModel.knownFeatures, + ) + collector.handleCollectedTree(analyseThenLogStructure) + MLModelPrediction(mlModel, collector) + } + else { + val collector = RootCollector.build( + callParameters, task, sessionDetails.descriptionComputer, notUsedDescription, + permanentSessionEnvironment, sessionLevels.first().additional, mlModel, apiPlatform, sessionLevels, mlModel.knownFeatures, + ) + collector.handleCollectedTree(analyseThenLogStructure) + MLModelPredictionBranching(mlModel, collector) + } + + sessionListener = approachListener.onStartedSession(session) + + return Session.StartOutcome.Success(session) + } + + private fun validateNotUsedDescription(notUsedDescription: PerTier) { + val maybeDuplicatedTaskTiers = sessionLevels.flatMap { it.main + it.additional } + val taskTiers = maybeDuplicatedTaskTiers.toSet() + require(notUsedDescription.keys == taskTiers) { + "Selectors for those and only those tiers must be represented in the 'notUsedDescription' that are present in the task. " + + "Missing: ${taskTiers - notUsedDescription.keys}, " + + "Redundant: ${notUsedDescription.keys - taskTiers}" + } + } +} + +/** + * An exception that indicates that for some reason, it was not possible to provide an ML model + * when calling [com.intellij.platform.ml.impl.model.MLModel.Provider.provideModel]. + * The session's start is considered as failed. + */ +@ApiStatus.Internal +class ModelNotAcquiredOutcome

: Session.StartOutcome.Failure

() { + override val failureDetails: String = "ML Model was not provided" +} diff --git a/platform/ml-impl/src/com/intellij/platform/ml/impl/MLTask.kt b/platform/ml-impl/src/com/intellij/platform/ml/impl/MLTask.kt index b999f047a316..55f17a58e58f 100644 --- a/platform/ml-impl/src/com/intellij/platform/ml/impl/MLTask.kt +++ b/platform/ml-impl/src/com/intellij/platform/ml/impl/MLTask.kt @@ -2,8 +2,8 @@ package com.intellij.platform.ml.impl import com.intellij.openapi.extensions.ExtensionPointName -import com.intellij.openapi.progress.runBlockingCancellable import com.intellij.platform.ml.* +import com.intellij.platform.ml.ScopeEnvironment.Companion.restrictedBy import com.intellij.platform.ml.impl.apiPlatform.MLApiPlatform import com.intellij.platform.ml.impl.apiPlatform.ReplaceableIJPlatform import com.intellij.platform.ml.impl.session.AdditionalTierScheme @@ -38,6 +38,11 @@ abstract class MLTask protected constructor( val predictionClass: Class ) { init { + require(levels.isNotEmpty()) { + """ + Task $this should contain at least one level + """.trimIndent() + } require(callParameters.size == levels.size) { """ Task $this has ${levels.size} levels, but in call parameters there are ${callParameters.size} levels. @@ -51,10 +56,10 @@ abstract class MLTask protected constructor( * A method of approaching an ML task. * Usually, it is inferencing an ML model and collecting logs. * - * Each [MLTaskApproach] is initialized once by the corresponding [MLTaskApproachInitializer], + * Each [MLTaskApproach] is initialized once by the corresponding [MLTaskApproachBuilder], * then the [apiPlatform] is fixed. * - * @see [com.intellij.platform.ml.impl.approach.LogDrivenModelInference] for currently used approach. + * @see [com.intellij.platform.ml.impl.LogDrivenModelInference] for currently used approach. */ @ApiStatus.Internal interface MLTaskApproach

{ @@ -65,15 +70,10 @@ interface MLTaskApproach

{ val task: MLTask

/** - * The platform, this approach is called within, that was provided by [MLTaskApproachInitializer] + * The platform, this approach is called within, that was provided by [MLTaskApproachBuilder] */ val apiPlatform: MLApiPlatform - /** - * A static declaration of the features, used in the approach. - */ - val approachDeclaration: Declaration - /** * Acquire the ML model and start the session. * @@ -87,31 +87,54 @@ interface MLTaskApproach

{ */ suspend fun startSession(callParameters: Environment, permanentSessionEnvironment: Environment): Session.StartOutcome

- data class Declaration( + data class SessionDeclaration( val sessionFeatures: Map>>, val levelsScheme: List ) companion object { - fun

findMlApproach(task: MLTask

, apiPlatform: MLApiPlatform = ReplaceableIJPlatform): MLTaskApproach

{ - return apiPlatform.accessApproachFor(task) + fun

findMlTaskApproach(task: MLTask

, apiPlatform: MLApiPlatform): MLTaskApproachBuilder

{ + val taskApproachBuilder = requireNotNull(apiPlatform.taskApproaches.find { it.task == task }) { + """ + No approach for task $task was registered in $apiPlatform. + Available approaches: ${apiPlatform.taskApproaches} + """.trimIndent() + } + @Suppress("UNCHECKED_CAST") + return taskApproachBuilder as MLTaskApproachBuilder

} suspend fun

startMLSession(task: MLTask

, apiPlatform: MLApiPlatform, callParameters: Environment, permanentSessionEnvironment: Environment): Session.StartOutcome

{ - val approach = findMlApproach(task, apiPlatform) - return approach.startSession(callParameters, permanentSessionEnvironment) + validateEnvironment(callParameters, task.callParameters.first(), permanentSessionEnvironment, task.levels.first()) + val approach = findMlTaskApproach(task, apiPlatform).buildApproach(apiPlatform) + val safelyAccessibleCallParameters = callParameters.restrictedBy(task.callParameters.first()) + return approach.startSession(safelyAccessibleCallParameters, permanentSessionEnvironment) } suspend fun

MLTask

.startMLSession(callParameters: Environment, permanentSessionEnvironment: Environment, apiPlatform: MLApiPlatform = ReplaceableIJPlatform): Session.StartOutcome

{ return startMLSession(this@startMLSession, apiPlatform, callParameters, permanentSessionEnvironment) } - fun

MLTask

.startCoroutineAndMLSession(callParameters: Environment, permanentSessionEnvironment: Environment, apiPlatform: MLApiPlatform = ReplaceableIJPlatform): Session.StartOutcome

{ - return runBlockingCancellable { - startMLSession(callParameters, permanentSessionEnvironment, apiPlatform) + private fun validateEnvironment(callParameters: Environment, + expectedCallParameters: Set>, + permanentSessionEnvironment: Environment, + expectedPermanentTiers: Set>) { + require(callParameters.tiers == expectedCallParameters) { + """ + Invalid call parameters passed. + Missing: ${expectedCallParameters - callParameters.tiers} + Redundant: ${callParameters.tiers - expectedCallParameters} + """.trimIndent() + } + require(permanentSessionEnvironment.tiers == expectedPermanentTiers) { + """ + Invalid main environment passed. + Missing: ${expectedCallParameters - callParameters.tiers} + Redundant: ${callParameters.tiers - expectedCallParameters} + """.trimIndent() } } } @@ -121,7 +144,7 @@ interface MLTaskApproach

{ * Initializes an [MLTaskApproach] */ @ApiStatus.Internal -interface MLTaskApproachInitializer

{ +interface MLTaskApproachBuilder

{ /** * The task, that the created [MLTaskApproach] is dedicated to solve. */ @@ -129,22 +152,25 @@ interface MLTaskApproachInitializer

{ /** * Initializes the approach. - * It is called only once during the application's runtime. - * So it is crucial that this function will accept the [MLApiPlatform] you want it to. + * It is called each time, when another ml session is started. * - * To access the API to build event validator statically, - * FUS uses the actual [com.intellij.platform.ml.impl.apiPlatform.IJPlatform], which could be problematic if you - * want to test FUS logs. - * So make sure that you will replace it with your test platform in - * time via [com.intellij.platform.ml.impl.apiPlatform.ReplaceableIJPlatform.replacingWith]. + * @param apiPlatform The platform, that is used to initialize approach's components. + * All [TierDescriptor]s, [EnvironmentExtender]s should be already final when this method is called. */ - fun initializeApproachWithin(apiPlatform: MLApiPlatform): MLTaskApproach

+ fun buildApproach(apiPlatform: MLApiPlatform): MLTaskApproach

+ + /** + * Builds a scheme for the finished event. + * Such events are registered with [com.intellij.platform.ml.impl.logs.events.MLSessionFinishedLogger] + */ + fun buildApproachSessionDeclaration(apiPlatform: MLApiPlatform): MLTaskApproach.SessionDeclaration companion object { - val EP_NAME = ExtensionPointName>("com.intellij.platform.ml.impl.approach") + val EP_NAME = ExtensionPointName>("com.intellij.platform.ml.impl.approach") } } +@ApiStatus.Internal data class LevelSignature( val main: M, val additional: A diff --git a/platform/ml-impl/src/com/intellij/platform/ml/impl/apiPlatform/IJPlatform.kt b/platform/ml-impl/src/com/intellij/platform/ml/impl/apiPlatform/IJPlatform.kt index 7a705955f797..72585fb5a0d8 100644 --- a/platform/ml-impl/src/com/intellij/platform/ml/impl/apiPlatform/IJPlatform.kt +++ b/platform/ml-impl/src/com/intellij/platform/ml/impl/apiPlatform/IJPlatform.kt @@ -2,16 +2,13 @@ package com.intellij.platform.ml.impl.apiPlatform import com.intellij.openapi.diagnostic.thisLogger -import com.intellij.openapi.util.registry.Registry import com.intellij.platform.ml.EnvironmentExtender import com.intellij.platform.ml.Feature import com.intellij.platform.ml.ObsoleteTierDescriptor import com.intellij.platform.ml.TierDescriptor -import com.intellij.platform.ml.impl.MLTaskApproachInitializer +import com.intellij.platform.ml.impl.MLTaskApproachBuilder import com.intellij.platform.ml.impl.apiPlatform.MLApiPlatform.ExtensionController import com.intellij.platform.ml.impl.apiPlatform.ReplaceableIJPlatform.replacingWith -import com.intellij.platform.ml.impl.logger.MLEvent -import com.intellij.platform.ml.impl.monitoring.MLApiStartupListener import com.intellij.platform.ml.impl.monitoring.MLTaskGroupListener import com.intellij.util.application import com.intellij.util.messages.Topic @@ -35,25 +32,9 @@ fun interface MessagingProvider { } } -fun interface MLTaskListenerProvider : MessagingProvider { +fun interface MLTaskGroupListenerProvider : MessagingProvider { companion object { - val TOPIC = MessagingProvider.createTopic("ml.task") - } -} - -fun interface MLEventProvider : MessagingProvider { - companion object { - val TOPIC = MessagingProvider.createTopic("ml.event") - } -} - -fun interface MLApiStartupListenerProvider : MessagingProvider { - companion object { - val TOPIC = MessagingProvider.createTopic("ml.startup") - } - - interface ThisProvider : MLApiStartupListenerProvider, MLApiStartupListener { - override fun provide(collector: (MLApiStartupListener) -> Unit) = collector(this) + val TOPIC = MessagingProvider.createTopic("ml.task") } } @@ -73,41 +54,22 @@ private data object IJPlatform : MLApiPlatform() { override val environmentExtenders: List> get() = EnvironmentExtender.EP_NAME.extensionList - override val taskApproaches: List> - get() = MLTaskApproachInitializer.EP_NAME.extensionList + override val taskApproaches: List> + get() = MLTaskApproachBuilder.EP_NAME.extensionList override val taskListeners: List - get() = MessagingProvider.collect(MLTaskListenerProvider.TOPIC) - - override val events: List - get() = MessagingProvider.collect(MLEventProvider.TOPIC) - - override val startupListeners: List - get() = MessagingProvider.collect(MLApiStartupListenerProvider.TOPIC) - - override fun addStartupListener(listener: MLApiStartupListener): ExtensionController { - val connection = application.messageBus.connect() - connection.subscribe(MLApiStartupListenerProvider.TOPIC, MLApiStartupListenerProvider { collector -> collector(listener) }) - return ExtensionController { connection.disconnect() } - } + get() = MessagingProvider.collect(MLTaskGroupListenerProvider.TOPIC) override fun addTaskListener(taskListener: MLTaskGroupListener): ExtensionController { val connection = application.messageBus.connect() - connection.subscribe(MLTaskListenerProvider.TOPIC, MLTaskListenerProvider { it(taskListener) }) - return ExtensionController { connection.disconnect() } - } - - override fun addEvent(event: MLEvent): ExtensionController { - val connection = application.messageBus.connect() - connection.subscribe(MLEventProvider.TOPIC, MLEventProvider { it(event) }) + connection.subscribe(MLTaskGroupListenerProvider.TOPIC, MLTaskGroupListenerProvider { it(taskListener) }) return ExtensionController { connection.disconnect() } } override fun manageNonDeclaredFeatures(descriptor: ObsoleteTierDescriptor, nonDeclaredFeatures: Set) { - if (!Registry.`is`("ml.description.logMissing")) return val printer = CodeLikePrinter() val codeLikeMissingDeclaration = printer.printCodeLikeString(nonDeclaredFeatures.map { it.declaration }) - thisLogger().info("${descriptor::class.java} is missing declaration: setOf($codeLikeMissingDeclaration)") + thisLogger().debug("${descriptor::class.java} is missing declaration: setOf($codeLikeMissingDeclaration)") } } @@ -115,11 +77,11 @@ private data object IJPlatform : MLApiPlatform() { * Also a "real-life" [MLApiPlatform], but it can be replaced with another one any time. * * We always want to test [com.intellij.platform.ml.impl.MLTaskApproach]es. - * But after they are initialized by [com.intellij.platform.ml.impl.MLTaskApproachInitializer], + * But after they are initialized by [com.intellij.platform.ml.impl.MLTaskApproachBuilder], * the passed [MLApiPlatform] could already spread all the way within the API. * But the user-defined instances of the api could be overridden for testing sake. * - * To replace all [TierDescriptor], [EnvironmentExtender] and [MLTaskApproachInitializer] + * To replace all [TierDescriptor], [EnvironmentExtender] and [MLTaskApproachBuilder] * to test your code, you may call [replacingWith] and pass the desired environment, * that contains all the objects you need for your test. */ @@ -137,32 +99,17 @@ object ReplaceableIJPlatform : MLApiPlatform() { override val environmentExtenders: List> get() = platform.environmentExtenders - override val taskApproaches: List> + override val taskApproaches: List> get() = platform.taskApproaches override val taskListeners: List get() = platform.taskListeners - override val events: List - get() = platform.events - - override val startupListeners: List - get() = platform.startupListeners - - - override fun addStartupListener(listener: MLApiStartupListener): ExtensionController { - return extend(listener) { platform -> platform.addStartupListener(listener) } - } - override fun addTaskListener(taskListener: MLTaskGroupListener): ExtensionController { return extend(taskListener) { platform -> platform.addTaskListener(taskListener) } } - override fun addEvent(event: MLEvent): ExtensionController { - return extend(event) { platform -> platform.addEvent(event) } - } - override fun manageNonDeclaredFeatures(descriptor: ObsoleteTierDescriptor, nonDeclaredFeatures: Set) = platform.manageNonDeclaredFeatures(descriptor, nonDeclaredFeatures) diff --git a/platform/ml-impl/src/com/intellij/platform/ml/impl/apiPlatform/MLApiPlatform.kt b/platform/ml-impl/src/com/intellij/platform/ml/impl/apiPlatform/MLApiPlatform.kt index a01264fb9a47..76f9e49f1ce7 100644 --- a/platform/ml-impl/src/com/intellij/platform/ml/impl/apiPlatform/MLApiPlatform.kt +++ b/platform/ml-impl/src/com/intellij/platform/ml/impl/apiPlatform/MLApiPlatform.kt @@ -1,15 +1,11 @@ // Copyright 2000-2023 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.apiPlatform -import com.intellij.openapi.progress.ProcessCanceledException import com.intellij.platform.ml.* -import com.intellij.platform.ml.impl.MLTask -import com.intellij.platform.ml.impl.MLTaskApproach -import com.intellij.platform.ml.impl.MLTaskApproachInitializer -import com.intellij.platform.ml.impl.logger.MLEvent -import com.intellij.platform.ml.impl.logger.MLEventsLogger -import com.intellij.platform.ml.impl.monitoring.* +import com.intellij.platform.ml.impl.MLTaskApproachBuilder +import com.intellij.platform.ml.impl.monitoring.MLApproachListener import com.intellij.platform.ml.impl.monitoring.MLApproachListener.Companion.asJoinedListener +import com.intellij.platform.ml.impl.monitoring.MLTaskGroupListener import com.intellij.platform.ml.impl.monitoring.MLTaskGroupListener.Companion.onAttemptedToStartSession import com.intellij.platform.ml.impl.monitoring.MLTaskGroupListener.Companion.targetedApproaches import org.jetbrains.annotations.ApiStatus @@ -18,7 +14,7 @@ import org.jetbrains.annotations.ApiStatus * Represents an environment, that provides extendable parts of the ML API. * * Each entity inside the API could access the platform, it is running within, - * as everything happens after [com.intellij.platform.ml.impl.MLTaskApproachInitializer.initializeApproachWithin], + * as everything happens after [com.intellij.platform.ml.impl.MLTaskApproachBuilder.buildApproach], * where the platform is acknowledged. * * All usages of the ij platform functionality (extension points, registry keys, etc.) shall be @@ -26,47 +22,23 @@ import org.jetbrains.annotations.ApiStatus */ @ApiStatus.Internal abstract class MLApiPlatform { - private val finishedInitialization = lazy { MLApiPlatformInitializationProcess() } - private var initializationStage: InitializationStage = InitializationStage.NotStarted - - /** - * Each [MLTaskApproach] is initialized only once during the application's lifetime. - * This function keeps track of all approaches that were initialized already, and initializes - * them when they are first needed. - */ - fun

accessApproachFor(task: MLTask

): MLTaskApproach

{ - return finishedInitialization.value.getApproachFor(task) - } - - - /** - * The extendable static state of the platform that must be fixed to create FUS group's validator. - * - * These values must be static, as they define the FUS scheme - */ - val staticState: StaticState - get() = StaticState(tierDescriptors, environmentExtenders, taskApproaches) - /** * The descriptors that are available in the platform. - * This value is interchangeable during the application runtime, - * see [staticState]. + * This value is interchangeable during the application runtime. */ abstract val tierDescriptors: List /** * The complete list of environment extenders, available in the platform. - * This value is interchangeable during the application runtime, - * see [staticState]. + * This value is interchangeable during the application runtime. */ abstract val environmentExtenders: List> /** * The complete list of the approaches for ML tasks, available in the platform. - * This value is interchangeable during the application runtime, - * see [staticState]. + * This value is interchangeable during the application runtime. */ - abstract val taskApproaches: List> + abstract val taskApproaches: List> /** @@ -84,194 +56,29 @@ abstract class MLApiPlatform { */ abstract fun addTaskListener(taskListener: MLTaskGroupListener): ExtensionController - /** - * ML events that will be written to FUS logs. - * As FUS is initialized only once, on the application's startup, they all must be registered - * before that via [addMLEventBeforeFusInitialized]. - * - * This value could be mutable, however, only during a short period of time: after the application's startup, - * and before FUS logs initialization. - */ - abstract val events: List - - /** - * Adds another ML event dynamically. - * The event could be removed via the corresponding [ExtensionController.removeExtension] call. - */ - fun addMLEventBeforeFusInitialized(event: MLEvent): ExtensionController { - when (val stage = initializationStage) { - is InitializationStage.Failed -> throw Exception("Initialization of ML Api Platform has failed, events could not be added.", stage.asException) - is InitializationStage.Cancelled -> throw Exception("Initialization of ML Api Platform has been cancelled") - is InitializationStage.PotentiallySuccessful -> { - require(stage.order <= InitializationStage.InitializingApproaches.order) { - "FUS group initialization has already been started, not allowed to register more ML Events" - } - return addEvent(event) - } - } - } - - /** - * The complete list of the listeners that are listening to the process of an MLApiPlatform's initialization. - * The initialization is performed in [finishedInitialization]'s init block. - * - * This value could be mutable, so - * additional listeners could be added via [addStartupListener]. - * - * If a listener was added after a certain initialization stage, - * only callbacks of those stages will be triggered later that have not happened yet. - */ - abstract val startupListeners: List - - /** - * Adds another startup listener. - * The listener could be removed via the corresponding [ExtensionController.removeExtension] call. - */ - abstract fun addStartupListener(listener: MLApiStartupListener): ExtensionController - /** * Declares how the computed but non-declared features will be handled. + * + * TODO: Remove, it is useless, or add a generic "logDebug" method instead of this. */ abstract fun manageNonDeclaredFeatures(descriptor: ObsoleteTierDescriptor, nonDeclaredFeatures: Set) - - internal abstract fun addEvent(event: MLEvent): ExtensionController - fun interface ExtensionController { fun removeExtension() } - data class StaticState( - val tierDescriptors: List, - val environmentExtenders: List>, - val taskApproaches: List>, - ) - - private sealed class InitializationStage( - val callListener: (MLApiStartupProcessListener) -> Unit, - val allowStartup: Boolean - ) { - sealed class PotentiallySuccessful(val order: Int, callListener: (MLApiStartupProcessListener) -> Unit, allowStartup: Boolean) : InitializationStage(callListener, allowStartup) - - data object Cancelled : InitializationStage({ it.onCanceled() }, allowStartup = true) - - class Failed(lastStage: InitializationStage, nextStage: InitializationStage, exception: Throwable, callListener: (MLApiStartupProcessListener) -> Unit) : InitializationStage(callListener, allowStartup = false) { - val asException = Exception("Failed to proceed from the initialization stage $lastStage to $nextStage", exception) - } - - data object NotStarted : PotentiallySuccessful(0, {}, allowStartup = true) - data object InitializingApproaches : PotentiallySuccessful(1, { it.onStartedInitializingApproaches() }, allowStartup = true) - data class InitializingFUS(val initializedApproaches: Collection>) : PotentiallySuccessful(2, { - it.onStartedInitializingFus(initializedApproaches) - }, allowStartup = false) - - data object Finished : PotentiallySuccessful(3, { it.onFinished() }, allowStartup = false) - } - - private inner class MLApiPlatformInitializationProcess { - val approachPerTask: Map, MLTaskApproach<*>> - private val completeInitializersList: List> = taskApproaches.toMutableList() - - init { - require(initializationStage.allowStartup) { - """ - ML API Platform's initialization should not be run twice. Current state is $initializationStage - """.trimIndent() - } - - fun currentStartupListeners(): List = startupListeners.map { it.onBeforeStarted(this@MLApiPlatform) } - - fun proceedToNextStage(nextStage: InitializationStage.PotentiallySuccessful, action: () -> T): T { - return try { - action().also { - currentStartupListeners().forEach { nextStage.callListener(it) } - initializationStage = nextStage - } - } - catch (ex: ProcessCanceledException) { - initializationStage = InitializationStage.Cancelled - throw ex - } - catch (ex: Throwable) { - val failure = InitializationStage.Failed(initializationStage, nextStage, ex) { it.onFailed(ex) } - initializationStage = failure - throw failure.asException - } - } - - proceedToNextStage(InitializationStage.InitializingApproaches) {} - - val initializedApproachPerTask = mutableListOf>() - - approachPerTask = proceedToNextStage(InitializationStage.InitializingFUS(initializedApproachPerTask)) { - completeInitializersList.validate() - - fun initializeApproach(approachInitializer: MLTaskApproachInitializer) { - initializedApproachPerTask.add(InitializerAndApproach( - approachInitializer, - approachInitializer.initializeApproachWithin(this@MLApiPlatform) - )) - } - completeInitializersList.forEach { initializeApproach(it) } - initializedApproachPerTask.associate { it.initializer.task to it.approach } - } - - proceedToNextStage(InitializationStage.Finished) { - MLEventsLogger.Manager.ensureInitialized(okIfInitializing = true, this@MLApiPlatform) - } - } - - fun

getApproachFor(task: MLTask

): MLTaskApproach

{ - val taskApproach = requireNotNull(approachPerTask[task]) { - val mainMessage = "No approach for task $task was found" - val lateRegistrationMessage = getLateApproachRegistrationAssumption(task) - if (lateRegistrationMessage != null) "$mainMessage. $lateRegistrationMessage" else mainMessage - } - - @Suppress("UNCHECKED_CAST") - return taskApproach as MLTaskApproach

- } - - private fun getLateApproachRegistrationAssumption(task: MLTask<*>): String? { - val currentInitializersList = taskApproaches.toMutableList() - if (completeInitializersList == currentInitializersList) return null - val taskApproaches = currentInitializersList.filter { it.task == task } - if (taskApproaches.isEmpty()) return null - require(taskApproaches.size == 1) { "More than one approach for task $task: $taskApproaches" } - return "Approach ${taskApproaches.first()} for task ${task.name} was registered after the ML API Platform was initialized" - } - - private fun List>.validate() { - val duplicateInitializerPerTask = this.groupBy { it.task }.filter { it.value.size > 1 } - require(duplicateInitializerPerTask.isEmpty()) { - "Found more than one approach for the following tasks: ${duplicateInitializerPerTask}" - } - } - } - companion object { fun MLApiPlatform.getDescriptorsOfTiers(tiers: Set>): PerTier> { val descriptorsPerTier = tierDescriptors.groupBy { it.tier } return tiers.associateWith { descriptorsPerTier[it] ?: emptyList() } } - fun MLApiPlatform.ensureApproachesInitialized() { - when (val stage = initializationStage) { - is InitializationStage.Failed -> throw Exception("Unable to ensure that approaches are initialized", stage.asException) - InitializationStage.NotStarted -> finishedInitialization.value - InitializationStage.Cancelled -> finishedInitialization.value - InitializationStage.InitializingApproaches -> throw Exception("Recursion detected while initializing approaches") - is InitializationStage.InitializingFUS -> return - InitializationStage.Finished -> return - } - } - - fun MLApiPlatform.getJoinedListenerForTask(taskApproach: MLTaskApproach

, + fun MLApiPlatform.getJoinedListenerForTask(taskApproachBuilder: MLTaskApproachBuilder

, permanentSessionEnvironment: Environment): MLApproachListener { - val relevantGroupListeners = taskListeners.filter { taskApproach.javaClass in it.targetedApproaches } + val relevantGroupListeners = taskListeners.filter { taskApproachBuilder.javaClass in it.targetedApproaches } val approachListeners = relevantGroupListeners.mapNotNull { - it.onAttemptedToStartSession(taskApproach, permanentSessionEnvironment) + it.onAttemptedToStartSession(taskApproachBuilder, permanentSessionEnvironment) } return approachListeners.asJoinedListener() } diff --git a/platform/ml-impl/src/com/intellij/platform/ml/impl/approach/AnalysisMethod.kt b/platform/ml-impl/src/com/intellij/platform/ml/impl/approach/AnalysisMethod.kt deleted file mode 100644 index e156dc93a663..000000000000 --- a/platform/ml-impl/src/com/intellij/platform/ml/impl/approach/AnalysisMethod.kt +++ /dev/null @@ -1,31 +0,0 @@ -// Copyright 2000-2023 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.approach - -import com.intellij.platform.ml.FeatureDeclaration -import com.intellij.platform.ml.PerTier -import com.intellij.platform.ml.impl.model.MLModel -import com.intellij.platform.ml.impl.session.AnalysedRootContainer -import com.intellij.platform.ml.impl.session.DescribedRootContainer -import org.jetbrains.annotations.ApiStatus -import java.util.concurrent.CompletableFuture - -/** - * Represents the method, that is utilized for the [LogDrivenModelInference]'s analysis. - */ -@ApiStatus.Internal -interface AnalysisMethod, P : Any> { - /** - * Static declaration of the features, that are used in the session tree's analysis. - */ - val structureAnalysisDeclaration: PerTier>> - - /** - * Static declaration of the session's entities, that are not tiers. - */ - val sessionAnalysisDeclaration: Map>> - - /** - * Perform the completed session's analysis. - */ - fun analyseTree(treeRoot: DescribedRootContainer): CompletableFuture> -} diff --git a/platform/ml-impl/src/com/intellij/platform/ml/impl/approach/LogDrivenModelInference.kt b/platform/ml-impl/src/com/intellij/platform/ml/impl/approach/LogDrivenModelInference.kt deleted file mode 100644 index f9e1bc66a7f4..000000000000 --- a/platform/ml-impl/src/com/intellij/platform/ml/impl/approach/LogDrivenModelInference.kt +++ /dev/null @@ -1,223 +0,0 @@ -// Copyright 2000-2023 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.approach - -import com.intellij.openapi.diagnostic.thisLogger -import com.intellij.platform.ml.* -import com.intellij.platform.ml.ScopeEnvironment.Companion.narrowedTo -import com.intellij.platform.ml.impl.* -import com.intellij.platform.ml.impl.MLTaskApproach.Companion.startMLSession -import com.intellij.platform.ml.impl.apiPlatform.MLApiPlatform -import com.intellij.platform.ml.impl.apiPlatform.MLApiPlatform.Companion.getDescriptorsOfTiers -import com.intellij.platform.ml.impl.apiPlatform.MLApiPlatform.Companion.getJoinedListenerForTask -import com.intellij.platform.ml.impl.model.MLModel -import com.intellij.platform.ml.impl.monitoring.MLApproachListener -import com.intellij.platform.ml.impl.monitoring.MLSessionListener -import com.intellij.platform.ml.impl.session.* -import org.jetbrains.annotations.ApiStatus - -/** - * The main way to apply classical machine learning approaches: run the ML model, collect the logs, retrain the model, repeat. - * - * @param task The task that is solved by this approach. - * @param apiPlatform The platform, that the approach will be running within. - */ -@ApiStatus.Internal -abstract class LogDrivenModelInference, P : Any>( - override val task: MLTask

, - override val apiPlatform: MLApiPlatform -) : MLTaskApproach

{ - /** - * The method that is used to analyze sessions. - * - * [StructureAndModelAnalysis] is currently used analysis method - * that is dedicated to analyze session's tree-like structure, - * and the ML model. - */ - abstract val analysisMethod: AnalysisMethod - - /** - * Provides an ML model to use during session's lifetime. - */ - abstract val mlModelProvider: MLModel.Provider - - /** - * Declares features, that are not used by the ML model, but must be computed anyway, - * so they make it to logs. - * - * A feature cannot be simultaneously declared as "not used description" and as used by the [mlModelProvider]'s - * provided model. - * If a feature is not declared as "not used but still computed" or as "used by the model", then it will be computed. - * - * It must contain explicitly declared selectors for each tier used in [task], as well as in [additionallyDescribedTiers]. - * - * @param callParameters Contains any additional parameters that you passed in [com.intellij.platform.ml.impl.MLTaskApproach.startMLSession] - * The tier set is equal to the one that was declared in [com.intellij.platform.ml.impl.MLTask.callParameters]'s first level. - * @param mlModel The model that was acquired by the [mlModelProvider] - */ - abstract fun getNotUsedDescription(callParameters: Environment, mlModel: M): PerTier - - /** - * Performs description's computation. - * Could perform caching mechanisms to avoid recomputing features every time. - */ - abstract val descriptionComputer: DescriptionComputer - - /** - * Tiers that do not make a part of te [task], but they could be described and passed to the ML model. - * - * The size of this list must correspond to the number of levels in the solved [task]. - */ - abstract val additionallyDescribedTiers: List>> - - private val levels: List by lazy { - (task.levels zip additionallyDescribedTiers).map { LevelSignature(it.first, it.second) } - } - - private val approachValidation: Unit by lazy { validateApproach() } - - override suspend fun startSession(callParameters: Environment, permanentSessionEnvironment: Environment): Session.StartOutcome

{ - return startSessionMonitoring(callParameters, permanentSessionEnvironment) - } - - private suspend fun startSessionMonitoring(callParameters: Environment, permanentSessionEnvironment: Environment): Session.StartOutcome

{ - val approachListener = apiPlatform.getJoinedListenerForTask(this, permanentSessionEnvironment) - try { - return acquireModelAndStartSession(callParameters, permanentSessionEnvironment, approachListener) - } - catch (e: Throwable) { - approachListener.onFailedToStartSessionWithException(e) - return Session.StartOutcome.UncaughtException(e) - } - } - - private suspend fun acquireModelAndStartSession(unsafeCallParameters: Environment, - permanentSessionEnvironment: Environment, - approachListener: MLApproachListener): Session.StartOutcome

{ - approachValidation - - val callParameters = unsafeCallParameters.narrowedTo(task.callParameters.first()) - val extendedPermanentSessionEnvironment = Environment.joined(callParameters, permanentSessionEnvironment) - - val mlModel: M = run { - val nullableMlModel = mlModelProvider.provideModel(callParameters, extendedPermanentSessionEnvironment, levels) - if (nullableMlModel == null) { - val failure = ModelNotAcquiredOutcome

() - approachListener.onFailedToStartSession(failure) - return failure - } - nullableMlModel - } - - var sessionListener: MLSessionListener? = null - - val analyseThenLogStructure = SessionTreeHandler, M, P> { treeRoot -> - sessionListener?.onSessionDescriptionFinished(treeRoot) - analysisMethod.analyseTree(treeRoot).thenApplyAsync { analysedSession -> - sessionListener?.onSessionAnalysisFinished(analysedSession) - }.exceptionally { - thisLogger().error(it) - } - } - - val notUsedDescription = getNotUsedDescription(callParameters, mlModel) - validateNotUsedDescription(notUsedDescription) - - val session = if (levels.size == 1) { - val collector = SolitaryLeafCollector.build( - callParameters, task.callParameters, descriptionComputer, notUsedDescription, - permanentSessionEnvironment, levels.first().additional, mlModel, apiPlatform, levels.first(), mlModel.knownFeatures - ) - collector.handleCollectedTree(analyseThenLogStructure) - MLModelPrediction(mlModel, collector) - } - else { - val collector = RootCollector.build( - callParameters, task.callParameters, descriptionComputer, notUsedDescription, - permanentSessionEnvironment, levels.first().additional, mlModel, apiPlatform, levels, mlModel.knownFeatures - ) - collector.handleCollectedTree(analyseThenLogStructure) - MLModelPredictionBranching(mlModel, collector) - } - - sessionListener = approachListener.onStartedSession(session) - - return Session.StartOutcome.Success(session) - } - - override val approachDeclaration: MLTaskApproach.Declaration - get() { - approachValidation - - return MLTaskApproach.Declaration( - sessionFeatures = analysisMethod.sessionAnalysisDeclaration, - levelsScheme = levels.map { levelTiers -> - LevelSignature( - buildMainTiersScheme(levelTiers.main, apiPlatform), - buildAdditionalTiersScheme(levelTiers.additional, apiPlatform), - ) - } - ) - } - - private fun validateApproach() { - require(task.levels.size == additionallyDescribedTiers.size) { - "Task $task has ${task.levels.size} levels, when 'additionallyDescribedTiers' has ${additionallyDescribedTiers.size}" - } - - require(levels.isNotEmpty()) { "Task must declare at least one level" } - - val maybeDuplicatedTaskTiers = levels.flatMap { it.main + it.additional } - val taskTiers = maybeDuplicatedTaskTiers.toSet() - - require(maybeDuplicatedTaskTiers.size == taskTiers.size) { - "There are duplicated tiers in the declaration: ${maybeDuplicatedTaskTiers - taskTiers}" - } - } - - private fun validateNotUsedDescription(notUsedDescription: PerTier) { - val maybeDuplicatedTaskTiers = levels.flatMap { it.main + it.additional } - val taskTiers = maybeDuplicatedTaskTiers.toSet() - require(notUsedDescription.keys == taskTiers) { - "Selectors for those and only those tiers must be represented in the 'notUsedDescription' that are present in the task. " + - "Missing: ${taskTiers - notUsedDescription.keys}, " + - "Redundant: ${notUsedDescription.keys - taskTiers}" - } - } - - private fun buildTierDescriptionDeclaration(tierDescriptors: Collection): Set> { - return tierDescriptors.flatMap { - if (it is ObsoleteTierDescriptor) it.partialDescriptionDeclaration else it.descriptionDeclaration - }.toSet() - } - - private fun buildMainTiersScheme(tiers: Set>, apiEnvironment: MLApiPlatform): PerTier { - val tiersDescriptors = apiEnvironment.getDescriptorsOfTiers(tiers) - - return tiers.associateWith { tier -> - val tierDescriptors = tiersDescriptors.getValue(tier) - val descriptionDeclaration = buildTierDescriptionDeclaration(tierDescriptors) - val analysisDeclaration = analysisMethod.structureAnalysisDeclaration[tier] ?: emptySet() - MainTierScheme(descriptionDeclaration, analysisDeclaration) - } - } - - private fun buildAdditionalTiersScheme(tiers: Set>, apiEnvironment: MLApiPlatform): PerTier { - val tiersDescriptors = apiEnvironment.getDescriptorsOfTiers(tiers) - - return tiers.associateWith { tier -> - val tierDescriptors = tiersDescriptors.getValue(tier) - val descriptionDeclaration = buildTierDescriptionDeclaration(tierDescriptors) - AdditionalTierScheme(descriptionDeclaration) - } - } -} - -/** - * An exception that indicates that for some reason, it was not possible to provide an ML model - * when calling [com.intellij.platform.ml.impl.model.MLModel.Provider.provideModel]. - * The session's start is considered as failed. - */ -@ApiStatus.Internal -class ModelNotAcquiredOutcome

: Session.StartOutcome.Failure

() { - override val failureDetails: String = "ML Model was not provided" -} diff --git a/platform/ml-impl/src/com/intellij/platform/ml/impl/environment/ExtendedEnvironment.kt b/platform/ml-impl/src/com/intellij/platform/ml/impl/environment/ExtendedEnvironment.kt index 56956e8f7cb2..72ed8a047aa2 100644 --- a/platform/ml-impl/src/com/intellij/platform/ml/impl/environment/ExtendedEnvironment.kt +++ b/platform/ml-impl/src/com/intellij/platform/ml/impl/environment/ExtendedEnvironment.kt @@ -1,6 +1,8 @@ // Copyright 2000-2023 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.environment +import com.intellij.openapi.diagnostic.debug +import com.intellij.openapi.diagnostic.logger import com.intellij.platform.ml.* import com.intellij.platform.ml.EnvironmentExtender.Companion.extendTierInstance import com.intellij.platform.ml.ScopeEnvironment.Companion.accessibleSafelyByOrNull @@ -45,7 +47,7 @@ class ExtendedEnvironment : Environment { tiersToExtend: Set>) { val alreadyExistingButRequestedTiers = tiersToExtend.filter { it in mainEnvironment } require(alreadyExistingButRequestedTiers.isEmpty()) { - """ + """ Requested to extend $alreadyExistingButRequestedTiers, but they already exist in the main environment (which contains ${mainEnvironment.tiers}) """.trimIndent() @@ -53,7 +55,8 @@ class ExtendedEnvironment : Environment { val nonOverridingExtenders = environmentExtenders.filter { it.extendingTier !in mainEnvironment } storage = buildExtendedEnvironment( tiersToExtend + mainEnvironment.tiers, - nonOverridingExtenders + mainEnvironment.separateIntoExtenders() + nonOverridingExtenders + mainEnvironment.separateIntoExtenders(), + mainEnvironment.tiers ) } @@ -69,7 +72,8 @@ class ExtendedEnvironment : Environment { val nonOverridingExtenders = environmentExtenders.filter { it.extendingTier !in mainEnvironment } storage = buildExtendedEnvironment( nonOverridingExtenders.map { it.extendingTier }.toSet() + mainEnvironment.tiers, - nonOverridingExtenders + mainEnvironment.separateIntoExtenders() + nonOverridingExtenders + mainEnvironment.separateIntoExtenders(), + mainEnvironment.tiers ) } @@ -87,20 +91,48 @@ class ExtendedEnvironment : Environment { * Creates an [Environment] that contents tiers from [tiers], that were successfully extended by [extenders] */ private fun buildExtendedEnvironment(tiers: Set>, - extenders: List>): Environment { + extenders: List>, + mainTiers: Set>): Environment { val validatedExtendersPerTier = validateExtenders(tiers, extenders) val extensionOrder = ENVIRONMENT_RESOLVER.resolve(validatedExtendersPerTier) val storage = TierInstanceStorage() - extensionOrder.map { - val safelyAccessibleEnvironment = storage.accessibleSafelyByOrNull(it) ?: return@map - it.extendTierInstance(safelyAccessibleEnvironment)?.let { extendedTierInstance -> + val extensionOutcome: List = extensionOrder.map { environmentExtender -> + val safelyAccessibleEnvironment = storage.accessibleSafelyByOrNull(environmentExtender) + ?: return@map ExtensionOutcome.InsufficientEnvironment(environmentExtender) + environmentExtender.extendTierInstance(safelyAccessibleEnvironment)?.let { extendedTierInstance -> storage.putTierInstance(extendedTierInstance) - } + ExtensionOutcome.Success(environmentExtender, extendedTierInstance) + } ?: ExtensionOutcome.NullReturned(environmentExtender) } + + logger().debug { + "Extending environment having ${mainTiers}\n" + + extensionOutcome + .filterNot { it.environmentExtender is ContainingExtender<*> } + .withIndex() + .joinToString("\n") { (index, outcome) -> " extender #$index: $outcome" } + } + return storage } + private sealed class ExtensionOutcome { + abstract val environmentExtender: EnvironmentExtender<*> + + data class Success(override val environmentExtender: EnvironmentExtender<*>, val instance: Any) : ExtensionOutcome() { + override fun toString() = "[success] ${environmentExtender.javaClass.simpleName} -> ${environmentExtender.extendingTier}" + } + + data class InsufficientEnvironment(override val environmentExtender: EnvironmentExtender<*>) : ExtensionOutcome() { + override fun toString() = "[insufficient environment] ${environmentExtender.javaClass.simpleName} " + } + + data class NullReturned(override val environmentExtender: EnvironmentExtender<*>) : ExtensionOutcome() { + override fun toString() = "[null returned] ${environmentExtender.javaClass.simpleName} " + } + } + private fun validateExtenders(tiers: Set>, extenders: List>): Map, EnvironmentExtender<*>> { val extendableTiers: Set> = extenders.map { it.extendingTier }.toSet() @@ -130,16 +162,16 @@ class ExtendedEnvironment : Environment { } } -private fun Environment.separateIntoExtenders(): List> { - class ContainingExtender(private val tier: Tier) : EnvironmentExtender { - override val extendingTier: Tier = tier +internal class ContainingExtender(private val containingEnvironment: Environment, private val tier: Tier) : EnvironmentExtender { + override val extendingTier: Tier = tier - override fun extend(environment: Environment): T { - return this@separateIntoExtenders[tier] - } - - override val requiredTiers: Set> = emptySet() + override fun extend(environment: Environment): T { + return containingEnvironment[tier] } - return this.tiers.map { tier -> ContainingExtender(tier) } + override val requiredTiers: Set> = emptySet() +} + +private fun Environment.separateIntoExtenders(): List> { + return this.tiers.map { tier -> ContainingExtender(this, tier) } } diff --git a/platform/ml-impl/src/com/intellij/platform/ml/impl/logger/MLApiPlatformLogger.kt b/platform/ml-impl/src/com/intellij/platform/ml/impl/logger/MLApiPlatformLogger.kt deleted file mode 100644 index 9e99ff03e259..000000000000 --- a/platform/ml-impl/src/com/intellij/platform/ml/impl/logger/MLApiPlatformLogger.kt +++ /dev/null @@ -1,36 +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.logger - -import com.intellij.internal.statistic.eventLog.events.BooleanEventField -import com.intellij.internal.statistic.eventLog.events.ClassEventField -import com.intellij.internal.statistic.eventLog.events.EventField -import com.intellij.internal.statistic.eventLog.events.EventPair -import com.intellij.platform.ml.impl.apiPlatform.MLApiPlatform -import com.intellij.platform.ml.impl.monitoring.MLApiStartupListener -import com.intellij.platform.ml.impl.monitoring.MLApiStartupProcessListener -import org.jetbrains.annotations.ApiStatus - -@ApiStatus.Internal -class MLApiPlatformStartupLogger : EventIdRecordingMLEvent(), MLApiStartupListener { - companion object { - val SUCCESS = BooleanEventField("success") - val EXCEPTION = ClassEventField("exception") - } - - override val eventName: String = "startup" - - override val declaration: Array> = arrayOf(SUCCESS) - - override fun onBeforeStarted(apiPlatform: MLApiPlatform): MLApiStartupProcessListener { - val eventId = getEventId(apiPlatform) - return object : MLApiStartupProcessListener { - override fun onFinished() = eventId.log(SUCCESS with true) - - override fun onFailed(exception: Throwable?) { - val fields = mutableListOf>(SUCCESS with false) - exception?.let { fields += EXCEPTION with it.javaClass } - eventId.log(*fields.toTypedArray()) - } - } - } -} diff --git a/platform/ml-impl/src/com/intellij/platform/ml/impl/logger/MLEventsLogger.kt b/platform/ml-impl/src/com/intellij/platform/ml/impl/logger/MLEventsLogger.kt deleted file mode 100644 index 0ef9c2fe1e5b..000000000000 --- a/platform/ml-impl/src/com/intellij/platform/ml/impl/logger/MLEventsLogger.kt +++ /dev/null @@ -1,129 +0,0 @@ -// Copyright 2000-2023 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.logger - -import com.intellij.internal.statistic.eventLog.EventLogGroup -import com.intellij.internal.statistic.eventLog.events.EventField -import com.intellij.internal.statistic.eventLog.events.VarargEventId -import com.intellij.internal.statistic.service.fus.collectors.CounterUsagesCollector -import com.intellij.openapi.progress.ProcessCanceledException -import com.intellij.platform.ml.impl.apiPlatform.MLApiPlatform -import com.intellij.platform.ml.impl.apiPlatform.MLApiPlatform.Companion.ensureApproachesInitialized -import com.intellij.platform.ml.impl.apiPlatform.ReplaceableIJPlatform -import org.jetbrains.annotations.ApiStatus - -@ApiStatus.Internal -interface MLEvent { - val eventName: String - - val declaration: Array> - - fun onEventGroupInitialized(eventId: VarargEventId) -} - -@ApiStatus.Internal -abstract class EventIdRecordingMLEvent : MLEvent { - private var providedEventId: VarargEventId? = null - - protected fun getEventId(apiPlatform: MLApiPlatform): VarargEventId { - MLEventsLogger.Manager.ensureInitialized(okIfInitializing = false, apiPlatform) - return requireNotNull(providedEventId) - } - - final override fun onEventGroupInitialized(eventId: VarargEventId) { - providedEventId = eventId - } -} - -/** - * It logs ML sessions with tiers' descriptions and the analysis. - * Each session is logged after it is finished, and all analyzers [com.intellij.platform.ml.impl.session.analysis.SessionAnalyser] - * have yielded the analysis. - * - * As the FUS logs' validators are initialized once during the application's start, - * before giving the [getGroup] for the FUS to use, we must first walk through all declared - * [com.intellij.platform.ml.impl.MLTaskApproach] in the platform and register them in the FUS scheme - * (in case they require FUS logging). - */ -@ApiStatus.Internal -class MLEventsLogger : CounterUsagesCollector() { - override fun getGroup(): EventLogGroup = Initializer.GROUP.also { Manager.ensureInitialized(okIfInitializing = false) } - - object Manager { - private val defaultPlatform = ReplaceableIJPlatform - - internal fun ensureInitialized(okIfInitializing: Boolean, apiPlatform: MLApiPlatform = defaultPlatform): State { - when (val state = Initializer.state) { - is State.FailedToInitialize -> throw Exception("ML Event Log already has failed to initialize", state.exception) - State.Initializing -> if (!okIfInitializing) throw IllegalStateException("Initialization recursion") - State.NonInitialized -> Initializer.initializeGroup(apiPlatform) - State.Cancelled -> Initializer.initializeGroup(apiPlatform) - is State.Initialized -> { - val currentApiPlatformState = apiPlatform.staticState - require(currentApiPlatformState == state.context.staticState) { - """ - FUS ML Logger was initialized from ${state.context.apiPlatform} with presumably immutable state ${state.context.staticState}, - but it is used from ${apiPlatform} with state ${currentApiPlatformState}, which differs from the initial state. - Hence, something that was expected to be logged will not be. - """.trimIndent() - } - } - } - return Initializer.state - } - } - - internal data class InitializationContext( - val apiPlatform: MLApiPlatform, - val staticState: MLApiPlatform.StaticState - ) - - internal sealed interface State { - data object NonInitialized : State - - data object Initializing : State - - data object Cancelled : State - - class FailedToInitialize(val exception: Throwable) : State - - class Initialized(val context: InitializationContext) : State - } - - internal object Initializer { - val GROUP = EventLogGroup("ml", 3) - - var state: State = State.NonInitialized - - fun initializeGroup(apiPlatform: MLApiPlatform) { - state = State.Initializing - try { - apiPlatform.ensureApproachesInitialized() - val apiPlatformState = apiPlatform.staticState - - apiPlatform.addDefaultEventsAndListeners() - - apiPlatform.events.forEach { event -> - val eventId = GROUP.registerVarargEvent(event.eventName, *event.declaration) - event.onEventGroupInitialized(eventId) - } - state = State.Initialized(InitializationContext(apiPlatform, apiPlatformState)) - } - catch (e: ProcessCanceledException) { - state = State.Cancelled - throw e - } - catch (e: Throwable) { - state = State.FailedToInitialize(e) - } - (state as? State.FailedToInitialize)?.let { - throw Exception("Failed to initialize FUS ML Logger", it.exception) - } - } - - private fun MLApiPlatform.addDefaultEventsAndListeners() { - val startupLogger = MLApiPlatformStartupLogger() - addEvent(startupLogger) - addStartupListener(startupLogger) - } - } -} diff --git a/platform/ml-impl/src/com/intellij/platform/ml/impl/logger/MLSessionFailedLogger.kt b/platform/ml-impl/src/com/intellij/platform/ml/impl/logger/MLSessionFailedLogger.kt deleted file mode 100644 index ad4aae17eb05..000000000000 --- a/platform/ml-impl/src/com/intellij/platform/ml/impl/logger/MLSessionFailedLogger.kt +++ /dev/null @@ -1,98 +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.logger - -import com.intellij.internal.statistic.eventLog.events.EventField -import com.intellij.internal.statistic.eventLog.events.ObjectEventData -import com.intellij.internal.statistic.eventLog.events.ObjectEventField -import com.intellij.platform.ml.Session -import com.intellij.platform.ml.impl.MLTaskApproach -import com.intellij.platform.ml.impl.apiPlatform.MLApiPlatform -import com.intellij.platform.ml.impl.monitoring.* -import com.intellij.platform.ml.impl.monitoring.MLTaskGroupListener.ApproachListeners.Companion.monitoredBy -import com.intellij.platform.ml.impl.session.analysis.ShallowSessionAnalyser -import com.intellij.platform.ml.impl.session.analysis.ShallowSessionAnalyser.Companion.declarationObjectDescription -import org.jetbrains.annotations.ApiStatus - -/** - * Logs to FUS information about an ML session, that has failed to start. - * - * A way to start logging, is to create a [FailedSessionLoggerRegister] and add it via [com.intellij.platform.ml.impl.apiPlatform.ReplaceableIJPlatform.addStartupListener] - * - * @param taskApproach The approach that is monitored - * @param exceptionalAnalysers Analyzers, that are triggered, when [MLTaskApproach.startSession] fails with an unhandled exception - * @param normalFailureAnalysers Analyzers, that aer triggered, when [MLTaskApproach.startSession] returns a [Session.StartOutcome.Failure] - */ -@ApiStatus.Internal -class MLSessionFailedLogger( - private val taskApproach: MLTaskApproach

, - private val exceptionalAnalysers: Collection>, - private val normalFailureAnalysers: Collection>>, - private val apiPlatform: MLApiPlatform, -) : EventIdRecordingMLEvent(), MLTaskGroupListener { - override val eventName: String = "${taskApproach.task.name}.failed" - - private val fields: Map = ( - exceptionalAnalysers.map { ObjectEventField(it.name, it.declarationObjectDescription) } + - normalFailureAnalysers.map { ObjectEventField(it.name, it.declarationObjectDescription) }) - .associateBy { it.name } - - override val declaration: Array> = fields.values.toTypedArray() - - override val approachListeners: Collection> - get() { - val eventId = getEventId(apiPlatform) - return listOf( - taskApproach.javaClass monitoredBy MLApproachInitializationListener { permanentSessionEnvironment -> - object : MLApproachListener { - override fun onFailedToStartSessionWithException(exception: Throwable) { - val analysis = exceptionalAnalysers.map { - val analyserField = fields.getValue(it.name) - analyserField with ObjectEventData(it.analyse(permanentSessionEnvironment, exception)) - } - eventId.log(*analysis.toTypedArray()) - } - - override fun onFailedToStartSession(failure: Session.StartOutcome.Failure

) { - val analysis = normalFailureAnalysers.map { - val analyserField = fields.getValue(it.name) - analyserField with ObjectEventData(it.analyse(permanentSessionEnvironment, failure)) - } - eventId.log(*analysis.toTypedArray()) - } - - override fun onStartedSession(session: Session

) = null - } - } - ) - } -} - -/** - * Registers failed sessions' logging as a separate FUS event. - * - * See [MLSessionFailedLogger] - */ -@ApiStatus.Internal -open class FailedSessionLoggerRegister( - private val targetApproachClass: Class>, - private val exceptionalAnalysers: Collection>, - private val normalFailureAnalysers: Collection>>, -) : MLApiStartupListener { - override fun onBeforeStarted(apiPlatform: MLApiPlatform): MLApiStartupProcessListener { - return object : MLApiStartupProcessListener { - override fun onStartedInitializingFus(initializedApproaches: Collection>) { - @Suppress("UNCHECKED_CAST") - val targetInitializedApproach: MLTaskApproach

= requireNotNull( - initializedApproaches.find { it.approach.javaClass == targetApproachClass }) { - """ - Could not create logger for failed sessions for $targetApproachClass in platform $apiPlatform, because the corresponding approach was not initialized. - Initialized approaches: $initializedApproaches - """.trimIndent() - }.approach as MLTaskApproach

- val finishedEventLogger = MLSessionFailedLogger(targetInitializedApproach, exceptionalAnalysers, normalFailureAnalysers, apiPlatform) - apiPlatform.addMLEventBeforeFusInitialized(finishedEventLogger) - apiPlatform.addTaskListener(finishedEventLogger) - } - } - } -} diff --git a/platform/ml-impl/src/com/intellij/platform/ml/impl/logger/MLSessionFinishedLogger.kt b/platform/ml-impl/src/com/intellij/platform/ml/impl/logger/MLSessionFinishedLogger.kt deleted file mode 100644 index 248ce7e91031..000000000000 --- a/platform/ml-impl/src/com/intellij/platform/ml/impl/logger/MLSessionFinishedLogger.kt +++ /dev/null @@ -1,90 +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.logger - -import com.intellij.internal.statistic.eventLog.events.EventField -import com.intellij.platform.ml.Environment -import com.intellij.platform.ml.Session -import com.intellij.platform.ml.impl.MLTaskApproach -import com.intellij.platform.ml.impl.apiPlatform.MLApiPlatform -import com.intellij.platform.ml.impl.apiPlatform.MLApiStartupListenerProvider -import com.intellij.platform.ml.impl.monitoring.* -import com.intellij.platform.ml.impl.monitoring.MLTaskGroupListener.ApproachListeners.Companion.monitoredBy -import com.intellij.platform.ml.impl.session.AnalysedRootContainer -import com.intellij.platform.ml.impl.session.DescribedRootContainer -import org.jetbrains.annotations.ApiStatus - -/** - * Logs to FUS information about a finished ML session, tier instances' descriptions, analysis. - * - * A way to start logging, is to create a [FinishedSessionLoggerRegister] and add it with [com.intellij.platform.ml.impl.apiPlatform.ReplaceableIJPlatform.addStartupListener]. - * If you are using a regression model, see [InplaceFeaturesScheme.FusScheme.Companion.DOUBLE]. - * - * @param approach The approach whose sessions are monitored - * @param configuration The particular scheme that is used to serialize sessions - */ -@ApiStatus.Internal -class MLSessionFinishedLogger( - approach: MLTaskApproach

, - configuration: FusSessionEventBuilder.FusScheme

, - private val apiPlatform: MLApiPlatform, -) : EventIdRecordingMLEvent(), MLTaskGroupListener { - private val loggingScheme: FusSessionEventBuilder

= configuration.createEventBuilder(approach.approachDeclaration) - private val fusDeclaration: SessionFields

= loggingScheme.buildFusDeclaration() - override val declaration: Array> = fusDeclaration.getFields() - - override val eventName: String = "${approach.task.name}.finished" - - override val approachListeners: Collection> = listOf( - approach.javaClass monitoredBy InitializationLogger() - ) - - inner class InitializationLogger : MLApproachInitializationListener { - override fun onAttemptedToStartSession(permanentSessionEnvironment: Environment): MLApproachListener = ApproachLogger() - } - - inner class ApproachLogger : MLApproachListener { - override fun onFailedToStartSessionWithException(exception: Throwable) {} - - override fun onFailedToStartSession(failure: Session.StartOutcome.Failure

) {} - - override fun onStartedSession(session: Session

): MLSessionListener = SessionLogger() - } - - inner class SessionLogger : MLSessionListener { - override fun onSessionDescriptionFinished(sessionTree: DescribedRootContainer) {} - - override fun onSessionAnalysisFinished(sessionTree: AnalysedRootContainer

) { - val eventId = getEventId(apiPlatform) - val fusEventData = loggingScheme.buildRecord(sessionTree, fusDeclaration) - eventId.log(*fusEventData) - } - } -} - -/** - * Registers successful ML sessions' logging as a separate FUS event. - * - * See [MLSessionFinishedLogger] - */ -@ApiStatus.Internal -open class FinishedSessionLoggerRegister( - private val targetApproachClass: Class>, - private val fusScheme: FusSessionEventBuilder.FusScheme

, -) : MLApiStartupListener, MLApiStartupListenerProvider.ThisProvider { - override fun onBeforeStarted(apiPlatform: MLApiPlatform): MLApiStartupProcessListener { - return object : MLApiStartupProcessListener { - override fun onStartedInitializingFus(initializedApproaches: Collection>) { - @Suppress("UNCHECKED_CAST") - val targetInitializedApproach: MLTaskApproach

= requireNotNull(initializedApproaches.find { it.approach.javaClass == targetApproachClass }) { - """ - Could not create logger for finished sessions of $targetApproachClass in platform $apiPlatform, because the corresponding approach was not initialized. - Initialized approaches: $initializedApproaches - """.trimIndent() - }.approach as MLTaskApproach

- val finishedEventLogger = MLSessionFinishedLogger(targetInitializedApproach, fusScheme, apiPlatform) - apiPlatform.addMLEventBeforeFusInitialized(finishedEventLogger) - apiPlatform.addTaskListener(finishedEventLogger) - } - } - } -} diff --git a/platform/ml-impl/src/com/intellij/platform/ml/impl/logger/FusSessionEventBuilder.kt b/platform/ml-impl/src/com/intellij/platform/ml/impl/logs/EventSessionEventBuilder.kt similarity index 66% rename from platform/ml-impl/src/com/intellij/platform/ml/impl/logger/FusSessionEventBuilder.kt rename to platform/ml-impl/src/com/intellij/platform/ml/impl/logs/EventSessionEventBuilder.kt index f2a13c183247..5ba2c6a7286b 100644 --- a/platform/ml-impl/src/com/intellij/platform/ml/impl/logger/FusSessionEventBuilder.kt +++ b/platform/ml-impl/src/com/intellij/platform/ml/impl/logs/EventSessionEventBuilder.kt @@ -1,5 +1,5 @@ -// Copyright 2000-2023 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.logger +// 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 import com.intellij.internal.statistic.eventLog.events.EventPair import com.intellij.internal.statistic.eventLog.events.ObjectDescription @@ -10,7 +10,7 @@ import com.intellij.platform.ml.impl.session.AnalysedSessionTree import org.jetbrains.annotations.ApiStatus /** - * Represents FUS fields of a session's subtree. + * Represents FUS event fields of a session's subtree. */ @ApiStatus.Internal abstract class SessionFields

: ObjectDescription() { @@ -25,25 +25,24 @@ abstract class SessionFields

: ObjectDescription() { * @param P The type of the ML task's prediction */ @ApiStatus.Internal -interface FusSessionEventBuilder

{ +interface EventSessionEventBuilder

{ /** - * Configuration of a [FusSessionEventBuilder], that builds it when accepts approach's declaration. + * Configuration of a [EventSessionEventBuilder], that builds it when accepts approach's declaration. */ - interface FusScheme

{ - fun createEventBuilder(approachDeclaration: MLTaskApproach.Declaration): FusSessionEventBuilder

+ interface EventScheme

{ + fun createEventBuilder(approachDeclaration: MLTaskApproach.SessionDeclaration): EventSessionEventBuilder

} /** * Builds declaration of all features, that will be logged for the session tiers' description and analysis. - * It is required because FUS logs validators are built 'statically'. */ - fun buildFusDeclaration(): SessionFields

+ fun buildSessionFields(): SessionFields

/** - * Builds a concrete FUS record, that contains fields that were built by [buildFusDeclaration]. + * Builds a concrete log record, that contains fields that were built by [buildSessionFields]. * * @param sessionStructure A session tree that been already analyzed and is ready to be logged. - * @param sessionFields Session fields that were built by [buildFusDeclaration] earlier. + * @param sessionFields Session fields that were built by [buildSessionFields] earlier. */ fun buildRecord(sessionStructure: AnalysedRootContainer

, sessionFields: SessionFields

): Array> } diff --git a/platform/ml-impl/src/com/intellij/platform/ml/impl/logger/InplaceFeaturesScheme.kt b/platform/ml-impl/src/com/intellij/platform/ml/impl/logs/InplaceFeaturesScheme.kt similarity index 97% rename from platform/ml-impl/src/com/intellij/platform/ml/impl/logger/InplaceFeaturesScheme.kt rename to platform/ml-impl/src/com/intellij/platform/ml/impl/logs/InplaceFeaturesScheme.kt index b9d1988e2bf9..e75d4f63671d 100644 --- a/platform/ml-impl/src/com/intellij/platform/ml/impl/logger/InplaceFeaturesScheme.kt +++ b/platform/ml-impl/src/com/intellij/platform/ml/impl/logs/InplaceFeaturesScheme.kt @@ -1,5 +1,5 @@ -// Copyright 2000-2023 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.logger +// 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 import com.intellij.internal.statistic.eventLog.FeatureUsageData import com.intellij.internal.statistic.eventLog.events.* @@ -20,13 +20,13 @@ import org.jetbrains.annotations.ApiStatus class InplaceFeaturesScheme

internal constructor( private val predictionField: EventField, private val predictionTransformer: (P?) -> F?, - private val approachDeclaration: MLTaskApproach.Declaration -) : FusSessionEventBuilder

{ + private val approachDeclaration: MLTaskApproach.SessionDeclaration +) : EventSessionEventBuilder

{ class FusScheme

( private val predictionField: EventField, private val predictionTransformer: (P?) -> F?, - ) : FusSessionEventBuilder.FusScheme

{ - override fun createEventBuilder(approachDeclaration: MLTaskApproach.Declaration): FusSessionEventBuilder

= InplaceFeaturesScheme( + ) : EventSessionEventBuilder.EventScheme

{ + override fun createEventBuilder(approachDeclaration: MLTaskApproach.SessionDeclaration): EventSessionEventBuilder

= InplaceFeaturesScheme( predictionField, predictionTransformer, approachDeclaration @@ -37,7 +37,7 @@ class InplaceFeaturesScheme

internal constructor( } } - override fun buildFusDeclaration(): SessionFields

{ + override fun buildSessionFields(): SessionFields

{ require(approachDeclaration.levelsScheme.isNotEmpty()) return if (approachDeclaration.levelsScheme.size == 1) PredictionSessionFields(approachDeclaration.levelsScheme.first(), predictionField, predictionTransformer, diff --git a/platform/ml-impl/src/com/intellij/platform/ml/impl/logs/events/sessionFailed.kt b/platform/ml-impl/src/com/intellij/platform/ml/impl/logs/events/sessionFailed.kt new file mode 100644 index 000000000000..9aadaf47dbc4 --- /dev/null +++ b/platform/ml-impl/src/com/intellij/platform/ml/impl/logs/events/sessionFailed.kt @@ -0,0 +1,86 @@ +// 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.events + +import com.intellij.internal.statistic.eventLog.EventLogGroup +import com.intellij.internal.statistic.eventLog.events.ObjectEventData +import com.intellij.internal.statistic.eventLog.events.ObjectEventField +import com.intellij.internal.statistic.eventLog.events.VarargEventId +import com.intellij.platform.ml.Session +import com.intellij.platform.ml.impl.MLTask +import com.intellij.platform.ml.impl.MLTaskApproach +import com.intellij.platform.ml.impl.MLTaskApproach.Companion.findMlTaskApproach +import com.intellij.platform.ml.impl.MLTaskApproachBuilder +import com.intellij.platform.ml.impl.apiPlatform.MLApiPlatform +import com.intellij.platform.ml.impl.apiPlatform.ReplaceableIJPlatform +import com.intellij.platform.ml.impl.monitoring.MLApproachInitializationListener +import com.intellij.platform.ml.impl.monitoring.MLApproachListener +import com.intellij.platform.ml.impl.monitoring.MLTaskGroupListener +import com.intellij.platform.ml.impl.monitoring.MLTaskGroupListener.ApproachListeners.Companion.monitoredBy +import com.intellij.platform.ml.impl.session.analysis.ShallowSessionAnalyser +import com.intellij.platform.ml.impl.session.analysis.ShallowSessionAnalyser.Companion.declarationObjectDescription +import org.jetbrains.annotations.ApiStatus + +/** + * For the given [this] event group, registers an event, that will be corresponding to a failure of [task]'s start attempt. + * The event will be logged automatically. + * + * @param eventName Event's name in the event group + * @param task Task whose finish will be recorded + * @param exceptionalAnalysers Analyzers, that are triggered, when [MLTaskApproach.startSession] fails with an unhandled exception + * @param normalFailureAnalysers Analyzers, that aer triggered, when [MLTaskApproach.startSession] returns a [Session.StartOutcome.Failure] + * @param apiPlatform Platform that will be used to build event validators: it should include desired analysis and description features. + * + * @return Controller, that could disable event's logging by calling [MLApiPlatform.ExtensionController.removeExtension] + */ +@ApiStatus.Internal +fun EventLogGroup.registerEventSessionFailed( + eventName: String, + task: MLTask

, + exceptionalAnalysers: Collection>, + normalFailureAnalysers: Collection>>, + apiPlatform: MLApiPlatform = ReplaceableIJPlatform +): MLApiPlatform.ExtensionController { + val approachBuilder = findMlTaskApproach(task, apiPlatform) + val fields: Map = (normalFailureAnalysers + exceptionalAnalysers) + .map { ObjectEventField(it.name, it.declarationObjectDescription) } + .associateBy { it.name } + val eventId = registerVarargEvent(eventName, *fields.values.toTypedArray()) + val logger = MLSessionFailedLogger(approachBuilder, exceptionalAnalysers, normalFailureAnalysers, eventId, fields) + return apiPlatform.addTaskListener(logger) +} + +private class MLSessionFailedLogger( + private val taskApproachBuilder: MLTaskApproachBuilder

, + private val exceptionalAnalysers: Collection>, + private val normalFailureAnalysers: Collection>>, + private val eventId: VarargEventId, + private val fields: Map, +) : MLTaskGroupListener { + + override val approachListeners: Collection> + get() { + return listOf( + taskApproachBuilder.javaClass monitoredBy MLApproachInitializationListener { permanentSessionEnvironment -> + object : MLApproachListener { + override fun onFailedToStartSessionWithException(exception: Throwable) { + val analysis = exceptionalAnalysers.map { + val analyserField = fields.getValue(it.name) + analyserField with ObjectEventData(it.analyse(permanentSessionEnvironment, exception)) + } + eventId.log(*analysis.toTypedArray()) + } + + override fun onFailedToStartSession(failure: Session.StartOutcome.Failure

) { + val analysis = normalFailureAnalysers.map { + val analyserField = fields.getValue(it.name) + analyserField with ObjectEventData(it.analyse(permanentSessionEnvironment, failure)) + } + eventId.log(*analysis.toTypedArray()) + } + + override fun onStartedSession(session: Session

) = null + } + } + ) + } +} diff --git a/platform/ml-impl/src/com/intellij/platform/ml/impl/logs/events/sessionFinished.kt b/platform/ml-impl/src/com/intellij/platform/ml/impl/logs/events/sessionFinished.kt new file mode 100644 index 000000000000..e0bbd7e777c7 --- /dev/null +++ b/platform/ml-impl/src/com/intellij/platform/ml/impl/logs/events/sessionFinished.kt @@ -0,0 +1,81 @@ +// 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.events + +import com.intellij.internal.statistic.eventLog.EventLogGroup +import com.intellij.internal.statistic.eventLog.events.VarargEventId +import com.intellij.platform.ml.Environment +import com.intellij.platform.ml.Session +import com.intellij.platform.ml.impl.MLTask +import com.intellij.platform.ml.impl.MLTaskApproach.Companion.findMlTaskApproach +import com.intellij.platform.ml.impl.MLTaskApproachBuilder +import com.intellij.platform.ml.impl.apiPlatform.MLApiPlatform +import com.intellij.platform.ml.impl.apiPlatform.ReplaceableIJPlatform +import com.intellij.platform.ml.impl.logs.EventSessionEventBuilder +import com.intellij.platform.ml.impl.logs.SessionFields +import com.intellij.platform.ml.impl.monitoring.MLApproachInitializationListener +import com.intellij.platform.ml.impl.monitoring.MLApproachListener +import com.intellij.platform.ml.impl.monitoring.MLSessionListener +import com.intellij.platform.ml.impl.monitoring.MLTaskGroupListener +import com.intellij.platform.ml.impl.monitoring.MLTaskGroupListener.ApproachListeners.Companion.monitoredBy +import com.intellij.platform.ml.impl.session.AnalysedRootContainer +import com.intellij.platform.ml.impl.session.DescribedRootContainer +import org.jetbrains.annotations.ApiStatus + +/** + * For the given [this] event group, registers an event, that will be corresponding to a finish of [task]'s session. + * The event will be logged automatically. + * + * @param eventName Event's name in the event group + * @param task Task whose finish will be recorded + * @param scheme Chosen scheme for the session structure + * @param apiPlatform Platform that will be used to build event validators: it should include desired analysis and description features. + * + * @return Controller, that could disable event's logging by calling [MLApiPlatform.ExtensionController.removeExtension] + */ +@ApiStatus.Internal +fun EventLogGroup.registerEventSessionFinished( + eventName: String, + task: MLTask

, + scheme: EventSessionEventBuilder.EventScheme

, + apiPlatform: MLApiPlatform = ReplaceableIJPlatform +): MLApiPlatform.ExtensionController { + val approachBuilder = findMlTaskApproach(task, apiPlatform) + val sessionDeclaration = approachBuilder.buildApproachSessionDeclaration(apiPlatform) + val finishedSessionEventBuilder: EventSessionEventBuilder

= scheme.createEventBuilder(sessionDeclaration) + val sessionFields: SessionFields

= finishedSessionEventBuilder.buildSessionFields() + val eventId = registerVarargEvent(eventName, *sessionFields.getFields()) + val logger = MLSessionFinishedLogger(approachBuilder, eventId, finishedSessionEventBuilder, sessionFields) + return apiPlatform.addTaskListener(logger) +} + +private class MLSessionFinishedLogger( + approachBuilder: MLTaskApproachBuilder

, + private val eventId: VarargEventId, + private val finishedSessionEventBuilder: EventSessionEventBuilder

, + private val sessionFields: SessionFields

+) : MLTaskGroupListener { + override val approachListeners: Collection> = listOf( + approachBuilder.javaClass monitoredBy InitializationLogger() + ) + + inner class InitializationLogger : MLApproachInitializationListener { + override fun onAttemptedToStartSession(permanentSessionEnvironment: Environment): MLApproachListener = ApproachLogger() + } + + inner class ApproachLogger : MLApproachListener { + override fun onFailedToStartSessionWithException(exception: Throwable) {} + + override fun onFailedToStartSession(failure: Session.StartOutcome.Failure

) {} + + override fun onStartedSession(session: Session

): MLSessionListener = SessionLogger() + } + + inner class SessionLogger : MLSessionListener { + override fun onSessionDescriptionFinished(sessionTree: DescribedRootContainer) {} + + override fun onSessionAnalysisFinished(sessionTree: AnalysedRootContainer

) { + val fusEventData = finishedSessionEventBuilder.buildRecord(sessionTree, sessionFields) + eventId.log(*fusEventData) + } + } +} diff --git a/platform/ml-impl/src/com/intellij/platform/ml/impl/monitoring/MLApiStartupProcessListener.kt b/platform/ml-impl/src/com/intellij/platform/ml/impl/monitoring/MLApiStartupProcessListener.kt deleted file mode 100644 index f19be7f01751..000000000000 --- a/platform/ml-impl/src/com/intellij/platform/ml/impl/monitoring/MLApiStartupProcessListener.kt +++ /dev/null @@ -1,31 +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.monitoring - -import com.intellij.platform.ml.impl.MLTaskApproach -import com.intellij.platform.ml.impl.MLTaskApproachInitializer -import com.intellij.platform.ml.impl.apiPlatform.MLApiPlatform -import org.jetbrains.annotations.ApiStatus - -@ApiStatus.Internal -fun interface MLApiStartupListener { - fun onBeforeStarted(apiPlatform: MLApiPlatform): MLApiStartupProcessListener -} - -@ApiStatus.Internal -data class InitializerAndApproach( - val initializer: MLTaskApproachInitializer, - val approach: MLTaskApproach -) - -@ApiStatus.Internal -interface MLApiStartupProcessListener { - fun onStartedInitializingApproaches() {} - - fun onStartedInitializingFus(initializedApproaches: Collection>) {} - - fun onFinished() {} - - fun onFailed(exception: Throwable?) {} - - fun onCanceled() {} -} diff --git a/platform/ml-impl/src/com/intellij/platform/ml/impl/monitoring/MLApproachListener.kt b/platform/ml-impl/src/com/intellij/platform/ml/impl/monitoring/MLApproachListener.kt index 5020d357a91a..35fd46f200da 100644 --- a/platform/ml-impl/src/com/intellij/platform/ml/impl/monitoring/MLApproachListener.kt +++ b/platform/ml-impl/src/com/intellij/platform/ml/impl/monitoring/MLApproachListener.kt @@ -4,6 +4,8 @@ package com.intellij.platform.ml.impl.monitoring import com.intellij.platform.ml.Environment import com.intellij.platform.ml.Session import com.intellij.platform.ml.impl.MLTaskApproach +import com.intellij.platform.ml.impl.MLTaskApproachBuilder +import com.intellij.platform.ml.impl.apiPlatform.MLTaskGroupListenerProvider import com.intellij.platform.ml.impl.monitoring.MLApproachInitializationListener.Companion.asJoinedListener import com.intellij.platform.ml.impl.monitoring.MLApproachListener.Companion.asJoinedListener import com.intellij.platform.ml.impl.monitoring.MLSessionListener.Companion.asJoinedListener @@ -15,8 +17,8 @@ import org.jetbrains.annotations.ApiStatus /** * Provides listeners for a set of [com.intellij.platform.ml.impl.MLTaskApproach] * - * Only [com.intellij.platform.ml.impl.approach.LogDrivenModelInference] and the subclasses, that are - * calling [com.intellij.platform.ml.impl.approach.LogDrivenModelInference.startSession] are monitored. + * Only [com.intellij.platform.ml.impl.LogDrivenModelInference] and the subclasses, that are + * calling [com.intellij.platform.ml.impl.LogDrivenModelInference.startSession] are monitored. */ @ApiStatus.Internal interface MLTaskGroupListener { @@ -35,27 +37,34 @@ interface MLTaskGroupListener { * A proper way to create it is to use [monitoredBy] */ data class ApproachListeners internal constructor( - val taskApproach: Class>, + val taskApproachBuilder: Class>, val approachListener: Collection> ) { companion object { - infix fun Class>.monitoredBy(approachListener: MLApproachInitializationListener) = ApproachListeners( + infix fun Class>.monitoredBy(approachListener: MLApproachInitializationListener) = ApproachListeners( this, listOf(approachListener)) - infix fun Class>.monitoredBy(approachListeners: Collection>) = ApproachListeners( + infix fun Class>.monitoredBy(approachListeners: Collection>) = ApproachListeners( this, approachListeners) } } - companion object { - internal val MLTaskGroupListener.targetedApproaches: Set>> - get() = approachListeners.map { it.taskApproach }.toSet() + @ApiStatus.Internal + interface Default : MLTaskGroupListener, MLTaskGroupListenerProvider { + override fun provide(collector: (MLTaskGroupListener) -> Unit) { + return collector(this) + } + } - internal fun

MLTaskGroupListener.onAttemptedToStartSession(taskApproach: MLTaskApproach

, + companion object { + internal val MLTaskGroupListener.targetedApproaches: Set>> + get() = approachListeners.map { it.taskApproachBuilder }.toSet() + + internal fun

MLTaskGroupListener.onAttemptedToStartSession(taskApproachBuilder: MLTaskApproachBuilder

, permanentSessionEnvironment: Environment): MLApproachListener? { @Suppress("UNCHECKED_CAST") val approachListeners: List> = approachListeners - .filter { it.taskApproach == taskApproach.javaClass } + .filter { it.taskApproachBuilder == taskApproachBuilder.javaClass } .flatMap { it.approachListener } as List> return approachListeners.asJoinedListener().onAttemptedToStartSession(permanentSessionEnvironment) } @@ -65,11 +74,14 @@ interface MLTaskGroupListener { /** * Listens to the attempt of starting new [Session] of the [MLTaskApproach], that this listener was put * into correspondence to via [com.intellij.platform.ml.impl.monitoring.MLTaskGroupListener.ApproachListeners.Companion.monitoredBy] + * + * @param R Type of the [com.intellij.platform.ml.impl.model.MLModel] + * @param P Prediction's type */ @ApiStatus.Internal fun interface MLApproachInitializationListener { /** - * Called each time, when [com.intellij.platform.ml.impl.approach.LogDrivenModelInference.startSession] is invoked + * Called each time, when [com.intellij.platform.ml.impl.LogDrivenModelInference.startSession] is invoked * * @return A listener, that will be monitoring how successful the start was. If it is not needed, null is returned. */ @@ -85,25 +97,25 @@ fun interface MLApproachInitializationListener { } /** - * Listens to the process of starting new [Session] of [com.intellij.platform.ml.impl.approach.LogDrivenModelInference]. + * Listens to the process of starting new [Session] of [com.intellij.platform.ml.impl.LogDrivenModelInference]. */ @ApiStatus.Internal interface MLApproachListener { /** * Called if the session was not started, * on exceptionally rare occasions, - * when the [com.intellij.platform.ml.impl.approach.LogDrivenModelInference.startSession] failed with an exception + * when the [com.intellij.platform.ml.impl.LogDrivenModelInference.startSession] failed with an exception */ - fun onFailedToStartSessionWithException(exception: Throwable) + fun onFailedToStartSessionWithException(exception: Throwable) {} /** * Called if the session was not started, * but the failure is 'ordinary'. */ - fun onFailedToStartSession(failure: Session.StartOutcome.Failure

) + fun onFailedToStartSession(failure: Session.StartOutcome.Failure

) {} /** - * Called when a new [com.intellij.platform.ml.impl.approach.LogDrivenModelInference]'s session was started successfully. + * Called when a new [com.intellij.platform.ml.impl.LogDrivenModelInference]'s session was started successfully. * * @return A listener for tracking the session's progress, null if the session will not be tracked. */ @@ -131,7 +143,7 @@ interface MLApproachListener { } /** - * Listens to session events of a [com.intellij.platform.ml.impl.approach.LogDrivenModelInference] + * Listens to session events of a [com.intellij.platform.ml.impl.LogDrivenModelInference] */ @ApiStatus.Internal interface MLSessionListener { @@ -139,14 +151,14 @@ interface MLSessionListener { * All tier instances were established (the tree will not be growing further), * described, and predictions in the [sessionTree] were finished. */ - fun onSessionDescriptionFinished(sessionTree: DescribedRootContainer) + fun onSessionDescriptionFinished(sessionTree: DescribedRootContainer) {} /** * Called only after [onSessionDescriptionFinished] * * All tree nodes were analyzed. */ - fun onSessionAnalysisFinished(sessionTree: AnalysedRootContainer

) + fun onSessionAnalysisFinished(sessionTree: AnalysedRootContainer

) {} companion object { fun Collection>.asJoinedListener(): MLSessionListener { diff --git a/platform/ml-impl/src/com/intellij/platform/ml/impl/session/LevelDescriptor.kt b/platform/ml-impl/src/com/intellij/platform/ml/impl/session/LevelDescriptor.kt index 8955f016848a..e358aa120a9d 100644 --- a/platform/ml-impl/src/com/intellij/platform/ml/impl/session/LevelDescriptor.kt +++ b/platform/ml-impl/src/com/intellij/platform/ml/impl/session/LevelDescriptor.kt @@ -1,12 +1,16 @@ // Copyright 2000-2023 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.session +import com.intellij.openapi.diagnostic.debug +import com.intellij.openapi.diagnostic.thisLogger import com.intellij.platform.ml.* +import com.intellij.platform.ml.Feature.Companion.toCompactString import com.intellij.platform.ml.ScopeEnvironment.Companion.restrictedBy import com.intellij.platform.ml.TierRequester.Companion.fulfilledBy import com.intellij.platform.ml.impl.DescriptionComputer import com.intellij.platform.ml.impl.FeatureSelector import com.intellij.platform.ml.impl.FeatureSelector.Companion.or +import com.intellij.platform.ml.impl.MLTask import com.intellij.platform.ml.impl.apiPlatform.MLApiPlatform import com.intellij.platform.ml.impl.apiPlatform.MLApiPlatform.Companion.getDescriptorsOfTiers import com.intellij.platform.ml.impl.environment.ExtendedEnvironment @@ -17,7 +21,8 @@ data class LevelDescriptor( val apiPlatform: MLApiPlatform, val descriptionComputer: DescriptionComputer, val usedFeaturesSelectors: PerTier, - val notUsedFeaturesSelectors: PerTier, + val notUsedFeaturesFilters: PerTier, + val mlTask: MLTask<*> ) { suspend fun describe( nextLevelCallParameters: Environment, @@ -25,6 +30,15 @@ data class LevelDescriptor( upperLevels: List, nextLevelAdditionalTiers: Set> ): DescribedLevel { + thisLogger().debug { + """ + [${mlTask.name}] Describing level environment: + Call parameters: ${nextLevelCallParameters.tiers} + Main environment: ${nextLevelMainEnvironment.tiers} + Additional environment: $nextLevelAdditionalTiers + """.trimIndent() + } + val mainEnvironment = Environment.joined(listOf( Environment.of(upperLevels.flatMap { it.mainInstances.keys + it.additionalInstances.keys }), nextLevelMainEnvironment @@ -33,14 +47,28 @@ data class LevelDescriptor( val availableEnvironment = ExtendedEnvironment(apiPlatform.environmentExtenders, mainEnvironment) val extendedEnvironment = availableEnvironment.restrictedBy(nextLevelMainEnvironment.tiers + nextLevelAdditionalTiers) + thisLogger().debug { + """ + [${mlTask.name}] Built the environment for description + Available environment: ${availableEnvironment.tiers} + Tier that will be described: ${extendedEnvironment.tiers} + """.trimIndent() + } + val runnableDescriptorsPerTier = apiPlatform.getDescriptorsOfTiers(extendedEnvironment.tiers) .mapValues { (_, descriptors) -> descriptors.fulfilledBy(availableEnvironment) } - val extendedEnvironmentDescription = runnableDescriptorsPerTier + val extendedEnvironmentDescription: Map, Declaredness>>> = runnableDescriptorsPerTier .mapValues { (tier, tierDescriptors) -> describeTier(tier, tierDescriptors, availableEnvironment) } + thisLogger().debug { + "[${mlTask.name}] Described level environment:\n" + extendedEnvironmentDescription.map { (tier, description) -> + "\t - $tier: $description" + }.joinToString("\n") + } + val nextLevel = DescribedLevel( mainInstances = nextLevelMainEnvironment.tierInstances.associateWith { mainTierInstance -> DescribedTierData(extendedEnvironmentDescription.getValue(mainTierInstance.tier)) @@ -59,7 +87,7 @@ data class LevelDescriptor( return FeatureFilter.ACCEPT_ALL } - val tierFeaturesSelector = (usedFeaturesSelectors[tier] ?: FeatureSelector.NOTHING) or notUsedFeaturesSelectors.getValue(tier) + val tierFeaturesSelector = (usedFeaturesSelectors[tier] ?: FeatureSelector.NOTHING) or notUsedFeaturesFilters.getValue(tier) val tierComputableFeatures = tierDescriptors.flatMap { it.descriptionDeclaration }.toSet() val tierToComputeSelection = tierFeaturesSelector.select(tierComputableFeatures) @@ -80,7 +108,7 @@ data class LevelDescriptor( private fun removeRedundantDescription(descriptor: TierDescriptor, computedDescriptionWithRedundancies: Set): Set { val usedFeaturesSelector = usedFeaturesSelectors[descriptor.tier] ?: FeatureSelector.NOTHING - val notUsedFeaturesSelector = notUsedFeaturesSelectors.getValue(descriptor.tier) + val notUsedFeaturesSelector = notUsedFeaturesFilters.getValue(descriptor.tier) return computedDescriptionWithRedundancies.mapNotNull { computedFeature -> val featureIsUsed = usedFeaturesSelector.select(computedFeature.declaration) @@ -88,7 +116,13 @@ data class LevelDescriptor( if (!featureIsUsed && !featureIsNotUsed) { if (!descriptor.descriptionPolicy.tolerateRedundantDescription) throw IllegalArgumentException( - "Feature $computedFeature of $descriptor must not have been computed. It is not used by the ML model or marked as not used" + """ + Feature $computedFeature of $descriptor must not have been computed. It is not used by the ML model or marked as not used. + + You could set DescriptionPolicy.tolerateRedundantDescription to true, if this descriptor computes lightweight features, + and redundantly computed features could be tolerated. + + """.trimIndent() ) else return@mapNotNull null } @@ -134,12 +168,15 @@ data class LevelDescriptor( require(expectedButNotComputedDescriptionDeclaration.isEmpty()) { """ - Features ${expectedButNotComputedDescriptionDeclaration.map { it.name }} - were expected to be computed by $descriptor, - because was declared and accepted by the feature filter. - If the features are nullable, you should mark the declarations as .nullable(), - and then either put the 'null' to the result set explicitly, or mark the corresponding - descriptor's descriptionPolicy as 'putNullImplicitly'. + Some expected features were not computed + + Features ${expectedButNotComputedDescriptionDeclaration.map { it.name }} + were expected to be computed by $descriptor, because was declared and accepted by the feature filter. + + If the features are nullable, you should mark the declarations as .nullable(), + and then either put the 'null' to the result set explicitly, or mark the corresponding + descriptor's descriptionPolicy as 'putNullImplicitly'. + """.trimIndent() } } @@ -174,12 +211,23 @@ data class LevelDescriptor( private suspend fun describeTier(tier: Tier<*>, tierDescriptors: List, environment: Environment): DescriptionPartition { val toComputeFilter = createFilterOfFeaturesToCompute(tier, tierDescriptors) val usefulTierDescriptors = tierDescriptors.filter { it.couldBeUseful(toComputeFilter) } + thisLogger().debug { + """ + [${mlTask.name}] Describing $tier + - Usable descriptors: $usefulTierDescriptors + - Not usable descriptors: ${tierDescriptors.filterNot { it in usefulTierDescriptors }} + """.trimIndent() + } val description: Map> = descriptionComputer.computeDescription( tier, usefulTierDescriptors, environment, toComputeFilter ) + thisLogger().debug { + "[${mlTask.name}] Computed description of $tier\n" + + description.entries.joinToString("\n") { (descriptor, features) -> "\t- ${descriptor.javaClass.simpleName}: ${features.map { it.toCompactString() }}" } + } val usedByModelFilter = createFilterOfUsedFeatures(tier) @@ -187,7 +235,7 @@ data class LevelDescriptor( description.entries .map { (tierDescriptor, computedDescription) -> val redundanciesFreeDescription = removeRedundantDescription(tierDescriptor, computedDescription) - validateDescription(tierDescriptor, redundanciesFreeDescription, usedByModelFilter) + validateDescription(tierDescriptor, redundanciesFreeDescription, toComputeFilter) makeDescriptionPartition(tierDescriptor, redundanciesFreeDescription, usedByModelFilter) } .reduceOrNull { first, second -> diff --git a/platform/ml-impl/src/com/intellij/platform/ml/impl/session/SessionStructureCollector.kt b/platform/ml-impl/src/com/intellij/platform/ml/impl/session/SessionStructureCollector.kt index 8eb955d7ec1e..647fb0c4c5d9 100644 --- a/platform/ml-impl/src/com/intellij/platform/ml/impl/session/SessionStructureCollector.kt +++ b/platform/ml-impl/src/com/intellij/platform/ml/impl/session/SessionStructureCollector.kt @@ -1,10 +1,14 @@ // Copyright 2000-2023 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.session +import com.intellij.openapi.diagnostic.debug +import com.intellij.openapi.diagnostic.thisLogger import com.intellij.platform.ml.* +import com.intellij.platform.ml.Feature.Companion.toCompactString import com.intellij.platform.ml.impl.DescriptionComputer import com.intellij.platform.ml.impl.FeatureSelector import com.intellij.platform.ml.impl.LevelTiers +import com.intellij.platform.ml.impl.MLTask import com.intellij.platform.ml.impl.apiPlatform.MLApiPlatform import com.intellij.platform.ml.impl.model.MLModel import org.jetbrains.annotations.ApiStatus @@ -21,8 +25,9 @@ class RootCollector, P : Any> private constructor( levelsTiers: List, usedFeaturesSelectors: PerTier, override val levelDescriptor: LevelDescriptor, - override val levelPositioning: LevelPositioning -) : NestableStructureCollector, M, P>(usedFeaturesSelectors) { + override val levelPositioning: LevelPositioning, + mlTask: MLTask

+) : NestableStructureCollector, M, P>(usedFeaturesSelectors, mlTask) { init { validateSuperiorCollector(initializationCallParameters, callParameters.first(), levelsTiers, levelMainEnvironment, levelAdditionalTiers, mlModel, notUsedFeaturesSelectors) @@ -36,7 +41,7 @@ class RootCollector, P : Any> private constructor( companion object { suspend fun , P : Any> build( initializationCallParameters: Environment, - callParameters: List>>, + mlTask: MLTask

, descriptionComputer: DescriptionComputer, notUsedFeaturesSelectors: PerTier, levelMainEnvironment: Environment, @@ -44,11 +49,11 @@ class RootCollector, P : Any> private constructor( mlModel: M, apiPlatform: MLApiPlatform, levelsTiers: List, - usedFeaturesSelectors: PerTier + usedFeaturesSelectors: PerTier, ): RootCollector { - val levelDescriptor = LevelDescriptor(apiPlatform, descriptionComputer, mlModel.knownFeatures, notUsedFeaturesSelectors) - val levelPositioning = LevelPositioning.superior(initializationCallParameters, callParameters, levelsTiers, levelMainEnvironment, levelDescriptor, levelAdditionalTiers) - return RootCollector(initializationCallParameters, callParameters, notUsedFeaturesSelectors, levelMainEnvironment, levelAdditionalTiers, mlModel, levelsTiers, usedFeaturesSelectors, levelDescriptor, levelPositioning) + val levelDescriptor = LevelDescriptor(apiPlatform, descriptionComputer, mlModel.knownFeatures, notUsedFeaturesSelectors, mlTask) + val levelPositioning = LevelPositioning.superior(initializationCallParameters, mlTask.callParameters, levelsTiers, levelMainEnvironment, levelDescriptor, levelAdditionalTiers) + return RootCollector(initializationCallParameters, mlTask.callParameters, notUsedFeaturesSelectors, levelMainEnvironment, levelAdditionalTiers, mlModel, levelsTiers, usedFeaturesSelectors, levelDescriptor, levelPositioning, mlTask) } } } @@ -64,8 +69,9 @@ class SolitaryLeafCollector, P : Any> private constructor( levelScheme: LevelTiers, usedFeaturesSelectors: PerTier, override val levelDescriptor: LevelDescriptor, - override val levelPositioning: LevelPositioning -) : PredictionCollector, M, P>(usedFeaturesSelectors) { + override val levelPositioning: LevelPositioning, + mlTask: MLTask

, +) : PredictionCollector, M, P>(usedFeaturesSelectors, mlTask) { init { validateSuperiorCollector(initializationCallParameters, callParameters.first(), listOf(levelScheme), levelMainEnvironment, levelAdditionalTiers, mlModel, notUsedFeaturesSelectors) @@ -78,7 +84,7 @@ class SolitaryLeafCollector, P : Any> private constructor( companion object { suspend fun , P : Any> build( initializationCallParameters: Environment, - callParameters: List>>, + mlTask: MLTask

, descriptionComputer: DescriptionComputer, notUsedFeaturesSelectors: PerTier, levelMainEnvironment: Environment, @@ -86,12 +92,12 @@ class SolitaryLeafCollector, P : Any> private constructor( mlModel: M, apiPlatform: MLApiPlatform, levelScheme: LevelTiers, - usedFeaturesSelectors: PerTier + usedFeaturesSelectors: PerTier, ): SolitaryLeafCollector { - val levelDescriptor = LevelDescriptor(apiPlatform, descriptionComputer, mlModel.knownFeatures, notUsedFeaturesSelectors) - val levelPositioning = LevelPositioning.superior(initializationCallParameters, callParameters, listOf(levelScheme), levelMainEnvironment, + val levelDescriptor = LevelDescriptor(apiPlatform, descriptionComputer, mlModel.knownFeatures, notUsedFeaturesSelectors, mlTask) + val levelPositioning = LevelPositioning.superior(initializationCallParameters, mlTask.callParameters, listOf(levelScheme), levelMainEnvironment, levelDescriptor, levelAdditionalTiers) - return SolitaryLeafCollector(initializationCallParameters, callParameters, notUsedFeaturesSelectors, levelMainEnvironment, levelAdditionalTiers, mlModel, levelScheme, usedFeaturesSelectors, levelDescriptor, levelPositioning) + return SolitaryLeafCollector(initializationCallParameters, mlTask.callParameters, notUsedFeaturesSelectors, levelMainEnvironment, levelAdditionalTiers, mlModel, levelScheme, usedFeaturesSelectors, levelDescriptor, levelPositioning, mlTask) } } } @@ -100,8 +106,9 @@ class SolitaryLeafCollector, P : Any> private constructor( class BranchingCollector, P : Any>( override val levelDescriptor: LevelDescriptor, override val levelPositioning: LevelPositioning, - usedFeaturesSelectors: PerTier -) : NestableStructureCollector, M, P>(usedFeaturesSelectors) { + usedFeaturesSelectors: PerTier, + mlTask: MLTask

, +) : NestableStructureCollector, M, P>(usedFeaturesSelectors, mlTask) { override fun createTree(thisLevel: DescribedLevel, collectedNestedStructureTrees: List>): SessionTree.Branching { return SessionTree.Branching(thisLevel, collectedNestedStructureTrees) @@ -110,7 +117,8 @@ class BranchingCollector, P : Any>( @ApiStatus.Internal abstract class PredictionCollector, M : MLModel

, P : Any>( - private val usedFeaturesSelectors: PerTier + private val usedFeaturesSelectors: PerTier, + private val mlTask: MLTask

, ) : StructureCollector() { private var predictionSubmitted = false private var submittedPrediction: P? = null @@ -127,9 +135,37 @@ abstract class PredictionCollector { + return listOf("================ LEVEL #$index ================") + + listOf("Main:") + + mainInstances.map { (instance, description) -> listOf("* ${instance.tier}").withIndent() + description.description.prettyPrinted().withIndent(2) }.flatten() + + listOf("Additional:") + + additionalInstances.map { (instance, description) -> listOf("* ${instance.tier}").withIndent() + description.description.prettyPrinted().withIndent(2) }.flatten() + + listOf("Call parameters: ${this.callParameters.tiers}") + } + + private fun DescriptionPartition.prettyPrinted(): List { + return listOf( + "declared, used: ${this.declared.used.map { it.toCompactString() }}", + "declared, not used: ${this.declared.notUsed.map { it.toCompactString() }}", + "non declared, used: ${this.nonDeclared.used.map { it.toCompactString() }}", + "non declared, not used: ${this.nonDeclared.notUsed.map { it.toCompactString() }}" + ) + } + + private fun List.withIndent(n: Int = 1): List { + return map { "\t".repeat(n) + it } + } + private fun PerTierInstance.extractDescriptionForModel(): PerTier> { return this.entries.filter { it.key.tier in usedFeaturesSelectors }.associate { (tierInstance, data) -> tierInstance.tier to data.description.declared.used + data.description.nonDeclared.used @@ -151,8 +187,9 @@ abstract class PredictionCollector, P : Any>( override val levelDescriptor: LevelDescriptor, override val levelPositioning: LevelPositioning, - usedFeaturesSelectors: PerTier -) : PredictionCollector, M, P>(usedFeaturesSelectors) { + usedFeaturesSelectors: PerTier, + mlTask: MLTask

, +) : PredictionCollector, M, P>(usedFeaturesSelectors, mlTask) { override fun createTree(thisLevel: DescribedLevel, prediction: P?): SessionTree.Leaf { return SessionTree.Leaf(levelPositioning.thisLevel, prediction) @@ -255,7 +292,8 @@ sealed class StructureCollector, M : MLModel

, @ApiStatus.Internal abstract class NestableStructureCollector, M : MLModel

, P : Any>( - private val usedFeaturesSelectors: PerTier + private val usedFeaturesSelectors: PerTier, + private val mlTask: MLTask

) : StructureCollector() { private val nestedSessionsStructures: MutableList>> = mutableListOf() private var nestingFinished = false @@ -264,7 +302,8 @@ abstract class NestableStructureCollector, verifyNestedLevelEnvironment(levelMainEnvironment, levelAdditionalTiers) return BranchingCollector(levelDescriptor, levelPositioning.nestNextLevel(callParameters, levelMainEnvironment, levelAdditionalTiers, levelDescriptor), - usedFeaturesSelectors) + usedFeaturesSelectors, + mlTask) .also { it.trackCollectedStructure() } } @@ -273,7 +312,8 @@ abstract class NestableStructureCollector, return LeafCollector( levelDescriptor, levelPositioning.nestNextLevel(callParameters, levelMainEnvironment, levelAdditionalTiers, levelDescriptor), - usedFeaturesSelectors + usedFeaturesSelectors, + mlTask ).also { it.trackCollectedStructure() } } diff --git a/platform/ml-impl/src/com/intellij/platform/ml/impl/approach/GroupedAnalysis.kt b/platform/ml-impl/src/com/intellij/platform/ml/impl/session/analysis/GroupedAnalysis.kt similarity index 85% rename from platform/ml-impl/src/com/intellij/platform/ml/impl/approach/GroupedAnalysis.kt rename to platform/ml-impl/src/com/intellij/platform/ml/impl/session/analysis/GroupedAnalysis.kt index 708827a6b63c..711ce8285a8f 100644 --- a/platform/ml-impl/src/com/intellij/platform/ml/impl/approach/GroupedAnalysis.kt +++ b/platform/ml-impl/src/com/intellij/platform/ml/impl/session/analysis/GroupedAnalysis.kt @@ -1,11 +1,10 @@ -// Copyright 2000-2023 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.approach +// 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.session.analysis import com.intellij.platform.ml.Feature import com.intellij.platform.ml.FeatureDeclaration import com.intellij.platform.ml.impl.model.MLModel import com.intellij.platform.ml.impl.session.DescribedRootContainer -import com.intellij.platform.ml.impl.session.analysis.* import org.jetbrains.annotations.ApiStatus import java.util.concurrent.CompletableFuture @@ -14,7 +13,7 @@ import java.util.concurrent.CompletableFuture */ @ApiStatus.Internal data class GroupedAnalysisDeclaration( - val structureAnalysis: StructureAnalysisDeclaration, + val sessionStructureAnalysis: StructureAnalysisDeclaration, val mlModelAnalysis: Set> ) @@ -36,7 +35,7 @@ class JoinedGroupedSessionAnalyser, P : Any>( private val mlModelAnalysers: Collection>, ) : SessionAnalyser, M, P> { override val analysisDeclaration = GroupedAnalysisDeclaration( - structureAnalysis = SessionStructureAnalysisJoiner().joinDeclarations(structureAnalysers.map { it.analysisDeclaration }), + sessionStructureAnalysis = SessionStructureAnalysisJoiner().joinDeclarations(structureAnalysers.map { it.analysisDeclaration }), mlModelAnalysis = MLModelAnalysisJoiner().joinDeclarations(mlModelAnalysers.map { it.analysisDeclaration }) ) diff --git a/platform/ml-impl/src/com/intellij/platform/ml/impl/approach/LanguageSpecific.kt b/platform/ml-impl/src/com/intellij/platform/ml/impl/session/analysis/LanguageSpecific.kt similarity index 86% rename from platform/ml-impl/src/com/intellij/platform/ml/impl/approach/LanguageSpecific.kt rename to platform/ml-impl/src/com/intellij/platform/ml/impl/session/analysis/LanguageSpecific.kt index bc7197aa188f..c2a1d9580c11 100644 --- a/platform/ml-impl/src/com/intellij/platform/ml/impl/approach/LanguageSpecific.kt +++ b/platform/ml-impl/src/com/intellij/platform/ml/impl/session/analysis/LanguageSpecific.kt @@ -1,12 +1,11 @@ -// Copyright 2000-2023 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.approach +// 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.session.analysis import com.intellij.lang.Language import com.intellij.platform.ml.Feature import com.intellij.platform.ml.FeatureDeclaration import com.intellij.platform.ml.impl.model.MLModel import com.intellij.platform.ml.impl.session.DescribedRootContainer -import com.intellij.platform.ml.impl.session.analysis.MLModelAnalyser import org.jetbrains.annotations.ApiStatus import java.util.concurrent.CompletableFuture diff --git a/platform/ml-impl/src/com/intellij/platform/ml/impl/approach/StructureAndModelAnalysis.kt b/platform/ml-impl/src/com/intellij/platform/ml/impl/session/analysis/MLSessionAnalyser.kt similarity index 79% rename from platform/ml-impl/src/com/intellij/platform/ml/impl/approach/StructureAndModelAnalysis.kt rename to platform/ml-impl/src/com/intellij/platform/ml/impl/session/analysis/MLSessionAnalyser.kt index 7f670f52e298..0881de9149cd 100644 --- a/platform/ml-impl/src/com/intellij/platform/ml/impl/approach/StructureAndModelAnalysis.kt +++ b/platform/ml-impl/src/com/intellij/platform/ml/impl/session/analysis/MLSessionAnalyser.kt @@ -1,13 +1,10 @@ -// Copyright 2000-2023 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.approach +// 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.session.analysis import com.intellij.platform.ml.FeatureDeclaration import com.intellij.platform.ml.PerTierInstance import com.intellij.platform.ml.impl.model.MLModel import com.intellij.platform.ml.impl.session.* -import com.intellij.platform.ml.impl.session.analysis.MLModelAnalyser -import com.intellij.platform.ml.impl.session.analysis.StructureAnalyser -import com.intellij.platform.ml.impl.session.analysis.StructureAnalysisDeclaration import org.jetbrains.annotations.ApiStatus import java.util.concurrent.CompletableFuture @@ -20,26 +17,32 @@ import java.util.concurrent.CompletableFuture * @param sessionAnalysisKeyModel Key that will be used in logs, to write ML model's features into. */ @ApiStatus.Internal -class StructureAndModelAnalysis, P : Any>( +class MLSessionAnalyser, P : Any>( structureAnalysers: Collection>, mlModelAnalysers: Collection>, private val sessionAnalysisKeyModel: String = DEFAULT_SESSION_KEY_ML_MODEL -) : AnalysisMethod { +) { private val groupedAnalyser = JoinedGroupedSessionAnalyser(structureAnalysers, mlModelAnalysers) - override val structureAnalysisDeclaration: StructureAnalysisDeclaration - get() = groupedAnalyser.analysisDeclaration.structureAnalysis - - override val sessionAnalysisDeclaration: Map>> = mapOf( - sessionAnalysisKeyModel to groupedAnalyser.analysisDeclaration.mlModelAnalysis - ) - - override fun analyseTree(treeRoot: DescribedRootContainer): CompletableFuture> { + fun analyseSessionStructure(treeRoot: DescribedRootContainer): CompletableFuture> { return groupedAnalyser.analyse(treeRoot).thenApply { buildAnalysedSessionTree(treeRoot, it) as AnalysedRootContainer

} } + /** + * Static declaration of the features, that are used in the session tree's analysis. + */ + val structureAnalysisDeclaration: StructureAnalysisDeclaration + get() = groupedAnalyser.analysisDeclaration.sessionStructureAnalysis + + /** + * Static declaration of the session's entities, that are not tiers. + */ + val sessionAnalysisDeclaration: Map>> = mapOf( + sessionAnalysisKeyModel to groupedAnalyser.analysisDeclaration.mlModelAnalysis + ) + private fun buildAnalysedSessionTree(tree: DescribedSessionTree, analysis: GroupedAnalysis): AnalysedSessionTree

{ val treeAnalysisPerInstance: PerTierInstance = tree.levelData.mainInstances.entries .associate { (tierInstance, data) -> diff --git a/platform/ml-impl/src/com/intellij/platform/ml/impl/approach/Versioned.kt b/platform/ml-impl/src/com/intellij/platform/ml/impl/session/analysis/Versioned.kt similarity index 85% rename from platform/ml-impl/src/com/intellij/platform/ml/impl/approach/Versioned.kt rename to platform/ml-impl/src/com/intellij/platform/ml/impl/session/analysis/Versioned.kt index 89130dd75d4f..78a8b707c038 100644 --- a/platform/ml-impl/src/com/intellij/platform/ml/impl/approach/Versioned.kt +++ b/platform/ml-impl/src/com/intellij/platform/ml/impl/session/analysis/Versioned.kt @@ -1,12 +1,11 @@ -// Copyright 2000-2023 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.approach +// 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.session.analysis import com.intellij.openapi.util.Version import com.intellij.platform.ml.Feature import com.intellij.platform.ml.FeatureDeclaration import com.intellij.platform.ml.impl.model.MLModel import com.intellij.platform.ml.impl.session.DescribedRootContainer -import com.intellij.platform.ml.impl.session.analysis.MLModelAnalyser import org.jetbrains.annotations.ApiStatus import java.util.concurrent.CompletableFuture diff --git a/platform/ml-impl/src/com/intellij/platform/ml/impl/util.kt b/platform/ml-impl/src/com/intellij/platform/ml/impl/util.kt new file mode 100644 index 000000000000..cbb89a24b7cb --- /dev/null +++ b/platform/ml-impl/src/com/intellij/platform/ml/impl/util.kt @@ -0,0 +1,58 @@ +// 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.platform.ml.FeatureDeclaration +import org.jetbrains.annotations.ApiStatus +import kotlin.reflect.KClass +import kotlin.reflect.full.declaredMemberProperties +import kotlin.reflect.full.isSubclassOf +import kotlin.reflect.jvm.isAccessible + +/** + * Extracts all declarations of [FeatureDeclaration] from [declarationContainer]'s declaration + * + * The common usage is to write extractFieldsAsFeatureDeclarations(Companion) when the companion + * object contains the declarations. + * + * @param declarationContainer Object whose declaration contains features' descriptions + * @param extractors Types of properties that are desired to be extracted from. + */ +@ApiStatus.Internal +inline fun extractFieldsAsFeatureDeclarations(declarationContainer: T, + extractors: Collection = listOf(FeatureDeclarationsExtractor.OfProperty, + FeatureDeclarationsExtractor.ofIterable, + FeatureDeclarationsExtractor.ofMapValues) +): Set> { + return T::class.declaredMemberProperties + .flatMap { property -> + property.isAccessible = true + val propertyValue = property.call(declarationContainer) + val returnTypeClassifier = property.returnType.classifier + if (returnTypeClassifier !is KClass<*>) return@flatMap emptyList() + val extractedDeclarationsByExtractor = extractors.associateWith { it.extract(returnTypeClassifier, propertyValue) } + val workedExtractors = extractedDeclarationsByExtractor.entries.filter { it.value != null } + if (workedExtractors.isEmpty()) return@flatMap emptyList() + require(workedExtractors.size == 1) { "Ambiguous extraction of $property: ${workedExtractors}" } + return@flatMap requireNotNull(workedExtractors.first().value) + } + .toSet() +} + +@ApiStatus.Internal +fun interface FeatureDeclarationsExtractor { + fun extract(returnClass: KClass<*>, propertyValue: Any?): Collection>? + + companion object { + val OfProperty = FeatureDeclarationsExtractor { returnClass, propertyValue -> + if (returnClass.isSubclassOf(FeatureDeclaration::class)) listOf(propertyValue as FeatureDeclaration<*>) else null + } + + val ofIterable = FeatureDeclarationsExtractor { returnClass, propertyValue -> + if (returnClass.isSubclassOf(Iterable::class)) (propertyValue as Iterable<*>).filterIsInstance(FeatureDeclaration::class.java) else null + } + + val ofMapValues = FeatureDeclarationsExtractor { returnClass, propertyValue -> + if (returnClass.isSubclassOf(Map::class)) (propertyValue as Map<*, *>).values.filterIsInstance(FeatureDeclaration::class.java) else null + } + } +} diff --git a/platform/ml-impl/test/com/intellij/platform/ml/impl/Demo.kt b/platform/ml-impl/test/com/intellij/platform/ml/impl/Demo.kt index 7d513a465360..af73b886ecfd 100644 --- a/platform/ml-impl/test/com/intellij/platform/ml/impl/Demo.kt +++ b/platform/ml-impl/test/com/intellij/platform/ml/impl/Demo.kt @@ -2,31 +2,34 @@ package com.intellij.platform.ml.impl import com.intellij.internal.statistic.FUCollectorTestCase +import com.intellij.internal.statistic.eventLog.EventLogGroup import com.intellij.internal.statistic.eventLog.events.ClassEventField import com.intellij.internal.statistic.eventLog.events.EventField import com.intellij.internal.statistic.eventLog.events.EventPair +import com.intellij.internal.statistic.service.fus.collectors.CounterUsageCollectorEP +import com.intellij.internal.statistic.service.fus.collectors.CounterUsagesCollector +import com.intellij.internal.statistic.service.fus.collectors.UsageCollectors.COUNTER_EP_NAME import com.intellij.lang.Language import com.intellij.openapi.fileTypes.PlainTextLanguage import com.intellij.openapi.util.Version import com.intellij.platform.ml.* -import com.intellij.platform.ml.impl.MLTaskApproach.Companion.startCoroutineAndMLSession +import com.intellij.platform.ml.impl.MLTaskApproach.Companion.startMLSession import com.intellij.platform.ml.impl.apiPlatform.CodeLikePrinter import com.intellij.platform.ml.impl.apiPlatform.MLApiPlatform import com.intellij.platform.ml.impl.apiPlatform.ReplaceableIJPlatform -import com.intellij.platform.ml.impl.approach.* -import com.intellij.platform.ml.impl.logger.FailedSessionLoggerRegister -import com.intellij.platform.ml.impl.logger.FinishedSessionLoggerRegister -import com.intellij.platform.ml.impl.logger.InplaceFeaturesScheme -import com.intellij.platform.ml.impl.logger.MLEvent +import com.intellij.platform.ml.impl.logs.InplaceFeaturesScheme +import com.intellij.platform.ml.impl.logs.events.registerEventSessionFailed +import com.intellij.platform.ml.impl.logs.events.registerEventSessionFinished import com.intellij.platform.ml.impl.model.MLModel -import com.intellij.platform.ml.impl.monitoring.* +import com.intellij.platform.ml.impl.monitoring.MLApproachInitializationListener +import com.intellij.platform.ml.impl.monitoring.MLApproachListener +import com.intellij.platform.ml.impl.monitoring.MLSessionListener +import com.intellij.platform.ml.impl.monitoring.MLTaskGroupListener import com.intellij.platform.ml.impl.monitoring.MLTaskGroupListener.ApproachListeners.Companion.monitoredBy import com.intellij.platform.ml.impl.session.* -import com.intellij.platform.ml.impl.session.analysis.MLModelAnalyser -import com.intellij.platform.ml.impl.session.analysis.ShallowSessionAnalyser -import com.intellij.platform.ml.impl.session.analysis.StructureAnalyser -import com.intellij.platform.ml.impl.session.analysis.StructureAnalysis +import com.intellij.platform.ml.impl.session.analysis.* import com.intellij.testFramework.fixtures.BasePlatformTestCase +import com.intellij.util.application import com.jetbrains.fus.reporting.model.lion3.LogEvent import kotlinx.coroutines.runBlocking import java.net.URI @@ -70,7 +73,7 @@ object TierItem : Tier() object TierGit : Tier() -class CompletionSessionFeatures1 : TierDescriptor { +class CompletionSessionFeatures1 : TierDescriptor.Default(TierCompletionSession) { companion object { val CALL_ORDER = FeatureDeclaration.int("call_order") val LANGUAGE_ID = FeatureDeclaration.categorical("language_id", Language.getRegisteredLanguages().map { it.id }.toSet()) @@ -78,8 +81,6 @@ class CompletionSessionFeatures1 : TierDescriptor { val COMPLETION_TYPE = FeatureDeclaration.enum("completion_type").nullable() } - override val tier: Tier<*> = TierCompletionSession - override val descriptionPolicy: DescriptionPolicy = DescriptionPolicy(false, false) override val additionallyRequiredTiers: Set> = setOf(TierGit) @@ -100,14 +101,12 @@ class CompletionSessionFeatures1 : TierDescriptor { } } -class ItemFeatures1 : TierDescriptor { +class ItemFeatures1 : TierDescriptor.Default(TierItem) { companion object { val DECORATIONS = FeatureDeclaration.int("decorations") val LENGTH = FeatureDeclaration.int("length") } - override val tier: Tier<*> = TierItem - override val descriptionPolicy: DescriptionPolicy = DescriptionPolicy(false, false) override val descriptionDeclaration: Set> = setOf( @@ -123,14 +122,12 @@ class ItemFeatures1 : TierDescriptor { } } -class GitFeatures1 : TierDescriptor { +class GitFeatures1 : TierDescriptor.Default(TierGit) { companion object { val N_COMMITS = FeatureDeclaration.int("n_commits") val HAS_USER = FeatureDeclaration.boolean("has_user") } - override val tier: Tier<*> = TierGit - override val descriptionPolicy: DescriptionPolicy = DescriptionPolicy(false, false) override val descriptionDeclaration: Set> = setOf( @@ -218,11 +215,11 @@ class RandomModelSeedAnalyser : MLModelAnalyser { class RandomModel(val seed: Int) : MLModel, Versioned, LanguageSpecific { private val generator = Random(seed) - class Provider : MLModel.Provider { + object Provider : MLModel.Provider { private val generator = Random(1) override fun provideModel(callParameters: Environment, environment: Environment, sessionTiers: List): RandomModel? { return if (generator.nextBoolean()) { - if (generator.nextBoolean()) RandomModel(generator.nextInt()) else throw IllegalStateException() + if (generator.nextBoolean()) RandomModel(generator.nextInt()) else throw IllegalStateException("A random error was encountered") } else null } @@ -246,7 +243,7 @@ class RandomModel(val seed: Int) : MLModel, Versioned, LanguageSpecific class SomeListener(private val name: String) : MLTaskGroupListener { override val approachListeners = listOf( - MockTaskApproach::class.java monitoredBy InitializationListener() + MockTaskApproachBuilder::class.java monitoredBy InitializationListener() ) private fun log(message: String) = println("[Listener $name says] $message") @@ -260,7 +257,7 @@ class SomeListener(private val name: String) : MLTaskGroupListener { inner class ApproachListener : MLApproachListener { override fun onFailedToStartSessionWithException(exception: Throwable) { - log("failed to start session with exception: $exception") + log("failed to start session with exception: $exception, trace: ${exception.stackTraceToString()}") } override fun onFailedToStartSession(failure: Session.StartOutcome.Failure) { @@ -309,6 +306,17 @@ object FailureLogger : ShallowSessionAnalyser("failed", MockTask, listOf(ExceptionLogger), listOf(FailureLogger)) + it.registerEventSessionFinished("finished", MockTask, InplaceFeaturesScheme.FusScheme.DOUBLE) + } + } + + override fun getGroup() = GROUP +} + object ThisTestApiPlatform : TestApiPlatform() { override val tierDescriptors = listOf( CompletionSessionFeatures1(), @@ -321,20 +329,7 @@ object ThisTestApiPlatform : TestApiPlatform() { ) override val taskApproaches = listOf( - MockTaskApproach.Initializer() - ) - - - override val initialStartupListeners: List = listOf( - FinishedSessionLoggerRegister( - MockTaskApproach::class.java, - InplaceFeaturesScheme.FusScheme.DOUBLE - ), - FailedSessionLoggerRegister( - MockTaskApproach::class.java, - exceptionalAnalysers = listOf(ExceptionLogger), - normalFailureAnalysers = listOf(FailureLogger) - ) + MockTaskApproachBuilder() ) override val initialTaskListeners: List = listOf( @@ -342,8 +337,6 @@ object ThisTestApiPlatform : TestApiPlatform() { SomeListener("Alex"), ) - override val initialEvents: List = listOf() - override fun manageNonDeclaredFeatures(descriptor: ObsoleteTierDescriptor, nonDeclaredFeatures: Set) { val printer = CodeLikePrinter() @@ -366,27 +359,25 @@ object MockTask : MLTask( ) ) -class MockTaskApproach( - apiPlatform: MLApiPlatform, - task: MLTask -) : LogDrivenModelInference(task, apiPlatform) { +class MockTaskApproachDetails : LogDrivenModelInference.SessionDetails { override val additionallyDescribedTiers: List>> = listOf( setOf(TierGit), setOf(), setOf(), ) - override val analysisMethod: AnalysisMethod = StructureAndModelAnalysis( - structureAnalysers = listOf(SomeStructureAnalyser()), - mlModelAnalysers = listOf( + override val mlModelAnalysers: Collection> + get() = listOf( RandomModelSeedAnalyser(), ModelVersionAnalyser(), ModelLanguageAnalyser() ) - ) - override val mlModelProvider = RandomModel.Provider() + override val structureAnalysers: Collection> + get() = listOf(SomeStructureAnalyser()) + + override val mlModelProvider = RandomModel.Provider override fun getNotUsedDescription(callParameters: Environment, mlModel: RandomModel) = mapOf( TierCompletionSession to FeatureFilter.REJECT_ALL, @@ -397,12 +388,15 @@ class MockTaskApproach( override val descriptionComputer: DescriptionComputer = StateFreeDescriptionComputer - class Initializer : MLTaskApproachInitializer { - override val task: MLTask = MockTask - override fun initializeApproachWithin(apiPlatform: MLApiPlatform) = MockTaskApproach(apiPlatform, task) + class Builder : LogDrivenModelInference.SessionDetails.Builder { + override fun build(apiPlatform: MLApiPlatform): LogDrivenModelInference.SessionDetails { + return MockTaskApproachDetails() + } } } +class MockTaskApproachBuilder : LogDrivenModelInference.Builder(MockTask, MockTaskApproachDetails.Builder()) + class TestTask : BasePlatformTestCase() { fun `test demo ml task`() { // After the session is finished, it will be logged to community/platform/ml-impl/testResources/ml_logs.js @@ -416,26 +410,30 @@ class TestTask : BasePlatformTestCase() { //MLEventLogger.Manager.ensureNotInitialized() ReplaceableIJPlatform.replacingWith(ThisTestApiPlatform) { + val logger = MockTaskFusLogger() + val loggerEP = CounterUsageCollectorEP() + application.extensionArea.getExtensionPoint(COUNTER_EP_NAME).registerExtension(loggerEP, this.testRootDisposable) + FUCollectorTestCase.listenForEvents("FUS", this.testRootDisposable, collectLogs) { repeat(3) { sessionIndex -> println("Demo session #$sessionIndex has started") - val startOutcome = MockTask.startCoroutineAndMLSession( - callParameters = Environment.of(), - permanentSessionEnvironment = Environment.of( - TierCompletionSession with CompletionSession( - language = PlainTextLanguage.INSTANCE, - callOrder = 1, - completionType = CompletionType.SMART + runBlocking { + + val startOutcome = MockTask.startMLSession( + callParameters = Environment.of(), + permanentSessionEnvironment = Environment.of( + TierCompletionSession with CompletionSession( + language = PlainTextLanguage.INSTANCE, + callOrder = 1, + completionType = CompletionType.SMART + ) ) ) - ) - val completionSession = startOutcome.session ?: return@repeat - - runBlocking { + val completionSession = startOutcome.session ?: return@runBlocking completionSession.withNestedSessions { lookupSessionCreator -> lookupSessionCreator.nestConsidering(Environment.of(), Environment.of(TierLookup with LookupImpl(true, 1))) diff --git a/platform/ml-impl/test/com/intellij/platform/ml/impl/TestApiPlatform.kt b/platform/ml-impl/test/com/intellij/platform/ml/impl/TestApiPlatform.kt index f7d1030b5e2e..e29f1e56bce5 100644 --- a/platform/ml-impl/test/com/intellij/platform/ml/impl/TestApiPlatform.kt +++ b/platform/ml-impl/test/com/intellij/platform/ml/impl/TestApiPlatform.kt @@ -3,42 +3,20 @@ package com.intellij.platform.ml.impl import com.intellij.platform.ml.impl.apiPlatform.MLApiPlatform import com.intellij.platform.ml.impl.apiPlatform.MLApiPlatform.ExtensionController -import com.intellij.platform.ml.impl.logger.MLEvent -import com.intellij.platform.ml.impl.monitoring.MLApiStartupListener import com.intellij.platform.ml.impl.monitoring.MLTaskGroupListener abstract class TestApiPlatform : MLApiPlatform() { private val dynamicTaskListeners: MutableList = mutableListOf() - private val dynamicStartupListeners: MutableList = mutableListOf() - private val dynamicEvents: MutableList = mutableListOf() abstract val initialTaskListeners: List - abstract val initialStartupListeners: List - - abstract val initialEvents: List - - final override val events: List - get() = initialEvents + dynamicEvents - final override val taskListeners: List get() = initialTaskListeners + dynamicTaskListeners - final override val startupListeners: List - get() = initialStartupListeners + dynamicStartupListeners - override fun addTaskListener(taskListener: MLTaskGroupListener): ExtensionController { return extend(taskListener, dynamicTaskListeners) } - override fun addEvent(event: MLEvent): ExtensionController { - return extend(event, dynamicEvents) - } - - override fun addStartupListener(listener: MLApiStartupListener): ExtensionController { - return extend(listener, dynamicStartupListeners) - } - private fun extend(obj: T, collection: MutableCollection): ExtensionController { collection.add(obj) return ExtensionController { collection.remove(obj) } diff --git a/platform/ml-impl/testResources/ml_logs.js b/platform/ml-impl/testResources/ml_logs.js index 3b731d64cd83..6e21f5e98531 100644 --- a/platform/ml-impl/testResources/ml_logs.js +++ b/platform/ml-impl/testResources/ml_logs.js @@ -1,6 +1,58 @@ let logs = [ -{"eventId": "startup", "data": {"success": true}}, -{"eventId": "mock.failed", "data": {"normal_failure": {"reason": "com.intellij.platform.ml.impl.approach.ModelNotAcquiredOutcome"}}}, -{"eventId": "mock.finished", "data": {"session": {"ml_model": {"random_seed": -1598672019, "language_id": "TEXT", "version": "0.0"}}, "additional": {"TierGit": {"description": {"not_used": {}, "used": {"has_user": true, "n_commits": 3}}, "id": -401817549}}, "main": {"TierCompletionSession": {"description": {"not_used": {}, "used": {"language_id": "TEXT", "git_user_is_Glebanister": true, "completion_type": "SMART", "call_order": 1}}, "id": 2005562610, "analysis": {"very_good_session": true}}}, "nested": [{"additional": {}, "main": {"TierLookup": {"description": {"not_used": {}, "used": {}}, "id": 38162, "analysis": {"lookup_index": 1}}}, "nested": [{"additional": {}, "prediction": 0.8091002299577453, "main": {"TierItem": {"description": {"not_used": {}, "used": {"length": 5, "decorations": 0}}, "id": -1220935314, "analysis": {}}}}, {"additional": {}, "prediction": 0.3358094493527286, "main": {"TierItem": {"description": {"not_used": {}, "used": {"length": 5, "decorations": 0}}, "id": -782084434, "analysis": {}}}}]}, {"additional": {}, "main": {"TierLookup": {"description": {"not_used": {}, "used": {}}, "id": 38349, "analysis": {"lookup_index": 2}}}, "nested": [{"additional": {}, "prediction": 0.6959769118328721, "main": {"TierItem": {"description": {"not_used": {}, "used": {"length": 8, "decorations": 1}}, "id": 1198136891, "analysis": {}}}}, {"additional": {}, "prediction": 0.45402567637488944, "main": {"TierItem": {"description": {"not_used": {}, "used": {"length": 7, "decorations": 1}}, "id": 122089051, "analysis": {}}}}, {"additional": {}, "main": {"TierItem": {"description": {"not_used": {}, "used": {"length": 17, "decorations": 1}}, "id": 1444769257, "analysis": {}}}}]}]}}, -{"eventId": "mock.failed", "data": {"exception": {"throwable_class": "java.lang.IllegalStateException"}}} -]; \ No newline at end of file + {"eventId": "failed", "data": {"normal_failure": {"reason": "com.intellij.platform.ml.impl.approach.ModelNotAcquiredOutcome"}}}, + { + "eventId": "finished", "data": { + "session": {"ml_model": {"random_seed": -1598672019, "language_id": "TEXT", "version": "0.0"}}, + "additional": {"TierGit": {"description": {"not_used": {}, "used": {"has_user": true, "n_commits": 3}}, "id": -401817549}}, + "main": { + "TierCompletionSession": { + "description": { + "not_used": {}, + "used": {"language_id": "TEXT", "git_user_is_Glebanister": true, "completion_type": "SMART", "call_order": 1} + }, "id": -1585807756, "analysis": {"very_good_session": true} + } + }, + "nested": [{ + "additional": {}, + "main": {"TierLookup": {"description": {"not_used": {}, "used": {}}, "id": 38162, "analysis": {"lookup_index": 1}}}, + "nested": [{ + "additional": {}, + "prediction": 0.8091002299577453, + "main": { + "TierItem": { + "description": {"not_used": {}, "used": {"length": 5, "decorations": 0}}, + "id": -1220935314, + "analysis": {} + } + } + }, { + "additional": {}, + "prediction": 0.3358094493527286, + "main": {"TierItem": {"description": {"not_used": {}, "used": {"length": 5, "decorations": 0}}, "id": -782084434, "analysis": {}}} + }] + }, { + "additional": {}, + "main": {"TierLookup": {"description": {"not_used": {}, "used": {}}, "id": 38349, "analysis": {"lookup_index": 2}}}, + "nested": [{ + "additional": {}, + "prediction": 0.6959769118328721, + "main": {"TierItem": {"description": {"not_used": {}, "used": {"length": 8, "decorations": 1}}, "id": 1198136891, "analysis": {}}} + }, { + "additional": {}, + "prediction": 0.45402567637488944, + "main": {"TierItem": {"description": {"not_used": {}, "used": {"length": 7, "decorations": 1}}, "id": 122089051, "analysis": {}}} + }, { + "additional": {}, + "main": { + "TierItem": { + "description": {"not_used": {}, "used": {"length": 17, "decorations": 1}}, + "id": 1444769257, + "analysis": {} + } + } + }] + }] + } + }, + {"eventId": "failed", "data": {"exception": {"throwable_class": "java.lang.IllegalStateException"}}} +] \ No newline at end of file diff --git a/platform/platform-impl/intellij.platform.ide.impl.iml b/platform/platform-impl/intellij.platform.ide.impl.iml index 1e3bbb6df30f..ae3bb3afd347 100644 --- a/platform/platform-impl/intellij.platform.ide.impl.iml +++ b/platform/platform-impl/intellij.platform.ide.impl.iml @@ -143,5 +143,6 @@ + \ No newline at end of file diff --git a/platform/platform-impl/src/com/intellij/codeInsight/inline/completion/logs/InlineContextFeatures.kt b/platform/platform-impl/src/com/intellij/codeInsight/inline/completion/logs/InlineContextFeatures.kt index 789e1d22746a..7631e8e5448a 100644 --- a/platform/platform-impl/src/com/intellij/codeInsight/inline/completion/logs/InlineContextFeatures.kt +++ b/platform/platform-impl/src/com/intellij/codeInsight/inline/completion/logs/InlineContextFeatures.kt @@ -20,7 +20,7 @@ import org.jetbrains.annotations.ApiStatus Please make sure the feature you have added is also present in org.jetbrains.completion.full.line.mlApi.description.InlineContextFeatures """, - replaceWith = ReplaceWith("org.jetbrains.completion.full.line.mlApi.description.InlineContextFeatures") + replaceWith = ReplaceWith("org.jetbrains.completion.full.line.platform.mlApi.description.InlineContextFeatures") ) object InlineContextFeatures { @@ -105,12 +105,16 @@ object InlineContextFeatures { .forEach { (i, element) -> add(PARENT_FEATURES[i].with(element.javaClass)) } } + /** + * Make sure you have updated [com.intellij.codeInsight.inline.completion.ml.TypingFeatures], when added new features here + */ + @Deprecated("Obsolete way to compute typing features, the new ml api should do that") private fun MutableList>.addTypingFeatures() { val typingSpeedTracker = TypingSpeedTracker.getInstance() val timeSinceLastTyping = typingSpeedTracker.getTimeSinceLastTyping() if (timeSinceLastTyping != null) { add(TIME_SINCE_LAST_TYPING.with(timeSinceLastTyping)) - addAll(typingSpeedTracker.getTypingSpeedEventPairs()) + addAll(typingSpeedTracker.getTypingSpeedEventPairs().map { it.first }) } } diff --git a/platform/platform-impl/src/com/intellij/codeInsight/inline/completion/logs/TypingSpeedTracker.kt b/platform/platform-impl/src/com/intellij/codeInsight/inline/completion/logs/TypingSpeedTracker.kt index 6720b80d6401..3b145060bf38 100644 --- a/platform/platform-impl/src/com/intellij/codeInsight/inline/completion/logs/TypingSpeedTracker.kt +++ b/platform/platform-impl/src/com/intellij/codeInsight/inline/completion/logs/TypingSpeedTracker.kt @@ -7,6 +7,8 @@ import com.intellij.internal.statistic.eventLog.events.EventPair import com.intellij.openapi.components.Service import com.intellij.openapi.components.service import com.intellij.openapi.diagnostic.thisLogger +import com.intellij.platform.ml.Feature +import com.intellij.platform.ml.FeatureDeclaration import org.jetbrains.annotations.ApiStatus import org.jetbrains.annotations.TestOnly import java.awt.event.KeyAdapter @@ -41,9 +43,9 @@ class TypingSpeedTracker { if (decayDuration == Duration.ZERO) 0F else 0.5.pow(duration / decayDuration.toDouble(DurationUnit.MILLISECONDS)).toFloat() - fun getTypingSpeedEventPairs(): Collection> = DECAY_DURATIONS.mapNotNull { (decayDuration, eventField) -> + fun getTypingSpeedEventPairs(): Collection, Feature>> = DECAY_DURATIONS.mapNotNull { (decayDuration, eventFieldAndFeature) -> typingSpeeds[decayDuration]?.let { - eventField.with(it) + (eventFieldAndFeature.first with it) to (eventFieldAndFeature.second with it) } } @@ -71,9 +73,10 @@ class TypingSpeedTracker { } companion object { - private val DECAY_DURATIONS = listOf(1, 2, 5, 30).associate { it.seconds to EventFields.Float("typing_speed_${it}s") } + private val DECAY_DURATIONS = listOf(1, 2, 5, 30).associate { it.seconds to Pair(EventFields.Float("typing_speed_${it}s"), FeatureDeclaration.float("typing_speed_${it}s")) } fun getInstance(): TypingSpeedTracker = service() - fun getEventFields(): Array> = DECAY_DURATIONS.values.toTypedArray() + fun getEventFields(): Array> = DECAY_DURATIONS.values.map { it.first }.toTypedArray() + fun getFeatures(): Set> = DECAY_DURATIONS.values.map { it.second }.toSet() } } diff --git a/platform/platform-impl/src/com/intellij/codeInsight/inline/completion/ml/TypingFeatures.kt b/platform/platform-impl/src/com/intellij/codeInsight/inline/completion/ml/TypingFeatures.kt new file mode 100644 index 000000000000..0a9a5369409a --- /dev/null +++ b/platform/platform-impl/src/com/intellij/codeInsight/inline/completion/ml/TypingFeatures.kt @@ -0,0 +1,23 @@ +// 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.intellij.platform.ml.* + +internal class TypingFeatures : TierDescriptor.Default(TierTyping) { + companion object { + val TIME_SINCE_LAST_TYPING = FeatureDeclaration.long("time_since_last_typing") + } + + override val descriptionDeclaration: Set> = TypingSpeedTracker.getFeatures() + TIME_SINCE_LAST_TYPING + + override suspend fun describe(environment: Environment, usefulFeaturesFilter: FeatureFilter): Set = buildSet { + with(environment[TierTyping]) { + val timeSinceLastTyping = getTimeSinceLastTyping() + if (timeSinceLastTyping != null) { + add(TIME_SINCE_LAST_TYPING.with(timeSinceLastTyping)) + addAll(getTypingSpeedEventPairs().map { it.second }) + } + } + } +} diff --git a/platform/platform-impl/src/com/intellij/codeInsight/inline/completion/ml/TypingSpeedProvider.kt b/platform/platform-impl/src/com/intellij/codeInsight/inline/completion/ml/TypingSpeedProvider.kt new file mode 100644 index 000000000000..6d03baab7c29 --- /dev/null +++ b/platform/platform-impl/src/com/intellij/codeInsight/inline/completion/ml/TypingSpeedProvider.kt @@ -0,0 +1,17 @@ +// 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.intellij.platform.ml.Environment +import com.intellij.platform.ml.EnvironmentExtender +import com.intellij.platform.ml.Tier + +internal class TypingSpeedProvider : EnvironmentExtender { + override val extendingTier: Tier = TierTyping + + override fun extend(environment: Environment): TypingSpeedTracker { + return TypingSpeedTracker.getInstance() + } + + override val requiredTiers: Set> = emptySet() +} diff --git a/platform/platform-impl/src/com/intellij/codeInsight/inline/completion/ml/tiers.kt b/platform/platform-impl/src/com/intellij/codeInsight/inline/completion/ml/tiers.kt new file mode 100644 index 000000000000..f04308064a7b --- /dev/null +++ b/platform/platform-impl/src/com/intellij/codeInsight/inline/completion/ml/tiers.kt @@ -0,0 +1,23 @@ +// 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.intellij.lang.Language +import com.intellij.openapi.editor.Editor +import com.intellij.platform.ml.Tier +import com.intellij.psi.PsiFile +import org.jetbrains.annotations.ApiStatus + +@ApiStatus.Internal +object TierTyping : Tier() + +@ApiStatus.Internal +object TierCaretLocation : Tier() + +@ApiStatus.Internal +data class CaretLocation( + val psiFile: PsiFile, + val editor: Editor, + val offset: Int, + val language: Language +) diff --git a/platform/platform-resources/src/META-INF/PlatformExtensions.xml b/platform/platform-resources/src/META-INF/PlatformExtensions.xml index f9730c823735..e25c750e957c 100644 --- a/platform/platform-resources/src/META-INF/PlatformExtensions.xml +++ b/platform/platform-resources/src/META-INF/PlatformExtensions.xml @@ -1627,6 +1627,10 @@ + + + +