【问题标题】:Why must use DataParallel when testing?测试时为什么一定要使用DataParallel?
【发布时间】:2020-05-18 00:14:15
【问题描述】:

在GPU上训练,num_gpus设置为1:

device_ids = list(range(num_gpus))
model = NestedUNet(opt.num_channel, 2).to(device)
model = nn.DataParallel(model, device_ids=device_ids)

CPU测试:

model = NestedUNet_Purn2(opt.num_channel, 2).to(dev)
device_ids = list(range(num_gpus))
model = torch.nn.DataParallel(model, device_ids=device_ids)
model_old = torch.load(path, map_location=dev)
pretrained_dict = model_old.state_dict()
model_dict = model.state_dict()
pretrained_dict = {k: v for k, v in pretrained_dict.items() if k in model_dict}
model_dict.update(pretrained_dict)
model.load_state_dict(model_dict)

这样会得到正确的结果,但是当我删除时:

device_ids = list(range(num_gpus))
model = torch.nn.DataParallel(model, device_ids=device_ids)

结果错误。

【问题讨论】:

    标签: deep-learning pytorch


    【解决方案1】:

    nn.DataParallel 包装模型,其中实际模型分配给module 属性。这也意味着状态字典中的键具有module. 前缀。

    让我们看一个非常简化的版本,只有一个卷积,看看有什么区别:

    class NestedUNet(nn.Module):
        def __init__(self):
            super().__init__()
            self.conv1 = nn.Conv2d(3, 16, kernel_size=3, padding=1)
    
    model = NestedUNet()
    
    model.state_dict().keys() # => odict_keys(['conv1.weight', 'conv1.bias'])
    
    # Wrap the model in DataParallel
    model_dp = nn.DataParallel(model, device_ids=range(num_gpus))
    
    model_dp.state_dict().keys() # => odict_keys(['module.conv1.weight', 'module.conv1.bias'])
    

    您使用nn.DataParallel 保存的状态字典与常规模型的状态不一致。您正在将当前状态字典与加载状态字典合并,这意味着加载状态被忽略,因为模型没有任何属于键的属性,而是留下随机初始化的模型。

    为了避免犯这个错误,你不应该合并状态字典,而是直接将它应用到模型中,在这种情况下,如果键不匹配就会出错。

    RuntimeError: Error(s) in loading state_dict for NestedUNet:
            Missing key(s) in state_dict: "conv1.weight", "conv1.bias".
            Unexpected key(s) in state_dict: "module.conv1.weight", "module.conv1.bias".
    

    为了让你保存的状态字典兼容,你可以去掉module.前缀:

    pretrained_dict = {key.replace("module.", ""): value for key, value in pretrained_dict.items()}
    model.load_state_dict(pretrained_dict)
    

    您也可以在将来通过在保存其状态之前从nn.DataParallel 解包模型来避免此问题,即保存model.module.state_dict()。因此,如果您想使用多个 GPU,您始终可以先加载模型及其状态,然后再决定将其放入 nn.DataParallel

    【讨论】:

      【解决方案2】:

      您使用 DataParallel 训练了您的模型并保存了它。因此,模型权重以module. 前缀存储。现在,当您在没有DataParallel 的情况下加载时,您基本上不会加载任何模型权重(模型具有随机权重)。结果,模型预测是错误的。

      我举个例子。

      model = nn.Linear(2, 4)
      model = torch.nn.DataParallel(model, device_ids=device_ids)
      model.state_dict().keys() # => odict_keys(['module.weight', 'module.bias'])
      

      另一方面,

      another_model = nn.Linear(2, 4)
      another_model.state_dict().keys() # => odict_keys(['weight', 'bias'])
      

      查看OrderedDict 键的区别。

      因此,在您的代码中,以下三行代码有效,但未加载模型权重。

      pretrained_dict = model_old.state_dict()
      model_dict = model.state_dict()
      pretrained_dict = {k: v for k, v in pretrained_dict.items() if k in model_dict}
      

      这里,model_dict 有没有module. 前缀的键,但pretrained_dict 在你不使用DataParalle 时有。所以,当不使用 DataParallel 时,pretrained_dict 本质上是空的。


      解决方案:如果你想避免使用DataParallel,或者你可以加载权重文件,创建一个不带模块前缀的新OrderedDict,然后加载回来。

      在不使用DataParallel 的情况下,以下内容适用于您的情况。

      # original saved file with DataParallel
      model_old = torch.load(path, map_location=dev)
      
      # create new OrderedDict that does not contain `module.`
      from collections import OrderedDict
      
      new_state_dict = OrderedDict()
      for k, v in model_old.items():
          name = k[7:] # remove `module.`
          new_state_dict[name] = v
      
      # load params
      model.load_state_dict(new_state_dict)
      

      【讨论】:

        猜你喜欢
        • 2015-05-19
        • 1970-01-01
        • 1970-01-01
        • 1970-01-01
        • 1970-01-01
        • 2020-09-30
        • 2011-11-08
        • 2018-04-11
        • 1970-01-01
        相关资源
        最近更新 更多