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

