You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch中NamedTuple无法加入safe_globals导致加载失败的问题

问题:PyTorch加载NamedTuple时出现WeightsUnpickler错误

问题详情

我定义了如下NamedTuple类:

class checkpoint_t(NamedTuple):
    epoch: int
    model_state_dict: Dict[str, Any]
    optimizer_state_dict: Dict[str, Any]
    model_name: str | None = None

保存后,尝试通过以下代码加载该NamedTuple:

import torch
from train import checkpoint_t
with torch.serialization.safe_globals([checkpoint_t]):
    print("safe globals: ", torch.serialization.get_safe_globals())
    checkpoint: checkpoint_t = torch.load(parsed_args.checkpoint, weights_only=True)

但依然报错:

WeightsUnpickler error: Unsupported global: GLOBAL __main__.checkpoint_t was not an allowed global by default. Please use torch.serialization.add_safe_globals([checkpoint_t]) or the torch.serialization.safe_globals([checkpoint_t]) context manager to allowlist this global if you trust this class/function.

原因分析

核心问题是类的全局名称不匹配:

  • 当用python -m package.subpackage.train运行训练脚本时,train模块的__name__会被替换为__main__,因此保存checkpoint时,pickle会把checkpoint_t记录为__main__.checkpoint_t。
  • 加载时,你从train模块导入的checkpoint_t,其全局名称是train.checkpoint_t(或完整包路径名),你添加到safe_globals里的是这个版本,但pickle要找的是__main__.checkpoint_t,两者不匹配导致报错。

解决方法

方法1:将NamedTuple移到独立模块(推荐)

把checkpoint_t的定义从train.py移到一个独立的工具模块,确保类的全局名称始终固定:

  1. 在utils/checkpoint.py中定义类:
from typing import NamedTuple, Dict, Any

class checkpoint_t(NamedTuple):
    epoch: int
    model_state_dict: Dict[str, Any]
    optimizer_state_dict: Dict[str, Any]
    model_name: str | None = None
  1. 训练保存时从该模块导入使用:
from utils.checkpoint import checkpoint_t
# ... 训练逻辑,创建checkpoint实例
torch.save(checkpoint_instance, "model_checkpoint.pth")
  1. 加载时同样从该模块导入:
import torch
from utils.checkpoint import checkpoint_t

with torch.serialization.safe_globals([checkpoint_t]):
    checkpoint: checkpoint_t = torch.load(parsed_args.checkpoint, weights_only=True)

方法2:加载时手动映射__main__模块的类

如果不想调整模块结构,可在加载时把train模块的checkpoint_t注册到__main__模块,让pickle能匹配到对应类:

import torch
from train import checkpoint_t

# 将train模块的checkpoint_t赋值给__main__模块
import __main__
__main__.checkpoint_t = checkpoint_t

# 同时添加两个版本到安全列表
with torch.serialization.safe_globals([checkpoint_t, __main__.checkpoint_t]):
    checkpoint: checkpoint_t = torch.load(parsed_args.checkpoint, weights_only=True)

方法3:修改训练脚本的运行方式

避免用python -m package.subpackage.train启动训练,改用以下方式之一,让train模块的__name__保持为实际包路径:

  • 直接运行脚本:python package/subpackage/train.py
  • 使用项目配置的console_scripts入口点启动训练

这样保存时checkpoint_t的全局名称会是package.subpackage.train.checkpoint_t,与加载时导入的类名称完全匹配,添加到safe_globals后即可正常加载。

内容的提问来源于stack exchange,提问作者Wang

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.14 11:34:53