处理含缺失值的大规模PySpark DataFrame:移除高相关列
大规模PySpark DataFrame高效计算带缺失值的相关矩阵
针对600万行、2000列且含大量高缺失率列的场景,直接用原生Correlation或循环列对的方法完全不可行,以下是高效解决方案:
第一步:先过滤高缺失率列
先移除缺失占比超90%的列,减少计算量——这类列信息密度极低,对后续分析无意义,还会拖慢计算速度。
from pyspark.sql import functions as F # 计算每列的缺失率 missing_rates = db.agg(*[ (1 - F.count(c)/F.count(F.lit(1))).alias(f"{c}_missing_rate") for c in db.columns ]).collect()[0].asDict() # 筛选缺失率<=90%的列 filtered_cols = [c for c in db.columns if missing_rates[f"{c}_missing_rate"] <= 0.9] db_filtered = db.select(filtered_cols)
第二步:高效计算Pearson相关系数(基于统计量推导)
利用Pearson相关系数的数学公式,通过分布式聚合计算所有列对的核心统计量,避免循环:
2.1 生成所有唯一列对
# 生成i<j的列对,避免重复计算 col_pairs = [] cols = db_filtered.columns for i in range(len(cols)): for j in range(i+1, len(cols)): col_pairs.append( (cols[i], cols[j]) )
2.2 定义聚合函数计算列对统计量
对每个列对,计算共同非缺失行的n(计数)、sum_x、sum_y、sum_x2、sum_y2、sum_xy:
# 构建聚合表达式 agg_exprs = [] for col1, col2 in col_pairs: # 标记当前列对的有效行(两列都非空) valid_flag = F.col(col1).isNotNull() & F.col(col2).isNotNull() # 计算该列对的统计量 exprs = [ F.count(F.when(valid_flag, 1)).alias(f"{col1}_{col2}_n"), F.sum(F.when(valid_flag, col1)).alias(f"{col1}_{col2}_sum_x"), F.sum(F.when(valid_flag, col2)).alias(f"{col1}_{col2}_sum_y"), F.sum(F.when(valid_flag, col1*col1)).alias(f"{col1}_{col2}_sum_x2"), F.sum(F.when(valid_flag, col2*col2)).alias(f"{col1}_{col2}_sum_y2"), F.sum(F.when(valid_flag, col1*col2)).alias(f"{col1}_{col2}_sum_xy"), ] agg_exprs.extend(exprs) # 执行聚合,得到所有列对的统计量 stats = db_filtered.agg(*agg_exprs).collect()[0].asDict()
2.3 计算相关系数并构建矩阵
import numpy as np num_cols = len(cols) corr_matrix = np.eye(num_cols, dtype=np.float32) # 对角线初始化为1 # 填充矩阵 for idx_i, col1 in enumerate(cols): for idx_j, col2 in enumerate(cols): if idx_i >= idx_j: continue # 跳过对角线和已计算的对称位置 # 获取当前列对的统计量 n = stats[f"{col1}_{col2}_n"] sum_x = stats[f"{col1}_{col2}_sum_x"] sum_y = stats[f"{col1}_{col2}_sum_y"] sum_x2 = stats[f"{col1}_{col2}_sum_x2"] sum_y2 = stats[f"{col1}_{col2}_sum_y2"] sum_xy = stats[f"{col1}_{col2}_sum_xy"] # 计算Pearson相关系数 numerator = n * sum_xy - sum_x * sum_y denominator_x = n * sum_x2 - sum_x **2 denominator_y = n * sum_y2 - sum_y **2 denominator = np.sqrt(denominator_x * denominator_y) if denominator == 0: corr = 0.0 # 若分母为0(如列值全相同),设为0 else: corr = numerator / denominator # 填充对称位置 corr_matrix[idx_i][idx_j] = corr corr_matrix[idx_j][idx_i] = corr
第三步:Spearman相关系数的高效计算
Spearman相关本质是对列的秩计算Pearson相关,只需先对每列计算秩(忽略缺失值),再复用上述Pearson计算逻辑:
from pyspark.sql.window import Window # 计算每列的秩(忽略缺失值,相同值取平均秩) ranked_cols = [] for c in cols: ranked_col = F.percent_rank().over( Window.orderBy(c) ).alias(f"{c}_rank") ranked_cols.append(ranked_col) db_ranked = db_filtered.select(*ranked_cols) # 复用上述Pearson相关系数的计算代码,把db_filtered换成db_ranked即可
关键优化点
- 提前过滤高缺失列:直接减少列数,大幅降低后续计算的列对数量
- 分布式聚合替代循环:利用Spark的分布式计算能力,一次性计算所有列对的统计量,避免O(n²)的循环开销
- 避免重复计算:只计算i<j的列对,填充对称矩阵,减少一半计算量
内容的提问来源于stack exchange,提问作者Hadij
相关产品推荐
相关产品推荐

