如何高效为大型DataFrame每个分组内的行随机打标签
优化方案
原代码性能瓶颈分析
- 每个自定义函数内执行
df.group_id == i会全表遍历4000万行做布尔匹配,2000次分组累计产生800亿次比较操作,是最大的性能瓶颈 - 使用
dummy.Pool线程池处理CPU密集型运算,受Python GIL全局解释器锁限制无法实现真正并行,反而额外增加了线程调度开销 - 采用Python原生
random.shuffle和列表操作,执行效率远低于C实现的numpy向量化运算 - 没有预先统计分组信息,重复计算每个分组的长度
最优实现方案(适配你已经对group_id排序的场景)
你已经提前将df按group_id排序,同分组行连续存储,仅需一次遍历即可完成计算,4000万行数据普通消费级CPU即可在数秒内跑完:
import pandas as pd import numpy as np # 你的原有数据生成、排序逻辑保持不变 N = int(4e7) M = int(2e3) + 1 col_1 = np.random.randint(1, M, N) col_2 = np.random.uniform(low = 1, high = 5, size = N) df = pd.DataFrame({'group_id': col_1, 'value': col_2}) df.sort_values(by = 'group_id', inplace = True) df.reset_index(inplace = True, drop = True) # 以下为优化后的打标签逻辑 # 一次统计所有分组的长度 group_sizes = df.groupby('group_id', sort=False).size().values # 预分配结果数组,避免动态扩容开销 batch_arr = np.zeros(len(df), dtype=np.int32) ptr = 0 for size in group_sizes: # 调用numpy原生随机排列,C实现效率极高 batch_arr[ptr:ptr+size] = np.random.permutation(size) + 1 ptr += size df['batch'] = batch_arr
通用场景方案(未提前排序也可使用)
如果你的df没有提前按group_id排序,可以直接用groupby+transform实现,性能也远高于原有多线程方案:
df['batch'] = df.groupby('group_id')['group_id'].transform(lambda x: np.random.permutation(len(x)) + 1)
内容的提问来源于stack exchange,提问作者Akira
相关产品推荐
相关产品推荐

