Spark如何在Java中对分组数据集进行自定义整体修改?
解决方案:分组移除第二条记录并拼接剩余文本
方案一:使用窗口函数+分组聚合(无需自定义UDAF)
这个方法无需编写自定义聚合逻辑,利用Spark内置函数即可快速实现需求,步骤清晰易维护:
- 为每个
Names分组内的记录添加行号,定位需要移除的第二条记录; - 过滤掉行号为2的记录;
- 按
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
相关产品推荐
相关产品推荐

