PyTorch中NamedTuple无法加入safe_globals导致加载失败的问题
问题详情
我定义了如下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_twas not an allowed global by default. Please usetorch.serialization.add_safe_globals([checkpoint_t])or thetorch.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移到一个独立的工具模块,确保类的全局名称始终固定:
- 在
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
- 训练保存时从该模块导入使用:
from utils.checkpoint import checkpoint_t # ... 训练逻辑,创建checkpoint实例 torch.save(checkpoint_instance, "model_checkpoint.pth")
- 加载时同样从该模块导入:
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

