Spark如何查找同主键记录中值不一致的列(Scala/Python)
找出DataFrame同主键组合下的差异列
给定一个DataFrame,需要找出相同主键组合对应的记录中,哪些列的值不匹配。例如:
- 主键为
lastName + firstName时,smith+john组合的savingRate列值不匹配(29.99 vs 0); chen+kodo组合的paymentRate和balance列值不匹配。
理论上同主键不应存在重复记录,但当前数据存在同主键特征列值不一致的情况,需要定位这些差异列。
Scala 解决方案
你的思路方向可行,但用均值和最小值判断存在漏洞(比如不同值的均值可能恰好等于最小值),更可靠的方式是统计每个特征列的不同值数量,或者比较最小值和最大值是否相等。以下是修正后的实现:
import org.apache.spark.sql.functions._ import org.apache.spark.sql.Column // 定义主键列和特征列 val primaryKeys = Seq("lastName", "firstName") val features_list = df.columns.filter(!primaryKeys.contains(_)) // 对每个特征列,计算同主键下的不同值数量 val distinctCountCols = features_list.map(colName => countDistinct(col(colName)).as(s"distinct_$colName") ) // 按主键分组,计算每个特征列的不同值计数 val groupedDf = df.groupBy(primaryKeys.head, primaryKeys.tail:_*).agg( distinctCountCols.head, distinctCountCols.tail:_* ) // 生成判断逻辑:找出不同值计数>1的特征列 val diffColumns = features_list.map(colName => when(col(s"distinct_$colName") > 1, lit(colName)).otherwise(lit(null)) ) // 合并差异列结果,得到每个主键对应的差异列列表 val result = groupedDf.withColumn("diff_columns", array_remove(array(diffColumns:_*), null)) .filter(size(col("diff_columns")) > 0) // 只保留有差异的主键组合 result.show(false)
代码说明
- 用
countDistinct统计每个特征列在同主键下的不同值数量,比均值/最小值的判断更准确; - 通过
when函数标记出有差异的列名,再用array和array_remove整理成差异列列表; - 最后过滤出存在差异的主键组合,直接展示哪些列有问题。
如果坚持用你的均值+最小值思路,修正后的代码如下(仅适用于数值型列):
import org.apache.spark.sql.functions._ val primaryKeys = Seq("lastName", "firstName") val features_list = df.columns.filter(!primaryKeys.contains(_)) // 同时计算每个特征列的min和avg,避免两次groupBy val aggExprs = features_list.flatMap(colName => Seq( min(col(colName)).cast("double").as(s"min_$colName"), avg(col(colName)).cast("double").as(s"avg_$colName") ) ) val groupedDf = df.groupBy(primaryKeys.head, primaryKeys.tail:_*).agg(aggExprs.head, aggExprs.tail:_*) // 生成差异列判断 val diffColumns = features_list.map(colName => when(col(s"min_$colName") =!= col(s"avg_$colName"), lit(colName)).otherwise(lit(null)) ) val result = groupedDf.withColumn("diff_columns", array_remove(array(diffColumns:_*), null)) .filter(size(col("diff_columns")) > 0) result.show(false)
Python 解决方案
如果用Python实现,逻辑和Scala一致,代码如下:
from pyspark.sql import functions as F # 定义主键和特征列 primary_keys = ["lastName", "firstName"] features_list = [col for col in df.columns if col not in primary_keys] # 计算每个特征列的不同值计数 distinct_count_cols = [F.countDistinct(F.col(col_name)).alias(f"distinct_{col_name}") for col_name in features_list] grouped_df = df.groupBy(*primary_keys).agg(*distinct_count_cols) # 生成差异列列表 diff_columns = [F.when(F.col(f"distinct_{col_name}") > 1, F.lit(col_name)).otherwise(F.lit(None)) for col_name in features_list] result = grouped_df.withColumn("diff_columns", F.array_remove(F.array(*diff_columns), None)) result = result.filter(F.size(F.col("diff_columns")) > 0) result.show(truncate=False)
内容的提问来源于stack exchange,提问作者helloworld
相关产品推荐
相关产品推荐

