Python带权重随机打乱数组 保持同值分组完整的优化方案问询
数组按组加权打乱的优化实现
原逻辑问题修正
- 原概率计算存在错误:
len(item)/len(groups)的总和不等于1,会触发rng.choice参数校验报错,正确分母应为原数组总长度len(my_array)。 - 分组逻辑冗余:无需手动替换NaN为唯一负数值,可利用原生方法自动完成分组。
优化实现方案
核心思路
- 利用
np.unique的equal_nan=False参数自动完成分组:非NaN的相同值归为同一组,每个NaN自动识别为独立分组,省去手动替换步骤。 - 用指数分布采样生成组排序key的经典trick,替代加权choice实现按组长度加权打乱,效率更高且代码更简洁,完全满足「长度越长的分组越大概率排在靠前位置」的需求。
完整代码实现
import numpy as np rng = np.random.default_rng() # 输入示例数组 my_array = np.array([np.nan, 1, 1, np.nan, np.nan, 2, np.nan, 3, 3, 3, np.nan]) # 一步获取分组值、分组长度 _, group_values, group_counts = np.unique( my_array, return_counts=True, equal_nan=False # 核心参数:每个NaN视为独立分组 ) # 生成排序key实现加权打乱:分组越长,key越小,排序越靠前 group_keys = -rng.uniform(size=len(group_values)) ** (1 / group_counts) # 按key对分组排序后拼接为最终数组 shuffled_array = np.concatenate([np.full(cnt, val) for val, cnt in zip(group_values[np.argsort(group_keys)], group_counts)])
旧版本numpy兼容方案
如果使用的numpy版本低于1.23,不支持equal_nan参数,可增加预处理步骤:
# 预处理:将所有NaN替换为唯一负数值 nan_mask = np.isnan(my_array) my_array[nan_mask] = np.arange(-1, -np.sum(nan_mask)-1, -1) # 执行上述分组、打乱逻辑后,替换回NaN shuffled_array[shuffled_array < 0] = np.nan
该方案无pandas依赖,纯numpy操作的时间复杂度为O(n log n),数组规模越大性能优势越明显。
内容的提问来源于stack exchange,提问作者AkariYukari
相关产品推荐
相关产品推荐

