Pytorch卷积网络序列帧预测排除边框计算损失及相关问题咨询
1 损失计算前裁剪边框的正确实现
你原有代码存在两个核心逻辑问题:
- 频繁做张量、PIL、OpenCV的格式转换,直接打断了PyTorch的计算图,导致梯度无法正常回传
- 裁剪逻辑写反了:你用模型输出
imgs_pred裁剪得到所谓的真值croppedImages_gt,用真值imgs_gt裁剪得到所谓的预测结果croppedImages_pred,最终计算损失还使用了未裁剪的原始张量,裁剪操作完全没有生效
你不需要做格式转换,直接用PyTorch原生张量操作即可完成裁剪+缩放的全流程,全程保留计算图:
import torch import torch.nn.functional as F imgs_pred = model(imgs_input) # 复用裁剪逻辑,输入输出均为[N, C, H, W]格式的张量 def crop_border(imgs_batch): # 缩放到76*76 resized = F.interpolate(imgs_batch, size=(76,76), mode='bilinear', align_corners=False) # 裁剪上下左右各6像素,对应你原来的[6:70,6:70] cropped = resized[:, :, 6:70, 6:70] # 缩放到目标尺寸256*448 final = F.interpolate(cropped, size=(256,448), mode='bilinear', align_corners=False) return final cropped_pred = crop_border(imgs_pred) cropped_gt = crop_border(imgs_gt) # 直接用裁剪后的张量计算损失 loss = criterion_mse(cropped_pred, cropped_gt)
2 梯度报错与内存不足问题原因
- 梯度报错:你做格式转换的时候,原有张量和模型输出的关联被切断,转换后得到的新张量没有梯度回传路径
grad_fn,所以会触发报错。你强行加requires_grad=True相当于新建了独立的可训练变量,和模型参数完全没有关联,梯度根本传不到模型里,属于无效操作 - 内存不足:你手动给所有中间裁剪张量加了梯度属性,这些张量本来不需要存储梯度信息,额外的梯度存储占用了大量显存/内存,最终触发OOM
- 额外说明:PyTorch 0.4版本之后
Variable已经被弃用,张量本身就支持梯度属性,不需要再用Variable封装
用上面的纯张量裁剪方案可以完全解决这两个问题,不需要手动设置梯度属性,计算图全程连贯,也不会有额外的内存开销。
3 .tar格式数据集构建与读取方案
小文件零散存储的IO开销是读取效率低的核心原因,把数据集打包成tar格式可以大幅减少随机IO次数,提升读取效率,操作流程如下:
- 构建tar包:直接用系统tar命令保留原有目录结构打包即可,不需要做特殊处理,比如打包训练集执行命令
tar -cf train.tar train/ - 读取tar包:用PyTorch官方推荐的WebDataset工具读取,示例逻辑如下:
import webdataset as wds from torchvision.transforms import ToTensor dataset = ( wds.WebDataset("train.tar") .decode("pil") # 解码图片为PIL格式 # 按你的序列命名规则调整匹配规则,比如6帧命名为0.jpg~5.jpg就填"0.jpg;1.jpg;2.jpg;3.jpg;4.jpg;5.jpg" .to_tuple("0.jpg", "1.jpg", "2.jpg", "3.jpg", "4.jpg", "5.jpg") .map(lambda frames: torch.stack([ToTensor()(f) for f in frames])) ) # 正常用DataLoader加载即可 dataloader = torch.utils.data.DataLoader(dataset, batch_size=8, num_workers=4)
该方案相比直接读取零散小文件,读取效率通常可以提升3~10倍。
内容的提问来源于stack exchange,提问作者Pia Lüdemann
相关产品推荐
相关产品推荐

