import torch
import torch.onnx
from model import MyModel
model = MyModel()
model.load_state_dict(torch.load("model.pth"))
model.eval()
dummy_input = torch.randn(1, 3, 224, 224)
torch.onnx.export(
model,
dummy_input,
"model.onnx",
input_names=["input"],
output_names=["output"],
dynamic_axes={"input": {0: "batch_size"}, "output": {0: "batch_size"}},
opset_version=12
)model.onnx 就是模型部署文件import torch
model.eval()
example_input = torch.randn(1, 3, 224, 224)
traced_script = torch.jit.trace(model, example_input)
traced_script.save("model.pt")scripted = torch.jit.script(model)
scripted.save("model.pt")model.save("saved_model/")导出结构:
saved_model/
├── saved_model.pb
└── variables/PyTorch → ONNX → TensorRTtrtexec --onnx=model.onnx --saveEngine=model.trtmodel.trt 就是 TensorRT 部署文件| 平台 | 常用格式 |
|---|---|
| Android / iOS | ONNX |
| iOS | CoreML (.mlmodel) |
| 边缘 GPU | TensorRT |
| 嵌入式 | NCNN / OpenVINO |
import coremltools as ct
model = ct.convert(
"model.onnx",
inputs=[ct.ImageType(shape=(1, 3, 224, 224))]
)
model.save("model.mlmodel")一个完整部署包通常包含:
model.onnx / model.pt / saved_model/
config.yaml # 模型配置
labels.txt # 类别标签
preprocess.py # 预处理
requirements.txt # 依赖| 场景 | 推荐导出格式 |
|---|---|
| 通用部署 | ONNX |
| PyTorch C++ | TorchScript |
| TF Serving | SavedModel |
| NVIDIA GPU | TensorRT |
| 移动端 | ONNX / CoreML |
| 云推理 | ONNX / Triton |
如果你愿意,可以直接告诉我:
我可以给你一步一步的具体命令和文件结构。