已开启
【26.2.0众测】快速入门示例用 torch.save({'arch': CNN, ...}) 保存权重,换进程 torch.load 会报 AttributeError: Can't get attribute 'CNN' #4781
fpwdh创建于  11 天前
fpwdh
fpwdh
11 天前 创建

问题描述(文档示例可用性)

《快速入门》的示例在训练结束时这样保存权重(原文):

torch.save({
               'epoch': 10,
               'arch': CNN,
               'state_dict': model.state_dict(),
               'optimizer' : optimizer.state_dict(),
            },'checkpoint.pth.tar')

其中 'arch': CNN 把模型类对象一起 pickle 进了 checkpoint。这样保存出来的文件只能在“CNN 这个类可被导入”的进程里加载;换一个进程/脚本去 torch.load 就会失败。

复现步骤

  1. 按《快速入门》跑通 train.py(实测 10 epoch,loss 0.0222,正常生成 checkpoint.pth.tar,125,487 字节);
  2. 在另一个进程里加载:
import torch
ck = torch.load('/root/t5/checkpoint.pth.tar', map_location='cpu', weights_only=False)
  1. 报错:
AttributeError: Can't get attribute 'CNN' on <module '__main__' (<class '_frozen_importlib.BuiltinImporter'>)>

(因为保存时 CNN 属于 __main__,新进程的 __main__ 里没有这个类。)

期望结果

文档示例保存的 checkpoint 可以被独立进程正常加载,用户拿到权重文件就能直接用。

建议

  1. 示例改为只保存可移植内容,例如:
torch.save({
    'epoch': 10,
    'state_dict': model.state_dict(),
    'optimizer': optimizer.state_dict(),
}, 'checkpoint.pth.tar')
  1. 如果确实需要保存模型结构,建议把 CNN 放到可导入的模块里(例如 model.py),或在文档中明确说明
    “加载时必须先 from <模块> import CNN,否则会报 AttributeError”。

(本反馈来自 TorchNPU 26.1.0 众测,材料见 pytorch-ecosystem 仓 01_tasks/2026/torch_npu_public_beta/fpwdh/)

likedislike
TorchNPU-BotTorchNPU-Bot成员
11 天前 添加了label:triage-review
TorchNPU-Bot
TorchNPU-Bot成员
11 天前 评论:

issue待分派,添加triage-review标签

likedislike