如何将safetensors模型转换为ONNX模型?PyTorch场景实操疑问
将safetensors模型转换为ONNX格式的解决方案
核心说明
safetensors是仅存储模型权重的格式,不包含模型架构定义。因此你必须先拥有对应模型的PyTorch结构代码(比如自定义模型类、开源预训练模型的结构),才能完成后续转换。
具体操作步骤
1. 定义/导入模型结构
- 自定义模型:写出和原模型完全一致的PyTorch类定义
import torch import torch.nn as nn class CustomModel(nn.Module): def __init__(self): super().__init__() self.conv = nn.Conv2d(3, 64, kernel_size=3) self.fc = nn.Linear(64 * 30 * 30, 10) def forward(self, x): x = torch.relu(self.conv(x)) x = x.flatten(1) x = self.fc(x) return x - 开源预训练模型(比如Hugging Face模型):直接导入模型结构
from transformers import ViTModel # 初始化空结构,后续加载权重 model = ViTModel.from_pretrained("google/vit-base-patch16-224", state_dict={})
2. 从safetensors加载权重到模型
使用safetensors.torch.load_file直接读取权重字典,再赋值给模型实例:
from safetensors.torch import load_file # 读取safetensors中的权重 weight_dict = load_file("model.safetensors") # 初始化模型 model = CustomModel() # 或上面的ViTModel实例 # 加载权重 model.load_state_dict(weight_dict) model.eval() # 转换前切换到评估模式,避免影响结果
3. 转换为ONNX格式
创建与模型输入维度匹配的示例张量,用torch.onnx.export完成导出:
# 生成示例输入(需和模型实际输入维度一致,比如这里是1张3通道224x224的图片) dummy_input = torch.randn(1, 3, 224, 224) # 导出ONNX文件 torch.onnx.export( model, dummy_input, "converted_model.onnx", opset_version=17, # 根据需求选择opset版本 input_names=["model_input"], output_names=["model_output"], dynamic_axes={"model_input": {0: "batch_size"}, "model_output": {0: "batch_size"}} # 支持动态批量 )
常见问题处理
- 权重与模型结构不匹配:检查模型类的层数、维度、参数名称是否和原模型完全一致
- Hugging Face模型快捷方式:若目录下有模型配置文件,可直接用
model = ViTModel.from_pretrained("./", weight_name="model.safetensors")加载
内容的提问来源于stack exchange,提问作者AndreaF
相关产品推荐
相关产品推荐

