PyTorch oneDNN Graph API 桥接:JIT 图融合器(LLGA)原理与实战指南
【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch
导读
本文聚焦 PyTorch 官方仓库中基于 oneDNN Graph API 实现的 JIT 图融合器——即位于torch/csrc/jit/codegen/onednn的 oneDNN Graph API Bridge(其前身名为 LLGA,Low Level Graph API)。该模块通过一组自定义 JIT 优化 Pass 将 TorchScript 图中的 PyTorch 算子映射到 oneDNN Graph 后端算子,实现 Conv、Eltwise 等算子的激进式融合,并支持 Float 与 BFloat16 两种精度的 CPU 推理。读完本文,你将掌握 oneDNN Graph 融合的完整流水线(突变消除 → 图构建 → 分区 → 图重写 → 布局传播 → 形状守卫 → 内核执行)、如何开启融合与查看日志、如何编写 Float/BFloat16 推理代码,以及如何向桥接层扩展新算子。
oneDNN Graph API Bridge 是什么
PyTorch 的 JIT 图融合器基于 oneDNN Graph API 构建。oneDNN Graph API 提供了灵活的分区(partition)机制,允许后端(oneDNN 库)对算子进行激进式融合(aggressive fusion),即把多个 PyTorch 算子(如aten::conv2d与aten::relu)合并成一个后端算子一次性执行,从而减少内核启动开销与中间数据读写。
当前实现具备以下关键特性:
- 支持 Float 与 BFloat16 推理。BFloat16 仅在 Intel Xeon Cooper Lake 及更新平台上有原生支持(这些平台内置 BFloat16 指令),因此只有在这些平台上 BFloat16 才能发挥性能优势;
- JIT 与 eager 模式的 AMP 支持存在差异:PyTorch 在 JIT 与 eager 两种模式下对自动混合精度(AMP)的支持是分叉的。因此要使用 BFloat16,应禁用 JIT 的 AMP 支持,转而利用 eager 模式的 AMP 支持(即
torch.cpu.amp.autocast()),详见下文 BFloat16 示例; - 当前仅在静态形状下获得加速:因为融合内核依赖形状特化与形状守卫保证缓存编译结果有效(见 interface.cpp 中的注释),动态形状支持仍在规划中;
- 权重缓存:启用 oneDNN Graph 后,推理期间权重保持不变(constant),因此会被缓存(
dnnl::graph::set_constant_tensor_cache(true),见 interface.cpp),避免重复编译。
从源码结构看,整个桥接层被组织在 torch/csrc/jit/codegen/onednn 目录下,包含interface、graph_helper、graph_fuser、graph_rewriter、layout_propagation、prepare_binary、defer_size_check、guard_shape、kernel、LlgaTensorImpl等模块。代码库中大量出现的 "LLGA" 字样即对应 oneDNN Graph(README 明确说明:oneDNN Graph was formerly known as LLGA,and thus LLGA in the codebase corresponds to oneDNN Graph)。
开启与关闭:torch.jit.enable_onednn_fusion
一切融合的前提是先全局开启 oneDNN Graph 融合开关:
torch.jit.enable_onednn_fusion(True) # 开启 torch.jit.enable_onednn_fusion(False) # 关闭该开关的底层实现在 interface.h 的RegisterLlgaFuseGraph中:
- 开启时通过
registerPrePass把fuser::onednn::fuseGraph注册为JIT 自定义 pre-pass(PassManager<RegisterLlgaFuseGraph>,注册的是registerPrePass而非默认 pass); - 关闭时通过
clearPrePass清除; - 注意
setEnabled中有TORCH_CHECK(AT_MKLDNN_ENABLED(), "Running oneDNN Graph fuser is only supported with MKLDNN builds.")——仅在启用 MKLDNN 的构建中支持,这一点在测试文件里也有对应守卫(见下文「测试」一节)。
开关状态保存在fuser::onednn::onednn_enabled这个std::atomic<bool>中(interface.h)。
Graph Optimization:五大优化 Pass 与真实流水线
README 概述了注册在 PyTorch 自定义 pre-pass 集合中的优化 Pass,核心分为五步:
别名与突变消除(Alias and mutation reduction)oneDNN Graph 的算子都是纯函数式(functional)的,而 PyTorch 算子存在 in-place 形式或通过创建视图(view)共享 buffer。为弥合后端算子与 PyTorch 算子之间的语义差距,流水线一开始会尽力消除突变(mutation)。
图传递(Graph passing)拿到 PyTorch TorchScript 图后,桥接层把图上的 PyTorch 算子映射为对应的 oneDNN Graph 算子,构建后端图。
分区(Partitioning)后端在图中选择可融合的区域并返回分区列表,每个分区对应一组被融合的算子。
图重写(Graph rewriting)基于后端返回的分区,重写原始 PyTorch JIT 图:同一分区内的算子被归组为一个 JIT 算子,即 oneDNN Graph 融合组(
prim::oneDNNFusionGroup)。布局传播(Layout propagation)消除分区边界上不必要的布局转换。桥接层为分区的输出设置不同的格式,让后端在内部完成布局转换:当设置为
ANY时,边界布局完全由后端决定;否则后端需遵循 PyTorch 指定的布局。目前实现中,当一个张量既是一个 oneDNN Graph 分区的输出、又是另一个分区的输入时,为其设置ANY布局(README 原话)。
需要强调的是,README 的五步是精简描述,真实代码中的 Pass 链更细。查看 interface.cpp 的fuseGraph,完整流水线为:
RemoveProfileNodesAndSpecializeTypes → RemoveTensorMutation(含白名单)→ RemoveListMutation → DecomposeSiluForLLGA(针对 silu 的分解) → PrepareBinaryForLLGA(二元算子标量输入处理) → DeferSizeCheck(延后大小检查) → CreateLlgaSubgraphs(构建并分区融合子图) → PropagateLayout(布局传播) → prepareFusionGroupAndGuardOutputs(添加形状守卫) → RemoveTensorTypeSpecializations(清除 IR 中的张量类型特化)其中每个步骤前后都有GRAPH_DUMP日志输出,这正是 Quick Start 中通过PYTORCH_JIT_LOG_LEVEL观察的流水线节点。
几个值得注意的实现细节:
- 只在 profiling 模式下运行:
fuseGraph一开始就检查getProfilingMode(),注释明确说明「依赖形状特化与形状守卫保证 kernel 中缓存编译的有效性,因此仅支持 profiling 模式」; - 突变消除白名单:
RemoveTensorMutation传入的 lambda 白名单包含aten::add_、aten::mul_、aten::tanh_、aten::elu_、aten::relu_、aten::relu6_、aten::gelu_、aten::sqrt_、aten::sigmoid_、aten::hardtanh_、aten::abs_、aten::square_、aten::pow_、aten::leaky_relu_、aten::round_、aten::exp_、aten::hardswish_、aten::silu_(interface.cpp); - 子图构建与清理:
CreateLlgaSubgraphs(graph_fuser.cpp)通过AliasDb维持别名信息,先递归构建全部子图、再递归清理并合并过小的子图,最后做一次全局 CSE 与死代码消除; - 融合守卫算子:
prim::oneDNNFusionGuard在运行期逐个检查输入张量类型是否与编译时的TensorType匹配(interface.cpp)。有趣的是,对于来自上游 LLGA 分区的 mkldnn 张量会直接放行(is_mkldnn检查),因为其形状已在源头校验过——这正是LlgaTensorImpl包装器存在的意义之一。
分区与算子映射
LlgaGraphHelper(graph_helper.h)负责算子→oneDNN Graph 的映射与分区:
createOperator(Node* node)把单个 PyTorch 算子转换为 oneDNN Graph 算子;shouldMerge/shouldConsiderForMerge决定哪些节点可以并入同一分区;OpPartitionMap维护「算子 ID → 分区 ID」的映射,供图重写阶段使用;- 头文件中还定义了
STRIDED_LAYOUT (0)与OPAQUE_LAYOUT (1)两种布局模式常量。
布局与张量描述:LlgaTensorDesc 与 LlgaTensorImpl
张量相关代码位于:
- LlgaTensorImpl.h
- LlgaTensorImpl.cpp
其中LlgaTensorDesc封装了 oneDNN Graph 的logical_tensor描述(tid、sizes、strides、dtype、property_type),并支持四种布局状态:strided、opaque、any与维度未知(DNNL_GRAPH_UNKNOWN_DIM)。它还记录了compute_inplace与input_tensor_index,用于支持分区输出复用输入张量内存的 in-place 计算。
LlgaTensorImpl继承自c10::TensorImpl,把 PyTorch 张量与 oneDNN Graph 张量描述绑定在一起。头文件注释说明了其历史作用:早期 oneDNN Graph 在分区之间使用 blocked 布局,该包装器用于绕过守卫检查;后来 oneDNN Graph 改为在分区之间使用 strided 张量,但包装器仍然有用——因为分区之间张量的 strides 与守卫预期不同,仍需要它绕过守卫检查。Engine与Stream均为单例(当前只有 CPU engine),编译后的分区被提交到 stream 上执行。
Graph Executor:运行期执行路径
运行期,被重写后的 PyTorch JIT 图会把 oneDNN Graph 分区派发给 oneDNN graph JIT 变参算子(prim::oneDNNFusionGroup,注册代码见 interface.cpp)。其执行核心是LlgaKernel(kernel.h):
- 输入映射:把每个分区的输入 PyTorch 张量映射为 oneDNN Graph 张量(
llga_from_aten_tensor); - 编译:对分区调用
compiled_partition compile(...); - 执行:把编译后的分区提交到 stream 执行;
- 输出映射:把输出的 oneDNN Graph 张量映射回 PyTorch 张量,交给 JIT 图上的下一个算子。
LlgaKernel内部的初始化逻辑值得说明(kernel.h 的成员函数注释):
- 由于 PyTorch 会把常量复制到子图内部而非引用,分区的常量输入不再出现在
graph->inputs()中,因此initializeConstantInputs需要借助从分区取回的 tensor id 找回缺失的常量输入; nPartitionInputs_ = nGraphInputs_ + constantInputs_.size():分区实际输入数等于图输入数加上常量输入数;initialize使用c10::once_flag保证只初始化一次,并缓存compilation_——这与 README 所述「推理期间权重被缓存」相呼应;- 每个内核的调试名形如
LlgaPartition_<id>(genDebugName),配合RECORD_FUNCTION可用于性能剖析。
测试:如何验证融合行为
README 给出运行 LLGA 融合器测试套件的命令:
pytest test/test_jit_llga_fuser.py测试文件 test/test_jit_llga_fuser.py 中的关键设定:
LLGA_FUSION_GROUP = 'prim::oneDNNFusionGroup',测试通过断言图中该节点数量来验证融合是否正确发生;LLGA_NOT_ENABLED = not torch.backends.mkldnn.is_available() or IS_WINDOWS or IS_MACOS——MKLDNN 不可用或 Windows/macOS 平台会跳过测试;setUp中执行torch._C._jit_set_autocast_mode(False)与torch.jit.enable_onednn_fusion(True),tearDown恢复原状——再次印证「JIT 与 eager AMP 分叉,BF16 需禁用 JIT AMP」这一约束;- BF16 分支使用
torch.autocast(device_type="cpu", cache_enabled=False, dtype=torch.bfloat16)做 eager 模式 AMP 追踪,随后torch.jit.freeze冻结模型并做 warmup; assertFused(graph, patterns)断言某些 aten 算子(如aten::relu)在融合后不再出现在图中(被吸收进融合组);- 测试类上方有
@unittest.skipIf(IS_AVX512_UNSUPPORTED, "This test fails for BF16 on machines without AVX512.")——BF16 测试需要支持 AVX512 的机器,这与 README 所述「BF16 需要 Cooper Lake 及更新平台」一致。
例如test_conv2d_eltwise(test/test_jit_llga_fuser.py)构建 Conv → Eltwise → Conv → Eltwise 模型,分别用relu / leaky_relu / sigmoid / square / abs / exp / hardswish / tanh / hardtanh(含 in-place 变体)验证:
- 断言图中恰好出现 2 个
oneDNNFusionGroup; - 断言
relu_等 in-place 算子已被突变消除 Pass 替换为普通形式(assertFused(graph, ['aten::' + eltwise_fn_name])); - 断言 eltwise 被融合进融合组(
assertFused(graph, ['aten::' + eltwise]))。
Quick Start:查看完整融合流水线日志
README 提供了一个「级联 Conv-Relu」示例(对应测试test_conv2d_eltwise),并建议开启日志输出以熟悉整条流水线。README 给出的流水线为:
Mutation Removal → Prepare Binary → Defer Size Check → Graph Fuser → Layout Propagation → Type Guard → Kernel Execution
启动命令:
DNNL_VERBOSE=1 PYTORCH_JIT_LOG_LEVEL=">>graph_helper:>>graph_fuser:>>kernel:>>interface" python -u test/test_jit_llga_fuser.py -k test_conv2d_eltwise各环境变量的作用:
DNNL_VERBOSE=1:开启 oneDNN(原 DNNL)库自身的 verbose 输出,可以看到每个被融合后执行的 oneDNN 原语及其耗时;PYTORCH_JIT_LOG_LEVEL=">>graph_helper:>>graph_fuser:>>kernel:>>interface":开启 PyTorch JIT 日志系统中对应模块的GRAPH_DUMP输出(这些GRAPH_DUMP调用分散在fuseGraph的每一步前后,以及graph_helper、kernel的运行路径上),用于观察每一步 Pass 前后图的变化;-k test_conv2d_eltwise:pytest 按关键字过滤,只运行该融合用例;-u:关闭 pytest 的捕获,让上述日志实时打印到终端。
代码库结构与算子扩展指南
README 给出了清晰的源码地图:
- 桥接层主体源码位于
torch/csrc/jit/codegen/onednn/*; - 张量相关代码位于 LlgaTensorImpl.h 与 LlgaTensorImpl.cpp;
- 桥接代码的 CMake 入口:
caffe2/CMakeLists.txt; - oneDNN Graph 子模块及其查找/依赖配置:
third_party/ideep/mkl-dnn、cmake/public/mkldnn.cmake、cmake/Modules/FindMKLDNN.cmake、cmake/Dependencies.cmake。
如何把一个新算子映射到 oneDNN Graph(README 明确给出的三步):
- 在 graph_helper.cpp 的
createOperator中为该算子添加映射条目; - 若该算子有 in-place 变体,需把它加入 interface.cpp 中传给
RemoveTensorMutation的 lambda 白名单; - 若希望该算子参与融合决策,还应把它加入 register_interface.cpp 的
canFuseNode。
需要说明的是,除这三步外,实际还涉及算子是否支持静形状/特定 dtype 的判断逻辑,扩展前建议先对照graph_helper.cpp中已有算子(如 Conv、Eltwise、Binary、Pool、LayerNorm、Cat 等)的映射写法,并运行上文测试套件验证。
实战示例一:Float 推理
README 给出的 Float 推理模板如下(完整继承):
# 全局开启 oneDNN graph 融合 torch.jit.enable_onednn_fusion(True) # 定义模型 def MyModel(torch.nn.Module): ... # 构造模型 model = MyModel(...) with torch.no_grad(): model.eval() model = torch.jit.trace(model, torch.rand(args.batch_size, 3, 224, 224)) # 运行模型 with torch.no_grad(): # oneDNN graph 融合将在运行时被触发 output = model(images)要点:
- 必须先调用
torch.jit.enable_onednn_fusion(True); - 使用
torch.jit.trace得到 TorchScript 图(融合 Pass 注册在 JIT pre-pass 中,trace/freeze 过程中即会执行); - 推理放在
torch.no_grad()与model.eval()下执行,这与 README 所述「权重作为常量被缓存、推理期间不变」的前提一致。
实战示例二:BFloat16 推理
README 给出的 BFloat16 模板如下(完整继承,含其关键注释):
# 假设已有一个名为 'model' 的模型 example_input = torch.rand(1, 3, 224, 224) # 开启 oneDNN Graph torch.jit.enable_onednn_fusion(True) # 禁用 JIT 的 AMP torch._C._jit_set_autocast_mode(False) with torch.no_grad(), torch.cpu.amp.autocast(): model = torch.jit.trace(model, (example_input)) model = torch.jit.freeze(model) # 2 次 warm-up(带示例输入做 trace/script 时为 2 次,无示例输入时为 3 次) model(example_input) model(example_input) # 后续运行中可观察到加速 model(example_input)解读与注意事项:
- 为什么禁用 JIT AMP:PyTorch 对 JIT 与 eager 模式的 AMP 支持存在分叉(divergent),测试文件的注释引用了 PyTorch issue #75956,说明必须关闭 JIT 的 autocast(
torch._C._jit_set_autocast_mode(False)),改用 eager 模式的torch.cpu.amp.autocast()来产生 BFloat16 计算; - 平台限制:BFloat16 只有在 Intel Xeon Cooper Lake 及更新平台(具备原生 BFloat16 支持)上才表现良好;测试套件中 BF16 用例也会在无 AVX512 的机器上跳过;
- warm-up 次数:带示例输入做 trace/script 需要 2 次 warm-up,不带示例输入需要 3 次——第一次运行会触发编译与缓存,加速效果在后续运行中体现;
- 静态形状前提:由于编译结果按形状特化并被缓存,请保证推理输入形状与 trace 时的形状一致,否则会命中形状守卫回退逻辑(测试
test_typecheck演示了改变输入形状后traced(x)仍与 eager 结果一致的回退路径)。
小结
oneDNN Graph API Bridge(LLGA)是 PyTorch JIT 在 CPU 推理侧的一条重要融合路径:它以自定义 pre-pass 的形式接管 TorchScript 图的优化,通过突变消除、算子映射、后端分区、图重写与布局传播五类核心 Pass(实际实现为更细的九步流水线)把可融合算子归组为oneDNNFusionGroup,运行期由LlgaKernel完成 oneDNN Graph 分区的编译与执行,并以形状守卫 + 常量缓存保证编译结果可复用。使用上记住三个关键约束即可快速上手:开启torch.jit.enable_onednn_fusion(True)、BFloat16 场景下禁用 JIT AMP 并改用 eager autocast、保持推理输入为静态形状。
【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考