如何在Spark Java中对自定义类Dataset按key分组
问题:使用Spark Java API按id分组Dataset并返回对象列表
我定义了如下Employee类:
public class Employee { int id; String name; String address; }
现在拥有一个Spark Dataset<Employee>实例,示例数据如下:
Employee(1,"test1","test1") Employee(1,"test2","test2") Employee(2,"test3","test3")
期望得到的分组结果为:
1---> [Employee(1,"test1","test1"),Employee(1,"test2","test2")] 2---> Employee(2,"test3","test3")
即需要按id字段对Dataset进行分组,并将结果以对象列表形式返回,请问如何通过Spark Java API实现?
解决方案
可以通过Spark SQL的groupBy结合collect_list聚合函数实现需求,具体实现如下:
完整代码示例
import org.apache.spark.sql.Dataset; import org.apache.spark.sql.Row; import org.apache.spark.sql.SparkSession; import static org.apache.spark.sql.functions.*; import java.util.List; import java.util.Map; import java.util.stream.Collectors; public class EmployeeGrouping { public static void main(String[] args) { // 初始化SparkSession SparkSession spark = SparkSession.builder() .appName("EmployeeGroupById") .master("local[*]") // 本地测试用,生产环境移除该配置 .getOrCreate(); // 模拟生成示例Dataset<Employee> Dataset<Employee> employeeDs = spark.createDataset( List.of( new Employee(1, "test1", "test1"), new Employee(1, "test2", "test2"), new Employee(2, "test3", "test3") ), org.apache.spark.sql.Encoders.bean(Employee.class) ); // 按id分组,收集每组的Employee对象为列表 Dataset<Row> groupedResult = employeeDs.groupBy("id") .agg(collect_list(struct("id", "name", "address")).as("employee_list")); // 转换为Map格式,匹配期望输出的键值结构 Map<Integer, List<Employee>> resultMap = groupedResult.javaRDD() .map(row -> { Integer id = row.getAs("id"); // 将Row列表转为Employee对象列表 List<Employee> empList = row.getAs("employee_list") .stream() .map(r -> { Employee emp = new Employee(); emp.id = r.getAs("id"); emp.name = r.getAs("name"); emp.address = r.getAs("address"); return emp; }) .collect(Collectors.toList()); return new scala.Tuple2<>(id, empList); }) .collectAsMap(); // 打印结果 resultMap.forEach((id, emps) -> { System.out.println(id + "---> " + emps); }); spark.stop(); } }
关键代码说明
struct("id", "name", "address"):将Employee的所有字段打包成一个结构体,确保聚合时能保留完整的对象信息。collect_list(...):将每组的结构体数据收集为一个列表,作为分组后的结果字段。- RDD转换为Map:通过JavaRDD的
collectAsMap方法,将分组后的Row数据转为键为id、值为Employee列表的Map,完全匹配期望的输出格式。
注:如果Employee类重写了toString()方法,打印列表时会自动输出和示例一致的对象格式。
内容的提问来源于stack exchange,提问作者Programmer
相关产品推荐
相关产品推荐

