如何将Spark DataFrame某列的随机样本添加至另一列?
解决方案
步骤说明
要给每行target_id生成对应随机样本列,可通过广播全局ID列表+自定义UDF的方式实现,兼顾分布式计算效率与抽样灵活性。
完整代码实现
from pyspark.sql.types import IntegerType, ArrayType from pyspark.sql.functions import udf, broadcast import random # 初始化用户提供的DataFrame target_id = [3733345, 3725312, 3717114, 3408996, 3354970] test_df = spark.createDataFrame(target_id, IntegerType()).withColumnRenamed("value", "target_id") # 收集所有target_id并广播,避免分布式计算中重复传输数据 all_target_ids = test_df.select("target_id").rdd.flatMap(lambda x: x).collect() broadcast_ids = spark.sparkContext.broadcast(all_target_ids) # 定义自定义UDF:生成排除当前id的随机样本 def get_random_samples(current_id, sample_size=5): # 过滤掉当前行的id,避免抽到自身 available_ids = [id for id in broadcast_ids.value if id != current_id] # 抽样,若可用id数量不足指定大小则取全部 return random.sample(available_ids, min(sample_size, len(available_ids))) # 注册UDF,指定返回类型为整数数组 sample_udf = udf(get_random_samples, ArrayType(IntegerType())) # 生成新列并查看结果 result_df = test_df.withColumn("other_target_ids", sample_udf(test_df["target_id"])) result_df.show(truncate=True)
关键细节
- 若不需要排除当前
target_id,直接删除available_ids的过滤逻辑即可 sample_size参数可按需调整抽样数量- 广播变量的使用能大幅提升分布式场景下的计算效率,避免每个Task重复加载全局ID列表
- 由于是随机抽样,每次运行结果会有差异,符合随机样本的需求
内容的提问来源于stack exchange,提问作者Arshad
相关产品推荐
相关产品推荐

