如何高效打乱numpy三维数组中指定比例沿axis=2的one-hot子数组?
高效实现方案
核心逻辑完全采用numpy向量化操作替代Python层级的嵌套循环,随机采样逻辑使用numpy原生无放回均匀采样,符合均匀分布要求,处理大数组时效率远高于嵌套循环实现。
完整实现代码
import numpy as np def shuffle_subarray(arr: np.ndarray, rate: float) -> np.ndarray: # 不修改原数组,先拷贝 res = arr.copy() # 计算前两维总子数组数量 total_count = res.shape[0] * res.shape[1] # 计算待采样数量,避免为负 sample_count = max(0, int(np.round(total_count * rate))) if sample_count == 0: return res # 前两维展平,方便批量索引 flat_res = res.reshape(-1, res.shape[2]) # 无放回均匀采样选中的子数组索引,保证均匀分布 selected_idx = np.random.choice(total_count, size=sample_count, replace=False) # 沿axis=1(原数组的axis=2)批量打乱选中的子数组 flat_res[selected_idx] = np.random.permuted(flat_res[selected_idx], axis=1) # 恢复原形状返回 return flat_res.reshape(res.shape) # 测试示例 if __name__ == "__main__": arr = np.array([ [ [0, 1, 0], [1, 0, 0], ], [ [0, 0, 1], [0, 1, 0], ], ]) # 打乱25%的子数组 shuffled_arr = shuffle_subarray(arr, rate=0.25) print(shuffled_arr)
说明
- 如果你使用的numpy版本低于1.20没有
np.random.permuted接口,可以替换为如下实现打乱逻辑:
该写法仅对选中的少量子数组做循环,远快于遍历所有子数组的嵌套循环。# 替代np.random.permuted的兼容实现 for idx in selected_idx: np.random.shuffle(flat_res[idx]) np.random.choice无放回采样原生保证每个子数组被选中的概率相等,完全符合均匀分布的要求。
内容的提问来源于stack exchange,提问作者Synthaze
相关产品推荐
相关产品推荐

