如何将含9000+记录的CSV文件转换为RDD<Vector>(Spark Java)
手把手教你实现CSV转RDD并计算列相似度
嘿,这个需求我熟,咱们一步步来搞定!首先得明确:RowMatrix需要的是数值型的向量RDD,而你的CSV里的Category1和Category2是分类变量(字符串类型),得先把它们转换成数值特征,才能构建符合要求的矩阵。下面是具体的Java版Spark实现步骤,代码直接就能用:
1. 先初始化Spark环境
Spark 2.x之后推荐用SparkSession,比旧的SparkContext更方便,本地调试用local[*]就行,生产环境换成集群地址:
import org.apache.spark.SparkConf; import org.apache.spark.api.java.JavaSparkContext; import org.apache.spark.sql.SparkSession; // 配置Spark应用 SparkConf conf = new SparkConf().setAppName("CSVToRowMatrixDemo").setMaster("local[*]"); JavaSparkContext sc = new JavaSparkContext(conf); SparkSession spark = SparkSession.builder().config(conf).getOrCreate();
2. 读取你的CSV文件
用Spark SQL读取CSV,要是你的CSV第一行是表头(比如写了id,Category1,Category2),就加上option("header", "true");没有表头的话,就用默认的_c0、_c1、_c2指代列:
import org.apache.spark.sql.Dataset; import org.apache.spark.sql.Row; // 替换成你自己的CSV文件路径 Dataset<Row> rawDF = spark.read() .option("header", "true") // 有表头就保留,没有就删掉这行 .option("inferSchema", "false") // 手动控制类型,避免自动推断出错 .csv("/path/to/your/9000_records.csv");
3. 把分类变量转成数值特征
因为Vector只认数值,所以得把Category1和Category2这两个字符串分类转成数值。这里分两种常用场景,你按需选:
场景A:直接用分类索引值作为向量元素
这种最简单,适合基于分类的索引(比如把每个分类值映射成0、1、2...)计算相似度:
import org.apache.spark.ml.feature.StringIndexer; import org.apache.spark.ml.feature.StringIndexerModel; // 给Category1做索引映射 StringIndexer c1Indexer = new StringIndexer() .setInputCol("Category1") .setOutputCol("Category1_Index"); StringIndexerModel c1Model = c1Indexer.fit(rawDF); Dataset<Row> indexedDF = c1Model.transform(rawDF); // 给Category2做同样的索引映射 StringIndexer c2Indexer = new StringIndexer() .setInputCol("Category2") .setOutputCol("Category2_Index"); StringIndexerModel c2Model = c2Indexer.fit(indexedDF); Dataset<Row> finalIndexedDF = c2Model.transform(indexedDF);
然后把DataFrame转成RDD<Vector>:
import org.apache.spark.api.java.JavaRDD; import org.apache.spark.mllib.linalg.Vectors; import org.apache.spark.mllib.linalg.Vector; JavaRDD<Vector> vectorRDD = finalIndexedDF.javaRDD().map(row -> { // 取出两个索引列的数值 double c1Val = row.getDouble(row.fieldIndex("Category1_Index")); double c2Val = row.getDouble(row.fieldIndex("Category2_Index")); // 构建密集向量 return Vectors.dense(c1Val, c2Val); });
场景B:用独热编码生成高维度特征向量
如果你需要把每个分类的不同取值作为独立的特征列(比如Category1有A、B、C三个值,就生成三个二进制特征),就用独热编码:
import org.apache.spark.ml.feature.OneHotEncoder; import org.apache.spark.ml.feature.VectorAssembler; // 先做StringIndexer(和场景A一样,先把分类转成索引) StringIndexer c1Indexer = new StringIndexer().setInputCol("Category1").setOutputCol("Category1_Index"); StringIndexer c2Indexer = new StringIndexer().setInputCol("Category2").setOutputCol("Category2_Index"); Dataset<Row> indexedDF = c1Indexer.fit(rawDF).transform(rawDF); indexedDF = c2Indexer.fit(indexedDF).transform(indexedDF); // 对索引列做独热编码 OneHotEncoder c1Encoder = new OneHotEncoder() .setInputCol("Category1_Index") .setOutputCol("Category1_Vec"); OneHotEncoder c2Encoder = new OneHotEncoder() .setInputCol("Category2_Index") .setOutputCol("Category2_Vec"); Dataset<Row> encodedDF = c1Encoder.fit(indexedDF).transform(indexedDF); encodedDF = c2Encoder.fit(encodedDF).transform(encodedDF); // 把两个独热向量合并成一个完整的特征向量 VectorAssembler assembler = new VectorAssembler() .setInputCols(new String[]{"Category1_Vec", "Category2_Vec"}) .setOutputCol("features"); Dataset<Row> featureDF = assembler.transform(encodedDF); // 转成RDD<Vector> JavaRDD<Vector> vectorRDD = featureDF.javaRDD().map(row -> (Vector) row.getAs("features"));
4. 创建RowMatrix并计算列相似度
现在有了RDD<Vector>,就可以创建RowMatrix,然后调用columnSimilarities()方法。这个方法可以传入一个阈值(比如0.1),过滤掉相似度低于这个值的结果,提升性能:
import org.apache.spark.mllib.linalg.distributed.RowMatrix; import org.apache.spark.mllib.linalg.distributed.MatrixEntry; // 构建RowMatrix RowMatrix rowMatrix = new RowMatrix(vectorRDD.rdd()); // 计算列相似度,传入0.1作为阈值,不需要过滤就传0.0 JavaRDD<MatrixEntry> similarityResult = rowMatrix.columnSimilarities(0.1).toJavaRDD(); // 打印结果看看 similarityResult.foreach(entry -> { System.out.printf("列%d 和 列%d 的相似度:%.4f%n", entry.i(), entry.j(), entry.value()); });
几个要注意的小细节
- 如果你的CSV里有缺失值,记得先处理!比如用
rawDF.na().drop()删掉缺失行,或者用rawDF.na().fill("unknown")给缺失的分类填充默认值,不然构建向量的时候会报错。 - 你的数据只有9000条,就算用精确计算(传0.0作为阈值)也不会有性能问题,放心用。
- 要是你用的Spark版本比较旧(比如1.x),API会有一点不一样,但核心逻辑是一致的。
内容的提问来源于stack exchange,提问作者Jitendra Patil
相关产品推荐
相关产品推荐

