保存google/vit-base-patch16-224-in21k模型遇AttributeError错误求助
解决PyTorch版本过低导致的
trunc_normal_属性错误 问题原因
你遇到的AttributeError: module 'torch.nn.init' has no attribute 'trunc_normal_'错误,本质是当前SageMaker内核的PyTorch版本低于1.7.0——torch.nn.init.trunc_normal_是PyTorch 1.7.0才新增的初始化API,仅调整transformers库版本无法解决底层PyTorch的API缺失问题。
解决方案
方案一:升级PyTorch版本
- 先检查当前环境的PyTorch版本:
import torch print(torch.__version__)
- 升级到1.7.0及以上版本,执行以下命令:
pip install torch>=1.7.0 torchvision torchaudio --upgrade
注:如果是SageMaker官方内核,升级前建议确认依赖兼容性,也可以创建独立的conda环境避免影响原有配置。
方案二:添加兼容的trunc_normal_实现(无需升级PyTorch)
在加载模型前,手动实现trunc_normal_并挂载到torch.nn.init模块:
import torch import math import os def trunc_normal_(tensor, mean=0., std=1., a=-2., b=2.): def norm_cdf(x): return (1. + math.erf(x / math.sqrt(2.))) / 2. with torch.no_grad(): l = norm_cdf((a - mean) / std) u = norm_cdf((b - mean) / std) tensor.uniform_(2 * l - 1, 2 * u - 1) tensor.erfinv_() tensor.mul_(std * math.sqrt(2.)) tensor.add_(mean) tensor.clamp_(min=a, max=b) return tensor # 给旧版本PyTorch补全API if not hasattr(torch.nn.init, 'trunc_normal_'): torch.nn.init.trunc_normal_ = trunc_normal_ # 后续正常加载并保存模型 from transformers import AutoImageProcessor, ViTModel processor = AutoImageProcessor.from_pretrained("google/vit-base-patch16-224-in21k") model = ViTModel.from_pretrained("google/vit-base-patch16-224-in21k") model_path = "model/" if not os.path.exists(model_path): os.mkdir(model_path) model.save_pretrained(save_directory=model_path) processor.save_pretrained(save_directory=model_path)
方案三:更换SageMaker内核
选择SageMaker提供的更新版本PyTorch内核,比如conda_pytorch_p38或指定高版本的内核(如PyTorch 1.12+),这类内核自带的PyTorch版本满足ViT模型的依赖要求,无需额外调整。
额外注意事项
- 部署SageMaker模型时,必须保存
AutoImageProcessor(取消你代码中对应行的注释),否则推理阶段无法正确预处理输入图像。 - 上传到S3时,确保将整个
model/目录完整上传,包含模型权重、配置文件及processor的相关文件。
内容的提问来源于stack exchange,提问作者iamabhaykmr
相关产品推荐
相关产品推荐

