Metadata-Version: 2.4
Name: onnx-rewriter
Version: 0.1.1
Summary: DAG subgraph isomorphism matcher and rewriter for ONNX IR
Author-email: zhengankun <zhengankun@example.com>
Maintainer-email: zhengankun <zhengankun@example.com>
License: MIT
Project-URL: Homepage, https://gitee.com/zhengankun/onnx-rewriter
Project-URL: Repository, https://gitee.com/zhengankun/onnx-rewriter
Project-URL: Issues, https://gitee.com/zhengankun/onnx-rewriter/issues
Classifier: Programming Language :: Python :: 3
Classifier: Programming Language :: Python :: 3.9
Classifier: Programming Language :: Python :: 3.10
Classifier: Programming Language :: Python :: 3.11
Classifier: Programming Language :: Python :: 3.12
Classifier: License :: OSI Approved :: MIT License
Classifier: Operating System :: OS Independent
Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
Classifier: Topic :: Software Development :: Libraries :: Python Modules
Requires-Python: >=3.9
Description-Content-Type: text/markdown
License-File: LICENSE
Requires-Dist: networkx>=2.5
Requires-Dist: onnx-ir>=0.2.1
Requires-Dist: loguru>=0.7.0
Requires-Dist: numpy>=1.24.0
Provides-Extra: test
Requires-Dist: pytest>=7.0; extra == "test"
Requires-Dist: pytest-cov>=3.0; extra == "test"
Requires-Dist: onnx>=1.15.0; extra == "test"
Requires-Dist: onnxruntime>=1.17.0; extra == "test"
Provides-Extra: dev
Requires-Dist: pytest>=7.0; extra == "dev"
Requires-Dist: pytest-cov>=3.0; extra == "dev"
Requires-Dist: black>=23.0; extra == "dev"
Requires-Dist: flake8>=6.0; extra == "dev"
Requires-Dist: mypy>=1.0; extra == "dev"
Requires-Dist: onnx>=1.15.0; extra == "dev"
Requires-Dist: onnxruntime>=1.17.0; extra == "dev"
Dynamic: license-file

# 🚀 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)` 方法，即可添加自定义优化规则。

---

## 📦 安装

```bash
pip install onnx-rewriter
```

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

---

## 🎬 快速开始

### 1️⃣ 导出并量化 ResNet18（示例模型）

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

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

---

### 2️⃣ QDQ → QOperator 转换

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

将所有 `DequantizeLinear → Op → QuantizeLinear` 模式转为对应的 `QLinear` 算子（如 `QLinearConv`），同时处理偏置。

**转换前后对比**：

**QDQ 格式**（量化/反量化显式）  
![resnet18_qdq](https://fastly.jsdelivr.net/gh/bucketio/img17@main/2026/08/08/1786167198257-42b1afe7-7c53-446a-a295-2b001f059bef.png)

**QOperator 格式**（QLinear 算子内嵌量化参数）  
![resnet18_qop](https://fastly.jsdelivr.net/gh/bucketio/img11@main/2026/08/08/1786167153430-ec5b493a-f6d8-441d-835b-b644fd56b357.png)

---

### 3️⃣ Einsum 分解

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

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

**分解前后对比**：

**原始 Einsum 图**  
![einsum_original](https://fastly.jsdelivr.net/gh/bucketio/img3@main/2026/08/08/1786167262162-4d9bdac5-ba14-41ce-9b8c-6465b48dd4f5.png)

**分解后图**  
![einsum_decompose](https://fastly.jsdelivr.net/gh/bucketio/img18@main/2026/08/08/1786167227758-93678a2c-9736-4ac7-b6ff-1a433b6d601d.png)

---

### 4️⃣ 融合激活层到 QuantizeLinear

```python
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`。偏置不作为模式输入，而是在匹配后检查。

```python
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`。模式可定义为：

```python
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`。

### 完整重写器模板

```python
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 的命名空间避免冲突。

---

## 📐 核心设计与流程

### 类图

```mermaid
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
```

### 匹配与改写流程图

```mermaid
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
