如何为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
相关产品推荐
相关产品推荐

