博客

为何 PyTorch 编译如此之快:算子融合 (Kernel Fusion)

作者: 2026年5月27日2026年6月2日暂无评论

当你使用 PyTorch 编译器时,你的模型运行速度会大幅提升,最高可达 10 倍。但这到底发生了什么?在未编译的情况下,GPU 会为你代码中的每个 torch 操作运行一个内核(GPU 上的函数)。这造成了两个主要的减速因素:内存数据移动耗时,以及启动每个新内核的开销。每当 GPU 启动一个内核,都需要支付一定的开销成本;而每一个中间结果意味着都需要进行内存的写入和读取。

这就是“融合”发挥作用的地方。PyTorch 的 Inductor 编译器会自动将相关的操作组合成高效的单一 Triton 内核。这使得数据能够保留在靠近寄存器的快速内存中,并减少了内核开销。在本文中,我们将看一个具体的融合示例,并概述进一步阅读的主题。你将直观地看到 torch.compile 是如何将你的 PyTorch 操作转换为优化后的 GPU 代码的。

为了从本文中获得最大收获,你应该具备 PyTorch 的基本熟悉度,并对 GPU 编程概念有大致的了解。

什么是垂直融合 (Vertical Fusion)?

可以将垂直融合想象成一种“链接”步骤的方法,使得一个步骤的输出直接进入下一个步骤。之所以称为“垂直”,是因为如果你想象计算图,这些操作是垂直堆叠的——每一步都依赖于前一步的结果。

这是深度学习中最常见的融合模式,因为神经网络本质上是一连串的操作:归一化、线性层、激活函数等等。它的巨大优势在于消除了中间结果——那些临时张量永远不需要被写入全局内存或从中读取。它们保留在快速寄存器中,GPU 可以更快地访问它们。

让我们深入了解一个垂直融合的示例,即逐点融合 (Pointwise Fusion)。

逐点融合示例

逐点操作是作用于每个元素的简单数学内核:加法、乘法、激活函数等。让我们看看在神经网络层中可能看到的模式:

逐点 PyTorch 示例

import torch

def pointwise_example(x, w, b):
    # Multiple element-wise operations
    tmp = x * w        # multiply
    tmp = tmp + b      # add
    tmp = tmp.sigmoid() # sigmoid activation
    return tmp

未融合:三个独立的内核

在没有融合的情况下,Inductor 会创建三个独立的 Triton 内核。如果 Triton 语法看起来令人生畏,请不要担心。重点不在于死记硬背语法,而是理解其模式:每个内核加载数据、执行一项操作并写入结果。

内核 1:乘法

@triton.jit
def mul_kernel(in_ptr0, in_ptr1, out_ptr0, xnumel, XBLOCK: tl.constexpr):
    xoffset = tl.program_id(0) * XBLOCK
    xindex = xoffset + tl.arange(0, XBLOCK)[:]
    xmask = xindex < xnumel
    x0 = xindex
    tmp0 = tl.load(in_ptr0 + x0, xmask)
    tmp1 = tl.load(in_ptr1 + x0, xmask)
    tmp2 = tmp0 * tmp1
    tl.store(out_ptr0 + x0, tmp2, xmask)

为简洁起见,我们仅包含后续内核的签名,因为它们几乎完全相同。完整源代码请参阅我们的 Git 仓库

内核 2:加法

@triton.jit
def add_kernel(in_ptr0, in_ptr1, out_ptr0, xnumel, XBLOCK: tl.constexpr)

内核 3:Sigmoid

@triton.jit
def sigmoid_kernel(in_ptr0, out_ptr0, xnumel, XBLOCK: tl.constexpr)

在这三个内核中,你总共执行了八次内存操作:为乘法读取两次输入,为加法读取乘法结果和偏置,为 Sigmoid 读取加法结果,并写入所有三个结果。这会产生大量的内存流量。

已融合:一个内核

通过融合,torch.compile 创建了一个单一内核:

内核 4:已融合

