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

如何在LibTorch中使用collate_fn解决多尺寸图像批处理报错问题

LibTorch可变尺寸图像批处理填充解决方案

错误原因

你当前的报错是因为提前使用了torch::data::transforms::Stack<>转换,该转换会直接尝试堆叠张量,要求所有输入张量尺寸完全一致,而你的数据集图像尺寸不统一,因此批大小大于1时触发错误:

what(): stack expects each tensor to be equal size, but got [3, 1264, 532] at entry 0 and [3, 299, 294] at entry 1

LibTorch的设计与PyTorch不同,没有单独的collate_fn配置项,它将批聚合逻辑作为数据集转换的一部分,通过自定义Batch转换即可实现填充逻辑。

具体实现步骤

1. 自定义填充堆叠转换

首先实现自定义转换类,完成单批样本的尺寸统计、填充、堆叠操作:

#include <torch/torch.h>

// 自定义批转换:将同批次不同尺寸图像填充到统一尺寸后堆叠
struct PadAndStack {
  torch::data::Example<> operator()(std::vector<torch::data::Example<>> examples) {
    // 统计批次内图像的最大高度、宽度(默认图像格式为[C, H, W])
    int max_h = 0, max_w = 0;
    for (auto& ex : examples) {
      max_h = std::max(max_h, ex.data.size(1));
      max_w = std::max(max_w, ex.data.size(2));
    }

    std::vector<torch::Tensor> padded_imgs, padded_targets;
    for (auto& ex : examples) {
      // 计算填充量,pad参数顺序为[左, 右, 上, 下],默认用0填充,可按需修改
      int pad_right = max_w - ex.data.size(2);
      int pad_bottom = max_h - ex.data.size(1);
      auto padded_img = torch::pad(ex.data, {0, pad_right, 0, pad_bottom});
      padded_imgs.push_back(padded_img);
      // 回归任务的target一般为固定长度,直接收集即可,若target尺寸可变按相同逻辑填充
      padded_targets.push_back(ex.target);
    }

    // 堆叠为批次张量返回
    return {
      torch::stack(padded_imgs),
      torch::stack(padded_targets)
    };
  }
};

2. 修改数据集与DataLoader创建逻辑

移除原来的Stack<>转换,改用自定义的批转换:

// 1. 创建原始数据集,不需要提前加Stack转换
auto raw_dataset = MyDataSet(pathToData);

// 2. 绑定自定义批转换,指定批次大小
auto batched_dataset = raw_dataset.map(
  torch::data::transforms::Batch<PadAndStack>(batchSize)
);

// 3. 创建DataLoader,注意此处DataLoader的batch_size要设为1,因为上游已经完成批聚合
auto dataLoader = torch::data::make_data_loader(
    std::move(batched_dataset),
    torch::data::DataLoaderOptions().batch_size(1).workers(numWorkersDataLoader)
);

3. 训练循环适配

训练循环不需要修改,返回的batch.data形状为[batchSize, 3, max_h, max_w],可直接送入模型训练。

可选优化项

  • 若需要固定输入尺寸(如统一缩放填充到224*224),可直接将max_h和max_w设置为固定值,不需要按批次动态计算
  • 可调整填充逻辑为居中填充,将差值平均分配到左右、上下侧,避免图像全部对齐左上角
  • 可修改torch::pad的填充值与填充模式,适配你的数据集特性

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.02 08:45:05