如何在PySpark DataFrame中按公司计算递推Balance列?
问题:为PySpark DataFrame按公司计算累计Balance列
现有如下PySpark DataFrame,需为每个公司新增Balance列:每个公司的首行Balance取值自Check列,后续每行的Balance按公式**上一行Balance值 × (1 + interest值)**计算。
原DataFrame
+----------+--------+------------+---------+ | Date|interest|Check | Company| +----------+--------+------------+---------+ |2023-12-31| 7| 704|Company A| |2024-12-31| 7| 704|Company A| |2025-12-31| 7| 704|Company A| |2026-12-31| 7| 704|Company A| |2023-12-31| 7| 704|Company B| |2024-12-31| 7| 704|Company B| |2025-12-31| 7| 704|Company B| |2026-12-31| 7| 704|Company B| |2023-12-31| 7| 704|Company C| |2024-12-31| 7| 704|Company C| |2025-12-31| 7| 704|Company C| |2026-12-31| 7| 704|Company C| +----------+--------+------------+---------+
预期输出
+----------+--------+------------+---------+-------------------+ | Date|interest|Check | Company| Balance | +----------+--------+------------+---------+-------------------+ |2023-12-31| 7| 704|Company A| 704 | |2024-12-31| 7| 704|Company A| 1507.680000000000| |2025-12-31| 7| 704|Company A| 2321.797600000000| |2026-12-31| 7| 704|Company A| 3306.447232000000| |2023-12-31| 7| 704|Company B| 704 | |2024-12-31| 7| 704|Company B| 1507.680000000000| |2025-12-31| 7| 704|Company B| 2321.797600000000| |2026-12-31| 7| 704|Company B| 3306.447232000000| |2023-12-31| 7| 704|Company C| 704 | |2024-12-31| 7| 704|Company C| 1507.680000000000| |2025-12-31| 7| 704|Company C| 2321.797600000000| |2026-12-31| 7| 704|Company C| 3306.447232000000| +----------+--------+------------+---------+-------------------+
解决方案
由于需要按公司分组进行递推计算,可使用分组Pandas UDF实现,该方式逻辑直观且适合处理这类依赖上一行结果的计算场景。
步骤与代码示例
- 导入必要依赖并创建Spark会话
from pyspark.sql import SparkSession from pyspark.sql.types import StructType, StructField, StringType, IntegerType, DoubleType from pyspark.sql.functions import pandas_udf, PandasUDFType # 初始化Spark会话 spark = SparkSession.builder.appName("BalanceCalculation").getOrCreate()
- 创建示例DataFrame(若已有DataFrame可跳过此步)
# 原数据 data = [ ("2023-12-31", 7, 704, "Company A"), ("2024-12-31", 7, 704, "Company A"), ("2025-12-31", 7, 704, "Company A"), ("2026-12-31", 7, 704, "Company A"), ("2023-12-31", 7, 704, "Company B"), ("2024-12-31", 7, 704, "Company B"), ("2025-12-31", 7, 704, "Company B"), ("2026-12-31", 7, 704, "Company B"), ("2023-12-31", 7, 704, "Company C"), ("2024-12-31", 7, 704, "Company C"), ("2025-12-31", 7, 704, "Company C"), ("2026-12-31", 7, 704, "Company C") ] # 定义schema schema = StructType([ StructField("Date", StringType()), StructField("interest", IntegerType()), StructField("Check", IntegerType()), StructField("Company", StringType()) ]) df = spark.createDataFrame(data, schema)
- 定义分组Pandas UDF计算Balance列
# 定义输出schema,在原schema基础上新增Balance列 output_schema = StructType(schema.fields + [StructField("Balance", DoubleType())]) @pandas_udf(output_schema, PandasUDFType.GROUPED_MAP) def calculate_balance(pdf): # 按日期排序,确保每行按时间顺序计算 pdf = pdf.sort_values("Date") # 首行Balance取Check列的值 pdf["Balance"] = pdf["Check"].iloc[0] # 从第二行开始,按公式递推计算Balance for i in range(1, len(pdf)): # 注:若interest为百分比(如7代表7%),需改为 (1 + pdf["interest"].iloc[i]/100) pdf["Balance"].iloc[i] = pdf["Balance"].iloc[i-1] * (1 + pdf["interest"].iloc[i]/100) return pdf # 按Company分组应用UDF result_df = df.groupBy("Company").apply(calculate_balance)
- 查看结果
result_df.show(truncate=False)
关键说明
- 必须先按
Date排序,确保每个公司的行按时间顺序排列,否则递推计算会出错。 - 若
interest字段为百分比数值(如7代表7%),需在公式中除以100;若为实际倍数(如7代表7倍),则直接使用(1 + pdf["interest"].iloc[i])。
内容的提问来源于stack exchange,提问作者cnns
相关产品推荐
相关产品推荐

