文章目录
-
- 前言
- [1. 先别急着写代码:默认值背后的那笔账](#1. 先别急着写代码:默认值背后的那笔账)
-
- [1.1 第一步:先定 baseline,现代版和老版差在哪](#1.1 第一步:先定 baseline,现代版和老版差在哪)
- [1.2 RMSNorm 和去 bias:小动作为什么能省出真金白银](#1.2 RMSNorm 和去 bias:小动作为什么能省出真金白银)
- [1.3 FFN 的演化:SwiGLU 不是"换个激活函数"那么简单](#1.3 FFN 的演化:SwiGLU 不是“换个激活函数”那么简单)
- [1.4 RoPE:把相对位置直接写进内积](#1.4 RoPE:把相对位置直接写进内积)
- [1.5 超参数共识:默认值背后是预算分配](#1.5 超参数共识:默认值背后是预算分配)
- [1.6 稳定性技巧:softmax 是最容易出事的地方](#1.6 稳定性技巧:softmax 是最容易出事的地方)
- [1.7 GQA / MQA:瓶颈怎么从"算不动"变成"搬不动"](#1.7 GQA / MQA:瓶颈怎么从“算不动”变成“搬不动”)
- [1.8 滑动窗口和混合注意力:用图结构控制长上下文成本](#1.8 滑动窗口和混合注意力:用图结构控制长上下文成本)
- [2. Embedding:从 token 到向量的第一步](#2. Embedding:从 token 到向量的第一步)
- [3. FFN 与 PreNorm:残差主干必须保持干净](#3. FFN 与 PreNorm:残差主干必须保持干净)
-
- [3.1 pre-norm block:残差主干必须保持干净](#3.1 pre-norm block:残差主干必须保持干净)
- [3.2 RMSNorm](#3.2 RMSNorm)
- [3.3 Linear 层](#3.3 Linear 层)
- [3.4 SiLU:ReLU 的平滑升级版](#3.4 SiLU:ReLU 的平滑升级版)
- [3.5 GLU 和 SwiGLU:门控不是换激活函数](#3.5 GLU 和 SwiGLU:门控不是换激活函数)
- [4. RoPE:旋转位置编码,把相对距离写进内积](#4. RoPE:旋转位置编码,把相对距离写进内积)
-
- [4.1 把 head 维度两两看成二维平面](#4.1 把 head 维度两两看成二维平面)
- [4.2 相对位置来自旋转矩阵的群性质](#4.2 相对位置来自旋转矩阵的群性质)
- [4.3 二维展开后能直接看到 sin 和 cos](#4.3 二维展开后能直接看到 sin 和 cos)
- [4.4 为什么要用一组频率](#4.4 为什么要用一组频率)
- [4.5 具体实现:rotate_half 只是矩阵乘法的向量化写法](#4.5 具体实现:rotate_half 只是矩阵乘法的向量化写法)
- [4.6 为什么 RoPE 只作用在 query 和 key 上](#4.6 为什么 RoPE 只作用在 query 和 key 上)
- [5. Attention 与 Transformer:把积木搭起来](#5. Attention 与 Transformer:把积木搭起来)
-
- [5.1 ScaledDotProductAttention](#5.1 ScaledDotProductAttention)
- [5.2 MultiHeadAttention](#5.2 MultiHeadAttention)
- [5.3 MultiHeadAttentionWithRoPE](#5.3 MultiHeadAttentionWithRoPE)
- [5.4 TransformerBlock](#5.4 TransformerBlock)
- [5.5 Transformer](#5.5 Transformer)
- [6. 训练优化:让模型真正学会说话](#6. 训练优化:让模型真正学会说话)
-
- [6.1 Loss:预测下一个 token](#6.1 Loss:预测下一个 token)
- [6.2 Optimizer:AdamW,以及它吃掉的内存](#6.2 Optimizer:AdamW,以及它吃掉的内存)
- [6.3 Cosine LR Scheduler](#6.3 Cosine LR Scheduler)
- [6.4 Gradient Clip](#6.4 Gradient Clip)
- [7. 资源核算:钱都花在矩阵乘法上](#7. 资源核算:钱都花在矩阵乘法上)

P.S. 推荐一个大神的教程给想要了解或者学习人工智能知识的读者,这个教程里内容讲解通俗易懂且风趣幽默,对我帮助很大。我想与大家分享这个宝藏教程,请点击下方链接查看, 传送门https://blog.csdn.net/qq_74013365
前言
想从零手写一个 decoder-only Transformer?先别急着抄代码。把模块照着架构图堆起来,那是拼积木的活儿;真正难的,是面对 pre-norm、RMSNorm、SwiGLU、RoPE、GQA 这一大排"现代默认值"时,你得搞清楚:它们凭什么站在这里?
我当年的做法很朴素:全都要。然后花三天训了一个 loss 纹丝不动的模型,最后发现是 norm 放错了位置。这篇就把我踩过的坑和算明白的账一次讲清楚。代码我全留着,一个不删。
1. 先别急着写代码:默认值背后的那笔账
现代大模型架构里那一堆变体,看着像时装周走秀:pre-norm、RMSNorm、SwiGLU、RoPE、GQA、滑动窗口、QK norm、z-loss、logit soft-capping......名字一个比一个唬人。
但扒开看,它们全在回答同一个问题:模型变大、上下文变长、训练规模上升之后,怎么让梯度、算力、显存和推理延迟别一起爆炸。
不是每个默认值都有严格的数学证明。很多选择是被大量模型和工程经验反复验证出来的------就像你妈炖肉"适量"放盐,看着不科学,但几十锅肉下来,它就是对的。
更重要的是别只记结论,要知道结论是从什么约束推出来的:残差路径为什么要保持干净、RMSNorm 为什么比 FLOPs 数字看起来更重要、GLU 为什么要配缩小的 d_ff、RoPE 为什么要作用在 query/key 上、GQA/MQA 到底在解决什么。
用一句话概括这组默认值:现代 decoder-only LLM 的架构优化,不是把 Transformer 改得更花,而是让残差主干更稳、FFN 更有表达力、位置关系更贴近 attention、推理缓存更轻------把实验预算留给数据、训练和系统实现。翻译成人话:该省的省,该花的花,别把钱花在梳妆打扮上。
1.1 第一步:先定 baseline,现代版和老版差在哪
一个最朴素的 Transformer block 可以写成两段:先做 self-attention,再做 FFN,中间都带残差连接。早期常见的是 post-norm,先把子层输出加回残差,再做 LayerNorm。符号写出来是这样:
h' = LayerNorm(h + Attention(h))
h_next = LayerNorm(h' + FFN(h'))
现代 decoder-only LLM 更常见的是 pre-norm,先归一化输入,再把子层输出加回原始残差流:
h' = h + Attention(Norm(h))
h_next = h' + FFN(Norm(h'))
两个式子只差一个括号的位置,训练含义天差地别。post-norm 里,残差相加之后马上被 LayerNorm 整体改写------说好的 identity path,直接变成"薛定谔的恒等路径"。pre-norm 里,h 可以沿着残差分支一路传到下一层,子层只是往里追加一个增量。
深层网络为什么难训练?因为每一层都可能放大、缩小或扭曲梯度。residual 的使命,就是给信息和梯度留一条近似恒等的高速公路:就算某层子模块这轮没学好,模型也能先把输入原样传下去,不至于整条路堵死。
pre-norm 的价值正在这里:norm 只服务子层计算,别来碰瓷残差主干。打个比方,post-norm 是"先合流再安检",pre-norm 是"先安检再合流"。队伍一长,区别就出来了------前者总有人在安检口被扣下来重新排队。
大模型层数加深以后,这个差异会放大成训练稳定性的差异:pre-norm 更容易避免梯度衰减和尖峰,也更容易扛得住较大的学习率。
当然,这不是说 post-norm 是原罪。BERT 这种老前辈就是 post-norm,活得好好的;有些新模型也会在残差之外再加一层 non-residual 的 norm。但我会把它们看成"保护残差主干"这个原则下的变体,而不是回到早期结构的理由。历史可以怀念,钱包不能乱花。
1.2 RMSNorm 和去 bias:小动作为什么能省出真金白银
LayerNorm 要干的事不少:算均值、算方差、归一化、再缩放平移:
mean = average(x)
var = average((x - mean)^2)
y = (x - mean) / sqrt(var + eps) * gamma + beta
RMSNorm 把减均值那步直接砍了,只留均方根缩放:
rms = sqrt(average(x^2) + eps)
y = x / rms * gamma
单看 FLOPs,这省不了多少------Transformer 里矩阵乘法才是吞算力的大户。但大模型跑起来,瓶颈经常不是"算",而是"搬":数据在显存里搬来搬去。norm 每层都做、每个 token 都做,算术量不大,但要读写 activation、参数、中间统计量。RMSNorm 少读写一条均值路径,kernel 更简单,wall-clock 上就能看到肉眼可见的收益。一句话:少搬一趟砖,就多一分流畅。
去 bias 也是同一个逻辑。线性层 y = xW + b 里那个 b,看着参数不多,但它要存储、要加载、要进优化器状态,而且在带 norm 的网络里,平移自由度本来就会被归一化抵消大半。
现代 LLM 于是纷纷在线性层和 norm 层把 bias 请了出去。不是抠门,是算过账的:这钱花得不值。就像公司里那个可有可无的岗位,裁了之后大家才发现,之前每个月都在为它付房租。
| 选择 | 常见现代做法 | 我会怎么理解 |
|---|---|---|
| Norm 位置 | pre-norm / non-residual norm | 让残差主干更接近 identity path,改善深层训练稳定性 |
| Norm 类型 | RMSNorm | 保留缩放稳定性,减少均值相关计算和数据搬运 |
| Bias | 多数线性层和 norm 不带 bias | 小参数收益有限,但会增加状态、读写和优化复杂度 |
1.3 FFN 的演化:SwiGLU 不是"换个激活函数"那么简单
Transformer block 里,FFN 通常占掉很大一部分参数和 FLOPs。最普通的 FFN 三件套:升维、激活、降维。
FFN(x) = activation(xW_up) W_down
如果 x 的维度是 d_model,中间维度是 d_ff,那么两层矩阵大约有 2 * d_model * d_ff 个参数。传统经验会取 d_ff = 4 * d_model,所以 FFN 参数量大约是 8 * d_model²。FFN 宽度从来不是小超参数------它直接决定每层一大半的预算花在哪。
GLU 系列换了个玩法:"内容分支 × 门控分支"。以 SwiGLU 为例:
FFN_GLU(x) = (activation(xW_gate) * xW_up) W_down
它多了一组上投影矩阵,要是还硬撑着 d_ff = 4 * d_model,参数和计算都会明显膨胀。怎么办?做个简单换算:普通 FFN 是 2 * d_model * 4d_model = 8d_model²;GLU 有三个矩阵,约为 3 * d_model * d_ff_glu。令二者接近,就得到 d_ff_glu ≈ 8/3 * d_model。
这就是为什么很多使用 SwiGLU/GeGLU 的模型,会把 FFN hidden size 调到约 2.66 倍 d_model。不是玄学,是预算对账对出来的。看看真实模型的比例:
| 模型 | d_ff / d_model |
|---|---|
| PaLM | 4 |
| Mistral 7B | 3.5 |
| LLaMA-2 70B | 3.5 |
| LLaMA 70B | 2.68 |
| Qwen 14B | 2.67 |
| DeepSeek 67B | 2.68 |
| Yi 34B | 2.85 |
| T5 v1.1 | 2.5 |
大部分模型都落在 2.5~2.7 附近,只有 PaLM、LLaMA-2 和 Mistral 稍微放飞了一点------看来不是所有学霸都遵守考场纪律。
从功能上看,GLU 不是把 ReLU 换成 Swish 就完事。门控分支让模型可以按通道决定哪些信息通过------相当于给 FFN 装了门禁,每个维度自己刷卡决定放不放行。经验上,SwiGLU/GeGLU 在很多模型中比 ReLU/GeLU 更稳定地带来收益;从工程上看,只要把中间维度按 8/3 调整,预算也能守得住。
这里有个对初学者很阴险的坑:如果做 ablation 比较 SwiGLU 和普通 SiLU FFN,却让两者参数量差一截,最后看到的 loss 差异根本说不清是门控的功劳,还是参数量在作弊。正确的姿势:SiLU baseline 用 4 * d_model,SwiGLU 用约 8/3 * d_model,匹配了参数规模再看曲线。科研界的公平竞争,从对齐参数量开始。
1.4 RoPE:把相对位置直接写进内积
位置编码要解决的问题很朴素:同一个词出现在不同位置,attention 得知道它俩隔了多远。绝对位置 embedding 是给 token 加个位置向量;正弦位置编码也是加法,只是位置向量带固定频率结构。加法方案简单,但 query/key 内积里会混入 token 内容、绝对位置和交叉项------一锅乱炖,很难保证打分只依赖相对距离。
我想要的形式是这样的:
score(i, j) = f(x_i, i)^T f(x_j, j) = g(x_i, x_j, i - j)
RoPE 的做法:把向量坐标两两配对,在二维平面里按位置角度旋转。设位置 i 对应旋转矩阵 R_i,attention 里实际参与内积的是 R_i q 和 R_j k,于是:
(R_i q)^T (R_j k) = q^T R_i^T R_j k = q^T R_{j-i} k
这一步就是 RoPE 的核心。位置 i 和 j 没有以两个独立标签进入打分,而是合并成 j - i 再进场。所以实现时,RoPE 要放在每一层 attention 的 query/key 路径上,而不是只在 embedding 层加一次位置向量。
它不是"给词贴位置标签",而是让 attention 打分天然携带相对位置信息。一句话:别给词编门牌号,要让词自己算出"我们隔了几个门"。
1.5 超参数共识:默认值背后是预算分配
看模型表格时,很容易把超参数当口诀背:FFN 四倍、head 多少、词表多大......背得比乘法表还熟。但更有用的问法是:这个数字到底在管哪笔预算?
FFN 宽度管每层参数和 MLP FLOPs;head 数和 head_dim 管 attention 表达和 kernel 形态;模型深宽比影响并行、延迟和每层吞吐;词表大小影响序列长度、embedding/softmax 参数和多语覆盖。
| 问题 | 保守默认值 | 推导和取舍 |
|---|---|---|
| FFN 要多宽 | 非 GLU:4 * d_model;GLU:约 8/3 * d_model | 普通 FFN 有两个矩阵,GLU 有三个矩阵;为了维持接近的预算,GLU 中间维度需要缩小 |
| head 怎么配 | head_dim * num_heads ≈ d_model | 保持投影维度和模型宽度接近,方便实现、并行和 checkpoint 迁移 |
| 深还是宽 | d_model / n_layer 常落在约 100-200 | 更深会增加串行层数和推理延迟;更宽会改变单层 GEMM、通信和显存压力 |
| 词表多大 | 单语约 30k-50k;多语/生产系统约 100k-250k | 大词表减少 token 数,但会增加 embedding 和输出 softmax 成本,也会改变稀有词覆盖 |
| 预训练要不要 dropout | 新模型常不用 dropout,但保留 weight decay | 海量 token 下过拟合不是唯一矛盾,weight decay 更多影响优化动力学和学习率 schedule |
以 FFN 宽度为例,4 * d_model 不是不可违反的律法。T5 的 11B 版本就曾用非常大的 FFN multiplier,照样训起来了。但后续许多模型还是回到更保守的比例,原因很现实:每次扩大 FFN 都在吃参数、算力和通信预算。收益要是没稳定盖过成本,默认值就该保守。成年人的世界,全是取舍。
词表也一样。小词表让 embedding 和 softmax 更便宜,但可能把一个词切得更碎、把序列拉得更长;大词表能减少 token 数、改善多语和特殊领域覆盖,但输出层更大、低频 token 更稀疏。单语模型常见 30k-50k,多语或生产系统常见 100k-250k。这不是谁更高级,是面向语种、数据分布和服务成本的不同折中------跟选手机内存一样:128G 够用,1T 更爽,但钱是你的。
至于 dropout 和 weight decay,别把"regularization"只理解成防止过拟合。预训练数据通常有数万亿 token,模型很少在同一批数据上反复练到纯记忆,dropout 的必要性自然下降;但 weight decay 依然常见,因为它会和学习率、cosine schedule、参数范数演化相互作用,影响训练动力学。老话说得好:有些朋友看着是来帮忙的,其实是来调整你节奏的。
1.6 稳定性技巧:softmax 是最容易出事的地方
模型变大以后,很多训练不稳定会集中暴露在 softmax 附近。softmax 先对 logit 做指数,再按总和归一化。logit 尺度一大,指数立刻把差异放大;分布一尖,梯度全挤到少数位置;attention score 或 vocab logit 一旦失控,就是 loss spike,严重时直接训练崩溃。
你看那种对比图:一条曲线 loss 低但抖得像帕金森,一条曲线 loss 高一点但稳如老狗。记住那句话:别训练出蓝色曲线那种模型------那不是训练,是蹦极。
输出端的 z-loss 可以这么理解:设 z = log(sum(exp(logits))),就是 softmax 归一化分母的 log 形式。z-loss 惩罚过大的 z,给输出 logit 的整体尺度加一个软约束。它不改变语言建模目标的主方向,但防止数值尺度无限膨胀------相当于给熊孩子规定:随便玩,但别把天花板拆了。
attention 端的 QK norm 处理另一个位置:在 query 和 key 做内积之前,先各自 norm 一下,让 QK^T 的分布别因为某个向量范数变大就走向极端。logit soft-capping 更直接:用 tanh 等函数把 logit 压进一个上限。
三者都不是为了让结构更复杂,而是为了避免 softmax 把尺度问题放大成训练事故。就像厨房里装了三个烟雾报警器------不是事多,是厨房真的容易着火。
1.7 GQA / MQA:瓶颈怎么从"算不动"变成"搬不动"
训练时,attention 可以并行处理整段序列,矩阵乘法足够大,GPU 容易保持高利用率。自回归生成不一样:第 t 个 token 生成时,只能使用前面已经生成的 KV cache;第 t+1 个 token 又要在新缓存上继续算。逐 token 推进,batch 和序列形状还老变。
标准 multi-head attention 会给每个 head 都保存一套 K/V。假设有 h 个 heads、每个 head_dim 是 k,那么每个 token 的 KV cache 大小就跟 h * k 相关,接近 d_model。上下文越长,缓存按 token 数线性增长;并发 batch 越大,缓存再按 batch 乘一遍。生成时每一步都要读过去的 K/V,瓶颈很容易从"算不动"变成"搬不动"------就像搬家,货没多少,光搬箱子就累死你。
MQA 的做法:让多个 query heads 共享同一组 K/V heads。query 照旧可以多头,但 K/V cache 的份数显著减少。GQA 是折中:不是所有 query 都共享一组 K/V,而是一组 query heads 共享一组 K/V heads。这个旋钮很实用------它让你在表达能力和推理成本之间连续调节,而不是只能在 full MHA 和 MQA 两端二选一。
所以 GQA/MQA 更像推理系统优化,而不是架构审美。训练 FLOPs 上它们未必惊艳,但在长上下文、较大 batch、服务端持续生成的场景里,KV cache 读写量直接决定吞吐和成本。一句话:训练时你在算,推理时你在搬。
1.8 滑动窗口和混合注意力:用图结构控制长上下文成本
full attention 可以看成一个完全图:每个 token 连到所有历史 token。表达能力强,但边数是 O(n²)。滑动窗口 attention 把图变成一条带宽有限的局部图:每个 token 只看附近窗口。每层成本低了,但远距离信息没法在一层内直接传播。
如果所有层都只做局部窗口,长程依赖就得靠多层逐步接力,路径越传越长。所以现代长上下文模型常用混合结构:大部分层用 sliding window 或 local attention,少数层保留 full attention,让远距离信息周期性地重新连通。核心不是某个固定间隔,而是在成本和表达能力之间选一张合适的 attention graph。选图如选路:不是每条路都通机场,但你要保证偶尔能上高速。
注意别跟 GQA/MQA 混了:GQA/MQA 主要减少每个 token 的 K/V 状态大小,滑动窗口主要减少每层 attention 连接数。前者压缩缓存宽度,后者压缩上下文连接密度。两者可以叠加,也会共同影响推理系统的 batch、cache 和 kernel 设计。
2. Embedding:从 token 到向量的第一步
tokenizer 之后,模型入口变成整数张量 (batch_size, sequence_length)。Embedding 层把每个 token ID 查成一个 d_model 维向量,输出 (batch_size, sequence_length, d_model)。
从这一步开始,我建议你强制自己用形状检查每个模块:最后一维是特征维,前面的维度都可以看作 batch-like 维度。这样写 Linear、RMSNorm、FFN、attention 时,代码自然支持 batch、sequence、head 等额外维度。就像强迫症患者整理衣柜:每个格子都有它的位置,拿衣服不迷路。
PyTorch 默认是 row-major 内存布局,常见线性层权重存成 (out_features, in_features)。数学上如果写列向量会得到 y = Wx,但代码里更常见的是对最后一维做 x @ W.T。这不是符号洁癖,是实现 bug 的高发区:权重按什么形状存、forward 里转不转置、输入的最后一维是不是等于 in_features,三者必须一致。不一致的下场就是:跑起来了,结果全错,你还以为是运气问题。
整个 decoder-only Transformer LM 可以写成下面这条路径:
token_ids: (B, T)
token_embedding: (B, T, D)
for each block:
residual stream stays (B, T, D)
final RMSNorm: (B, T, D)
LM head: (B, T, V)
logits: next-token scores for every position
Embedding 层实现:
python
class Embedding(torch.nn.Module):
def __init__(self, num_embeddings: int, embedding_dim: int, device=None, dtype=None):
super().__init__()
self.embeddings = torch.nn.Parameter(
torch.empty(num_embeddings, embedding_dim, device=device, dtype=dtype)
)
torch.nn.init.trunc_normal_(self.embeddings, mean=0.0, std=1 / math.sqrt(embedding_dim))
def forward(self, token_ids: torch.Tensor) -> torch.Tensor:
return self.embeddings[token_ids]
3. FFN 与 PreNorm:残差主干必须保持干净
3.1 pre-norm block:残差主干必须保持干净
现代 decoder-only LM 通常使用 pre-norm,而不是原始 Transformer 的 post-norm。计算顺序:
z = x + MultiHeadSelfAttention(RMSNorm(x))
y = z + FFN(RMSNorm(z))
这两行公式的重点是 residual stream。主干 x -> z -> y 上没有被 normalization 直接截断;RMSNorm 只放在进入子层之前。残差路径就是一条干净的信息高速路,attention 和 FFN 只往里追加更新量。训练深层 Transformer 时,这种结构通常比 post-norm 更稳定,因为梯度可以沿着残差路径更直接地往回传。路修得好不好,直接决定车堵不堵。
3.2 RMSNorm
LayerNorm 会减均值再除标准差,RMSNorm 只用均方根缩放:
RMS(a) = sqrt(mean(a_i^2) + eps)
RMSNorm(a_i) = a_i / RMS(a) * g_i
实现时先把输入 upcast 到 float32 再平方求和,避免低精度下 overflow 或精度损失,最后再 cast 回原 dtype。别偷懒直接在半精度上算------你会收获一个 loss 乱飞的模型,和一个怀疑人生的自己。
python
class RmsNorm(torch.nn.Module):
def __init__(self, d_model: int, eps=1e-5, device=None, dtype=None):
super().__init__()
self.eps = eps
self.d_model = d_model
self.g = torch.nn.Parameter(torch.ones(d_model, device=device, dtype=dtype))
def forward(self, x: torch.Tensor) -> torch.Tensor:
input_dtype = x.dtype
x = x.to(torch.float32)
variance = x.pow(2).mean(-1, keepdim=True)
x = x * torch.rsqrt(variance + self.eps)
return (self.g * x).to(input_dtype)
3.3 Linear 层
python
class Linear(torch.nn.Module):
def __init__(
self, in_features, out_features,
weights: Float[Tensor, "out in"] | None = None,
device=None, dtype=None
):
super().__init__()
if weights is None:
sigma = math.sqrt(2.0 / (in_features + out_features))
self.w = torch.nn.Parameter(
torch.empty(out_features, in_features, device=device, dtype=dtype)
)
torch.nn.init.trunc_normal_(self.w, mean=0.0, std=sigma, a=-3 * sigma, b=3 * sigma)
else:
self.w = torch.nn.Parameter(weights)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return einx.dot("... [in], out [in] -> ... out", x, self.w)
权重矩阵存成 out_features, in_features 是有讲究的:做矩阵乘法时,weight 作为右矩阵可以按行取数,对缓存更友好。调优这事,有时候就是"换个姿势存取数据",跟程序员换显示器支架一个道理。
3.4 SiLU:ReLU 的平滑升级版
python
class SiLu(torch.nn.Module):
def __init__(self):
super().__init__()
def forward(self, x: torch.Tensor) -> torch.Tensor:
return torch.sigmoid(x) * x
SiLU 又叫 Swish,2017 年提出的。它用微小的计算代价,换梯度稳定性、表达能力和现代架构的深度契合。说得通俗点:ReLU 在零点有个折角,像青春期叛逆的小孩,说不干就不干;SiLU 在零点平滑过渡,像被社会磨平了棱角的中年人,温和但依然有态度。LLaMA、Mistral 这些顶尖大模型集体选它,不是跟风,是它在资源允许时确实是更现代化的替代方案。
3.5 GLU 和 SwiGLU:门控不是换激活函数
GLU(Gated Linear Units)是一类带门控机制的激活函数,数学形式如下:
# 标准 GLU(以输入 x 为例)
a = W₁x + b₁ # 线性变换(内容路径)
b = W₂x + b₂ # 线性变换(门控路径)
output = a ⊙ σ(b) # ⊙ = 逐元素乘,σ = sigmoid(或其他激活函数)
本质:用门控信号 σ(b) 动态调节内容信号 a 的通过强度。打个比方,内容想进酒吧,门控就是门口保安------保安点头多少,内容就进去多少。
关键特性:
- 非线性增强:门控机制提供比单激活函数更强的表达能力
- 梯度友好:门控路径保留梯度流,缓解梯度消失
- 参数可控:通过调整门控激活函数衍生多种高效变体
常用变体(实际工程中更主流):
| 变体名称 | 公式 | 特点 |
|---|---|---|
| SwiGLU | SwiGLU(x) = a ⊗ SiLU(b) | 用 SiLU 替代 Sigmoid,梯度更优,GPT/LLaMA 均采用 |
| ReGLU | ReGLU(x) = a ⊗ ReLU(b) | 计算更快,适合轻量级模型 |
原始 Transformer 的 FFN 大致是 W2 ReLU(W1 x),中间维度常取 4 * d_model。现代 LLM 更常用 SwiGLU,并且通常去掉线性 bias。SwiGLU 可以写成:
SiLU(u) = u * sigmoid(u)
FFN(x) = W2( SiLU(W1 x) * W3 x )
这不是单纯把 ReLU 换成 SiLU。这里有两条投影分支:一条经过 SiLU 产生平滑的门控信号,另一条产生被门控的值;两者逐元素相乘,再投影回 d_model。可以理解成"让每个 hidden dimension 自己决定信息通过多少"------每个维度都是自己的流量控制员。
因为多了一组矩阵,为了让参数量和计算量可比,中间维度通常不是 4 * d_model,而是接近 8/3 * d_model,再向 64 的倍数取整以适配硬件。记住这个数,面试时能救命。
GLU 实现:
python
class Glu(torch.nn.Module):
def __init__(self, in_features: int, out_features: int, device=None, dtype=None):
super().__init__()
self.w1 = Linear(in_features, out_features, device=device, dtype=dtype)
self.w2 = Linear(in_features, out_features, device=device, dtype=dtype)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return torch.sigmoid(self.w1(x)) * self.w2(x)
SwiGLU FFN 实现:
python
class SwiGluFFN(torch.nn.Module):
def __init__(self, d_in: int, d_hidden: int, d_out: int, device=None, dtype=None) -> None:
super().__init__()
self.w1 = Linear(d_in, d_hidden, device=device, dtype=dtype)
self.w3 = Linear(d_in, d_hidden, device=device, dtype=dtype)
self.w2 = Linear(d_hidden, d_out, device=device, dtype=dtype)
self.silu = SiLu()
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.w2(self.silu(self.w1(x)) * self.w3(x))
4. RoPE:旋转位置编码,把相对距离写进内积
在 decoder-only Transformer 里,attention score 通常写成:
s_{m,n} = q_m^T k_n, q_m = W_q x_m, k_n = W_k x_n
这里的 m 是当前 query token 的位置,n 是被看的 key token 的位置。如果不加位置编码,两个 token 的相对顺序不会进入这个内积------模型只知道"这俩向量像不像",不知道"它们隔了多远"。就像你只知道两个同事脸长得像,不知道他们是不是同一个部门。
直接加绝对位置向量也行,但加法会把 token 内容、绝对位置和交叉项混在一起,很难保证打分以稳定方式依赖相对距离。
RoPE 的目标更明确:找到一个位置相关变换 f(x, m),让打分可以写成:
f_q(q, m)^T f_k(k, n) = g(q, k, n - m)
位置最终通过 n - m 进入 attention,而不是作为两个互不相关的绝对标签进入。这就是我要的全部。
4.1 把 head 维度两两看成二维平面
RoPE 的具体做法:设单个 attention head 的维度为 d_h,并且 d_h 是偶数。把向量按相邻维度两两分组:(x0, x1), (x2, x3), ..., (x_{d_h-2}, x_{d_h-1})。
第 i 个二维子空间使用一个固定频率:
ω_i = θ^(-2i/d_h), i = 0, ..., d_h/2 - 1
常见实现里 θ = 10000。位置 m 对应的旋转角度就是 m * ω_i,二维旋转矩阵为:
R_i(m) = [ cos(m ω_i) -sin(m ω_i) ]
[ sin(m ω_i) cos(m ω_i) ]
整个 head 上的旋转矩阵是 block diagonal 结构。实际工程不会显式构造这个大矩阵,只缓存每个 position、每个维度 pair 对应的 cos 和 sin,再用向量化操作完成旋转。数学给你理论,工程给你效率,两边都要兼顾。
4.2 相对位置来自旋转矩阵的群性质
RoPE 不直接改 token embedding,而是在每一层 attention 里旋转 query 和 key:
q̃_m = R_m q_m, k̃_n = R_n k_n
旋转矩阵有两个关键性质:
R_m^T = R_-m, R_a R_b = R_{a+b}
于是:
(R_m q_m)^T (R_n k_n) = q_m^T R_{n-m} k_n
这一步就是 RoPE 的数学核心。绝对位置 m 和 n 在推导中相消,只留下相对位移 n - m。所以 RoPE 不是"把位置信息塞进向量",而是把相对距离写进了 query/key 的内积结构。
妙不妙?妙。为什么妙?因为它靠群性质白嫖出来的------你要的只是旋转,旋转自己会带上"差值"这个礼物。
4.3 二维展开后能直接看到 sin 和 cos
只看第 i 个二维子空间,令 q_i = a, b,k_i = c, d,Δ = (n - m) ω_i。旋转后的二维内积可以展开为:
(R_i(m) q_i)^T (R_i(n) k_i) = (ac + bd) cos Δ + (bc - ad) sin Δ
这个式子说明两个细节。第一,原始内容相似度 ac + bd 仍然保留,只是被相对距离的余弦项调制;第二,bc - ad 是二维方向关系,被相对距离的正弦项调制。于是:同一对 token 内容,距离不同,打分不同;同一距离下,不同频率的二维子空间给出不同尺度的位置响应。内容管内容,距离管距离,各司其职。
4.4 为什么要用一组频率
如果所有二维子空间都用同一个频率,位置模式很快会周期性重复,表达能力也有限。RoPE 沿用了 sinusoidal position embedding 的多频率设计:小 i 对应较高频率,对短距离变化更敏感;大 i 对应较低频率,对长距离变化更平滑。
| 频率范围 | 数学效果 | 直观作用 |
|---|---|---|
| 高频 | ω_i 大,m ω_i 随位置变化快 | 更容易区分近距离 token 的顺序差异 |
| 低频 | ω_i 小,旋转角度变化慢 | 更适合给长距离依赖提供平滑的位置线索 |
多频率不是装饰,而是在不同二维子空间里提供不同"波长"的相对位置特征。模型后续通过注意力头和线性层学习如何组合这些尺度。就像调音台:高频旋钮管细节,低频旋钮管氛围,混音师决定怎么配。
4.5 具体实现:rotate_half 只是矩阵乘法的向量化写法
以相邻维度配对的实现为例,二维旋转可以写成:
x' = x cos α + rotate_half(x) sin α
rotate_half([x0, x1]) = [-x1, x0]
下面保留完整实现,和上面的数学一一对应:inv_freq 对应每个二维子空间的 ω_i,cos_cached / sin_cached 对应位置 m 的旋转角,x_rotated 对应 rotate_half。
python
class RoPE(torch.nn.Module):
def __init__(self, dim: int, max_seq_len: int = 2048, theta: float = 10000, device=None, dtype=None):
super().__init__()
self.dim = dim
self.max_seq_len = max_seq_len
# inv_freq: (dim//2,)
inv_freq = 1.0 / (
theta ** (torch.arange(0, dim, 2, device=device, dtype=torch.float32) / dim)
)
# t: (seq_len,)
t = torch.arange(max_seq_len, device=device, dtype=torch.float32)
# freqs: (seq_len, dim//2)
freqs = torch.einsum("i,j->ij", t, inv_freq) # outer product
emb = freqs.repeat_interleave(2, dim=-1) # (seq_len, dim)
self.register_buffer("cos_cached", emb.cos().to(dtype)) # (seq_len, dim)
self.register_buffer("sin_cached", emb.sin().to(dtype)) # (seq_len, dim)
def forward(self, x: Float[Tensor, "... seq d_k"], token_positions: Float[Tensor, "... seq"]) -> torch.Tensor:
# token_positions: (..., seq_len) 任意前缀维度
# x: (..., seq_len, dim)
cos = self.cos_cached[token_positions] # (..., seq_len, dim)
sin = self.sin_cached[token_positions] # (..., seq_len, dim)
x_reshaped = x.view(*x.shape[:-1], -1, 2) # (..., seq_len, dim//2, 2)
x_rotated = torch.stack((-x_reshaped[..., 1], x_reshaped[..., 0]), dim=-1) # rotate: (a,b) -> (-b,a)
x_rotated = x_rotated.view(*x.shape) # (..., seq_len, dim)
if x.ndim == 4:
cos = cos.unsqueeze(1)
sin = sin.unsqueeze(1)
x_rot = x * cos + x_rotated * sin
return x_rot
工程实现要注意两件事。第一,cos 和 sin 的缓存布局必须和维度配对方式一致;有些实现是相邻维度配对,有些实现把前半维当实部、后半维当虚部。对不上,就等着 cos 和 sin 各跳各的舞。第二,token_positions 用来索引缓存,实际类型应是整数张量;对带 KV cache 的增量推理,它通常不是简单的 0...seq_len-1,而是当前上下文里的真实位置。用错了,模型会以为自己在时间旅行。
4.6 为什么 RoPE 只作用在 query 和 key 上
attention 的输出是:
Attention(Q, K, V) = softmax(QK^T / √d_h) V
位置关系主要决定"当前 token 应该看哪些历史 token",也就是 softmax 前的打分矩阵 QK^T。因此 RoPE 作用在 query/key 上,让权重计算携带相对位置。value 承载的是被聚合的内容本身,通常不需要被同样旋转------否则会把位置信号进一步混进被读取的信息,内容聚合反而更难解释。
翻译一下:Q 和 K 负责"眼光",V 负责"内容"。你只需要让眼光知道距离,没必要让内容也跟着转圈。
5. Attention 与 Transformer:把积木搭起来
5.1 ScaledDotProductAttention
python
class ScaledDotProductAttention(torch.nn.Module):
def __init__(self):
super().__init__()
self.softmax = Softmax()
def forward(
self,
q: Float[Tensor, "... s d"],
k: Float[Tensor, "... s d"],
v: Float[Tensor, "... s d"],
mask: torch.Tensor | None = None, # true 表示该位置要被覆盖,不参与 softmax 计算
) -> torch.Tensor:
d_model = q.shape[-1]
# 计算 attention scores
att = einx.dot("... s_q [d], ... s_k [d] -> ... s_q s_k", q, k)
att_scale = att / math.sqrt(d_model)
if mask is not None:
if mask.ndim < att_scale.ndim:
mask = mask.reshape((1,) * (att_scale.ndim - mask.ndim) + mask.shape)
att_scale = att_scale.masked_fill(mask, -1e9)
att_score = self.softmax(att_scale)
return einx.dot("... s_q [s], ... [s] d -> ... s_q d", att_score, v)
注意 mask 的处理:mask 为 True 表示该位置要被覆盖,不参与 softmax 计算。实现里用 masked_fill 填成 -1e9 而不是 -inf-------inf 在某些精度下会搞出 NaN,-1e9 是工程界的"差不多得了"。
5.2 MultiHeadAttention
实际实现通常用一个投影矩阵完成多个 head 的投影,用来提升计算速度:
python
class MultiHeadAttention(torch.nn.Module):
def __init__(self, d_model: int, num_head: int, max_seq_len=2048, device=None, dtype=None):
super().__init__()
self.d_model = d_model
self.num_head = num_head
self.out_linear = Linear(d_model, d_model, device=device, dtype=dtype)
self.project = Linear(in_features=d_model, out_features=3 * d_model, device=device, dtype=dtype)
self.dot_product_att = ScaledDotProductAttention()
# 缓存 causal mask
causal_mask = torch.triu(
torch.ones(max_seq_len, max_seq_len, dtype=torch.bool, device=device),
diagonal=1
)
self.register_buffer("causal_mask", causal_mask)
def forward(self, x: Float[Tensor, "b s d"]) -> torch.Tensor:
seq_len = x.shape[1]
mask = self.causal_mask[:seq_len, :seq_len]
qkv = self.project(x)
q, k, v = einx.rearrange("b s (n h d) -> n b h s d", qkv, n=3, h=self.num_head)
output = self.dot_product_att(q, k, v, mask)
output = einx.rearrange("b h s d -> b s (h d)", output)
return self.out_linear(output)
causal mask 用 torch.triu 生成上三角矩阵缓存起来,每轮直接切片。diagonal=1 表示屏蔽未来位置------自己和自己之前都能看,未来不行。这规矩比公司考勤还严格。
5.3 MultiHeadAttentionWithRoPE
multi-head attention 只是把 d_model 拆成 num_heads * d_head。head 维度应该像 batch 维度一样独立处理:每个 head 都用自己的 Q/K/V 切片做 attention,但 RoPE 的位置旋转规则对所有 head 一样。
一个稳妥的实现路径:先做 Q、K、V 投影(一个大矩阵投影再拆),把形状从 (B, T, D) rearrange 成 (B, H, T, d_head),对 Q/K 应用 RoPE,所有 head 并行算 masked attention,最后 concat 回 (B, T, D) 并过 output projection。
python
class MultiHeadAttentionWithRoPE(MultiHeadAttention):
def __init__(self, d_model: int, num_head: int, theta: float = 10000, max_seq_len=2048, device=None, dtype=None):
super().__init__(d_model=d_model, num_head=num_head, max_seq_len=max_seq_len, device=device, dtype=dtype)
self.rope = RoPE(
d_model // num_head, max_seq_len=max_seq_len,
theta=theta, device=device, dtype=dtype
)
def forward(self, x: torch.Tensor, token_positions: torch.Tensor | None = None) -> torch.Tensor:
seq_len = x.shape[1]
batch_size = x.shape[0]
if token_positions is None:
token_positions = torch.arange(seq_len, device=x.device).unsqueeze(0).expand(batch_size, -1)
mask = self.causal_mask[:seq_len, :seq_len]
qkv = self.project(x)
q, k, v = einx.rearrange("b s (n h d) -> n b h s d", qkv, n=3, h=self.num_head)
# 对 q 和 k 应用 RoPE
q = self.rope(q, token_positions)
k = self.rope(k, token_positions)
output = self.dot_product_att(q, k, v, mask)
output = einx.rearrange("b h s d -> b s (h d)", output)
return self.out_linear(output)
5.4 TransformerBlock
把上面的模块堆叠起来,就是 TransformerBlock。跟原始 Transformer 架构不同的是,现代 LLM 通常用 pre-norm、用更简单的 RmsNorm、用 SwiGLU 作为 FFN:
python
class TransformerBlock(torch.nn.Module):
def __init__(
self,
d_model: int,
num_heads: int,
d_ff: int,
max_seq_len: int = 2048,
theta: float = 10000,
device=None,
dtype=None,
) -> None:
super().__init__()
self.rms_norm1 = RmsNorm(d_model, device=device, dtype=dtype)
self.rms_norm2 = RmsNorm(d_model, device=device, dtype=dtype)
self.mult_head_atten = MultiHeadAttentionWithRoPE(
d_model, num_heads, theta, max_seq_len=max_seq_len, device=device, dtype=dtype
)
self.ffe = FFN(d_model, d_ff, d_model, device=device, dtype=dtype)
def forward(self, x: torch.Tensor, token_positions: torch.Tensor | None = None) -> torch.Tensor:
x_norm = self.rms_norm1(x)
x_atten = self.mult_head_atten(x_norm, token_positions)
x = x + x_atten
x_norm = self.rms_norm2(x)
x_ffe = self.ffe(x_norm)
return x + x_ffe
5.5 Transformer
最终的 Transformer 就是由多个 block 叠加起来:
python
class Transformer(torch.nn.Module):
def __init__(
self,
d_model: int,
num_heads: int,
d_ff: int,
vocab_size: int,
num_layers: int,
max_seq_len=2048,
rope_theta: float = 10000,
device=None,
dtype=None,
):
super().__init__()
self.embedding = Embedding(num_embeddings=vocab_size, embedding_dim=d_model, device=device, dtype=dtype)
self.blocks = torch.nn.ModuleList(
[
TransformerBlock(
d_model=d_model,
num_heads=num_heads,
d_ff=d_ff,
max_seq_len=max_seq_len,
theta=rope_theta,
device=device,
dtype=dtype,
)
for _ in range(num_layers)
]
)
self.norm = RmsNorm(d_model=d_model, device=device, dtype=dtype)
self.out_linear = Linear(d_model, vocab_size, device=device, dtype=dtype)
self.max_seq_len = max_seq_len
def forward(self, token_ids: torch.Tensor, token_positions: torch.Tensor | None = None) -> torch.Tensor:
x = self.embedding(token_ids)
if token_positions is None:
batch_size, seq_len = token_ids.shape
token_positions = torch.arange(seq_len).unsqueeze(0).expand(batch_size, -1)
for block in self.blocks:
x = block(x, token_positions)
x_norm = self.norm(x)
logits = self.out_linear(x_norm)
return logits
注意 forward 里 token_positions 的默认处理:没显式传入时,用 arange 生成 0...seq_len-1 并广播到 batch。这是训练时的常规路径;推理带 KV cache 时,把真实位置传进来就行。
6. 训练优化:让模型真正学会说话
6.1 Loss:预测下一个 token
语言模型的训练目标很朴素:给定前缀 x_1 ... x_i,预测下一个 token x_{i+1}。模型一次前向会对每个位置都输出一个 logits 向量,所以一个长度为 T 的输入可以同时产生 T 个 next-token 预测。训练 batch 里,输入 x 和目标 y 的关系就是右移一位------输入的第 i 个位置,对应目标的第 i+1 个位置。
对单个位置,cross-entropy 可以写成:
loss_i = -log softmax(logits_i)[target_i]
= log(sum_j exp(logits_i[j] - max_i)) + max_i - logits_i[target_i]
要避免先显式算 softmax 再取 log------那会把两个容易溢出或下溢的操作连在一起。用上面的 logsumexp 形式,最大值被减掉,指数项更稳定;目标 token 的 logit 单独取出,最后对所有 batch-like 位置求平均。
perplexity = exp(mean_loss),可以理解成模型平均每一步还在多少个候选 token 之间困惑。perplexity 30,相当于模型每走一步都在 30 个词之间犹豫不决------像选择困难症患者点外卖,翻半小时菜单,最后点了跟上次一样的。
python
class CrossEntropyLoss(torch.nn.Module):
def __init__(self) -> None:
super().__init__()
def forward(self, logits: torch.Tensor, targets: torch.Tensor) -> torch.Tensor:
# 处理任意形状的 logits
logits = einx.rearrange("... c -> (...) c", logits)
targets = einx.rearrange("... -> (...)", targets)
log_probs = torch.nn.functional.log_softmax(logits, dim=-1)
correct_log_probs = log_probs[torch.arange(len(log_probs)), targets]
nll = -correct_log_probs
mean_loss = torch.mean(nll)
return mean_loss
6.2 Optimizer:AdamW,以及它吃掉的内存
SGD 的更新很好理解:参数沿着负梯度方向走一步。但现代 Transformer 通常用 AdamW,因为它会为每个参数维护一阶和二阶 moment,分别估计梯度均值和梯度平方均值。每个参数的步长会根据历史梯度尺度自适应调整,训练通常更稳。
AdamW 里的 W 很重要:weight decay 和梯度更新是解耦的。参数先按 theta = theta - lr * weight_decay * theta 往 0 拉,再用 moment-adjusted gradient 更新。这样 weight decay 更接近直接的参数正则,而不是混进 Adam 的梯度统计里。
一句话:原始 Adam 把权重衰减当"奖金"发,AdamW 把它当"税"单独收。奖金看心情,税是规矩------后者靠谱多了。
完整更新流程:
- 计算梯度:g_t = ∇θ L(θ{t-1})
- 一阶矩(动量)更新:m_t = β1 m_{t-1} + (1 - β1) g_t
- 二阶矩(梯度平方)更新:v_t = β2 v_{t-1} + (1 - β2) g_t²
- 偏置校正:m̂_t = m_t / (1 - β1^t),v̂_t = v_t / (1 - β2^t)(初始累计值较小需要放大,随着 step 增加放大倍数逐步减小)
- AdamW 核心更新(解耦权重衰减):θ_t = θ_{t-1} - η · m̂_t / (√v̂_t + ε) - η · λ · θ_{t-1}
和原始 Adam 的关键区别:原始 Adam 把权重衰减项包含在梯度内,受自适应学习率缩放,正则效果不稳定;AdamW 让权重衰减独立于梯度,直接作用于参数,与 SGD 的 L2 正则行为一致。
python
class AdamW(torch.optim.Optimizer):
def __init__(
self,
params: ParamsT,
lr=1e-3,
betas: tuple[float, float] = (0.9, 0.999),
weight_decay=1e-3,
eps=1e-8,
):
if lr < 0:
raise ValueError(f"invalid learning rate: {lr}")
beta1, beta2 = betas
defaults = {
"lr": lr,
"beta1": beta1,
"beta2": beta2,
"weight_decay": weight_decay,
"eps": eps,
}
super().__init__(params, defaults)
@overload
def step(self, closure: None = None) -> None: ...
@overload
def step(self, closure: Callable[[], float]) -> float: ...
def step(self, closure: Callable[[], float] | None = None) -> float | None:
loss = None if closure is None else closure()
for group in self.param_groups:
lr = group["lr"]
beta1 = group["beta1"]
beta2 = group["beta2"]
weight_decay = group["weight_decay"]
eps = group["eps"]
for p in group["params"]:
if p.grad is None:
continue
state = self.state[p]
# 初始化状态
if len(state) == 0:
state["t"] = 0
state["m"] = torch.zeros_like(p.data)
state["sm"] = torch.zeros_like(p.data)
m, sm = state["m"], state["sm"]
t = state["t"] + 1
grad = p.grad.data
# 更新有偏一阶矩估计
m.mul_(beta1).add_(grad, alpha=1.0 - beta1)
# 更新有偏二阶原始矩估计
sm.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)
# 偏置校正
m_hat = m / (1.0 - beta1**t)
sm_hat = sm / (1.0 - beta2**t)
# 更新参数
p.data.addcdiv_(m_hat, torch.sqrt(sm_hat) + eps, value=-lr)
# 解耦的权重衰减
if weight_decay != 0:
p.data.add_(p.data, alpha=-lr * weight_decay)
state["t"] = t
return loss
这个优化器的代价是内存。假设参数用 float32,单份参数需要 4 字节;梯度还要一份;AdamW 的 m 和 v 又各要一份。只看参数相关状态,就已经是参数量的约 4 倍内存,再加上 activation、logits、临时张量和 checkpoint。
理解这个成本,才能解释那个经典场景:同一个模型,推理时能装下,训练时爆显存。老板看着你,你看着 CUDA OOM。以后有人问"这模型不是挺小的吗",你就把这份账拍他脸上。
6.3 Cosine LR Scheduler
学习率决定每一步走多远,它通常比很多结构细节更先影响训练成败。一个实用策略是 warmup 加 cosine decay:前几个 step 从 0 线性升到最大学习率,让 moment state 和模型激活先进入稳定范围;之后用 cosine 逐步衰减到最小学习率,让训练后期更细地收敛。
前期像军训,节奏要猛,先把队伍带起来;后期像退休生活,动作要稳,别整幺蛾子。这里的调度器只是一个纯函数:输入当前 step 和几个超参数,输出本 step 的学习率。
python
def cos_lr_scheduler(it: int, warmup_iters: int, cos_cycle_iters: int, lr_min: float, lr_max: float) -> float:
if it <= warmup_iters:
return lr_max * it / warmup_iters
elif warmup_iters < it < cos_cycle_iters:
return lr_min + 0.5 * (lr_max - lr_min) * (
1 + math.cos(math.pi * (it - warmup_iters) / (cos_cycle_iters - warmup_iters))
)
else:
return lr_min
6.4 Gradient Clip
梯度裁剪解决的是另一类问题:偶发 batch 可能产生非常大的梯度,把参数一步推到坏区域。就像马路上突然窜出个醉汉司机,方向盘得赶紧抢回来。
L2 范数梯度裁剪(也叫 max-norm 裁剪):计算所有参数梯度的全局 L2 范数,若范数超过设定的 max_norm 阈值,就按比例缩小所有梯度,确保整体范数不超过阈值,避免梯度爆炸。
python
def gradient_clip(params: Iterable[torch.nn.Parameter], max_norm: float, delta=1e-6):
with torch.no_grad():
grads = [p.grad for p in params if p.grad is not None]
total_norm = torch.linalg.norm(
torch.stack([torch.linalg.norm(g.detach()) for g in grads])
)
if total_norm > max_norm:
clip_coef = max_norm / (total_norm + delta)
for g in grads:
g.detach().mul_(clip_coef)
7. 资源核算:钱都花在矩阵乘法上
实现完模型,除了能不能跑,还要能算它为什么贵。Transformer 里的主要 FLOPs 来自矩阵乘法:A ∈ R^{m×n} 乘 B ∈ R^{n×p} 大约需要 2mnp FLOPs。根据这个规则,可以把模型前向拆成一张账表:
| 组件 | 主要矩阵乘法 | 增长直觉 |
|---|---|---|
| Q/K/V projections | (BT, D) × (D, D) 三次 | 随 token 数和 D² 线性增长 |
| Attention scores | QK^T,每个 head 是 (T, d_head) × (d_head, T) | 随 T² 增长,长上下文时变重 |
| Attention values | softmax(QK^T) V | 同样随 T² 增长 |
| Output projection | (BT, D) × (D, D) | 随 token 数和 D² 线性增长 |
| SwiGLU FFN | W1、W3 上投影和 W2 下投影 | 通常是 block 内最大 dense compute 来源之一 |
| LM head | (BT, D) × (D, V) | 词表很大时不可忽略 |
这张账表会直接影响后续实验判断:增大 d_model 会放大大多数 dense projection;增大 context_length 会让 attention score/value 的 T² 项快速抬头;增大词表 V 会抬高 LM head 和 cross-entropy 的成本。
也正因为这样,训练一个看似"小"的模型,如果数据读取、验证 loss 或 checkpoint 写入不当,真实 wall-clock 也可能被非模型部分拖住。显存都够,时间不够------说的就是这种情况。
最后把现代默认值这套组合拳收个尾:pre-norm 保残差主干、RMSNorm 省搬运、无 bias 省钱、SwiGLU 上表达力、RoPE 管位置、GQA 管缓存。它们不是时尚单品,而是每一笔都算过账的预算分配。
从零手写 Transformer,写完你会发现:难的不是代码,是理解每一行代码为什么长这样。把"为什么"想明白了,代码自己就长出来了------就像你终于知道盐为什么"适量"之后,炖肉就再也不会咸了。
P.S. 推荐一个大神的教程给想要了解或者学习人工智能知识的读者,这个教程里内容讲解通俗易懂且风趣幽默,对我帮助很大。我想与大家分享这个宝藏教程,请点击下方链接查看,传送门https://blog.csdn.net/qq_74013365