PySpark中计算行的均值/标准差/总和是否有更简洁的实现方式?
PySpark行级统计量(均值、标准差)的简洁实现方式
PySpark内置了least、greatest这类行级极值计算函数,但没有直接提供行均值、标准差、求和的函数。尝试用pivot方法实现时内存占用过高,因此手动实现了相关逻辑,这里提供更简洁的替代方案:
利用PySpark数组函数简化实现(适用于PySpark 3.0+)
PySpark 3.0及以上版本提供了数组统计类函数,能直接用来实现行级统计,无需手动编写求和、平方差计算逻辑:
行均值(与原逻辑一致,将null转为0计算)
from pyspark.sql import functions as f def row_mean(*cols): # 把目标列转为数组,将null替换为0后计算均值 processed_arr = f.array(*[f.when(f.col(c).isNull(), 0).otherwise(f.col(c)) for c in cols]) return f.array_mean(processed_arr)
行总体标准差(与原逻辑一致,将null转为0计算)
注意:PySpark内置的array_stddev默认计算样本标准差(分母为N-1),若要和原代码一致的总体标准差(分母为N),需要做转换:
def row_stddev(*cols): processed_arr = f.array(*[f.when(f.col(c).isNull(), 0).otherwise(f.col(c)) for c in cols]) n = len(cols) # 转换为总体标准差 return f.array_stddev(processed_arr) * f.sqrt((n - 1) / n)
如果不需要将null转为0,而是直接忽略null值计算统计量,代码可以更简洁:
# 忽略null的行均值 def row_mean_ignore_null(*cols): return f.array_mean(f.array(*cols)) # 忽略null的行总体标准差 def row_stddev_ignore_null(*cols): n = len(cols) return f.array_stddev(f.array(*cols)) * f.sqrt((n - 1) / n)
低版本PySpark兼容方案(3.0以下)
如果你的PySpark版本低于3.0,没有array_mean、array_stddev,可以用aggregate函数实现求和逻辑:
def row_mean(*cols): processed_arr = f.array(*[f.when(f.col(c).isNull(), 0).otherwise(f.col(c)) for c in cols]) # 用aggregate累加数组元素求和,再除以列数 total = f.aggregate(processed_arr, f.lit(0.0), lambda acc, x: acc + x) return total / len(cols) def row_stddev(*cols): n = len(cols) mu = row_mean(*cols) processed_arr = f.array(*[f.when(f.col(c).isNull(), 0).otherwise(f.col(c)) for c in cols]) # 累加平方差后计算总体标准差 sum_sq_diff = f.aggregate(processed_arr, f.lit(0.0), lambda acc, x: acc + f.pow(x - mu, 2)) return f.sqrt(sum_sq_diff / n)
使用示例
调用方式和原代码完全一致:
day_stats = data.select( f.least(*data.columns[:-1]).alias("min"), f.greatest(*data.columns[:-1]).alias("max"), row_mean(*data.columns[:-1]).alias("mean"), row_stddev(*data.columns[:-1]).alias("stddev"), data.columns[-1], ).show()
内容的提问来源于stack exchange,提问作者Axeltherabbit
相关产品推荐
相关产品推荐

