[Pytorch进阶技巧(一)] 使用add_module替换部分模型

为什么要用add_module()函数

  1. 某些pytorch项目,需要动态调整结构。比如简单的三层全连接 l 1 , l 2 , l 3 l1, l2, l3 l1,l2,l3,在训练几个epoch后根据loss选择将全连接 l 2 l2 l2替换为其它结构 l 2 ′ l2 l2′。
  2. 使用了别人编写的pytorch代码,希望快速地将模型中的特定结构替换掉而不改动别人的源码。

什么是add_module()函数

具体函数定义直接查看

可以看到,是Module类的成员函数,输入参数为Module.add_module(name: str, module: Module)。功能为,为Module添加一个子module,对应名字为name。

怎么用add_module()函数

现在回忆一下,一般定义模型时,Module A的子module都是在A.init(self)中定义的,比如A中一个卷积子模块self.conv1 = torch.nn.Conv2d(…)。此时,这个卷积模块在A的名字其实是’conv1’。

对比之下,add_module()函数就可以在A.init(self)以外定义A的子模块。如定义同样的卷积子模块,可以通过A.add_module(‘conv1’, torch.nn.Conv2d(…))。

以上是给A添加一个子模块。那么删除就是del A.conv1。而替换同样采用add_module()函数,只要name与被替换模块相同即可完成替换。如之前已经定义了A.conv1,但此时希望将其替换为新的自定义模块NewOne(torch.nn.Module),只需要A.add_module(‘conv1’, NewOne())即可。

注意事项

  1. 如果是替换,只要保证前后(如这里的torch.nn.Conv2d与NewOne)的forward输入输出维度一致,就可以不用改写A.forward()。如果是增减,则需要考虑重写A.forward()。
  2. 如果使用了cuda,并且多卡,需要将model放回cpu后进行结构修改。
model: torch.nn.Module

para_model = torch.nn.DataParallel(model).cuda()

train_or_validate_or_something_else(para_model) 

model.cpu()

model.add_module(conv1, NewOne())

para_model = torch.nn.DataParallel(model).cuda()
经验分享 程序员 微信小程序 职场和发展