Glm5NextForConditionalGeneration model_type · glm5_next KDA × MLA 混合 mHC 四路残差 NoPE k-pool DSA

GLM-5.3-Flash 架构解剖

它与旗舰 GLM-5.3 同名不同构。旗舰是 78 层的 glm_moe_dsa,在 ATOM 里没有自己的模型文件; Flash 是 GLM-5 家族里唯一拥有独立模型文件的成员—— atom/models/glm5_next.py,1257 行:45 层里 34 层线性注意力、11 层稀疏 MLA, 残差是 4 路宽的,全模型没有位置编码,稀疏索引按 4 个 token 一池打分。

320.8 B文本侧总参数
17.38 B单 token 激活
34 + 11KDA 层 + MLA 层
4096hidden_size
×4hc_mult 残差路数
288 / 8专家数 / 每 token
2048 / 4index_topk / kpool
1 Mmax_position

先分清对象

三处 ATOM 从未 serve 过的结构

它新在哪

  • 残差是 4 路的。inputs_embeds 在 embedding 处就被展开成 [T, 4, 4096],一路带到第 45 层,最后用无权重的均值塌回 [T, 4096]。每个子层进出都要收一次、放一次。
  • 整个文本模型没有位置编码。qk_rope_head_dim == 0, MLA 不转 rope,indexer 也不转;位置信息来自 KDA 层的因果卷积与递推。
  • 注意力是混合的。34 层 KDA 线性注意力(无 KV cache,只有递推状态) + 11 层稀疏 MLA,比例 3∶1。

与旗舰 GLM-5.3 的关系

只有名字是共享的。旗舰 GLM-5.3(GlmMoeDsaForCausalLM,78 层 / 6144 / 256 专家 / 753.3 B)的 config 与 GLM-5.2 逐字段相同,挂在 deepseek_v2.py 上; Flash 的 45 层混合结构、四路残差与 k-pool 索引在那边一样都没有。

→ GLM-5.3 架构解剖(旗舰)

checkpoint 是原生多模态的(另带 24 层视觉塔),ATOM 只 serve 文本路径, model.visual.* 在加载时按前缀跳过。


规模

321.3 B 参数落盘,单 token 只碰 17.38 B

下面每个数字都来自 62 个 safetensors 分片的头部(读 shape 相加,fp8 的 weight_scale_inv 单列,不计入参数量)。官方标称 “320B 总 / 18B 激活” 对应文本侧含 MTP 的 320.78 B 与实测 17.38 B。

MoE routed experts288 专家 × 42 层 × 25.17 M
304.41 B94.7 %
MTP 层 45DSA + 完整 MoE,ATOM 未加载
7.43 B2.3 %
KDA 注意力34 层 × 137.7 M
4.68 B1.5 %
MLA 注意力11 层 × 117.4 M
1.29 B0.4 %
MoE shared expert1 个 × 42 层
1.06 B0.3 %
embed_tokens / lm_head各 154880 × 4096 BF16
1.27 B0.4 %
vision tower24 层,ATOM 直接跳过
0.56 B0.2 %
dense MLP层 0–2,d_ff 12288
0.45 B0.1 %
DSA indexer11 层 × 7.47 M
0.082 B0.03 %
MoE router + mHC + normsgate 288×4096 · hc_*_fn 24×16384
0.086 B0.03 %
checkpoint 合计 321.34 B ATOM 实际加载 313.35 B(去掉 vision 与 MTP) 单 token 激活 17.38 B

激活量是怎么凑出来的

  • MoE 专家 9.51 B — 42 层 ×(8 routed + 1 shared)× 25.17 M
  • 注意力 6.06 B — 34 层 KDA 全激活 + 11 层 MLA/indexer
  • embed + lm_head 1.27 B
  • dense MLP 0.45 B — 只有层 0–2
  • mHC + norms 0.036 B

注意力占激活量的 35 %,远高于旗舰的 31 %—— 因为 KDA 层每层 137.7 M 参数全部参与,没有稀疏可言。

量化布局

  • FP8 e4m3 block 128×128,动态激活量化:MoE 专家、dense MLP、 MLA 的 q_a/q_b/kv_a/o_proj
  • BF16 保留:全部 KDA 投影、整个 indexer、kv_b_projlm_head、embedding、所有 norm 与 mHC 参数。
  • weight_scale_inv 的形状即分块数,例如专家 gate_proj [2048, 4096] → scale [16, 32]

