ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

深入解析 Attention、Mask 与多头机制:从数学原理到 train-llm-from-scratch 的源码实现

深入解析 Attention、Mask 与多头机制:从数学原理到 train-llm-from-scratch 的源码实现 深入解析 Attention、Mask 与多头机制从数学原理到 train-llm-from-scratch 的源码实现【免费下载链接】train-llm-from-scratchA straightforward method for training your LLM, from downloading data to generating text.项目地址: https://gitcode.com/GitHub_Trending/tr/train-llm-from-scratch导读本文围绕 docs/foundations/attention.md 展开系统讲解解码器decoder-only语言模型中自注意力self-attention的数学定义、因果掩码causal mask的作用、为什么需要除以 √D、多头注意力的组合方式及其计算代价。结合本仓库 src/models/attention.py 的逐行实现、src/models/transformer_block.py 与 src/models/transformer.py 的调用上下文以及 configs/base.json 的真实超参数你将掌握从公式到可运行 PyTorch 代码的完整映射并学会用一套自查问题定位注意力相关的训练异常。自注意力每个 token 决定关注哪些历史 token自注意力是 Transformer 中最核心的运算。它的作用是让序列中的每一个 token 自行决定哪些历史 token 对它当前时刻的预测更重要。在解码器-only 的语言模型中位置 (t) 只能使用位置 (0..t) 的信息而不能看到位置 (t1..T-1) 的未来信息。这一限制正是下一个 token 预测next-token prediction训练目标能够保持诚实honest的原因——模型在任何时刻都无法通过偷看未来来作弊它只能基于已看到的前缀来推断下一个 token 的概率分布。该限制与仓库整体的训练目标一致如 docs/foundations/README.md 所述基础阶段 LLM 学习的是条件概率模型[ p_\theta(x_1, x_2, \ldots, x_T) \prod_{t1}^{T} p_\theta(x_t \mid x_{t}) ]注意力机制正是实现这种只能看过去的条件依赖关系的核心算子。注意力公式Q、K、V 三个线性投影对于输入张量 (X \in \mathbb{R}^{B \times T \times C})其中 (B) 为 batch size(T) 为序列长度(C) 为嵌入宽度 (n_embed)一个注意力头学习三个线性投影[ Q X W_Q,\quad K X W_K,\quad V X W_V ]其中 (Q,K,V \in \mathbb{R}^{B \times T \times D})(D) 为单头宽度head size。缩放点积注意力scaled dot-product attention定义为[ \text{Attention}(Q,K,V) \text{softmax}\left(\frac{QK^T}{\sqrt{D}} M\right)V ]掩码 (M) 在允许关注的位置为0在未来位置为 (-\infty)。加入 (M) 之后再经过 softmax被掩码位置的权重就会趋近于 0从而禁止信息从未来流向当前。仓库中的逐行实现上述公式在 src/models/attention.py 的Head类中得到了忠实且可读的实现def __init__(self, head_size: int, n_embed: int, context_length: int) - None: super().__init__() self.key nn.Linear(n_embed, head_size, biasFalse) # Key projection self.query nn.Linear(n_embed, head_size, biasFalse) # Query projection self.value nn.Linear(n_embed, head_size, biasFalse) # Value projection # Lower triangular matrix for causal masking self.register_buffer(tril, torch.tril(torch.ones(context_length, context_length))) def forward(self, x: torch.Tensor) - torch.Tensor: B, T, C x.shape head_size self.key.out_features k self.key(x) # (B, T, head_size) q self.query(x) # (B, T, head_size) scale_factor 1 / math.sqrt(head_size) # (B, T, head_size) (B, head_size, T) - (B, T, T) attn_weights q k.transpose(-2, -1) * scale_factor # Apply causal masking attn_weights attn_weights.masked_fill(self.tril[:T, :T] 0, float(-inf)) attn_weights F.softmax(attn_weights, dim-1) v self.value(x) # (B, T, head_size) # (B, T, T) (B, T, head_size) - (B, T, head_size) out attn_weights v return out值得注意的实现细节Q、K、V 三个投影均为无偏置的nn.Linear简化了参数结构head_size self.key.out_features直接从投影层读取保证缩放系数与投影宽度严格一致scale_factor 1 / math.sqrt(head_size)对应公式中的 (1/\sqrt{D})掩码缓冲区tril通过register_buffer注册会随模型一起移动设备、不参与梯度更新代码中通过self.tril[:T, :T] 0切片允许实际输入序列长度T小于等于context_length。Query、Key、Value 的直观理解对于每个 token三个向量的语义可以这样理解query查询我在寻找什么信息——当前 token 想要从其他位置获取什么key键我包含什么信息——当前 token 能为其他 token 提供什么value值如果被选中我应该传递什么内容——真正被加权求和搬运到下游的内容。点积 (q_t \cdot k_s) 衡量 token (t) 对 token (s) 信息的渴望程度。分数越高token (s) 的 value 对 token (t) 输出的贡献越大。这种查询—键—值的检索式结构是有意设计的注意力头本质上是在学一组可微的软检索规则而不是硬编码的离散规则。为什么除以 √D假设 queries 和 keys 的各分量方差约为 1那么它们的点积方差与维度 (D) 成正比。如果头尺寸变大而不做缩放点积 logits 会变得很大导致 softmax 输出过于尖锐接近 one-hot梯度变得极弱甚至消失。缩放因子将注意力 logits 保持在一个更稳定的范围[ \frac{q_t \cdot k_s}{\sqrt{D}} ]这正是scaled dot-product attention中 scaled 一词的来源。在本仓库实现里缩放以乘法形式出现* scale_factor其中scale_factor 1 / math.sqrt(head_size)等价于除以 (\sqrt{D})。因果掩码Causal Masking以 (T5) 为例允许的注意力模式是一个下三角矩阵[ \begin{bmatrix} 1 0 0 0 0 \ 1 1 0 0 0 \ 1 1 1 0 0 \ 1 1 1 1 0 \ 1 1 1 1 1 \end{bmatrix} ]矩阵中第 (i) 行第 (j) 列表示第 (i) 个位置是否可以关注第 (j) 个位置。仓库将这个模式存储为一个下三角缓冲区self.register_buffer(tril, torch.tril(torch.ones(context_length, context_length)))然后在 softmax 之前掩蔽未来位置attn_weights q k.transpose(-2, -1) * scale_factor attn_weights attn_weights.masked_fill(self.tril[:T, :T] 0, float(-inf)) attn_weights F.softmax(attn_weights, dim-1) out attn_weights v因为未来位置的 logits 被设为 (-\infty)其 softmax 概率自动变为 0。整个流程可用下图概括在训练teacher-forced阶段样本是连续窗口的 token 序列见 data_loader/data_loader.py模型在每一步都基于前缀预测下一个 token因果掩码保证位置 (t) 的预测绝不使用位置 (t1) 及之后的信息从而与下一个 token监督信号严格对齐。多头注意力Multi-Head Attention一个头只能学会一种注意力模式。多头机制让模型并行学习多种不同的注意力模式例如句法依赖syntax dependencies重复的名字或实体repeated names or entities局部短语结构local phrase structure分隔符与格式追踪delimiter and format tracking算术或类代码依赖arithmetic or code-like dependencies。仓库通过创建n_head个独立的Head模块实现多头见 src/models/attention.pyself.heads nn.ModuleList([ Head(n_embed // n_head, n_embed, context_length) for _ in range(n_head) ]) self.proj nn.Linear(n_embed, n_embed)前向传播时拼接各头输出再过最终投影x torch.cat([h(x) for h in self.heads], dim-1) x self.proj(x)如果 (H) 个头每个输出宽度 (DC/H)拼接后宽度恢复为 (C)[ \text{Concat}(\text{head}_1,\ldots,\text{head}_H) \in \mathbb{R}^{B \times T \times C} ]最终投影self.proj负责在头之间混合信息。整体结构如下图所示一个关键约束head_size 必须是整数Head(n_embed // n_head, ...)意味着n_embed必须能被n_head整除。本仓库 configs/base.json 的默认配置中n_embed 1024、n_head 16因此每个头宽度为1024 / 16 64。若修改配置导致无法整除会破坏张量形状一致性详见下文调试清单第 3 条。注意力在 Transformer 块中的位置注意力不是孤立存在的。在 src/models/transformer_block.py 中每个 Transformer 块采用 pre-norm 残差结构self.ln1 nn.LayerNorm(n_embed) self.attn MultiHeadAttention(n_head, n_embed, context_length) self.ln2 nn.LayerNorm(n_embed) self.mlp MLP(n_embed) def forward(self, x: torch.Tensor) - torch.Tensor: # Apply multi-head attention with residual connection x x self.attn(self.ln1(x)) # Apply MLP with residual connection x x self.mlp(self.ln2(x)) return x数学上即[ u x \text{MHA}(\text{LN}(x)),\qquad y u \text{MLP}(\text{LN}(u)) ]每个块承担两个互补的职责注意力跨位置搬运信息MLP 对每个位置独立做非线性变换本仓库的 MLP 使用 ReLU将宽度从 (C) 扩展到 (4C) 再投影回 (C)见 src/models/mlp.py。位置信息由 src/models/transformer.py 中的绝对位置嵌入提供pos_embedding self.position_embed(self.pos_idxs[:T])因为注意力本身是对位置置换等变的permutation-equivariant没有位置编码模型就无法区分 token 出现的先后。注意力代价为什么上下文长度如此昂贵每个头的注意力分数矩阵形状为[ B \times T \times T ]有 (H) 个头时核心分数存储量约为[ O(BHT^2) ]这正是上下文长度代价高昂的根本原因将序列长度 (T) 翻倍注意力矩阵规模约变为原来的 4 倍(T^2) 增长。本仓库作为教学实现刻意保持可读性直接物化materialize了这些矩阵没有使用 Flash Attention 之类的分块技巧因此在实际训练超长序列时需要留意显存预算。仓库在 configs/pretrain.json 中默认使用batch_size 8、grad_accum 12配合 1024 上下文长度通过梯度累积见 scripts/pretrain_base.py在受限显存下模拟更大的有效批次。注意力能做什么不能做什么注意力在位置之间混合信息但它本身不能生成词表上的概率分布——这由最后的语言模型头lm_headnn.Linear(n_embed, vocab_size)见 src/models/transformer.py负责在缺少位置信息的情况下感知 token 顺序——这依赖位置嵌入执行加权平均之外的非线性变换——这依赖 MLP。这些职责分别由位置嵌入、MLP、LayerNorm、残差路径和最终 LM head 承担。整个模型是一个分工明确的分层系统注意力解决关注谁其余组件解决如何变换与如何输出。训练异常时的注意力自查清单如果训练表现异常从注意力的角度依次检查这些问题掩码是否因果掩码是否真的禁止了 token 看到答案未来 token如果掩码失效训练损失会异常偏低但泛化极差q、k、v 是否来自同一个归一化输入在 pre-norm 结构中Q/K/V 投影的输入是ln1(x)而非原始x若不同源统计分布会不一致head_size n_embed // n_head是否为整数不能整除会直接导致形状错误或隐式截断拼接所有头后是否恰好恢复n_embed个通道torch.cat(..., dim-1)后宽度必须等于n_embed否则proj的输入维度不匹配序列长度T是否小于等于context_lengthself.tril[:T, :T]的切片要求T context_length超长输入需要裁剪生成阶段的idx[:, -self.context_length:]正是为此见 src/models/transformer.py。此外训练时若出现损失异常还可以对比训练集与验证集损失在 scripts/pretrain_base.py 的estimate_loss中模型分别在train与dev两个 split 上评估平均损失若训练损失低而验证损失高往往意味着模型依赖了不该有的捷径例如掩码泄露而不是真正的泛化。从基础注意力到后续阶段注意力产生的是隐藏状态而损失函数决定了隐藏状态如何转化为学习信号。理解注意力之后建议按学习路径继续阅读Objectives, Losses Perplexity——损失如何驱动注意力学到的表征Decoder-Only Transformer——注意力的完整上下文嵌入、块结构、参数量Optimization Training Systems——训练循环如何保持稳定Generation Sampling——注意力学到的分布如何生成文本。本仓库的一个重要设计是同一套骨干backbone在所有阶段被复用见 docs/foundations/README.md——SFT 保留 next-token 预测但掩码到助手回复奖励模型把输出从词表 logits 换成标量分数DPO 比较 chosen/rejected 的序列对数概率PPO/GRPO 采样补全并做受限 RL 更新。src/models/transformer.py 的forward_hidden专门把最终 LayerNorm 后的隐藏状态暴露给后训练阶段的辅助头value head、reward head这正是共享骨干设计的具体体现。扎实理解注意力机制是看懂这些后续所有阶段的基石。【免费下载链接】train-llm-from-scratchA straightforward method for training your LLM, from downloading data to generating text.项目地址: https://gitcode.com/GitHub_Trending/tr/train-llm-from-scratch创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表