如何构建PyTorch DataLoader适配CNN与Encoder的1:10多输入分类任务
关于CNN+Encoder组合分类模型的数据匹配与多DataLoader同步问题
我现在在做一个CNN与Encoder结合的分类任务,两者的输入数据规模是1:10——每一步里CNN接收1份图像数据,Encoder接收10份序列数据,模型最终输出10份分类结果。我不确定是否需要重复CNN的输入数据来匹配规模(毕竟单份CNN输入要对应10份输出),目前我的Dataset和DataLoader代码如下:
当前Dataset代码
def dataset(x_train ,y_train, x_eval, y_eval, x_image_train, x_image_eval): print("TensorDataset") # encoder input and model labels x_train = torch.from_numpy(x_train.astype(np.float32)) y_train = torch.from_numpy(y_train.astype(np.float32)) x_eval = torch.from_numpy(x_eval.astype(np.float32)) y_eval = torch.from_numpy(y_eval.astype(np.float32)) # CNN input x_image_train = torch.from_numpy(x_image_train.astype(np.float32)) x_image_eval = torch.from_numpy(x_image_eval.astype(np.float32)) train_data = torch.utils.data.TensorDataset(x_train, x_image_train, y_train) eval_data = torch.utils.data.TensorDataset(x_eval,x_image_eval, y_eval) return train_data, eval_data
当前DataLoader代码
train_sampler = torch.utils.data.distributed.DistributedSampler(train_data) train_batch_sampler = torch.utils.data.BatchSampler(train_sampler, batch_size, drop_last=True) train_loader = torch.utils.data.DataLoader(train_data, batch_sampler=train_batch_sampler, pin_memory=True, num_workers=nw)
另外,如果用两个DataLoader分别给CNN和Encoder喂数据,怎么保证两者的输入顺序一致?
一、是否需要重复CNN数据?
必须要,因为你的模型逻辑是单份CNN特征要对应10份Encoder输入和输出。当前代码直接将规模不匹配的x_train(10份)和x_image_train(1份)打包进TensorDataset,会直接触发样本数不匹配的错误,无法正常运行。提供两种可行方案:
方案1:预重复CNN数据(小数据场景)
在转换张量阶段,直接将CNN输入重复10次,匹配Encoder输入的样本规模:
def dataset(x_train ,y_train, x_eval, y_eval, x_image_train, x_image_eval): print("TensorDataset") # encoder input and model labels x_train = torch.from_numpy(x_train.astype(np.float32)) y_train = torch.from_numpy(y_train.astype(np.float32)) x_eval = torch.from_numpy(x_eval.astype(np.float32)) y_eval = torch.from_numpy(y_eval.astype(np.float32)) # CNN input:重复10次匹配Encoder规模 x_image_train = torch.from_numpy(x_image_train.astype(np.float32)) # 假设原shape为[N, C, H, W],重复后变为[N*10, C, H, W] x_image_train = x_image_train.repeat_interleave(10, dim=0) x_image_eval = torch.from_numpy(x_image_eval.astype(np.float32)) x_image_eval = x_image_eval.repeat_interleave(10, dim=0) train_data = torch.utils.data.TensorDataset(x_train, x_image_train, y_train) eval_data = torch.utils.data.TensorDataset(x_eval,x_image_eval, y_eval) return train_data, eval_data
方案2:自定义Dataset动态匹配(大数据场景)
如果数据量过大,预重复会占用过多内存,可自定义Dataset在取样本时动态匹配对应CNN数据:
class CustomDataset(torch.utils.data.Dataset): def __init__(self, encoder_data, cnn_data, labels): self.encoder_data = encoder_data self.cnn_data = cnn_data self.labels = labels self.ratio = 10 # 校验数据规模是否符合1:10 assert len(encoder_data) == len(cnn_data) * self.ratio def __len__(self): return len(self.encoder_data) def __getitem__(self, idx): # 根据Encoder样本索引找到对应CNN样本 cnn_idx = idx // self.ratio return self.encoder_data[idx], self.cnn_data[cnn_idx], self.labels[idx] # 构建Dataset def dataset(x_train ,y_train, x_eval, y_eval, x_image_train, x_image_eval): print("CustomDataset") x_train = torch.from_numpy(x_train.astype(np.float32)) y_train = torch.from_numpy(y_train.astype(np.float32)) x_eval = torch.from_numpy(x_eval.astype(np.float32)) y_eval = torch.from_numpy(y_eval.astype(np.float32)) x_image_train = torch.from_numpy(x_image_train.astype(np.float32)) x_image_eval = torch.from_numpy(x_image_eval.astype(np.float32)) train_data = CustomDataset(x_train, x_image_train, y_train) eval_data = CustomDataset(x_eval, x_image_eval, y_eval) return train_data, eval_data
二、双DataLoader如何保证输入顺序一致?
如果必须用两个独立DataLoader,需严格做到以下三点:
- 共享采样器实例:两个DataLoader必须使用同一个
Sampler(比如你的DistributedSampler),确保生成的索引序列完全一致; - 关闭随机打乱:若不用分布式采样器,需将两个DataLoader的
shuffle参数都设为False; - 同步迭代:训练循环中用
zip同时遍历两个DataLoader,保证每次取到的是对应批次。
示例代码:
# 共享同一个采样器 train_sampler = torch.utils.data.distributed.DistributedSampler(train_encoder_data) train_batch_sampler = torch.utils.data.BatchSampler(train_sampler, batch_size, drop_last=True) # 编码器DataLoader encoder_loader = torch.utils.data.DataLoader(train_encoder_data, batch_sampler=train_batch_sampler, pin_memory=True, num_workers=nw) # CNN DataLoader:需提前保证样本数与Encoder匹配 cnn_loader = torch.utils.data.DataLoader(train_cnn_data, batch_sampler=train_batch_sampler, pin_memory=True, num_workers=nw) # 同步迭代取数 for (encoder_x, encoder_y), cnn_x in zip(encoder_loader, cnn_loader): # 模型前向传播等操作 pass
注意:双DataLoader方案容易出现同步误差,优先推荐将数据打包进同一个Dataset的方案。
内容的提问来源于stack exchange,提问作者Slowpoke_james
相关产品推荐
相关产品推荐