核心机制 · 层调度

45 层里只有 11 层有 KV cache

layer_types 把 45 层排成一个严格的 4 拍循环:3 层 KDA,1 层 DSA, 从层 3 起每 4 层一个 DSA,最后一层(44)是 KDA。mlp_layer_types 是另一条独立的 schedule:前 3 层稠密 MLP,之后全是 MoE。

图 1 · 层调度

45 层的两条独立时刻表 layer_types 决定注意力种类 · mlp_layer_types 决定 FFN 种类 0 5 10 15 20 25 30 35 40 44 attention mlp 45 KDA · 34 层 线性注意力,无 KV cache 每请求持一份递推状态 137.7 M / 层 DSA · 11 层 层 3, 7, 11, … 43 MLA + k-pool 稀疏索引 117.4 M + 7.47 M / 层 dense MLP · 3 层 层 0–2 d_ff 12288 151.0 M / 层 MoE · 42 层 层 3–44 288 routed top-8 + 1 shared 7.27 B / 层,激活 227.7 M
两条时刻表互不对齐:注意力按 4 拍循环,FFN 只在层 3 切换一次。层 3 因此是第一个既是 DSA 又是 MoE 的层,调试新算子时优先看它;层 44 是 KDA + MoE。

checkpoint 里的直接证据

  • layers.0.self_attn(KDA):q/k/v_projq/k/v_conv1dA_logdt_biasf_a/f_b_projg_a/g_b_projb_projo_normo_proj——没有任何 q_a/kv_a
  • layers.3.self_attn(DSA):q_a_projq_b_projkv_a_proj_with_mqakv_b_projo_projindexer.*——没有任何 conv1d
  • 两类层的 hc_attn_* / hc_ffn_* 形状完全一致, 所以 mHC 与注意力种类无关。

这条 schedule 省了什么

KV cache 只由 11 层承担而不是 45 层。BF16 下每 token 11 × 576 × 2 B = 12 672 B;若 45 层全是 MLA 则是 51 840 B。 代价是 34 层各自持有一份每请求的递推状态(见“缓存账”一节), 它随并发数而不是随上下文长度增长,也正因如此 prefix caching 必须关掉


核心机制 · 残差

mHC:残差不是一条,是四条互相混合的

这是与 ATOM 现有模型差别最大的一处。Glm5NextHyperConnection 与 DeepSeek-V4 的 Block 在数学上完全一致——同样的 sigmoid 门、 同样的 Sinkhorn 日程(含特殊的第一轮)、同样的 HC_POST_MULT = 2.0—— 连 checkpoint 的张量名(hc_attn_fn / hc_attn_base / hc_attn_scale)都是 Block 期待的名字, hc_attn_fn[24, 16384] 正好是 hc_split_sinkhornmixes 布局。

图 2 · mHC 四路残差

mHC:每个子层进出都要把 4 路残差收一次、放一次 hc_mult = 4 · hc_sinkhorn_iters = 20 · 每层两个站点 stream 0 · 4096 stream 1 · 4096 stream 2 · 4096 stream 3 · 4096 residual [T, 4, 4096] × pre[T,4] 求和 x [T, 4096] 子层自己的 RMSNorm 子层 KDA / MLA 或 MoE / dense MLP comb[T,4,4] 双随机 行和 = 列和 = 1 stream 0′ stream 1′ stream 2′ stream 3′ 新的 residual [T, 4, 4096] post[T,4] × 子层输出 ① 一次线性,出 24 个数 flatten → [T, 16384] RMSNorm(无权重,over 16384) F.linear(hc_attn_fn [24, 16384]) 24 = (2 + hc_mult) × hc_mult ② 拆三份,各配一个激活 pre = sigmoid(·× scale₀ + base) + ε post = 2 · sigmoid(·× scale₁ + base) comb = ·× scale₂ + base → Sinkhorn scale [3] · base [24],都是 fp32 ③ Sinkhorn-Knopp 投影 第 1 轮:行 softmax + ε,再列归一 第 2–20 轮:行归一、列归一交替 落到 Birkhoff 多面体上的双随机矩阵 aiter mhc_pre 把这一整段融成一个 kernel
关键在最右边那条回路:新残差不是 x + sublayer(x),而是子层输出与旧残差按两组学出来的权重整体替换comb 是每 token 一个的 4×4 双随机矩阵,所以四条 stream 在每一层都被重新混合一次——它们不是四份副本,而是四条各自演化、持续交换信息的通路。

