资讯详情

资讯详情

Transformers 与语言模型完全指南:从 Pre-Norm、RoPE 到 BERT/GPT/T5、LoRA、Scaling Laws 与 LLM 评估体系

Transformers 与语言模型完全指南从 Pre-Norm、RoPE 到 BERT/GPT/T5、LoRA、Scaling Laws 与 LLM 评估体系【免费下载链接】maths-cs-ai-compendiumBecome a cracked AI/ML researcher/engineer with this unconventional textbook covering maths, computing, and ML with intuition.项目地址: https://gitcode.com/GitHub_Trending/mat/maths-cs-ai-compendium本文围绕 Transformer 在自然语言处理中的三大范式编码器 BERT、解码器 GPT、编码器-解码器 T5展开完整覆盖位置编码正弦、RoPE、ALiBi的相对位置数学推导、预训练目标MLM、CLM、span corruption、参数高效微调Adapters、LoRA、Prefix Tuning、Chinchilla scaling laws、Mixture of Experts以及从 BLEU/ROUGE 到 LLM-as-judge 的全套评估指标与基准。读完后你将能够从零实现 Pre-Norm Transformer 编码器块与因果注意力掩码、用 LoRA 以极小参数代价适配大模型、并按 Chinchilla 法则规划训练预算——这正是现代大语言模型的完整蓝图。1. Transformer 架构的关键设计选择Transformer 用自注意力取代了循环结构成为语言理解与生成的主流架构。回顾核心运算缩放点积注意力计算 $\text{softmax}(QK^T / \sqrt{d_k}) V$其中查询 Q、键 K、值 V 是输入的线性投影多头注意力并行运行 $h$ 个注意力头每个头使用不同的学习投影最后拼接结果。Transformer 块在此基础上叠加残差连接、层归一化LayerNorm与逐位置前馈网络FFN。1.1 Pre-Norm 与 Post-Norm一个微妙但重要的架构选择是层归一化的位置。原始 Transformer 采用post-norm残差与归一化位于子层之后即 $\text{LayerNorm}(x \text{Sublayer}(x))$。而现代模型大多采用pre-norm先归一化再过子层即 $x \text{Sublayer}(\text{LayerNorm}(x))$。Pre-norm 在训练中更稳定因为残差连接让梯度可以沿恒等路径直接穿过不受归一化的影响——这使得训练非常深的模型时不需要非常小心地设置学习率 warmup。在 Chapter 06 深度学习 中可以看到LayerNorm 沿每个样本的特征维度归一化、不依赖 batch 内其他样本这正是 Transformer 选用它的原因。1.2 FFN 子层块的关键-值记忆每个 Transformer 块中的前馈子层是逐 token 位置独立应用的两层 MLP$$\text{FFN}(x) W_2 \cdot \text{GELU}(W_1 x b_1) b_2$$内层维度通常取模型维度的 4 倍例如 $d_{\text{model}} 768$ 时 $d_{\text{ff}} 3072$。该 FFN 约占每个块参数的三分之二并被认为是存储训练中学得的事实性知识的关键-值key-value记忆。2. 位置编码让注意力看见顺序注意力本身是置换等变的把输入当作集合而非序列因此必须显式注入顺序信息。正弦编码原始 Transformer使用不同频率的固定正弦/余弦函数学习式位置嵌入为每个位置添加一个可训练向量BERT 和 GPT-2 采用。两者都是绝对编码位置 5 无论上下文如何都获得同一个向量。2.1 RoPE旋转位置编码**RoPERotary Position Embedding**通过在二维子空间中旋转查询与键向量来编码位置。对维度对 $(q_{2i}, q_{2i1})$位置 $m$ 以角度 $m\theta_i$其中 $\theta_i 10000^{-2i/d}$施加旋转RoPE 的精髓在于旋转后的查询与键的点积 $q^T k$ 只依赖相对位置$m - n$而非绝对位置。推导如下——把旋转写成 $q R_m q$、$k R_n k$$R_m$ 为块对角旋转矩阵注意力分数变为$$q^T k (R_m q)^T (R_n k) q^T R_m^T R_n , k q^T R_{n-m} , k$$最后一步来自旋转群性质$R_m^T R_n R_{n-m}$先逆旋转 $m$ 再正旋转 $n$等价于旋转 $n-m$。因此注意力分数只依赖相对距离 $n-m$。模型由此获得了天然的距离概念——无需任何学习的位置参数就能泛化到训练中未见过的序列长度。2.2 ALiBi更简单的线性偏置**ALiBiAttention with Linear Biases**更简洁直接给注意力分数加一个基于距离的固定线性惩罚$$\text{score}_{ij} q_i^T k_j - m \cdot |i - j|$$其中 $m$ 是头特定的斜率。不同头使用不同斜率使一些头聚焦局部、另一些头看全局。ALiBi 不需要任何学习的位置参数且对长于训练长度的序列泛化良好。3. 三大范式编码器、解码器、编码器-解码器Transformer 语言模型的三大范式是encoder-only、decoder-only与encoder-decoder。它们的区别在于模型能看到什么注意力掩码以及如何训练。3.1 BERT编码器专用与 MLM/NSP 预训练BERTBidirectional Encoder Representations from TransformersDevlin et al., 2019是经典的编码器专用模型。它用完全双向注意力处理文本每个 token 都可以注意左右所有 token。这赋予 BERT 丰富的上下文表示但意味着它无法自回归地生成文本。BERT 的预训练有两个目标掩码语言建模MLM随机掩码 15% 的输入 token 并让模型预测它们。被选中的 token 中80% 替换为 [MASK]10% 替换为随机词10% 保持不变防止模型只学会看到 [MASK] 才预测。训练目标$$\mathcal{L}{\text{MLM}} -\sum{i \in \mathcal{M}} \log P(w_i \mid w_{\backslash \mathcal{M}})$$其中 $\mathcal{M}$ 是掩码位置集合$w_{\backslash \mathcal{M}}$ 是这些位置被掩码后的句子。这是一种去噪目标模型学习重建被破坏的输入。下一句预测NSP训练 BERT 判断两句是否在原文中相邻。句首的 [CLS] token 用于这个二分类。NSP 旨在帮助需要理解句间关系的任务如问答但后续工作RoBERTa表明它贡献甚微、可以丢弃。BERT 的适配方式在预训练表示上叠加任务专用头一个简单的线性层并**微调fine-tuning**整个模型。分类任务用 [CLS] 表示token 级任务NER、词性标注用每个 token 的表示。这种方式把预训练学到的语言知识以较少的标注数据迁移到新任务。3.2 GPT解码器专用与因果语言建模GPTGenerative Pre-trained TransformerRadford et al., 2018是经典的解码器专用模型使用因果自回归注意力每个 token 只能注意之前位置含自身的 token。实现上是在注意力矩阵中把未来位置的分值置为 $-\infty$在 softmax 之前。训练目标是简单的因果语言建模CLM$$\mathcal{L}{\text{CLM}} -\sum{i1}^{n} \log P(w_i \mid w_1, \ldots, w_{i-1})$$这与 Chapter 07 文本处理 中 n-gram 语言模型的目标相同但用 Transformer 参数化后可以条件于整个前文而不只是最后 $k-1$ 个 token。GPT-215 亿参数展示了强零样本能力无需任何微调仅凭自然语言提示Translate English to French: …就能执行任务GPT-31750 亿参数证明单纯规模即可带来上下文学习in-context learning在提示中给出若干输入-输出示例模型无需任何梯度更新即可执行新任务。3.3 T5 与 BART编码器-解码器与 span corruptionT5Text-to-Text Transfer TransformerRaffel et al., 2020把所有 NLP 任务都框定为 text-to-text输入是文本串可带任务前缀如 translate English to German:输出也是文本串。编码器用双向注意力处理输入解码器用交叉注意力对编码器自回归地生成输出。T5 的预训练用span corruption随机选取连续 token 片段替换为哨兵 token模型必须生成原始 token。例如 The cat sat on the mat 变成输入 The [X] on [Y]目标是 [X] cat sat [Y] the mat。这是 BERT MLM 从单 token 到片段的推广。BARTLewis et al., 2020同样是去噪目标预训练的编码器-解码器模型但采用更广的破坏策略集合token 掩码、token 删除、span 掩码、句子重排、文档旋转。破坏的多样性迫使模型学习更鲁棒的表示。4. 参数高效微调PEFT随着模型变大全量微调更新所有参数变得不现实一个 175B 参数的模型仅存优化器状态就需要数百 GB。**参数高效微调PEFT**只适配极小部分参数Adapters在既有 Transformer 层之间插入小的瓶颈层通常是下投影到小维度 → 非线性 → 上投影回原维度的两层线性只训练 adapter 权重、冻结原模型。新增参数不到 5%在多数任务上匹配全量微调的效果。LoRALow-Rank Adaptation不新增层而是修改权重矩阵本身。不更新完整权重矩阵 $W$而是学习其更新的低秩分解$W W BA$其中 $B$ 为 $d \times r$、$A$ 为 $r \times d$且 $r \ll d$典型 $r 4 \sim 64$。原始 $W$ 冻结只训练 $A$ 和 $B$。推理时更新可直接合并回原权重零额外延迟。Prefix tuning在每个注意力层的键和值矩阵前拼接一串可学习的虚拟 token序列。模型把这些前缀向量当作真实 token 来注意只训练前缀参数。它类似 prompt tuning但工作在激活空间而非嵌入空间。5. 提示工程与上下文学习提示工程是设计输入文本以从预训练模型中引出期望行为、而不更新任何参数的艺术Zero-shot prompting用自然语言描述任务Classify the sentiment of the following review:Few-shot prompting在真正的问题之前给出若干输入-输出示例Chain-of-thoughtCoT加上 Lets think step by step 或在示例中给出推理过程通过引导模型分解问题在算术与逻辑推理任务上大幅提升表现。**上下文学习ICL**指大模型能仅凭提示中的示例学会执行任务、权重完全不改变——示例充当了某种隐式任务说明。其机制仍是开放的研究问题一种假说是注意力层在前向传播中实现了某种梯度下降实际上在上下文示例上训练了自己。6. Scaling Laws 与 Mixture of Experts6.1 从 Kaplan 到 ChinchillaScaling laws描述了模型规模、数据规模、计算预算与性能损失之间的可预测关系。Kaplan et al. (2020) 发现损失对每个变量都遵循幂律$$L(N) \propto N^{-\alpha_N}, \quad L(D) \propto D^{-\alpha_D}, \quad L(C) \propto C^{-\alpha_C}$$其中 $N$ 是参数量、$D$ 是数据规模、$C$ 是计算预算。这些幂律跨越多个数量级成立暗示单纯放大即可获得可预测的改进。Chinchilla scaling lawsHoffmann et al., 2022指出大多数大模型训练不足。对固定计算预算 $C$最优分配使模型规模与训练数据等比例增长$$N_{\text{opt}} \propto C^{0.5}, \quad D_{\text{opt}} \propto C^{0.5}$$即计算预算翻倍时模型规模与数据规模都应乘 $\sqrt{2}$而不是只把模型做大。Kaplan 建议 $N$ 比 $D$ 增长更快导致了很大但训练不足undertrained的模型。Chinchilla70B 参数、1.4T token在相同计算预算下追平了 Gopher280B 参数、300B token证明早期模型严重数据饥渴。实用经验法则每个参数约训练 20 个 token。6.2 Mixture of Experts参数涨、计算不涨MoE通过增加容量而不等比例增加计算来扩展模型。它用多个专家expertFFN 层加一个**门控网络router**替代单一大的前馈层为每个 token 选择激活哪些专家。门控函数计算每个专家的 routing 分数并选 top-$k$典型 $k 1$ 或 $k 2$$$G(x) \text{TopK}(\text{softmax}(W_g x))$$只有被选中的专家处理该 token所以计算成本随 $k$激活专家数而非总专家数 $E$ 增长8 个专家、top-2 路由的模型参数是密集模型的 4 倍计算只多 2 倍。MoE 的关键挑战是负载均衡若路由器把大部分 token 发给少数热门专家其余专家就被浪费。训练时加入辅助的负载均衡损失鼓励专家利用率均匀$$\mathcal{L}{\text{balance}} E \cdot \sum{i1}^{E} f_i \cdot p_i$$其中 $f_i$ 是分配给专家 $i$ 的 token 比例$p_i$ 是专家 $i$ 的平均路由概率。当两者都均匀各为 $1/E$时乘积最小。**专家并行expert parallelism**把不同专家放到不同加速器上前向传播中通过 all-to-all 通信把 token 路由到承载其专家的设备再路由回来——这个通信成本就是大规模 MoE 的主要工程挑战。Switch Transformer、Mixtral、GShard 都用 MoE 以可控的推理成本取得了强性能。7. 语言模型评估指标、基准与陷阱构建模型只是一半工作衡量其是否有效是另一半。NLP 评估尤其困难翻译可以有多种正确形式、摘要可以不复用参考原文、有帮助且无害的回复在合理的人类之间也会有分歧。7.1 精确匹配与 token 级指标Exact matchEM输出是否精确等于金标准答案用于抽取式问答SQuAD等短答案、无歧义任务。EM 很严格——New York City 与 new york city 不做归一化就不匹配——但简单无歧义。Token 级指标把 NLP 当作 token 级分类问题使用精确率、召回率与 F1定义见 Chapter 06 经典 ML 及本教材的分类评估部分。精确率$P \text{TP} / (\text{TP} \text{FP})$预测 token 中正确的比例——只预测很少实体但全对精确率很高召回率$R \text{TP} / (\text{TP} \text{FN})$金标准 token 被找到的比例——把所有 token 都预测为实体则召回率完美但精确率极差F1是调和平均 $F_1 \frac{2PR}{P R}$。用调和而非算术平均是因为它会惩罚不平衡$P$、$R$ 任一偏低$F_1$ 就偏低。对 NERF1 按实体类型分别计算再宏平均词性标注更常用 token 级准确率因为每个 token 都有标签。Span-level F1SQuAD 使用比较预测 span 与金标准 span 的 token 集合比 EM 宽容。金标准是 the Eiffel Tower、模型预测 Eiffel Tower 时EM 为零但 span F1 很高5 个 token 中 4 个重叠。7.2 机器翻译与摘要BLEU、ROUGE、METEOR、ChrFBLEUBilingual Evaluation UnderstudyPapineni et al., 2002是机器翻译的经典指标衡量候选译文与一个或多个参考译文的 n-gram 重叠结合多个 n-gram 层级的精确率与简短惩罚$$\text{BLEU} \text{BP} \cdot \exp!\left(\sum_{n1}^{N} w_n \log p_n\right)$$其中 $p_n$ 是修正 n-gram 精确率候选中每个 n-gram 的计数被截断到任一参考中的最大计数防止 the the the the 这类退化候选得分偏高权重 $w_n$ 通常均匀$w_n 1/N$$N 4$。简短惩罚$\text{BP} \min(1, \exp(1 - r/c))$$c$ 为候选长度、$r$ 为参考长度惩罚比参考短得多的候选——否则模型输出极少而安全的词就能拿到高精确率。BLEU 在语料级跨多句平均与人类判断相关性尚可句级则较差它奖励精确 n-gram 匹配而漏掉合法改写——the cat is on the mat 与 a feline sits atop the rug 的双词组重叠为零尽管意思相同且它完全忽略召回率。ROUGELin, 2004是摘要标准指标与强调精确率的 BLEU 相反ROUGE 强调召回率ROUGE-Nn-gram 的召回率 $\text{ROUGE-N} \frac{|\text{n-grams}{\text{ref}} \cap \text{n-grams}{\text{cand}}|}{|\text{n-grams}_{\text{ref}}|}$ROUGE-1、ROUGE-2 最常用ROUGE-L用候选与参考的最长公共子序列LCS捕捉句级词序而不要求连续匹配。LCS 长度除以参考长度得召回、除以候选长度得精确、F 值组合两者$$R_{\text{LCS}} \frac{\text{LCS}(X, Y)}{m}, \quad P_{\text{LCS}} \frac{\text{LCS}(X, Y)}{n}, \quad F_{\text{LCS}} \frac{(1 \beta^2) R_{\text{LCS}} P_{\text{LCS}}}{R_{\text{LCS}} \beta^2 P_{\text{LCS}}}$$其中 $m$、$n$ 是参考与候选长度$\beta$ 通常取为偏向召回$\beta \to \infty$ 即纯召回。LCS 用动态规划计算时间 $O(mn)$类似 文本处理 中讲过的编辑距离。METEORBanerjee and Lavie, 2005针对 BLEU 的弱点纳入同义词、词干还原与词序先通过精确匹配、词干匹配Porter 词干还原与同义词匹配WordNet对齐候选与参考的词再计算偏向召回的 unigram 精确率与召回率的调和平均并施加碎片化惩罚fragmentation penalty惩罚匹配词顺序与参考不一致的候选。ChrF在字符 n-gram 而非词 n-gram 上计算 F 值对形态变化稳健对黏着语至关重要并部分处理分词差异。ChrF 在字符 n-gram 之外加入词双词组已成为机器翻译中与 BLEU 并用的推荐指标尤其适用于形态丰富的语言。7.3 内在指标困惑度、BPB、BERTScore、BLEURT、COMETPerplexity困惑度衡量模型对留出测试集的预测能力是语言模型的标准内在指标$\text{PPL} \exp(-\frac{1}{N} \sum_{i} \log P(w_i \mid w_{i}))$越低越好。注意困惑度只在同分词器之间可比因为不同分词器对同一段文本产生不同的序列长度 $N$词汇量大的模型单 token 困惑度偏低但每句处理的 token 更少。Bits-per-byteBPB按 UTF-8 字节数而非 token 数归一化从而与分词无关\text{BPB} \frac{-\sum_{i} \log_2 P(w_i \mid w_{i})}{\text{number of UTF-8 bytes}}BERTScoreZhang et al., 2020跳出表面 n-gram 匹配在嵌入空间中计算相似度候选的每个 token 用上下文嵌入通常来自预训练 BERT的余弦相似度匹配参考中最相似的 token聚合为精确率、召回率与 F1$$R_{\text{BERT}} \frac{1}{|r|} \sum_{r_i \in r} \max_{c_j \in c} \cos(r_i, c_j), \quad P_{\text{BERT}} \frac{1}{|c|} \sum_{c_j \in c} \max_{r_i \in r} \cos(c_j, r_i)$$其中 $r_i$、$c_j$ 是参考与候选 token 的上下文嵌入。它捕捉 n-gram 指标漏掉的语义相似automobile 与 car 共享字符为零但 BERT 嵌入相近得分很高。BLEURTSellam et al., 2020更进一步直接在人类质量评分上微调 BERT给定参考-候选对输出标量质量分。它先在合成数据对参考翻译做随机扰动并用 BLEU/METEOR 等打分上训练再用人类评分微调与人类判断的相关性优于任何表面级指标。COMETRei et al., 2020以源句、参考、候选三者为条件而非仅参考与候选的学习式机器翻译指标用多语言编码器XLM-R嵌入三者并预测质量分。因为能看到源句COMET 可以发现只看参考的指标漏掉的意义错误如通顺但事实错误的译文。LLM-as-judge大规模评估的现代做法。不计算对参考的指标而是提示一个强语言模型评估输出质量。裁判接收输入、模型响应以及可选的参考答案输出评分如 1–5 分或成对偏好响应 A 优于 B。成对比较Chatbot Arena 所用是最可靠的 LLM 裁判形式裁判看两个响应选更好的避免绝对打分的校准问题。结果聚合成Elo 分源自国际象棋。模型 $A$ 对 $B$ 的期望胜率$$P(A \succ B) \frac{1}{1 10^{(R_B - R_A) / 400}}$$每次比较后更新$R_A R_A K(S - P(A \succ B))$$S \in {0, 1}$ 是实际结果、$K$ 控制更新幅度。持续战胜强对手的模型上升快输给弱对手的下降。位置偏差裁判倾向于偏好先呈现的响应某些模型则是后呈现的。交换两个顺序各评一次取平均可缓解。冗长偏差裁判倾向偏好更长、更详细的响应即使简洁答案更好。自一致性检查裁判对同一输入多次评估是否给出相同评分高方差说明评估信号噪声大。标注者间一致性Cohens kappa 或 Krippendorffs alpha衡量多个裁判是否一致是评估可靠性的上限。数据污染Contamination若评估数据出现在模型训练集中基准分数虚高且失去意义——对网络爬取数据训练的 LLM 尤为危险流行基准很可能已被收入训练集。缓解手段使用未公开发布的留出测试集、周期性重新生成问题的动态基准、嵌入基准数据中的金丝雀字符串canary strings检测泄漏、对比污染子集与干净子集上的表现。7.4 标准基准理解、推理、代码与安全标准 NLU 基准GLUE / SuperGLUE覆盖情感SST-2、文本相似度STS-B、自然语言推理MNLI、RTE、共指WSC、问答BoolQ等多任务基准。GLUE 现在被认为已饱和多数任务超过人类水平SuperGLUE 仍然更难。MMLU57 个学科数学、历史、法律、医学、计算机等的多选题检验模型在预训练中是否吸收了广博知识MMLU-Pro 更难10 个选项而非 4 个并加入多步推理题。HellaSwag选择场景最合理的续写错误选项由模型对抗生成——表面合理、语义错误。WinoGrande仅一词之差的最小对比对测试常识共指消解。ARC小学科学题easy/challenge 两套测试事实与推理能力。推理与数学基准GSM8K约 8500 道需要多步算术推理的小学数学应用题是基础数学推理与 CoT 提示本文第 5 节的标准基准。MATH竞赛级数学题代数、数论、几何、计数、概率需要多步符号推理MATH-500 是常被报告的 500 题子集。AIME美国数学邀请赛级题目需要多步深层数学推理。文档中举例 DeepSeek-R1 在 AIME 2024 上得 79.8%说明经强化学习的推理模型已可接近强人类选手。HumanEval / MBPP代码生成检查模型代码是否通过单元测试。HumanEval 含 164 道 Python 题函数签名 docstring需生成函数体。指标是passk——$k$ 个生成解中至少一个通过全部测试的概率单样本估计$$\text{pass}k 1 - \frac{\binom{n-c}{k}}{\binom{n}{k}}$$$n$ 为总样本数、$c$ 为通过数该公式修正了取 $k$ 个中最好的偏差。SWE-bench更进一步考察模型能否修改现有代码库来解决真实的 GitHub issue——对实际软件工程能力的更硬考验。GPQAGraduate-Level Google-Proof QA生物、物理、化学的专家级问题连领域专家都答不好检验的是真正理解而非模式匹配Diamond 子集最难。安全与对齐基准TruthfulQA测试模型是否复述常见误解——题目设计成网络上最流行的答案恰好是错的如吞下的口香糖会停留 7 年真相是会正常排出。记忆了流行错误说法的模型得分差。BBQBias Benchmark for QA年龄、性别、种族、宗教等类别的社会偏见Toxigen模型对特定人群群体生成毒性内容的倾向。MT-Bench80 道精心设计的问题写作、角色扮演、推理、数学、代码、抽取、STEM、人文评估多轮对话LLM 裁判GPT-4按 1–10 打分多轮格式测试模型能否追问、保持上下文、处理澄清请求。Chatbot ArenaLMSYS真实用户做盲测成对比较不知道模型身份即投票。其 Elo 排行榜因反映真实用户在多样、未精选提示上的偏好被视为通用 LLM 质量最生态有效的评估。AlpacaEval自动化成对评估在固定指令集上将模型输出与参考模型GPT-4对比由裁判模型判定胜率AlpacaEval 2.0用长度控制的胜率length-controlled win rate修正冗长偏差。任务专用指标WER语音识别$\text{WER} (S D I) / N$$S$、$D$、$I$ 为替换、删除、插入错误$N$ 为参考词数——即按参考长度归一化的词级编辑距离Slot F1任务导向对话是否正确抽取结构化信息例如从 Book me a flight to Paris tomorrow 中抽出 destination: Paris 与 date: tomorrow引用准确性RAG 系统检查模型生成的引用是否真正支撑其声明——将声明与检索到的段落核对统计完全/部分/未获支持的比例。7.5 评估陷阱应试训练Teaching to the test优化基准分数而非真实能力。在 MMLU 式多选题上微调的模型 MMLU 分高但同样的问题以开放式提问就可能失败。指标博弈Metric gaming模型可被优化出在自动指标上得分高高 BLEU、低困惑度的输出而不真正更好。BLEU 最优的翻译往往是安全、平淡的意译而非自然流畅的译文。基准饱和Benchmark saturation模型逼近或超过人类水平后基准不再提供信息。GLUE、SQuAD 1.1 等已饱和。领域不断造更难的基准但创建—饱和—替换的循环让纵向比较困难。人类评估仍是金标准但昂贵、缓慢、难复现。不同标注者群体众包 vs 领域专家、不同文化、不同语言会给出不同判断——报告标注者间一致性与标注者背景信息对可复现性至关重要。8. 编程任务从零实现核心组件本章配套的三道编码任务建议用 Colab 或 notebook 运行把上述概念落到代码上全部使用 JAX。8.1 任务一从零实现 Transformer 编码器块实现完整的 Transformer 编码器块多头注意力、前馈、残差连接、层归一化并可视化每个注意力头的模式。注意实现采用的是pre-norm结构import jax import jax.numpy as jnp import matplotlib.pyplot as plt def layer_norm(x, gamma, beta, eps1e-5): mean x.mean(axis-1, keepdimsTrue) var x.var(axis-1, keepdimsTrue) return gamma * (x - mean) / jnp.sqrt(var eps) beta def multi_head_attention(Q, K, V, W_q, W_k, W_v, W_o, n_heads): B, T, D Q.shape head_dim D // n_heads q Q W_q # (B, T, D) k K W_k v V W_v # Reshape to (B, n_heads, T, head_dim) q q.reshape(B, T, n_heads, head_dim).transpose(0, 2, 1, 3) k k.reshape(B, T, n_heads, head_dim).transpose(0, 2, 1, 3) v v.reshape(B, T, n_heads, head_dim).transpose(0, 2, 1, 3) scores q k.transpose(0, 1, 3, 2) / jnp.sqrt(head_dim) weights jax.nn.softmax(scores, axis-1) out (weights v).transpose(0, 2, 1, 3).reshape(B, T, D) return out W_o, weights def transformer_block(x, params): # Pre-norm multi-head self-attention normed layer_norm(x, params[ln1_g], params[ln1_b]) attn_out, weights multi_head_attention( normed, normed, normed, params[W_q], params[W_k], params[W_v], params[W_o], n_heads4 ) x x attn_out # Pre-norm feed-forward normed layer_norm(x, params[ln2_g], params[ln2_b]) ff jax.nn.gelu(normed params[W1] params[b1]) ff ff params[W2] params[b2] x x ff return x, weights # Initialise parameters d_model, d_ff, n_heads 32, 128, 4 key jax.random.PRNGKey(42) keys jax.random.split(key, 10) params { W_q: jax.random.normal(keys[0], (d_model, d_model)) * 0.05, W_k: jax.random.normal(keys[1], (d_model, d_model)) * 0.05, W_v: jax.random.normal(keys[2], (d_model, d_model)) * 0.05, W_o: jax.random.normal(keys[3], (d_model, d_model)) * 0.05, ln1_g: jnp.ones(d_model), ln1_b: jnp.zeros(d_model), ln2_g: jnp.ones(d_model), ln2_b: jnp.zeros(d_model), W1: jax.random.normal(keys[4], (d_model, d_ff)) * 0.05, b1: jnp.zeros(d_ff), W2: jax.random.normal(keys[5], (d_ff, d_model)) * 0.05, b2: jnp.zeros(d_model), } # Test with random input x jax.random.normal(keys[6], (2, 8, d_model)) # batch2, seq_len8 out, attn_weights transformer_block(x, params) print(fInput shape: {x.shape}) print(fOutput shape: {out.shape}) print(fAttention weights shape: {attn_weights.shape}) # (B, n_heads, T, T) # Visualise attention patterns for each head fig, axes plt.subplots(1, 4, figsize(16, 3.5)) for h in range(4): im axes[h].imshow(attn_weights[0, h], cmapBlues, vmin0) axes[h].set_title(fHead {h}) axes[h].set_xlabel(Key pos); axes[h].set_ylabel(Query pos) plt.suptitle(Multi-Head Attention Patterns) plt.tight_layout(); plt.show()要点multi_head_attention中(B, T, D)张量先按头切分为(B, n_heads, T, head_dim)再算注意力分数除以 $\sqrt{\text{head_dim}}$ 实现缩放transformer_block严格按 pre-norm 顺序——先layer_norm再进子层残差加法在归一化之外——这正是第 1.1 节稳定性论证的代码体现。8.2 任务二因果注意力掩码 vs 双向注意力实现因果自回归注意力掩码并与双向注意力对比验证掩码确实阻止了未来信息流向过去import jax import jax.numpy as jnp import matplotlib.pyplot as plt def attention(Q, K, V, maskNone): d_k Q.shape[-1] scores Q K.T / jnp.sqrt(d_k) if mask is not None: scores jnp.where(mask, scores, -1e9) weights jax.nn.softmax(scores, axis-1) return weights V, weights seq_len, d_model 6, 8 key jax.random.PRNGKey(0) k1, k2, k3 jax.random.split(key, 3) Q jax.random.normal(k1, (seq_len, d_model)) K jax.random.normal(k2, (seq_len, d_model)) V jax.random.normal(k3, (seq_len, d_model)) # Bidirectional (encoder-style): all positions visible bidir_mask jnp.ones((seq_len, seq_len), dtypebool) bidir_out, bidir_weights attention(Q, K, V, bidir_mask) # Causal (decoder-style): only past and current positions visible causal_mask jnp.tril(jnp.ones((seq_len, seq_len), dtypebool)) causal_out, causal_weights attention(Q, K, V, causal_mask) fig, axes plt.subplots(1, 3, figsize(14, 4)) tokens [ft{i} for i in range(seq_len)] axes[0].imshow(bidir_weights, cmapBlues, vmin0, vmax0.5) axes[0].set_title(Bidirectional Attention\n(BERT-style)) axes[0].set_xticks(range(seq_len)); axes[0].set_xticklabels(tokens) axes[0].set_yticks(range(seq_len)); axes[0].set_yticklabels(tokens) axes[1].imshow(causal_mask.astype(float), cmapGreys, vmin0, vmax1) axes[1].set_title(Causal Mask\n(1 allowed, 0 blocked)) axes[1].set_xticks(range(seq_len)); axes[1].set_xticklabels(tokens) axes[1].set_yticks(range(seq_len)); axes[1].set_yticklabels(tokens) axes[2].imshow(causal_weights, cmapBlues, vmin0, vmax0.5) axes[2].set_title(Causal Attention\n(GPT-style)) axes[2].set_xticks(range(seq_len)); axes[2].set_xticklabels(tokens) axes[2].set_yticks(range(seq_len)); axes[2].set_yticklabels(tokens) for ax in axes: ax.set_xlabel(Key); ax.set_ylabel(Query) plt.tight_layout(); plt.show() # Verify: in causal attention, output at position i depends only on positions i print(Causal attention weight at position 2 (should only attend to 0, 1, 2):) print(f Weights: {causal_weights[2]}) print(f Sum of future weights (should be ~0): {causal_weights[2, 3:].sum():.6f})要点jnp.tril生成下三角布尔掩码jnp.where(mask, scores, -1e9)把被屏蔽位置的分值压到 $-10^9$近似 $-\infty$softmax 后其权重趋于 0——最后一行打印验证了位置 2 对未来位置的注意力权重之和约为 0。这就是第 3.2 节 GPT 因果注意力的最小实现。8.3 任务三LoRA 低秩适配实现 LoRA展示它如何用远少于全量微调的可训练参数修改权重矩阵import jax import jax.numpy as jnp d_model 256 rank 4 # LoRA rank (much smaller than d_model) key jax.random.PRNGKey(42) k1, k2, k3 jax.random.split(key, 3) # Original frozen weight matrix W_frozen jax.random.normal(k1, (d_model, d_model)) * 0.02 # LoRA matrices (only these are trainable) B jnp.zeros((d_model, rank)) # initialised to zero A jax.random.normal(k2, (rank, d_model)) * 0.01 # random init # Forward pass: W_effective W_frozen B A x jax.random.normal(k3, (8, d_model)) # Without LoRA y_original x W_frozen.T # With LoRA W_effective W_frozen B A y_lora x W_effective.T # Parameter counts full_params d_model * d_model lora_params d_model * rank rank * d_model # B A print(fModel dimension: {d_model}) print(fLoRA rank: {rank}) print(fFull fine-tuning parameters: {full_params:,}) print(fLoRA parameters: {lora_params:,}) print(fParameter reduction: {full_params / lora_params:.1f}x) print(f\nSince B is initialised to zeros, initial LoRA output matches original:) print(f Max difference: {jnp.abs(y_original - y_lora).max():.2e}) # Simulate training: only update A and B def lora_forward(A, B, W_frozen, x): return x (W_frozen B A).T def dummy_loss(A, B, W_frozen, x, target): pred lora_forward(A, B, W_frozen, x) return jnp.mean((pred - target) ** 2) # Target: some transformation of x target x jax.random.normal(jax.random.PRNGKey(99), (d_model, d_model)).T * 0.02 grad_fn jax.jit(jax.grad(dummy_loss, argnums(0, 1))) lr 0.01 for step in range(200): gA, gB grad_fn(A, B, W_frozen, x, target) A A - lr * gA B B - lr * gB loss_before dummy_loss(jnp.zeros_like(A), jnp.zeros_like(B), W_frozen, x, target) loss_after dummy_loss(A, B, W_frozen, x, target) print(f\nLoss before LoRA: {loss_before:.6f}) print(fLoss after LoRA: {loss_after:.6f}) print(fEffective weight change rank: {jnp.linalg.matrix_rank(B A)})要点$B$ 初始化为零保证训练起点 $W BA W$不破坏预训练行为只训练 $A$、$B$$d \cdot r r \cdot d$ 个参数 vs 全量的 $d^2$ 个$d 256$、$r 4$ 时约 32 倍参数缩减最后的jnp.linalg.matrix_rank(B A)验证有效权重更新的秩不超过 $r$与第 4 节的低秩分解定义吻合。结语Transformer 用自注意力取代循环结构后成为语言理解与生成的主导架构而 BERT、GPT、T5 分别确立了编码器、解码器、编码器-解码器三条路线MLM、CLM、span corruption 三大预训练目标则决定了各路线的能力边界。本文覆盖的 pre-norm 稳定性、RoPE 的相对位置数学、LoRA 的低秩适配、Chinchilla 的 20 token/参数法则、MoE 的负载均衡以及从 BLEU 到 LLM-as-judge 的评估体系共同构成现代大语言模型的完整蓝图。若需继续深入可参考同系列的 嵌入与序列模型seq2seq 如何过渡到 Transformer与 进阶文本生成生成技术、RAG 与推理模型。【免费下载链接】maths-cs-ai-compendiumBecome a cracked AI/ML researcher/engineer with this unconventional textbook covering maths, computing, and ML with intuition.项目地址: https://gitcode.com/GitHub_Trending/mat/maths-cs-ai-compendium创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
觉得有用,分享给同行:

为您的企业打造数字门面

稳重轻奢商务风格,端正雅致视觉,长效耐看不易过时。

立即咨询 →