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.
234 lines
12 KiB
YAML
234 lines
12 KiB
YAML
schema: superpaper.ledger/v1
|
||
retired_ids: []
|
||
paper:
|
||
id: "kda-project"
|
||
title: "KDA 训练→推理 手写实现 — 完整笔记"
|
||
authors: ["dela"]
|
||
notes_language: zh
|
||
source:
|
||
kind: markdown
|
||
|
||
coverage:
|
||
mode: full
|
||
sections_in:
|
||
- "KDA 递归核心"
|
||
- "Gate 激活"
|
||
- "分块并行计算"
|
||
- "GVA 分组值注意力"
|
||
- "KDAAttention 层"
|
||
- "Gated MLA 矩阵吸收版"
|
||
- "SiTU-GLU 与 Stable LatentMoE"
|
||
- "K3 混合架构"
|
||
- "Attention Residual 深度残差"
|
||
- "反向传播推导"
|
||
sections_skipped:
|
||
- "Triton kernel 细节"
|
||
- "Docker 部署"
|
||
- "AttnRes 论文的 kernel 级调度与 pipeline 重叠"
|
||
|
||
questions:
|
||
- id: Q1
|
||
text: "KDA 的状态更新如何避免 softmax、实现线性复杂度?"
|
||
- id: Q2
|
||
text: "safe gate 与 standard gate 的区别是什么?"
|
||
- id: Q3
|
||
text: "分块并行如何在保持递归等价的同时利用 GPU 并行?"
|
||
- id: Q4
|
||
text: "GVA 的 repeat_interleave + sum 反向是怎么回事?"
|
||
- id: Q5
|
||
text: "MLA 矩阵吸收如何避免解压 K/V?"
|
||
- id: Q6
|
||
text: "SiTU-GLU 为什么比 SwiGLU 更稳定?"
|
||
- id: Q7
|
||
text: "AttnRes 如何把残差流从等权累加换成按内容选择?"
|
||
- id: Q8
|
||
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
|
||
text: "KDA 用 delta rule 更新 KV 状态矩阵,不需要 softmax,复杂度 O(T·K·V)"
|
||
kind: methodological
|
||
status: core
|
||
- id: C2
|
||
text: "safe gate = lower_bound · σ(rate · input),保证 gate 值在 [lower_bound, 0] 范围内"
|
||
kind: methodological
|
||
status: core
|
||
- id: C3
|
||
text: "分块计算:chunk 内用下三角解,chunk 间用状态递推,数值等价于 naive recurrent"
|
||
kind: methodological
|
||
status: core
|
||
- id: C4
|
||
text: "MLA 矩阵吸收:q 吸收 W_UK 后直接与 latent c 内积,永不解压 K/V"
|
||
kind: methodological
|
||
status: core
|
||
- id: C5
|
||
text: "LatentMoE 通过 latent 接口把 routed 专家限制在半宽空间 ℓ=d/2"
|
||
kind: methodological
|
||
status: core
|
||
- id: C6
|
||
text: "AttnRes 用逐 token 的深度维 softmax 代替等权残差累加:打分在 RMS 归一化后做,加权和在原始张量上做"
|
||
kind: methodological
|
||
status: core
|
||
- id: C7
|
||
text: "Block AttnRes 块内退化为普通求和、只让块输出进入源列表,源数从 O(N) 降到 O(N/S)"
|
||
kind: methodological
|
||
status: core
|
||
- id: C8
|
||
text: "两阶段算法 = inter 块间批量 einsum + intra online-softmax 增量合并,与 naive 逐层实现数值等价 (atol 1e-5)"
|
||
kind: methodological
|
||
status: core
|
||
- id: C9
|
||
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"}
|
||
- {name: T, latex: "T", meaning: "序列长度", kind: "shape parameter"}
|
||
- {name: H, latex: "H", meaning: "query/key 头数", kind: "shape parameter"}
|
||
- {name: HV, latex: "H_V", meaning: "value 头数 (GVA)", kind: "shape parameter"}
|
||
- {name: G, latex: "G", meaning: "GVA 组数 = HV/H", kind: "shape parameter"}
|
||
- {name: K, latex: "K", meaning: "key/query 头维度", kind: "shape parameter"}
|
||
- {name: V, latex: "V", meaning: "value 头维度 (= K)", kind: "shape parameter"}
|
||
- {name: D, latex: "D", meaning: "hidden_size", kind: "shape parameter"}
|
||
- {name: C, latex: "C", meaning: "chunk_size", kind: "shape parameter"}
|
||
- {name: r, latex: "r", meaning: "KV latent rank (kv_lora_rank)", kind: "shape parameter"}
|
||
- {name: ell, latex: "\\ell", meaning: "MoE latent 宽度 = d/2", kind: "shape parameter"}
|
||
- {name: S, latex: "S", meaning: "KV 状态矩阵", domain: "[B, HV, K, V]", kind: value}
|
||
- {name: q, latex: "q", meaning: "query", domain: "[B, T, H, K]", kind: value}
|
||
- {name: k, latex: "k", meaning: "key", domain: "[B, T, H, K]", kind: value}
|
||
- {name: v, latex: "v", meaning: "value", domain: "[B, T, HV, V]", kind: value}
|
||
- {name: g, latex: "g", meaning: "gate (log-space decay)", domain: "[B, T, HV, K]", kind: value}
|
||
- {name: beta, latex: "\\beta", meaning: "学习率/写入强度", domain: "[B, T, HV]", kind: value}
|
||
- {name: A_log, latex: "A_{\\log}", meaning: "head-wise 衰减参数 (log-space)", domain: "[HV]", kind: value}
|
||
- {name: dt_bias, latex: "\\Delta_b", meaning: "per-dim gate bias", domain: "[HV, K]", kind: value}
|
||
- {name: c, latex: "c", meaning: "KV latent 向量", domain: "[B, T, r]", kind: value}
|
||
- {name: W_UK, latex: "W_{UK}", meaning: "Key 解压矩阵 (MLA)", domain: "[H, d_q, r]", kind: value}
|
||
- {name: W_UV, latex: "W_{UV}", meaning: "Value 解压矩阵 (MLA)", domain: "[H, d_v, r]", kind: value}
|
||
- {name: N, latex: "N", meaning: "AttnRes 原子层数 = 2L", kind: "shape parameter"}
|
||
- {name: S, latex: "S", meaning: "AttnRes 块大小(原子层)", kind: "shape parameter"}
|
||
- {name: v_i, latex: "v_i", meaning: "AttnRes 第 i 个源(v_0 = embedding 输出)", domain: "[B, T, D]", kind: value}
|
||
- {name: w_l, latex: "w_l", meaning: "第 l 层 depth query(零初始化)", domain: "[D]", kind: value}
|
||
- {name: alpha, latex: "\\alpha_{l,i}", meaning: "深度维 softmax 权重", domain: "[n, B, T]", kind: value}
|
||
- {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", "键衰减注意力"]}
|
||
- {canonical: "GVA", aliases: ["Grouped Value Attention", "分组值注意力"]}
|
||
- {canonical: "MLA", aliases: ["Multi-head Latent Attention", "多头隐变量注意力"]}
|
||
- {canonical: "MoE", aliases: ["Mixture of Experts", "混合专家"]}
|
||
- {canonical: "SiTU-GLU", aliases: ["Sigmoid Tanh Unit GLU"]}
|
||
- {canonical: "delta rule", aliases: ["δ 规则"]}
|
||
- {canonical: "safe gate", aliases: ["安全门控"]}
|
||
- {canonical: "matrix absorption", aliases: ["矩阵吸收"]}
|
||
- {canonical: "AttnRes", aliases: ["Attention Residual", "注意力残差", "深度残差"]}
|
||
- {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
|
||
claim: C1
|
||
title: "KDA 递归状态更新推导"
|
||
expand: true
|
||
figure: null
|
||
steps:
|
||
- {id: "1", from: "S_{t-1}", to: "S_{\\mathrm{dec}} = \\exp(g_t) \\odot S_{t-1}", rule: scale}
|
||
- {id: "2", from: "S_{\\mathrm{dec}}", to: "r_t = v_t - k_t \\cdot S_{\\mathrm{dec}}", rule: definition}
|
||
- {id: "3", from: "r_t", to: "S_t = S_{\\mathrm{dec}} + (\\beta_t k_t) \\otimes r_t", rule: definition}
|
||
- {id: "4", from: "S_t", to: "o_t = (q_t \\cdot \\text{scale}) \\cdot S_t", rule: definition}
|
||
- id: DER2
|
||
claim: C4
|
||
title: "MLA 矩阵吸收推导"
|
||
expand: true
|
||
figure: null
|
||
steps:
|
||
- {id: "1", from: "q \\in [B,T,H,d_q]", to: "q_{\\mathrm{abs}} = q \\cdot W_{UK} \\in [B,T,H,r]", rule: substitute}
|
||
- {id: "2", from: "q_{\\mathrm{abs}}, c", to: "\\text{score} = q_{\\mathrm{abs}} \\cdot c^T \\in [B,H,T,T]", rule: definition}
|
||
- {id: "3", from: "\\text{attn}, c", to: "\\tilde{o}_{\\mathrm{lat}} = \\text{attn} \\cdot c \\in [B,H,T,r]", rule: definition}
|
||
- {id: "4", from: "\\tilde{o}_{\\mathrm{lat}}", to: "\\tilde{o} = \\tilde{o}_{\\mathrm{lat}} \\cdot W_{UV}^T \\in [B,H,T,d_v]", rule: substitute}
|
||
- id: DER3
|
||
claim: C8
|
||
title: "AttnRes 两阶段 online softmax 合并推导"
|
||
expand: true
|
||
figure: null
|
||
steps:
|
||
- {id: "1", from: "s_{l,i} = \\tilde{w}_l^T \\mathrm{RMS}(v_i)", to: "(m, n, d) = (\\max_i s_i, \\sum_i e^{s_i - m} v_i, \\sum_i e^{s_i - m})", rule: definition}
|
||
- {id: "2", from: "inter sources b_0..b_{j-1} 固定", to: "一次批量 einsum 'q d, n b t d -> q n b t' 得块内全部 query 的 (m,n,d)", rule: substitute}
|
||
- {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
|
||
claim: C1
|
||
title: "KDA 递归状态更新张量图"
|
||
grammar: tensor-face
|
||
toolkit: supertensor
|
||
signals: [shape, contraction, broadcast]
|
||
status: planned
|
||
- id: F2
|
||
claim: C4
|
||
title: "MLA 矩阵吸收计算流"
|
||
grammar: tensor-face
|
||
toolkit: supertensor
|
||
signals: [shape, contraction, transpose]
|
||
status: planned
|
||
- id: F3
|
||
claim: C7
|
||
title: "Full vs Block AttnRes 的源列表增长"
|
||
grammar: tensor-face
|
||
toolkit: supertensor
|
||
signals: [shape, contraction]
|
||
status: planned
|