CANN/ge:基于Pattern匹配实现Pass

发布时间:2026/9/10 20:39:18
CANN/ge:基于Pattern匹配实现Pass 基于Pattern匹配实现Pass【免费下载链接】geGEGraph Engine是面向昇腾的图编译器和执行器提供了计算图优化、多流并行、内存复用和模型下沉等技术手段加速模型执行效率减少模型内存占用。 GE 提供对 PyTorch、TensorFlow 前端的友好接入能力并同时支持 onnx、pb 等主流模型格式的解析与编译。项目地址: https://gitcode.com/cann/ge为提高自定义融合Pass的开发效率本章提供了一组基于图结构的匹配与替换接口用于实现Pass的构建。整个融合行为的逻辑可分为以下三个步骤匹配通过Pattern来定义一种图结构进行子图结构的匹配查找。决策根据匹配结果结合更具体的条件判断该匹配是否满足融合要求。替换在确认可融合后执行图结构的替换操作完成优化。图 1逻辑架构图 ![图示](https://raw.gitcode.com/cann/ge/raw/243ea8d2d8f7623dd210c0867dde5782ec5594a9/docs/zh/user_guides/graph_dev/figures/logical_architecture.png 逻辑架构图?utm_sourcegitcode_repo_files)逻辑架构如上图所示相关核心概念解释如下Pattern模式在图匹配过程中Pattern用于描述子图结构特征的模板或规则通过图匹配算法在Graph中查找符合特定规则的子图。PatternMatcher匹配器执行匹配算法的核心对象它接收一个Pattern并按照Pattern中定义的子图结构去Graph中查找符合定义的子图。GraphRewriter重写器执行改图的核心对象接收匹配到的子图边界及目标Replacement图替换后将原子图节点替换为Replacement中的节点结构完成图的重构。开发者通过继承GE提供的基类并重写其方法来实现自定义融合Pass 接着调用注册宏将Pass注册到指定阶段根据应用场景的不同GE提供两类基类供继承通用子图融合1:1或复杂拓扑替换场景适用于需匹配完整子图结构并整体替换为另一子图的场景继承PatternFusionPass类实现自定义融合Pass类并通过REG_FUSION_PASS注册宏将Pass注册到指定阶段。单节点替换单节点替换为N个节点场景继承DecomposePass类实现自定义融合Pass类并通过REG_DECOMPOSE_PASS注册宏将Pass注册到指定阶段。两种场景在使用方法上类似区别在于匹配对象的粒度子图/单节点与替换方式复杂拓扑替换/多节点展开下面将详细展开介绍。场景介绍通用子图融合1:1或复杂拓扑替换场景融合Pass开发本节首先对该场景下涉及的核心数据结构PatternFusionPass进行介绍其次介绍开发过程中需要重写的3个函数Patterns、MeetRequirements与Replacement最后介绍如何将Pass注册到指定阶段。PatternFusionPass类声明如下class PatternFusionPass : public FusionBasePass { public: Status Run(GraphPtr graph, CustomPassContext pass_context) override; protected: virtual std::vectorPatternUniqPtr Patterns() 0; virtual bool MeetRequirements(const std::unique_ptrMatchResult match_result); virtual GraphUniqPtr Replacement(const std::unique_ptrMatchResult match_result) 0; };Run函数调用Patterns获取模板的拓扑Pattern将Pattern在目标Graph中逐一匹配调用MeetRequirements对匹配到的Pattern作出是否需要被替换的判断 最后通过Replacement获取目标结构将满足替换条件的Pattern进行替换。开发者通过继承PatternFusionPass类并重写Patterns、MeetRequirements与Replacement函数实现自定义融合Pass的开发。函数介绍如下|函数|说明|是否必须重写| |--|--|--| |Patterns|定义在目标图中匹配的模板拓扑返回一个或多个图结构指针。|是| |MeetRequirements|对Patterns匹配到的图结构按条件进行过滤输入匹配结果返回布尔值。|否默认直接返回true| |Replacement|定义替换结构输入匹配结果返回图指针。|是|PatternsPatterns用于定义目标图中匹配的一个或多个模板拓扑使用EsGraphBuilder构建一张DAGDirected Acyclic Graph有向无环图图来表达Pattern如下所示std::vectorPatternUniqPtr Patterns() override { std::vectorPatternUniqPtr patterns; // 使用EsGraphBuilder构建pattern auto graph_builder es::EsGraphBuilder(pattern); // 此处定义pattern // ... // 初始化Pattern对象xxx请替换为实际输出节点 auto graph graph_builder.BuildAndReset({xxx}); auto pattern std::make_uniquePattern(std::move(*graph)); patterns.emplace_back(std::move(pattern)); // 可以继续向patterns中添加多个pattern // ... return patterns; }其中EsGraphBuilder为图构建器类用于构建计算图。 推荐开发者使用ES接口进行Pattern的定义其提供了定义输入、常量与算子等接口以下是使用ES API定义一个ReLu单算子pattern的示例std::vectorPatternUniqPtr patterns; // 创建一个EsGraphBuilder实例用于构建计算图图的名称为pattern auto graph_builder es::EsGraphBuilder(pattern); auto data graph_builder.CreateInput(0); auto relu es::Relu(data); // 构建并重置图将{relu}作为输出节点 auto graph graph_builder.BuildAndReset({relu}); // 将graph移动到Pattern构造函数中创建pattern对象 auto pattern std::make_uniquePattern(std::move(*graph)); patterns.emplace_back(std::move(pattern));[!NOTE]说明用于匹配的Pattern需要满足自包含除了边界的输出算子边界内所有算子的数据输出消费者都要在边界内非自包含的Pattern不会被匹配。除了上文中提到的使用EsGraphBuilder构建PatternGE还提供了两种接口实现对Pattern更细粒度的定义CaptureTensor定义过程中可以捕获Pattern中的一个Tensor从而在MatchResult中可以按序获取。方法声明如下入参node_output为NodeIo类型由节点与索引组成表示为某个节点的某个输出。// CaptureTensor声明 Pattern CaptureTensor(const NodeIo node_output); // NodeIo结构体 struct NodeIo { GNode node; int64_t index; };调用CaptureTensor捕获data示例如下std::vectorPatternUniqPtr patterns; // 创建一个EsGraphBuilder实例用于构建计算图图的名称为pattern auto graph_builder es::EsGraphBuilder(pattern); auto data graph_builder.CreateInput(0); auto relu es::Relu(data); // 构建计算图构建的图仅包含data - ReLU(relu)的结构 auto graph graph_builder.BuildAndReset({relu}); // 创建一个Pattern实例用构建好的图初始化 auto pattern std::make_uniquePattern(std::move(*graph)); // 调用CaptureTensor捕获data pattern-CaptureTensor({*relu.GetProducer(), 0}) patterns.emplace_back(std::move(pattern));PatternMatcherConfig构造自定义Pass可以传入PatternMatcherConfig以开启Const值匹配功能以及IR属性及其值的匹配能力。基类PatternFusionPass构造函数如下explicit PatternFusionPass(std::unique_ptrPatternMatcherConfig match_config);使用PatternMatcherConfigBuilder来构造PatternMatcherConfig类PatternMatcherConfigBuilder提供两个函数作为匹配能力的开关EnableConstValueMatch开启Const值匹配在匹配过程中将对Pattern中定义的Const/Constant进行值的匹配值相等才认为匹配成功。EnableIrAttrMatch开启IR属性及其值匹配Pass将在Pattern匹配过程中对Pattern中节点上携带的IR属性的数量和值进行匹配。以下为名为CustomFusionPass的自定义Pass类打开Const值匹配的构造函数示例explicit CustomFusionPass() : PatternFusionPass(PatternMatcherConfigBuilder().EnableConstValueMatch().Build()) {}MeetRequirements对于Patterns获取到的匹配结果在MeetRequirements中进行筛选。 从上文Run函数的实现中可以看到每个MatchResult类型的匹配结果作为MeetRequirements的入参通过MatchResult开发者可以获取匹配结果的信息进行筛选最后返回的布尔值作为是否替换该匹配结果的依据如下所示bool MeetRequirements(const std::unique_ptrMatchResult match_result) override { // 可以使用传入的match_result对匹配结果进行筛选 // 满足条件返回true if (IsSatisfy(match_result)) { return true; } // 不满足条件返回false return false; }MatchResult是匹配结果类包含匹配结果的节点、连边等信息。开发者可以使用MatchResult成员函数获取匹配结果的相关信息以进行筛选以下是使用GetCapturedTensor成员函数校验ReLu输出是否为动态shape的示例NodeIo relu_output; // 尝试从match_result中获取第一个捕获的输出张量存储到relu_output if(match_result-GetCapturedTensor(0,relu_output) ! GRAPH_SUCCESS){ return false; } TensorDesc relu_out_tensor_desc; // 从relu_output中获取输出张量描述信息 relu_output.node.GetOutputDesc(relu_output.index, relu_out_tensor_desc); if (relu_out_tensor_desc.GetShape().GetShapeSize() ! -1){ return false; } return true;ReplacementReplacement中定义目标结构替换与Patterns中匹配且MeetRequirements为true的部分。与Patterns一样使用EsGraphBuilder定义结构此处不再赘述GraphUniqPtr Replacement(const std::unique_ptrMatchResult match_result) override { auto replacement_graph_builder es::EsGraphBuilder(replacement); // 此处定义替换结构 // ... return replacement_graph_builder.BuildAndReset({r_a}); }[!NOTE]说明如果Pass注册阶段在InferShape后需要在Replacement中自行调用GeUtils::InferShape此外如果要使用GeUtils::CheckNodeSupportOnAicore判断目标结构是否支持该函数的调用需要在InferShape之后。注册自定义融合Pass完成对融合pass的定义后需要使用注册宏REG_FUSION_PASS将其注册到对应阶段如下是将名为CustomFusionPass的自定义Pass注册到kBeforeInferShape阶段的示例REG_FUSION_PASS(CustomFusionPass).Stage(CustomPassStage::kBeforeInferShape);各阶段详细说明请参见Stage。单节点替换单节点替换为N个节点场景融合Pass开发该场景下的Pass继承的基类为DecomposePass。由于被替换结构是单个节点此处Pattern不再需要通过Patterns定义而是在构造函数中直接传入算子类型如下所示class CustomOne2NPass : public DecomposePass { public: CustomOne2NPass(const std::vectorAscendString op_types) : DecomposePass(op_types) {} };与一般场景类似继承自DecomposePass的Pass也需要重写MeetRequirements与Replacement但两方法的入参类型不再是MatchResult而是GNode即通过构造时传入的op_types在图中匹配到的节点。bool MeetRequirements(const GNode matched_node) override { ... } GraphUniqPtr Replacement(const GNode matched_node) override { ... }注册自定义融合Pass如下是使用注册宏REG_DECOMPOSE_PASS将Conv2D作为op_types初始化CustomOne2NPass并将其注册在kAfterInferShape的示例REG_DECOMPOSE_PASS(CustomOne2NPass, {Conv2D}).Stage(CustomPassStage::kAfterInferShape);开发示例此处以MatMulAdd结构融合为GEMM自定义Pass为例对应上述一般场景下的融合Pass开发详细介绍如何通过自定义融合Pass修改Graph详细可以参见样例源码。样例仓还提供了更多样例用户可以单击融合Pass样例进行查看。修改前后的图结构如下本例识别图中左边的MatMulAdd结构并通过图修改接口替换为右边的单个GEMM节点// |o----------------------------------- // |o a b // |o \ / a b c // |o MatMul c \ | / // |o \ / GEMM // |o Add // |o-----------------------------------包含的头文件。#include iostream // 自定义融合Pass接口头文件 #include ge/fusion/pass/pattern_fusion_pass.h // ES接口头文件 #include es_all_ops.h使用自定义Pass修改Graph。class FuseMatMulAndAddPass : public PatternFusionPass { protected: // 重写Patterns std::vectorPatternUniqPtr Patterns() override { std::cout Define pattern for FuseMatMulAndAddPass std::endl; std::vectorPatternUniqPtr patterns; // 创建一个EsGraphBuilder实例用于构建计算图图的名称为pattern0 auto graph_builder0 es::EsGraphBuilder(pattern0); auto [a0, b0, c0] graph_builder0.CreateInputs3(); auto matmul0 es::MatMul(a0, b0); auto add0 es::Add(matmul0, c0); // 构建并重置图 auto graph0 graph_builder0.BuildAndReset({add0}); auto pattern0 std::make_uniquePattern(std::move(*graph0)); patterns.emplace_back(std::move(pattern0)); return patterns; } // 重写Replacement GraphUniqPtr Replacement(const std::unique_ptrMatchResult match_result) override { std::cout Define replacement for FuseMatMulAndAddPass std::endl; // 构建替换后的图 auto replace_graph_builder es::EsGraphBuilder(replacement); auto [r_a, r_b, r_c] replace_graph_builder.CreateInputs3(); auto alpha_const replace_graph_builder.CreateScalar(1); auto beta_const replace_graph_builder.CreateScalar(1); auto gemm es::GEMM(r_a, r_b, r_c, alpha_const, beta_const); // 构建并重置图 return replace_graph_builder.BuildAndReset({gemm}); } };注册自定义融合Pass。// 使用REG_FUSION_PASS注册宏进行改图Pass注册并指定被调用的阶段 REG_FUSION_PASS(FuseMatMulAndAddPass).Stage(CustomPassStage::kBeforeInferShape);如何使用自定义Pass完成上述自定义Pass后本节简单介绍如何把改图函数编译成动态库插件方式以便注册的Pass在图编译阶段被框架调用。详细使用说明请参见样例使用指导。把开发示例中的改图函数编译成仅以.so结尾的动态库文件。编译成功后执行make install命令将上述.so动态库文件安装到${INSTALL_DIR}/opp/vendors/xxx/custom_fusion_passes/目录下。支持设置软链接的方式.so文件对执行用户需要有可读权限多个${INSTALL_DIR}/opp/vendors/xxx目录按照字母序排序后遍历寻找custom_fusion_passes/子目录单个子目录内的.so按照字母序加载非.so结尾的文件在加载时跳过。其中${INSTALL_DIR}请替换为CANN软件安装后文件存储路径。以root用户安装为例安装后文件默认存储路径为/usr/local/Ascend/cann。xxx有且仅有一层自定义目录。custom_fusion_passes该目录下不能有子目录。支持但不限于如下几种入口编译模型文件如果要查看上述自定义Pass有没有生效在编译模型前需要dump图进行查看在执行之前设置DUMP_GE_GRAPH详细说明请参见《环境变量参考》环境变量然后使用如下入口编译模型使用ATC工具进行模型转换。ATC工具使用方法请参见《ATC离线模型编译工具》。编译Graph为离线模型。编译并运行Graph。结果验证请参见样例使用指导程序运行查看运行结果。设置了dump环境变量后程序执行完毕会在当前路径生成ge_onnx*.pbtxt等图文件用户可以获取如下两张图以指定Pass执行阶段在InferShape之前为例然后使用Netron等可视化软件查看ge_onnx_xxxx_PreRunBegin.pbtxt融合前的图ge_onnx_xxxx_RunCustomPassBeforeInfershape.pbtxt融合后的图查看融合前的图结构为通过自定义Pass修改后的图结构如下所示可以看出MatMulAdd结构已经替换为单个GEMM节点。【免费下载链接】geGEGraph Engine是面向昇腾的图编译器和执行器提供了计算图优化、多流并行、内存复用和模型下沉等技术手段加速模型执行效率减少模型内存占用。 GE 提供对 PyTorch、TensorFlow 前端的友好接入能力并同时支持 onnx、pb 等主流模型格式的解析与编译。项目地址: https://gitcode.com/cann/ge创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考