如何在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| +-------+------+-------------------+----------+
灵活解决方案
思路
- 计算累积概率数组:将原始概率依次累加,得到每个取值对应的区间上限(示例中累积概率为
[0.3, 0.6, 1.0]) - 动态构建抽样逻辑:根据累积概率和取值的对应关系,自动生成
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
相关产品推荐
相关产品推荐

