PyTorch如何为特定批次应用变换?多worker下idx异常解惑
问题解答
1. 为何使用num_workers时idx的表现异常?
当num_workers > 0时,PyTorch会启动多个子进程并行加载数据。每个子进程独立执行Dataset.__getitem__方法,它们的print输出会直接同步到主进程控制台,导致不同进程的打印内容相互交织,看起来顺序混乱。
你看到的奇怪数字(比如964、57)是多进程输出的干扰,同时因为数据是并行预取的,子进程会提前读取后续批次的idx,所以idx的打印顺序和主进程迭代批次的顺序完全不一致。但实际上,DataLoader返回的批次数据是正确的,只是打印输出混乱了。
2. 如何为特定批次(或特定idx对应的数据)应用变换?
针对最后一批(或指定批次)的方案
最可靠的方式是在主进程迭代DataLoader时跟踪批次序号,判断是否是目标批次后再应用变换,不受多进程加载的影响:
import torch class test(torch.utils.data.Dataset): def __init__(self): self.source = [i for i in range(10)] def __len__(self): return len(self.source) def __getitem__(self, idx): return self.source[idx] ds = test() batch_size = 3 dl = torch.utils.data.DataLoader(dataset=ds, batch_size=batch_size, shuffle=False, num_workers=5) # 计算总批次数量 total_batches = len(ds) // batch_size if len(ds) % batch_size != 0: total_batches += 1 for batch_idx, batch_data in enumerate(dl, 1): if batch_idx == total_batches: # 对最后一批应用自定义变换,示例:乘以2 transformed_batch = batch_data * 2 print("最后一批变换后:", transformed_batch) else: print("普通批次:", batch_data)
针对特定idx的方案
如果需要对单个/特定几个idx对应的数据应用变换,可以直接在Dataset的__getitem__里判断,虽然多进程下打印混乱,但数据映射是准确的:
import torch class test(torch.utils.data.Dataset): def __init__(self): self.source = [i for i in range(10)] # 定义需要变换的idx集合 self.target_ids = {9} def __len__(self): return len(self.source) def __getitem__(self, idx): data = self.source[idx] if idx in self.target_ids: # 应用变换,示例:加10 data += 10 return data ds = test() dl = torch.utils.data.DataLoader(dataset=ds, batch_size=3, shuffle=False, num_workers=5) for batch in dl: print(batch)
内容的提问来源于stack exchange,提问作者Jake
相关产品推荐
相关产品推荐

