You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Java Spark:分类数据下approxNearestNeighbor的键向量创建方法

解决Spark MinHashLSH中分类键转查询向量的问题

核心思路很简单:复用你已经训练好的plModel(特征转换Pipeline模型)来处理单个分类键,这样能保证转换逻辑和训练数据集完全一致,避免手动构造向量时出现的映射错误。

具体步骤如下:

1. 包装查询的分类键为DataFrame

因为Spark的Pipeline只能处理Dataset<Row>,所以我们需要把要查询的分类键(比如"banana")包装成和原始输入数据结构一致的小DataFrame:

// 假设我们要查询的分类键是"banana"
List<Row> queryData = Arrays.asList(RowFactory.create(99, "banana")); // id随便填一个不冲突的值就行
Dataset<Row> queryDf = spark.createDataFrame(queryData, schema); // 复用原始数据的schema

2. 用训练好的Pipeline转换得到特征向量

直接调用plModel.transform()处理这个查询DataFrame,就能得到和训练数据格式一致的features列:

Dataset<Row> transformedQuery = plModel.transform(queryDf);

3. 提取查询向量

从转换后的DataFrame中提取出features向量,作为approxNearestNeighbor的key参数:

Vector queryKey = transformedQuery.first().getAs<Vector>("features");

完整的查询代码整合

把上面的步骤整合到你的原有代码中,完整示例如下:

List<Row> dataA = Arrays.asList(RowFactory.create(0, "apple"), RowFactory.create(1, "banana"), RowFactory.create(2, "coconut")); 
StructType schema = new StructType( new StructField[] { 
    new StructField("id", DataTypes.IntegerType, false, Metadata.empty()), 
    new StructField("fruits", DataTypes.StringType, false, Metadata.empty()) 
}); 
Dataset<Row> dfA = spark.createDataFrame(dataA, schema); 

StringIndexer stringIndexer = new StringIndexer().setInputCol("fruits").setOutputCol("fruitIndex").setHandleInvalid("keep"); 
OneHotEncoderEstimator encoder = new OneHotEncoderEstimator().setInputCols(new String[]{"fruitIndex"}).setOutputCols(new String[]{"fruitVec"}); 
String[] featuredCols = new String[] {"fruitIndex","fruitVec"}; 
VectorAssembler assembler = new VectorAssembler().setInputCols(featuredCols).setOutputCol("features"); 

Pipeline sovPipeline = new Pipeline().setStages(new PipelineStage[]{stringIndexer, encoder, assembler}); 
PipelineModel plModel = sovPipeline.fit(dfA); 
Dataset<Row> dfT = plModel.transform(dfA); 

MinHashLSH mh = new MinHashLSH().setNumHashTables(5).setInputCol("features").setOutputCol("hashes"); 
MinHashLSHModel model = mh.fit(dfT); 

// --- 新增的查询向量生成部分 ---
// 要查询的目标分类键
String targetFruit = "banana";
// 包装成符合输入格式的DataFrame
List<Row> queryData = Arrays.asList(RowFactory.create(99, targetFruit));
Dataset<Row> queryDf = spark.createDataFrame(queryData, schema);
// 通过训练好的Pipeline转换得到特征向量
Dataset<Row> transformedQuery = plModel.transform(queryDf);
Vector queryKey = transformedQuery.first().getAs<Vector>("features");
// 执行近似最近邻查询
model.approxNearestNeighbors(dfT, queryKey, 2).show();

关键注意事项

  • 绝对不要手动构造向量:StringIndexer的索引映射是基于训练数据的统计结果生成的,手动编写比如apple→0这类映射很容易出错,尤其是当训练数据更新后,复用Pipeline能始终保证映射逻辑一致。
  • 如果查询的分类键不在训练数据中,因为你设置了setHandleInvalid("keep"),StringIndexer会给它分配一个新的索引,后续的Encoder和Assembler也能正常处理,但这种情况下的近似邻接结果可能没有实际业务意义,你可以根据需求调整handleInvalid的策略。

内容的提问来源于stack exchange,提问作者MSS

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.28 09:27:36