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
相关产品推荐
相关产品推荐

