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

使用torch.save保存Hugging Face ViT模型失败的问题及正确方案咨询

Hugging Face ViT模型的正确保存与加载方法

报错原因说明

你用torch.save保存ViT模型时出现的Stale file handle及后续RuntimeError,大概率和DDP环境下的文件操作、pickle序列化的兼容性问题有关。虽然torch.save理论上可保存模型状态字典,但Hugging Face模型有专门的序列化机制,能更稳定地避免这类问题。

保存ViT模型的正确方式

推荐方式:使用模型自带的save_pretrained()

这是Hugging Face官方推荐的方法,自动保存模型权重和配置文件,无需依赖pickle,兼容性更强:

from torch.nn.parallel import DistributedDataParallel as DDP

# 若模型被DDP包装,先取出原始模型
if isinstance(args.model, DDP):
    model = args.model.module
else:
    model = args.model

# 保存到指定目录,会生成pytorch_model.bin(权重)和config.json(配置)
model.save_pretrained("./vit_saved_model")

如果需要保存训练相关的额外信息(如优化器状态、训练步数),可单独用torch.save保存:

torch.save({
    'it': args.it,
    'epoch_num': args.epoch_num,
    'opt_state_dict': args.opt.state_dict(),
    'scheduler_state_dict': try_to_get_scheduler_state_dict(args.scheduler)
}, "./training_states.pt")

备选方式:保存模型状态字典(兼容torch.save)

如果坚持用torch.save,必须确保保存的是**原始模型(非DDP包装)**的状态字典:

# 取出DDP包装内的原始模型
model = get_model_from_ddp(args.model)
# 保存状态字典及其他训练信息
torch.save({
    'training_mode': args.training_mode,
    'it': args.it,
    'epoch_num': args.epoch_num,
    'args_dict': vars(uutils.make_args_pickable(args)),
    'model_state_dict': model.state_dict(),
    'opt_state_dict': args.opt.state_dict(),
    'scheduler_state_dict': try_to_get_scheduler_state_dict(args.scheduler)
}, args.log_root / ckpt_filename)

加载ViT模型的正确方式

对应save_pretrained的加载方法

直接从保存的目录加载,自动读取权重和配置:

from transformers import ViTForImageClassification

# 加载完整模型
model = ViTForImageClassification.from_pretrained("./vit_saved_model")
# 若需要加载训练状态
training_states = torch.load("./training_states.pt")
args.opt.load_state_dict(training_states['opt_state_dict'])

对应torch.save状态字典的加载方法

先初始化和保存时结构一致的模型,再加载状态字典:

from transformers import ViTForImageClassification

# 初始化模型(需和保存时的模型结构、预训练权重一致)
model = ViTForImageClassification.from_pretrained("google/vit-base-patch16-224")
# 加载状态字典
ckpt = torch.load(args.log_root / ckpt_filename)
model.load_state_dict(ckpt['model_state_dict'])
# 加载其他训练状态
args.it = ckpt['it']
args.epoch_num = ckpt['epoch_num']
args.opt.load_state_dict(ckpt['opt_state_dict'])

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.15 14:20:23