使用nn.ModuleList时未传入优化器的权重也被训练的原因及解决方法
现象原因
该问题和nn.ModuleList本身的功能无关,是初始化ModuleList时的写法错误导致的:
使用[nn.Linear(3,3)]*2创建列表时,程序只会实例化1个nn.Linear对象,之后将该对象的内存引用复制2次存入列表,这就导致layers[0]和layers[1]实际指向同一个模块实例,二者的权重、偏置参数共享同一块内存地址。
你虽然仅将layers[1]的参数传入了优化器,但由于layers[0]的参数和layers[1]完全绑定,优化器更新layers[1]参数时,layers[0]的参数也会同步变化,所以才会出现未传入优化器的层也被训练的假象。
解决方法
初始化nn.ModuleList时,要为每个位置单独实例化子模块,不要用列表乘倍数的写法复用同一个实例引用,正确写法有两种:
- 显式逐个声明子模块
self.layers = nn.ModuleList([nn.Linear(3,3), nn.Linear(3,3)])
- 用推导式生成独立实例
self.layers = nn.ModuleList([nn.Linear(3,3) for _ in range(2)])
修改后每个层都是独立的实例,参数互不干扰,此时仅将指定层的参数传入优化器,就只会更新对应层的权重,不会影响其他未传入优化器的层。
内容的提问来源于stack exchange,提问作者Moran Reznik
相关产品推荐
相关产品推荐