mHC 参数为什么平铺在 layer 上

参数的名字来自它的属性路径。checkpoint 里它们是 layers.N.hc_attn_fn 这样平铺在层上的, 所以 ATOM 也必须把它们声明在 Glm5NextDecoderLayer 上, 而不是塞进 Glm5NextHyperConnection 那个 helper—— 埋进子模块会把名字改成 hc.fn,然后静默地停留在初始值


结构全图

从 input_ids 到 logits,再放大其中一层

图 3 · 结构全图

完整模型栈 Glm5NextForConditionalGeneration · 313.3 B 已加载 input_ids [T] embed_tokens [154880, 4096] BF16 634.4 M · 每 rank 全表 unsqueeze(-2).expand(-1,4,-1) [T, 4096] → [T, 4, 4096] layers 0 – 44 残差全程保持 [T, 4, 4096] 34 × KDA · 137.7 M / 层 11 × MLA + k-pool · 124.9 M 3 × dense MLP · 151.0 M 42 × MoE · 7.27 B / 层 每层 2 个 mHC 站点 · 786 K KV cache 只由 11 层承担 312.1 B · 激活 16.1 B HyperHead · residual.mean(-2) [T,4,4096] → [T,4096] · 无权重 model.norm · RMSNorm(4096) lm_head [154880, 4096] BF16 · 634.4 M logits [T, 154880] layer 45 · MTP draft 完整 DSA + MoE · 7.43 B 加载时按层号 ≥ 45 丢弃 DecoderLayer 展开 两个 mHC 站点 + 一个注意力分支 + 一个 FFN 分支 residual · [T, 4, 4096] ① hc.pre(attention 站点) [T,4,4096] → x [T,4096],input_layernorm 融在同一个 kernel 里 分支 A · KDA(34 层) in_proj [32960, 4096] 一次 GEMM 32960 = 4×8192 + 64 + 128 conv1d(k=4) → 64 头 × 128 delta-rule 递推 → o_norm → o_proj 状态 conv [3,24576] + ssm [64,128,128] 分支 B · MLA + DSA(11 层) q_a 4096→1536 · q_b 1536→64×256 kv_a 4096→512 · kv_b 512→64×512 k_pe = zeros[T, 64] ← NoPE 零填充 KV entry 512 + 64 = 576 indexer 先跑,写 top-k 槽位 ② hc.post_expand post × 子层输出 + comb × 旧残差 → [T, 4, 4096] ③ hc.pre(FFN 站点) 同一段代码,换 hc_ffn_* 一组参数与 post_attention_layernorm 分支 C · dense MLP(层 0–2) gate_up_proj [2×12288, 4096] swiglu_oai_split(limit = 10.0) down_proj [4096, 12288] FP8 e4m3 · block 128×128 分支 D · MoE(层 3–44) gate nn.Linear[288, 4096] · fp32 sigmoid + noaux_tc → top-8 / 288 expert d_ff 2048 · + 1 shared × routed_scaling_factor 2.5 ④ hc.post_expand → 下一层的 residual [T, 4, 4096] 这一层为什么这么排 · 四个站点里只有两个带权重矩阵,另外两个是 mHC 的收放,共 786 K · 分支 A/B 由 layer_types 决定,分支 C/D 由 first_k_dense_replace=3 决定 · 子层的 RMSNorm 作用在收拢后的 [T,4096] 上,被 aiter mhc_pre 融进同一 kernel · 层 44 是 KDA + MoE;层 3 是第一个 DSA + MoE 层
看右边时请盯住那条竖脊——它不是常见的 [T, 4096] 残差,而是 [T, 4, 4096]。四个编号站点里,① 和 ③ 是同一段代码,只是换了一组 hc_* 参数;收放的机制见图 2。

分支 A · 34 层的主力

KDA:一次 GEMM 出六路,再走 delta-rule 递推

Glm5NextKDAAttention 直接继承 Kimi-K3 的 KimiKDAAttention。 两者结构上是同一层——同样把 q/k/v_conv1d 分开存、同样的每头 A_log、 每通道 dt_bias、低秩遗忘门——唯一的差别是输出门: Kimi 一步投影(g_proj),GLM 拆成秩 128 的两步(g_b_proj @ g_a_proj)。

图 4 · KDA 数据通路

