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

如何在PySpark中用scipy.stats生成正态分布数组(含UDF方案)

在PySpark中生成指定均值、标准差的正态分布随机数数组

问题原因

你遇到的NameError: 'st'未定义,通常是因为在UDF外部导入了scipy.stats(比如import scipy.stats as st),但PySpark普通UDF的序列化机制无法将外部模块引用传递到Executor节点,导致执行时找不到该模块。


方案一:修复普通UDF

解决核心是在UDF内部导入依赖模块,确保Executor执行时能正确加载:

from pyspark.sql import SparkSession
from pyspark.sql.functions import udf
from pyspark.sql.types import ArrayType, DoubleType

spark = SparkSession.builder.appName("NormalDistGenerator").getOrCreate()

def generate_normal(mean, std, size):
    # 必须在UDF内部导入scipy.stats
    import scipy.stats as st
    return st.norm.rvs(loc=mean, scale=std, size=size).tolist()

# 注册UDF并指定返回类型为Double数组
normal_udf = udf(generate_normal, ArrayType(DoubleType()))

# 测试示例
df = spark.createDataFrame([(0, 1, 5), (2, 0.5, 3)], ["mean", "std", "size"])
df.withColumn("normal_array", normal_udf("mean", "std", "size")).show(truncate=False)

注意:需确保所有Spark Executor节点都安装了scipy和numpy,否则会出现模块缺失错误。


方案二:使用Pandas UDF(性能更优)

如果处理大规模数据,推荐用Pandas UDF,它的执行模式能更好地处理外部模块引用,且性能优于普通UDF:

from pyspark.sql import SparkSession
from pyspark.sql.functions import pandas_udf
from pyspark.sql.types import ArrayType, DoubleType
import scipy.stats as st
import pandas as pd

spark = SparkSession.builder.appName("NormalDistPandasUDF").getOrCreate()

@pandas_udf(ArrayType(DoubleType()))
def generate_normal_pd(mean: pd.Series, std: pd.Series, size: pd.Series) -> pd.Series:
    return pd.Series([
        st.norm.rvs(loc=m, scale=s, size=int(z)).tolist() 
        for m, s, z in zip(mean, std, size)
    ])

# 测试示例
df = spark.createDataFrame([(0, 1, 5), (2, 0.5, 3)], ["mean", "std", "size"])
df.withColumn("normal_array", generate_normal_pd("mean", "std", "size")).show(truncate=False)

方案三:使用Spark内置函数(无第三方依赖)

如果无法在集群安装scipy,可以用Spark原生的randn()生成标准正态分布,再通过线性变换得到目标分布,最后收集为数组:

from pyspark.sql import SparkSession
from pyspark.sql.functions import randn, collect_list, expr
from pyspark.sql.window import Window

spark = SparkSession.builder.appName("NormalDistBuiltin").getOrCreate()

# 构造测试数据
df = spark.createDataFrame([(0, 1, 5), (2, 0.5, 3)], ["mean", "std", "size"])

# 生成唯一ID用于分组
df = df.withColumn("row_id", expr("monotonically_increasing_id()"))

# 按size列值生成重复行
exploded_df = df.withColumn("dummy", expr("explode(array_repeat(1, size))"))

# 将标准正态分布转换为指定均值、标准差的分布
transformed_df = exploded_df.withColumn("normal_val", expr("mean + std * randn()"))

# 分组收集为数组
result_df = transformed_df.groupBy("row_id", "mean", "std", "size")\
    .agg(collect_list("normal_val").alias("normal_array"))\
    .drop("row_id")

result_df.show(truncate=False)

这个方案无需第三方库,兼容性强,但逻辑相对繁琐,适合受限环境。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.25 11:24:24