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

如何构建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,需严格做到以下三点:

  1. 共享采样器实例:两个DataLoader必须使用同一个Sampler(比如你的DistributedSampler),确保生成的索引序列完全一致;
  2. 关闭随机打乱:若不用分布式采样器,需将两个DataLoader的shuffle参数都设为False;
  3. 同步迭代:训练循环中用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.11 17:35:24