ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

从逐头循环到 FlashAttention/FlexAttention:LLMs-from-scratch 九种多头注意力实现与性能对比

从逐头循环到 FlashAttention/FlexAttention:LLMs-from-scratch 九种多头注意力实现与性能对比 示例工程大模型人工智能【免费下载链接】LLMs-from-scratchImplement a ChatGPT-like LLM in PyTorch from scratch, step by step项目地址https://gitcode.com/GitHub_Trending/ll/LLMs-from-scratch点击查看免费下载本文围绕仓库 ch03/02_bonus_efficient-multihead-attention 的配套实验展开逐项拆解其中 9 种因果多头注意力Multi-Head Attention, MHA的 PyTorch 实现涵盖第 3 章教学版、合并 QKV、Einsum、PyTorchscaled_dot_product_attention含/不含 FlashAttention、nn.MultiheadAttention含/不含need_weights以及 PyTorch 2.5 引入的 FlexAttention。读完本文你将掌握每种实现的代码结构、参数语义、底层原理并理解仓库在 M3 MacBook Air CPU 与 NVIDIA A100 GPU 上记录的基准测试方法论与结果能够在自己的 GPT/Llama 类解码器项目中做出合理的实现选型。1. 为什么需要“更高效”的多头注意力实现因果自注意力causal self-attention是 GPT、Llama 等 decoder-only LLM 的核心组件它把每个 token 的输入向量分别投影为 Query、Key、Value计算缩放点积注意力分数并通过上三角因果掩码禁止 token 关注未来位置最后按注意力权重聚合 Value。第 3 章主代码 ch03/01_main-chapter-code/ch03.ipynb 给出了最直观、最易读的教学版实现但它远不是最快的实现显式构造并存储(b, num_heads, num_tokens, num_tokens)的完整注意力分数矩阵在长序列下内存开销为 O(n²)且大量矩阵搬运可以进一步优化。仓库中的补充目录 ch03/02_bonus_efficient-multihead-attention 正是为了解决“教学正确但不够高效”的问题它通过 mha-implementations.ipynb 一次性对比 9 种实现并用统一的输入规模batch8、序列长度1024、嵌入维度768在 CPU 与 A100 GPU 上做基准测试帮助读者在“可读性”“参数合并”“内存优化”“融合算子”之间做量化取舍。README 末尾的三张性能汇总图前向、前向反向、torch.compile后的前向反向均“越低越好”即来自该 notebook 的绘图代码。2. 环境准备与统一基准配置notebook 的第一步是固定随机种子并自动选择设备import torch torch.manual_seed(123) if torch.backends.mps.is_available(): device torch.device(mps) # Apple Silicon GPU (Metal) elif torch.cuda.is_available(): device torch.device(cuda) # NVIDIA GPU else: device torch.device(cpu) # CPU fallback print(fUsing device: {device}) print(fPyTorch version: {torch.__version__})随后构造统一的输入张量所有实现都在同一份数据上比较batch_size 8 context_len 1024 embed_dim 768 embeddings torch.randn((batch_size, context_len, embed_dim), devicedevice)即输入形状为(8, 1024, 768)。实验中每个模块实例化时统一使用num_heads12因此head_dim d_out / num_heads 768 / 12 64与 GPT-2 small 规模的配置一致。2.1 版本与安装要求运行全部代码尤其是 FlexAttention需要PyTorch 2.5 及以上FlexAttention 在更早版本中不存在notebook 中用packaging解析版本并设置MIN_TORCH_VERSION 2.5.0做运行前检查。PyTorch 2.5 要求 Python 3.9 或更高。CPU 机器安装pip install torch torchvision torchaudioGPU 机器安装notebook 给出的 CUDA 版本安装示例pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu1242.2 统一构造参数语义所有 9 个类的构造参数高度一致含义如下表参数含义本实验取值d_in输入嵌入维度768d_out输出维度必须能被num_heads整除768wrapper 版为每头 64再由out_proj拼回num_heads注意力头数12head_dim每头维度d_out // num_heads64context_length因果掩码缓冲区的最大序列长度1024dropoutsoftmax 后对注意力权重的 dropout 概率基准时为 0.0qkv_biasQ/K/V 投影是否带偏置False注意GPT-2 实际使用qkv_biasTrue按需开启每种实现输出形状均被验证为torch.Size([8, 1024, 768])保证“功能等价”后再比较性能。3. 九种实现逐项拆解3.1 第 3 章 CausalAttention 包装器逐头实现这是最“忠实于讲解”的版本每个注意力头是一个独立的CausalAttention模块各持有自己的 Q/K/V 投影外层用nn.ModuleList逐头循环再拼接import torch.nn as nn class CausalAttention(nn.Module): def __init__(self, d_in, d_out, context_length, dropout, qkv_biasFalse): super().__init__() self.d_out d_out self.W_query nn.Linear(d_in, d_out, biasqkv_bias) self.W_key nn.Linear(d_in, d_out, biasqkv_bias) self.W_value nn.Linear(d_in, d_out, biasqkv_bias) self.dropout nn.Dropout(dropout) self.register_buffer(mask, torch.triu(torch.ones(context_length, context_length), diagonal1)) def forward(self, x): b, num_tokens, d_in x.shape keys self.W_key(x) queries self.W_query(x) values self.W_value(x) attn_scores queries keys.transpose(1, 2) attn_scores.masked_fill_( self.mask.bool()[:num_tokens, :num_tokens], -torch.inf) attn_weights torch.softmax(attn_scores / keys.shape[-1]**0.5, dim-1) attn_weights self.dropout(attn_weights) context_vec attn_weights values return context_vec class Ch03_MHA_Wrapper(nn.Module): def __init__(self, d_in, d_out, context_length, dropout, num_heads, qkv_biasFalse): super().__init__() self.heads nn.ModuleList( [CausalAttention(d_in, d_out, context_length, dropout, qkv_bias) for _ in range(num_heads)] ) self.out_proj nn.Linear(d_out*num_heads, d_out*num_heads) def forward(self, x): context_vec torch.cat([head(x) for head in self.heads], dim-1) return self.out_proj(context_vec)关键点因果掩码通过register_buffer注册torch.triu(..., diagonal1)上三角矩阵forward中按当前num_tokens截取并masked_fill_为-inf。包装器实例化时d_out embed_dim // 12 6412 个头拼成 768 维后经out_proj输出。其性能瓶颈在于逐头 Python 循环和 12 份独立投影参数。3.2 第 3 章 MultiHeadAttention 类张量并行多头主章节的正式版本不再循环而是用一次Linear(d_in, d_out)投影出全部头的 Q/K/V再通过viewtranspose隐式拆分为多头class Ch03_MHA(nn.Module): def __init__(self, d_in, d_out, context_length, dropout, num_heads, qkv_biasFalse): super().__init__() assert d_out % num_heads 0, d_out must be divisible by num_heads self.d_out d_out self.num_heads num_heads self.head_dim d_out // num_heads self.W_query nn.Linear(d_in, d_out, biasqkv_bias) self.W_key nn.Linear(d_in, d_out, biasqkv_bias) self.W_value nn.Linear(d_in, d_out, biasqkv_bias) self.out_proj nn.Linear(d_out, d_out) self.dropout nn.Dropout(dropout) self.register_buffer(mask, torch.triu(torch.ones(context_length, context_length), diagonal1)) def forward(self, x): b, num_tokens, d_in x.shape keys self.W_key(x) # (b, num_tokens, d_out) queries self.W_query(x) values self.W_value(x) keys keys.view(b, num_tokens, self.num_heads, self.head_dim) values values.view(b, num_tokens, self.num_heads, self.head_dim) queries queries.view(b, num_tokens, self.num_heads, self.head_dim) keys keys.transpose(1, 2) queries queries.transpose(1, 2) values values.transpose(1, 2) attn_scores queries keys.transpose(2, 3) mask_bool self.mask.bool()[:num_tokens, :num_tokens] attn_scores.masked_fill_(mask_bool, -torch.inf) attn_weights torch.softmax(attn_scores / keys.shape[-1]**0.5, dim-1) attn_weights self.dropout(attn_weights) context_vec (attn_weights values).transpose(1, 2) context_vec context_vec.contiguous().view(b, num_tokens, self.d_out) context_vec self.out_proj(context_vec) return context_vec该版本用矩阵运算并行处理全部 12 个头是后续所有优化版本的“正确性基准”例如测试用例test_mha_einsum_matches_ch03就是拿它和 Einsum 版逐参数比对。3.3 合并 QKV 权重MultiHeadAttentionCombinedQKV第 3 章的三个独立nn.Linear(d_in, d_out)被合并为一个nn.Linear(d_in, 3 * d_out)一次矩阵乘法同时算完 Q/K/V减少内核启动与权重搬运class MultiHeadAttentionCombinedQKV(nn.Module): def __init__(self, d_in, d_out, num_heads, context_length, dropout0.0, qkv_biasFalse): super().__init__() assert d_out % num_heads 0, d_out is indivisible by num_heads self.num_heads num_heads self.context_length context_length self.head_dim d_out // num_heads self.qkv nn.Linear(d_in, 3 * d_out, biasqkv_bias) self.proj nn.Linear(d_out, d_out) self.dropout nn.Dropout(dropout) self.register_buffer( mask, torch.triu(torch.ones(context_length, context_length), diagonal1) ) def forward(self, x): batch_size, num_tokens, embed_dim x.shape qkv self.qkv(x) # (b, n, 3 * embed_dim) qkv qkv.view(batch_size, num_tokens, 3, self.num_heads, self.head_dim) qkv qkv.permute(2, 0, 3, 1, 4) # (3, b, num_heads, n, head_dim) queries, keys, values qkv.unbind(0) attn_scores queries keys.transpose(-2, -1) attn_scores attn_scores.masked_fill( self.mask.bool()[:num_tokens, :num_tokens], -torch.inf ) attn_weights torch.softmax(attn_scores / keys.shape[-1]**0.5, dim-1) attn_weights self.dropout(attn_weights) context_vec attn_weights values context_vec context_vec.transpose(1, 2) context_vec context_vec.contiguous().view(batch_size, num_tokens, embed_dim) context_vec self.proj(context_vec) return context_vec此实现基于仓库讨论区第 51 号讨论由 Rayed Bin Wahed 分享的代码改写而来其核心差异正如 notebook 所注用self.qkv nn.Linear(d_in, 3 * d_out, biasqkv_bias)替换self.W_query / W_key / W_value三个层再以q, k, v qkv.unbind(0)拆分。这种“合并 QKV”结构也是 GPT/Llama 等主流开源模型实际采用的布局。3.4 基于 Einsum 的 MHAMHAEinsum把矩阵乘法改写为 Einstein 求和标记投影权重用nn.Parameter直接声明并手动初始化import math class MHAEinsum(nn.Module): def __init__(self, d_in, d_out, context_length, dropout, num_heads, qkv_biasFalse): super().__init__() assert d_out % num_heads 0, d_out must be divisible by num_heads self.d_out d_out self.num_heads num_heads self.head_dim d_out // num_heads self.W_query nn.Parameter(torch.randn(d_in, d_out)) self.W_key nn.Parameter(torch.randn(d_in, d_out)) self.W_value nn.Parameter(torch.randn(d_in, d_out)) if qkv_bias: self.bias_q nn.Parameter(torch.zeros(d_out)) self.bias_k nn.Parameter(torch.zeros(d_out)) self.bias_v nn.Parameter(torch.zeros(d_out)) else: self.register_parameter(bias_q, None) self.register_parameter(bias_k, None) self.register_parameter(bias_v, None) self.out_proj nn.Linear(d_out, d_out) self.dropout nn.Dropout(dropout) self.register_buffer(mask, torch.triu(torch.ones(context_length, context_length), diagonal1)) self.reset_parameters() def reset_parameters(self): nn.init.kaiming_uniform_(self.W_query, amath.sqrt(5)) nn.init.kaiming_uniform_(self.W_key, amath.sqrt(5)) nn.init.kaiming_uniform_(self.W_value, amath.sqrt(5)) if self.bias_q is not None: fan_in, _ nn.init._calculate_fan_in_and_fan_out(self.W_query) bound 1 / math.sqrt(fan_in) nn.init.uniform_(self.bias_q, -bound, bound) nn.init.uniform_(self.bias_k, -bound, bound) nn.init.uniform_(self.bias_v, -bound, bound) def forward(self, x): b, n, _ x.shape Q torch.einsum(bnd,do-bno, x, self.W_query) K torch.einsum(bnd,do-bno, x, self.W_key) V torch.einsum(bnd,do-bno, x, self.W_value) if self.bias_q is not None: Q self.bias_q K self.bias_k V self.bias_v Q Q.view(b, n, self.num_heads, self.head_dim).transpose(1, 2) K K.view(b, n, self.num_heads, self.head_dim).transpose(1, 2) V V.view(b, n, self.num_heads, self.head_dim).transpose(1, 2) scores torch.einsum(bhnd,bhmd-bhnm, Q, K) / (self.head_dim ** 0.5) mask self.mask[:n, :n] scores scores.masked_fill(mask.bool(), -torch.inf) attn_weights torch.softmax(scores, dim-1) attn_weights self.dropout(attn_weights) context_vec torch.einsum(bhnm,bhmd-bhnd, attn_weights, V) context_vec context_vec.transpose(1, 2).reshape(b, n, self.d_out) context_vec self.out_proj(context_vec) return context_vec三个 einsum 式子分别对应投影bnd,do-bno、缩放点积bhnd,bhmd-bhnm与加权聚合bhnm,bhmd-bhnd。该版本主要用于展示“同一数学运算的不同表达方式”仓库测试证明它与第 3 章版数值一致。3.5 PyTorch scaled_dot_product_attention FlashAttentionPyTorch 官方的nn.functional.scaled_dot_product_attentionSDPA在硬件与输入满足条件时自动派发到内存优化的FlashAttention内核见原论文 “FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness”无需显式物化完整注意力矩阵class MHAPyTorchScaledDotProduct(nn.Module): def __init__(self, d_in, d_out, num_heads, context_length, dropout0.0, qkv_biasFalse): super().__init__() assert d_out % num_heads 0, d_out is indivisible by num_heads self.num_heads num_heads self.context_length context_length self.head_dim d_out // num_heads self.d_out d_out self.qkv nn.Linear(d_in, 3 * d_out, biasqkv_bias) self.proj nn.Linear(d_out, d_out) self.dropout dropout def forward(self, x): batch_size, num_tokens, embed_dim x.shape qkv self.qkv(x) # (b, n, 3 * embed_dim) qkv qkv.view(batch_size, num_tokens, 3, self.num_heads, self.head_dim) qkv qkv.permute(2, 0, 3, 1, 4) # (3, b, num_heads, n, head_dim) queries, keys, values qkv use_dropout 0. if not self.training else self.dropout context_vec nn.functional.scaled_dot_product_attention( queries, keys, values, attn_maskNone, dropout_puse_dropout, is_causalTrue) context_vec context_vec.transpose(1, 2).contiguous().view(batch_size, num_tokens, self.d_out) context_vec self.proj(context_vec) return context_vec要点is_causalTrue让内核直接采用因果掩码语义无需用户再传掩码矩阵。dropout_p只在训练模式生效use_dropout 0. if not self.training else self.dropout。这是把“手写注意力”替换为生产级融合算子的最小改动路径也是 GPT 实现中常见的做法。3.6 关闭 FlashAttention 的 SDPA显式因果掩码与 3.5 完全相同的 QKV 结构区别在于显式传入attn_mask并把is_causal设为False从而禁用 FlashAttention 派发路径用于隔离对比“内核融合”与“普通 SDPA”的差距class MHAPyTorchSDPAWithoutFlash(nn.Module): def __init__(self, d_in, d_out, num_heads, context_length, dropout0.0, qkv_biasFalse): super().__init__() assert d_out % num_heads 0, d_out is indivisible by num_heads self.num_heads num_heads self.context_length context_length self.head_dim d_out // num_heads self.d_out d_out self.qkv nn.Linear(d_in, 3 * d_out, biasqkv_bias) self.proj nn.Linear(d_out, d_out) self.dropout dropout self.register_buffer(mask, torch.triu(torch.ones(context_length, context_length), diagonal1).bool()) def forward(self, x): batch_size, num_tokens, embed_dim x.shape qkv self.qkv(x) qkv qkv.view(batch_size, num_tokens, 3, self.num_heads, self.head_dim) qkv qkv.permute(2, 0, 3, 1, 4) queries, keys, values qkv use_dropout 0. if not self.training else self.dropout if self.context_length num_tokens: attn_mask self.mask[:num_tokens, :num_tokens] else: attn_mask self.mask[:self.context_length, :self.context_length] # SDPA uses True for positions that may participate in attention context_vec nn.functional.scaled_dot_product_attention( queries, keys, values, attn_mask~attn_mask, dropout_puse_dropout, is_causalFalse) context_vec context_vec.transpose(1, 2).contiguous().view(batch_size, num_tokens, self.d_out) context_vec self.proj(context_vec) return context_vec掩码语义注意点缓冲区mask是triu(diagonal1)未来位置为TrueSDPA 的attn_mask约定True表示“允许参与注意力”因此这里传入~attn_mask得到“非未来位置允许”的因果掩码。仓库测试test_sdpa_without_flash_does_not_attend_to_future_tokens专门验证了这一点。3.7 torch.nn.MultiheadAttention默认 need_weightsTrue直接包装 PyTorch 高层 APInn.MultiheadAttentionbatch_firstTrue使其输入输出布局与仓库其他实现一致class MHAPyTorchClass(nn.Module): def __init__(self, d_in, d_out, num_heads, context_length, dropout0.0, qkv_biasFalse, need_weightsTrue): super().__init__() self.context_length context_length self.multihead_attn nn.MultiheadAttention( embed_dimd_out, num_headsnum_heads, dropoutdropout, biasqkv_bias, add_bias_kvqkv_bias, batch_firstTrue, ) self.need_weights need_weights self.proj nn.Linear(d_out, d_out) self.register_buffer(mask, torch.triu(torch.ones(context_length, context_length), diagonal1).bool()) def forward(self, x): batch_size, num_tokens, _ x.shape if self.context_length num_tokens: attn_mask self.mask[:num_tokens, :num_tokens] else: attn_mask self.mask[:self.context_length, :self.context_length] attn_output, _ self.multihead_attn( x, x, x, attn_maskattn_mask, need_weightsself.need_weights ) output self.proj(attn_output) return output注意add_bias_kvqkv_bias当qkv_biasTrue时nn.MultiheadAttention还会额外添加可学习的bias_k、bias_v参数这与前面手写实现仅“是否带偏置”的含义略有不同。默认need_weightsTrue会计算并返回注意力权重矩阵属于开销更高的路径。3.8 torch.nn.MultiheadAttention need_weightsFalse同一类只需把need_weights设为False即可按 PyTorch 官方文档切换到优化后的scaled_dot_product_attention。notebook 引用了官方文档说明need_weights: If specified, returnsattn_output_weightsin addition toattn_outputs. Setneed_weightsFalseto use the optimizedscaled_dot_product_attentionand achieve the best performance for MHA. Default:Truemha_pytorch_class_noweights MHAPyTorchClass( d_inembed_dim, d_outembed_dim, context_lengthcontext_len, dropout0.0, num_heads12, qkv_biasFalse, need_weightsFalse # NEW! ).to(device)这组对照3.7 vs 3.8非常直观地展示了“多计算一步注意力权重”带来的性能代价。3.9 PyTorch FlexAttentionFlexAttentionPyTorch 2.5 引入把 FlashAttention 的 IO 优化内核与用户自定义注意力掩码结合起来掩码不再是一张稠密矩阵而是一个可编程的BlockMask与掩码函数。notebook 用q_idx kv_idx直接表达因果性from packaging.version import parse as parse_version def normalize_version(version): parsed_version parse_version(version) return parse_version(f{parsed_version.major}.{parsed_version.minor}.{parsed_version.micro}) current_version normalize_version(torch.__version__) MIN_TORCH_VERSION 2.5.0 required_version parse_version(MIN_TORCH_VERSION) if current_version required_version and torch.cuda.is_available(): from torch.nn.attention.flex_attention import flex_attention, create_block_mask def causal(b, h, q_idx, kv_idx): return q_idx kv_idx class MHAPyTorchFlexAttention(nn.Module): def __init__(self, d_in, d_out, num_heads, context_length, dropout0.0, qkv_biasFalse): super().__init__() assert d_out % num_heads 0, d_out is indivisible by num_heads self.num_heads num_heads self.context_length context_length self.head_dim d_out // num_heads self.d_out d_out self.qkv nn.Linear(d_in, 3 * d_out, biasqkv_bias) self.proj nn.Linear(d_out, d_out) self.dropout dropout def forward(self, x): batch_size, num_tokens, embed_dim x.shape qkv self.qkv(x) qkv qkv.view(batch_size, num_tokens, 3, self.num_heads, self.head_dim) qkv qkv.permute(2, 0, 3, 1, 4) queries, keys, values qkv attn_mask create_block_mask(causal, BNone, HNone, Q_LENnum_tokens, KV_LENnum_tokens, devicex.device) context_vec flex_attention(queries, keys, values, block_maskattn_mask) context_vec context_vec.transpose(1, 2).contiguous().view(batch_size, num_tokens, self.d_out) context_vec self.proj(context_vec) return context_vec两个需要特别注意的限制notebook 均明确标注FlexAttention 目前不支持 dropout代码中use_dropout一行被注释掉因此训练时如需 dropout 不能直接使用该路径。仅支持 PyTorch 2.5 且需要 CUDAnotebook 用if current_version required_version and torch.cuda.is_available():保护定义与实例化。代码注释还提示由于 PyTorch 2.10 及更新版本不再支持对BlockMask做切片forward每次按当前num_tokens重新create_block_mask而非在__init__中预建后复用。4. 性能基准方法与结果4.1 快速对比M3 MacBook Air CPU 与 NVIDIA A100 GPUnotebook 第 10 节先用 IPython 的%timeit做快速对比。M3 MacBook Air CPUPyTorch 2.4.0输入(8, 1024, 768)仅前向实现每循环耗时1) CausalAttention MHA wrapper逐头179 ms2) Ch03_MHA张量并行多头166 ms3) 合并 QKV 权重190 ms4) Einsum196 ms5) SDPAFlashAttention 路径110 ms6) SDPA显式掩码无 Flash99.5 ms7) nn.MultiheadAttention默认198 ms8) nn.MultiheadAttentionneed_weightsFalse168 msNVIDIA A100 GPUPyTorch 2.6.0cu124前向该组测量前先执行torch.set_float32_matmul_precision(high)以启用 Tensor Core实现每循环耗时1) CausalAttention MHA wrapper逐头4.68 ms2) Ch03_MHA张量并行多头3.08 ms3) 合并 QKV 权重3.81 ms4) Einsum4.11 ms5) SDPAFlashAttention 路径1.1 ms6) SDPA显式掩码无 Flash1.8 ms7) nn.MultiheadAttention默认3.04 ms8) nn.MultiheadAttentionneed_weightsFalse2.13 ms9) FlexAttention13.9 ms需要说明快速对比是单次短时测量FlexAttention 单元因耗时较长仅以少量循环测量且 FlexAttention 的数值很可能包含首次调用时掩码构造与内核编译的一次性开销结论应以 README 汇总图带预热与多轮重复的正式基准为准。4.2 正式基准预热 CUDA Event前向、前向反向针对 CUDA 异步执行的特性notebook 使用torch.cuda.Event计时并先做 5 次预热、torch.cuda.synchronize()同步后再循环 1000 次取均值与标准差import numpy as np def time_pytorch_function(func, *input, num_repeats1_000): start torch.cuda.Event(enable_timingTrue) end torch.cuda.Event(enable_timingTrue) # Warmup for _ in range(5): func(*input) torch.cuda.synchronize() times [] for _ in range(num_repeats): start.record() func(*input) end.record() torch.cuda.synchronize() times.append(start.elapsed_time(end)) return np.mean(times), np.std(times)前向反向的测量额外包一层forward_backward清空梯度、前向、loss output.sum()、loss.backward()从而把反向传播纳入耗时def forward_backward(func, embeddings): if embeddings.grad is not None: embeddings.grad.zero_() output func(embeddings) loss output.sum() loss.backward()绘图函数plot_execution_times输出柱状图带标准差误差棒、柱顶标注毫秒数值深色主题分别保存为1_forward-only.pdf、2_forward-and-backward.pdf、3_forward-and-backward-compiled.pdf——这正是 README 中三张汇总图对应的三个场景。4.3 torch.compile 后的前向反向第三组基准在 4.2 的基础上用torch.compile编译每个模块后再做前向反向计时并设置torch._dynamo.config.suppress_errors True避免个别图模式编译失败中断流程import torch._dynamo torch._dynamo.config.suppress_errors True def prepare_function(fn): fn torch.compile(fn) return fnfunctions字典把 8 个实现以及满足版本与 CUDA 条件时加入的 FlexAttention统一注册三组基准共用同一套名字与实例保证横向可比。5. 仓库测试对实现的正确性验证“跑得快”的前提是“算得对”。仓库在 ch03/02_bonus_efficient-multihead-attention/tests/test_mha_implementations.py 中用 pytest 锁定了两个关键性质1Einsum 版与第 3 章版数值一致test_mha_einsum_matches_ch03测试通过llms_from_scratch.utils.import_definitions_from_notebook见 pkg/llms_from_scratch/utils.py直接从 notebook 导入类定义然后执行copy_weights由于Ch03_MHA的W_query是nn.Linear权重形状(d_out, d_in)而MHAEinsum的W_query是形状(d_in, d_out)的nn.Parameter复制时需转置to_mha.W_query.copy_(from_mha.W_query.weight.T)。随后在 3 组参数配置下比对输出pytest.mark.parametrize( d_in,d_out,batch,seq_len,num_heads,seed, [ (768, 768, 2, 4, 12, 123), # d_in d_out (768, 1536, 2, 4, 12, 456), # d_in ! d_out (1024, 512, 2, 4, 8, 789), # d_in d_out ], ) def test_mha_einsum_matches_ch03(...): ... assert out_linear.shape out_einsum.shape torch.Size([batch, seq_len, d_out]) assert torch.allclose(out_linear, out_einsum, atol1e-5)覆盖d_in d_out、d_in ! d_out、d_in d_out三种形态从侧面说明优化实现保持了数学等价性。2无 FlashAttention 的 SDPA 仍满足因果性test_sdpa_without_flash_does_not_attend_to_future_tokens构造 12 维小模型context_length4, num_heads3把第 2 个 token 之后的输入整体加 10 制造扰动断言前两个 token 的输出完全不变x_with_changed_future x.clone() x_with_changed_future[:, 2:] 10 out model(x) out_with_changed_future model(x_with_changed_future) assert torch.allclose(out[:, :2], out_with_changed_future[:, :2], atol1e-5)这直接验证了attn_mask~attn_mask的正确语义——未来位置的扰动不会泄漏到当前输出从而保证“关闭 FlashAttention 的 SDPA”与手写因果掩码行为一致。6. 结果解读与工程选型建议基于上述 9 种实现与仓库记录的数据可以梳理出几条清晰的选型线索以下结论均以本仓库实验配置与版本为前提不同硬件与 PyTorch 版本下数值会变化教学优先第 3 章的两个类Ch03_MHA_Wrapper、Ch03_MHA可读性最好适合理解“逐头”与“张量并行多头”两种视角其中张量并行版本在快速基准中A100 上 3.08 ms反而优于逐头 wrapper4.68 ms说明消除 Python 循环收益明显。合并 QKV从 A100 快速基准看合并 QKV3.81 ms比三投影分离的 Ch03_MHA3.08 ms并无优势甚至略慢——这提醒我们“参数合并”的意义更多在于权重布局与工程实现与主流模型 checkpoint 对齐不能想当然认为一定更快。生产首选nn.functional.scaled_dot_product_attentionis_causalTrue是接入 FlashAttention 的最短路径A100 快速基准中 1.1 ms 的耗时显著领先手写实现nn.MultiheadAttention在need_weightsFalse时同样走 SDPA 优化路径2.13 ms而默认need_weightsTrue3.04 ms会因额外计算权重矩阵变慢。FlexAttention 的定位它把“自定义掩码如滑动窗口、前缀、稀疏结构 Flash 级内核”组合起来适合需要灵活注意力模式的场景但快速基准中 13.9 ms 的数字说明其收益依赖预热/复用BlockMask本实验每次 forward 重建掩码且不支持 dropout、要求 PyTorch 2.5 与 CUDA训练阶段需谨慎评估。正确性优先无论选择哪种优化路径都应像仓库测试那样先做“与教学版逐权重对齐 因果性检查”再谈性能。7. 相关文件速览ch03/02_bonus_efficient-multihead-attention/README.md本专题入口含三张性能汇总图前向 / 前向反向 / 编译后前向反向。ch03/02_bonus_efficient-multihead-attention/mha-implementations.ipynb9 种实现的完整代码、快速基准与正式基准预热 CUDA Event 绘图。ch03/02_bonus_efficient-multihead-attention/tests/test_mha_implementations.pyEinsum 等价性测试与 SDPA 因果性测试。ch03/01_main-chapter-code/ch03.ipynb第 3 章主代码教学版注意力的出处。pkg/llms_from_scratch/utils.pyimport_definitions_from_notebook等测试工具支持从 notebook 直接导入类定义进行验证。赞分享示例工程大模型人工智能【免费下载链接】LLMs-from-scratchImplement a ChatGPT-like LLM in PyTorch from scratch, step by step项目地址https://gitcode.com/GitHub_Trending/ll/LLMs-from-scratch点击查看免费下载相关推荐Python PDF处理库PyPDF5分钟快速入门与实战指南Python PDF处理库PyPDF5分钟快速入门与实战指南 PyPDF是一个功能强大的纯Python PDF处理库专门用于PDF文件的拆分、合并、裁剪和页后端在Windows上完美体验Mac触控板mac-precision-touchpad完整配置指南在Windows上完美体验Mac触控板mac precision touchpad完整配置指南 你是否在Windows电脑上使用MacBook或Magic T驱动开发系统底层硬件开发FLUX注意力机制多头自注意力实现FLUX注意力机制多头自注意力实现 引言从Transformer到FLUX的注意力革命 在深度学习领域注意力机制Attention Mechanism人工智能大模型本地部署媒体生成上一篇BlockNote无障碍动画打造对运动偏好友好的现代编辑器下一篇Unity雨雪天气系统终极指南如何用粒子效果打造沉浸式环境交互创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表