如何在PySpark DataFrame上实现np.fill_diagonal对角线填充功能
报错原因
np.fill_diagonal是numpy专属方法,仅支持操作numpy数组对象,而PySpark DataFrame是分布式数据集,不属于numpy数组类型,自然没有ndim这类numpy数组才有的属性,所以触发报错。
可行实现方案
方案1:纯PySpark原生实现(推荐,支持大数据量)
透视表的对角线本质是行维度的product_1值与列名(即pivot生成的product_2枚举值)相等的单元格,直接按这个规则逐列判断替换即可,代码如下:
# 拿到所有pivot生成的列名(排除第一列product_1) pivot_cols = [col for col in df_spark_res.columns if col != 'product_1'] # 逐个列处理:对角线位置的空值替换为0 select_expr = [F.col('product_1')] for col_name in pivot_cols: select_expr.append( F.when( (F.col('product_1') == F.lit(col_name)) & (F.col(col_name).isNull()), 0 ).otherwise(F.col(col_name)).alias(col_name) ) # 生成替换后的结果表 df_filled = df_spark_res.select(*select_expr)
该方案全程走分布式计算,不会把数据拉到单节点内存,适合大数据量场景。
方案2:小数据量下转Pandas处理
如果你的数据量很小,可以把PySpark DataFrame转成Pandas DataFrame后再用numpy方法填充,最后转回PySpark DataFrame即可:
import pandas as pd # 转pandas pd_df = df_spark_res.toPandas() # 填充对角线,第一列是product_1,所以从第二列开始算数值区域 np.fill_diagonal(pd_df.values[:, 1:], 0) # 转回PySpark DataFrame df_filled = spark.createDataFrame(pd_df)
注意该方案会把全量数据拉到driver节点内存,数据量过大会触发OOM,仅适合小数据集使用。
内容的提问来源于stack exchange,提问作者madfrostx
相关产品推荐
相关产品推荐