KDA:一次 GEMM 出六路,再走 delta-rule 递推 形状按 TP1 标注;TP8 下每 rank 只留 8 个头 x [T, 4096] (mHC 收拢后) in_proj · 一次 GEMM [32960, 4096] BF16 → [T, 32960] checkpoint 只给 q/k/v;g 由 g_b @ g_a 折叠写入,b_proj 与 f_a_proj 在 load 后追加为两条尾巴 1 按列切 6 段(末两段宽度画大了,真实比例是 64 / 128 对 8192) q · 8192 64 头 × 128 k · 8192 64 头 × 128 v · 8192 64 头 × 128 g · 8192 折叠得到的输出门 b · 64 每头一个 f_a · 128 低秩 causal_conv1d · kernel 4 · SiLU conv_weight [24576, 4](q|k|v 三份预拼) 读写每请求 conv 状态 [3, 24576] BF16 2 beta b.float() → [1, T, 64] 必须 fp32 f_b_proj [8192, 128] gate [1,T,64,128] delta-rule 递推 prefill: aiter chunk_kimi_delta_attn(q, k, v, g, beta, A_log[64], dt_bias[8192]) decode: fused_sigmoid_gating_delta_rule_update(就地更新最终状态) 递推状态 ssm_state [64, 128, 128] fp32 —— 每层每请求 512 KiB use_qk_l2norm_in_kernel · 每通道遗忘门下界 gate_lower_bound = −5.0 3 o_norm · KimiRMSNormGated(128) 用 g 分片门控 [T, 64, 128] 可把输出量化融进同一个 kernel o_proj [4096, 8192] BF16 · RowParallel → x [T, 4096] 交回 hc.post_expand out [T, 4096] g 绕过递推,直达 o_norm
切片顺序 q|k|v|g|b|f_a 是 forward 里硬编码的边界,与 process_weights_after_loading 的拼接顺序必须一致。TP8 下 in_proj 变成 [4232, 4096]4×1024 + 8 + 128),ssm_state 变成 [8, 128, 128]

为什么折叠输出门是精确的

g = g_b_proj(g_a_proj(x)),两步之间没有激活函数, 所以 W_g = W_b @ W_a 在 fp32 下乘一次就等价。ATOM 在 process_weights_after_loading 里把结果写进 in_proj.weight[3·lp : 4·lp],再把两个因子的存储清空。 这样父类的状态缓存、TP 切分与 CUDA graph 处理全部原样复用。 加载完整性靠 _loaded_input_shards 记账—— q/k/v 三个 checkpoint 张量指向同一个参数, 参数级的加载报告分辨不出“q 到了”和“q/k/v 都到了”。

transformers 侧的一个真 bug

checkpoint 的 modules_to_not_convert 写的是 model.layers.N.self_attn.f_a_proj,但真实 key 前缀是 model.language_model.layers.,而且 glm5_next 的转换表会在 FP8 量化器跑之前把它改名成 self_attn.forget_gate.f_a_proj。 两处都对不上 → 全部 68 个遗忘门线性层(34 层 × 2)被包进 FP8Linear 却仍持 BF16 权重、配一个新初始化的 weight_scale_inv。用 transformers 直接跑这个 checkpoint 的人, KDA 的衰减是静默损坏的。ATOM 不走这条路,不受影响。


分支 B · 11 层的 MLA

NoPE 不是“少算一步 rope”,是要凑出 64 lane 的零

config 写的是 qk_rope_head_dim = 0mla_use_nope = true。 最自然的实现是让 rope 那一半是零宽切片——这也是本次移植最初的选择, 而它在两个互相独立的方向上都是错的。

图 5 · NoPE 的两种实现

