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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 04:21:13