自定义RNN类使用DataParallel报错:'DataParallel'对象无'init_hidden'属性
这个错误我太熟悉了!当你用nn.DataParallel包装自定义RNN模型后,原来的模型实例会被封装到DataParallel对象的module属性中——直接调用包装后的模型的init_hidden()自然会找不到方法,因为这个方法是你自定义RNN类的成员,而非DataParallel类的。
下面给你两种可行的解决思路:
思路一:直接访问module属性调用方法
这是最简单直接的方式,初始化隐藏层时,不要直接调用包装后的模型,而是调用它的module属性对应的方法:
# 假设你原来的代码是这样的: model = RNN(input_size, hidden_size, ...) model = nn.DataParallel(model) # 现在改成这样调用init_hidden: hidden = model.module.init_hidden(batch_size)
这样就能直接调用到你自定义RNN类里的init_hidden方法,完美避开DataParallel的封装问题。
思路二:让DataParallel支持init_hidden调用(更优雅)
如果你不想每次都写.module,可以通过动态绑定或自定义子类的方式,让DataParallel对象也能直接调用init_hidden:
方式1:动态绑定方法
在包装模型后,给DataParallel实例绑定你的init_hidden方法:
model = RNN(input_size, hidden_size, ...) model = nn.DataParallel(model) # 定义适配DataParallel的init_hidden方法 def dp_init_hidden(self, batch_size): return self.module.init_hidden(batch_size) # 给DataParallel实例动态添加方法 import types model.init_hidden = types.MethodType(dp_init_hidden, model) # 现在可以直接调用了: hidden = model.init_hidden(batch_size)
方式2:自定义DataParallel子类
如果经常需要这个功能,不如自己写一个支持init_hidden的DataParallel子类:
class MyDataParallel(nn.DataParallel): def __init__(self, module, device_ids=None, output_device=None, dim=0): super().__init__(module, device_ids, output_device, dim) def init_hidden(self, batch_size): return self.module.init_hidden(batch_size) # 使用时用这个子类包装模型: model = RNN(input_size, hidden_size, ...) model = MyDataParallel(model) # 直接调用init_hidden: hidden = model.init_hidden(batch_size)
额外提醒:隐藏层的设备同步
另外要注意,init_hidden返回的张量要和模型在同一设备上。如果模型被DataParallel分配到多个GPU,建议在RNN类的init_hidden方法里自动匹配设备:
def init_hidden(self, batch_size): # 获取模型参数所在的设备 device = next(self.parameters()).device return (torch.zeros(self.num_layers, batch_size, self.hidden_size).to(device), torch.zeros(self.num_layers, batch_size, self.hidden_size).to(device))
这样不管模型在哪个设备上,隐藏层张量都会自动对应,避免出现设备不匹配的错误。
内容的提问来源于stack exchange,提问作者 hidden_size
相关产品推荐
相关产品推荐

