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

Spark如何在Java中对分组数据集进行自定义整体修改?

解决方案:分组移除第二条记录并拼接剩余文本

方案一:使用窗口函数+分组聚合(无需自定义UDAF)

这个方法无需编写自定义聚合逻辑,利用Spark内置函数即可快速实现需求,步骤清晰易维护:

  1. 为每个Names分组内的记录添加行号,定位需要移除的第二条记录;
  2. 过滤掉行号为2的记录;
  3. 按Names分组,将剩余的Random_Text拼接成字符串。

Java代码示例

import org.apache.spark.sql.Dataset;
import org.apache.spark.sql.Row;
import org.apache.spark.sql.SparkSession;
import org.apache.spark.sql.expressions.Window;
import org.apache.spark.sql.expressions.WindowSpec;
import static org.apache.spark.sql.functions.*;

public class GroupProcessExample {
    public static void main(String[] args) {
        SparkSession spark = SparkSession.builder()
                .appName("GroupRemoveSecondAndConcat")
                .master("local[*]")
                .getOrCreate();

        // 构造示例数据集
        Dataset<Row> dtf = spark.createDataFrame(
                new Object[][]{
                        {"Michael", "Hello"},
                        {"Jim", "Good"},
                        {"Bob", "How"},
                        {"Michael", "Good"},
                        {"Michael", "Morning"},
                        {"Bob", "Are"},
                        {"Bob", "You"},
                        {"Bob", "Doing"},
                        {"Jim", "Bye"}
                },
                new String[]{"Names", "Random_Text"}
        );

        // 定义窗口:按Names分组,用自增ID保证原始顺序(有明确排序字段可替换)
        WindowSpec windowSpec = Window.partitionBy(col("Names")).orderBy(monotonically_increasing_id());
        Dataset<Row> numberedDf = dtf.withColumn("row_num", row_number().over(windowSpec));

        // 过滤每组的第二条记录
        Dataset<Row> filteredDf = numberedDf.filter(col("row_num").notEqual(2));

        // 分组拼接剩余文本
        Dataset<Row> resultDf = filteredDf.groupBy(col("Names"))
                .agg(concat_ws(" ", collect_list(col("Random_Text"))).alias("Random_Text"));

        resultDf.show();
        spark.stop();
    }
}

方案二:自定义UserDefinedAggregateFunction(UDAF)

如果业务需求更复杂,需要完全自定义聚合逻辑,可以实现Java版UDAF,在聚合过程中跳过每组第二条记录,最终拼接剩余文本:

自定义UDAF实现类

import org.apache.spark.sql.Row;
import org.apache.spark.sql.expressions.MutableAggregationBuffer;
import org.apache.spark.sql.expressions.UserDefinedAggregateFunction;
import org.apache.spark.sql.types.DataType;
import org.apache.spark.sql.types.DataTypes;
import org.apache.spark.sql.types.StructField;
import org.apache.spark.sql.types.StructType;

import java.util.ArrayList;
import java.util.List;

public class SkipSecondConcatUDAF extends UserDefinedAggregateFunction {

    // 输入字段:单个字符串类型的Random_Text
    @Override
    public StructType inputSchema() {
        return DataTypes.createStructType(
                new StructField[]{DataTypes.createStructField("text", DataTypes.StringType, true)}
        );
    }

    // 缓冲区:存储文本列表和当前计数
    @Override
    public StructType bufferSchema() {
        return DataTypes.createStructType(new StructField[]{
                DataTypes.createStructField("textList", DataTypes.createArrayType(DataTypes.StringType), true),
                DataTypes.createStructField("count", DataTypes.IntegerType, false)
        });
    }

    // 输出类型:拼接后的字符串
    @Override
    public DataType dataType() {
        return DataTypes.StringType;
    }

    // 确定性:相同输入返回相同结果
    @Override
    public boolean deterministic() {
        return true;
    }

    // 初始化缓冲区
    @Override
    public void initialize(MutableAggregationBuffer buffer) {
        buffer.update(0, new ArrayList<String>());
        buffer.update(1, 0);
    }

    // 更新缓冲区:跳过第二条记录,其余加入列表
    @Override
    public void update(MutableAggregationBuffer buffer, Row input) {
        List<String> textList = (List<String>) buffer.get(0);
        int count = buffer.getInt(1) + 1;
        String text = input.getString(0);
        if (count != 2) {
            textList.add(text);
        }
        buffer.update(0, textList);
        buffer.update(1, count);
    }

    // 合并分区缓冲区
    @Override
    public void merge(MutableAggregationBuffer buffer1, Row buffer2) {
        List<String> list1 = (List<String>) buffer1.get(0);
        List<String> list2 = (List<String>) buffer2.get(0);
        list1.addAll(list2);
        buffer1.update(0, list1);
        buffer1.update(1, buffer1.getInt(1) + buffer2.getInt(1));
    }

    // 生成最终拼接结果
    @Override
    public Object evaluate(Row buffer) {
        List<String> textList = (List<String>) buffer.get(0);
        return String.join(" ", textList);
    }
}

主程序中使用UDAF

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.*;

public class UDAFExample {
    public static void main(String[] args) {
        SparkSession spark = SparkSession.builder()
                .appName("SkipSecondConcatUDAF")
                .master("local[*]")
                .getOrCreate();

        // 注册自定义UDAF
        spark.udf().register("skipSecondConcat", new SkipSecondConcatUDAF());

        // 构造示例数据集
        Dataset<Row> dtf = spark.createDataFrame(
                new Object[][]{
                        {"Michael", "Hello"},
                        {"Jim", "Good"},
                        {"Bob", "How"},
                        {"Michael", "Good"},
                        {"Michael", "Morning"},
                        {"Bob", "Are"},
                        {"Bob", "You"},
                        {"Bob", "Doing"},
                        {"Jim", "Bye"}
                },
                new String[]{"Names", "Random_Text"}
        );

        // 调用UDAF进行分组聚合
        Dataset<Row> resultDf = dtf.groupBy(col("Names"))
                .agg(callUDF("skipSecondConcat", col("Random_Text")).alias("Random_Text"));

        resultDf.show();
        spark.stop();
    }
}

注意事项

  • 窗口函数方案中,monotonically_increasing_id()用于保证原始输入顺序,若数据集有明确排序字段(如时间戳),建议替换为该字段以确保行号准确性;
  • 自定义UDAF方案中,若需要严格保证文本顺序,建议先为数据添加全局排序字段,再在UDAF中按顺序处理。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.03 05:21:57