mirror of
https://gitflic.ru/project/openide/openide.git
synced 2026-09-27 10:03:11 +07:00
[ML-2843] ML API: Migrate more FLCC features
Conceal platform api Add "Typing" tier & add more features Remove empty file Add debug information to environment extension Automate description's declaration maximally Refactor events logging Add more features from FLCC Add debug information to ML API Add utility function for shorter features declarations Merge-request: IJ-MR-127545 Merged-by: Gleb Marin <Gleb.Marin@jetbrains.com> GitOrigin-RevId: 8ada4b7e75302f7e0a8d222030e443939487f493
This commit is contained in:
committed by
intellij-monorepo-bot
parent
543850b0bf
commit
0f048de5e1
@@ -87,4 +87,10 @@ sealed class Feature {
|
||||
: TypedFeature<com.intellij.openapi.util.Version>(name, value) {
|
||||
override val valueType = FeatureValueType.Version
|
||||
}
|
||||
|
||||
companion object {
|
||||
fun Feature.toCompactString(): String {
|
||||
return "${this.declaration.name}=${this.value}"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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<TierDescriptor> = ExtensionPointName.create("com.intellij.platform.ml.descriptor")
|
||||
}
|
||||
|
||||
abstract class Default(final override val tier: Tier<*>) : TierDescriptor {
|
||||
override val additionallyRequiredTiers: Set<Tier<*>>
|
||||
get() = emptySet()
|
||||
|
||||
override val descriptionPolicy: DescriptionPolicy
|
||||
get() = DescriptionPolicy(false, false)
|
||||
|
||||
override fun couldBeUseful(usefulFeaturesFilter: FeatureFilter): Boolean {
|
||||
return descriptionDeclaration.any { usefulFeaturesFilter.accept(it) }
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@ApiStatus.Internal
|
||||
|
||||
@@ -47,5 +47,6 @@
|
||||
<orderEntry type="library" name="kotlinx-serialization-core" level="project" />
|
||||
<orderEntry type="module" module-name="intellij.platform.ml" />
|
||||
<orderEntry type="module" module-name="intellij.platform.statistics" />
|
||||
<orderEntry type="library" name="kotlin-reflect" level="project" />
|
||||
</component>
|
||||
</module>
|
||||
@@ -6,9 +6,6 @@
|
||||
|
||||
<projectService serviceInterface="com.intellij.internal.statistic.local.LanguageUsageStatisticsProvider"
|
||||
serviceImplementation="com.intellij.internal.statistic.local.LanguageUsageStatisticsProviderImpl"/>
|
||||
<statistics.counterUsagesCollector implementationClass="com.intellij.platform.ml.impl.logger.MLEventsLogger"/>
|
||||
<registryKey defaultValue="false" description="Log features that were computed but missing in an ObsoleteTierDescriptor's declaration"
|
||||
key="ml.description.logMissing"/>
|
||||
</extensions>
|
||||
|
||||
<projectListeners>
|
||||
@@ -23,7 +20,7 @@
|
||||
interface="com.intellij.platform.ml.impl.turboComplete.SmartPipelineRunner"
|
||||
dynamic="true"/>
|
||||
<extensionPoint qualifiedName="com.intellij.platform.ml.impl.approach"
|
||||
interface="com.intellij.platform.ml.impl.MLTaskApproachInitializer"
|
||||
interface="com.intellij.platform.ml.impl.MLTaskApproachBuilder"
|
||||
dynamic="true"/>
|
||||
|
||||
<extensionPoint name="mlCompletionCorrectnessSupporter"
|
||||
|
||||
@@ -16,7 +16,7 @@ import org.jetbrains.annotations.ApiStatus
|
||||
*
|
||||
* Two types of instances have the power to select features:
|
||||
* - [com.intellij.platform.ml.impl.model.MLModel] to tell which features are known.
|
||||
* - [com.intellij.platform.ml.impl.approach.LogDrivenModelInference] to tell which features are not known by the ML model,
|
||||
* - [com.intellij.platform.ml.impl.LogDrivenModelInference] to tell which features are not known by the ML model,
|
||||
* but they are still must be computed and then logged.
|
||||
*/
|
||||
@ApiStatus.Internal
|
||||
|
||||
@@ -0,0 +1,258 @@
|
||||
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
|
||||
package com.intellij.platform.ml.impl
|
||||
|
||||
import com.intellij.openapi.diagnostic.thisLogger
|
||||
import com.intellij.platform.ml.*
|
||||
import com.intellij.platform.ml.ScopeEnvironment.Companion.narrowedTo
|
||||
import com.intellij.platform.ml.impl.LogDrivenModelInference.SessionDetails.Companion.getLevels
|
||||
import com.intellij.platform.ml.impl.LogDrivenModelInference.SessionDetails.Companion.validate
|
||||
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 com.intellij.platform.ml.impl.session.analysis.MLModelAnalyser
|
||||
import com.intellij.platform.ml.impl.session.analysis.MLSessionAnalyser
|
||||
import com.intellij.platform.ml.impl.session.analysis.StructureAnalyser
|
||||
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.
|
||||
* @param sessionDetails Preferences on how the session will be executed
|
||||
* @param thisBuilder The builder that has created this class.
|
||||
* It is essential to put it, otherwise [com.intellij.platform.ml.impl.monitoring.MLTaskGroupListener] will not manage to find the
|
||||
* correct task listener.
|
||||
*/
|
||||
@ApiStatus.Internal
|
||||
open class LogDrivenModelInference<M : MLModel<P>, P : Any>(
|
||||
final override val task: MLTask<P>,
|
||||
final override val apiPlatform: MLApiPlatform,
|
||||
private val sessionDetails: SessionDetails<M, P>,
|
||||
private val thisBuilder: MLTaskApproachBuilder<P>,
|
||||
) : MLTaskApproach<P> {
|
||||
private val sessionLevels = sessionDetails.getLevels(task)
|
||||
|
||||
interface SessionDetails<M : MLModel<P>, P : Any> {
|
||||
fun interface Builder<M : MLModel<P>, P : Any> {
|
||||
fun build(apiPlatform: MLApiPlatform): SessionDetails<M, P>
|
||||
}
|
||||
|
||||
/**
|
||||
* 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<StructureAnalyser<M, P>>
|
||||
|
||||
/**
|
||||
* The analyzers, that are giving features to the model, that was used during the session.
|
||||
*/
|
||||
val mlModelAnalysers: Collection<MLModelAnalyser<M, P>>
|
||||
|
||||
/**
|
||||
* Provides an ML model to use during session's lifetime.
|
||||
*/
|
||||
val mlModelProvider: MLModel.Provider<M, P>
|
||||
|
||||
/**
|
||||
* 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<Set<Tier<*>>>
|
||||
|
||||
/**
|
||||
* 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<FeatureFilter>
|
||||
|
||||
companion object {
|
||||
fun <M : MLModel<P>, P : Any> SessionDetails<M, P>.getLevels(task: MLTask<P>): List<LevelSignature<Set<Tier<*>>, Set<Tier<*>>>> {
|
||||
return (task.levels zip additionallyDescribedTiers).map { LevelSignature(it.first, it.second) }
|
||||
}
|
||||
|
||||
fun <M : MLModel<P>, P : Any> SessionDetails<M, P>.validate(task: MLTask<P>) {
|
||||
require(task.levels.size == additionallyDescribedTiers.size) {
|
||||
"Task $task has ${task.levels.size} levels, when 'additionallyDescribedTiers' has ${additionallyDescribedTiers.size}"
|
||||
}
|
||||
|
||||
val maybeDuplicatedTaskTiers: List<Tier<*>> = getLevels(task).flatMap { it.main.toList() + it.additional }
|
||||
val taskTiers: Set<Tier<*>> = 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<M : MLModel<P>, P : Any>(
|
||||
override val task: MLTask<P>,
|
||||
private val details: SessionDetails.Builder<M, P>,
|
||||
) : MLTaskApproachBuilder<P> {
|
||||
final override fun buildApproach(apiPlatform: MLApiPlatform): MLTaskApproach<P> {
|
||||
val sessionDetails = details.build(apiPlatform)
|
||||
sessionDetails.validate(task)
|
||||
val modelInference = buildModelInference(apiPlatform, sessionDetails)
|
||||
return modelInference
|
||||
}
|
||||
|
||||
open fun buildModelInference(apiPlatform: MLApiPlatform, sessionDetails: SessionDetails<M, P>): LogDrivenModelInference<M, P> {
|
||||
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<TierDescriptor>): Set<FeatureDeclaration<*>> {
|
||||
return tierDescriptors.flatMap {
|
||||
if (it is ObsoleteTierDescriptor) it.partialDescriptionDeclaration else it.descriptionDeclaration
|
||||
}.toSet()
|
||||
}
|
||||
|
||||
private fun buildMainTiersScheme(sessionAnalyser: MLSessionAnalyser<M, P>, tiers: Set<Tier<*>>, apiEnvironment: MLApiPlatform): PerTier<MainTierScheme> {
|
||||
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<Tier<*>>, apiEnvironment: MLApiPlatform): PerTier<AdditionalTierScheme> {
|
||||
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<P> {
|
||||
return startSessionMonitoring(callParameters, permanentSessionEnvironment)
|
||||
}
|
||||
|
||||
private suspend fun startSessionMonitoring(callParameters: Environment, permanentSessionEnvironment: Environment): Session.StartOutcome<P> {
|
||||
val approachListener = apiPlatform.getJoinedListenerForTask<M, P>(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<M, P>): Session.StartOutcome<P> {
|
||||
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<P>()
|
||||
approachListener.onFailedToStartSession(failure)
|
||||
return failure
|
||||
}
|
||||
nullableMlModel
|
||||
}
|
||||
|
||||
var sessionListener: MLSessionListener<M, P>? = null
|
||||
|
||||
val sessionAnalyser = MLSessionAnalyser(sessionDetails.structureAnalysers, sessionDetails.mlModelAnalysers)
|
||||
|
||||
val analyseThenLogStructure = SessionTreeHandler<DescribedRootContainer<M, P>, 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<FeatureFilter>) {
|
||||
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<P : Any> : Session.StartOutcome.Failure<P>() {
|
||||
override val failureDetails: String = "ML Model was not provided"
|
||||
}
|
||||
@@ -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<T : Any> protected constructor(
|
||||
val predictionClass: Class<T>
|
||||
) {
|
||||
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<T : Any> 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<P : Any> {
|
||||
@@ -65,15 +70,10 @@ interface MLTaskApproach<P : Any> {
|
||||
val task: MLTask<P>
|
||||
|
||||
/**
|
||||
* 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<P : Any> {
|
||||
*/
|
||||
suspend fun startSession(callParameters: Environment, permanentSessionEnvironment: Environment): Session.StartOutcome<P>
|
||||
|
||||
data class Declaration(
|
||||
data class SessionDeclaration(
|
||||
val sessionFeatures: Map<String, Set<FeatureDeclaration<*>>>,
|
||||
val levelsScheme: List<LevelScheme>
|
||||
)
|
||||
|
||||
companion object {
|
||||
fun <P : Any> findMlApproach(task: MLTask<P>, apiPlatform: MLApiPlatform = ReplaceableIJPlatform): MLTaskApproach<P> {
|
||||
return apiPlatform.accessApproachFor(task)
|
||||
fun <P : Any> findMlTaskApproach(task: MLTask<P>, apiPlatform: MLApiPlatform): MLTaskApproachBuilder<P> {
|
||||
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<P>
|
||||
}
|
||||
|
||||
suspend fun <P : Any> startMLSession(task: MLTask<P>,
|
||||
apiPlatform: MLApiPlatform,
|
||||
callParameters: Environment,
|
||||
permanentSessionEnvironment: Environment): Session.StartOutcome<P> {
|
||||
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 <P : Any> MLTask<P>.startMLSession(callParameters: Environment, permanentSessionEnvironment: Environment, apiPlatform: MLApiPlatform = ReplaceableIJPlatform): Session.StartOutcome<P> {
|
||||
return startMLSession(this@startMLSession, apiPlatform, callParameters, permanentSessionEnvironment)
|
||||
}
|
||||
|
||||
fun <P : Any> MLTask<P>.startCoroutineAndMLSession(callParameters: Environment, permanentSessionEnvironment: Environment, apiPlatform: MLApiPlatform = ReplaceableIJPlatform): Session.StartOutcome<P> {
|
||||
return runBlockingCancellable {
|
||||
startMLSession(callParameters, permanentSessionEnvironment, apiPlatform)
|
||||
private fun validateEnvironment(callParameters: Environment,
|
||||
expectedCallParameters: Set<Tier<*>>,
|
||||
permanentSessionEnvironment: Environment,
|
||||
expectedPermanentTiers: Set<Tier<*>>) {
|
||||
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<P : Any> {
|
||||
* Initializes an [MLTaskApproach]
|
||||
*/
|
||||
@ApiStatus.Internal
|
||||
interface MLTaskApproachInitializer<P : Any> {
|
||||
interface MLTaskApproachBuilder<P : Any> {
|
||||
/**
|
||||
* The task, that the created [MLTaskApproach] is dedicated to solve.
|
||||
*/
|
||||
@@ -129,22 +152,25 @@ interface MLTaskApproachInitializer<P : Any> {
|
||||
|
||||
/**
|
||||
* 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<P>
|
||||
fun buildApproach(apiPlatform: MLApiPlatform): MLTaskApproach<P>
|
||||
|
||||
/**
|
||||
* 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<MLTaskApproachInitializer<*>>("com.intellij.platform.ml.impl.approach")
|
||||
val EP_NAME = ExtensionPointName<MLTaskApproachBuilder<*>>("com.intellij.platform.ml.impl.approach")
|
||||
}
|
||||
}
|
||||
|
||||
@ApiStatus.Internal
|
||||
data class LevelSignature<M, A>(
|
||||
val main: M,
|
||||
val additional: A
|
||||
|
||||
@@ -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<T> {
|
||||
}
|
||||
}
|
||||
|
||||
fun interface MLTaskListenerProvider : MessagingProvider<MLTaskGroupListener> {
|
||||
fun interface MLTaskGroupListenerProvider : MessagingProvider<MLTaskGroupListener> {
|
||||
companion object {
|
||||
val TOPIC = MessagingProvider.createTopic<MLTaskGroupListener, MLTaskListenerProvider>("ml.task")
|
||||
}
|
||||
}
|
||||
|
||||
fun interface MLEventProvider : MessagingProvider<MLEvent> {
|
||||
companion object {
|
||||
val TOPIC = MessagingProvider.createTopic<MLEvent, MLEventProvider>("ml.event")
|
||||
}
|
||||
}
|
||||
|
||||
fun interface MLApiStartupListenerProvider : MessagingProvider<MLApiStartupListener> {
|
||||
companion object {
|
||||
val TOPIC = MessagingProvider.createTopic<MLApiStartupListener, MLApiStartupListenerProvider>("ml.startup")
|
||||
}
|
||||
|
||||
interface ThisProvider : MLApiStartupListenerProvider, MLApiStartupListener {
|
||||
override fun provide(collector: (MLApiStartupListener) -> Unit) = collector(this)
|
||||
val TOPIC = MessagingProvider.createTopic<MLTaskGroupListener, MLTaskGroupListenerProvider>("ml.task")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -73,41 +54,22 @@ private data object IJPlatform : MLApiPlatform() {
|
||||
override val environmentExtenders: List<EnvironmentExtender<*>>
|
||||
get() = EnvironmentExtender.EP_NAME.extensionList
|
||||
|
||||
override val taskApproaches: List<MLTaskApproachInitializer<*>>
|
||||
get() = MLTaskApproachInitializer.EP_NAME.extensionList
|
||||
override val taskApproaches: List<MLTaskApproachBuilder<*>>
|
||||
get() = MLTaskApproachBuilder.EP_NAME.extensionList
|
||||
|
||||
override val taskListeners: List<MLTaskGroupListener>
|
||||
get() = MessagingProvider.collect(MLTaskListenerProvider.TOPIC)
|
||||
|
||||
override val events: List<MLEvent>
|
||||
get() = MessagingProvider.collect(MLEventProvider.TOPIC)
|
||||
|
||||
override val startupListeners: List<MLApiStartupListener>
|
||||
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<Feature>) {
|
||||
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<EnvironmentExtender<*>>
|
||||
get() = platform.environmentExtenders
|
||||
|
||||
override val taskApproaches: List<MLTaskApproachInitializer<*>>
|
||||
override val taskApproaches: List<MLTaskApproachBuilder<*>>
|
||||
get() = platform.taskApproaches
|
||||
|
||||
|
||||
override val taskListeners: List<MLTaskGroupListener>
|
||||
get() = platform.taskListeners
|
||||
|
||||
override val events: List<MLEvent>
|
||||
get() = platform.events
|
||||
|
||||
override val startupListeners: List<MLApiStartupListener>
|
||||
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<Feature>) =
|
||||
platform.manageNonDeclaredFeatures(descriptor, nonDeclaredFeatures)
|
||||
|
||||
|
||||
@@ -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 <P : Any> accessApproachFor(task: MLTask<P>): MLTaskApproach<P> {
|
||||
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<TierDescriptor>
|
||||
|
||||
/**
|
||||
* 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<EnvironmentExtender<*>>
|
||||
|
||||
/**
|
||||
* 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<MLTaskApproachInitializer<*>>
|
||||
abstract val taskApproaches: List<MLTaskApproachBuilder<*>>
|
||||
|
||||
|
||||
/**
|
||||
@@ -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<MLEvent>
|
||||
|
||||
/**
|
||||
* 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<MLApiStartupListener>
|
||||
|
||||
/**
|
||||
* 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<Feature>)
|
||||
|
||||
|
||||
internal abstract fun addEvent(event: MLEvent): ExtensionController
|
||||
|
||||
fun interface ExtensionController {
|
||||
fun removeExtension()
|
||||
}
|
||||
|
||||
data class StaticState(
|
||||
val tierDescriptors: List<TierDescriptor>,
|
||||
val environmentExtenders: List<EnvironmentExtender<*>>,
|
||||
val taskApproaches: List<MLTaskApproachInitializer<*>>,
|
||||
)
|
||||
|
||||
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<InitializerAndApproach<*>>) : PotentiallySuccessful(2, {
|
||||
it.onStartedInitializingFus(initializedApproaches)
|
||||
}, allowStartup = false)
|
||||
|
||||
data object Finished : PotentiallySuccessful(3, { it.onFinished() }, allowStartup = false)
|
||||
}
|
||||
|
||||
private inner class MLApiPlatformInitializationProcess {
|
||||
val approachPerTask: Map<MLTask<*>, MLTaskApproach<*>>
|
||||
private val completeInitializersList: List<MLTaskApproachInitializer<*>> = taskApproaches.toMutableList()
|
||||
|
||||
init {
|
||||
require(initializationStage.allowStartup) {
|
||||
"""
|
||||
ML API Platform's initialization should not be run twice. Current state is $initializationStage
|
||||
""".trimIndent()
|
||||
}
|
||||
|
||||
fun currentStartupListeners(): List<MLApiStartupProcessListener> = startupListeners.map { it.onBeforeStarted(this@MLApiPlatform) }
|
||||
|
||||
fun <T> 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<InitializerAndApproach<*>>()
|
||||
|
||||
approachPerTask = proceedToNextStage(InitializationStage.InitializingFUS(initializedApproachPerTask)) {
|
||||
completeInitializersList.validate()
|
||||
|
||||
fun <T : Any> initializeApproach(approachInitializer: MLTaskApproachInitializer<T>) {
|
||||
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 <P : Any> getApproachFor(task: MLTask<P>): MLTaskApproach<P> {
|
||||
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<P>
|
||||
}
|
||||
|
||||
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<MLTaskApproachInitializer<*>>.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<Tier<*>>): PerTier<List<TierDescriptor>> {
|
||||
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 <R, P : Any> MLApiPlatform.getJoinedListenerForTask(taskApproach: MLTaskApproach<P>,
|
||||
fun <R, P : Any> MLApiPlatform.getJoinedListenerForTask(taskApproachBuilder: MLTaskApproachBuilder<P>,
|
||||
permanentSessionEnvironment: Environment): MLApproachListener<R, P> {
|
||||
val relevantGroupListeners = taskListeners.filter { taskApproach.javaClass in it.targetedApproaches }
|
||||
val relevantGroupListeners = taskListeners.filter { taskApproachBuilder.javaClass in it.targetedApproaches }
|
||||
val approachListeners = relevantGroupListeners.mapNotNull {
|
||||
it.onAttemptedToStartSession<P, R>(taskApproach, permanentSessionEnvironment)
|
||||
it.onAttemptedToStartSession<P, R>(taskApproachBuilder, permanentSessionEnvironment)
|
||||
}
|
||||
return approachListeners.asJoinedListener()
|
||||
}
|
||||
|
||||
@@ -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<M : MLModel<P>, P : Any> {
|
||||
/**
|
||||
* Static declaration of the features, that are used in the session tree's analysis.
|
||||
*/
|
||||
val structureAnalysisDeclaration: PerTier<Set<FeatureDeclaration<*>>>
|
||||
|
||||
/**
|
||||
* Static declaration of the session's entities, that are not tiers.
|
||||
*/
|
||||
val sessionAnalysisDeclaration: Map<String, Set<FeatureDeclaration<*>>>
|
||||
|
||||
/**
|
||||
* Perform the completed session's analysis.
|
||||
*/
|
||||
fun analyseTree(treeRoot: DescribedRootContainer<M, P>): CompletableFuture<AnalysedRootContainer<P>>
|
||||
}
|
||||
-223
@@ -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<M : MLModel<P>, P : Any>(
|
||||
override val task: MLTask<P>,
|
||||
override val apiPlatform: MLApiPlatform
|
||||
) : MLTaskApproach<P> {
|
||||
/**
|
||||
* 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<M, P>
|
||||
|
||||
/**
|
||||
* Provides an ML model to use during session's lifetime.
|
||||
*/
|
||||
abstract val mlModelProvider: MLModel.Provider<M, P>
|
||||
|
||||
/**
|
||||
* 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<FeatureFilter>
|
||||
|
||||
/**
|
||||
* 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<Set<Tier<*>>>
|
||||
|
||||
private val levels: List<LevelTiers> 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<P> {
|
||||
return startSessionMonitoring(callParameters, permanentSessionEnvironment)
|
||||
}
|
||||
|
||||
private suspend fun startSessionMonitoring(callParameters: Environment, permanentSessionEnvironment: Environment): Session.StartOutcome<P> {
|
||||
val approachListener = apiPlatform.getJoinedListenerForTask<M, P>(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<M, P>): Session.StartOutcome<P> {
|
||||
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<P>()
|
||||
approachListener.onFailedToStartSession(failure)
|
||||
return failure
|
||||
}
|
||||
nullableMlModel
|
||||
}
|
||||
|
||||
var sessionListener: MLSessionListener<M, P>? = null
|
||||
|
||||
val analyseThenLogStructure = SessionTreeHandler<DescribedRootContainer<M, P>, 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<FeatureFilter>) {
|
||||
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<TierDescriptor>): Set<FeatureDeclaration<*>> {
|
||||
return tierDescriptors.flatMap {
|
||||
if (it is ObsoleteTierDescriptor) it.partialDescriptionDeclaration else it.descriptionDeclaration
|
||||
}.toSet()
|
||||
}
|
||||
|
||||
private fun buildMainTiersScheme(tiers: Set<Tier<*>>, apiEnvironment: MLApiPlatform): PerTier<MainTierScheme> {
|
||||
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<Tier<*>>, apiEnvironment: MLApiPlatform): PerTier<AdditionalTierScheme> {
|
||||
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<P : Any> : Session.StartOutcome.Failure<P>() {
|
||||
override val failureDetails: String = "ML Model was not provided"
|
||||
}
|
||||
+49
-17
@@ -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<Tier<*>>) {
|
||||
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<Tier<*>>,
|
||||
extenders: List<EnvironmentExtender<*>>): Environment {
|
||||
extenders: List<EnvironmentExtender<*>>,
|
||||
mainTiers: Set<Tier<*>>): 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<ExtensionOutcome> = 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<ExtendedEnvironment>().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<Tier<*>>, extenders: List<EnvironmentExtender<*>>): Map<Tier<*>, EnvironmentExtender<*>> {
|
||||
val extendableTiers: Set<Tier<*>> = extenders.map { it.extendingTier }.toSet()
|
||||
|
||||
@@ -130,16 +162,16 @@ class ExtendedEnvironment : Environment {
|
||||
}
|
||||
}
|
||||
|
||||
private fun Environment.separateIntoExtenders(): List<EnvironmentExtender<*>> {
|
||||
class ContainingExtender<T : Any>(private val tier: Tier<T>) : EnvironmentExtender<T> {
|
||||
override val extendingTier: Tier<T> = tier
|
||||
internal class ContainingExtender<T : Any>(private val containingEnvironment: Environment, private val tier: Tier<T>) : EnvironmentExtender<T> {
|
||||
override val extendingTier: Tier<T> = tier
|
||||
|
||||
override fun extend(environment: Environment): T {
|
||||
return this@separateIntoExtenders[tier]
|
||||
}
|
||||
|
||||
override val requiredTiers: Set<Tier<*>> = emptySet()
|
||||
override fun extend(environment: Environment): T {
|
||||
return containingEnvironment[tier]
|
||||
}
|
||||
|
||||
return this.tiers.map { tier -> ContainingExtender(tier) }
|
||||
override val requiredTiers: Set<Tier<*>> = emptySet()
|
||||
}
|
||||
|
||||
private fun Environment.separateIntoExtenders(): List<EnvironmentExtender<*>> {
|
||||
return this.tiers.map { tier -> ContainingExtender(this, tier) }
|
||||
}
|
||||
|
||||
@@ -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<EventField<*>> = 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<EventPair<*>>(SUCCESS with false)
|
||||
exception?.let { fields += EXCEPTION with it.javaClass }
|
||||
eventId.log(*fields.toTypedArray())
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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<EventField<*>>
|
||||
|
||||
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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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<M, P : Any>(
|
||||
private val taskApproach: MLTaskApproach<P>,
|
||||
private val exceptionalAnalysers: Collection<ShallowSessionAnalyser<Throwable>>,
|
||||
private val normalFailureAnalysers: Collection<ShallowSessionAnalyser<Session.StartOutcome.Failure<P>>>,
|
||||
private val apiPlatform: MLApiPlatform,
|
||||
) : EventIdRecordingMLEvent(), MLTaskGroupListener {
|
||||
override val eventName: String = "${taskApproach.task.name}.failed"
|
||||
|
||||
private val fields: Map<String, ObjectEventField> = (
|
||||
exceptionalAnalysers.map { ObjectEventField(it.name, it.declarationObjectDescription) } +
|
||||
normalFailureAnalysers.map { ObjectEventField(it.name, it.declarationObjectDescription) })
|
||||
.associateBy { it.name }
|
||||
|
||||
override val declaration: Array<EventField<*>> = fields.values.toTypedArray()
|
||||
|
||||
override val approachListeners: Collection<MLTaskGroupListener.ApproachListeners<*, *>>
|
||||
get() {
|
||||
val eventId = getEventId(apiPlatform)
|
||||
return listOf(
|
||||
taskApproach.javaClass monitoredBy MLApproachInitializationListener { permanentSessionEnvironment ->
|
||||
object : MLApproachListener<M, P> {
|
||||
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<P>) {
|
||||
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<P>) = null
|
||||
}
|
||||
}
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Registers failed sessions' logging as a separate FUS event.
|
||||
*
|
||||
* See [MLSessionFailedLogger]
|
||||
*/
|
||||
@ApiStatus.Internal
|
||||
open class FailedSessionLoggerRegister<M, P : Any>(
|
||||
private val targetApproachClass: Class<out MLTaskApproach<P>>,
|
||||
private val exceptionalAnalysers: Collection<ShallowSessionAnalyser<Throwable>>,
|
||||
private val normalFailureAnalysers: Collection<ShallowSessionAnalyser<Session.StartOutcome.Failure<P>>>,
|
||||
) : MLApiStartupListener {
|
||||
override fun onBeforeStarted(apiPlatform: MLApiPlatform): MLApiStartupProcessListener {
|
||||
return object : MLApiStartupProcessListener {
|
||||
override fun onStartedInitializingFus(initializedApproaches: Collection<InitializerAndApproach<*>>) {
|
||||
@Suppress("UNCHECKED_CAST")
|
||||
val targetInitializedApproach: MLTaskApproach<P> = 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<P>
|
||||
val finishedEventLogger = MLSessionFailedLogger<M, P>(targetInitializedApproach, exceptionalAnalysers, normalFailureAnalysers, apiPlatform)
|
||||
apiPlatform.addMLEventBeforeFusInitialized(finishedEventLogger)
|
||||
apiPlatform.addTaskListener(finishedEventLogger)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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<R, P : Any>(
|
||||
approach: MLTaskApproach<P>,
|
||||
configuration: FusSessionEventBuilder.FusScheme<P>,
|
||||
private val apiPlatform: MLApiPlatform,
|
||||
) : EventIdRecordingMLEvent(), MLTaskGroupListener {
|
||||
private val loggingScheme: FusSessionEventBuilder<P> = configuration.createEventBuilder(approach.approachDeclaration)
|
||||
private val fusDeclaration: SessionFields<P> = loggingScheme.buildFusDeclaration()
|
||||
override val declaration: Array<EventField<*>> = fusDeclaration.getFields()
|
||||
|
||||
override val eventName: String = "${approach.task.name}.finished"
|
||||
|
||||
override val approachListeners: Collection<MLTaskGroupListener.ApproachListeners<*, *>> = listOf(
|
||||
approach.javaClass monitoredBy InitializationLogger()
|
||||
)
|
||||
|
||||
inner class InitializationLogger : MLApproachInitializationListener<R, P> {
|
||||
override fun onAttemptedToStartSession(permanentSessionEnvironment: Environment): MLApproachListener<R, P> = ApproachLogger()
|
||||
}
|
||||
|
||||
inner class ApproachLogger : MLApproachListener<R, P> {
|
||||
override fun onFailedToStartSessionWithException(exception: Throwable) {}
|
||||
|
||||
override fun onFailedToStartSession(failure: Session.StartOutcome.Failure<P>) {}
|
||||
|
||||
override fun onStartedSession(session: Session<P>): MLSessionListener<R, P> = SessionLogger()
|
||||
}
|
||||
|
||||
inner class SessionLogger : MLSessionListener<R, P> {
|
||||
override fun onSessionDescriptionFinished(sessionTree: DescribedRootContainer<R, P>) {}
|
||||
|
||||
override fun onSessionAnalysisFinished(sessionTree: AnalysedRootContainer<P>) {
|
||||
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<M, P : Any>(
|
||||
private val targetApproachClass: Class<out MLTaskApproach<P>>,
|
||||
private val fusScheme: FusSessionEventBuilder.FusScheme<P>,
|
||||
) : MLApiStartupListener, MLApiStartupListenerProvider.ThisProvider {
|
||||
override fun onBeforeStarted(apiPlatform: MLApiPlatform): MLApiStartupProcessListener {
|
||||
return object : MLApiStartupProcessListener {
|
||||
override fun onStartedInitializingFus(initializedApproaches: Collection<InitializerAndApproach<*>>) {
|
||||
@Suppress("UNCHECKED_CAST")
|
||||
val targetInitializedApproach: MLTaskApproach<P> = 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<P>
|
||||
val finishedEventLogger = MLSessionFinishedLogger<M, P>(targetInitializedApproach, fusScheme, apiPlatform)
|
||||
apiPlatform.addMLEventBeforeFusInitialized(finishedEventLogger)
|
||||
apiPlatform.addTaskListener(finishedEventLogger)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
+10
-11
@@ -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<P : Any> : ObjectDescription() {
|
||||
@@ -25,25 +25,24 @@ abstract class SessionFields<P : Any> : ObjectDescription() {
|
||||
* @param P The type of the ML task's prediction
|
||||
*/
|
||||
@ApiStatus.Internal
|
||||
interface FusSessionEventBuilder<P : Any> {
|
||||
interface EventSessionEventBuilder<P : Any> {
|
||||
/**
|
||||
* 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<P : Any> {
|
||||
fun createEventBuilder(approachDeclaration: MLTaskApproach.Declaration): FusSessionEventBuilder<P>
|
||||
interface EventScheme<P : Any> {
|
||||
fun createEventBuilder(approachDeclaration: MLTaskApproach.SessionDeclaration): EventSessionEventBuilder<P>
|
||||
}
|
||||
|
||||
/**
|
||||
* 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<P>
|
||||
fun buildSessionFields(): SessionFields<P>
|
||||
|
||||
/**
|
||||
* 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<P>, sessionFields: SessionFields<P>): Array<EventPair<*>>
|
||||
}
|
||||
+7
-7
@@ -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<P : Any, F> internal constructor(
|
||||
private val predictionField: EventField<F>,
|
||||
private val predictionTransformer: (P?) -> F?,
|
||||
private val approachDeclaration: MLTaskApproach.Declaration
|
||||
) : FusSessionEventBuilder<P> {
|
||||
private val approachDeclaration: MLTaskApproach.SessionDeclaration
|
||||
) : EventSessionEventBuilder<P> {
|
||||
class FusScheme<P : Any, F>(
|
||||
private val predictionField: EventField<F>,
|
||||
private val predictionTransformer: (P?) -> F?,
|
||||
) : FusSessionEventBuilder.FusScheme<P> {
|
||||
override fun createEventBuilder(approachDeclaration: MLTaskApproach.Declaration): FusSessionEventBuilder<P> = InplaceFeaturesScheme(
|
||||
) : EventSessionEventBuilder.EventScheme<P> {
|
||||
override fun createEventBuilder(approachDeclaration: MLTaskApproach.SessionDeclaration): EventSessionEventBuilder<P> = InplaceFeaturesScheme(
|
||||
predictionField,
|
||||
predictionTransformer,
|
||||
approachDeclaration
|
||||
@@ -37,7 +37,7 @@ class InplaceFeaturesScheme<P : Any, F> internal constructor(
|
||||
}
|
||||
}
|
||||
|
||||
override fun buildFusDeclaration(): SessionFields<P> {
|
||||
override fun buildSessionFields(): SessionFields<P> {
|
||||
require(approachDeclaration.levelsScheme.isNotEmpty())
|
||||
return if (approachDeclaration.levelsScheme.size == 1)
|
||||
PredictionSessionFields(approachDeclaration.levelsScheme.first(), predictionField, predictionTransformer,
|
||||
@@ -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 <R, P : Any> EventLogGroup.registerEventSessionFailed(
|
||||
eventName: String,
|
||||
task: MLTask<P>,
|
||||
exceptionalAnalysers: Collection<ShallowSessionAnalyser<Throwable>>,
|
||||
normalFailureAnalysers: Collection<ShallowSessionAnalyser<Session.StartOutcome.Failure<P>>>,
|
||||
apiPlatform: MLApiPlatform = ReplaceableIJPlatform
|
||||
): MLApiPlatform.ExtensionController {
|
||||
val approachBuilder = findMlTaskApproach(task, apiPlatform)
|
||||
val fields: Map<String, ObjectEventField> = (normalFailureAnalysers + exceptionalAnalysers)
|
||||
.map { ObjectEventField(it.name, it.declarationObjectDescription) }
|
||||
.associateBy { it.name }
|
||||
val eventId = registerVarargEvent(eventName, *fields.values.toTypedArray())
|
||||
val logger = MLSessionFailedLogger<R, P>(approachBuilder, exceptionalAnalysers, normalFailureAnalysers, eventId, fields)
|
||||
return apiPlatform.addTaskListener(logger)
|
||||
}
|
||||
|
||||
private class MLSessionFailedLogger<M, P : Any>(
|
||||
private val taskApproachBuilder: MLTaskApproachBuilder<P>,
|
||||
private val exceptionalAnalysers: Collection<ShallowSessionAnalyser<Throwable>>,
|
||||
private val normalFailureAnalysers: Collection<ShallowSessionAnalyser<Session.StartOutcome.Failure<P>>>,
|
||||
private val eventId: VarargEventId,
|
||||
private val fields: Map<String, ObjectEventField>,
|
||||
) : MLTaskGroupListener {
|
||||
|
||||
override val approachListeners: Collection<MLTaskGroupListener.ApproachListeners<*, *>>
|
||||
get() {
|
||||
return listOf(
|
||||
taskApproachBuilder.javaClass monitoredBy MLApproachInitializationListener { permanentSessionEnvironment ->
|
||||
object : MLApproachListener<M, P> {
|
||||
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<P>) {
|
||||
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<P>) = null
|
||||
}
|
||||
}
|
||||
)
|
||||
}
|
||||
}
|
||||
@@ -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 <R, P : Any> EventLogGroup.registerEventSessionFinished(
|
||||
eventName: String,
|
||||
task: MLTask<P>,
|
||||
scheme: EventSessionEventBuilder.EventScheme<P>,
|
||||
apiPlatform: MLApiPlatform = ReplaceableIJPlatform
|
||||
): MLApiPlatform.ExtensionController {
|
||||
val approachBuilder = findMlTaskApproach(task, apiPlatform)
|
||||
val sessionDeclaration = approachBuilder.buildApproachSessionDeclaration(apiPlatform)
|
||||
val finishedSessionEventBuilder: EventSessionEventBuilder<P> = scheme.createEventBuilder(sessionDeclaration)
|
||||
val sessionFields: SessionFields<P> = finishedSessionEventBuilder.buildSessionFields()
|
||||
val eventId = registerVarargEvent(eventName, *sessionFields.getFields())
|
||||
val logger = MLSessionFinishedLogger<R, P>(approachBuilder, eventId, finishedSessionEventBuilder, sessionFields)
|
||||
return apiPlatform.addTaskListener(logger)
|
||||
}
|
||||
|
||||
private class MLSessionFinishedLogger<R, P : Any>(
|
||||
approachBuilder: MLTaskApproachBuilder<P>,
|
||||
private val eventId: VarargEventId,
|
||||
private val finishedSessionEventBuilder: EventSessionEventBuilder<P>,
|
||||
private val sessionFields: SessionFields<P>
|
||||
) : MLTaskGroupListener {
|
||||
override val approachListeners: Collection<MLTaskGroupListener.ApproachListeners<*, *>> = listOf(
|
||||
approachBuilder.javaClass monitoredBy InitializationLogger()
|
||||
)
|
||||
|
||||
inner class InitializationLogger : MLApproachInitializationListener<R, P> {
|
||||
override fun onAttemptedToStartSession(permanentSessionEnvironment: Environment): MLApproachListener<R, P> = ApproachLogger()
|
||||
}
|
||||
|
||||
inner class ApproachLogger : MLApproachListener<R, P> {
|
||||
override fun onFailedToStartSessionWithException(exception: Throwable) {}
|
||||
|
||||
override fun onFailedToStartSession(failure: Session.StartOutcome.Failure<P>) {}
|
||||
|
||||
override fun onStartedSession(session: Session<P>): MLSessionListener<R, P> = SessionLogger()
|
||||
}
|
||||
|
||||
inner class SessionLogger : MLSessionListener<R, P> {
|
||||
override fun onSessionDescriptionFinished(sessionTree: DescribedRootContainer<R, P>) {}
|
||||
|
||||
override fun onSessionAnalysisFinished(sessionTree: AnalysedRootContainer<P>) {
|
||||
val fusEventData = finishedSessionEventBuilder.buildRecord(sessionTree, sessionFields)
|
||||
eventId.log(*fusEventData)
|
||||
}
|
||||
}
|
||||
}
|
||||
-31
@@ -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<T : Any>(
|
||||
val initializer: MLTaskApproachInitializer<T>,
|
||||
val approach: MLTaskApproach<T>
|
||||
)
|
||||
|
||||
@ApiStatus.Internal
|
||||
interface MLApiStartupProcessListener {
|
||||
fun onStartedInitializingApproaches() {}
|
||||
|
||||
fun onStartedInitializingFus(initializedApproaches: Collection<InitializerAndApproach<*>>) {}
|
||||
|
||||
fun onFinished() {}
|
||||
|
||||
fun onFailed(exception: Throwable?) {}
|
||||
|
||||
fun onCanceled() {}
|
||||
}
|
||||
+31
-19
@@ -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<R, P : Any> internal constructor(
|
||||
val taskApproach: Class<out MLTaskApproach<P>>,
|
||||
val taskApproachBuilder: Class<out MLTaskApproachBuilder<P>>,
|
||||
val approachListener: Collection<MLApproachInitializationListener<R, P>>
|
||||
) {
|
||||
companion object {
|
||||
infix fun <R, P : Any> Class<out MLTaskApproach<P>>.monitoredBy(approachListener: MLApproachInitializationListener<R, P>) = ApproachListeners(
|
||||
infix fun <R, P : Any> Class<out MLTaskApproachBuilder<P>>.monitoredBy(approachListener: MLApproachInitializationListener<R, P>) = ApproachListeners(
|
||||
this, listOf(approachListener))
|
||||
|
||||
infix fun <R, P : Any> Class<out MLTaskApproach<P>>.monitoredBy(approachListeners: Collection<MLApproachInitializationListener<R, P>>) = ApproachListeners(
|
||||
infix fun <R, P : Any> Class<out MLTaskApproachBuilder<P>>.monitoredBy(approachListeners: Collection<MLApproachInitializationListener<R, P>>) = ApproachListeners(
|
||||
this, approachListeners)
|
||||
}
|
||||
}
|
||||
|
||||
companion object {
|
||||
internal val MLTaskGroupListener.targetedApproaches: Set<Class<out MLTaskApproach<*>>>
|
||||
get() = approachListeners.map { it.taskApproach }.toSet()
|
||||
@ApiStatus.Internal
|
||||
interface Default : MLTaskGroupListener, MLTaskGroupListenerProvider {
|
||||
override fun provide(collector: (MLTaskGroupListener) -> Unit) {
|
||||
return collector(this)
|
||||
}
|
||||
}
|
||||
|
||||
internal fun <P : Any, R> MLTaskGroupListener.onAttemptedToStartSession(taskApproach: MLTaskApproach<P>,
|
||||
companion object {
|
||||
internal val MLTaskGroupListener.targetedApproaches: Set<Class<out MLTaskApproachBuilder<*>>>
|
||||
get() = approachListeners.map { it.taskApproachBuilder }.toSet()
|
||||
|
||||
internal fun <P : Any, R> MLTaskGroupListener.onAttemptedToStartSession(taskApproachBuilder: MLTaskApproachBuilder<P>,
|
||||
permanentSessionEnvironment: Environment): MLApproachListener<R, P>? {
|
||||
@Suppress("UNCHECKED_CAST")
|
||||
val approachListeners: List<MLApproachInitializationListener<R, P>> = approachListeners
|
||||
.filter { it.taskApproach == taskApproach.javaClass }
|
||||
.filter { it.taskApproachBuilder == taskApproachBuilder.javaClass }
|
||||
.flatMap { it.approachListener } as List<MLApproachInitializationListener<R, P>>
|
||||
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<R, P : Any> {
|
||||
/**
|
||||
* 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<R, P : Any> {
|
||||
}
|
||||
|
||||
/**
|
||||
* 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<R, P : Any> {
|
||||
/**
|
||||
* 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<P>)
|
||||
fun onFailedToStartSession(failure: Session.StartOutcome.Failure<P>) {}
|
||||
|
||||
/**
|
||||
* 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<R, P : Any> {
|
||||
}
|
||||
|
||||
/**
|
||||
* 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<R, P : Any> {
|
||||
@@ -139,14 +151,14 @@ interface MLSessionListener<R, P : Any> {
|
||||
* All tier instances were established (the tree will not be growing further),
|
||||
* described, and predictions in the [sessionTree] were finished.
|
||||
*/
|
||||
fun onSessionDescriptionFinished(sessionTree: DescribedRootContainer<R, P>)
|
||||
fun onSessionDescriptionFinished(sessionTree: DescribedRootContainer<R, P>) {}
|
||||
|
||||
/**
|
||||
* Called only after [onSessionDescriptionFinished]
|
||||
*
|
||||
* All tree nodes were analyzed.
|
||||
*/
|
||||
fun onSessionAnalysisFinished(sessionTree: AnalysedRootContainer<P>)
|
||||
fun onSessionAnalysisFinished(sessionTree: AnalysedRootContainer<P>) {}
|
||||
|
||||
companion object {
|
||||
fun <R, P : Any> Collection<MLSessionListener<R, P>>.asJoinedListener(): MLSessionListener<R, P> {
|
||||
|
||||
@@ -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<FeatureSelector>,
|
||||
val notUsedFeaturesSelectors: PerTier<FeatureFilter>,
|
||||
val notUsedFeaturesFilters: PerTier<FeatureFilter>,
|
||||
val mlTask: MLTask<*>
|
||||
) {
|
||||
suspend fun describe(
|
||||
nextLevelCallParameters: Environment,
|
||||
@@ -25,6 +30,15 @@ data class LevelDescriptor(
|
||||
upperLevels: List<DescribedLevel>,
|
||||
nextLevelAdditionalTiers: Set<Tier<*>>
|
||||
): 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<Tier<*>, Declaredness<Usage<Set<Feature>>>> = 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<Feature>): Set<Feature> {
|
||||
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<TierDescriptor>, 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<TierDescriptor, Set<Feature>> = 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 ->
|
||||
|
||||
+62
-22
@@ -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<M : MLModel<P>, P : Any> private constructor(
|
||||
levelsTiers: List<LevelTiers>,
|
||||
usedFeaturesSelectors: PerTier<FeatureSelector>,
|
||||
override val levelDescriptor: LevelDescriptor,
|
||||
override val levelPositioning: LevelPositioning
|
||||
) : NestableStructureCollector<SessionTree.ComplexRoot<M, DescribedTierData, P>, M, P>(usedFeaturesSelectors) {
|
||||
override val levelPositioning: LevelPositioning,
|
||||
mlTask: MLTask<P>
|
||||
) : NestableStructureCollector<SessionTree.ComplexRoot<M, DescribedTierData, P>, M, P>(usedFeaturesSelectors, mlTask) {
|
||||
|
||||
init {
|
||||
validateSuperiorCollector(initializationCallParameters, callParameters.first(), levelsTiers, levelMainEnvironment, levelAdditionalTiers, mlModel, notUsedFeaturesSelectors)
|
||||
@@ -36,7 +41,7 @@ class RootCollector<M : MLModel<P>, P : Any> private constructor(
|
||||
companion object {
|
||||
suspend fun <M : MLModel<P>, P : Any> build(
|
||||
initializationCallParameters: Environment,
|
||||
callParameters: List<Set<Tier<*>>>,
|
||||
mlTask: MLTask<P>,
|
||||
descriptionComputer: DescriptionComputer,
|
||||
notUsedFeaturesSelectors: PerTier<FeatureFilter>,
|
||||
levelMainEnvironment: Environment,
|
||||
@@ -44,11 +49,11 @@ class RootCollector<M : MLModel<P>, P : Any> private constructor(
|
||||
mlModel: M,
|
||||
apiPlatform: MLApiPlatform,
|
||||
levelsTiers: List<LevelTiers>,
|
||||
usedFeaturesSelectors: PerTier<FeatureSelector>
|
||||
usedFeaturesSelectors: PerTier<FeatureSelector>,
|
||||
): RootCollector<M, P> {
|
||||
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<M : MLModel<P>, P : Any> private constructor(
|
||||
levelScheme: LevelTiers,
|
||||
usedFeaturesSelectors: PerTier<FeatureSelector>,
|
||||
override val levelDescriptor: LevelDescriptor,
|
||||
override val levelPositioning: LevelPositioning
|
||||
) : PredictionCollector<SessionTree.SolitaryLeaf<M, DescribedTierData, P>, M, P>(usedFeaturesSelectors) {
|
||||
override val levelPositioning: LevelPositioning,
|
||||
mlTask: MLTask<P>,
|
||||
) : PredictionCollector<SessionTree.SolitaryLeaf<M, DescribedTierData, P>, M, P>(usedFeaturesSelectors, mlTask) {
|
||||
|
||||
init {
|
||||
validateSuperiorCollector(initializationCallParameters, callParameters.first(), listOf(levelScheme), levelMainEnvironment, levelAdditionalTiers, mlModel, notUsedFeaturesSelectors)
|
||||
@@ -78,7 +84,7 @@ class SolitaryLeafCollector<M : MLModel<P>, P : Any> private constructor(
|
||||
companion object {
|
||||
suspend fun <M : MLModel<P>, P : Any> build(
|
||||
initializationCallParameters: Environment,
|
||||
callParameters: List<Set<Tier<*>>>,
|
||||
mlTask: MLTask<P>,
|
||||
descriptionComputer: DescriptionComputer,
|
||||
notUsedFeaturesSelectors: PerTier<FeatureFilter>,
|
||||
levelMainEnvironment: Environment,
|
||||
@@ -86,12 +92,12 @@ class SolitaryLeafCollector<M : MLModel<P>, P : Any> private constructor(
|
||||
mlModel: M,
|
||||
apiPlatform: MLApiPlatform,
|
||||
levelScheme: LevelTiers,
|
||||
usedFeaturesSelectors: PerTier<FeatureSelector>
|
||||
usedFeaturesSelectors: PerTier<FeatureSelector>,
|
||||
): SolitaryLeafCollector<M, P> {
|
||||
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<M : MLModel<P>, P : Any> private constructor(
|
||||
class BranchingCollector<M : MLModel<P>, P : Any>(
|
||||
override val levelDescriptor: LevelDescriptor,
|
||||
override val levelPositioning: LevelPositioning,
|
||||
usedFeaturesSelectors: PerTier<FeatureSelector>
|
||||
) : NestableStructureCollector<SessionTree.Branching<M, DescribedTierData, P>, M, P>(usedFeaturesSelectors) {
|
||||
usedFeaturesSelectors: PerTier<FeatureSelector>,
|
||||
mlTask: MLTask<P>,
|
||||
) : NestableStructureCollector<SessionTree.Branching<M, DescribedTierData, P>, M, P>(usedFeaturesSelectors, mlTask) {
|
||||
override fun createTree(thisLevel: DescribedLevel,
|
||||
collectedNestedStructureTrees: List<DescribedSessionTree<M, P>>): SessionTree.Branching<M, DescribedTierData, P> {
|
||||
return SessionTree.Branching(thisLevel, collectedNestedStructureTrees)
|
||||
@@ -110,7 +117,8 @@ class BranchingCollector<M : MLModel<P>, P : Any>(
|
||||
|
||||
@ApiStatus.Internal
|
||||
abstract class PredictionCollector<T : SessionTree.PredictionContainer<M, DescribedTierData, P>, M : MLModel<P>, P : Any>(
|
||||
private val usedFeaturesSelectors: PerTier<FeatureSelector>
|
||||
private val usedFeaturesSelectors: PerTier<FeatureSelector>,
|
||||
private val mlTask: MLTask<P>,
|
||||
) : StructureCollector<T, M, P>() {
|
||||
private var predictionSubmitted = false
|
||||
private var submittedPrediction: P? = null
|
||||
@@ -127,9 +135,37 @@ abstract class PredictionCollector<T : SessionTree.PredictionContainer<M, Descri
|
||||
require(!predictionSubmitted)
|
||||
submittedPrediction = prediction
|
||||
predictionSubmitted = true
|
||||
thisLogger().debug {
|
||||
"[${mlTask.name}] Prediction collected for task ${mlTask.name}\n" +
|
||||
"\tPrediction value: $prediction\n" +
|
||||
"\tDescription:\n" +
|
||||
levelPositioning.levels.withIndex().flatMap { it.value.prettyPrinted(it.index) }.withIndent().joinToString("\n")
|
||||
}
|
||||
submitTreeToHandlers(createTree(levelPositioning.thisLevel, submittedPrediction))
|
||||
}
|
||||
|
||||
private fun DescribedLevel.prettyPrinted(index: Int): List<String> {
|
||||
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<String> {
|
||||
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<String>.withIndent(n: Int = 1): List<String> {
|
||||
return map { "\t".repeat(n) + it }
|
||||
}
|
||||
|
||||
private fun PerTierInstance<DescribedTierData>.extractDescriptionForModel(): PerTier<Set<Feature>> {
|
||||
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<T : SessionTree.PredictionContainer<M, Descri
|
||||
class LeafCollector<M : MLModel<P>, P : Any>(
|
||||
override val levelDescriptor: LevelDescriptor,
|
||||
override val levelPositioning: LevelPositioning,
|
||||
usedFeaturesSelectors: PerTier<FeatureSelector>
|
||||
) : PredictionCollector<SessionTree.Leaf<M, DescribedTierData, P>, M, P>(usedFeaturesSelectors) {
|
||||
usedFeaturesSelectors: PerTier<FeatureSelector>,
|
||||
mlTask: MLTask<P>,
|
||||
) : PredictionCollector<SessionTree.Leaf<M, DescribedTierData, P>, M, P>(usedFeaturesSelectors, mlTask) {
|
||||
|
||||
override fun createTree(thisLevel: DescribedLevel, prediction: P?): SessionTree.Leaf<M, DescribedTierData, P> {
|
||||
return SessionTree.Leaf(levelPositioning.thisLevel, prediction)
|
||||
@@ -255,7 +292,8 @@ sealed class StructureCollector<T : DescribedSessionTree<M, P>, M : MLModel<P>,
|
||||
|
||||
@ApiStatus.Internal
|
||||
abstract class NestableStructureCollector<T : DescribedChildrenContainer<M, P>, M : MLModel<P>, P : Any>(
|
||||
private val usedFeaturesSelectors: PerTier<FeatureSelector>
|
||||
private val usedFeaturesSelectors: PerTier<FeatureSelector>,
|
||||
private val mlTask: MLTask<P>
|
||||
) : StructureCollector<T, M, P>() {
|
||||
private val nestedSessionsStructures: MutableList<CompletableFuture<DescribedSessionTree<M, P>>> = mutableListOf()
|
||||
private var nestingFinished = false
|
||||
@@ -264,7 +302,8 @@ abstract class NestableStructureCollector<T : DescribedChildrenContainer<M, P>,
|
||||
verifyNestedLevelEnvironment(levelMainEnvironment, levelAdditionalTiers)
|
||||
return BranchingCollector<M, P>(levelDescriptor,
|
||||
levelPositioning.nestNextLevel(callParameters, levelMainEnvironment, levelAdditionalTiers, levelDescriptor),
|
||||
usedFeaturesSelectors)
|
||||
usedFeaturesSelectors,
|
||||
mlTask)
|
||||
.also { it.trackCollectedStructure() }
|
||||
}
|
||||
|
||||
@@ -273,7 +312,8 @@ abstract class NestableStructureCollector<T : DescribedChildrenContainer<M, P>,
|
||||
return LeafCollector<M, P>(
|
||||
levelDescriptor,
|
||||
levelPositioning.nestNextLevel(callParameters, levelMainEnvironment, levelAdditionalTiers, levelDescriptor),
|
||||
usedFeaturesSelectors
|
||||
usedFeaturesSelectors,
|
||||
mlTask
|
||||
).also { it.trackCollectedStructure() }
|
||||
}
|
||||
|
||||
|
||||
+4
-5
@@ -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<FeatureDeclaration<*>>
|
||||
)
|
||||
|
||||
@@ -36,7 +35,7 @@ class JoinedGroupedSessionAnalyser<M : MLModel<P>, P : Any>(
|
||||
private val mlModelAnalysers: Collection<MLModelAnalyser<M, P>>,
|
||||
) : SessionAnalyser<GroupedAnalysisDeclaration, GroupedAnalysis<M, P>, M, P> {
|
||||
override val analysisDeclaration = GroupedAnalysisDeclaration(
|
||||
structureAnalysis = SessionStructureAnalysisJoiner<M, P>().joinDeclarations(structureAnalysers.map { it.analysisDeclaration }),
|
||||
sessionStructureAnalysis = SessionStructureAnalysisJoiner<M, P>().joinDeclarations(structureAnalysers.map { it.analysisDeclaration }),
|
||||
mlModelAnalysis = MLModelAnalysisJoiner<M, P>().joinDeclarations(mlModelAnalysers.map { it.analysisDeclaration })
|
||||
)
|
||||
|
||||
+2
-3
@@ -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
|
||||
|
||||
+18
-15
@@ -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<M : MLModel<P>, P : Any>(
|
||||
class MLSessionAnalyser<M : MLModel<P>, P : Any>(
|
||||
structureAnalysers: Collection<StructureAnalyser<M, P>>,
|
||||
mlModelAnalysers: Collection<MLModelAnalyser<M, P>>,
|
||||
private val sessionAnalysisKeyModel: String = DEFAULT_SESSION_KEY_ML_MODEL
|
||||
) : AnalysisMethod<M, P> {
|
||||
) {
|
||||
private val groupedAnalyser = JoinedGroupedSessionAnalyser(structureAnalysers, mlModelAnalysers)
|
||||
|
||||
override val structureAnalysisDeclaration: StructureAnalysisDeclaration
|
||||
get() = groupedAnalyser.analysisDeclaration.structureAnalysis
|
||||
|
||||
override val sessionAnalysisDeclaration: Map<String, Set<FeatureDeclaration<*>>> = mapOf(
|
||||
sessionAnalysisKeyModel to groupedAnalyser.analysisDeclaration.mlModelAnalysis
|
||||
)
|
||||
|
||||
override fun analyseTree(treeRoot: DescribedRootContainer<M, P>): CompletableFuture<AnalysedRootContainer<P>> {
|
||||
fun analyseSessionStructure(treeRoot: DescribedRootContainer<M, P>): CompletableFuture<AnalysedRootContainer<P>> {
|
||||
return groupedAnalyser.analyse(treeRoot).thenApply {
|
||||
buildAnalysedSessionTree(treeRoot, it) as AnalysedRootContainer<P>
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 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<String, Set<FeatureDeclaration<*>>> = mapOf(
|
||||
sessionAnalysisKeyModel to groupedAnalyser.analysisDeclaration.mlModelAnalysis
|
||||
)
|
||||
|
||||
private fun buildAnalysedSessionTree(tree: DescribedSessionTree<M, P>, analysis: GroupedAnalysis<M, P>): AnalysedSessionTree<P> {
|
||||
val treeAnalysisPerInstance: PerTierInstance<AnalysedTierData> = tree.levelData.mainInstances.entries
|
||||
.associate { (tierInstance, data) ->
|
||||
+2
-3
@@ -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
|
||||
|
||||
@@ -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 <reified T : Any> extractFieldsAsFeatureDeclarations(declarationContainer: T,
|
||||
extractors: Collection<FeatureDeclarationsExtractor> = listOf(FeatureDeclarationsExtractor.OfProperty,
|
||||
FeatureDeclarationsExtractor.ofIterable,
|
||||
FeatureDeclarationsExtractor.ofMapValues)
|
||||
): Set<FeatureDeclaration<*>> {
|
||||
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<FeatureDeclaration<*>>?
|
||||
|
||||
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
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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<LookupItem>()
|
||||
|
||||
object TierGit : Tier<GitRepository>()
|
||||
|
||||
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<CompletionType>("completion_type").nullable()
|
||||
}
|
||||
|
||||
override val tier: Tier<*> = TierCompletionSession
|
||||
|
||||
override val descriptionPolicy: DescriptionPolicy = DescriptionPolicy(false, false)
|
||||
|
||||
override val additionallyRequiredTiers: Set<Tier<*>> = 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<FeatureDeclaration<*>> = 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<FeatureDeclaration<*>> = setOf(
|
||||
@@ -218,11 +215,11 @@ class RandomModelSeedAnalyser : MLModelAnalyser<RandomModel, Double> {
|
||||
class RandomModel(val seed: Int) : MLModel<Double>, Versioned, LanguageSpecific {
|
||||
private val generator = Random(seed)
|
||||
|
||||
class Provider : MLModel.Provider<RandomModel, Double> {
|
||||
object Provider : MLModel.Provider<RandomModel, Double> {
|
||||
private val generator = Random(1)
|
||||
override fun provideModel(callParameters: Environment, environment: Environment, sessionTiers: List<LevelTiers>): 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<Double>, 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<RandomModel, Double> {
|
||||
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<Double>) {
|
||||
@@ -309,6 +306,17 @@ object FailureLogger : ShallowSessionAnalyser<Session.StartOutcome.Failure<Doubl
|
||||
}
|
||||
}
|
||||
|
||||
class MockTaskFusLogger : CounterUsagesCollector() {
|
||||
companion object {
|
||||
val GROUP = EventLogGroup("mock-task", 1).also {
|
||||
it.registerEventSessionFailed<RandomModel, Double>("failed", MockTask, listOf(ExceptionLogger), listOf(FailureLogger))
|
||||
it.registerEventSessionFinished<RandomModel, Double>("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<MLApiStartupListener> = listOf(
|
||||
FinishedSessionLoggerRegister<RandomModel, Double>(
|
||||
MockTaskApproach::class.java,
|
||||
InplaceFeaturesScheme.FusScheme.DOUBLE
|
||||
),
|
||||
FailedSessionLoggerRegister<RandomModel, Double>(
|
||||
MockTaskApproach::class.java,
|
||||
exceptionalAnalysers = listOf(ExceptionLogger),
|
||||
normalFailureAnalysers = listOf(FailureLogger)
|
||||
)
|
||||
MockTaskApproachBuilder()
|
||||
)
|
||||
|
||||
override val initialTaskListeners: List<MLTaskGroupListener> = listOf(
|
||||
@@ -342,8 +337,6 @@ object ThisTestApiPlatform : TestApiPlatform() {
|
||||
SomeListener("Alex"),
|
||||
)
|
||||
|
||||
override val initialEvents: List<MLEvent> = listOf()
|
||||
|
||||
|
||||
override fun manageNonDeclaredFeatures(descriptor: ObsoleteTierDescriptor, nonDeclaredFeatures: Set<Feature>) {
|
||||
val printer = CodeLikePrinter()
|
||||
@@ -366,27 +359,25 @@ object MockTask : MLTask<Double>(
|
||||
)
|
||||
)
|
||||
|
||||
class MockTaskApproach(
|
||||
apiPlatform: MLApiPlatform,
|
||||
task: MLTask<Double>
|
||||
) : LogDrivenModelInference<RandomModel, Double>(task, apiPlatform) {
|
||||
|
||||
class MockTaskApproachDetails : LogDrivenModelInference.SessionDetails<RandomModel, Double> {
|
||||
override val additionallyDescribedTiers: List<Set<Tier<*>>> = listOf(
|
||||
setOf(TierGit),
|
||||
setOf(),
|
||||
setOf(),
|
||||
)
|
||||
|
||||
override val analysisMethod: AnalysisMethod<RandomModel, Double> = StructureAndModelAnalysis(
|
||||
structureAnalysers = listOf(SomeStructureAnalyser()),
|
||||
mlModelAnalysers = listOf(
|
||||
override val mlModelAnalysers: Collection<MLModelAnalyser<RandomModel, Double>>
|
||||
get() = listOf(
|
||||
RandomModelSeedAnalyser(),
|
||||
ModelVersionAnalyser(),
|
||||
ModelLanguageAnalyser()
|
||||
)
|
||||
)
|
||||
|
||||
override val mlModelProvider = RandomModel.Provider()
|
||||
override val structureAnalysers: Collection<StructureAnalyser<RandomModel, Double>>
|
||||
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<Double> {
|
||||
override val task: MLTask<Double> = MockTask
|
||||
override fun initializeApproachWithin(apiPlatform: MLApiPlatform) = MockTaskApproach(apiPlatform, task)
|
||||
class Builder : LogDrivenModelInference.SessionDetails.Builder<RandomModel, Double> {
|
||||
override fun build(apiPlatform: MLApiPlatform): LogDrivenModelInference.SessionDetails<RandomModel, Double> {
|
||||
return MockTaskApproachDetails()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
class MockTaskApproachBuilder : LogDrivenModelInference.Builder<RandomModel, Double>(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)))
|
||||
|
||||
@@ -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<MLTaskGroupListener> = mutableListOf()
|
||||
private val dynamicStartupListeners: MutableList<MLApiStartupListener> = mutableListOf()
|
||||
private val dynamicEvents: MutableList<MLEvent> = mutableListOf()
|
||||
|
||||
abstract val initialTaskListeners: List<MLTaskGroupListener>
|
||||
|
||||
abstract val initialStartupListeners: List<MLApiStartupListener>
|
||||
|
||||
abstract val initialEvents: List<MLEvent>
|
||||
|
||||
final override val events: List<MLEvent>
|
||||
get() = initialEvents + dynamicEvents
|
||||
|
||||
final override val taskListeners: List<MLTaskGroupListener>
|
||||
get() = initialTaskListeners + dynamicTaskListeners
|
||||
|
||||
final override val startupListeners: List<MLApiStartupListener>
|
||||
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 <T> extend(obj: T, collection: MutableCollection<T>): ExtensionController {
|
||||
collection.add(obj)
|
||||
return ExtensionController { collection.remove(obj) }
|
||||
|
||||
@@ -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"}}}
|
||||
];
|
||||
{"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"}}}
|
||||
]
|
||||
@@ -143,5 +143,6 @@
|
||||
<orderEntry type="library" name="jetbrains.intellij.deps.rwmutex.idea" level="project" />
|
||||
<orderEntry type="module" module-name="intellij.platform.bootstrap.coroutine" />
|
||||
<orderEntry type="library" name="lz4-java" level="project" />
|
||||
<orderEntry type="module" module-name="intellij.platform.ml" />
|
||||
</component>
|
||||
</module>
|
||||
+6
-2
@@ -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<EventPair<*>>.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 })
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+7
-4
@@ -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<EventPair<*>> = DECAY_DURATIONS.mapNotNull { (decayDuration, eventField) ->
|
||||
fun getTypingSpeedEventPairs(): Collection<Pair<EventPair<*>, 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<EventField<*>> = DECAY_DURATIONS.values.toTypedArray()
|
||||
fun getEventFields(): Array<EventField<*>> = DECAY_DURATIONS.values.map { it.first }.toTypedArray()
|
||||
fun getFeatures(): Set<FeatureDeclaration<*>> = DECAY_DURATIONS.values.map { it.second }.toSet()
|
||||
}
|
||||
}
|
||||
|
||||
+23
@@ -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<FeatureDeclaration<*>> = TypingSpeedTracker.getFeatures() + TIME_SINCE_LAST_TYPING
|
||||
|
||||
override suspend fun describe(environment: Environment, usefulFeaturesFilter: FeatureFilter): Set<Feature> = buildSet {
|
||||
with(environment[TierTyping]) {
|
||||
val timeSinceLastTyping = getTimeSinceLastTyping()
|
||||
if (timeSinceLastTyping != null) {
|
||||
add(TIME_SINCE_LAST_TYPING.with(timeSinceLastTyping))
|
||||
addAll(getTypingSpeedEventPairs().map { it.second })
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
+17
@@ -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<TypingSpeedTracker> {
|
||||
override val extendingTier: Tier<TypingSpeedTracker> = TierTyping
|
||||
|
||||
override fun extend(environment: Environment): TypingSpeedTracker {
|
||||
return TypingSpeedTracker.getInstance()
|
||||
}
|
||||
|
||||
override val requiredTiers: Set<Tier<*>> = emptySet()
|
||||
}
|
||||
@@ -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<TypingSpeedTracker>()
|
||||
|
||||
@ApiStatus.Internal
|
||||
object TierCaretLocation : Tier<CaretLocation>()
|
||||
|
||||
@ApiStatus.Internal
|
||||
data class CaretLocation(
|
||||
val psiFile: PsiFile,
|
||||
val editor: Editor,
|
||||
val offset: Int,
|
||||
val language: Language
|
||||
)
|
||||
@@ -1627,6 +1627,10 @@
|
||||
<projectService serviceInterface="com.intellij.internal.performanceTests.ProjectInitializationDiagnosticService"
|
||||
serviceImplementation="com.intellij.internal.performanceTests.DummyProjectInitializationDiagnosticService"/>
|
||||
|
||||
<!-- ML API -->
|
||||
<platform.ml.descriptor implementation="com.intellij.codeInsight.inline.completion.ml.TypingFeatures"/>
|
||||
<platform.ml.environmentExtender implementation="com.intellij.codeInsight.inline.completion.ml.TypingSpeedProvider"/>
|
||||
|
||||
<!-- Presentation Assistant -->
|
||||
<applicationService serviceImplementation="com.intellij.platform.ide.impl.presentationAssistant.PresentationAssistant"/>
|
||||
<applicationConfigurable groupId="appearance"
|
||||
|
||||
Reference in New Issue
Block a user