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

如何将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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 22:49:51