PyTorch实现多并行图像编码器参数未识别问题排查
现有可生成输出特征的图像编码器,需要将输入图像拆分为若干图像块(约16个),每个图像块输入独立、参数互不共享的图像编码器。
- 初始处理流程:
Input -> Encoder -> Output - 目标改造后流程:
patch-input1 -> Encoder1 -> output1 patch-input2 -> Encoder2 -> output2 ... patch-inputN -> EncoderN -> outputN
实现时基于nn.Module构建模型类,由于图像块数量N不固定,在模型初始化阶段动态确定,因此最初在__init__()方法中通过普通Python列表存储多个编码器实例,在forward函数中循环调用对应编码器处理各图像块输入。
样本输入推理时未触发任何报错,但使用torchsummary模块统计可训练参数时发现异常:
- 单编码器管线的参数量约为100万
- 多编码器管线的参数量仅约20万
- 参数统计结果中未展示编码器各层的结构信息,编码器对应的参数完全未被统计
原实现代码片段如下:
class Patch(nn.Module): ''' takes image tensor input returns a list of patch tensors ''' class Encoder(nn.Module): ''' definition of the encoder ''' class Model(nn.Module): def __init__(self, patches=10, *kwargs): super().__init__() self.patch = Patch() self.enc = [] for i in range(patches): enc.append(Encoder()) def forward(self, x): ''' patches is a list of tensors formed using an image tensor ''' patches = self.patch(x) output = [] for i in range(patches): output.append(self.enc[i](patch[i])) output_feats = torch.cat(output, dim=0)
咨询问题:当前实现存在什么问题?是否有更规范合理的并行编码器实现方式?
核心问题是普通Python列表存储的nn.Module子类实例不会被PyTorch的模块管理机制识别:
PyTorch的参数追踪、设备迁移、权重存加载逻辑,只会识别两类子模块:一类是直接赋值给模型类属性的nn.Module实例,另一类是存放在nn.ModuleList、nn.ModuleDict这类官方提供的专用模块容器中的实例。你用普通Python列表存放编码器,这些编码器不会被注册为当前模型的子模块,直接导致三个问题:
- 调用
.parameters()、.named_parameters()时不会返回这些编码器的参数,因此torchsummary统计参数量时完全漏掉这部分,也不会展示编码器的层结构 - 调用
.to("cuda")、.to("cpu")迁移模型设备时,列表里的编码器不会自动跟随迁移,后续输入张量和模型参数不在同一设备时会触发报错 - 调用
torch.save()保存模型、load_state_dict()加载权重时,这部分编码器的参数不会被包含在状态字典中,保存加载都会丢权重。
你贴的示例代码里还有两处明显笔误,如果你本地运行推理没报错,说明实际运行的代码已经修正了这部分:
__init__循环中写的是enc.append(Encoder()),但你定义的列表属性是self.enc,直接写enc会触发未定义变量报错- forward循环中用
range(patches)遍历,这里的patches是初始化传入的整数参数,不是切分得到的图像块列表;且取输入时写的是patch[i],和前面定义的切分结果变量名patches不一致。
直接用PyTorch官方提供的nn.ModuleList替换普通Python列表即可,这个容器就是专门为存储动态数量的子模块设计的,会自动完成子模块注册,完全兼容PyTorch全量生态工具。修正后的可运行代码如下:
import torch import torch.nn as nn class Patch(nn.Module): ''' 输入图像张量,返回切分后的图像块张量列表 ''' class Encoder(nn.Module): ''' 图像编码器结构定义 ''' class Model(nn.Module): def __init__(self, patches=10, *kwargs): super().__init__() self.patch = Patch() # 替换普通列表为ModuleList,自动注册所有独立编码器 self.enc = nn.ModuleList() for i in range(patches): self.enc.append(Encoder()) def forward(self, x): patches = self.patch(x) output = [] # 修正循环逻辑与变量名 for patch, encoder in zip(patches, self.enc): output.append(encoder(patch)) output_feats = torch.cat(output, dim=0) return output_feats
这种实现下每个编码器参数完全独立、互不共享,torchsummary可以正常统计到全部编码器的参数量(总参数和单编码器参数量×编码器个数的量级匹配),模型设备迁移、权重保存加载也都能正常运行。如果后续需要做推理加速,也可以基于这个结构调整批处理逻辑,不需要改动模块注册部分。
内容的提问来源于stack exchange,提问作者NitishJaiswal

