【问题标题】:Load pytorch model from 0.4.1 to 0.4.0?将pytorch模型从0.4.1加载到0.4.0?
【发布时间】:2019-05-09 17:28:06
【问题描述】:

我使用 pytorch 0.4.1 (GPU) 训练了 DENSENET161 模型,在测试环境中我必须在 pytorch 版本 0.4.0 (CPU) 中加载它。我已经在使用model.cpu() 但是当我加载静态字典model.load_state_dict(checkpoint['state_dict'])

我收到以下错误:

RuntimeError:为 DenseNet 加载 state_dict 时出错:意外 state_dict 中的键:“features.norm0.num_batches_tracked”, “features.denseblock1.denselayer1.norm1.num_batches_tracked”, “features.denseblock1.denselayer1.norm2.num_batches_tracked”, "features.denseblock1.denselayer2.norm1.num_batches_tracked",...

【问题讨论】:

    标签: python deep-learning pytorch


    【解决方案1】:

    这似乎源于 PyTorch 0.4.1 和 0.4 之间规范化层实现的差异 - 前者跟踪一些名为 num_batches_tracked 的状态变量,而 pytorch 0.4 没有预料到。假设只有意外的键并且没有丢失的键(我无法确定,因为你已经剪掉了错误消息),你可以删除无关的键,希望模型能够加载。因此尝试

    model_dict = checkpoint['state_dict']
    filtered = {
        k: v for k, v in model_dict.items() if 'num_batches_tracked' not in k
    }
    model.load_state_dict(filtered)
    

    请注意,除了您在此处看到的内容外,规范化的内部结构可能已经发生了变化,因此即使此修复程序抑制了异常,该模型仍可能默默地行为不端。

    【讨论】:

    • 实际上我在加载 state_dict() 时使用了 strict=False 它解决了这个问题。
    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2021-08-14
    • 2020-05-24
    • 2021-12-15
    • 2019-09-26
    • 2021-10-09
    相关资源
    最近更新 更多