在PyTorch中使用zip()遍历两个nn.ModuleList构建计算图是否正确?
在PyTorch中使用zip()遍历两个nn.ModuleList构建计算图是否正确?
别担心,你这么做完全是正确的,不会对模型的训练和推理造成任何本质问题!
先给你吃个定心丸:nn.ModuleList的核心作用就是帮PyTorch追踪所有子模块的可训练参数——只要你把自定义的_ConvBlock和_DownsampleBlock都正确添加到了ModuleList里,不管是用一个还是两个ModuleList,PyTorch都会正常注册这些参数,计算图也能按你的逻辑正确构建。
你的代码逻辑非常清晰:
- 在
__init__里,把两组功能不同的模块分别存入两个ModuleList,方便管理不同类型的块; - 在
forward里用zip遍历,先执行卷积块并保存输出做跳连,再执行下采样块,这个流程完全符合你要的网络结构。
至于你看到的torchinfo展示差异,这只是工具的格式显示问题,和模型本身的正确性无关:
- 用两个
ModuleList时,工具会重复展开每个ModuleList的层级,所以看起来模块是交替从两个列表里取的; - 用单个
ModuleList时,所有模块按顺序平铺展示。
你仔细对比两个summary里的每个模块参数数量,会发现完全一致——这说明参数注册没有任何问题,模型的计算逻辑是完全相同的。
如果觉得两个ModuleList的summary看着不够整洁,这里给你一个小优化建议:把每一组「卷积块+下采样块」打包成一个独立的子模块,再把这些子模块放到一个ModuleList里,这样torchinfo的展示会更清晰,同时不影响你的跳连逻辑。示例代码如下:
class _EncoderStep(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.conv_block = _ConvBlock(in_channels, _CONV_BLOCK_LEN, _CONV_BLOCK_GROWTH_RATE) self.downsample_block = _DownsampleBlock(self.conv_block.out_channels, out_channels) def forward(self, x): conv_out = self.conv_block(x) down_out = self.downsample_block(conv_out) return conv_out, down_out # 在主模型的__init__中: def __init__(self, in_channels): super().__init__() self.encoder_steps = nn.ModuleList() for out_channels in _FILTERS: step = _EncoderStep(in_channels, out_channels) self.encoder_steps.append(step) in_channels = out_channels # 对应的forward: def forward(self, x): skip_connections = [] for step in self.encoder_steps: conv_out, x = step(x) skip_connections.append(conv_out)
这个优化只是让模型结构的展示更规整,本质上和你原来的实现是等价的,你可以根据自己的偏好选择。
备注:内容来源于stack exchange,提问作者Ivan Tishchenko
相关产品推荐
相关产品推荐

