ONNX Runtime中如何调用PyTorch模型非forward函数获取输出
核心原因
ONNX 仅存储静态计算图结构、模型权重与输入输出定义,不会保留 PyTorch 模型类的自定义成员方法逻辑。你之前导出 ONNX 时,仅通过默认forward方法跟踪计算路径,导出的计算图只覆盖从图像、文本输入到最终图文相似度logits的计算流程,中间生成的图像嵌入、文本嵌入张量没有被标记为模型输出,因此无法通过ONNX Runtime直接获取对应结果。
实现方案
不需要在ONNX Runtime中复刻PyTorch的自定义方法调用逻辑,只要把需要的嵌入张量暴露为ONNX模型的输出即可,有两种常用实现路径:
方案1:重新导出ONNX时直接加入嵌入输出(最稳妥,推荐)
写一个简单的模型包装类,在forward方法中同时返回你需要的图像嵌入、文本嵌入和原有的logits结果,导出时把这些返回值都加入输出列表即可。
import torch import clip class CLIPExportWrapper(torch.nn.Module): def __init__(self, base_clip_model): super().__init__() self.model = base_clip_model def forward(self, image_input, text_input): # 调用原模型的两个编码方法获取嵌入 image_emb = self.model.encode_image(image_input) text_emb = self.model.encode_text(text_input) # 复刻原CLIP forward的相似度计算逻辑 image_emb_norm = image_emb / image_emb.norm(dim=-1, keepdim=True) text_emb_norm = text_emb / text_emb.norm(dim=-1, keepdim=True) scale = self.model.logit_scale.exp() logits_img = scale * image_emb_norm @ text_emb_norm.t() logits_text = logits_img.t() # 按顺序返回所有需要的结果 return image_emb, text_emb, logits_img, logits_text # 加载原模型 device = "cuda" if torch.cuda.is_available() else "cpu" base_model, preprocess = clip.load("RN50", device=device) base_model.eval() wrapped_model = CLIPExportWrapper(base_model).to(device) # 准备dummy输入 img_size = base_model.visual.input_resolution dummy_img = torch.randn(10, 3, img_size, img_size).to(device) dummy_text = clip.tokenize(["quick brown fox", "lorem ipsum"]).to(device) # 导出ONNX torch.onnx.export( wrapped_model, (dummy_img, dummy_text), "clip_with_embedding.onnx", export_params=True, input_names=["IMAGE", "TEXT"], # 输出名顺序和forward返回顺序严格对应 output_names=["IMAGE_EMBEDDING", "TEXT_EMBEDDING", "LOGITS_PER_IMAGE", "LOGITS_PER_TEXT"], opset_version=14, dynamic_axes={ "IMAGE": {0: "image_batch"}, "TEXT": {0: "text_batch"}, "IMAGE_EMBEDDING": {0: "image_batch"}, "TEXT_EMBEDDING": {0: "text_batch"}, "LOGITS_PER_IMAGE": {0: "image_batch", 1: "text_batch"}, "LOGITS_PER_TEXT": {0: "text_batch", 1: "image_batch"} } )
导出完成后,ONNX Runtime推理时直接指定输出名就能拿到嵌入结果:
import onnxruntime as ort import numpy as np sess = ort.InferenceSession("clip_with_embedding.onnx") results = sess.run( ["IMAGE_EMBEDDING", "TEXT_EMBEDDING", "LOGITS_PER_IMAGE", "LOGITS_PER_TEXT"], {"IMAGE": dummy_img.cpu().numpy(), "TEXT": dummy_text.cpu().numpy()} ) image_embedding, text_embedding, logits_img, logits_text = results # 此时image_embedding形状为[batch_size, 1024],和PyTorch下encode_image返回结果完全一致 print(image_embedding.shape)
方案2:修改已导出的ONNX模型,追加中间张量为输出
如果不想重新导出模型,可以直接加载已有的ONNX文件,找到对应嵌入张量的节点名,将其标记为模型输出即可。
import onnx from onnx import helper # 加载已导出的原模型 onnx_model = onnx.load("model.onnx") # 先可视化模型结构,定位encode_image、encode_text最终输出的张量名称 # 图像嵌入对应visual分支最后一层的输出,文本嵌入对应文本编码分支最后一层的输出 # 追加输出定义,elem_type=1对应float32类型,和CLIP模型权重类型一致 onnx_model.graph.output.extend([ helper.ValueInfoProto( name="替换为定位到的图像嵌入张量名", type=helper.TensorProto(elem_type=1) ), helper.ValueInfoProto( name="替换为定位到的文本嵌入张量名", type=helper.TensorProto(elem_type=1) ) ]) # 保存修改后的模型 onnx.save(onnx_model, "clip_modified.onnx")
注意:定位中间张量名时建议通过模型可视化工具确认,不要手动猜测节点名称,避免取到错误的中间结果。
内容的提问来源于stack exchange,提问作者Deshwal
相关产品推荐
相关产品推荐

