ConvTransFormatPass Python 样例使用指导
【免费下载链接】geGE(Graph Engine)是面向昇腾的图编译器和执行器,提供了计算图优化、多流并行、内存复用和模型下沉等技术手段,加速模型执行效率,减少模型内存占用。 GE 提供对 PyTorch、TensorFlow 前端的友好接入能力,并同时支持 onnx、pb 等主流模型格式的解析与编译。项目地址: https://gitcode.com/cann/ge
本目录提供graph_base_pass/3_modify_conv_data_format_pass的纯 Python版本示例,主链路与 C++ConvTransFormatPass一致:
- 遍历图中
Conv2D/Conv2DV2,筛选data_format == NCHW的节点; - 将
data_format改为NHWC; - 从卷积输出做 BFS,按顺序匹配
perm == [0,2,3,1]与[0,3,1,2]的Transpose,删除对应Transpose与 perm 常量产点并重连数据边。
本样例继承FusionBasePass并重写run(),通过Graph.remove_edge/add_data_edge/remove_node与Node.set_attr完成改写,不使用SubgraphRewriter。
与 C++ 版本的差异
- 整图回滚:C++ 在
SetAttr或删边失败时用备份图恢复。当前 PythonGraph不支持整图深拷贝,本样例在失败时抛出异常,不保证与 C++ 相同的原子回滚语义。 - 读取 Transpose 的 perm:C++ 使用
GNode::GetInputConstData。Python 侧通过 perm 输入端的Const/Constant节点的value属性读取(与pattern_base_pass/4_add_zero_pass中 Const 校验方式一致)。若图中 perm 不以此形式出现,可能无法识别并删除Transpose,与 C++ 覆盖范围可能略有差别。
前置条件
- 已 source CANN 环境(
source ${ASCEND_PATH}/set_env.sh) - 可导入 GE Python 包(含
ge.graph、ge.passes及 pass 加载链路)
使用方式
- 通过环境变量让 GE 在编译期加载该 Python pass(在
3_modify_conv_data_format_pass目录下时):
export ASCEND_GE_PY_PASS_PATH=$PWD/python/src/python_modify_conv_data_format_pass.py- 复用上级目录 样例 README 中的ATC 离线编译或在线推理步骤(
data/torch_gen_onnx.py、data/torch_forward.py等)。
预期现象
日志中会出现类似打印:
PythonConvTransFormatPass is starting Remove output edges success Remove output edges success PythonConvTransFormatPass completed对比DUMP_GE_GRAPH导出的 pbtxt 时,应看到卷积data_format为NHWC,且目标Transpose被移除(与 C++ 样例说明一致)。
【免费下载链接】geGE(Graph Engine)是面向昇腾的图编译器和执行器,提供了计算图优化、多流并行、内存复用和模型下沉等技术手段,加速模型执行效率,减少模型内存占用。 GE 提供对 PyTorch、TensorFlow 前端的友好接入能力,并同时支持 onnx、pb 等主流模型格式的解析与编译。项目地址: https://gitcode.com/cann/ge
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考