PySpark如何按id分组对列计算带25阈值的重置式累计求和
实现思路
你需要的是按id分组的带重置阈值的累计求和,这类需要跨行传递状态的计算,普通逐行UDF和lag窗口函数没法直接实现:普通UDF是逐行执行的,无法保留前序行的累计状态;lag窗口函数只能获取前N行的固定值,无法传递动态变化的累加状态,因此都无法实现该需求。
这里推荐用Spark 3.0+支持的applyInPandas方法,在分组内直接遍历值计算,代码简洁易维护。
完整实现代码
from pyspark.sql import SparkSession from pyspark.sql.types import DoubleType, StructField import pandas as pd # 构造示例数据 spark = SparkSession.builder.appName("cumsum_reset").getOrCreate() df = spark.createDataFrame( [('Mark', 0.0), ('Mark', 1), ('Mark', 1), ('Mark', 1), ('Mark', 25), ('Mark', 1), ('Mark', 1),('Mark', 1),('Mark', 20), ('Mark', 1),('Mark', 1),('Mark', 1), ('Mark', 1),('Mark', 1),('John', 0), ('John', 1),('John', 1),('John', 1), ('John', 1),('John', 1),('John', 1), ('John', 1),('John', 9),('John', 1), ('John', 1),('John', 1),('John', 1), ('John', 1),('John', 1),('John', 1), ('John', 1),('John', 1),('John', 1), ('John', 7),('John', 1)], ('id', "V")) # 定义分组内计算逻辑 def reset_cumsum(pdf: pd.DataFrame) -> pd.DataFrame: threshold = 25 current_sum = 0 res = [] for v in pdf['V']: if current_sum >= threshold: current_sum = v else: current_sum += v res.append(current_sum) pdf['cumsum_reset'] = res return pdf # 分组应用计算 result = df.groupBy('id').applyInPandas( reset_cumsum, schema=df.schema.add(StructField("cumsum_reset", DoubleType())) ) # 查看结果 result.orderBy('id').show(50)
结果验证
运行后输出的累计和完全符合要求:比如Mark分组前5行的累计和依次为0、1、2、3、28,第6行因为上一次累计和28≥25,所以重置为1,后续累加依次为2、3、23、24、25、26、27,和预期效果一致。
低版本Spark替代方案
如果你使用的是Spark 2.x不支持applyInPandas,可以用aggregateByKey先收集分组内所有V值再计算:
# 先给每行加全局排序id,保证分组内顺序和输入一致 from pyspark.sql import Window from pyspark.sql.functions import row_number, monotonically_increasing_id df_with_idx = df.withColumn("idx", row_number().over(Window.orderBy(monotonically_increasing_id()))) # 按id分组收集有序的V值 grouped_rdd = df_with_idx.rdd.map(lambda x: (x['id'], (x['idx'], x['V'])))\ .groupByKey()\ .mapValues(lambda x: [v for (idx, v) in sorted(x, key=lambda i: i[0])]) # 计算带重置的累计和 def calc_cumsum_list(v_list): threshold = 25 current_sum = 0 res = [] for v in v_list: if current_sum >= threshold: current_sum = v else: current_sum += v res.append(current_sum) return res processed_rdd = grouped_rdd.flatMap(lambda x: [(x[0], v, cs) for v, cs in zip(x[1], calc_cumsum_list(x[1]))]) result = processed_rdd.toDF(["id", "V", "cumsum_reset"]) result.show(50)
内容的提问来源于stack exchange,提问作者K. Sante
相关产品推荐
相关产品推荐

