用 PyTorch Profiler 拆解注意力计算:从算子耗时到显存峰值

2026-07-10 36 预计阅读时间: 1 分钟
来源: huggingface.co AI 摘要 Original link

Disclaimer: This article is an AI-assisted summary. Read it together with the original source when precision matters. The summary may omit context, version differences, or edge cases and is not official documentation.

预计阅读时间:11 分钟

注意力层经常同时承受计算、显存和数据搬运压力。只看一次前向传播的总耗时,很难判断瓶颈究竟来自矩阵乘法、Softmax、中间张量,还是内核调度。更有效的办法是用 PyTorch Profiler 把注意力计算拆到算子级,并结合输入形状、显存分配和执行时间做判断。

本文基于“分析 PyTorch 中的注意力计算”这一主题给出一套可直接实践的方法。由于来源摘要没有提供具体实验数据,下面的代码和判断流程属于可复现的工程示例,不代表原文给出的基准结果。

注意力为什么值得单独分析

标准缩放点积注意力可以写成:

scores = Q @ K^T / sqrt(d)
probabilities = softmax(scores)
output = probabilities @ V

假设查询和键的序列长度都是 N,注意力分数张量的形状通常为 [batch, heads, N, N]。这意味着序列长度翻倍时,分数矩阵的元素数量约增至四倍。实际运行中需要关注的不只是 FLOPs,还包括:

  • Q @ K^TP @ V 两次矩阵乘法的设备时间。
  • Softmax 以及缩放、掩码等逐元素操作是否产生大量独立内核。
  • [N, N] 中间结果是否被显式物化,从而推高显存峰值。
  • 张量布局、数据类型和形状是否让框架选择了不同的执行路径。
  • 首次运行中的内核加载、缓存建立等一次性成本是否污染了测量。

因此,分析注意力不能只在函数两端放一个计时器。算子表适合回答“时间花在哪里”,时间线适合回答“这些工作以什么顺序发生、设备是否存在空闲”。

一个可运行的 Profiler 实验

下面的脚本对比手工实现的注意力与 torch.nn.functional.scaled_dot_product_attention。运行前按机器显存调整 BATCHHEADSSEQ_LENHEAD_DIM;没有 CUDA 时,脚本会自动使用 CPU。

import math
from pathlib import Path

import torch
import torch.nn.functional as F


def manual_attention(q, k, v):
    scale = 1.0 / math.sqrt(q.size(-1))
    scores = torch.matmul(q, k.transpose(-2, -1)) * scale
    probabilities = torch.softmax(scores, dim=-1)
    return torch.matmul(probabilities, v)


def sdpa_attention(q, k, v):
    return F.scaled_dot_product_attention(
        q,
        k,
        v,
        dropout_p=0.0,
        is_causal=False,
    )


def profile_attention(name, operation, q, k, v, device):
    trace_dir = Path('traces') / name
    trace_dir.mkdir(parents=True, exist_ok=True)

    activities = [torch.profiler.ProfilerActivity.CPU]
    if device.type == 'cuda':
        activities.append(torch.profiler.ProfilerActivity.CUDA)

    schedule = torch.profiler.schedule(
        wait=1,
        warmup=1,
        active=3,
        repeat=1,
    )

    with torch.profiler.profile(
        activities=activities,
        schedule=schedule,
        on_trace_ready=torch.profiler.tensorboard_trace_handler(str(trace_dir)),
        record_shapes=True,
        profile_memory=True,
        with_stack=False,
    ) as profiler:
        for _ in range(5):
            with torch.profiler.record_function(name):
                output = operation(q, k, v)
                output.square().mean().backward()
            profiler.step()

    sort_key = 'self_cuda_time_total' if device.type == 'cuda' else 'self_cpu_time_total'
    print(f'\n=== {name} ===')
    print(profiler.key_averages().table(sort_by=sort_key, row_limit=15))


def main():
    torch.manual_seed(7)
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    dtype = torch.float16 if device.type == 'cuda' else torch.float32

    BATCH = 2
    HEADS = 8
    SEQ_LEN = 512
    HEAD_DIM = 64
    shape = (BATCH, HEADS, SEQ_LEN, HEAD_DIM)

    print(f'device={device}, dtype={dtype}, shape={shape}')

    for name, operation in [
        ('manual_attention', manual_attention),
        ('scaled_dot_product_attention', sdpa_attention),
    ]:
        q = torch.randn(shape, device=device, dtype=dtype, requires_grad=True)
        k = torch.randn(shape, device=device, dtype=dtype, requires_grad=True)
        v = torch.randn(shape, device=device, dtype=dtype, requires_grad=True)
        profile_attention(name, operation, q, k, v, device)


