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
相关产品推荐
相关产品推荐

