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

Spark DataFrame聚合/去重:实现键唯一的机器学习数据集转换

处理Spark DataFrame重复键的无偏解决方案,适配机器学习需求

嘿,这个需求简直是机器学习数据预处理里的常客——数据源带重复键的情况太闹心了,搞不好就给模型喂进偏差数据。不过Spark的API足够灵活,完全能针对不同类型的列给出无偏的处理方案,我给你整理了几种常用的实现方式,直接就能用:


1. 数值型列:均值合并(保留统计信息)

如果你的重复行里是连续数值特征(比如用户点击量、商品交易金额),用均值合并是最稳妥的选择——既能保留数据的统计特征,又不会给模型引入人为偏差。

代码示例

Scala版本

import org.apache.spark.sql.functions._

// 假设键是user_id,需要合并的数值列是click_count、purchase_amount
val deduplicatedDf = df.groupBy("user_id")
  .agg(
    avg("click_count").alias("click_count"),
    avg("purchase_amount").alias("purchase_amount"),
    first("register_time").alias("register_time") // 非数值列可保留首/末次记录
  )

Python版本

from pyspark.sql import functions as F

deduplicated_df = df.groupBy("user_id") \
  .agg(
    F.avg("click_count").alias("click_count"),
    F.avg("purchase_amount").alias("purchase_amount"),
    F.first("register_time").alias("register_time")
  )

2. 字符串/标签型列:拼接或取众数(适配分类场景)

如果重复行里是分类标签(比如用户兴趣标签、商品类别),可以根据需求选两种处理方式:

2.1 字符串拼接(保留所有标签信息)

适合需要完整保留用户/物品所有标签的场景,比如推荐系统里的兴趣画像:

// Scala示例
val deduplicatedDf = df.groupBy("user_id")
  .agg(
    collect_set("interest_tags").alias("unique_tags"), // 自动去重后拼接
    concat_ws("|", collect_list("interest_tags")).alias("full_tags") // 保留所有重复标签
  )
# Python示例
deduplicated_df = df.groupBy("user_id") \
  .agg(
    F.collect_set("interest_tags").alias("unique_tags"),
    F.concat_ws("|", F.collect_list("interest_tags")).alias("full_tags")
  )

2.2 取众数(贴合主流标签,避免冗余)

如果标签有明确的主流值(比如用户最常浏览的类别),取频次最高的标签更贴合真实情况:

// Scala示例:先统计标签频次,再取每组频次最高的标签
val tagFrequencyDf = df.groupBy("user_id", "interest_tags")
  .count()
  .orderBy(desc("count"))

val deduplicatedDf = tagFrequencyDf.groupBy("user_id")
  .agg(first("interest_tags").alias("dominant_tag"))
# Python示例
tag_frequency_df = df.groupBy("user_id", "interest_tags") \
  .count() \
  .orderBy(F.desc("count"))

deduplicated_df = tag_frequency_df.groupBy("user_id") \
  .agg(F.first("interest_tags").alias("dominant_tag"))

3. 通用场景:随机取值(完全中立无偏)

如果某些列没有明确的统计规律,或者你不想引入任何统计假设,随机选取重复行中的一行是最中立的方式——绝对不会给模型带来偏向性。

代码示例

Scala版本

import org.apache.spark.sql.expressions.Window

val windowSpec = Window.partitionBy("user_id").orderBy(rand())

val deduplicatedDf = df.withColumn("rand_num", rand())
  .withColumn("row_num", row_number().over(windowSpec))
  .filter(col("row_num") === 1)
  .drop("rand_num", "row_num")

Python版本

from pyspark.sql.window import Window

window_spec = Window.partitionBy("user_id").orderBy(F.rand())

deduplicated_df = df.withColumn("rand_num", F.rand()) \
  .withColumn("row_num", F.row_number().over(window_spec)) \
  .filter(F.col("row_num") == 1) \
  .drop("rand_num", "row_num")

额外小提示

处理前建议先排查重复情况,明确哪些键有重复、重复次数多少,方便选择最合适的方案:

// Scala:查看重复键的数量
df.groupBy("user_id").count().filter(col("count") > 1).show()
# Python:查看重复键的数量
df.groupBy("user_id").count().filter(F.col("count") > 1).show()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.20 07:57:56