注意力层经常同时承受计算、显存和数据搬运压力。只看一次前向传播的总耗时,很难判断瓶颈究竟来自矩阵乘法、Softmax、中间张量,还是内核调度。更有效的办法是用 PyTorch Profiler 把注意力计算拆到算子级,并结合输入形状、显存分配和执行时间做判断。
本文基于“分析 PyTorch 中的注意力计算”这一主题给出一套可直接实践的方法。由于来源摘要没有提供具体实验数据,下面的代码和判断流程属于可复现的工程示例,不代表原文给出的基准结果。
注意力为什么值得单独分析
标准缩放点积注意力可以写成:
scores = Q @ K^T / sqrt(d)
probabilities = softmax(scores)
output = probabilities @ V
假设查询和键的序列长度都是 N,注意力分数张量的形状通常为 [batch, heads, N, N]。这意味着序列长度翻倍时,分数矩阵的元素数量约增至四倍。实际运行中需要关注的不只是 FLOPs,还包括:
Q @ K^T和P @ V两次矩阵乘法的设备时间。- Softmax 以及缩放、掩码等逐元素操作是否产生大量独立内核。
[N, N]中间结果是否被显式物化,从而推高显存峰值。- 张量布局、数据类型和形状是否让框架选择了不同的执行路径。
- 首次运行中的内核加载、缓存建立等一次性成本是否污染了测量。
因此,分析注意力不能只在函数两端放一个计时器。算子表适合回答“时间花在哪里”,时间线适合回答“这些工作以什么顺序发生、设备是否存在空闲”。
一个可运行的 Profiler 实验
下面的脚本对比手工实现的注意力与 torch.nn.functional.scaled_dot_product_attention。运行前按机器显存调整 BATCH、HEADS、SEQ_LEN 和 HEAD_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 total 或 CUDA total 包含子算子的时间,Self CPU 或 Self CUDA 更接近算子自身消耗。高层 record_function 区域的总时间很高是正常的,因为它包住了整个注意力调用。
定位底层瓶颈时,应优先查看自身设备时间较高的矩阵乘法、Softmax、复制和类型转换操作。如果 CPU 自身时间高、CUDA 时间却不高,问题可能在 Python 调度、张量准备或同步,而不一定在 GPU 算术吞吐。
2. 检查调用次数
一次大矩阵乘法和数百次小内核可能拥有相似的总耗时,但优化方向完全不同。调用次数异常高通常意味着计算被切得太碎,或者循环放在了 Python 层。注意力实现还可能因为掩码、布局转换和逐元素运算增加内核数量。
3. 结合输入形状判断
启用 record_shapes=True 后,可以把慢算子与具体张量形状关联起来。这里尤其要确认:
Q、K、V是否保持[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_shapes、profile_memory 和调用栈记录都会进一步放大影响。它适合解释执行过程,但最终延迟和吞吐仍应使用轻量基准工具,在关闭 Profiler 后重新测量。
上线前的分析清单
一次可信的注意力性能分析至少应满足这些条件:
- 固定 PyTorch 版本、设备型号、驱动环境、数据类型和输入形状。
- 明确测量的是前向、反向,还是完整训练步骤。
- 预热后再采样,并确保 GPU 测量点进行了必要同步。
- 同时查看算子自身时间、调用次数、输入形状和内存分配。
- 用多个代表性序列长度测试,而不是只测最理想的形状。
- 在关闭 Profiler 后复测吞吐、延迟和峰值显存。
- 对优化前后结果做数值正确性检查,尤其关注低精度和掩码语义。
注意力优化很少能靠单个耗时数字完成。把算子表、时间线、形状和显存放在一起,才能判断应该换实现、改布局、减少中间张量,还是仅仅修正测量方法。