@triton.jit
def triton_poi_fused_add_mul_sigmoid_0(in_ptr0, in_ptr1, in_ptr2,
                                        out_ptr0, xnumel, XBLOCK: tl.constexpr):
    xoffset = tl.program_id(0) * XBLOCK
    xindex = xoffset + tl.arange(0, XBLOCK)[:]
    xmask = xindex < xnumel
    x0 = xindex

    # Load all inputs once
    tmp0 = tl.load(in_ptr0 + (x0), xmask)
    tmp1 = tl.load(in_ptr1 + (x0), xmask)
    tmp3 = tl.load(in_ptr2 + (x0), xmask)

    # Fused pointwise operations: mul -> add -> sigmoid
    tmp2 = tmp0 * tmp1
    tmp4 = tmp2 + tmp3
    tmp5 = tl.sigmoid(tmp4)

    # Store final result only
    tl.store(out_ptr0 + (x0), tmp5, xmask)

注意区别:我们一次性加载所有输入,连续执行所有三个操作,并且仅存储最终结果。中间值(tmp2tmp4)保留在寄存器中——这是 GPU 上最快的内存。它们从未触及较慢的全局内存。

收益

  • 内核启动次数:从 3 次减少到 1 次
  • 中间缓冲区:消除了 2 个(乘法结果和加法结果)
  • 内存带宽:读取 5 个完整张量并写入 3 个完整张量(8 次内存操作)减少为读取 3 个张量并写入 1 个(4 次内存操作)——内存流量降低了 50%

其他融合类型

逐点融合只是垂直融合的一种类型。Inductor 使用其他形式的垂直融合来保持 GPU 的高效运行:

归约融合 (Reduction Fusion):将最大值、平均值或求和等归约操作与其前后的操作相结合。这对于批归一化 (Batch Normalization) 等操作至关重要。

GEMM + 后序融合 (Epilogue Fusion):将简单的数学运算附加到繁重的矩阵计算末尾。无需执行矩阵乘法、将结果写入内存,然后再读回以添加偏置并应用 ReLU,偏置和激活函数在矩阵乘法后立即在同一个内核中完成。

前序融合 (Prologue Fusion):后序融合的反向操作——在数据加载时进行预处理。例如,在矩阵乘法之前对输入进行归一化,可以在数据进入时实时完成。

除了作为最显著融合类型的垂直融合外,Inductor 还使用水平融合。

水平融合 (Horizontal Fusion):同时对相同输入运行多个独立操作。例如,在单个内核中同时计算 sin(x)cos(x),只需加载一次 x,而不是两次。

入门:在你的代码中查看融合

让我们通过一个使用归约模式的完整示例来演示。

第 1 步:创建一个简单的归约示例

创建一个名为 fusion_example.py 的文件:

import torch

def reduction_example(x):
    # Pointwise operation followed by reduction
    tmp = x * 2.0
    result = tmp.sum(dim=-1)
    result = result + 1.0
    return result

# Create test input
x = torch.randn(1024, 1024, device='cuda')

compiled_fn = torch.compile(reduction_example)
result_fused = compiled_fn(x)

第 2 步:查看生成的代码

使用 TORCH_LOGS 环境变量运行脚本,以查看 Inductor 生成的内容:

TORCH_LOGS="output_code" python fusion_example.py

这将把生成的 Triton 内核输出到你的终端。寻找名为类似 triton_per_fused_add_mul_sum_0 的内核。per 前缀表示“归约处理”内核,名称表明加法、乘法和求和都被融合在了一起。

结论

融合是 torch.compile 执行的最重要的优化之一。通过将相关操作链接到单一内核中,它减少了内存流量和内核开销——这些通常是 GPU 工作中的主要性能瓶颈。

尝试使用 torch.compile 加速你自己的代码吧。无需更改实现方式,只需添加一个 torch 编译器装饰器,剩下的就交给编译器去完成。

了解更多:PyTorch 文档 pytorch.org/docs/stable/torch.compiler.html 提供了关于编译和优化策略的完整指南。参考我们的 Git 仓库以获取完整源代码。