qk_rope_head_dim = 0 的两种实现方式,只有一种能跑 latent 侧要 576,per-head 侧要 ≤ 256 —— 两个约束作用在不同张量上 方案一 · 让 rope 半边是零宽切片 kv_lora_rank · 512 + 0 宽 KV entry = 512 × asm decode kernel 按 576 编译,只在 gfx1250 断言 × KV_PeDim = 0 → tl.arange(0,0),Triton 编译期拒绝 prefill 走 flash_attn_varlen,对 head-dim 通用 —— 所以它看起来是好的 方案二 · _ROPE_PAD = 64 条恒为零的通道 kv_lora_rank · 512 0 × 64 KV entry = 576 (config.mla_kv_entry_dim 声明给分配器) 零块对点积的贡献是 Σ(0×0) = 0 —— 这在数值上就是 NoPE 本身 rope_is_zero_pad 让 prefill 在每个 flash-attn 站点丢掉这 64 条 两侧各自看到它需要的宽度,且都是精确的 prefill 侧 flash_attn_varlen_func(CK) 丢掉 64 条零 每头 q/k 宽度 = 256 CK 的 head_dim 上限就是 256 decode 侧 aiter asm mla_decode_fwd 查询宽度必须 = 576 q_nope 被 kv_b 吸收 → 512 ⧺ 0 × 64 = 576 为什么不能直接把 per-head 也补到 320 qk_nope_head_dim 已经是 256,补完是 320 CK 的 flash-attention 把 head_dim 卡在 256 Kimi-K3 撞不到:它的 nope 是 128,128+64=192 所以 padding 只在 latent 侧有效,per-head 侧必须丢回去
这两个方案的差别只有一个数:_ROPE_PAD 是 0 还是 64。看起来更“干净”的那个是错的——而且错得不显眼:prefill 路径对 head-dim 是通用的,所以只做 prefill 的逐层比对会给一个全绿的结果。

零宽切片错在哪 · 其一:静默算错

paged MLA entry 按 kv_lora_rank + qk_rope_head_dim 算出来是 512, 而 aiter 的 asm decode kernel 是按 576 宽的查询编译的。 它只在 gfx1250 路径上断言这个 576,cfg_mla_asm 的分发表又从不按 head_size 选择, 于是 gfx950 上这次不匹配被算出来而不是被拒绝。 更麻烦的是它只影响 decode:prefill 走 flash_attn_varlen, 对 head-dim 是通用的,所以任何只覆盖 prefill 的逐层比对都会放行一个坏模型。

零宽切片错在哪 · 其二:直接崩

KV_PeDim == 0 让每一处 tl.arange(0, KV_PeDim) 变成 arange(0, 0),Triton 在编译期就拒绝;上游 aiter 的三个 gather_kv_b_proj kernel 对此没有任何保护。 它只在命中前缀缓存的 prefill 路径上触发, 所以单条 prompt 的 demo 能过,一上并发就死于 NameError('kv_pe_data is not defined')

两个容易漏掉的配套

其一,填充后的宽度必须用 config.mla_kv_entry_dim 声明给缓存分配器: KimiMLAGDNBackend 会遮蔽普通 MLA 分配器、直接按原始 config 算 pool 大小, 少了这个声明就会 pool 按 512 建、写入按 576 走,服务在启动时死于 shape '[..., -1, 576]' is invalid。 其二,这个 pad 必须由 _ZeroRopePad 追加,而它刻意不是 nn.Module——把 q_b_proj 包进一个 Module 会在参数路径里 插进一层,权重就再也加载不上了。


核心机制 · 稀疏索引

k-pool:4 个 token 压成一行,top-k 在池上选

DSA 层的 indexer 不给每个 token 存一行 key。连续 index_kpool = 4 个 token 压成 一个 cache entry, top-k 因此在池粒度上跑——index_topk / index_kpool = 512 个候选池—— 每个被选中的池再展开回它覆盖的 4 个 token 位置。 还没压满的“尾池”按 index_kpool_always_select_tail 永远入选, 所以最新的 token 绝不会被丢掉。

图 6 · k-pool 索引

