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

如何在PySpark中按指定概率为每行从多项分布抽样?

为PySpark数据框实现灵活的多项分布抽样

问题背景

我有一组取值及对应的抽样概率,希望给PySpark数据框的每一行从该多项分布中抽样。目前已经能用硬编码随机数的方式实现,但需要一种能适配任意数量取值与概率的灵活方案。

示例取值与概率:

assignment_values = ["foo", "buzz", "boo"]
value_probabilities = [0.3, 0.3, 0.4]

硬编码实现方式(仅适用于固定数量的取值):

from pyspark.sql import Row
import pyspark.sql.functions as F

data = [
    {"person": 1, "company": "5g"},
    {"person": 2, "company": "9s"},
    {"person": 3, "company": "1m"},
    {"person": 4, "company": "3l"},
    {"person": 5, "company": "2k"},
    {"person": 6, "company": "7c"},
    {"person": 7, "company": "3m"},
    {"person": 8, "company": "2p"},
    {"person": 9, "company": "4s"},
    {"person": 10, "company": "8y"},
]
df = spark.createDataFrame(Row(**x) for x in data)

(
    df
    .withColumn("rand", F.rand())
    .withColumn(
        "assignment", 
        F.when(F.col("rand") < F.lit(0.3), "foo")
        .when(F.col("rand") < F.lit(0.6), "buzz")
        .otherwise("boo")
    )
    .show()
)

输出结果:

+-------+------+-------------------+----------+
|company|person|               rand|assignment|
+-------+------+-------------------+----------+
|     5g|     1| 0.8020603266148111|       boo|
|     9s|     2| 0.1297179045352752|       foo|
|     1m|     3|0.05170251723736685|       foo|
|     3l|     4|0.07978240998283603|       foo|
|     2k|     5| 0.5931269297050258|      buzz|
|     7c|     6|0.44673560271164037|      buzz|
|     3m|     7| 0.1398027427612647|       foo|
|     2p|     8| 0.8281404801171598|       boo|
|     4s|     9|0.15568513681001817|       foo|
|     8y|    10| 0.6173220502731542|       boo|
+-------+------+-------------------+----------+

灵活解决方案

思路

  1. 计算累积概率数组:将原始概率依次累加,得到每个取值对应的区间上限(示例中累积概率为[0.3, 0.6, 1.0])
  2. 动态构建抽样逻辑:根据累积概率和取值的对应关系,自动生成when条件链,无需硬编码每个阈值和取值

实现代码

import pyspark.sql.functions as F
from pyspark.sql import Row

# 定义任意数量的取值和概率
assignment_values = ["foo", "buzz", "boo"]
value_probabilities = [0.3, 0.3, 0.4]

# 计算累积概率
cumulative_probs = []
current_sum = 0.0
for prob in value_probabilities:
    current_sum += prob
    cumulative_probs.append(current_sum)

# 初始化抽样条件:默认取最后一个值作为兜底
assignment_expr = F.lit(assignment_values[-1])

# 从倒数第二个取值开始,反向构建when条件链
for val, cum_prob in zip(reversed(assignment_values[:-1]), reversed(cumulative_probs[:-1])):
    assignment_expr = F.when(F.col("rand") < F.lit(cum_prob), val).otherwise(assignment_expr)

# 生成数据框并执行抽样
data = [
    {"person": i+1, "company": f"{i+1}x"} for i in range(10)
]
df = spark.createDataFrame(Row(**x) for x in data)

result_df = df.withColumn("rand", F.rand()).withColumn("assignment", assignment_expr)
result_df.show()

方案优势

  • 完全灵活:不管取值和概率的数量多少,只需修改assignment_values和value_probabilities两个列表即可,无需改动抽样逻辑代码
  • 逻辑精准:通过累积概率匹配随机数区间,和硬编码逻辑完全一致,保证抽样概率的准确性
  • 性能高效:基于PySpark内置函数构建表达式,比自定义UDF的执行效率更高

验证概率准确性(可选)

如果需要验证抽样结果是否符合预期概率,可以统计各取值的出现频率:

result_df.groupBy("assignment").count().withColumn(
    "frequency", F.col("count") / F.sum("count").over()
).show()

输出会接近预设的[0.3, 0.3, 0.4]概率分布。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.15 19:20:29