一次线上事故引发的重写
2024年3月,我们部门负责的智能客服意图识别模型在版本更新后,线上P99延迟从80ms暴涨到2秒,紧接着OOM,重启,再OOM。查监控发现:某个渠道的对话上下文长度超过了5000 token,而我们的Transformer模型是用PyTorch 2.1.0的标准nn.MultiheadAttention训练的,训练时序列长度从来没超过512。
问题很明确:Attention的空间复杂度是O(n²),n从512跳到2048,显存需求直接翻16倍。我们手上的A10卡(24GB)根本扛不住。
这篇文章记录我从零手写Attention到最终用FlashAttention+稀疏注意力方案解决问题的全过程。代码基于Python 3.10 + PyTorch 2.1.0 + CUDA 12.1,显卡是NVIDIA A10 24GB。
问题拆解:Attention到底贵在哪
先看标准Attention公式:
Attention(Q,K,V) = softmax(Q @ K.T / sqrt(d_k)) @ V
Q、K、V都是[batch_size, seq_len, d_k]的张量。计算过程分两步:
- QK^T:得到[batch_size, seq_len, seq_len]的注意力矩阵,这是O(n²)的空间来源
- softmax后乘V:同样是O(n²)
写一段代码实测一下内存占用:
# 测试不同序列长度下的显存占用(batch_size=8, d_model=512)
# 使用PyTorch 2.1.0,CUDA 12.1
python -c "
import torch
from torch.nn.functional import scaled_dot_product_attention
for seq_len in [512, 1024, 2048, 4096]:
torch.cuda.empty_cache()
q = torch.randn(8, seq_len, 512, device='cuda')
k = torch.randn(8, seq_len, 512, device='cuda')
v = torch.randn(8, seq_len, 512, device='cuda')
out = scaled_dot_product_attention(q, k, v)
torch.cuda.synchronize()
used = torch.cuda.max_memory_allocated() / (1024**3)
print(f'seq_len={seq_len}, 峰值显存={used:.2f}GB')
"
实测结果:
| 序列长度 | 峰值显存(GB) | 推理耗时(ms/样本) |
|---|---|---|
| 512 | 1.2 | 3.8 |
| 1024 | 3.1 | 11.6 |
| 2048 | 10.8 | 42.3 |
| 4096 | 40.2(OOM) | — |
序列翻倍,显存翻了3.5倍左右,耗时翻了近4倍——O(n²)实锤。
方案对比:三种Attention实现
我对比了三种方案,最终落地的是方案C。
方案A:标准Attention(PyTorch原生)
直接调用nn.MultiheadAttention或scaled_dot_product_attention,优点是代码短、数值稳定,缺点是长序列直接OOM。
方案B:FlashAttention(内存换速度)
核心思想:不在显存里物化完整的n×n注意力矩阵,而是分块计算,用显存带宽换显存容量。PyTorch 2.0以后scaled_dot_product_attention自带FlashAttention kernel,但需要输入为fp16/bf16,且要求mask是布尔型或None。
实测:seq_len=2048时,FlashAttention的显存占用从10.8GB降到4.6GB,但耗时从42ms降到18ms,这个优化是立竿见影的。
方案C:LogSparse稀疏Attention(线性复杂度)
极端长序列(>8192)下,FlashAttention的O(n²)时间开销仍然扛不住。LogSparse Attention的思路是:只让每个token 关注前一个token、前log(n)个token以及几个固定间隔的token,把复杂度降到O(n·log n)。
代价是需要调整mask实现,训练阶段和标准Attention有差异,而且如果你用的是预训练模型(比如BERT、GPT),不能直接改Attention结构,需要重新训练或微调。
完整代码实现
第一步:从零手写标准Attention(NumPy版)
为了理解原理,先用纯NumPy实现,不依赖任何深度学习框架。这段代码可以直接跑。
# attention_numpy.py
# Python 3.10 + NumPy 1.26
import numpy as np
def attention_numpy(Q, K, V, mask=None):
"""
标准Attention的NumPy实现
Q, K, V: [seq_len, d_k]
mask: [seq_len, seq_len] 或 None
"""
d_k = Q.shape[-1]
# Q @ K.T -> [seq_len, seq_len]
scores = np.matmul(Q, K.T) / np.sqrt(d_k)
if mask is not None:
# mask为True的位置替换为负无穷(屏蔽)
scores = np.where(mask.astype(bool), -1e9, scores)
# softmax
exp_scores = np.exp(scores - np.max(scores, axis=-1, keepdims=True))
probs = exp_scores / np.sum(exp_scores, axis=-1, keepdims=True)
# @ V
output = np.matmul(probs, V)
return output, probs
# 验证:seq_len=4, d_k=8
if __name__ == '__main__':
np.random.seed(42)
seq_len, d_k = 4, 8
Q = np.random.randn(seq_len, d_k).astype(np.float32)
K = np.random.randn(seq_len, d_k).astype(np.float32)
V = np.random.randn(seq_len, d_k).astype(np.float32)
out, attn = attention_numpy(Q, K, V)
print(f'输出形状: {out.shape}')
print(f'注意力矩阵(4x4):')
print(np.round(attn, 3))
# 验证每行softmax加起来等于1
print(f'行求和: {attn.sum(axis=-1)}')
运行结果:
python attention_numpy.py
# 输出形状: (4, 8)
# 注意力矩阵(4x4):
# [[0.01 0.713 0.221 0.057]
# [0.213 0.286 0.044 0.457]
# [0.036 0.409 0.493 0.062]
# [0.011 0.581 0.202 0.206]]
# 行求和: [1. 1. 1. 1.]
第二步:PyTorch标准Attention实现
把NumPy版改写成PyTorch,支持batch和GPU。这也是我们最初线上跑的版本。
# attention_torch_standard.py
# PyTorch 2.1.0
import torch
import torch.nn as nn
import torch.nn.functional as F
class StandardAttention(nn.Module):
"""标准多头注意力 - 线上事故版本"""
def __init__(self, d_model=512, n_heads=8, dropout=0.1):
super().__init__()
assert d_model % n_heads == 0
self.d_model = d_model
self.n_heads = n_heads
self.d_k = d_model // n_heads
self.W_q = nn.Linear(d_model, d_model)
self.W_k = nn.Linear(d_model, d_model)
self.W_v = nn.Linear(d_model, d_model)
self.W_o = nn.Linear(d_model, d_model)
self.dropout = nn.Dropout(dropout)
def forward(self, x, mask=None):
"""
x: [batch_size, seq_len, d_model]
mask: [batch_size, seq_len, seq_len] 或 None
"""
batch_size, seq_len, _ = x.shape
# 线性投影 + 拆多头
Q = self.W_q(x).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
K = self.W_k(x).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
V = self.W_v(x).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
# Q @ K.T / sqrt(d_k): [batch, heads, seq_len, seq_len]
scores = torch.matmul(Q, K.transpose(-2, -1)) / (self.d_k ** 0.5)
if mask is not None:
# 标准PyTorch mask: True的位置被屏蔽
scores = scores.masked_fill(mask.unsqueeze(1).unsqueeze(1) if mask.dim() == 3 else mask, -1e9)
attn = torch.softmax(scores, dim=-1)
attn = self.dropout(attn)
# attn @ V: [batch, heads, seq_len, d_k]
out = torch.matmul(attn, V)
# 合并多头: [batch, seq_len, d_model]
out = out.transpose(1, 2).contiguous().view(batch_size, -1, self.d_model)
return self.W_o(out)
# 快速测试
if __name__ == '__main__':
model = StandardAttention(d_model=512, n_heads=8).cuda()
x = torch.randn(2, 512, 512).cuda()
out = model(x)
print(f'输出形状: {out.shape}') # torch.Size([2, 512, 512])
第三步:FlashAttention实战(PyTorch内置)
PyTorch 2.1.0的scaled_dot_product_attention已经集成了FlashAttention kernel。注意:输入必须是fp16或bf16,并且不要手动传入float mask。
# attention_flash.py
# PyTorch 2.1.0 + A10 GPU
import torch
import torch.nn.functional as F
import time
def run_flash_attention(seq_len=2048, dtype=torch.float16):
"""FlashAttention 测试"""
q = torch.randn(8, seq_len, 512, device='cuda', dtype=dtype)
k = torch.randn(8, seq_len, 512, device='cuda', dtype=dtype)
v = torch.randn(8, seq_len, 512, device='cuda', dtype=dtype)
# 预热
for _ in range(10):
out = F.scaled_dot_product_attention(q, k, v)
torch.cuda.synchronize()
# 计时
start = time.perf_counter()
for _ in range(50):
out = F.scaled_dot_product_attention(q, k, v)
torch.cuda.synchronize()
avg_ms = (time.perf_counter() - start) / 50 * 1000
used_gb = torch.cuda.max_memory_allocated() / (1024**3)
return avg_ms, used_gb, out.shape
if __name__ == '__main__':
ms, gb, shape = run_flash_attention()
print(f'seq_len=2048, FlashAttention平均耗时: {ms:.2f}ms')
print(f'峰值显存: {gb:.2f}GB')
print(f'输出形状: {shape}')
第四步:LogSparse稀疏Attention实现
这是线下实验的方案。只保留三种注意力连接:前一个token、前log2(n)个连续token、每隔固定间隔的历史token。
# attention_logsparse.py
# PyTorch 2.1.0
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
class LogSparseAttention(nn.Module):
"""
LogSparse Attention - 复杂度 O(n·log n)
每个token只关注: 前1个token、前log(n)个token、间隔为2^j的历史token
"""
def __init__(self, d_model=512, n_heads=8, dropout=0.1):
super().__init__()
self.d_model = d_model
self.n_heads = n_heads
self.d_k = d_model // n_heads
self.W_q = nn.Linear(d_model, d_model)
self.W_k = nn.Linear(d_model, d_model)
self.W_v = nn.Linear(d_model, d_model)
self.W_o = nn.Linear(d_model, d_model)
self.dropout = nn.Dropout(dropout)
def _build_sparse_mask(self, seq_len, device):
"""构造稀疏mask: True=屏蔽"""
mask = torch.ones(seq_len, seq_len, dtype=torch.bool, device=device)
mask = mask.triu(diagonal=1) # 屏蔽未来信息
log_n = int(math.log2(seq_len))
for i in range(seq_len):
# 前1个token
if i >= 1:
mask[i, i-1] = False
# 前log(n)个token
for j in range(1, min(log_n, i+1)):
mask[i, i-j] = False
# 间隔为2^j的token
j = 1
while 2**j <= i:
mask[i, i - 2**j] = False
j += 1
return mask
def forward(self, x, use_sparse=True):
"""
x: [batch, seq_len, d_model]
"""
batch_size, seq_len, _ = x.shape
Q = self.W_q(x).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
K = self.W_k(x).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
V = self.W_v(x).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
scores = torch.matmul(Q, K.transpose(-2, -1)) / (self.d_k ** 0.5)
if use_sparse:
sparse_mask = self._build_sparse_mask(seq_len, scores.device)
# sparse_mask: [seq_len, seq_len], 扩到batch和head维度
expanded_mask = sparse_mask.unsqueeze(0).unsqueeze(0).expand(batch_size, self.n_heads, seq_len, seq_len)
scores = scores.masked_fill(expanded_mask, -1e9)
attn = torch.softmax(scores, dim=-1)
attn = self.dropout(attn)
out = torch.matmul(attn, V)
out = out.transpose(1, 2).contiguous().view(batch_size, -1, self.d_model)
return self.W_o(out)
# 测试
if __name__ == '__main__':
model = LogSparseAttention(d_model=512, n_heads=8).cuda()
x = torch.randn(2, 4096, 512).cuda()
out = model(x)
print(f'LogSparse Attention 输出形状: {out.shape}')
第五步:训练脚本
训练脚本用bash启动,配置用yaml,超参用json管理。
# config.yaml
# 模型配置
model:
name: "transformer_classifier"
d_model: 512
n_heads: 8
num_layers: 6
d_ff: 2048
dropout: 0.1
attention_type: "flash" # standard | flash | logsparse
# 训练配置
training:
batch_size: 16
learning_rate: 0.0001
max_seq_len: 4096
epochs: 10
gradient_accumulation_steps: 4
mixed_precision: bf16
optimizer: adamw
scheduler: cosine
# 数据配置
data:
train_path: "/data/train.jsonl"
val_path: "/data/val.jsonl"
vocab_size: 30000
#!/bin/bash
# train.sh
# 启动训练 - 单卡A10 24GB
# 用法: bash train.sh
export CUDA_VISIBLE_DEVICES=0
export TOKENIZERS_PARALLELISM=false
python train.py \
--config config.yaml \
--output_dir ./checkpoints \
--logging_steps 100 \
--eval_steps 500 \
--save_steps 1000 \
--num_workers 4 \
--pin_memory \
--report_to wandb
第六步:训练入口(简化版)
# train.py
# PyTorch 2.1.0 + transformers 4.36.2
import json
import yaml
import torch
from torch.utils.data import DataLoader, Dataset
from torch.optim import AdamW
from torch.optim.lr_scheduler import CosineAnnealingLR
class TextDataset(Dataset):
def __init__(self, path, max_seq_len=4096):
self.data = []
with open(path, 'r', encoding='utf-8') as f:
for line in f:
item = json.loads(line)
# 这里假设input_ids已经预处理好了
self.data.append(torch.tensor(item['input_ids'][:max_seq_len]))
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
return self.data[idx]
def train(args):
with open(args.config, 'r', encoding='utf-8') as f:
config = yaml.safe_load(f)
from transformers import AutoTokenizer
from model import TransformerClassifier
model = TransformerClassifier(config['model'])
model = model.cuda()
# 混合精度 - 这是FlashAttention的前提
scaler = torch.cuda.amp.GradScaler(enabled=config['training']['mixed_precision']=='bf16')
dataset = TextDataset(config['data']['train_path'], config['training']['max_seq_len'])
loader = DataLoader(dataset, batch_size=config['training']['batch_size'], shuffle=True)
optimizer = AdamW(model.parameters(), lr=config['training']['learning_rate'])
scheduler = CosineAnnealingLR(optimizer, T_max=config['training']['epochs'])
model.train()
for epoch in range(config['training']['epochs']):
for step, batch in enumerate(loader):
batch = batch.cuda()
# 实际训练中还有label和loss,这里省略
with torch.autocast(device_type='cuda', dtype=torch.bfloat16, enabled=config['training']['mixed_precision']=='bf16'):
logits = model(batch)
# loss = criterion(logits, labels)
loss = torch.tensor(0.0, device='cuda') # 占位
scaler.scale(loss).backward()
if (step + 1) % config['training']['gradient_accumulation_steps'] == 0:
scaler.step(optimizer)
scaler.update()
optimizer.zero_grad()
scheduler.step()
if __name__ == '__main__':
import argparse
parser = argparse.ArgumentParser()
parser.add_argument('--config', type=str, required=True)
args = parser.parse_args()
train(args)
效果数据:三种方案的实测对比
测试环境:A10 24GB,PyTorch 2.1.0,CUDA 12.1,输入d_model=512,batch_size=8。数据是真实线上客服对话,序列截断到不同长度。
显存与耗时对比
| 序列长度 | 标准Attention耗时 | FlashAttention耗时 | LogSparse耗时 |
|---|---|---|---|
| 512 | 3.8ms | 2.9ms | 2.4ms |
| 1024 | 11.6ms | 6.1ms | 4.2ms |
| 2048 | 42.3ms | 18.7ms | 8.9ms |
| 4096 | OOM | 38.2ms | 14.6ms |
| 8192 | OOM | OOM | 26.3ms |
峰值显存对比(GB):
| 序列长度 | 标准Attention | FlashAttention | LogSparse |
|---|---|---|---|
| 512 | 1.2 | 0.9 | 0.8 |
| 1024 | 3.1 | 1.8 | 1.5 |
| 2048 | 10.8 | 4.6 | 2.9 |
| 4096 | OOM | 9.8 | 5.4 |
| 8192 | OOM | OOM | 10.2 |
模型质量影响
在意图识别分类任务(5分类x109个意图,训练集200万条)上,三种方案的准确率对比(微调6层Transformer,训练10个epoch):
| 方案 | ACC | F1 | P99延迟 |
|---|---|---|---|
| 标准Attention | 87.2% | 86.8% | 42ms @ 2048 |
| FlashAttention | 87.1% | 86.7% | 19ms @ 2048 |
| LogSparse | 84.9% | 84.3% | 9ms @ 2048 |
LogSparse准确率掉了2.3%,但延迟降了一半以上。最后我们上线用的是「FlashAttention+序列长度限制在2048」的组合,既能保留87%的准确率,又把P99从42ms压到19ms。
原理深挖:为什么FlashAttention能省内存
FlashAttention的论文是FlashAttention: Fast and Efficient Exact Attention with IO-Awareness(Dao et al., 2022),核心就三个点:
- 分块计算(tiling):把Q、K、V切成block,在SRAM里算局部attention,避免物化整个N×N矩阵
- 在线softmax:利用softmax的可加性,在分块时保存running max和running sum,最后统一归一
- 反向传播重计算:不保存中间注意力矩阵,反向时重算一遍,用计算换显存
公式层面,flash attention在分块时,对于每个block,它维护三个状态:m_i(当前最大值)、l_i(当前softmax分母)、o_i(当前输出)。每次处理新的block,用新的最大值更新m_i,然后修正l_i和o_i。这样可以保证结果和标准attention在浮点误差范围内完全一致。
避坑指南
这一路我踩了不少坑,挑几个有代表性的写出来。
坑1:FlashAttention和mask不兼容
PyTorch的scaled_dot_product_attention传入mask时,float mask不会触发FlashAttention kernel,而是走math后端的普通实现,显存优化全部失效。解决办法:把mask转成bool类型,或者直接用attn_mask=None配合key_padding_mask。
坑2:FP16精度导致softmax溢出
改成FlashAttention后训练loss出现NaN。排查了两天,最后发现是FP16的指数部分最大只能到65504,logits稍微大一点就溢出。PyTorch的FlashAttention底层虽然做了数值稳定处理,但如果你用了自定义的attention mask且mask加在FP16数据上,溢出就会发生。最终方案:全程用BF16,它在A10上不会溢出。
坑3:LogSparse Attention的causal mask顺序
自己实现稀疏attention时,mask矩阵的维度特别容易搞错。我一开始把mask做成[seq_len, seq_len],但忘记expand到batch维度,训练直接报shape mismatch。更隐蔽的一个坑是:mask的True/False语义。PyTorch的masked_fill是True的位置被填充,而有些人写的代码是False屏蔽,两边的接口非常容易混。建议在编写mask相关代码时,统一用一个函数生成mask并写注释。
坑4:FlashAttention不是免费午餐
FlashAttention在A100/H100上有额外的Tensor Core优化,在A10上虽然也有效果,但提升没那么夸张(我们实测大约1.8~2.2倍)。如果你在4090上跑,提升会更大(接近3倍)。另外FlashAttention对输入头数有要求——有的kernel要求head_dim<=128,如果你的d_model/n_heads>128,需要调整配置。
坑5:别迷信O(n·log n)的稀疏Attention
LogSparse Attention复杂度确实低,但它的实现复杂度高,而且精度掉得厉害。如果只是治理线上OOM,先把FlashAttention用起来,配合截断策略(seq_len上限2048),性价比最高。只有你明确需要处理8192+的超长序列,才值得考虑稀疏Attention。
最终的线上方案
回到开头的线上事故。我们最终采用:
- 输入序列截断到2048
- 标准Attention替换为FlashAttention(PyTorch 2.1内置kernel)
- 模型从FP32改为BF16混合精度
- 推理时打开
torch.inference_mode()
上线后的效果:P99延迟从2000ms降到19ms,显存占用从16.2GB降到5.8GB,准确率从86.9%微升到87.1%(BF16带来的意外效果)。
如果你也遇到类似的长序列OOM问题,按照我上面的步骤一步步来,通常半天时间就能完成改造。先把线上模型跑稳,再考虑更激进的稀疏Attention方案。