🚀 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 格式(量化/反量化显式)
QOperator 格式(QLinear 算子内嵌量化参数)
3️⃣ Einsum 分解
python examples/run_einsum_decomposer.py --input einsum_original.onnx --output einsum_decomposed.onnx
将复杂的 Einsum 算子分解为 MatMul、Mul、Transpose、ReduceSum 等基本算子,便于后续优化或部署。
分解前后对比:
原始 Einsum 图
分解后图
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中其他OpPattern的name字符串,顺序需与算子实际输入顺序一致。
如果某个输入不需参与模式匹配(例如偏置,只需检查其为常量),可以不在模式中声明,匹配后手动检查。 -
SubgraphPattern:由多个OpPattern组成,通过输入名称引用声明节点间边关系,支持任意复杂的 DAG 结构,包括 多个节点共享同一个前驱。 -
SubgraphPatterns:可同时提供多个候选模式,匹配时依次尝试,返回带pattern_index的MatchResult,便于区分命中哪个模式。 -
安全替换工具:使用
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:匹配共享输入的子图(所有输入均在模式中声明)
设想一个场景:Conv 和 Add 使用同一个数据源(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_mapping和output_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
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
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
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
89fc4661dab1d0898a82fdac2863d752ac74b7744abd2c5d34856e40c4cd8847
|
|
| MD5 |
38d48c6cd7361361881ab833d0428b0c
|
|
| BLAKE2b-256 |
a5aee8808086a0ba183b8fcd2340cc2061f212484f6a951e900a62acbca653b2
|