[stats-collector] Extract interface ContextFeatures

GitOrigin-RevId: 5c9b94bed1172d416a9a7f3be9d3c097fb2dbe7e
This commit is contained in:
Vitaliy.Bibaev
2019-08-27 13:41:59 +00:00
committed by intellij-monorepo-bot
parent 4f5498a006
commit 404233f94a
5 changed files with 55 additions and 37 deletions
@@ -0,0 +1,16 @@
// Copyright 2000-2019 JetBrains s.r.o. Use of this source code is governed by the Apache 2.0 license that can be found in the LICENSE file.
package com.intellij.codeInsight.completion.ml;
import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
public interface ContextFeatures {
@Nullable
Boolean binaryValue(@NotNull String name);
@Nullable
Double floatValue(@NotNull String name);
@Nullable
String categoricalValue(@NotNull String name);
}
@@ -1,32 +0,0 @@
// Copyright 2000-2019 JetBrains s.r.o. Use of this source code is governed by the Apache 2.0 license that can be found in the LICENSE file.
package com.intellij.codeInsight.completion.ml
import com.intellij.openapi.util.Key
import com.intellij.psi.PsiElement
import java.util.*
class ContextFeatures(private val featuresSnapshot: Map<String, MLFeatureValue>) {
companion object {
private val CONTEXT_FEATURES_KEY = Key.create<ContextFeatures>("com.intellij.completion.ml.context_features")
private val EMPTY = ContextFeatures(emptyMap())
fun extract(element: PsiElement): ContextFeatures {
return element.getUserData(CONTEXT_FEATURES_KEY) ?: EMPTY
}
fun setContextFeatures(element: PsiElement, features: Map<String, MLFeatureValue>) {
element.putUserData(CONTEXT_FEATURES_KEY, ContextFeatures(Collections.unmodifiableMap(features)))
}
fun clear(element: PsiElement) {
element.putUserData(CONTEXT_FEATURES_KEY, null)
}
}
fun binaryValue(name: String): Boolean? = (featuresSnapshot[name])?.asBinary()
fun floatValue(name: String): Double? = (featuresSnapshot[name])?.asFloat()
fun categoricalValue(name: String): String? = (featuresSnapshot[name])?.asCategorical()
}
@@ -0,0 +1,35 @@
// Copyright 2000-2019 JetBrains s.r.o. Use of this source code is governed by the Apache 2.0 license that can be found in the LICENSE file.
package com.intellij.completion.ml
import com.intellij.codeInsight.completion.ml.ContextFeatures
import com.intellij.codeInsight.completion.ml.MLFeatureValue
import com.intellij.openapi.util.Key
import com.intellij.psi.PsiElement
import java.util.*
class ContextFeaturesStorage(private val featuresSnapshot: Map<String, MLFeatureValue>): ContextFeatures {
companion object {
private val CONTEXT_FEATURES_KEY = Key.create<ContextFeaturesStorage>("com.intellij.completion.ml.context_features")
private val EMPTY = ContextFeaturesStorage(emptyMap())
fun extract(element: PsiElement): ContextFeaturesStorage {
return element.getUserData(CONTEXT_FEATURES_KEY) ?: EMPTY
}
fun setContextFeatures(element: PsiElement, features: Map<String, MLFeatureValue>) {
element.putUserData(CONTEXT_FEATURES_KEY,
ContextFeaturesStorage(Collections.unmodifiableMap(features)))
}
fun clear(element: PsiElement) {
element.putUserData(CONTEXT_FEATURES_KEY, null)
}
}
override fun binaryValue(name: String): Boolean? = (featuresSnapshot[name])?.asBinary()
override fun floatValue(name: String): Double? = (featuresSnapshot[name])?.asFloat()
override fun categoricalValue(name: String): String? = (featuresSnapshot[name])?.asCategorical()
}
@@ -3,14 +3,13 @@ package com.intellij.completion.ml
import com.intellij.codeInsight.completion.CompletionLocation
import com.intellij.codeInsight.completion.CompletionWeigher
import com.intellij.codeInsight.completion.ml.ContextFeatures
import com.intellij.codeInsight.completion.ml.ElementFeatureProvider
import com.intellij.codeInsight.lookup.LookupElement
class MLCompletionWeigher : CompletionWeigher() {
override fun weigh(element: LookupElement, location: CompletionLocation): Comparable<Nothing>? {
val psiFile = location.completionParameters.originalFile
val contextFeatures = ContextFeatures.extract(psiFile)
val contextFeatures = ContextFeaturesStorage.extract(psiFile)
val result = mutableMapOf<String, Any>()
for (provider in ElementFeatureProvider.forLanguage(psiFile.language)) {
val name = provider.name
@@ -4,7 +4,7 @@ package com.intellij.stats.completion
import com.intellij.codeInsight.lookup.LookupManager
import com.intellij.codeInsight.lookup.impl.LookupImpl
import com.intellij.codeInsight.completion.ml.ContextFeatureProvider
import com.intellij.codeInsight.completion.ml.ContextFeatures
import com.intellij.completion.ml.ContextFeaturesStorage
import com.intellij.codeInsight.completion.ml.MLFeatureValue
import com.intellij.completion.settings.CompletionMLRankingSettings
import com.intellij.completion.tracker.PositionTrackingListener
@@ -107,8 +107,8 @@ class CompletionTrackerInitializer(experimentHelper: WebServiceStatus) {
}
}
ContextFeatures.setContextFeatures(file, result)
Disposer.register(lookup, Disposable { ContextFeatures.clear(file) })
ContextFeaturesStorage.setContextFeatures(file, result)
Disposer.register(lookup, Disposable { ContextFeaturesStorage.clear(file) })
lookupStorage.contextFactors = result.mapValues { it.value.toString() }
}
}