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

如何实现PyTorch自定义哈希表查找OP的ONNX导出symbolic方法?

实现TableLookup OP的symbolic方法以导出为ONNX的LabelEncoder

核心方案

你的整数表查找逻辑完全适配ONNX ML规范中版本2的LabelEncoder OP的map模式,该模式支持将整数输入通过预定义键值对映射到目标整数输出。只需在symbolic方法中提取哈希表的键值对,构造符合ONNX要求的节点即可。

具体实现

  1. 为哈希表封装添加键值对提取接口
    给你的table对象新增一个方法(比如get_key_value_pairs()),返回所有预定义键值对的Python列表,格式为[(key1, value1), (key2, value2), ...]。

  2. 实现symbolic方法
    修改TableLookup类的symbolic静态方法,生成ai.onnx.ml.LabelEncoder节点:

    class TableLookup(torch.autograd.Function):
        @staticmethod
        def forward(ctx, table, x):
            y = table.lookup(x)
            return y
        
        @staticmethod
        def symbolic(g, table, x):
            # 提取键值对并扁平化(ONNX要求map_int64为[key, val, key, val,...]格式)
            kv_pairs = table.get_key_value_pairs()
            flat_kv = []
            for key, val in kv_pairs:
                flat_kv.append(key)
                flat_kv.append(val)
            
            # 构造LabelEncoder节点,指定map模式和键值对
            return g.op(
                "ai.onnx.ml.LabelEncoder",
                x,
                mode_s="map",
                map_int64_i=flat_kv,
                domain_="ai.onnx.ml",
                version_=2
            )
    

导出与验证

  • 导出ONNX模型
    导出时需指定足够高的opset版本(建议13及以上),确保PyTorch支持ai.onnx.ml域的OP:

    import torch
    
    # 初始化哈希表和测试输入
    table = YourTableWrapper()
    x = torch.tensor([1, 2, 3], dtype=torch.int64)
    inference_fn = lambda x: TableLookup.apply(table, x)
    
    torch.onnx.export(
        inference_fn,
        (x,),
        "table_lookup.onnx",
        opset_version=13,
        input_names=["x"],
        output_names=["y"]
    )
    
  • 验证导出结果
    使用ONNX Runtime加载模型,对比输出是否与PyTorch一致:

    import onnxruntime as ort
    import numpy as np
    
    sess = ort.InferenceSession("table_lookup.onnx")
    ort_out = sess.run(["y"], {"x": x.numpy()})[0]
    torch_out = TableLookup.apply(table, x).numpy()
    
    assert np.array_equal(ort_out, torch_out)
    

内容的提问来源于stack exchange,提问作者user416983

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 11:45:10