如何使用Spark JavaRDD的map与filter函数读取CSV并按指定列筛选
使用Spark JavaRDD结合map和filter处理CSV并筛选特定列
我来一步步教你怎么用Spark的JavaRDD,通过map和filter函数读取CSV文件,还能根据特定列做筛选。咱们就用你给的出租车数据CSV来举例。
步骤1:初始化SparkContext
首先得先创建SparkContext,这是Spark应用的入口:
import org.apache.spark.SparkConf; import org.apache.spark.api.java.JavaRDD; import org.apache.spark.api.java.JavaSparkContext; public class CsvProcessing { public static void main(String[] args) { // 配置Spark,本地调试用local[*],生产环境替换成集群master地址 SparkConf conf = new SparkConf().setAppName("CsvFilterExample").setMaster("local[*]"); JavaSparkContext sc = new JavaSparkContext(conf);
步骤2:读取CSV文件为JavaRDD
用textFile方法读取CSV,得到每行字符串的RDD:
// 替换成你的CSV文件实际路径 JavaRDD<String> csvRdd = sc.textFile("path/to/your/taxi_data.csv");
步骤3:分离表头和数据行
CSV第一行是表头,我们先把它单独拎出来,避免后续处理时把表头当成数据:
// 提取表头,然后过滤掉表头行,只保留数据行 String header = csvRdd.first(); JavaRDD<String> dataRdd = csvRdd.filter(line -> !line.equals(header));
步骤4:用map解析每行数据
接下来用map把每行字符串分割成可操作的结构,这里提供两种常用方式:
方式一:转换成字符串数组(快速上手)
直接按逗号分割每行,得到字段数组:
JavaRDD<String[]> parsedRdd = dataRdd.map(line -> line.split(","));
方式二:自定义JavaBean(推荐,代码更易维护)
先定义一个对应CSV字段的Java类,比如TaxiTrip:
public class TaxiTrip { private int vendorId; private String pickupDatetime; private int passengerCount; private double tripDistance; // 可以根据需求添加其他字段 // 从字符串数组初始化对象的构造方法 public TaxiTrip(String[] fields) { this.vendorId = Integer.parseInt(fields[0]); this.pickupDatetime = fields[1]; this.passengerCount = Integer.parseInt(fields[3]); this.tripDistance = Double.parseDouble(fields[4]); } // 提供getter方法供后续筛选使用 public int getPassengerCount() { return passengerCount; } public double getTripDistance() { return tripDistance; } }
然后用map把每行转换成TaxiTrip对象:
JavaRDD<TaxiTrip> tripRdd = dataRdd.map(line -> new TaxiTrip(line.split(",")));
步骤5:用filter根据特定列筛选数据
现在就可以用filter实现你需要的筛选逻辑了,举两个常见例子:
示例1:筛选乘客数大于1的记录(用字符串数组方式)
乘客数是CSV的第5列(索引为3):
JavaRDD<String[]> filteredRdd = parsedRdd.filter(fields -> Integer.parseInt(fields[3]) > 1);
示例2:筛选行程距离超过2公里的记录(用JavaBean方式)
这种方式代码可读性更高,不用记索引:
JavaRDD<TaxiTrip> filteredTripRdd = tripRdd.filter(trip -> trip.getTripDistance() > 2.0);
步骤6:输出或保存筛选结果
最后可以把结果打印出来,或者保存到文件:
// 打印筛选后的前10条数据 filteredTripRdd.take(10).forEach(trip -> System.out.println("乘客数:" + trip.getPassengerCount() + ",行程距离:" + trip.getTripDistance()) ); // 关闭SparkContext sc.close(); } }
小提示
- 如果你的CSV字段里包含逗号(比如带引号的字段),直接用
split(",")会出错,这时候可以用OpenCSV这类专门的CSV解析库来处理。 - 生产环境中记得把
setMaster("local[*]")替换成你的Spark集群master地址。
内容的提问来源于stack exchange,提问作者vinodh
相关产品推荐
相关产品推荐

