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

Palantir Foundry中如何获取df.withColumn的UDF内累加器最终值?

在Palantir Foundry中获取UDF内累加器的最终值

我尝试在Palantir Foundry中使用自定义累加器,在df.withColumn(...)调用的用户自定义函数(UDF)里聚合DataFrame每行数据。从生成的DataFrame能看到累加器值按预期递增,但主脚本里的累加器变量值未更新,且主脚本中累加器的Python ID与UDF内的不一致。想知道执行完成后,如何从主脚本获取DataFrame列中显示的累加器最终值?


Foundry环境测试代码

from transforms.api import transform_df, Input, Output
import numpy as np
from pyspark.accumulators import AccumulatorParam
from pyspark.sql.functions import udf, struct

global accum

@transform_df(
    Output("ri.foundry.main.dataset.xxxxxxxx-xxxx-xxxx-xxxx-xxxxxxxxxxxx"),
)
def compute(ctx):

    from pyspark.sql.types import StructType, StringType, IntegerType,  StructField

    data2 = [("James","","Smith","36636","M",3000),
        ("Michael","Rose","","40288","M",4000),
        ("Robert","","Williams","42114","M",4000),
        ("Maria","Anne","Jones","39192","F",4000),
        ("Jen","Mary","Brown","","F",-1)
    ]

    schema = StructType([ \
        StructField("firstname",StringType(),True), \
        StructField("middlename",StringType(),True), \
        StructField("lastname",StringType(),True), \
        StructField("id", StringType(), True), \
        StructField("gender", StringType(), True), \
        StructField("salary", IntegerType(), True) \
    ])

    df = ctx.spark_session.createDataFrame(data=data2, schema=schema)

    ####################################

    class AccumulatorNumpyArray(AccumulatorParam):
        def zero(self, zero: np.ndarray):
            return zero

        def addInPlace(self, v1, v2):
            return v1 + v2

    sc = ctx.spark_session.sparkContext

    shape = 3

    global accum
    accum = sc.accumulator(
            np.zeros(shape, dtype=np.int64),
            AccumulatorNumpyArray(),
            )

    def func(row):
        global accum
        accum += np.ones(shape)
        return str(accum) + '_' + str(id(accum))

    user_defined_function = udf(func, StringType())

    new = df.withColumn("processed", user_defined_function(struct([df[col] for col in df.columns])))
    new.show(2)

    print(accum)

    return df

Foundry环境执行结果

DataFrame输出:

+---------+----------+--------+-----+------+------+--------------------+
|firstname|middlename|lastname|   id|gender|salary|           processed|
+---------+----------+--------+-----+------+------+--------------------+
|    James|          |   Smith|36636|     M|  3000|[1. 1. 1.]_140388...|
|  Michael|      Rose|        |40288|     M|  4000|[2. 2. 2.]_140388...|
+---------+----------+--------+-----+------+------+--------------------+
only showing top 2 rows

主脚本输出:

> accum
 Accumulator<id=0, value=[0 0 0]>
> id(accum)
 140574405092256

普通PySpark环境测试代码(移除Foundry模板)

import numpy as np
from pyspark.accumulators import AccumulatorParam
from pyspark.sql.functions import udf, struct
from pyspark.sql.types import StructType, StringType, IntegerType, StructField
from pyspark.sql import SparkSession
from pyspark.context import SparkContext

spark = (
    SparkSession.builder.appName("Python Spark SQL basic example")
    .config("spark.some.config.option", "some-value")
    .getOrCreate()
)

data2 = [
    ("James", "", "Smith", "36636", "M", 3000),
    ("Michael", "Rose", "", "40288", "M", 4000),
    ("Robert", "", "Williams", "42114", "M", 4000),
    ("Maria", "Anne", "Jones", "39192", "F", 4000),
    ("Jen", "Mary", "Brown", "", "F", -1),
]

schema = StructType(
    [
        StructField("firstname", StringType(), True),
        StructField("middlename", StringType(), True),
        StructField("lastname", StringType(), True),
        StructField("id", StringType(), True),
        StructField("gender", StringType(), True),
        StructField("salary", IntegerType(), True),
    ]
)

df = spark.createDataFrame(data=data2, schema=schema)

####################################

class AccumulatorNumpyArray(AccumulatorParam):
    def zero(self, zero: np.ndarray):
        return zero

    def addInPlace(self, v1, v2):
        return v1 + v2

sc = SparkContext.getOrCreate()

shape = 3

global accum
accum = sc.accumulator(
    np.zeros(shape, dtype=np.int64),
    AccumulatorNumpyArray(),
)

def func(row):
    global accum
    accum += np.ones(shape)
    return str(accum) + "_" + str(id(accum))

user_defined_function = udf(func, StringType())

new = df.withColumn(
    "processed", user_defined_function(struct([df[col] for col in df.columns]))
)
new.show(2, False)

print(id(accum))
print(accum)

普通PySpark环境执行结果(符合预期)

+---------+----------+--------+-----+------+------+--------------------------+
|firstname|middlename|lastname|id   |gender|salary|processed                 |
+---------+----------+--------+-----+------+------+--------------------------+
|James    |          |Smith   |36636|M     |3000  |[1. 1. 1.]_139642682452576|
|Michael  |Rose      |        |40288|M     |4000  |[1. 1. 1.]_139642682450224|
+---------+----------+--------+-----+------+------+--------------------------+
only showing top 2 rows

140166944013424
[3. 3. 3.]

问题原因与解决方案

原因分析

Foundry的Transforms框架对PySpark执行环境做了封装,导致累加器对象在Driver和Executor间的序列化/反序列化逻辑与普通PySpark不同:

  1. 普通PySpark中,累加器的更新会通过RPC同步回Driver端的原始对象;但在Foundry里,UDF中使用的累加器是原始对象的序列化副本,更新仅发生在Executor的副本上,Driver端的原始累加器不会同步更新。
  2. 主脚本与UDF内的累加器ID不同,正是因为UDF里的实例是Driver端对象的反序列化副本,二者并非同一个对象。

解决方案

不要依赖主脚本中的累加器变量获取值,而是直接从生成的DataFrame中提取最终的累加器值:

  1. 提取processed列的最后一行数据,解析其中的累加器字符串部分。
  2. 将解析后的字符串转换回numpy数组。

修改Foundry代码中的相关逻辑:

# 原有代码
new = df.withColumn("processed", user_defined_function(struct([df[col] for col in df.columns])))
new.show(2)

print(accum)

# 修改后代码
new = df.withColumn("processed", user_defined_function(struct([df[col] for col in df.columns])))
new.show(2)

# 提取最后一行的累加器值(注意:collect()会将数据拉到Driver,大数据集慎用,可改用分布式方式提取)
last_processed = new.select("processed").orderBy("salary").last()
if last_processed:
    accum_str = last_processed.processed.split('_')[0]
    final_accum_value = np.fromstring(accum_str.strip('[]'), sep=' ')
    print("最终累加器值:", final_accum_value)

替代方案(更规范的Foundry用法)

如果你的核心需求是聚合数据,优先使用PySpark原生聚合函数(如sum、agg),而非自定义累加器+UDF:

  • 原生聚合是分布式安全的,性能更优。
  • 避免Foundry环境中序列化带来的对象不一致问题。

示例代码:

from pyspark.sql.functions import sum as spark_sum, array, lit

# 每行生成全1的数组
df = df.withColumn("ones", array([lit(1)]*3))
# 聚合求和
total = df.agg(spark_sum("ones")).first()[0]
print("最终总和:", total)

内容的提问来源于stack exchange,提问作者JulianKarlBauer

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 10:17:23