如何在Polars DataFrame中为group_id与date的每个组合随机抽取n个ID并解决抽样重复问题
如何在Polars DataFrame中为group_id与date的每个组合随机抽取n个ID并解决抽样重复问题
我太懂你这个困扰了!之前我也碰到过类似的情况——在Polars里按group_id和date分组抽样n个ID,结果每个组合抽出来的ID居然一模一样,后来才发现是全局种子在搞鬼。既然你还需要结果可复现,那咱们就给每个分组组合配个专属的随机种子,完美解决这个问题!
为什么会出现重复抽样?
如果你之前是这么写的:
# 错误示例:全局种子导致每个分组抽样结果重复 result = df.group_by(["group_id", "date"]).agg( pl.col("id").sample(n=2, seed=42) )
问题就出在这个全局的seed=42上——每个分组在抽样时都会重置到同一个随机状态,抽出来的自然就是完全一样的ID集合。
解决方法1:分组后用自定义函数生成唯一种子
咱们可以给每个group_id+date的组合生成一个唯一的种子,比如把组的键值转成哈希值,这样既保证每个组的随机性独立,又能让结果可复现。
import polars as pl import numpy as np # 先构造个测试数据,模拟你的场景 df = pl.DataFrame({ "group_id": np.repeat([1, 2, 3], 10), "date": np.tile(["2024-01-01", "2024-01-02"], 15), "id": np.arange(30) }) # 定义每个分组的抽样函数 def sample_single_group(group: pl.DataFrame, sample_size: int = 2) -> pl.DataFrame: # 取当前组的group_id和date作为唯一标识 group_key = (group["group_id"][0], group["date"][0]) # 把组键转成哈希值当种子,限制在32位整数范围内避免溢出 unique_seed = hash(group_key) % (2**32) # 用专属种子抽样 return group.sample(sample_size, seed=unique_seed) # 分组应用这个抽样函数 final_result = df.group_by(["group_id", "date"], maintain_order=True).apply(sample_single_group) print(final_result)
解决方法2:用向量化操作替代apply(更高效)
如果你的数据量很大,apply的效率可能不够,咱们可以用Polars的向量化操作,先给每个分组生成种子列,再在分组内打乱抽样:
final_result = ( df # 给每个group_id+date组合生成唯一种子 .with_columns( unique_seed=pl.struct(["group_id", "date"]).hash() % (2**32) ) # 按组合分组 .group_by(["group_id", "date"], maintain_order=True) # 用组内的专属种子打乱id列,再取前n个 .agg( pl.col("id").shuffle(seed=pl.col("unique_seed").first()).head(2) ) # 把抽样出的id展开成行 .explode("id") ) print(final_result)
这个方法全程用Polars的内置向量化函数,比apply快很多,大数据场景下更友好。
小提示
- 哈希种子的时候用
% (2**32)是为了把哈希值限制在Polars支持的整数范围内,避免出现无效种子的问题。 - 不管用哪种方法,只要你的
group_id和date组合不变,生成的种子就不会变,完全能保证结果的可复现性。
备注:内容来源于stack exchange,提问作者pinpss
相关产品推荐
相关产品推荐