if __name__ == '__main__':
    main()

保存为 profile_attention.py 后,可以这样运行:

python profile_attention.py

查看时间线需要安装并启动 TensorBoard:

python -m pip install tensorboard
python -m tensorboard.main --logdir traces --port 6006

然后在浏览器中打开 http://localhost:6006。Profiler 页面能够展示算子统计和时间线;具体视图取决于本地安装的 PyTorch 与 TensorBoard 插件版本。

从统计表里读取真正的瓶颈

分析输出时,不要只盯着第一行。可以按下面的顺序检查。

1. 区分总时间与自身时间

CPU totalCUDA total 包含子算子的时间,Self CPUSelf CUDA 更接近算子自身消耗。高层 record_function 区域的总时间很高是正常的,因为它包住了整个注意力调用。

定位底层瓶颈时,应优先查看自身设备时间较高的矩阵乘法、Softmax、复制和类型转换操作。如果 CPU 自身时间高、CUDA 时间却不高,问题可能在 Python 调度、张量准备或同步,而不一定在 GPU 算术吞吐。

2. 检查调用次数

一次大矩阵乘法和数百次小内核可能拥有相似的总耗时,但优化方向完全不同。调用次数异常高通常意味着计算被切得太碎,或者循环放在了 Python 层。注意力实现还可能因为掩码、布局转换和逐元素运算增加内核数量。

3. 结合输入形状判断

启用 record_shapes=True 后,可以把慢算子与具体张量形状关联起来。这里尤其要确认:

  • QKV 是否保持 [B, H, N, D] 的预期布局。
  • 转置之后是否触发了额外的连续化或复制。
  • 短序列和小 batch 是否让内核启动成本占据更大比例。
  • 不同序列长度是否走了不同的执行路径。

record_shapes 会增加分析开销,因此适合诊断,不适合长期保持在生产性能测试中。

4. 检查内存而不只看耗时

手工实现通常会显式创建分数矩阵和 Softmax 输出。Profiler 的内存列可以帮助找到分配量较大的操作,但它不是完整的显存审计工具。需要准确比较峰值时,可以这样补充测量:

if torch.cuda.is_available():
    torch.cuda.reset_peak_memory_stats()
    output = sdpa_attention(q, k, v)
    output.sum().backward()
    torch.cuda.synchronize()
    peak_mb = torch.cuda.max_memory_allocated() / 1024**2
    print(f'peak allocated: {peak_mb:.2f} MiB')

比较不同实现时,应重新创建输入、清空上一轮引用,并在相同的前向与反向条件下测量。否则缓存分配器和仍被引用的张量会扭曲结果。

让实验结果能够指导改造

Profiler 的目标不是生成一张漂亮的火焰图,而是缩小决策范围。可以把观察结果映射到具体动作:

  • 如果分数矩阵的内存分配占主导,优先评估 scaled_dot_product_attention 等可能避免显式物化全部中间结果的路径,并在目标硬件上验证。
  • 如果大量时间花在布局转换或复制上,检查上游投影层输出的维度顺序和连续性。
  • 如果运行被许多小算子切碎,减少 Python 循环,并考虑框架提供的融合实现。
  • 如果不同输入长度差异很大,建立长度分桶基准,不要用单一形状代表全部流量。
  • 如果只在首次迭代变慢,把预热与稳定阶段分开;不要把初始化成本当成稳态延迟。

还要注意,Profiler 本身会引入开销,record_shapesprofile_memory 和调用栈记录都会进一步放大影响。它适合解释执行过程,但最终延迟和吞吐仍应使用轻量基准工具,在关闭 Profiler 后重新测量。

上线前的分析清单

一次可信的注意力性能分析至少应满足这些条件:

  • 固定 PyTorch 版本、设备型号、驱动环境、数据类型和输入形状。
  • 明确测量的是前向、反向,还是完整训练步骤。
  • 预热后再采样,并确保 GPU 测量点进行了必要同步。
  • 同时查看算子自身时间、调用次数、输入形状和内存分配。
  • 用多个代表性序列长度测试,而不是只测最理想的形状。
  • 在关闭 Profiler 后复测吞吐、延迟和峰值显存。
  • 对优化前后结果做数值正确性检查,尤其关注低精度和掩码语义。

注意力优化很少能靠单个耗时数字完成。把算子表、时间线、形状和显存放在一起,才能判断应该换实现、改布局、减少中间张量,还是仅仅修正测量方法。


相关推荐