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

如何在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();
    }
}

关键代码说明

  1. struct("id", "name", "address"):将Employee的所有字段打包成一个结构体,确保聚合时能保留完整的对象信息。
  2. collect_list(...):将每组的结构体数据收集为一个列表,作为分组后的结果字段。
  3. RDD转换为Map:通过JavaRDD的collectAsMap方法,将分组后的Row数据转为键为id、值为Employee列表的Map,完全匹配期望的输出格式。

注:如果Employee类重写了toString()方法,打印列表时会自动输出和示例一致的对象格式。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 11:33:30