4.8 构建onnx结构模型-Less

前言

构建onnx方式通常有两种:

1、通过代码转换成onnx结构,比如pytorch ---> onnx

2、通过onnx 自定义结点,图,生成onnx结构

本文主要是简单学习和使用两种不同onnx结构,

下面以 Less 结点进行分析

方式

方法一:pytorch --> onnx

暂缓,主要研究方式二

方法二: onnx

cpp 复制代码
import onnx 
from onnx import TensorProto, helper, numpy_helper
import numpy as np

def run():
    print("run start....\n")

    less = helper.make_node(
        "Less",
        name="Less_0",
        inputs=["input1", "input2"],
        outputs=["output1"],
    )
    input1_data = np.load("./tensor.npy") # 16, 397
    # input1_data = np.load("./data.npy")  # 16, 398 test
    # print(f"input1_data shape:{input1_data.shape}\n")
    # input1_data = np.zeros((16,398))
    initializer = [ 
        helper.make_tensor("input1", TensorProto.FLOAT, [16,397], input1_data)
    ]

    cast_nodel = helper.make_node(
            op_type="Cast",
            inputs=["output1"],
            outputs=["output2"],
            name="test_cast",
            to=TensorProto.FLOAT,
        )
    value_info = helper.make_tensor_value_info(
            "output2", TensorProto.BOOL, [16,397])

    graph = helper.make_graph(
        nodes=[less, cast_nodel],
        name="test_graph",
        inputs=[helper.make_tensor_value_info(
            "input2", TensorProto.FLOAT, [16,1]
        )],
        outputs=[helper.make_tensor_value_info(
            "output2",TensorProto.FLOAT, [16,397]
        )],
        initializer=initializer,
        value_info=[value_info],
    )

    op = onnx.OperatorSetIdProto()
    op.version = 11
    model = helper.make_model(graph, opset_imports=[op])
    model.ir_version = 8
    print("run done....\n")
    return model

if __name__ == "__main__":
    model = run()
    onnx.save(model, "./test_less_ori.onnx")

run

cpp 复制代码
import onnx
import onnxruntime
import numpy as np


# 检查onnx计算图
def check_onnx(mdoel):
    onnx.checker.check_model(model)
    # print(onnx.helper.printable_graph(model.graph))

def run(model):
    print(f'run start....\n')
    session = onnxruntime.InferenceSession(model,providers=['CPUExecutionProvider'])
    input_name1 = session.get_inputs()[0].name  
    input_data1= np.random.randn(16,1).astype(np.float32)
    print(f'input_data1 shape:{input_data1.shape}\n')

    output_name1 = session.get_outputs()[0].name

    pred_onx = session.run(
    [output_name1], {input_name1: input_data1})[0]

    print(f'pred_onx shape:{pred_onx.shape} \n')

    print(f'run end....\n')


if __name__ == '__main__':
    path = "./test_less_ori.onnx"
    model = onnx.load("./test_less_ori.onnx")
    check_onnx(model)
    run(path)
相关推荐
Freak嵌入式14 小时前
版本混乱 / 依赖缺失?uPyPi:MicroPython 版 PyPI,彻底解决库管理混乱
linux·服务器·数据库·单片机·嵌入式硬件·性能优化·依赖倒置原则
爱喝水的鱼丶19 小时前
SAP-ABAP:SAP 实战笔记:BAPI_OUTB_DELIVERY_CONFIRM_DEC 详解——外向交货单过账的利器
运维·性能优化·sap·abap·经验交流
咱入行浅20 小时前
慢查询日志在性能优化中的价值
性能优化
布兰妮甜2 天前
原子化 CSS 深度解析:Tailwind 原理、自定义配置、大型项目利弊
css·性能优化·tailwind·前端工程化·设计系统
要开心吖ZSH2 天前
一次 HikariCP 连接池耗尽导致的线上雪崩排查实录——问题概览
java·spring boot·mysql·性能优化·连接池·hikaricp
想你依然心痛2 天前
地图瓦片原理深度解析:切片、加载、缓存机制详解
性能优化·离线地图·金字塔模型·缓存机制·地图瓦片·切片规则·瓦片加载
数据库小学妹3 天前
MySQL联合索引失效怎么解决?从key_len逆向分析
数据库·mysql·性能优化·索引失效·联合索引·最左前缀·keylen
一只叫煤球的猫3 天前
ThreadForge 源码解读三:从任务执行到并发编排,ScopeJoiner 是怎么工作的?
后端·性能优化·开源
做前端的娜娜子3 天前
如何实现网页加载进度条?
性能优化·掘金·金石计划
小孔龙3 天前
Android GPU 渲染管线:一帧画面如何走上屏幕
android·性能优化·gpu