ML in SE: Compute ML weight when creating FoundElementInfo (IDEA-294763)

GitOrigin-RevId: 48695b856e8b80427f606df5bb9bedc61a933a17
This commit is contained in:
Adam Malek
2022-06-11 16:51:25 +00:00
committed by intellij-monorepo-bot
parent 5b1ef53f46
commit 6c9af3fbe8
12 changed files with 103 additions and 117 deletions
@@ -2,31 +2,12 @@
package com.intellij.ide.actions.searcheverywhere;
public class FoundItemDescriptor<I> {
private final static int MAX_WEIGHT = 10_000;
private final static int ML_WEIGHT_DISABLED_VALUE = -1;
private final I item;
private final int weight;
private final int mlWeight;
public FoundItemDescriptor(I item, int weight) {
this.item = item;
this.weight = weight;
mlWeight = ML_WEIGHT_DISABLED_VALUE; // not set
}
public FoundItemDescriptor(I item, int weight, double mlWeight) {
this.item = item;
this.weight = weight;
this.mlWeight = (int)(mlWeight * MAX_WEIGHT) * 100_000 + this.weight;
}
public int getMlWeight() {
return mlWeight;
}
public boolean isMlWeightSet() {
return mlWeight != ML_WEIGHT_DISABLED_VALUE;
}
public I getItem() {
@@ -34,6 +15,6 @@ public class FoundItemDescriptor<I> {
}
public int getWeight() {
return isMlWeightSet() ? mlWeight : weight;
return weight;
}
}
@@ -109,15 +109,7 @@ public class ActionSearchEverywhereContributor implements WeightedSearchEverywhe
return true;
}
final FoundItemDescriptor<GotoActionModel.MatchedValue> descriptor;
SearchEverywhereMlService mlService = SearchEverywhereMlService.getInstance();
if (mlService != null) {
descriptor = getMLWeightedItemDescriptor(mlService, element);
}
else {
descriptor = new FoundItemDescriptor<>(element, element.getMatchingDegree());
}
final var descriptor = new FoundItemDescriptor<GotoActionModel.MatchedValue>(element, element.getMatchingDegree());
return consumer.process(descriptor);
});
}, progressIndicator);
@@ -231,19 +223,6 @@ public class ActionSearchEverywhereContributor implements WeightedSearchEverywhe
});
}
private FoundItemDescriptor<GotoActionModel.MatchedValue> getMLWeightedItemDescriptor(@NotNull SearchEverywhereMlService service,
@NotNull GotoActionModel.MatchedValue element) {
double mlWeight = service.getMlWeight(this, element, element.getMatchingDegree());
if (mlWeight > 0 && service.shouldOrderByMl()) {
if (element.getType() == GotoActionModel.MatchedValueType.ABBREVIATION) {
return new FoundItemDescriptor<>(element, element.getMatchingDegree(), 1.0);
}
return new FoundItemDescriptor<>(element, element.getMatchingDegree(), mlWeight);
}
return new FoundItemDescriptor<>(element, element.getMatchingDegree());
}
public static class Factory implements SearchEverywhereContributorFactory<GotoActionModel.MatchedValue> {
@NotNull
@Override
@@ -107,15 +107,6 @@ public class FileSearchEverywhereContributor extends AbstractGotoSEContributor {
return true;
}
SearchEverywhereMlService mlService = SearchEverywhereMlService.getInstance();
if (mlService != null) {
double mlWeight = mlService.getMlWeight(this, element, degree);
if (mlWeight >= 0.0 && mlService.shouldOrderByMl()) {
return consumer.process(new FoundItemDescriptor<>(element, degree, mlWeight));
}
}
return consumer.process(new FoundItemDescriptor<>(element, degree));
}
@@ -275,7 +275,14 @@ class GroupedResultsSearcher implements SESearcher {
assert contributor == myExpandedContributor; // Only expanded contributor items allowed
Collection<SearchEverywhereFoundElementInfo> section = sections.get(contributor);
SearchEverywhereFoundElementInfo newElementInfo = new SearchEverywhereFoundElementInfo(element, priority, contributor);
final var mlService = SearchEverywhereMlService.getInstance();
final SearchEverywhereFoundElementInfo newElementInfo;
if (mlService == null) {
newElementInfo = new SearchEverywhereFoundElementInfo(element, priority, contributor);
}
else {
newElementInfo = mlService.createFoundElementInfo(contributor, element, priority);
}
if (section.size() >= myNewLimit) {
return false;
@@ -347,7 +354,15 @@ class GroupedResultsSearcher implements SESearcher {
@Override
public boolean addElement(Object element, SearchEverywhereContributor<?> contributor, int priority, ProgressIndicator indicator) throws InterruptedException {
SearchEverywhereFoundElementInfo newElementInfo = new SearchEverywhereFoundElementInfo(element, priority, contributor);
final var mlService = SearchEverywhereMlService.getInstance();
SearchEverywhereFoundElementInfo newElementInfo;
if (mlService == null) {
newElementInfo = new SearchEverywhereFoundElementInfo(element, priority, contributor);
}
else {
newElementInfo = mlService.createFoundElementInfo(contributor, element, priority);
}
Condition condition = conditionsMap.get(contributor);
Collection<SearchEverywhereFoundElementInfo> section = sections.get(contributor);
int limit = sectionsLimits.get(contributor);
@@ -268,7 +268,15 @@ class MixedResultsSearcher implements SESearcher {
}
public boolean addElement(Object element, SearchEverywhereContributor<?> contributor, int priority, ProgressIndicator indicator) throws InterruptedException {
SearchEverywhereFoundElementInfo newElementInfo = new SearchEverywhereFoundElementInfo(element, priority, contributor);
final var mlService = SearchEverywhereMlService.getInstance();
final SearchEverywhereFoundElementInfo newElementInfo;
if (mlService == null) {
newElementInfo = new SearchEverywhereFoundElementInfo(element, priority, contributor);
}
else {
newElementInfo = mlService.createFoundElementInfo(contributor, element, priority);
}
Condition condition = conditionsMap.get(contributor);
Collection<SearchEverywhereFoundElementInfo> section = mySections.get(contributor);
int limit = sectionsLimits.get(contributor);
@@ -1,18 +1,16 @@
// Copyright 2000-2021 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.ide.actions.searcheverywhere;
import com.intellij.util.concurrency.annotations.RequiresEdt;
import com.intellij.openapi.util.UserDataHolderBase;
import org.jetbrains.annotations.ApiStatus;
import org.jetbrains.annotations.NotNull;
import javax.swing.*;
import java.util.Comparator;
/**
* Class containing info about found elements
*/
@ApiStatus.Internal
public class SearchEverywhereFoundElementInfo {
public class SearchEverywhereFoundElementInfo extends UserDataHolderBase {
public final int priority;
public final Object element;
public final SearchEverywhereContributor<?> contributor;
@@ -5,6 +5,7 @@ import com.intellij.openapi.diagnostic.Logger
import com.intellij.openapi.extensions.ExtensionPointName
import com.intellij.openapi.project.Project
import org.jetbrains.annotations.ApiStatus
import org.jetbrains.annotations.Contract
@ApiStatus.Internal
abstract class SearchEverywhereMlService {
@@ -39,6 +40,11 @@ abstract class SearchEverywhereMlService {
abstract fun onSessionStarted(project: Project?, mixedListInfo: SearchEverywhereMixedListInfo)
@Contract("_, _, _ -> new")
abstract fun createFoundElementInfo(contributor: SearchEverywhereContributor<*>,
element: Any,
priority: Int): SearchEverywhereFoundElementInfo
abstract fun onSearchRestart(project: Project?, tabId: String, reason: SearchRestartReason,
keysTyped: Int, backspacesTyped: Int, searchQuery: String,
previousElementsProvider: () -> List<SearchEverywhereFoundElementInfo>)
@@ -124,7 +124,7 @@ internal class SearchEverywhereMLStatisticsCollector : CounterUsagesCollector()
data.addAll(additional)
data.addAll(context.features)
val elementData = getElementsData(selectedElements, elements, elementIdProvider, selectedItems, project, state)
val elementData = getElementsData(selectedElements, elements, elementIdProvider, selectedItems, project)
data.add(IS_PROJECT_DISPOSED_KEY.with(elementData == null))
if (elementData != null) {
data.addAll(elementData)
@@ -145,11 +145,10 @@ internal class SearchEverywhereMLStatisticsCollector : CounterUsagesCollector()
elements: List<SearchEverywhereFoundElementInfo>,
elementIdProvider: SearchEverywhereMlItemIdProvider,
selectedItems: List<Any>,
project: Project?,
state: SearchEverywhereMlSearchState): List<EventPair<*>>? {
project: Project?): List<EventPair<*>>? {
return ArrayList<EventPair<*>>().apply {
addAll(getSelectedElementsData(selectedElements, elements, elementIdProvider, selectedItems))
addAll(getCollectedElementsData(elements, project, elementIdProvider, state) ?: return null)
addAll(getCollectedElementsData(elements, project, elementIdProvider) ?: return null)
}
}
@@ -182,8 +181,7 @@ internal class SearchEverywhereMLStatisticsCollector : CounterUsagesCollector()
*/
private fun getCollectedElementsData(elements: List<SearchEverywhereFoundElementInfo>,
project: Project?,
elementIdProvider: SearchEverywhereMlItemIdProvider,
state: SearchEverywhereMlSearchState): ArrayList<EventPair<*>>? {
elementIdProvider: SearchEverywhereMlItemIdProvider): ArrayList<EventPair<*>>? {
val data = ArrayList<EventPair<*>>()
val actionManager = ActionManager.getInstance()
val value = elements.take(REPORTED_ITEMS_LIMIT).map {
@@ -195,7 +193,7 @@ internal class SearchEverywhereMLStatisticsCollector : CounterUsagesCollector()
CONTRIBUTOR_ID_KEY.with(it.contributor.searchProviderId)
)
addElementFeatures(elementIdProvider.getId(it.element), it, state, result, actionManager)
addElementFeatures(elementIdProvider.getId(it.element), it, result, actionManager)
ObjectEventData(result)
}
data.add(COLLECTED_RESULTS_DATA_KEY.with(value))
@@ -219,19 +217,19 @@ internal class SearchEverywhereMLStatisticsCollector : CounterUsagesCollector()
private fun addElementFeatures(elementId: Int?,
elementInfo: SearchEverywhereFoundElementInfo,
state: SearchEverywhereMlSearchState,
result: MutableList<EventPair<*>>,
actionManager: ActionManager) {
val itemInfo = state.getElementFeatures(elementId, elementInfo.element, elementInfo.contributor, elementInfo.priority)
if (itemInfo.features.isNotEmpty()) {
result.add(FEATURES_DATA_KEY.with(ObjectEventData(itemInfo.features)))
val mlFeatures = elementInfo.getUserData(SearchEverywhereMlSessionService.ML_FEATURES_KEY) ?: emptyList()
if (mlFeatures.isNotEmpty()) {
result.add(FEATURES_DATA_KEY.with(ObjectEventData(mlFeatures)))
}
state.getMLWeightIfDefined(elementId)?.let { score ->
result.add(ML_WEIGHT_KEY.with(roundDouble(score)))
val mlWeight = elementInfo.getUserData(SearchEverywhereMlSessionService.ML_WEIGHT_KEY) ?: -1.0
if (mlWeight >= 0.0) {
result.add(ML_WEIGHT_KEY.with(roundDouble(mlWeight)))
}
itemInfo.id?.let { result.add(ID_KEY.with(it)) }
elementId?.let { result.add(ID_KEY.with(it)) }
doWhenIsActionWrapper(elementInfo.element) {
val action = it.action
@@ -18,9 +18,6 @@ internal class SearchEverywhereMlSearchState(
private val modelProvider: SearchEverywhereModelProvider,
private val providersCache: FeaturesProviderCache?
) {
private val cachedElementsInfo: MutableMap<Int, SearchEverywhereMLItemInfo> = hashMapOf()
private val cachedMLWeight: MutableMap<Int, Double> = hashMapOf()
val searchStateFeatures = SearchEverywhereStateFeaturesProvider().getSearchStateFeatures(tabId, searchQuery)
private val model: SearchEverywhereRankingModel by lazy {
@@ -32,18 +29,6 @@ internal class SearchEverywhereMlSearchState(
element: Any,
contributor: SearchEverywhereContributor<*>,
priority: Int): SearchEverywhereMLItemInfo {
if (elementId == null) {
return computeElementFeatures(null, element, priority, contributor)
}
return cachedElementsInfo.computeIfAbsent(elementId) {
return@computeIfAbsent computeElementFeatures(elementId, element, priority, contributor)
}
}
private fun computeElementFeatures(elementId: Int?,
element: Any,
priority: Int,
contributor: SearchEverywhereContributor<*>): SearchEverywhereMLItemInfo {
val features = arrayListOf<EventPair<*>>()
val contributorId = contributor.searchProviderId
SearchEverywhereElementFeaturesProvider.getFeatureProvidersForContributor(contributorId).forEach { provider ->
@@ -53,31 +38,12 @@ internal class SearchEverywhereMlSearchState(
return SearchEverywhereMLItemInfo(elementId, contributorId, features)
}
@Synchronized
fun getMLWeightIfDefined(elementId: Int?): Double? {
return elementId?.let { cachedMLWeight[elementId] }
}
@Synchronized
fun getMLWeight(elementId: Int?,
element: Any,
contributor: SearchEverywhereContributor<*>,
context: SearchEverywhereMLContextInfo,
priority: Int): Double {
if (elementId == null) {
return computeMLWeight(context, null, element, contributor, priority)
}
return cachedMLWeight.computeIfAbsent(elementId) {
computeMLWeight(context, elementId, element, contributor, priority)
}
}
private fun computeMLWeight(context: SearchEverywhereMLContextInfo,
elementId: Int?,
element: Any,
contributor: SearchEverywhereContributor<*>,
priority: Int): Double {
val features = ArrayList<EventPair<*>>()
features.addAll(context.features)
features.addAll(getElementFeatures(elementId, element, contributor, priority).features)
@@ -3,13 +3,21 @@ package com.intellij.ide.actions.searcheverywhere.ml
import com.intellij.ide.actions.searcheverywhere.*
import com.intellij.ide.actions.searcheverywhere.ml.settings.SearchEverywhereMlSettings
import com.intellij.ide.util.gotoByName.GotoActionModel
import com.intellij.internal.statistic.eventLog.events.EventPair
import com.intellij.openapi.components.service
import com.intellij.openapi.project.Project
import com.intellij.openapi.util.Key
import java.util.concurrent.atomic.AtomicInteger
import java.util.concurrent.atomic.AtomicReference
internal class SearchEverywhereMlSessionService : SearchEverywhereMlService() {
companion object {
val ML_FEATURES_KEY = Key.create<List<EventPair<*>>>("se-ml-features")
val ML_WEIGHT_KEY = Key.create<Double>("se-ml-weight")
private const val MAX_ELEMENT_WEIGHT = 10_000
internal const val RECORDER_CODE = "MLSE"
fun getService() = EP_NAME.findExtensionOrFail(SearchEverywhereMlSessionService::class.java).takeIf { it.isEnabled() }
@@ -53,6 +61,42 @@ internal class SearchEverywhereMlSessionService : SearchEverywhereMlService() {
}
}
override fun createFoundElementInfo(contributor: SearchEverywhereContributor<*>,
element: Any,
priority: Int): SearchEverywhereFoundElementInfo {
val foundElementInfoWithoutMl = SearchEverywhereFoundElementInfo(element, priority, contributor)
if (!isEnabled()) return foundElementInfoWithoutMl
val session = getCurrentSession() ?: return foundElementInfoWithoutMl
val state = session.getCurrentSearchState() ?: return foundElementInfoWithoutMl
val elementId = session.itemIdProvider.getId(element)
val mlFeatures = state.getElementFeatures(elementId, element, contributor, priority).features
// We can only compute the ML weight if we are in a tab which has a ML model
val mlWeight = state.takeIf { SearchEverywhereTabWithMl.findById(it.tabId) != null }
?.getMLWeight(elementId, element, contributor, session.cachedContextInfo, priority)
val foundElementInfo = if (shouldOrderByMl() && mlWeight != null) {
val priorityWithMl = getPriorityWithMl(element, mlWeight, priority)
SearchEverywhereFoundElementInfo(element, priorityWithMl, contributor)
}
else {
foundElementInfoWithoutMl
}
foundElementInfo.putUserData(ML_FEATURES_KEY, mlFeatures)
foundElementInfo.putUserData(ML_WEIGHT_KEY, mlWeight)
return foundElementInfo
}
private fun getPriorityWithMl(element: Any, mlWeight: Double, priority: Int): Int {
val weight = if (element is GotoActionModel.MatchedValue && element.type == GotoActionModel.MatchedValueType.ABBREVIATION) 1.0 else mlWeight
return (weight * MAX_ELEMENT_WEIGHT).toInt() * 100_000 + priority
}
override fun onSearchRestart(project: Project?,
tabId: String,
reason: SearchRestartReason,
@@ -7,6 +7,8 @@ import com.intellij.ide.actions.searcheverywhere.SearchEverywhereUI
import com.intellij.ide.actions.searcheverywhere.ml.SearchEverywhereMlSessionService
import com.intellij.ide.actions.searcheverywhere.ml.SearchEverywhereTabWithMl
import com.intellij.ide.actions.searcheverywhere.ml.features.SearchEverywhereContributorFeaturesProvider
import com.intellij.ide.actions.searcheverywhere.ml.SearchEverywhereMlSessionService.Companion.ML_FEATURES_KEY
import com.intellij.ide.actions.searcheverywhere.ml.SearchEverywhereMlSessionService.Companion.ML_WEIGHT_KEY
import com.intellij.ide.scratch.ScratchFileCreationHelper
import com.intellij.ide.scratch.ScratchFileService
import com.intellij.ide.scratch.ScratchRootType
@@ -57,12 +59,14 @@ class OpenFeaturesInScratchFileAction : AnAction() {
val searchSession = mlSessionService.getCurrentSession()!!
val state = searchSession.getCurrentSearchState()
val tabId = searchEverywhereUI.selectedTabID
val features = searchEverywhereUI.foundElementsInfo.map { info ->
val rankingWeight = info.priority
val contributor = info.contributor.searchProviderId
val elementName = StringUtil.notNullize(info.element.toString(), "undefined")
val mlWeight = if (isTabWithMl(tabId)) mlSessionService.getMlWeight(info.contributor, info.element, rankingWeight) else null
val mlWeight = info.getUserData(ML_WEIGHT_KEY)
val mlFeatures: Map<String, Any> = info.getUserData(ML_FEATURES_KEY)
?.associate { it.field.name to it.data as Any }
?: emptyMap()
val elementId = searchSession.itemIdProvider.getId(info.element)
return@map ElementFeatures(
@@ -71,7 +75,7 @@ class OpenFeaturesInScratchFileAction : AnAction() {
mlWeight,
rankingWeight,
contributor,
state!!.getElementFeatures(elementId, info.element, info.contributor, rankingWeight).featuresAsMap().toSortedMap()
mlFeatures.toSortedMap()
)
}
@@ -87,10 +91,6 @@ class OpenFeaturesInScratchFileAction : AnAction() {
)
}
private fun isTabWithMl(tabId: String): Boolean {
return SearchEverywhereTabWithMl.findById(tabId) != null
}
private fun createScratchFileContext(json: String) = ScratchFileCreationHelper.Context().apply {
text = json
fileExtension = JsonFileType.DEFAULT_EXTENSION
@@ -25,8 +25,10 @@ internal abstract class SearchEverywhereRankingModelTest
protected fun performSearchFor(searchQuery: String, featuresProviderCache: FeaturesProviderCache? = null): RankingAssertion {
VirtualFileManager.getInstance().syncRefresh()
val rankedElements: List<FoundItemDescriptor<*>> = filterElements(searchQuery)
.map { it.withMlWeight(getMlWeight(it, searchQuery, featuresProviderCache)) }
.sortedByDescending { it.mlWeight }
.associateWith { getMlWeight(it, searchQuery, featuresProviderCache) }
.entries
.sortedByDescending { it.value }
.map { it.key }
assert(rankedElements.size > 1) { "Found ${rankedElements.size} which is unsuitable for ranking assertions" }
@@ -53,8 +55,6 @@ internal abstract class SearchEverywhereRankingModelTest
}.fold(emptyMap()) { acc, value -> acc + value }
}
private fun FoundItemDescriptor<*>.withMlWeight(mlWeight: Double) = FoundItemDescriptor(this.item, this.weight, mlWeight)
protected class RankingAssertion(private val results: List<FoundItemDescriptor<*>>) {
fun thenAssertElement(element: FoundItemDescriptor<*>) = ElementAssertion(element)
fun findElementAndAssert(predicate: (FoundItemDescriptor<*>) -> Boolean) = ElementAssertion(results.find(predicate)!!)