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

如何为PyTorch数据集设置可调节的打乱幅度参数?

轻度打乱PyTorch数据集(无重复)

你之前的方法会出现数据重复,原因是用torch.where对每个位置独立选择原数据或全打乱后的数据,这会导致同一个元素可能被多次选中,最终出现重复。下面是两种无重复、可通过参数p控制打乱程度的实现方案:

方案一:部分元素随机重排

这种方法会随机挑选p*len(X)个元素,仅对这部分元素进行打乱重排,其余元素保持原位置,既保证无重复,又能精准控制打乱程度:

import torch

def mild_shuffle(X, p):
    n = len(X)
    if p <= 0:
        return X.clone()
    if p >= 1:
        return X[torch.randperm(n)]
    
    # 确定要打乱的元素数量
    shuffle_count = int(p * n)
    # 随机选出shuffle_count个不重复的索引
    selected_idx = torch.randperm(n)[:shuffle_count]
    # 对选中的索引进行打乱
    shuffled_selected = selected_idx[torch.randperm(shuffle_count)]
    
    # 复制原数据,避免修改原张量
    shuffled_X = X.clone()
    # 将选中位置的元素替换为打乱后的对应元素
    shuffled_X[selected_idx] = X[shuffled_selected]
    return shuffled_X

方案二:随机交换元素对

通过随机交换若干组元素对来实现轻度打乱,交换次数由p控制,p=1时交换次数接近n/2,效果近似完全打乱:

import torch

def mild_shuffle_swap(X, p):
    n = len(X)
    if p <= 0:
        return X.clone()
    if p >= 1:
        return X[torch.randperm(n)]
    
    # 计算交换次数,p=1时约交换n/2次
    swap_times = int(p * n // 2)
    shuffled_X = X.clone()
    
    for _ in range(swap_times):
        # 随机选出两个不同的索引
        idx1, idx2 = torch.randint(0, n, (2,))
        # 交换对应位置的元素
        shuffled_X[idx1], shuffled_X[idx2] = shuffled_X[idx2].clone(), shuffled_X[idx1].clone()
    
    return shuffled_X

效果说明

  • 当p=0时,直接返回原数据集,无任何打乱
  • 当p=1时,等同于执行完全打乱操作
  • 当0<p<1时,仅对部分元素进行位置调整,所有数据点唯一无重复

内容的提问来源于stack exchange,提问作者Tiana Johnson

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 03:45:30