如何将ONNX模型转回pth格式或在fastai learner中加载ONNX模型
fastai加载ONNX模型实现方案
fastai原生Learner.load()方法仅支持加载fastai/PyTorch格式的.pth权重,无法直接读取ONNX格式模型,你期望的learn.load('model.onnx')调用没有原生支持,以下是两种可落地的实现路径:
方案1:封装ONNX推理逻辑适配Learner(无需转pth,仅支持推理场景)
如果只需要用fastai的推理、验证接口跑模型,不需要微调训练,可以把ONNX推理逻辑封装成PyTorch兼容的模块,直接挂载到Learner上使用,步骤如下:
- 先安装依赖:确保环境已安装
onnxruntime、对应版本的fastai和torch - 自定义ONNX推理包装类,兼容PyTorch张量流:
import torch import onnxruntime as ort from fastai.vision.all import * class ONNXInferenceWrapper(torch.nn.Module): def __init__(self, onnx_path): super().__init__() self.ort_session = ort.InferenceSession(onnx_path, providers=['CPUExecutionProvider']) # 自动读取ONNX的输入输出名,无需硬编码 self.input_name = self.ort_session.get_inputs()[0].name self.output_name = self.ort_session.get_outputs()[0].name def forward(self, x): # PyTorch张量转numpy供ONNX推理 x_np = x.cpu().numpy() ort_result = self.ort_session.run([self.output_name], {self.input_name: x_np})[0] # 推理结果转回PyTorch张量,适配fastai计算流程 return torch.tensor(ort_result).to(x.device)
- 初始化Learner后替换原有模型即可正常调用推理接口:
# 按导出ONNX时的预处理逻辑初始化你的DataLoaders dls = ... # 初始化基础Learner,损失函数、指标按实际任务配置 learn = Learner(dls, torch.nn.Identity(), loss_func=CrossEntropyLossFlat(), metrics=accuracy) # 替换为包装好的ONNX模型 learn.model = ONNXInferenceWrapper("model.onnx") # 后续可正常调用learn.validate()、learn.get_preds()、learn.predict()等推理接口
注意:该方案仅支持推理/验证,ONNX是静态计算图不支持反向传播,需要做微调训练必须使用方案2。
方案2:将ONNX权重转为PyTorch pth格式(支持训练+推理,兼容原生load)
目前没有一键完成ONNX转fastai可用pth的工具,需要手动对齐模型结构后迁移权重,步骤如下:
- 首先确认导出ONNX时使用的原始模型结构,必须和导出时的结构完全一致,否则权重加载会报错
- 读取ONNX权重并赋值给对应结构的PyTorch模型:
import onnx import torch from fastai.vision.all import * # 初始化和导出ONNX时完全一致的模型结构,替换为你自己的模型定义 model = timm.create_model("resnet18", pretrained=False, num_classes=10) # 加载ONNX模型文件 onnx_model = onnx.load("model.onnx") # 提取ONNX中的所有权重参数 onnx_weight_map = {} for init_tensor in onnx_model.graph.initializer: weight = torch.from_numpy(onnx.numpy_helper.to_array(init_tensor)) onnx_weight_map[init_tensor.name] = weight # 对齐权重到PyTorch模型的state_dict pth_state_dict = model.state_dict() for layer_name in pth_state_dict.keys(): if layer_name in onnx_weight_map: target_weight = onnx_weight_map[layer_name] # 处理ONNX和PyTorch的权重维度顺序差异(比如卷积层、全连接层可能需要转置) if pth_state_dict[layer_name].shape == target_weight.shape: pth_state_dict[layer_name] = target_weight else: # 以卷积层权重转置为例,根据实际模型结构调整维度顺序 pth_state_dict[layer_name] = target_weight.permute(3,2,0,1) # 加载对齐后的权重,strict设为True可检查是否有遗漏层 model.load_state_dict(pth_state_dict, strict=True)
- 将加载好权重的模型挂载到Learner,保存为fastai兼容的pth格式,后续即可用原生
load方法调用:
dls = ... # 你的DataLoaders对象 learn = Learner(dls, model, loss_func=CrossEntropyLossFlat(), metrics=accuracy) # 保存为fastai格式权重,后续直接调用learn.load('model_fastai')即可加载 learn.save('model_fastai')
注意:如果你的ONNX是从TensorFlow、PaddlePaddle等非PyTorch框架导出的,层命名、权重维度差异会更大,需要逐层核对调整。
内容的提问来源于stack exchange,提问作者dheo arokhim
相关产品推荐
相关产品推荐

