You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何高效打乱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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.29 14:54:02