如何将字典中模块的参数添加到PyTorch Module的参数集合中?
解决PyTorch字典中子模块参数不被追踪的问题
这个问题我之前也碰到过!PyTorch的nn.Module只会自动追踪注册在官方容器(比如ModuleList、Sequential)或者直接作为类属性赋值的子模块,普通Python字典里的子模块不会被自动纳入参数集合,所以你调用parameters()的时候看不到dict_block里的卷积层参数。
给你两个最直接的解决办法:
解决方案1:改用nn.ModuleDict(推荐)
这是PyTorch官方专门为字典形式存储子模块设计的容器,它会自动注册里面的所有子模块,参数会被正常追踪,用法和普通字典几乎一致:
import torch class MyModule(torch.nn.Module): def __init__(self): super(MyModule, self).__init__() self.conv_0=torch.nn.Conv2d(3,32,3,stride=1,padding=0) self.blocks=torch.nn.ModuleList([ torch.nn.Conv2d(3,32,3,stride=1,padding=0), torch.nn.Conv2d(32,64,3,stride=1,padding=0)]) # 把普通字典替换成ModuleDict self.dict_block=torch.nn.ModuleDict({ "key_1": torch.nn.Conv2d(64,128,3,1,0), "key_2": torch.nn.Conv2d(56,1024,3,1,0) }) if __name__=="__main__": my_module=MyModule() # 现在能看到dict_block里的参数了 for param in my_module.parameters(): print(param.shape)
这样修改后,你直接用my_module.parameters()给优化器传参就行,完全不用手动添加参数组,既保留了字典的键值对访问方式,又让PyTorch自动管理参数。
解决方案2:手动注册子模块
如果一定要保留普通字典的形式,可以在初始化的时候用self.add_module()手动把字典里的每个子模块注册到当前Module中:
import torch class MyModule(torch.nn.Module): def __init__(self): super(MyModule, self).__init__() self.conv_0=torch.nn.Conv2d(3,32,3,stride=1,padding=0) self.blocks=torch.nn.ModuleList([ torch.nn.Conv2d(3,32,3,stride=1,padding=0), torch.nn.Conv2d(32,64,3,stride=1,padding=0)]) self.dict_block={ "key_1": torch.nn.Conv2d(64,128,3,1,0), "key_2": torch.nn.Conv2d(56,1024,3,1,0) } # 遍历字典,逐个注册子模块 for key, module in self.dict_block.items(): # 给每个模块加个唯一名字,避免和其他属性冲突 self.add_module(f"dict_block_{key}", module) if __name__=="__main__": my_module=MyModule() print(my_module.parameters())
add_module会把传入的子模块加入到当前Module的追踪列表里,这样参数就会被包含在parameters()的结果中。不过这种方法不如ModuleDict简洁,毕竟官方容器已经帮你封装好了这些逻辑。
内容的提问来源于stack exchange,提问作者Ash
相关产品推荐
相关产品推荐

