Day 6 · 创新深读(上):AttnRes + Stable LatentMoE 三件套
信息怎么「跨层串门」?极端稀疏怎么「稳住」?
K3 靠两个架构创新「变大还变聪明」:AttnRes(注意力残差)让每一层都能翻前面所有层的笔记(信息沿深度流动);Stable LatentMoE 三件套(RMSNorm + SiTU-GLU + QB 直方图)在 896 专家这种极端稀疏下稳住训练。今天是「硬核但可以懂」的一天。
学完你能回答
标准残差连接有什么「瓶颈」?AttnRes 怎么治?
学完你能回答
896 个专家的两个「病」分别由哪两件套来治?
0起点:这些你早就会了
今天讲的还是「信息怎么传递」——但这次是深度方向(层与层之间),不是宽度方向(字与字之间)。先想几个场景:
| 你已经会的 | AI 世界里对应的东西 |
| 上课记笔记:新知识点写在旧笔记后面,只看上一页 | 标准残差连接:每一层只把上一层的结果加进来 |
| 老师翻出第一课的知识点来回答今天的提问 | AttnRes(注意力残差):每一层能「选择性」查看前面所有层的笔记 |
| 老师先问一句「这个问题该翻哪本笔记?」 | 伪查询 w:每层学出来的一个小问题,用它决定「该看重哪些旧笔记」 |
| 水烧到 100°C 就沸腾,不会无限升温 | SiTU-GLU(有界激活):数字涨到上限就被「封顶」,防止爆炸 |
| 全国人口普查不用逐个点名,按省汇总就能估计分布 | QB 直方图估计:不精确排序几百万个数字,用「分桶汇总」近似算分位数 |
核心就一句:AttnRes 管「跨层信息流动」,三件套管「极端稀疏下的稳定」——都是让 2.8T 的大模型「又强又不炸」的关键。
1标准残差:每层只「看见」前一层
现在绝大多数 AI 模型都用残差连接(residual connection):信息一层层往下传,每层把自己学到的「新理解」加到上一层传来的「旧理解」上。就像一个接力跑,每一棒只能从上一棒手里接棒。
听起来很合理,但它有个隐藏的瓶颈:信息在一层层传递中会被「压缩」——传到第 50 层时,第 1 层学到的东西可能已经被挤得面目全非。这就像传话游戏:第 1 个人说的悄悄话,传到第 50 个人耳朵里早就不是原话了。
图 1 · 标准残差 vs AttnRes:左边——标准残差像接力跑,第 3 层只能从第 2 层手里接棒,第 1 层的信息传到这里已经「打折」;右边——AttnRes 让第 3 层可以直接看第 1、2 层的原始结果(虚线箭头),按需取用,信息跨层流动不再走样。这就是报告说的:把「注意力」从序列维度(字与字)搬到深度维度(层与层)。
打个比方
标准残差 = 传话游戏(每传一次都可能失真);AttnRes = 会议室里每个人都把发言写在白板上,新来的人想查谁的话都可以直接看白板,不用听转述。白板就是「所有层的笔记」。
进阶小注 ·标准残差 h_l = h_{l−1} + f_l(h_{l−1}) 把所有先前信息沿深度压缩进单一状态 h_l——报告原话:「这一瓶颈令人联想到 RNN 在时间维度上的处境」。Transformer 用注意力取代了序列上的递归;AttnRes 把同样的方法论应用于深度:每一层以数据依赖的权重选择性访问所有先前层,而非均匀累积。
2AttnRes 怎么选「看谁」?靠一个「伪查询」
「能看所有层」不代表「每层都看」——那样信息会糊成一锅粥。AttnRes 的高明之处:每一层自己学一个「伪查询」w,用它去对前面所有层的笔记「打分」,再按分数加权取用——就像老师提问前先想「这个问题该翻哪本笔记」,翻到重点的那本多读几页。
图 2 · 伪查询机制:左边三份「旧笔记」:起点嵌入 h₀、第 1 层输出 f₁、第 2 层输出 f₂。第 3 层学了一个伪查询 w₃(一个小向量),用它给三份笔记分别打分,再归一化成权重 α₀、α₁、α₂(右上三个小标签),最后加权求和得到自己的输出 h₃。权重是学出来的、随任务变化的——「现在该多参考语法层,还是多参考语义层?」模型自己说了算。
别被 α 吓到:α 就是「每个旧笔记占多大比重」。三个 α 加起来等于 1(像三份材料各占 20%、50%、30%)。重点是这套打分机制靠学习得来,不是人手工设计的——K3 在训练中自己发现「翻哪本笔记最有用」。
进阶小注 ·AttnRes 形式化(式 8–9):层 l 的伪查询 q_l = w_l(可学习);键值对 k_i = v_i(i=0 为嵌入,i≥1 为 f_i(h_i));注意力权重 α_{i→l} = exp(q_lᵀ RMSNorm(k_i)) / Σ_j exp(q_lᵀ RMSNorm(k_j))(softmax 核);输出 h_l = Σ α·v。RMSNorm 防止输出幅度大的层主导权重。全量形式计算量 O(L²d),因深度 L<100 可承受。
3太贵怎么办?块级 AttnRes:先「同层归约」再「跨块查表」
「每层都看前面所有层」听起来很美,但存不下:每层的输出都得留着给后面所有层看,模型越深,留的「笔记」越多,内存吃不消。K3 的折中方案叫块级 AttnRes:把层分成几组(块),块内各层先合并成一份「块笔记」,跨块时只对「块笔记」做注意力——笔记数量一下少了很多。
图 3 · 块级 AttnRes:① 把层分成块(K3 每块 12 层,共 8 块,加上起点嵌入 = 9 个块);② 块内各层的输出先相加,归约成一份「块笔记」bₙ(第二行绿色小块);③ 跨块时只对这几份块笔记做注意力——需要保留的笔记从「每层一份」变成「每块一份」,内存大降,还几乎不损失效果。
进阶小注 ·块级 AttnRes(式 10):块 n 内各层输出求和归约为单一表示 b_n = Σ f_j(h_j)(含部分和 b_n^i),令 b₀ = 嵌入。跨块仅对 N 个块级表示施加全量注意力,内存/通信从 O(Ld) 降至 O(Nd);块结构同时为推理时的状态设界,可经 online softmax 与顺序块内部分和合并。经验上 N≈8 即恢复大部分收益;K3 采用每块 12 层 × 8 块(末块不完整)+ 嵌入层 = 9 块。
4承接 Day 4:896 个专家会得什么「病」?
Day 4 讲到 K3 把专家扩到 896 个、每 token 激活 16 个(稀疏度 56)。这个「极端稀疏」会放大两种失效模式——用大白话说,就是两个会「炸」的病:
- 病一 · 激活爆炸(数字炸掉):路由通路是「降投影 → 专家 → 升投影」的链条,像连了 4 个放大镜;896 个专家 + 2.8T 规模的叠加,让放大镜把数字越放越大——大到超出电脑能表示的范围,训练就崩了。
- 病二 · 负载失衡(专家打架):近 1000 个专家做负载均衡,超出了老方法(固定步长的偏置更新)能管好的区间——有的专家忙死、有的闲死(Day 4 的 QB 就是来治这个的)。
K3 的「处方」是三件套:RMSNorm(治数字尺度漂移)+ SiTU-GLU(治激活爆炸)+ QB 直方图(治负载失衡)。前三件今天逐个拆开讲。
一句话记住:前两件套(RMSNorm、SiTU-GLU)负责「别让数字爆掉」,第三件套(QB 直方图)负责「别让专家闲死」。三者合称 Stable LatentMoE——「Stable(稳定)」两个字,就是它们的目标。
5第一件套:在「聚合 → 升维」之间插一个 RMSNorm
先看路由通路的完整结构。一个 token 被选中后,要走:降投影(W↓)→ 专家计算(在潜空间里)→ 聚合 → 升投影(W↑)。老版本的 LatentMoE 在「聚合完」直接接「升投影」——问题在于:聚合结果的大小随选中的专家和权重变来变去,尺度不稳定。
K3 的改法极其简单:在聚合和升投影之间,插一个 RMSNorm(归一化)——把数字「压回标准尺度」再送去升维。就像包饺子前先给面团称重整形,保证每个饺子大小一致。
图 4 · Stable LatentMoE 的结构:输入 x 兵分两路——上方绿色:2 个共享专家全宽度直接处理(全科医生);下方蓝色:路由通路先降投影到低维潜空间,896 个专家里选 16 个干活,聚合后插一个 RMSNorm(深蓝色小块,K3 的独家改动)再升投影回全宽度。两路结果相加 = 输出 y。RMSNorm 让路由分支的「数字尺度」不再忽大忽小,与共享分支合并时也稳。
进阶小注 ·式(11):y = Σ E_j^shared(x) + W↑ · RMSNorm( Σ pᵢ E_i^routed(W↓x) )。原始 LatentMoE 直接对聚合表示 u 施加 W↑,而 u 的尺度随所选专家与路由权重变化;插入 RMSNorm 降低路由分支对尺度变化的敏感性。报告特别注明:该 RMSNorm 除稳定训练外,还持续改善验证损失与下游基准——既稳又好。
6第二件套:SiTU-GLU——给激活值装一个「限高杆」
病一是「数字爆炸」。为什么会爆?因为专家里的主力激活函数 SwiGLU 有个特性:输入越大,输出越大,没有上限。在 896 专家 + 2.8T 规模的链条里,几个大数一相乘,直接超出电脑能存的范围。
K3 的解法:换用自家发明的 SiTU-GLU——把 SwiGLU 里「没有上限」的那部分,换成「涨到上限就封顶」的版本。效果:小数字时跟 SwiGLU 几乎一样(保持聪明),大数字时被牢牢摁在 100 以内(不爆炸)。
图 5 · SiTU-GLU vs SwiGLU(示意曲线):横轴输入大小、纵轴输出大小。灰色虚线(SwiGLU)没有上限,输入越大输出越大——在 896 专家的链条里就是隐患;红色实线(SiTU-GLU)在输入不大时和灰色几乎一样(不损失能力),但涨到 β₁×β₂ = 4×25 = 100 就封顶,无论如何都超不过这个「限高杆」。小聪明保留,大爆炸杜绝。
两个参数记一下:β₁ = 4(门分支的封顶)、β₂ = 25(升分支的封顶),所以输出上限是 4×25 = 100。为什么不用「硬截断」(大于 100 直接砍掉)?因为硬截断在边界处「没坡度」,模型学不动;SiTU-GLU 用的是平滑封顶——接近 100 时渐渐变平,永远保留一点点坡度,训练照常进行。
打个比方
SwiGLU 像没装限速的电车,越开越快最后失控;SiTU-GLU 像装了「软限速器」的跑车——平时随便飙(小输入),接近 100 km/h 时油门自动变软(平滑封顶),既不会失控,也不像硬限速那样一脚刹死。
进阶小注 ·SiTU-GLU(式 12):[β₁·tanh(W_g x/β₁) ⊙ σ(W_g x)] ⊙ [β₂·tanh(W_u x/β₂)],β₁=4、β₂=25。缩放 tanh 在原点附近一阶近似线性(局部响应与 SwiGLU 吻合)、大幅处有界;输出逐坐标满足 |f(x)| ≤ β₁β₂ = 100(附录 B 严格证明)。相比对门控预激活硬截断,平滑封顶在远离饱和边界处保留非零梯度,训练行为更好。
7第三件套:QB 直方图——几百万个数字怎么算「分位数」?
Day 4 讲过 QB(分位数平衡)给每个专家发「偏置分」。但有个工程难题被一笔带过:偏置分要用「路由器分数的分位数」来算,而分数有几百万个,分散在几百张卡上——全部收集起来排序?不现实,太慢。
K3 的聪明解法:不精确排序,改用「直方图」。把分数范围切成几百个「桶」,每个桶记「落进来多少个分数」——几百万个分数,瞬间压缩成几百个数字。而且桶是可以相加的:每张卡算自己的桶,一次汇总(all-reduce)就把所有卡的桶加起来,等于拿到了全局分布,然后从桶里读出分位数。
图 6 · QB 直方图估计:① 每个专家有「路由器分数离入选线的差距」(边际),一共几百万个、散在几百张卡上;② 不精确排序——把分数范围切成几百个桶,本地各数各的(几百万 → 几百);③ 一次汇总把所有卡的桶加起来,得到全局分布,从中读出分位数 → 算出专家偏置。通信量只有「每个专家几百个桶」,这就是 896 个专家还能做负载均衡的原因。
进阶小注 ·QB 更新(式 14):b_j ← −quantile_{1−k/n}(s_{:,j} − α),其中 α 是 Top-(k+1) 截断值,随后减去公共均值(不改变 Top-k 选择)。大规模下精确收集全局边际不可行,改用直方图估计:一次 all-reduce 对分桶计数求和(计数可加),从汇总计数恢复分位数,误差以分桶宽度为界(附录 D)。因因果性,偏置更新仅下一步生效——一批绝不会用自己推的偏置路由。
8映射到 K3:原文里怎么说这些创新?
今天的四个关键词(AttnRes、RMSNorm、SiTU-GLU、QB 直方图)都能在报告里找到原话:
K3 原文(§2.2 Attention Residuals)
「标准残差连接将所有先前信息沿深度压缩进单一状态 h_l——这一瓶颈令人联想到 RNN 在时间维度上的处境……每一层选择性地从所有先前层检索表示,而非均匀地累积它们。」
逐词翻译:「瓶颈」= 每层只见前一层、信息被压缩;「选择性地检索」= 用伪查询 w 决定看谁、看多重(图 2 的 α)。AttnRes 就是把「注意力」从字与字之间,搬到了层与层之间。
K3 原文(§2.3.1 归一化 LatentMoE)
「Kimi K3 改为在专家聚合与升投影之间插入 RMSNorm……该归一化降低了路由分支对尺度变化的敏感性,使其在与全宽度共享分支合并之前保持稳定。除稳定训练外,额外的 RMSNorm 还持续改善了验证损失与下游基准测试。」
逐词翻译:「对尺度变化的敏感性」= 聚合结果忽大忽小(病一的前兆);插 RMSNorm 把它摁回标准尺度(图 4 深蓝色小块)。重点是最后一句:又稳又好——不只为稳定,还真的提升效果。
K3 原文(§2.3.2 Sigmoid Tanh Unit GLU)
「我们将封顶超参数设为门分支 β₁ = 4、升分支 β₂ = 25……SiTU-GLU 得以保留 SwiGLU 的局部响应,同时约束乘积中的两个因子。」
逐词翻译:「局部响应」= 小输入时跟 SwiGLU 一样聪明(图 5 两条线在小输入处重合);「约束乘积中的两个因子」= 给乘法的两边都装限高杆,所以最终输出被锁死在 4×25 = 100 以内。
K3 原文(§2.3.3 Quantile Balancing 直方图估计)
「我们改为从每个专家边际的直方图中读取其分位数:一次 all-reduce 对各个 rank 的分桶计数求和,再从汇总计数中恢复分位数。」
逐词翻译:「边际」= 每个专家离入选线差多少分;「分桶计数求和」= 桶可以相加,所有卡的桶一加就是全局分布;「从汇总计数恢复分位数」= 从压缩后的几百个桶里读出偏置该设多少。几百万 → 几百,这就是 896 个专家还能被管住的关键。
一句话总结今天的 K3 印象
K3 的「又大又稳」来自两套创新:AttnRes 让信息跨层自由流动(伪查询决定看谁,块级设计让它付得起账);Stable LatentMoE 三件套——RMSNorm 治尺度漂移、SiTU-GLU 治激活爆炸(上限 100)、QB 直方图治负载失衡(几百个桶代替几百万个数字)。明天是收官日:视觉编码器 MoonViT-V2、优化器 Per-Head Muon,然后看 K3 到底考了多少分、省钱省到什么程度!
9自测题(先自己答,再点开看解析)
Q1标准残差连接有什么「瓶颈」?AttnRes 怎么治它?
标准残差里每层只把前一层的结果加进来,信息层层接力会「变形」——像传话游戏,传到深层时最早的信息已被压缩走样(报告称之为 RNN 式的瓶颈)。AttnRes 让每一层能直接查看前面所有层的原始输出,按需取用,信息沿深度流动不再走样。
Q2「伪查询 w」是什么?它决定了什么?
w 是每一层专门学出来的一个查询向量(不是真实词向量,所以叫「伪」)。第 3 层用 w₃ 对前面各层的表示打分,分数归一化成权重 α,再加权求和得到自己的输出。它决定了「这一层该重点参考哪些旧笔记」——而且是训练中学出来的,不是人设计的。
Q3块级 AttnRes 为什么更省?K3 怎么配置的?
全量 AttnRes 要为后面所有层保留每一层的输出(O(Ld) 内存);块级 AttnRes 先把每块内的层输出相加成一份「块笔记」,跨块只对块笔记做注意力——内存从「每层一份」变成「每块一份」(O(Nd))。K3 配置:每块 12 层、共 8 块,加上嵌入共 9 个块,几乎不损失效果。
Q4896 个专家会得哪两个「病」?三件套各治什么?
病一:激活爆炸——路由链条(降投影→专家→升投影)放大数字,2.8T 规模下越放越大直到崩溃;病二:负载失衡——近 1000 个专家的均衡超出老方法能力,专家忙闲不均。三件套:RMSNorm 治尺度漂移、SiTU-GLU 治激活爆炸(上限 100)、QB 直方图治负载失衡。合称 Stable LatentMoE。
Q5QB 的「直方图估计」解决了什么问题?为什么能省这么多?
算偏置需要「路由器分数的分位数」,而分数有几百万个、散在几百张卡上——精确收集排序不现实。直方图把分数切成几百个桶、每桶记个数,几百万 → 几百;而且桶可以相加,各卡本地计数后一次 all-reduce 汇总,就等于拿到了全局分布,再从桶里读分位数。通信量只有几百个桶,误差以桶宽为界(桶越细越准)。
明日预告 · Day 7(收官)
创新深读(下)+ 大考成绩单:K3 到底考了多少分?
MoonViT-V2 视觉编码器 · Per-Head Muon 优化器 · 四大轴线评测 · 省钱对比 · 六个真实案例
开始 Day 7 →