Skip to main content

🚀 ONNX Rewriter

基于 onnx-ir 的通用 ONNX 图匹配与改写框架,内置常用量化优化 Pass,适用于模型剪枝、融合、量化格式转换等场景。

作者:zhengankun
许可证:MIT


✨ 特性

  • 🔍 通用图匹配引擎 —— 支持 子图模式(SubgraphPattern或模式(SubgraphPatterns,可匹配任意 DAG 结构(含共享输入、多输出、分支汇合)。
  • 🧩 内置重写器(开箱即用):
    • ConvAddRewriter:将 Conv + Add 折叠为带偏置的 Conv
    • QDQToQOperatorRewriter:将 DequantizeLinear → Op → QuantizeLinear 转换为对应的 QLinear 算子(Conv/MatMul/Add/Mul/Gemm)。
    • FuseReluClipToQuantizeRewriter:将 Relu/Relu6/Clip 融合到 QuantizeLinear(条件:量化范围被截断范围覆盖)。
    • EinsumDecomposerRewriter:将 Einsum 算子分解为 MatMul、Mul、Transpose、ReduceSum 等基本算子(支持常用模式)。
  • 🧹 自动死代码消除 —— 改写后自动调用 RemoveUnusedNodesPass,清理孤立节点。
  • 🧩 极易扩展 —— 继承 Rewriter 基类并实现 rewrite(model) 方法,即可添加自定义优化规则。

📦 安装

pip install onnx-rewriter

若希望从源码安装,可克隆仓库并执行 pip install -e .


🎬 快速开始

1️⃣ 导出并量化 ResNet18(示例模型)

# 导出原始 ResNet18
python examples/export_resnet18.py

# 静态量化生成 QDQ 格式模型
python examples/quantize_resnet18.py

2️⃣ QDQ → QOperator 转换

python examples/run_qdq_to_qoperator.py resnet18_qdq.onnx resnet18_qop.onnx

将所有 DequantizeLinear → Op → QuantizeLinear 模式转为对应的 QLinear 算子(如 QLinearConv),同时处理偏置。

转换前后对比

QDQ 格式(量化/反量化显式)
resnet18_qdq

QOperator 格式(QLinear 算子内嵌量化参数)
resnet18_qop


3️⃣ Einsum 分解

python examples/run_einsum_decomposer.py --input einsum_original.onnx --output einsum_decomposed.onnx

将复杂的 Einsum 算子分解为 MatMul、Mul、Transpose、ReduceSum 等基本算子,便于后续优化或部署。

分解前后对比

原始 Einsum 图
einsum_original

分解后图
einsum_decompose


4️⃣ 融合激活层到 QuantizeLinear

import onnx_ir as ir
from onnx_rewriter.rewriters import FuseReluClipToQuantizeRewriter

model = ir.load("resnet18_qop.onnx")
rewriter = FuseReluClipToQuantizeRewriter()
optimized_model = rewriter.rewrite(model)
ir.save(optimized_model, "resnet18_fused.onnx")

Relu/Relu6/Clip 的截断范围覆盖量化范围,则将其融合进 QuantizeLinear,进一步精简计算图。


🧑‍💻 自定义重写器

核心概念

  • OpPattern:描述单个算子的类型(支持 | 多选、通配符 *)及其输入名称(用于建立依赖)。
    重要OpPattern.inputs 必须是同一 SubgraphPattern 中其他 OpPatternname 字符串,顺序需与算子实际输入顺序一致。
    如果某个输入不需参与模式匹配(例如偏置,只需检查其为常量),可以不在模式中声明,匹配后手动检查。

  • SubgraphPattern:由多个 OpPattern 组成,通过输入名称引用声明节点间边关系,支持任意复杂的 DAG 结构,包括 多个节点共享同一个前驱

  • SubgraphPatterns:可同时提供多个候选模式,匹配时依次尝试,返回带 pattern_indexMatchResult,便于区分命中哪个模式。

  • 安全替换工具:使用 replace_subgraph(位于 onnx_rewriter.core.replace)安全地删除旧子图、插入新子图并修复连接关系,推荐所有重写器使用此函数

示例 1:Conv + Add 折叠(使用 replace_subgraph)

以下示例演示如何将 Conv + Add(其中 Add 的第二个输入是常量偏置)折叠为带偏置的 Conv。偏置不作为模式输入,而是在匹配后检查。

from onnx_rewriter.core import Rewriter, SubgraphPattern, OpPattern, GraphMatcher
from onnx_rewriter.core.replace import replace_subgraph
import onnx_ir as ir

class ConvAddFusionRewriter(Rewriter):
    def rewrite(self, model):
        graph = model.graph

        # 定义模式:Add 依赖 Conv 的输出(仅声明两个节点)
        pattern = SubgraphPattern([
            OpPattern("Conv", name="conv"),
            OpPattern("Add", name="add", inputs=["conv"]),   # Add 的第一个输入来自 Conv
        ])

        matcher = GraphMatcher(pattern)
        for match in matcher.match_graph(graph):
            conv_node = match.get_op("conv")
            add_node = match.get_op("add")

            # 检查 Add 的第二个输入是否为常量(initializer)
            bias_val = add_node.inputs[1]
            if bias_val.name not in graph.initializers:
                continue

            # 确保 Conv 的输出仅被 Add 使用
            conv_out = conv_node.outputs[0]
            consumers = [n for n in graph.nodes if conv_out in n.inputs]
            if len(consumers) != 1:
                continue

            # --- 构建替换子图 ---
            tape = ir.tape.Tape()
            # 子图输入:data, weight(bias 作为 initializer)
            data_val = ir.val("data", dtype=conv_node.inputs[0].dtype, shape=conv_node.inputs[0].shape)
            weight_val = ir.val("weight", dtype=conv_node.inputs[1].dtype, shape=conv_node.inputs[1].shape)
            # 复制 bias 常量
            bias_tensor = graph.initializers[bias_val.name].const_value
            bias_val_new = ir.val("bias", const_value=bias_tensor)

            # 创建新的 Conv 节点(3 个输入)
            new_out = tape.op(
                "Conv",
                inputs=[data_val, weight_val, bias_val_new],
                attributes=conv_node.attributes,
                name=f"{conv_node.name}_with_bias",
            )
            new_out.shape = add_node.outputs[0].shape
            new_out.dtype = add_node.outputs[0].dtype

            subgraph = ir.Graph(
                inputs=[data_val, weight_val],
                outputs=[new_out],
                nodes=tape.nodes,
                initializers=[bias_val_new],
                opset_imports=graph.opset_imports,
                name=f"{conv_node.name}_fused",
            )

            # 映射:子图输入 -> 主图实际值
            input_mapping = {
                data_val: conv_node.inputs[0],
                weight_val: conv_node.inputs[1],
            }
            output_mapping = {new_out: add_node.outputs[0]}

            # 执行替换(自动删除旧节点并修复连接)
            replace_subgraph(graph, subgraph, input_mapping, output_mapping)

        return model

示例 2:匹配共享输入的子图(所有输入均在模式中声明)

设想一个场景:ConvAdd 使用同一个数据源(Identity 的输出),且 Add 的另一个输入来自 Constant。模式可定义为:

pattern = SubgraphPattern([
    OpPattern("Identity", name="data"),                 # 数据节点
    OpPattern("Conv", name="conv", inputs=["data"]),    # Conv 使用 data
    OpPattern("Constant", name="bias"),                 # 偏置常量
    OpPattern("Add", name="add", inputs=["data", "bias"]), # Add 共享 data,并使用 bias
])

匹配后,你可以使用相同的方法构建新子图并调用 replace_subgraph

完整重写器模板

from onnx_rewriter.core import Rewriter, GraphMatcher, SubgraphPattern, OpPattern
from onnx_rewriter.core.replace import replace_subgraph
import onnx_ir as ir

class MyRewriter(Rewriter):
    def rewrite(self, model):
        graph = model.graph

        pattern = SubgraphPattern([
            OpPattern("OpA", name="a"),
            OpPattern("OpB", name="b", inputs=["a"]),
        ])

        matcher = GraphMatcher(pattern)
        for match in matcher.match_graph(graph):
            a_node = match.get_op("a")
            b_node = match.get_op("b")

            # 构建新子图(略),然后调用 replace_subgraph
            # subgraph = ...
            # replace_subgraph(graph, subgraph, input_mapping, output_mapping)

        return model

提示replace_subgraph 会依据 input_mappingoutput_mapping 自动识别旧子图的边界,无需手动指定 old_nodes。它还会处理 initializer 的命名空间避免冲突。


📐 核心设计与流程

类图

classDiagram
    class Rewriter {
        +rewrite(model: ir.Model) ir.Model
    }
    class GraphMatcher {
        +match_graph(graph: ir.Graph) Iterator[MatchResult]
        +match_ops(nodes, node_by_output) Iterator[MatchResult]
    }
    class SubgraphPattern {
        +ops: List[OpPattern]
        +build_pattern_graph() nx.DiGraph
    }
    class OpPattern {
        +op_type: str
        +name: str
        +inputs: List[Union[str, int]]
        +match_op_type(actual_op_type: str) bool
    }
    class MatchResult {
        +pattern_index: int
        +get_op(pattern_or_name) ir.Node
        +get_value(pattern_or_name) ir.Value
        +get_nodes() List[ir.Node]
    }
    class replace_subgraph {
        <<function>>
        +replace_subgraph(graph, subgraph, input_mapping, output_mapping, apply_namespace)
    }
    Rewriter <|-- ConvAddRewriter
    Rewriter <|-- QDQToQOperatorRewriter
    Rewriter <|-- FuseReluClipToQuantizeRewriter
    Rewriter <|-- EinsumDecomposerRewriter
    GraphMatcher --> SubgraphPattern : uses
    SubgraphPattern --> OpPattern : contains
    GraphMatcher --> MatchResult : returns
    replace_subgraph --> Graph : modifies

匹配与改写流程图

flowchart TD
    A[加载ONNX模型] --> B[应用重写器]
    B --> C[定义子图模式]
    C --> D[GraphMatcher匹配]
    D --> E{是否匹配?}
    E -->|是| F[执行替换]
    F --> G[清理未使用节点]
    G --> H[拓扑排序]
    H --> I[输出优化模型]
    E -->|否| I

🤝 贡献

欢迎提交 Issue 和 Pull Request!如果你觉得这个工具不错,请给个 ⭐ 支持~


📄 许可证

MIT License © 2026 zhengankun

Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distributions

No source distribution files available for this release.See tutorial on generating distribution archives.

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

onnx_rewriter-0.1.1-py3-none-any.whl (30.5 kB view details)

Uploaded Python 3

File details

Details for the file onnx_rewriter-0.1.1-py3-none-any.whl.

File metadata

  • Download URL: onnx_rewriter-0.1.1-py3-none-any.whl
  • Upload date:
  • Size: 30.5 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/7.0.0 CPython/3.12.13

File hashes

Hashes for onnx_rewriter-0.1.1-py3-none-any.whl
Algorithm Hash digest
SHA256 89fc4661dab1d0898a82fdac2863d752ac74b7744abd2c5d34856e40c4cd8847
MD5 38d48c6cd7361361881ab833d0428b0c
BLAKE2b-256 a5aee8808086a0ba183b8fcd2340cc2061f212484f6a951e900a62acbca653b2

See more details on using hashes here.

Release history Release notifications | RSS feed

This release

0.1.1 This release

1 file

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page