Document LatentMoE sigmoid routing, sparse dispatch, and K3 block figures

Ledger C10–C12 match the permute-pad-bmm path and Switch aux/z-loss.
Section 8 adds overview and component TikZ; MoE capacity is C_moe so it
does not collide with KDA chunk size.
This commit is contained in:
dela
2026-08-26 14:43:58 +08:00
parent ea7167b3f7
commit a2c4217dae
6 changed files with 239 additions and 18 deletions
+46
View File
@@ -45,6 +45,8 @@ questions:
text: "Block AttnRes 的两阶段算法为什么和 naive 逐层实现数值等价?"
- id: Q9
text: "深度残差接入 CausalLM 时怎样避免参数被重复注册?"
- id: Q10
text: "为什么不用 stack([e(z) for e in experts]) 稠密计算全部专家?稀疏 permute-dispatch 如何让每个 token 只算 k 个专家?"
claims:
- id: C1
@@ -83,6 +85,18 @@ claims:
text: "BorrowedSubLayer 用普通 tuple 持有 norm/fn,不注册为子模块,保证参数与 state_dict 键不重复"
kind: methodological
status: supporting
- id: C10
text: "LatentMoE 稀疏执行 = permute-dispatch + pad 到 [R, C, ℓ] + 三次 bmm + scatter-add,每个 token 只算 k 个专家(FLOPs R·C 而非 R·N)"
kind: methodological
status: core
- id: C11
text: "K3 路由 = s=σ(W_r x)、Top-k(s+b)、p_i = s_i/Σ_{j∈T}s_j;expert_bias 只进 TopK 选择、不进归一化权重"
kind: methodological
status: core
- id: C12
text: "负载均衡:Switch/GShard aux = n_r·Σ f_e·P_e 与 router z-loss = mean (logsumexp logits)^2,训练时加到 CE 上,只更新 router"
kind: methodological
status: core
symbols:
- {name: B, latex: "B", meaning: "batch size", kind: "shape parameter"}
@@ -115,6 +129,14 @@ symbols:
- {name: h_l, latex: "h_l", meaning: "深度注意力聚合出的层输入", domain: "[B, T, D]", kind: value}
- {name: b_j, latex: "b_j", meaning: "Block AttnRes 第 j 块的输出", domain: "[B, T, D]", kind: value}
- {name: p, latex: "p", meaning: "块内 running partial", domain: "[B, T, D]", kind: value}
- {name: s_moe, latex: "s", meaning: "router sigmoid 分数 σ(W_r x)", domain: "[B, T, n_r]", kind: value}
- {name: b, latex: "b", meaning: "expert bias(非持久 buffer,只进 TopK)", domain: "[n_r]", kind: value}
- {name: p_i, latex: "p_i", meaning: "sigmoid-L1 路由权重", domain: "[B, T, k]", kind: value}
- {name: C_moe, latex: "C_{\\mathrm{moe}}", meaning: "MoE 专家容量 = max 负载(pad 宽度)", kind: "shape parameter"}
- {name: f_e, latex: "f_e", meaning: "专家 e 被路由到的 token 占比", kind: value}
- {name: P_e, latex: "P_e", meaning: "专家 e 的平均 sigmoid 分数", kind: value}
- {name: L_aux, latex: "\\mathcal{L}_{aux}", meaning: "Switch/GShard 负载均衡损失", kind: value}
- {name: L_z, latex: "\\mathcal{L}_z", meaning: "router z-loss", kind: value}
terms:
- {canonical: "KDA", aliases: ["Key-Decayed Attention", "键衰减注意力"]}
@@ -129,6 +151,10 @@ terms:
- {canonical: "depth residual", aliases: ["DepthResidual", "深度维残差"]}
- {canonical: "online softmax", aliases: ["在线 softmax", "增量 softmax"]}
- {canonical: "atomic layer", aliases: ["原子层", "atomic sublayer"]}
- {canonical: "permute-dispatch", aliases: ["置换-分发", "专家分发", "dispatch"]}
- {canonical: "grouped GEMM", aliases: ["padded bmm", "分组矩阵乘", "batched GEMM"]}
- {canonical: "load balancing loss", aliases: ["负载均衡损失", "aux loss", "Switch/GShard aux"]}
- {canonical: "z-loss", aliases: ["router z-loss", "logit 正则"]}
derivations:
- id: DER1
@@ -162,6 +188,26 @@ derivations:
- {id: "3", from: "单源 partial p", to: "(m, n, d) = (s_p, p, 1),因为 e^{s_p - m} = 1", rule: definition}
- {id: "4", from: "(m_a,n_a,d_a), (m_b,n_b,d_b)", to: "m = \\max(m_a,m_b);\\ n = e^{m_a-m} n_a + e^{m_b-m} n_b;\\ d = e^{m_a-m} d_a + e^{m_b-m} d_b", rule: scale}
- {id: "5", from: "(m, n, d)", to: "h_l = n / d,与 forward_naive 逐位一致", rule: definition}
- id: DER4
claim: C11
title: "K3 sigmoid-TopK 路由推导"
expand: true
figure: null
steps:
- {id: "1", from: "l = W_r x", to: "s = \\sigma(l) \\in [B,T,n_r]", rule: definition}
- {id: "2", from: "s + b", to: "T = \\mathrm{TopK}(s+b, k)", rule: selection}
- {id: "3", from: "T, s", to: "p_i = s_i / \\sum_{j \\in T} s_j", rule: normalize}
- {id: "4", from: "p, z", to: "u = \\sum_{i \\in T} p_i E_i^{rt}(z)", rule: definition}
- id: DER5
claim: C10
title: "稀疏 dispatch 执行流推导"
expand: true
figure: null
steps:
- {id: "1", from: "tok 重复 k 次 + eid 扁平化", to: "order = argsort(eid),同专家 token 连续", rule: permute}
- {id: "2", from: "counts = bincount(eid)", to: "C = max(counts);padded = index_put(zeros[R,C,ℓ], (eid, local_pos), z[tok])", rule: pad}
- {id: "3", from: "padded + 堆叠权重 [R,...]", to: "三次 bmm 得 [R,C,ff] → [R,C,ℓ](grouped GEMM)", rule: substitute}
- {id: "4", from: "out[eid,local_pos] 加权", to: "u = index_add(0, tok, p ⊙ out),FLOPs R·C 而非 R·N", rule: scatter-add}
figures:
- id: F1