PySpark实现:统计非零值间零值单元格数量(大数据场景)
基于PySpark高效统计每行非零值间零值数量的方案
需求说明
给定带表头的表格,需统计每行中两个非零值单元格之间的零值单元格数量(示例输出为(3,2))。由于Pandas逐行逐单元格的处理方式在大数据集下效率低下,因此采用PySpark的分布式计算能力实现高效处理。
实现方案(优先原生Spark函数)
核心思路
通过数组转换、位置展开、窗口函数组合的方式,避免逐行遍历,利用Spark的分布式执行引擎处理大数据集。
代码示例
1. 初始化环境与测试数据
from pyspark.sql import SparkSession from pyspark.sql import functions as F from pyspark.sql.window import Window spark = SparkSession.builder.appName("ZeroCountBetweenNonZero").getOrCreate() # 测试数据: # 行1:[1,0,0,0,2,3] → 非零值间零值数量为3(1和2之间) # 行2:[4,0,0,5] → 非零值间零值数量为2(4和5之间) data = [(1, 0, 0, 0, 2, 3), (4, 0, 0, 5)] df = spark.createDataFrame(data, ["col1", "col2", "col3", "col4", "col5", "col6"])
2. 数据处理步骤
# 生成唯一行ID(若已有业务唯一键可跳过此步),将所有列转为数组 df_with_rowid = df.withColumn("row_id", F.monotonically_increasing_id()) \ .withColumn("values", F.array(*df.columns)) # 展开数组的索引与值,过滤掉零值 exploded_df = df_with_rowid.select("row_id", F.posexplode("values").alias("pos", "val")) \ .filter(F.col("val") != 0) # 用窗口函数获取下一个非零值的位置,计算中间零值数量 window_spec = Window.partitionBy("row_id").orderBy("pos") result_df = exploded_df.withColumn("next_pos", F.lead("pos").over(window_spec)) \ .filter(F.col("next_pos").isNotNull()) \ .withColumn("zero_count", F.col("next_pos") - F.col("pos") - 1) # 按行收集零值数量,得到最终结果 final_result = result_df.groupBy("row_id") \ .agg(F.collect_list("zero_count").alias("zero_counts")) \ .select("zero_counts") final_result.show(truncate=False)
输出结果
+------------+ |zero_counts | +------------+ |[3] | |[2] | +------------+
替代方案:Pandas向量化UDF
若原生函数无法满足复杂逻辑,可使用Pandas UDF(向量化处理,比普通Python UDF效率数倍):
from pyspark.sql.functions import pandas_udf from pyspark.sql.types import ArrayType, IntegerType import pandas as pd @pandas_udf(ArrayType(IntegerType())) def count_zeros_between_nonzero(arr: pd.Series) -> pd.Series: def process_row(row): non_zero_pos = [i for i, val in enumerate(row) if val != 0] return [non_zero_pos[i+1] - non_zero_pos[i] - 1 for i in range(len(non_zero_pos)-1)] return arr.apply(process_row) # 直接调用UDF计算 df.withColumn("zero_counts", count_zeros_between_nonzero(F.array(*df.columns))) \ .select("zero_counts") \ .show(truncate=False)
优化说明
- 优先使用原生Spark函数:原生函数基于JVM执行,避免Python-JVM通信开销,性能最优;
- 若使用UDF,务必选择Pandas向量化UDF,而非普通Python UDF;
- 若表中已有唯一行标识符(如业务主键),无需生成
row_id,直接用主键分组即可。
内容的提问来源于stack exchange,提问作者Tom Nguyen
相关产品推荐
相关产品推荐

