PyTorch中设置batch_size大于1时VVC图像块预测模型报错的解决问询
问题背景
我正在构建一个神经网络,用于预测VVC(Versatile Video Coding)压缩过程中图像的分区方式。模型输入是YUV420图像的单帧Y分量,使用包含真实块位置和大小的CSV文件进行训练。
输入与真值
- 输入:单帧10位YUV420图像的Y分量
- 真值:包含块位置、大小和额外分区标记的CSV文件
示例(388016_320x480_37.yuv):
示例(388016_320x480_37.csv):
问题描述
我实现了train.py和dataset.py,但在DataLoader中设置batch_size > 1时遇到错误。batch_size=1时模型运行正常,但增大batch_size会导致运行时错误。
代码摘要
以下是我的custom_collate_fn和DataLoader设置的简化版本:
def custom_collate_fn(batch): frames = [item[0] for item in batch] # Y-frame tensors blocks = [item[1] for item in batch] # Block information frames = torch.stack(frames, dim=0) # Stacking frames along batch dimension return frames, blocks dataloader = DataLoader( dataset, batch_size=batch_size, shuffle=True, collate_fn=custom_collate_fn )
观察结果
- 当
batch_size = 1时,训练循环中的blocks_batch是包含单组块数据的列表 - 当
batch_size > 1时,blocks_batch变成列表的列表,导致索引时出错
训练循环片段:
for i, (frame, blocks_batch) in enumerate(dataloader): frame = frame.to(device) # Shape: [batch_size, 1, H, W] blocks = blocks_batch[0] # Works with batch_size=1 but fails with larger sizes
报错信息
File "C:\Users\Administrator\Documents\VVC_fast\test4\train.py", line 91, in <module> loss1 = criterion(out_split, target_split) File "C:\ProgramData\miniconda3\envs\newenv3\lib\site-packages\torch\nn\modules\module.py", line 1518, in _wrapped_call_impl return self._call_impl(*args, **kwargs) File "C:\ProgramData\miniconda3\envs\newenv3\lib\site-packages\torch\nn\modules\module.py", line 1527, in _call_impl return forward_call(*args, **kwargs) File "C:\ProgramData\miniconda3\envs\newenv3\lib\site-packages\torch\nn\modules\loss.py", line 725, in forward return F.binary_cross_entropy_with_logits(input, target, File "C:\ProgramData\miniconda3\envs\newenv3\lib\site-packages\torch\nn\functional.py", line 3193, in binary_cross_entropy_with_logits raise ValueError(f"Target size ({target.size()}) must be the same as input size ({input.size()})") ValueError: Target size (torch.Size([1, 1])) must be the same as input size (torch.Size([2, 1]))
问题分析与解决思路
我来帮你拆解问题本质,再给出几个实用的解决方向:
核心问题根源
直接报错的原因很明确:当batch_size=2时,模型输出是对应整个batch的张量(比如形状[2,1]),但你在训练循环里只取了blocks_batch[0]——也就是第一个样本的真值(形状[1,1]),两者维度完全不匹配,自然触发了尺寸错误。
而更深层的核心问题是:不同图像的VVC分区块数量是可变的,你当前的custom_collate_fn只是把每个样本的块数据简单打包成列表的列表,这种结构没法和带batch维度的模型输出直接对齐,必须针对性处理这些变长的真值数据。
具体解决方法
方案1:拼接所有块数据+记录batch归属(适合逐块预测的模型)
如果你的模型是对每个真实存在的分区块做预测(比如预测每个块的分裂类型),可以把整个batch的所有块数据拼在一起,同时记录每个块属于哪个batch样本,这样就能和模型输出一一对应。
修改custom_collate_fn:
def custom_collate_fn(batch): frames = [] all_blocks = [] batch_indices = [] # 标记每个块属于第几个batch样本 for batch_idx, (frame, blocks) in enumerate(batch): frames.append(frame) # 把当前样本的块数据加入列表 all_blocks.append(blocks) # 给当前样本的每个块添加batch索引 batch_indices.extend([batch_idx] * len(blocks)) # 堆叠帧张量(batch维度在前) frames = torch.stack(frames, dim=0) # 拼接所有块数据成一个大张量 all_blocks = torch.cat(all_blocks, dim=0) batch_indices = torch.tensor(batch_indices, dtype=torch.long) return frames, all_blocks, batch_indices
调整训练循环:
for i, (frame, all_blocks, batch_indices) in enumerate(dataloader): frame = frame.to(device) all_blocks = all_blocks.to(device) batch_indices = batch_indices.to(device) # 模型前向传播 outputs = model(frame) # 假设模型输出是[batch_size, 所有可能块的数量, 预测维度],先展平成[总块数, 预测维度] out_split = outputs.flatten(0, 1) # 只保留对应真实块的预测结果(用batch_indices对齐) out_split = out_split[batch_indices] # 计算损失 loss1 = criterion(out_split, all_blocks)
方案2:用Padding统一块数据长度(适合块数差异不大的场景)
如果你的数据集里不同图像的块数差异不大,可以把每个样本的块数据补到当前batch的最大块数,同时用mask标记哪些是真实块、哪些是补的无效数据。
修改custom_collate_fn:
def custom_collate_fn(batch): frames = [item[0] for item in batch] blocks_list = [item[1] for item in batch] # 找到当前batch中块数最多的样本 max_block_count = max(len(blocks) for blocks in blocks_list) padded_blocks = [] for blocks in blocks_list: # 计算需要补的长度 pad_length = max_block_count - len(blocks) # 补0(如果你的块数据是张量,要保证padding的维度和原数据一致) padding = torch.zeros(pad_length, blocks.shape[1], dtype=blocks.dtype) padded = torch.cat([blocks, padding], dim=0) padded_blocks.append(padded) # 堆叠帧和补全后的块张量 frames = torch.stack(frames, dim=0) padded_blocks = torch.stack(padded_blocks, dim=0) # shape: [batch_size, max_block_count, 块特征维度] return frames, padded_blocks
调整训练循环(需要忽略padding部分的损失):
for i, (frame, padded_blocks) in enumerate(dataloader): frame = frame.to(device) padded_blocks = padded_blocks.to(device) outputs = model(frame) # shape: [batch_size, max_block_count, 预测维度] # 创建mask:标记哪些是真实块(非padding,这里假设padding是全0张量) mask = (padded_blocks.sum(dim=-1) != 0) # 只计算真实块的损失 loss1 = criterion(outputs[mask], padded_blocks[mask])
方案3:把真值转换成固定尺寸的掩码(推荐,从根源解决batch问题)
如果你的模型是对图像的每个位置/每个固定网格块做预测(比如预测每个位置是否是块边界,或者属于哪种块类型),那最好的方式是把CSV里的块信息转换成和输入图像同尺寸的掩码张量。
比如,根据CSV里的块位置和大小,生成一个HxW的张量:每个像素标记对应的块ID、块大小或者分裂类型。这样每个样本的真值张量尺寸都是固定的(和输入Y帧的尺寸匹配),DataLoader可以直接用默认的collate_fn堆叠,batch_size>1时完全不会有问题。
这种方式从根源上消除了变长数据的问题,是最稳妥的方案,推荐优先考虑。
备注:内容来源于stack exchange,提问作者조동건

