You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

求助:将Python到期值求和逻辑转换为PySpark代码

PySpark代码修正:基于ExpirationPeriod计算折扣总和新列

需求说明

现有包含float类型ExpirationPeriod列的数据集,需根据该列值的整数部分,计算每一年折扣值的总和并生成新列。折扣值计算公式为:100/1.02**i + 5,其中i从0遍历到「整数部分-1」。

可正常运行的Python实现

import pandas as pd

df = pd.DataFrame({'ExpirationPeriod':[1.2,2.0,3.0,4.5]})

def sum_data(row):
    return sum([100/1.02**i + 5 for i in range(int(row['ExpirationPeriod']))])

df['output'] = df.apply(sum_data, axis=1)
print(df)

执行后输出:

ExpirationPeriod      output
0               1.2  105.000000
1               2.0  208.921569
2               3.0  311.805459
3               4.5  413.662215

你的PySpark代码问题分析

你写的PySpark代码存在几个核心问题:

  • UDF逻辑完全偏离需求:当前lambda计算的是range(x)的求和,而非题目要求的折扣值总和
  • 参数传递冗余:用struct包裹单列没必要,直接传递列值即可
  • 返回类型不匹配:UDF声明返回DoubleType,但代码里返回的是列表,类型冲突

修正后的PySpark代码

from pyspark.sql.functions import udf, col
from pyspark.sql.types import IntegerType, DoubleType

# 创建测试数据集
df = spark.createDataFrame(
    [(1.2, ),
     (2.0, ),
     (3.0, ),
     (4.5, )],
    ["ExpirationPeriod"])

# 提取ExpirationPeriod的整数部分
df = df.withColumn("Counts", df["ExpirationPeriod"].cast(IntegerType()))

# 实现正确逻辑的UDF:计算折扣值总和
def calculate_discount_sum(n):
    total = 0.0
    for i in range(n):
        total += 100 / (1.02 ** i) + 5
    return total

discount_sum_udf = udf(calculate_discount_sum, DoubleType())

# 生成新列
new_df = df.withColumn("output", discount_sum_udf(col("Counts")))

# 查看结果
new_df.show(truncate=False)

执行后输出:

+----------------+------+-------------------+
|ExpirationPeriod|Counts|output             |
+----------------+------+-------------------+
|1.2             |1     |105.0              |
|2.0             |2     |208.92156862745098 |
|3.0             |3     |311.8054594386774  |
|4.5             |4     |413.66221513595823 |
+----------------+------+-------------------+

结果和Python版本完全一致。

进阶优化:避免UDF(推荐)

PySpark中UDF性能不如内置函数,可通过sequence+aggregate实现无UDF版本:

from pyspark.sql.functions import sequence, lit, aggregate, pow

new_df = df.withColumn("Counts", df["ExpirationPeriod"].cast(IntegerType())) \
           .withColumn(
               "output",
               aggregate(
                   sequence(lit(0), col("Counts")-1),
                   lit(0.0),
                   lambda acc, i: acc + (100 / pow(lit(1.02), i)) + lit(5)
               )
           )

new_df.show(truncate=False)

此版本无需自定义UDF,性能更优,结果一致。

内容的提问来源于stack exchange,提问作者Mike P.

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.25 11:58:19