Metadata-Version: 2.4
Name: torch-to-onnx
Version: 0.0.3
Summary: A tool for converting PyTorch models to ONNX using ir-py
Author-email: zhengankun <ankun.zheng@qq.com>
License: MIT
Requires-Python: >=3.8
Description-Content-Type: text/markdown
License-File: LICENSE
Requires-Dist: torch>=2.0.0
Requires-Dist: onnx>=1.14.0
Requires-Dist: onnx-ir>=0.1.0
Provides-Extra: test
Requires-Dist: pytest>=7.0.0; extra == "test"
Requires-Dist: torchvision>=0.15.0; extra == "test"
Dynamic: license-file

# torch-to-onnx

将 PyTorch 模型转换为 ONNX 格式，基于现代 `torch.export` API 和 `onnx_ir` 库。

---

## 为什么使用 `torch.export`？

- **模型结构可视化**：导出的 `ExportedProgram` 可以直接在 **Model Explorer**（如 Netron、PyTorch 官方可视化工具）中打开，清晰展示模块的嵌套层次和完整的计算图结构，极大方便了模型调试、结构理解和算子核对。

- **自定义算子零成本支持**：对于通过 `torch.library` 注册的自定义算子，`torch.export` 能直接捕获并保留在图中，**无需额外编写任何注册代码**即可顺利导出，大幅降低了扩展性门槛，使新算子的引入更加轻松。

1. **更准确的图捕获**  
   - 基于 FX 图（而非 JIT 追踪），可捕获完整的计算图，包括条件分支和循环。  
   - 输出的是**可重入、可序列化**的计算图，与模型权重分离，方便后续优化和部署。

2. **与 PyTorch 生态深度集成**  
   - 与 `torch.compile`、`torch._dynamo` 同源，兼容 PyTorch 2.0+ 的 torch 操作。  
   - 天然支持 `torch.export.export` 输出的 `ExportedProgram` 格式，可方便地进行图变换（如量化、剪枝）。

3. **分工明确**  
   - **算法研发人员**：只需关注模型定义，调用 `torch.export.export` 即可得到 `ExportedProgram`，无需关心 ONNX 细节。  
   - **工程/部署团队**：编写转换器，将 `ExportedProgram` 转为 ONNX，可集中精力处理算子映射和性能优化。

本工具正是基于这一思想，将 `ExportedProgram` 转换为标准 ONNX 模型，让前后端解耦。
## 特性

- ✅ 基于 `torch.export`（FX 图）精准追踪。  
- ✅ 完全控制 ONNX opset 和 IR 版本。  
- ✅ 轻量、易扩展 —— 可自由添加 ATen 算子映射。  
- ✅ 输出标准 ONNX 模型，兼容任何 ONNX Runtime（如 ONNX Runtime、TensorRT、OpenVINO）。

## 安装

```bash
pip install torch-to-onnx
```

若未安装 `onnx-ir`，请单独安装：

```bash
pip install onnx-ir   # 或根据你的 onnx_ir 来源安装
```

## 基本用法

```python
import torch
import torchvision.models as models
from torch_to_onnx import convert_exported_program_to_onnx
import onnx

# 1. 创建模型
model = models.resnet18(pretrained=True).eval()
dummy_input = torch.randn(1, 3, 224, 224)

# 2. 使用 torch.export 导出（算法研发负责此步）
ep = torch.export.export(model, (dummy_input,), strict=True)

# 3. 转换为 ONNX（工程团队负责此转换器）
ir_model = convert_exported_program_to_onnx(
    ep,
    opset_version=18,
)

# 4. 序列化并保存
ir.save(ir_model, "resnet18.onnx")
print("ONNX 模型已保存！")
```

## 自定义算子映射

若转换过程中遇到未支持的 ATen 算子，你可通过两种方式扩展：

### 方式一：修改映射表（简单场景）

在 `ops.py` 中的 `ATEN_TO_ONNX_OP` 字典中添加新条目，例如：

```python
from torch_to_onnx.ops import register_mapping
register_mapping("my_aten_op", "MyOnnxOp") # 添加映射
```

### 方式二：注册带特殊处理逻辑的算子（推荐）

使用 `@register_op` 装饰器，在 `converter.py` 或任何地方定义自己的处理器函数，统一注册 ATen 别名、ONNX 算子名和转换逻辑：

```python
from torch_to_onnx.ops import register_op

@register_op(["my_aten_op", "my_aten_op.default"], onnx_op="MyOnnxOp")
def handle_my_op(node, value_map, tape, opset_version):
    # node 是 FX 节点，你可自定义输入和属性提取
    inputs = []
    for arg in node.args:
        if isinstance(arg, torch.fx.Node):
            inputs.append(value_map[arg.name])
        else:
            # 处理常量...
    attrs = {"some_attr": node.kwargs.get("attr")}
    return inputs, attrs
```

然后重新运行转换，新算子将被正确处理。

## 支持的算子

当前支持的常见 ATen 算子包括：

- 卷积/归一化：`conv2d`, `batch_norm`, `relu`, `max_pool2d`, `avg_pool2d`, `adaptive_avg_pool2d`
- 算术：`add`, `sub`, `mul`, `div`, `matmul`, `mm`
- 全连接：`linear`（转为 Gemm）
- 形状变换：`reshape`, `view`, `transpose`, `permute`, `unsqueeze`, `squeeze`, `cat`, `stack`
- 归约：`mean`, `sum`
- 激活：`softmax`, `sigmoid`, `tanh`
- 其他：`dropout`, `slice`, `flatten`, `shape`

如果你遇到未覆盖的算子，可按照上述“自定义算子”一节轻松添加。

## 依赖

- PyTorch >= 1.12.0
- ONNX >= 1.12.0
- onnx-ir（或你自己的 onnx_ir 分支）

## 许可证

MIT

## 作者

zhengankun (ankun.zheng@qq.com)