写入侧 · 每 4 个 token 产出一行 k = k_norm(wk(hidden)) [T,128] gate = index_kpool_compress_gate(hidden) [T,128] 0 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 pool 0 pool 1 pool 2 pool 3 pool 4 未满的尾池 压缩:按维度独立的 slot softmax w[slot, d] = softmax over slot ( gate[p, slot, d] + ape[slot, d] ) pooled[p, d] = Σ_slot w[slot, d] · k[p, slot, d] ape = compress_ape [4, 128] softmax 沿 slot 轴、每个维度各做一次 —— 不是每 slot 一个标量门 1 row 0 · 144 B row 1 · 144 B row 2 · 144 B row 3 · 144 B row 4 · 144 B tail_cache [L, slot, 2, 4, 128] bf16 → Hadamard-128(正交,含 1/√128 归一)→ ue8m0 幂次缩放 FP8 → indexer_k_quant_and_cache 同一个 H 也作用在 query 上,⟨Hq, Hk⟩ == ⟨q, k⟩;而 1/√128 = 2^−3.5 不是 2 的幂,量化出的字节会不同 读取侧 · 池粒度打分,token 粒度交付 q = wq_b(q_c) [4096, 1536] → [T, 32, 128] fwht128_quant_fp8 2 mqa logits prefill: fp8_mqa_logits decode: paged_mqa_logits 3 top-k on pools k = 2048 / 4 = 512 → [T, 512] int32 4 expand ×4 tok = pid·4 + {0,1,2,3} → 2048 个 token id 5 append tail 最多再补 3 个 有效列 = 2051 6 物理行宽 topk_out_width = ceil((2048 + 4 − 1) / 128) × 128 = 2176 生产者与 MLA metadata 共用这一个宽度,由 model_ops/glm5_next/geometry.py 单点定义 两个 regime,边界在 max_seqlen_k = index_topk = 2048 · ≤ 2048:top-k 选中每一个池,展开后覆盖全部位置,与稠密因果 MLA 数值完全相同 · > 2048:池化打分与池化 top-k 真正决定注意力看什么 · 此时若 ATOM_GLM5_KPOOL=0,代码直接抛 NotImplementedError,而不是悄悄降级
尾池的原始 kgate 必须活过产生它们的那一步。ATOM 把它们塞进 KDA 已经拥有的每请求状态槽而不是另开一个分页缓存,于是免费继承了状态池的生命周期、fork 与搬迁语义。少了搬迁那一步,一个被重定位的请求就会读到前一个请求留在新槽里的半截池——只有在负载下池边界移动时才会暴露的那种损坏。

block size 为什么被改成 64

deepgemm_fp8_paged_mqa_logits 只在 preshuffle 布局下算得对, 而该布局要求每个 block 的行数是 16 的倍数。 一个 B token 的 block 需要 B / 4 个索引行, 所以最小可行的 B 就是 4 × 16 = 64Config 在架构判定处直接改写 kv_cache_block_size, 因为 BlockManager 与 slot_mapping 都假定全局只有一个 block size。

index cache 1584 B/token → 396 B/token

off-switch 为什么必须报错而不是降级

ATOM_GLM5_KPOOL=0 同时关掉池化写入、池化打分与池化选择, 让两套实现成为同一个选择的两个真正独立的实现—— 这正是短上下文下的 A/B 能成为检验而不是同义反复的原因。 但超过 index_topk 之后两者的选择会真正分叉, 此时返回 token 粒度的回退结果是静默地错,而不只是慢, 所以 _assert_kpool_regime 在那里直接抛 NotImplementedError


权重实测

从 62 个 safetensors 分片头部读到的真实形状

checkpoint 里所有张量都在 model.language_model.* / model.visual.* 之下, 只有 lm_head 在顶层。ATOM 的 weights_mapping 因此只有一条规则: "model.language_model." → "model."model.visual.*skip_weight_prefixes 丢弃,MTP 层 45 由 loader 的“层号越界”过滤自动丢弃。

