ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

CANN/GE BatchMatMul展平融合Pass样例

CANN/GE BatchMatMul展平融合Pass样例 样例使用指导【免费下载链接】geGEGraph Engine是面向昇腾的图编译器和执行器提供了计算图优化、多流并行、内存复用和模型下沉等技术手段加速模型执行效率减少模型内存占用。 GE 提供对 PyTorch、TensorFlow 前端的友好接入能力并同时支持 onnx、pb 等主流模型格式的解析与编译。项目地址: https://gitcode.com/cann/ge功能描述本样例以 BatchMatMulV2 展平融合 pass 为例介绍将 BatchMatMul 展平为 MatMul 的融合 pass 实现 提供 atc 工具离线编译模型验证方式pass 使用 eager style api 和融合接口实现。融合原理将 A[b,m,k] B[k,n] 的 BatchMatMulV2 运算通过 Reshape 展平为 [b*m,k][k,n] 的 MatMulV2 运算 再通过 Reshape 恢复为 [b,m,n] 的输出形状。本样例仅覆盖 A 为 3 维、B 为 2 维、两个输入 dtype 相同且为 float/float16/bfloat16、无 bias/offset_w、offset_x0、transpose 属性均为 false 的教学场景。 带 batch 维广播、bias、offset_w 或非 0 offset_x 的 BatchMatMul/BatchMatMulV2 不在本样例优化范围内。目录结构├── README.md // C 样例说明 ├── src │ ├──batch_matmul_flatten_pass.cpp // pass 实现文件 ├── CMakeLists.txt // 编译脚本 ├── data │ ├──gen_onnx.py // onnx 导出脚本用于 ATC 离线验证支持 --batch/--m/--k/--n 参数 │ ├──quick_verify.sh // 一键式验证脚本支持传入 shape 和执行次数 │ ├──benchmark_model.cpp // 性能测试程序源码 ├── gen_es_api │ ├──CMakeLists.txt // 生成 eager style api 的编译脚本环境要求编译器GCC 7.3.x使用 python 及其依赖库版本python3.9、onnx已完成环境准备。实现步骤定义类BatchMatmulFlattenPass继承PatternFusionPass。重写基类PatternFusionPass中的 3 个函数Patterns定义匹配模板用于在整图中获取与该模板相同的拓扑。pattern-CaptureTensor()捕获 BatchMatMul 节点用于读取属性和输入 shape。MeetRequirements对模板匹配到的拓扑进行筛选。检查仅存在 x1/x2 两个有效输入输入 A 为 3 维、输入 B 为 2 维、两个输入 dtype 相同且为浮点类型、offset_x0、transpose 属性为 false。Replacement定义替换部分。构建 ReshapeMatMulV2Reshape 替换图根据 shape 是否动态选择 Const 或动态计算 reshape 目标 shape。使用InferShapeAndCheckSupport验证替换图正确性。注册BatchMatmulFlattenPass为自定义融合 pass执行阶段为 AfterInferShape。程序编译假设 CANN 软件包的安装目录为 INSTALL_PATH例如/home/HwHiAiUser/Ascend/。配置环境变量。运行软件包中设置环境变量脚本命令如下source ${ASCEND_PATH}/set_env.sh${ASCEND_PATH}为 CANN 软件包安装目录下的 cann 路径。请替换相关软件包的实际安装路径例如${INSTALL_PATH}/cann。根据实际情况修改当前目录CMakeLists.txt文件中的如下信息。ASCEND_PATH可以设置默认的软件包路径如果通过set_env.sh设置了$ASCEND_HOME_PATH无需修改。PASS_SO_DIR可以设置自定义融合 pass 动态库安装目录名默认为pass_so_dir。target_include_directories需要包含的头文件对于本示例无需修改。如果是用户自行开发的代码当需要添加头文件时在示例下方直接增加行即可注意不要删除原有项目。如果网络中有自定义算子请增加自定义算子的原型定义头文件。target_link_libraries需要链接的库对于本示例无需修改。如果是用户自行开发的代码当需要添加链接库时在示例下方直接增加行即可注意不要删除原有项目。禁止链接软件包中的其他 so否则后续升级可能会导致兼容性问题。依次执行mkdir build cd build cmake ..执行后在build目录下产生的 es_all_build/generated_code 目录中包含 es 构图 api 的头文件及源码。执行make命令编译自定义 pass so成功编译后通过make install将动态库文件libbatch_matmul_flatten_pass.so安装到自定义融合 pass 目录下。 可以在make后增加可选参数-j$(nproc)用于并行执行构建任务$(nproc)动态获取 CPU 核心数。make -j$(nproc) batch_matmul_flatten_pass make install程序运行配置环境变量如已执行跳过。运行软件包中设置环境变量脚本命令如下source ${ASCEND_PATH}/set_env.sh${ASCEND_PATH}请替换相关软件包的实际安装路径。使用 ATC 离线推理。设置环境变量dump 出编译过程中的模型图export DUMP_GE_GRAPH1进入当前目录data目录执行.py文件导出 onnx文件中使用了 onnx 库运行前请确保安装python gen_onnx.py也可指定 shape 参数导出不同 shape 的模型python gen_onnx.py --batch 32 --m 64 --k 512 --n 256执行结束后在data目录下生成.onnx格式的模型文件名称为model.onnx。执行 ATC 工具命令关于 ATC 工具的详细说明请前往昇腾文档搜索文档ATC离线模型编译工具soc_version请根据实际环境修改atc --model./model.onnx --framework5 --soc_versionxxx --output./model_fused日志中出现如下打印Define pattern for BatchMatmulFlattenPass Define MeetRequirements for BatchMatmulFlattenPass Define replacement for BatchMatmulFlattenPass Created node: Reshape Created node: MatMulV2 Created node: Reshape InferShapeAndCheckSupport success一键式验证可选使用quick_verify.sh脚本可一键完成编译、ATC、dump 图检查和性能测试。脚本中的soc_version默认为Ascend910B3请根据实际环境修改脚本中 atc 命令的--soc_version参数参考使用 ATC 离线推理中的说明cd data ./quick_verify.sh [batch] [m] [k] [n] [test_rounds]默认参数batch32, m64, k512, n256, test_rounds3脚本会自动检查并编译 Pass如未编译检查并编译 benchmark_model生成 ONNX 模型ATC 编译融合模型检查 dump 图验证融合效果运行多轮性能测试清理所有中间文件若融合前 dump 图中没有 BatchMatMul 节点或融合后没有出现 ReshapeMatMulV2Reshape 替换结构脚本会直接退出失败。查看运行结果ATC 工具命令执行完成后目录下生成一系列.pbtxt和.txt文件。 对比以下 dump 图ge_proto_xxxxx_graph_x_PreRunBegin.txt执行前 dump 图应包含 BatchMatMulV2 节点ge_proto_xxxxx_graph_x_RunCustomPass_AfterInferShape.txt执行 InferShape 后的自定义 pass dump 图应包含 ReshapeMatMulV2Reshape 节点不再包含 BatchMatMulV2 节点可以发现模型已按预期优化即 BatchMatMulV2 被 ReshapeMatMulV2Reshape 替换。若未获得预期结果可设置如下环境变量如使用 atc 命令还需添加参数--logdebug让日志打印到屏幕来定位原因export ASCEND_SLOG_PRINT_TO_STDOUT1 export ASCEND_GLOBAL_LOG_LEVEL0【免费下载链接】geGEGraph Engine是面向昇腾的图编译器和执行器提供了计算图优化、多流并行、内存复用和模型下沉等技术手段加速模型执行效率减少模型内存占用。 GE 提供对 PyTorch、TensorFlow 前端的友好接入能力并同时支持 onnx、pb 等主流模型格式的解析与编译。项目地址: https://gitcode.com/cann/ge创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表