Transformer架构优化

文章目录

    • 前言
    • [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 把它当"税"单独收。奖金看心情,税是规矩------后者靠谱多了。

完整更新流程:

  1. 计算梯度:g_t = ∇θ L(θ{t-1})
  2. 一阶矩(动量)更新:m_t = β1 m_{t-1} + (1 - β1) g_t
  3. 二阶矩(梯度平方)更新:v_t = β2 v_{t-1} + (1 - β2) g_t²
  4. 偏置校正:m̂_t = m_t / (1 - β1^t),v̂_t = v_t / (1 - β2^t)(初始累计值较小需要放大,随着 step 增加放大倍数逐步减小)
  5. 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

相关推荐
西安圣木通2 小时前
智能体时代来临:重构企业生产力,开启商业效率新范式
大数据·人工智能·重构
阿里云基础软件2 小时前
一句话看透 JVM,SysOM 诊断 Skill 新增 Java 应用诊断能力
java·开发语言·jvm·人工智能·操作系统·sysom 诊断 skill
AI情绪识别开源2 小时前
检信ALLEMOTION 认知矫正能力语音识别模型拼音准确性测试报告
人工智能·语音识别
OreDev2 小时前
模型驾驭工程(Harness Engineering)与 DeepSeek Harness 的定位
人工智能
EEET智选2 小时前
科技晨间速报 | 苹果折叠屏deepseek新模型同日发布
人工智能·科技·ai·aigc·手机
智购科技无人售货机工厂2 小时前
2026 AI视觉货柜渗透率突破48%:从技术选型到规模化部署的工程实践~YH
人工智能
CIO_Alliance3 小时前
AI微调系列(1)| 全量微调、LoRA、QLoRA 三种方案怎么选
前端·人工智能·神经网络·机器学习·embedding·企业ai转型
真上帝的左手3 小时前
27. 数据产品- BI - AI 应用实战终篇-从 RAG 架构到企业级平台落地
大数据·人工智能·架构