checkpoint 张量(去掉前缀)dtypeshape说明
embed_tokens.weightBF16[154880, 4096]634.4 M,每 rank 全表
layers.N.hc_attn_fn / hc_ffn_fnBF16[24, 16384]24 = (2+4)×4,16384 = 4×4096;ATOM 侧声明为 fp32,加载时上转
layers.N.hc_*_base / hc_*_scaleF32[24] / [3]激活前的加性偏置 / pre·post·comb 三个缩放
KDA 层(0,1,2,4,5,6,8,…,44 共 34 层)
self_attn.q_proj / k_proj / v_projBF16[8192, 4096]各自映射到融合 in_proj 的分片 0/1/2
self_attn.q_conv1d / k_conv1d / v_conv1dBF16[8192, 1, 4]每通道因果卷积,kernel = 4
self_attn.g_a_proj / g_b_projBF16[128, 4096] / [8192, 128]秩 128 输出门,load 后折叠进 in_proj 分片 3
self_attn.f_a_proj / f_b_projBF16[128, 4096] / [8192, 128]低秩遗忘门
self_attn.A_log / dt_biasF32[64] / [8192]每头 / 每通道
self_attn.b_proj.weightBF16[64, 4096]beta,每头一个标量
self_attn.o_norm.weight / o_proj.weightBF16[128] / [4096, 8192]门控 RMSNorm + 输出投影
DSA 层(3,7,11,…,43 共 11 层)
self_attn.q_a_proj / q_b_projF8_E4M3[1536, 4096] / [16384, 1536]scale [12,32] / [128,12]
self_attn.kv_a_proj_with_mqaF8_E4M3[512, 4096]没有 rope 分量,输出宽度就是 kv_lora_rank
self_attn.kv_b_proj.weightBF16[32768, 512]64 × (256 + 256),不量化,被 MLA 吸收
self_attn.o_proj.weightF8_E4M3[4096, 16384]64 头 × v_head_dim 256
indexer.wq_b / wk / weights_projBF16[4096, 1536] / [128, 4096] / [32, 4096]32 头 × 128;整个 indexer 都在量化排除名单里
indexer.index_kpool_compress_gateBF16[128, 4096]checkpoint 存成裸矩阵,映射规则补上 .weight
indexer.index_kpool_compress_apeBF16[4, 128]每 slot 的可学习偏置,进 softmax 前相加
FFN
mlp.gate_proj / up_proj(层 0–2)F8_E4M3[12288, 4096]scale [96, 32]
mlp.gate.weight / e_score_correction_biasBF16 / F32[288, 4096] / [288]ATOM 用 nn.Linear(dtype=fp32),与参考实现一致
mlp.experts.{0..287}.{gate,up}_projF8_E4M3[2048, 4096]每层 288 份 · 单专家 25.17 M
mlp.experts.*.down_projF8_E4M3[4096, 2048]合并成 gate_up_proj / down_proj 两个大张量
mlp.shared_experts.*F8_E4M3同上1 个,与 routed 输出相加
不加载
layers.45.*(MTP)eh_proj [4096, 8192]完整 DSA + MoE + enorm/hnorm/shared_head.norm,7.43 B
model.visual.*24 层 · hidden 1024image 448 / patch 14 / merge 2 / 时序 patch 2 → out 4096,0.56 B

packed_modules_mapping 为什么可以是层无关的

Kimi-K3 必须逐 KDA 层枚举 .q_proj → .in_proj,因为它的全注意力层也有一个 g_proj 不能被折叠。GLM-5.3-Flash 不需要:它的 MLA 层用的是 q_a_proj / q_b_proj / kv_a_proj_with_mqa / kv_b_proj, 没有一个包含 .q_proj.k_proj.v_proj 子串(匹配锚定在前导的点上)。 于是一张层无关的表就够了——而且只有层无关的表才能挂在上, 让 model_runner 在模型构造之前读到它。 这个顺序很要紧:从 __init__ 里二次重映射量化配置会破坏它的 layer pattern, 并静默地把每一个注意力投影都标成已量化。


显存

缓存账:一半随上下文长,一半随并发长

混合架构把显存开销劈成两块。分页部分(MLA KV + 索引 key)只由 11 个 DSA 层贡献, 按 token 计费;每请求部分(KDA 递推状态 + 尾池)由 34 个 KDA 层贡献, 按并发数计费,与上下文长度无关。

缓存形状 / 计费单位大小备注
MLA KV(分页)11 层 × 576 × dtype / token 12 672 B/token
BF16 · FP8 时减半
若 45 层全是 MLA 则是 51 840 B
索引 key(分页)11 层 × 144 B / 4 token 396 B/token 未池化时 1 584 B;144 = (128 + 4) 向上对齐到 16
分页合计每 token 13 068 B ≈ 12.8 KiB 1 M 上下文 ≈ 12.8 GiB(BF16 KV)
一个 block(64 token)MLA 811 008 B + index 25 344 B 836 352 Bblock size 由 index_kpool × 16 定出
KDA conv 状态34 层 × [3, 24576/tp] BF16 0.60 MiB/请求
TP8
24576 = 128×64×2 + 128×64
KDA 递推状态34 层 × [64/tp, 128, 128] F32 17.0 MiB/请求
TP8;TP4 为 34.0
必须 fp32,aiter 内核原样读回
k-pool 尾池11 层 × [slot, 2, 4, 128] BF16 22 KiB/请求与 KDA 状态共用同一个槽 id
每请求合计TP8 17.6 MiB 并发 32 时约 0.55 GiB

prefix caching 必须关

KDA 的递推状态是每请求的,不是每 block 的,无法像分页 KV 那样在请求间共享。 带前缀命中的 prefill 会 fork 状态槽:读槽(state_fork_src)与写槽不同, 尾池缓存必须跟着一起搬——relocate_state_slots 就是干这个的。


ATOM 侧

这次移植是“组装 + 一个新算子”,不是从零写一个模型

