diff --git a/plugins/search-everywhere-ml/ranking/core/src/com/intellij/searchEverywhereMl/ranking/core/model/SearchEverywhereFeatureNameConverter.kt b/plugins/search-everywhere-ml/ranking/core/src/com/intellij/searchEverywhereMl/ranking/core/model/SearchEverywhereFeatureNameConverter.kt new file mode 100644 index 000000000000..88a95215f184 --- /dev/null +++ b/plugins/search-everywhere-ml/ranking/core/src/com/intellij/searchEverywhereMl/ranking/core/model/SearchEverywhereFeatureNameConverter.kt @@ -0,0 +1,41 @@ +// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license that can be found in the LICENSE file. +package com.intellij.searchEverywhereMl.ranking.core.model + +// This conversion is required while bundled models still expect legacy camelCase names. +// It can be removed after models are re-trained with snake_case feature names. +private val LEGACY_FEATURE_NAME_ALIASES = mapOf( + "action_plugin_type" to "pluginType", + "is_in_library" to "isFromLibrary", + "max_usage_se" to "maxUsageSE", + "min_usage_se" to "minUsageSE", +) + +internal fun withLegacyFeatureNameAliases(features: Map): Map { + val featuresWithLegacyAliases = HashMap(features.size * 2) + for ((featureName, value) in features) { + featuresWithLegacyAliases[featureName] = value + featuresWithLegacyAliases.putIfAbsent(toLegacyFeatureName(featureName), value) + } + return featuresWithLegacyAliases +} + +internal fun toLegacyFeatureName(featureName: String): String { + return LEGACY_FEATURE_NAME_ALIASES[featureName] ?: featureName.snakeCaseToCamelCase() +} + +private fun String.snakeCaseToCamelCase(): String { + if (indexOf('_') < 0) return this + + val result = StringBuilder(length) + var uppercaseNext = false + for (char in this) { + if (char == '_') { + uppercaseNext = true + continue + } + + result.append(if (uppercaseNext) char.uppercaseChar() else char) + uppercaseNext = false + } + return result.toString() +} diff --git a/plugins/search-everywhere-ml/ranking/core/src/com/intellij/searchEverywhereMl/ranking/core/model/SearchEverywhereRankingModel.kt b/plugins/search-everywhere-ml/ranking/core/src/com/intellij/searchEverywhereMl/ranking/core/model/SearchEverywhereRankingModel.kt index efde2cf47896..8b4edd0c535f 100644 --- a/plugins/search-everywhere-ml/ranking/core/src/com/intellij/searchEverywhereMl/ranking/core/model/SearchEverywhereRankingModel.kt +++ b/plugins/search-everywhere-ml/ranking/core/src/com/intellij/searchEverywhereMl/ranking/core/model/SearchEverywhereRankingModel.kt @@ -12,7 +12,7 @@ import com.intellij.searchEverywhereMl.ranking.core.features.SearchEverywhereFil internal abstract class SearchEverywhereRankingModel(protected val model: DecisionFunction) { abstract fun predict(features: Map): Double protected fun rawModelPredict(features: Map): Double { - return model.predict(buildArray(model.featuresOrder, features)) + return model.predict(buildArray(model.featuresOrder, withLegacyFeatureNameAliases(features))) } private fun buildArray(featuresOrder: Array, features: Map): DoubleArray { diff --git a/plugins/search-everywhere-ml/ranking/core/test/com/intellij/searchEverywhereMl/ranking/core/model/SearchEverywhereFeatureNameConverterTest.kt b/plugins/search-everywhere-ml/ranking/core/test/com/intellij/searchEverywhereMl/ranking/core/model/SearchEverywhereFeatureNameConverterTest.kt new file mode 100644 index 000000000000..681f0e6b4ce8 --- /dev/null +++ b/plugins/search-everywhere-ml/ranking/core/test/com/intellij/searchEverywhereMl/ranking/core/model/SearchEverywhereFeatureNameConverterTest.kt @@ -0,0 +1,38 @@ +package com.intellij.searchEverywhereMl.ranking.core.model + +import org.junit.Assert.assertEquals +import org.junit.Test +import org.junit.runner.RunWith +import org.junit.runners.JUnit4 + +@RunWith(JUnit4::class) +internal class SearchEverywhereFeatureNameConverterTest { + @Test + fun `snake case name is converted to camelCase`() { + assertEquals("queryLength", toLegacyFeatureName("query_length")) + } + + @Test + fun `non trivial legacy names are taken from lookup table`() { + assertEquals("pluginType", toLegacyFeatureName("action_plugin_type")) + assertEquals("isFromLibrary", toLegacyFeatureName("is_in_library")) + assertEquals("maxUsageSE", toLegacyFeatureName("max_usage_se")) + assertEquals("minUsageSE", toLegacyFeatureName("min_usage_se")) + } + + @Test + fun `aliases are added and existing keys are preserved`() { + val features = mapOf( + "query_length" to 7, + "queryLength" to 3, + "is_in_library" to true, + ) + + val featuresWithAliases = withLegacyFeatureNameAliases(features) + + assertEquals(7, featuresWithAliases["query_length"]) + assertEquals(3, featuresWithAliases["queryLength"]) + assertEquals(true, featuresWithAliases["is_in_library"]) + assertEquals(true, featuresWithAliases["isFromLibrary"]) + } +}