45 层里几乎每一块难的东西,ATOM 都已经因为别的模型存在了。真正新写的只有 k-pool 索引: model_ops/glm5_next/indexer.py 负责分发,kpool.py 负责内核, geometry.py 负责那几个 CPU 可测的几何契约。

GLM-5.3-Flash 的部件ATOM 里已有的东西贴合度
KDA 线性注意力kimi_k3.KimiKDAAttention + aiter kimi_delta_attn 非常近。同样分开存的 q/k/v_conv1d、每头 A_log、每通道 dt_bias、 低秩遗忘门,连 gate_lower_bound 都已经在读。仅需折叠输出门
mHC 超连接sparse_attn_v4.hc_split_sinkhorn + deepseek_v4.Block 数学上完全一致,连 checkpoint 的张量名都是它期待的名字; dim = 4096 也满足 aiter 融合 mhc_pre/mhc_post% 512 == 0 约束
k-pool DSA 索引新写 DeepSeek-V4 的 Compressor 也按 compress_ratio=4 池化并带 ape 项, 但它是重叠的、带 RoPE 的;GLM 这个不重叠、无 RoPE
MLAmodel_ops/attention_mla.py 经 MLAModules 需要 _ROPE_PAD = 64rope_is_zero_pad,见图 5
288 专家 sigmoid / noaux_tcmodel_ops/fused_moe、glm4_moe.py、deepseek_v2.py直接可用
Block FP8 128×128既有 DeepSeek block-FP8 路径直接可用
clamped SwiGLUmodel_ops/swiglu_oai.swiglu_oai_split alpha=1, beta=0, limit=10.0 等价于参考实现的 silu(clamp(gate))·clamp(up)

三个调试开关

  • ATOM_GLM5_KPOOL=0(默认 1)—— 同时关掉池化写入、池化打分与池化选择。 超过 2048 时直接报错,不静默降级。
  • ATOM_GLM5_FORCE_DENSE_MLA=1(默认 0)—— 关掉稀疏。 2048 以内本就等价,任何输出差异就定位在 indexer / top-k 而不是 MLA 本身。
  • ATOM_GLM5_DISABLE_FUSED_MHC=1(默认 0)—— 把 aiter 的 mhc_pre/mhc_post 换成 torch 参考路径,用于逐层对数。

明确拒绝的组合

Config 在识别到 Glm5Next* 架构后会直接抛错而不是降级: PCP > 1、DCP > 1、投机解码、TBO。理由是这些特性的池化索引 metadata / 状态布局 还没有对应的实现。kpool 的 dispatch 里另有两道同样性质的关卡: DCP/PCP world size > 1、以及 max_seqlen_q > 1 的投机解码。 MTP 草稿层与多模态输入同样尚未接通。


读码路线

按这个顺序看,一小时能过完

#文件 · 位置看什么
1atom/models/glm5_next.py · 模块 docstring 三个决定:池粒度选择、64 lane 零 rope、输出门折叠。先读这段,后面全是它的展开
2_normalize_glm5_next_config() config 别名怎么补齐;glm5_kda_layers0-based(Kimi-K3 是 1-based,别照抄); 三处 raise 分别守住 NoPE、index 几何、层划分的完整性
3Glm5NextDecoderLayer.forward 四个站点的顺序,以及 mHC 参数为什么平铺在 layer 上而不是塞进 helper
4Glm5NextHyperConnection.pre / post_expand 融合内核与 torch 参考两条路径并列,逐行对得上(图 2)
5Glm5NextMLAAttention.__init__ / forward _ZeroRopePadmla_kv_entry_dim、 以及 indexer 为什么要在 mla_attn 之前被显式调用
6model_ops/glm5_next/kpool.py · 前 115 行 三个 *_ref 参考实现就是正确性的定义:池化 softmax 的轴、 Hadamard 的 1/√128、ue8m0 的幂次缩放
7model_ops/glm5_next/indexer.py prefill 与 decode 两条分发;_kpool_write_completed_pools 如何从 tail cache 里 补上跨 chunk 的槽
8model_ops/attentions/kimi_mla_gdn_attn.py · 172–300 索引缓存行数、尾池 buffer 与 KDA 状态槽共享生命周期
9atom/config.py · glm5_kpool_block_size / 2035 附近 block size 为什么被改成 64,以及被明确拒绝的并行特性
10recipes/GLM-5.3-Flash.md bring-up 期发现的 4 个 bug(3 个在上游)与参考 oracle 的搭法