Initial K3 snapshot: 0.5B KDA/MLA/MoE train path
Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
This commit is contained in:
@@ -0,0 +1,187 @@
|
||||
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 时怎样避免参数被重复注册?"
|
||||
|
||||
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
|
||||
|
||||
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}
|
||||
|
||||
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"]}
|
||||
|
||||
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}
|
||||
|
||||
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
|
||||
@@ -0,0 +1,94 @@
|
||||
% KDA 笔记 preamble — 基于 superpaper/assets/notes-macros.tex
|
||||
\usepackage[fontset=fandol]{ctex}
|
||||
\usepackage{amsmath,amssymb}
|
||||
\usepackage{graphicx}
|
||||
\usepackage[margin=2.2cm]{geometry}
|
||||
\usepackage[most]{tcolorbox}
|
||||
\usepackage{etoolbox}
|
||||
\usepackage{listings}
|
||||
\usepackage{booktabs}
|
||||
\usepackage{subcaption}
|
||||
\usepackage{float}
|
||||
\usepackage{tikz}
|
||||
\usepackage{hyperref}
|
||||
\usepackage{xcolor}
|
||||
\usepackage{multicol}
|
||||
\usepackage{tabularx}
|
||||
\usepackage{array}
|
||||
\usepackage{enumitem}
|
||||
|
||||
% ---------- 颜色 ----------
|
||||
\definecolor{codebg}{HTML}{F7F7F7}
|
||||
\definecolor{codeframe}{HTML}{CCCCCC}
|
||||
\definecolor{shapecolor}{HTML}{2E86C1}
|
||||
\definecolor{einsumcolor}{HTML}{884EA0}
|
||||
\definecolor{notegreen}{HTML}{27AE60}
|
||||
\definecolor{warnorange}{HTML}{E67E22}
|
||||
|
||||
% ---------- 代码样式 ----------
|
||||
\lstdefinestyle{pycode}{
|
||||
language=Python,
|
||||
backgroundcolor=\color{codebg},
|
||||
frame=single,
|
||||
rulecolor=\color{codeframe},
|
||||
basicstyle=\ttfamily\small,
|
||||
keywordstyle=\color{blue!70!black}\bfseries,
|
||||
commentstyle=\color{gray},
|
||||
stringstyle=\color{red!60!black},
|
||||
showstringspaces=false,
|
||||
breaklines=true,
|
||||
tabsize=4,
|
||||
columns=flexible,
|
||||
xleftmargin=4pt,
|
||||
xrightmargin=4pt,
|
||||
aboveskip=6pt,
|
||||
belowskip=6pt,
|
||||
}
|
||||
\lstset{style=pycode}
|
||||
|
||||
% ---------- 形状标注命令 ----------
|
||||
\newcommand{\shape}[1]{{\color{shapecolor}\ensuremath{[#1]}}}
|
||||
\newcommand{\einsum}[1]{{\color{einsumcolor}\texttt{einsum(#1)}}}
|
||||
\newcommand{\note}[1]{{\color{notegreen}\textit{#1}}}
|
||||
|
||||
% ---------- 盒子 ----------
|
||||
\newtcolorbox{knowledgebox}[1]{
|
||||
enhanced, colback=blue!5!white, colframe=blue!75!black, colbacktitle=blue!75!black,
|
||||
coltitle=white, fonttitle=\bfseries, title=#1,
|
||||
attach boxed title to top left={yshift=-2mm, xshift=2mm},
|
||||
boxrule=1pt, sharp corners
|
||||
}
|
||||
\newtcolorbox{importantbox}[1]{
|
||||
enhanced, colback=yellow!10!white, colframe=yellow!80!black, colbacktitle=yellow!80!black,
|
||||
coltitle=black, fonttitle=\bfseries, title=#1, sharp corners
|
||||
}
|
||||
\newtcolorbox{warningbox}[1]{
|
||||
enhanced, colback=red!5!white, colframe=red!75!black, colbacktitle=red!75!black,
|
||||
coltitle=white, fonttitle=\bfseries, title=#1, sharp corners
|
||||
}
|
||||
\newtcolorbox{tensorbox}[1]{
|
||||
enhanced, colback=blue!3!white, colframe=blue!40!black, colbacktitle=blue!50!black,
|
||||
coltitle=white, fonttitle=\bfseries, title=#1, sharp corners,
|
||||
boxrule=0.8pt
|
||||
}
|
||||
|
||||
% 代码-公式并行盒子
|
||||
\newtcolorbox{codemathtop}[1]{
|
||||
enhanced, colback=codebg, colframe=codeframe,
|
||||
fonttitle=\bfseries\ttfamily\small, title=#1,
|
||||
sharp corners, boxrule=0.6pt,
|
||||
left=4pt, right=4pt, top=2pt, bottom=2pt,
|
||||
}
|
||||
\newtcolorbox{codemathbot}{
|
||||
enhanced, colback=white, colframe=blue!30!black,
|
||||
sharp corners, boxrule=0.6pt,
|
||||
left=4pt, right=4pt, top=4pt, bottom=4pt,
|
||||
}
|
||||
|
||||
% ---------- 文档元信息 ----------
|
||||
\newcommand{\notetitle}{KDA 训练→推理 完整笔记}
|
||||
\newcommand{\notesubtitle}{递归 · 分块 · Gate · MLA · MoE · AttnRes · K3 架构}
|
||||
\newcommand{\notedate}{\today}
|
||||
|
||||
\newcommand{\splabel}[1]{\hypertarget{sp:#1}{}\label{sp:#1}}
|
||||
\newcommand{\spref}[1]{\hyperlink{sp:#1}{\texttt{#1}}}
|
||||
Binary file not shown.
@@ -0,0 +1,39 @@
|
||||
\documentclass[a4paper,11pt]{article}
|
||||
\input{notes-macros}
|
||||
|
||||
\begin{document}
|
||||
|
||||
% ---------- 封面 ----------
|
||||
\begin{titlepage}
|
||||
\centering
|
||||
\vspace{2cm}
|
||||
{\huge\bfseries \notetitle\par}
|
||||
\vspace{0.6cm}
|
||||
{\Large \notesubtitle\par}
|
||||
\vspace{0.8cm}
|
||||
{\large \notedate\par}
|
||||
\vspace{1.5cm}
|
||||
\begin{tcolorbox}[width=0.88\textwidth, colback=black!2!white, colframe=black!60, sharp corners]
|
||||
\textbf{项目}:\texttt{projects/kda/} — KDA 手写实现(naive recurrent → chunked → Triton)\\
|
||||
\textbf{架构}:KDA + Gated MLA + Stable LatentMoE + AttnRes 深度残差 (K3-like)\\
|
||||
\textbf{参考}:KDA arXiv:2510.26692; AttnRes arXiv:2603.15031; Kimi K3 architecture notes\\
|
||||
\textbf{代码}:\texttt{kda/ops/}, \texttt{kda/layers/}, \texttt{kda/models/}
|
||||
\end{tcolorbox}
|
||||
\end{titlepage}
|
||||
|
||||
\tableofcontents
|
||||
\newpage
|
||||
|
||||
\input{sections/sec-01} % KDA 递归核心
|
||||
\input{sections/sec-02} % Gate 激活
|
||||
\input{sections/sec-03} % 分块并行计算
|
||||
\input{sections/sec-04} % GVA 分组值注意力
|
||||
\input{sections/sec-05} % KDAAttention 层
|
||||
\input{sections/sec-06} % Gated MLA 矩阵吸收版
|
||||
\input{sections/sec-07} % SiTU-GLU 与 Stable LatentMoE
|
||||
\input{sections/sec-08} % K3 混合架构
|
||||
\input{sections/sec-09} % Attention Residual 深度残差
|
||||
\input{sections/sec-10} % 反向传播推导
|
||||
\input{sections/sec-11} % 符号表
|
||||
|
||||
\end{document}
|
||||
@@ -0,0 +1,127 @@
|
||||
% teach:
|
||||
% gap: 读者知道 softmax attention 但不知道线性注意力怎么维护状态矩阵
|
||||
% takeaway: KDA 用 delta rule 逐步更新 [K,V] 状态矩阵, 写入=擦旧写新, 每步 O(KV)
|
||||
% jump: 为什么 r_t = v - k·S 而不是直接用 v?delta rule 的"先擦再写"
|
||||
% omit: KDA 论文的 related work、实验细节
|
||||
|
||||
\section{KDA 递归核心}
|
||||
\splabel{C1}
|
||||
|
||||
\subsection{动机:从 softmax 到状态矩阵}
|
||||
|
||||
标准 attention 每个 token 都要回看所有历史,复杂度 $O(T^2)$。
|
||||
线性注意力换掉 softmax,把 $\sum_j v_j k_j^T$ 压成一个 $K \times V$ 的状态矩阵 $S$,
|
||||
每步只做 $o_t = q_t \cdot S$,复杂度降到 $O(T \cdot K \cdot V)$。
|
||||
|
||||
但裸线性注意力的问题是:$S$ 只能加,不能改。写进去的信息永远在那里。
|
||||
KDA 的核心想法是给 $S$ 加两个操作:\textbf{衰减}(逐渐忘记旧信息)和
|
||||
\textbf{delta rule}(先擦旧的,再写新的)。
|
||||
|
||||
\begin{importantbox}{如果你只记一件事}
|
||||
KDA 的状态更新 = 衰减旧状态 + $\beta_t k_t$ 写入 $(v_t - k_t \cdot S_{\mathrm{dec}})$。
|
||||
减去 $k_t \cdot S_{\mathrm{dec}}$ 就是"先把 $k_t$ 方向的旧预测擦掉"。
|
||||
\end{importantbox}
|
||||
|
||||
\subsection{逐步公式}
|
||||
|
||||
\noindent\textbf{输入张量:}
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{lll}
|
||||
\toprule
|
||||
符号 & 形状 & 含义 \\
|
||||
\midrule
|
||||
$q_t$ & \shape{B, HV, K} & query(已经 repeat\_interleave 到 HV) \\
|
||||
$k_t$ & \shape{B, HV, K} & key(同上) \\
|
||||
$v_t$ & \shape{B, HV, V} & value \\
|
||||
$g_t$ & \shape{B, HV, K} & gate(log-space 衰减率,逐维) \\
|
||||
$\beta_t$ & \shape{B, HV} & 写入强度标量 \\
|
||||
$S_{t-1}$ & \shape{B, HV, K, V} & 上一步的 KV 状态 \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\noindent\textbf{四步更新:}
|
||||
|
||||
\begin{enumerate}[leftmargin=2em]
|
||||
\item \textbf{衰减旧状态}(逐元素,$g_t$ 是 log-space 所以取 exp):
|
||||
\[
|
||||
S_{\mathrm{dec}} = \exp(g_t) \odot S_{t-1}
|
||||
\qquad \shape{B, HV, K, V}
|
||||
\]
|
||||
|
||||
\item \textbf{计算残差}(先用 $k_t$ 查旧状态,得到"旧预测",再减掉):
|
||||
\[
|
||||
p_t = \sum_k k_{t,k} \cdot S_{\mathrm{dec},k,\cdot}
|
||||
= \texttt{einsum('bhk, bhkv -> bhv')}
|
||||
\qquad \shape{B, HV, V}
|
||||
\]
|
||||
\[
|
||||
r_t = v_t - p_t \qquad \shape{B, HV, V}
|
||||
\]
|
||||
|
||||
\item \textbf{写入状态}(外积 rank-1 更新):
|
||||
\[
|
||||
a_t = \beta_t \cdot k_t \qquad \shape{B, HV, K}
|
||||
\]
|
||||
\[
|
||||
S_t = S_{\mathrm{dec}} + a_t \otimes r_t
|
||||
= S_{\mathrm{dec}} + \texttt{einsum('bhk, bhv -> bhkv')}
|
||||
\qquad \shape{B, HV, K, V}
|
||||
\]
|
||||
|
||||
\item \textbf{读出}:
|
||||
\[
|
||||
o_t = \frac{1}{\sqrt{K}} \cdot q_t \cdot S_t
|
||||
= \texttt{einsum('bhk, bhkv -> bhv')}
|
||||
\qquad \shape{B, HV, V}
|
||||
\]
|
||||
\end{enumerate}
|
||||
|
||||
\subsection{代码对照}
|
||||
|
||||
\begin{codemathtop}{ops/reference/recurrent.py — naive\_kda\_fwd (核心循环)}
|
||||
\begin{lstlisting}
|
||||
for t in range(T):
|
||||
q_t = qe[:, t] # [B, HV, K]
|
||||
k_t = ke[:, t] # [B, HV, K]
|
||||
v_t = v[:, t] # [B, HV, V]
|
||||
g_t = g[:, t] # [B, HV, K]
|
||||
b_t = beta[:, t] # [B, HV]
|
||||
|
||||
# Step 1: decay
|
||||
S_dec = S * g_t.exp().unsqueeze(-1) # [B,HV,K,V]
|
||||
|
||||
# Step 2: residual (delta rule)
|
||||
p_t = einsum('bhk, bhkv -> bhv', k_t, S_dec)
|
||||
r_t = v_t - p_t # [B,HV,V]
|
||||
|
||||
# Step 3: write (rank-1 update)
|
||||
a_t = b_t.unsqueeze(-1) * k_t # [B,HV,K]
|
||||
S = S_dec + einsum('bhk, bhv -> bhkv', a_t, r_t)
|
||||
|
||||
# Step 4: read
|
||||
o[:, t] = einsum('bhk, bhkv -> bhv', q_t, S)
|
||||
\end{lstlisting}
|
||||
\end{codemathtop}
|
||||
|
||||
\begin{warningbox}{为什么 exp(g\_t) 要 unsqueeze(-1)?}
|
||||
$g_t$ 的形状是 \shape{B, HV, K},而 $S$ 是 \shape{B, HV, K, V}。
|
||||
衰减是在 $K$ 维上逐元素(同一 $k$ 索引的所有 $v$ 维度共享同一个衰减率),
|
||||
所以 \texttt{exp(g\_t).unsqueeze(-1)} 把 K 维 broadcast 到 $K \times V$。
|
||||
\end{warningbox}
|
||||
|
||||
\subsection{Delta rule 的直觉}
|
||||
|
||||
\begin{knowledgebox}{为什么减去 $k_t \cdot S_{\mathrm{dec}}$?}
|
||||
把 $S$ 想象成一个 $K \to V$ 的线性映射。用 $k_t$ 去查它,得到的 $p_t = k_t^T S$
|
||||
就是``旧状态对 $k_t$ 方向的预测''。如果 $p_t$ 已经很接近 $v_t$,说明这个方向的信息
|
||||
已经写好了,不需要再写。$r_t = v_t - p_t$ 就是``需要修正的量''。
|
||||
|
||||
这就是 Widrow-Hoff delta rule:不是盲目地加,而是只修正误差。
|
||||
\end{knowledgebox}
|
||||
|
||||
\subsection{本章小结}
|
||||
|
||||
KDA 的状态更新 = 衰减旧状态 + $\beta_t k_t$ 写入残差 $(v_t - k_t \cdot S_{\mathrm{dec}})$。
|
||||
每步复杂度 $O(K \cdot V)$(两次矩阵-向量乘 + 一次外积),不需要 softmax。
|
||||
@@ -0,0 +1,111 @@
|
||||
% teach:
|
||||
% gap: 读者知道 g_t 是 gate 但不知道它怎么从 raw projection 变成一个负的 log-space 衰减
|
||||
% takeaway: safe gate 用 sigmoid 把值夹在 [lower_bound, 0], standard gate 用 -softplus 保证负
|
||||
% jump: 论文没解释为什么需要 A_log 和 dt_bias 两层
|
||||
% omit: Triton gate kernel 的 fused 实现细节
|
||||
|
||||
\section{Gate 激活}
|
||||
\splabel{C2}
|
||||
|
||||
\subsection{Gate 的角色}
|
||||
|
||||
回顾 §1:$S_{\mathrm{dec}} = \exp(g_t) \odot S_{t-1}$。$g_t$ 必须 $\leq 0$
|
||||
才是衰减($\exp(g_t) \leq 1$),否则状态会指数增长爆炸。
|
||||
|
||||
\texttt{g\_raw} 是从 \texttt{g\_proj(x)} 出来的 raw 值,没有约束。
|
||||
Gate 激活函数的任务是:把 raw 值映射到一个保证 $\leq 0$ 的范围。
|
||||
|
||||
\subsection{两种 Gate}
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{p{3cm}p{5.5cm}p{5cm}}
|
||||
\toprule
|
||||
& \textbf{Standard gate} & \textbf{Safe gate} \\
|
||||
\midrule
|
||||
公式 &
|
||||
$g = -\mathrm{rate} \cdot \mathrm{softplus}(\mathrm{input})$ &
|
||||
$g = L \cdot \sigma(\mathrm{rate} \cdot \mathrm{input})$ \\
|
||||
值域 &
|
||||
$(-\infty, 0]$ &
|
||||
$[L, 0]$($L$ 是 lower\_bound,如 $-5$) \\
|
||||
衰减范围 &
|
||||
$\exp(g) \in (0, 1]$ &
|
||||
$\exp(g) \in [\exp(L), 1]$ \\
|
||||
稳定性 &
|
||||
衰减可以任意快 &
|
||||
衰减有下限,不会瞬间清零 \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\noindent 其中:
|
||||
\begin{itemize}[nosep]
|
||||
\item $\mathrm{input} = g_{\mathrm{raw}} + \Delta_b$ \quad($\Delta_b$
|
||||
是 \texttt{dt\_bias} \shape{HV, K})
|
||||
\item $\mathrm{rate} = \exp(A_{\log})$ \quad($A_{\log}$ 是
|
||||
\texttt{A\_log} \shape{HV},head-wise 可学习)
|
||||
\end{itemize}
|
||||
|
||||
\begin{importantbox}{如果你只记一件事}
|
||||
Safe gate = $L \cdot \sigma(\mathrm{rate} \cdot \mathrm{input})$,
|
||||
$L=-5$ 时 $\exp(g) \geq \exp(-5) \approx 0.0067$,
|
||||
状态永远不会被``一次性清零''。
|
||||
\end{importantbox}
|
||||
|
||||
\subsection{代码对照}
|
||||
|
||||
\begin{codemathtop}{ops/reference/gate.py — kda\_gate\_reference}
|
||||
\begin{lstlisting}
|
||||
def kda_gate_reference(g, A_log, dt_bias=None, *,
|
||||
safe_gate=False, lower_bound=None):
|
||||
HV, K = g.shape[-2:]
|
||||
gate_input = g if dt_bias is None else g + dt_bias.view(HV, K)
|
||||
rate = A_log.view(HV, 1).exp()
|
||||
|
||||
if safe_gate:
|
||||
# safe: g in [lower_bound, 0]
|
||||
return lower_bound * torch.sigmoid(rate * gate_input)
|
||||
# standard: g in (-inf, 0]
|
||||
return -rate * F.softplus(gate_input)
|
||||
\end{lstlisting}
|
||||
\end{codemathtop}
|
||||
|
||||
\subsection{初始化与默认值}
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{llp{7cm}}
|
||||
\toprule
|
||||
参数 & 初始值 & 效果 \\
|
||||
\midrule
|
||||
\texttt{A\_log} & $\mathbf{0}$ \shape{HV} & $\mathrm{rate} = \exp(0) = 1$,不缩放 \\
|
||||
\texttt{dt\_bias} & $-4.0$ \shape{HV, K} & 初始时 $\mathrm{input} \approx g_{\mathrm{raw}} - 4$,
|
||||
配合 safe gate ($L=-5$) 得到 $g \approx -5 \cdot \sigma(-4) \approx -0.09$,
|
||||
即 $\exp(g) \approx 0.91$(约 91\% 状态保留) \\
|
||||
\texttt{lower\_bound} & $-5.0$ & safe gate 的下限 \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\subsection{Gate 在 API 中的位置}
|
||||
|
||||
Gate 激活在 \texttt{ops/api.py} 的 \texttt{chunk\_kda} 中调用,
|
||||
在进入 chunkwise 或 recurrent 核心之前完成。
|
||||
当 \texttt{use\_gate\_in\_kernel=True} 时,\texttt{g\_raw} 进入 API,
|
||||
API 内部完成 gate 激活;否则调用方自己完成。
|
||||
|
||||
\begin{lstlisting}
|
||||
# ops/api.py (simplified)
|
||||
if use_gate_in_kernel:
|
||||
gate_input = g + dt_bias.view(g.shape[-2:])
|
||||
rate = A_log.exp().view(1, 1, -1, 1)
|
||||
if safe_gate:
|
||||
g = lower_bound * torch.sigmoid(rate * gate_input)
|
||||
else:
|
||||
g = -rate * F.softplus(gate_input)
|
||||
\end{lstlisting}
|
||||
|
||||
\subsection{本章小结}
|
||||
|
||||
Gate 把 raw projection 映射到 $\leq 0$ 的 log-space 衰减率。
|
||||
Safe gate 用 sigmoid 限制在 $[L, 0]$,防止瞬间清零;
|
||||
standard gate 用 softplus 不限制下限。默认配置下初始状态保留约 91\%。
|
||||
@@ -0,0 +1,148 @@
|
||||
% teach:
|
||||
% gap: 读者知道递归形式但不知道怎么在 GPU 上并行, 以为只能一步步算
|
||||
% takeaway: chunk 内用下三角线性系统并行求解, chunk 间递推状态, 数值等价于 naive recurrent
|
||||
% jump: 论文直接写了 triangular solve 但没解释为什么要 solve 而不是直接矩阵乘
|
||||
% omit: Triton 实现细节
|
||||
|
||||
\section{分块并行计算(Chunkwise)}
|
||||
\splabel{C3}
|
||||
|
||||
\subsection{为什么需要分块?}
|
||||
|
||||
Naive recurrent 一步一步算,$T$ 步串行,GPU 利用率低。
|
||||
分块的想法是把序列切成 $T/C$ 个长度为 $C$ 的 chunk:
|
||||
\begin{itemize}[nosep]
|
||||
\item \textbf{chunk 内}:$C$ 个 token 之间的依赖可以用矩阵运算并行处理
|
||||
\item \textbf{chunk 间}:状态 $S$ 从上一个 chunk 传到下一个,仍然是递推
|
||||
\end{itemize}
|
||||
|
||||
\begin{importantbox}{如果你只记一件事}
|
||||
Chunkwise = chunk 内并行 + chunk 间递推。数值结果与 naive recurrent 逐位一致。
|
||||
\end{importantbox}
|
||||
|
||||
\subsection{chunk 内的 cumsum 与下三角解}
|
||||
|
||||
在每个 chunk 内,先对 $g$ 做 cumsum(前缀和),这样衰减就变成了相对距离的函数:
|
||||
\[
|
||||
g_{\mathrm{cum},i} = \sum_{j=0}^{i} g_j, \qquad
|
||||
\text{token } i \text{ 对 token } j \text{ 的衰减} = \exp(g_{\mathrm{cum},i} - g_{\mathrm{cum},j})
|
||||
\]
|
||||
|
||||
定义 \textbf{decayed dot} 矩阵(chunk 内 $C \times C$):
|
||||
\[
|
||||
A_{ij} = \langle x_i, \exp(g_{\mathrm{cum},i} - g_{\mathrm{cum},j}) \cdot k_j \rangle
|
||||
\qquad \shape{..., C, C}
|
||||
\]
|
||||
|
||||
这个矩阵的构造是 chunk 内计算的核心。用它可以构造一个下三角线性系统:
|
||||
\[
|
||||
M = I + \mathrm{tril}(A_{kk} \cdot \beta, \text{diagonal}=-1)
|
||||
\qquad \shape{..., C, C}
|
||||
\]
|
||||
\[
|
||||
M \cdot W = \exp(g_{\mathrm{cum}}) \cdot k \qquad \Rightarrow \qquad
|
||||
W = M^{-1} (\exp(g_{\mathrm{cum}}) \cdot k)
|
||||
\]
|
||||
\[
|
||||
M \cdot U = v \qquad \Rightarrow \qquad U = M^{-1} v
|
||||
\]
|
||||
|
||||
\subsection{代码对照}
|
||||
|
||||
\begin{codemathtop}{ops/reference/chunkwise.py — naive\_chunk\_kda (核心)}
|
||||
\begin{lstlisting}
|
||||
# Rearrange: [B,T,H,K] -> [B,H,N,C,K] where N=T/C
|
||||
q, k = [rearrange(x, 'b (n c) h d -> b h n c d', c=C)
|
||||
.repeat_interleave(HV//H, dim=1) for x in (q, k)]
|
||||
v, g = [rearrange(x, 'b (n c) h d -> b h n c d', c=C)
|
||||
for x in (v, g)]
|
||||
beta = rearrange(beta, 'b (n c) h -> b h n c', c=C)
|
||||
q = q * scale
|
||||
g = g.cumsum(dim=-2) # chunk 内 cumsum
|
||||
|
||||
# Construct triangular system
|
||||
A_kk = _decayed_dot(k, k, g) # [B,HV,N,C,C]
|
||||
M = eye + (A_kk * beta[...,None,:]).masked_fill(mask_upper, 0)
|
||||
W = solve_triangular(M, g.exp() * k, upper=False)
|
||||
U = solve_triangular(M, v, upper=False)
|
||||
|
||||
# A_qk: query 对 key 的 decayed dot (含对角线)
|
||||
A_qk = (_decayed_dot(q, k, g) * beta[...,None,:])
|
||||
.masked_fill(mask_strict_upper, 0)
|
||||
\end{lstlisting}
|
||||
\end{codemathtop}
|
||||
|
||||
\subsection{chunk 间递推}
|
||||
|
||||
每个 chunk 内算完后,用 $W$ 和 $U$ 来处理跨 chunk 的状态:
|
||||
|
||||
\begin{codemathtop}{ops/reference/chunkwise.py — chunk 间循环}
|
||||
\begin{lstlisting}
|
||||
S = zeros(B, HV, K, V) # inter-chunk state
|
||||
for n in range(T // C):
|
||||
# r = "local residual, adjusted by cross-chunk state"
|
||||
r = U[:,:,n] - W[:,:,n] @ S # [B,HV,C,V]
|
||||
|
||||
# output: cross-chunk part + intra-chunk part
|
||||
o[:,:,n] = (q_n * g_n.exp()) @ S + A_qk[:,:,n] @ r
|
||||
|
||||
# update cross-chunk state
|
||||
decay = (g_n[:,:,-1:,:] - g_n).exp() # decay to chunk end
|
||||
S = S * g_n[:,:,-1,:,None].exp() # decay old state
|
||||
S = S + (decay * k_n).T @ (r * beta_n) # write new
|
||||
\end{lstlisting}
|
||||
\end{codemathtop}
|
||||
|
||||
\subsection{形状流水线}
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{lll}
|
||||
\toprule
|
||||
变量 & 形状 & 说明 \\
|
||||
\midrule
|
||||
\texttt{q, k} (chunked) & \shape{B, HV, N, C, K} & $N = T/C$ 个 chunk \\
|
||||
\texttt{v, g} (chunked) & \shape{B, HV, N, C, V/K} & \\
|
||||
\texttt{beta} (chunked) & \shape{B, HV, N, C} & \\
|
||||
\texttt{A\_kk} & \shape{B, HV, N, C, C} & key-key decayed dot \\
|
||||
\texttt{M} & \shape{B, HV, N, C, C} & 下三角系统 \\
|
||||
\texttt{W} & \shape{B, HV, N, C, K} & $M^{-1}(\exp(g) \cdot k)$ \\
|
||||
\texttt{U} & \shape{B, HV, N, C, V} & $M^{-1} v$ \\
|
||||
\texttt{A\_qk} & \shape{B, HV, N, C, C} & query-key decayed dot \\
|
||||
\texttt{S} & \shape{B, HV, K, V} & 跨 chunk 状态 \\
|
||||
\texttt{r} & \shape{B, HV, C, V} & 调整后的残差 \\
|
||||
\texttt{o (chunk n)} & \shape{B, HV, C, V} & 本 chunk 输出 \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\begin{warningbox}{为什么用 triangular solve 而不是直接矩阵乘?}
|
||||
Delta rule 的"先擦再写"引入了 chunk 内 token 之间的递归依赖:
|
||||
token $i$ 的写入依赖 token $j < i$ 的写入结果。
|
||||
这个依赖关系恰好形成一个下三角线性系统 $M \cdot x = b$,
|
||||
用 \texttt{solve\_triangular} 可以在 $O(C^2)$ 内并行求解,
|
||||
而展开递归需要 $O(C)$ 步串行。
|
||||
\end{warningbox}
|
||||
|
||||
\subsection{Decayed dot 函数}
|
||||
|
||||
\begin{codemathtop}{ops/reference/chunkwise.py — \_decayed\_dot}
|
||||
\begin{lstlisting}
|
||||
def _decayed_dot(x, k, g):
|
||||
"""A[..., i, j] = <x_i, exp(g_i - g_j) * k_j>"""
|
||||
C = x.shape[-2]
|
||||
A = empty(*x.shape[:-2], C, C)
|
||||
for i in range(C):
|
||||
decay = (g[..., i:i+1, :] - g).exp() # [.., 1, K] - [.., C, K]
|
||||
A[..., i, :] = einsum('...jk,...jk->...j',
|
||||
x[..., i, None, :] * decay, k)
|
||||
return A
|
||||
\end{lstlisting}
|
||||
\end{codemathtop}
|
||||
|
||||
\noindent 这是一个 $C \times C$ 的矩阵,每个元素 $(i,j)$ 是
|
||||
$x_i$ 和 $\exp(g_i - g_j) \cdot k_j$ 的内积。Triton 实现会把这个双循环融合成一个 kernel。
|
||||
|
||||
\subsection{本章小结}
|
||||
|
||||
分块把 $T$ 步串行拆成 $T/C$ 个 chunk,chunk 内用下三角 solve 并行处理 delta rule 依赖,
|
||||
chunk 间递推状态 $S$。最终输出与 naive recurrent 逐位相同。
|
||||
@@ -0,0 +1,114 @@
|
||||
% teach:
|
||||
% gap: 读者不知道 q/k 和 v 为什么可以有不同的头数, 以及 repeat_interleave 的反传怎么做
|
||||
% takeaway: GVA 让 G 组 value heads 共享一组 q/k, forward repeat_interleave, backward sum
|
||||
% jump: 论文没解释为什么反传是 sum 而不是 mean
|
||||
% omit: GQA 的历史
|
||||
|
||||
\section{GVA(分组值注意力)}
|
||||
\splabel{GVA}
|
||||
|
||||
\subsection{为什么头数不一样?}
|
||||
|
||||
标准 MHA 里 $H_q = H_k = H_v$。GQA(Grouped Query Attention)让多组 q/k 共享同一组 v/k,
|
||||
减少 KV cache。KDA 反过来做:$H$ 组 q/k 对应 $H_V = G \cdot H$ 组 value heads。
|
||||
|
||||
直觉:value 维度决定表达能力,多一点 value head 增加容量;
|
||||
q/k 主要负责路由(``看哪里''),可以共享。
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{lll}
|
||||
\toprule
|
||||
& 标准头数 & GVA \\
|
||||
\midrule
|
||||
$q, k$ & \shape{B, T, H, K} & \shape{B, T, H, K}(不变)\\
|
||||
$v$ & \shape{B, T, H, V} & \shape{B, T, HV, V}($H_V = G \cdot H$) \\
|
||||
$g, \beta$ & \shape{B, T, H, K/1} & \shape{B, T, HV, K/1} \\
|
||||
$S$ & \shape{B, H, K, V} & \shape{B, HV, K, V} \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\subsection{Forward: repeat\_interleave}
|
||||
|
||||
进入 KDA 核心前,$q$ 和 $k$ 从 $H$ 维复制到 $H_V$ 维:
|
||||
|
||||
\begin{lstlisting}
|
||||
G = HV // H
|
||||
qe = q.repeat_interleave(G, dim=2) * scale # [B,T,H,K] -> [B,T,HV,K]
|
||||
ke = k.repeat_interleave(G, dim=2) # [B,T,H,K] -> [B,T,HV,K]
|
||||
\end{lstlisting}
|
||||
|
||||
\noindent 例如 $H=4, G=2, H_V=8$:head 0 的 q/k 复制到 value head 0 和 1,
|
||||
head 1 复制到 value head 2 和 3,依此类推。
|
||||
|
||||
\subsection{Backward: view + sum}
|
||||
|
||||
反传时,$dq_e$ 和 $dk_e$ 的形状是 \shape{B, T, HV, K}(在 $H_V$ 维上计算的梯度)。
|
||||
因为 forward 是复制,反传就是求和:
|
||||
|
||||
\begin{lstlisting}
|
||||
# Backward: HV -> H
|
||||
dq_H = dq_e.view(B, T, H, G, K).sum(dim=3) # [B,T,HV,K] -> [B,T,H,K]
|
||||
dk_H = dk_e.view(B, T, H, G, K).sum(dim=3)
|
||||
\end{lstlisting}
|
||||
|
||||
\begin{warningbox}{为什么是 sum 不是 mean?}
|
||||
\texttt{repeat\_interleave} 是\textbf{复制}:$y_0 = x_0, y_1 = x_0, y_2 = x_1, \ldots$
|
||||
|
||||
对 $x_0$ 的梯度 = $\frac{\partial L}{\partial y_0} + \frac{\partial L}{\partial y_1}$
|
||||
= \textbf{sum}(不是 mean)。
|
||||
|
||||
这和 \texttt{.expand()} 的反传一样:复制的反传是求和。
|
||||
\end{warningbox}
|
||||
|
||||
\subsection{scale 的处理}
|
||||
|
||||
$q$ 在 repeat\_interleave 之后乘了 \texttt{scale = $1/\sqrt{K}$}。
|
||||
反传时 chain rule 要求 $dq_{\mathrm{orig}} = dq_e \cdot \texttt{scale}$:
|
||||
|
||||
\begin{lstlisting}
|
||||
# q 在 forward 内被乘过 scale, chain rule:
|
||||
dq_H = dq_H * scale
|
||||
\end{lstlisting}
|
||||
|
||||
\subsection{KDAAttention 层中的投影}
|
||||
|
||||
\begin{codemathtop}{layers/kda\_attn.py — forward}
|
||||
\begin{lstlisting}
|
||||
def forward(self, x): # x: [B, T, D]
|
||||
B, T, _ = x.shape
|
||||
H, HV, K, V = self.num_heads, self.num_value_heads, ...
|
||||
|
||||
q = self.q_proj(x).view(B, T, H, K) # [B,T,D] -> [B,T,H*K] -> [B,T,H,K]
|
||||
k = self.k_proj(x).view(B, T, H, K) # 同上
|
||||
v = self.v_proj(x).view(B, T, HV, V) # [B,T,D] -> [B,T,HV*V] -> [B,T,HV,V]
|
||||
g_raw = self.g_proj(x).view(B, T, HV, K)
|
||||
beta_raw = self.beta_proj(x).view(B, T, HV)
|
||||
|
||||
o, _ = chunk_kda(q, k, v, g_raw, beta_raw, ...)
|
||||
return self.o_proj(o.reshape(B, T, HV * V)) # [B,T,HV,V] -> [B,T,D]
|
||||
\end{lstlisting}
|
||||
\end{codemathtop}
|
||||
|
||||
\subsection{投影矩阵形状总览}
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{llll}
|
||||
\toprule
|
||||
投影 & 权重形状 & 输入 & 输出 \\
|
||||
\midrule
|
||||
\texttt{q\_proj} & \shape{H \cdot K, D} & \shape{B,T,D} & \shape{B,T,H,K} \\
|
||||
\texttt{k\_proj} & \shape{H \cdot K, D} & \shape{B,T,D} & \shape{B,T,H,K} \\
|
||||
\texttt{v\_proj} & \shape{HV \cdot V, D} & \shape{B,T,D} & \shape{B,T,HV,V} \\
|
||||
\texttt{g\_proj} & \shape{HV \cdot K, D} & \shape{B,T,D} & \shape{B,T,HV,K} \\
|
||||
\texttt{beta\_proj} & \shape{HV, D} & \shape{B,T,D} & \shape{B,T,HV} \\
|
||||
\texttt{o\_proj} & \shape{D, HV \cdot V} & \shape{B,T,HV \cdot V} & \shape{B,T,D} \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\subsection{本章小结}
|
||||
|
||||
GVA 让 $H_V = G \cdot H$ 组 value heads 共享 $H$ 组 q/k。
|
||||
Forward 用 \texttt{repeat\_interleave} 复制,backward 用 \texttt{view+sum} 归约。
|
||||
$v, g, \beta$ 直接在 $H_V$ 维投影,q/k 在 $H$ 维投影。
|
||||
@@ -0,0 +1,75 @@
|
||||
% teach:
|
||||
% gap: 读者已知各组件, 但不清楚它们怎么黏在一起成为一个层
|
||||
% takeaway: KDAAttention = 投影 → gate+norm → chunk_kda → output 投影, 整个层就是 x → y [B,T,D]
|
||||
% jump: none
|
||||
% omit: from_config 工厂方法细节
|
||||
|
||||
\section{KDAAttention 层}
|
||||
|
||||
\subsection{完整数据流}
|
||||
|
||||
\texttt{KDAAttention} 把投影、gate 激活、KDA 核心计算和输出投影封装成一个
|
||||
\texttt{[B,T,D] $\to$ [B,T,D]} 的模块。
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{rlll}
|
||||
\toprule
|
||||
步骤 & 操作 & 输入形状 & 输出形状 \\
|
||||
\midrule
|
||||
1 & \texttt{q\_proj(x)} & \shape{B,T,D} & \shape{B,T,H,K} \\
|
||||
2 & \texttt{k\_proj(x)} & \shape{B,T,D} & \shape{B,T,H,K} \\
|
||||
3 & \texttt{v\_proj(x)} & \shape{B,T,D} & \shape{B,T,HV,V} \\
|
||||
4 & \texttt{g\_proj(x)} & \shape{B,T,D} & \shape{B,T,HV,K} \\
|
||||
5 & \texttt{beta\_proj(x)} & \shape{B,T,D} & \shape{B,T,HV} \\
|
||||
6 & \texttt{chunk\_kda(...)} & 上述 5 项 + 参数 & \shape{B,T,HV,V} \\
|
||||
7 & \texttt{o.reshape(...)} & \shape{B,T,HV,V} & \shape{B,T,HV \cdot V} \\
|
||||
8 & \texttt{o\_proj(...)} & \shape{B,T,HV \cdot V} & \shape{B,T,D} \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\subsection{chunk\_kda 内部做了什么}
|
||||
|
||||
\texttt{chunk\_kda}(\texttt{ops/api.py})是统一入口,按 \texttt{backend} 分发:
|
||||
|
||||
\begin{enumerate}[nosep]
|
||||
\item 如果 \texttt{use\_qk\_l2norm\_in\_kernel}:$q, k \leftarrow \text{L2-normalize}(q), \text{L2-normalize}(k)$
|
||||
\item 如果 \texttt{use\_beta\_sigmoid\_in\_kernel}:$\beta \leftarrow \sigma(\beta_{\mathrm{raw}})$
|
||||
\item 如果 \texttt{use\_gate\_in\_kernel}:应用 gate 激活(§2)
|
||||
\item 调用 \texttt{naive\_chunk\_kda}(或 triton/fla 版本)
|
||||
\end{enumerate}
|
||||
|
||||
\begin{knowledgebox}{三个 ``in\_kernel'' 开关}
|
||||
\begin{itemize}[nosep]
|
||||
\item \texttt{use\_qk\_l2norm}:L2-norm 让 $\langle q, k \rangle$ 变成余弦相似度,
|
||||
稳定训练
|
||||
\item \texttt{use\_beta\_sigmoid}:sigmoid 把 $\beta$ 限制在 $(0,1)$,
|
||||
控制写入强度
|
||||
\item \texttt{use\_gate\_in\_kernel}:gate 激活在 API 内部完成(vs 调用方自己做)
|
||||
\end{itemize}
|
||||
默认三个都是 \texttt{True}。
|
||||
\end{knowledgebox}
|
||||
|
||||
\subsection{可学习参数清单}
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{lll}
|
||||
\toprule
|
||||
参数 & 形状 & 说明 \\
|
||||
\midrule
|
||||
\texttt{q\_proj.weight} & \shape{H \cdot K, D} & query 投影 \\
|
||||
\texttt{k\_proj.weight} & \shape{H \cdot K, D} & key 投影 \\
|
||||
\texttt{v\_proj.weight} & \shape{HV \cdot V, D} & value 投影 \\
|
||||
\texttt{g\_proj.weight} & \shape{HV \cdot K, D} & gate 投影 \\
|
||||
\texttt{beta\_proj.weight} & \shape{HV, D} & beta 投影 \\
|
||||
\texttt{o\_proj.weight} & \shape{D, HV \cdot V} & 输出投影 \\
|
||||
\texttt{A\_log} & \shape{HV} & head-wise 衰减率(log-space)\\
|
||||
\texttt{dt\_bias} & \shape{HV, K} & per-dim gate bias \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\subsection{本章小结}
|
||||
|
||||
KDAAttention 是一个完整的 mixing 模块:5 个线性投影 + gate 激活 + KDA 核心 + 输出投影。
|
||||
三个 ``in\_kernel'' 开关控制 L2-norm、sigmoid、gate 是否在 API 内部完成。
|
||||
@@ -0,0 +1,173 @@
|
||||
% teach:
|
||||
% gap: 读者知道标准 MHA 但不知道 MLA 怎么压缩 KV、矩阵吸收怎么避免解压
|
||||
% takeaway: MLA 把 KV 压成低秩 latent c, 通过吸收 W_UK 进 q 直接在 latent 空间算 attention
|
||||
% jump: 为什么可以先在 latent 加权再乘 W_UV?因为矩阵乘和加权求和可交换
|
||||
% omit: RoPE (K3 用 NoPE)
|
||||
|
||||
\section{Gated MLA(矩阵吸收版)}
|
||||
\splabel{C4}
|
||||
|
||||
\subsection{标准 MHA 的 KV cache 问题}
|
||||
|
||||
标准 MHA 推理时需要缓存所有历史 token 的 $K, V$,cache 大小 $\propto T \cdot H \cdot d$。
|
||||
MLA 的想法:把 $K, V$ 压缩成一个低秩 latent $c$,cache 大小 $\propto T \cdot r$,
|
||||
其中 $r \ll H \cdot d$。
|
||||
|
||||
\subsection{低秩压缩}
|
||||
|
||||
\[
|
||||
c = \mathrm{RMSNorm}(W_{\downarrow} \cdot x) \qquad \shape{B, T, r}
|
||||
\]
|
||||
|
||||
推理时只缓存 $c$,不缓存解压后的 $K, V$。
|
||||
解压矩阵 $W_{\mathrm{KV}\uparrow}$ 包含两部分:
|
||||
\[
|
||||
W_{\mathrm{KV}\uparrow} = \begin{bmatrix} W_{UK} \\ W_{UV} \end{bmatrix}
|
||||
\qquad \shape{H \cdot (d_q + d_v), r}
|
||||
\]
|
||||
拆开:$W_{UK} \in \mathbb{R}^{H \times d_q \times r}$(key 解压),
|
||||
$W_{UV} \in \mathbb{R}^{H \times d_v \times r}$(value 解压)。
|
||||
|
||||
\subsection{矩阵吸收的核心思路}
|
||||
|
||||
\textbf{不解压} $K$ 和 $V$。标准做法会先解压再算 attention:
|
||||
|
||||
\begin{center}
|
||||
\textit{标准}:$k_h = c \cdot W_{UK,h}^T$ \shape{B,T,d_q},
|
||||
$\mathrm{score} = q_h \cdot k_h^T$
|
||||
\end{center}
|
||||
|
||||
矩阵吸收反过来:把 $W_{UK}$ 吸收进 $q$:
|
||||
|
||||
\begin{center}
|
||||
\textit{吸收}:$q_{\mathrm{abs},h} = q_h \cdot W_{UK,h}$ \shape{B,T,r},
|
||||
$\mathrm{score} = q_{\mathrm{abs},h} \cdot c^T$
|
||||
\end{center}
|
||||
|
||||
\begin{importantbox}{如果你只记一件事}
|
||||
$(q \cdot W_{UK}^T) \cdot c^T = q \cdot (W_{UK}^T \cdot c^T) = q_{\mathrm{abs}} \cdot c^T$
|
||||
|
||||
吸收后,attention 直接在 latent 空间 $r$ 维上算,永不解压到 $H \cdot d_q$ 维。
|
||||
\end{importantbox}
|
||||
|
||||
\subsection{完整计算流(四步)}
|
||||
|
||||
\begin{enumerate}[leftmargin=2em]
|
||||
\item \textbf{Q 低秩路径}(NoPE,只有 nope 段):
|
||||
\[
|
||||
q = W_{q\uparrow} \cdot \mathrm{RMSNorm}(W_{q\downarrow} \cdot x)
|
||||
\qquad \shape{B, T, H, d_q}
|
||||
\]
|
||||
|
||||
\item \textbf{吸收 $W_{UK}$ + 打分}:
|
||||
\[
|
||||
q_{\mathrm{abs}} = q \cdot W_{UK} \quad
|
||||
\xrightarrow{\texttt{einsum('bthd,hdj->bthj')}} \quad \shape{B, T, H, r}
|
||||
\]
|
||||
\[
|
||||
\mathrm{score} = q_{\mathrm{abs}} \cdot c^T \quad
|
||||
\xrightarrow{\texttt{einsum('bthj,bsj->bhts')}} \quad \shape{B, H, T, T}
|
||||
\]
|
||||
\[
|
||||
\mathrm{attn} = \mathrm{softmax}(\mathrm{causal\_mask}(\mathrm{score}))
|
||||
\qquad \shape{B, H, T, T}
|
||||
\]
|
||||
|
||||
\item \textbf{先在 latent 加权,再乘 $W_{UV}^T$}:
|
||||
\[
|
||||
\tilde{o}_{\mathrm{lat}} = \mathrm{attn} \cdot c \quad
|
||||
\xrightarrow{\texttt{einsum('bhts,bsj->bhtj')}} \quad \shape{B, H, T, r}
|
||||
\]
|
||||
\[
|
||||
\tilde{o} = \tilde{o}_{\mathrm{lat}} \cdot W_{UV}^T \quad
|
||||
\xrightarrow{\texttt{einsum('bhtj,hvj->bhtv')}} \quad \shape{B, H, T, d_v}
|
||||
\]
|
||||
|
||||
\item \textbf{输出门 + 投影}:
|
||||
\[
|
||||
y = W_o \big[ \sigma(W_g \cdot x) \odot \tilde{o}_{\mathrm{flat}} \big]
|
||||
\qquad \shape{B, T, D}
|
||||
\]
|
||||
\end{enumerate}
|
||||
|
||||
\subsection{代码对照}
|
||||
|
||||
\begin{codemathtop}{layers/mla.py — GatedMLA.forward}
|
||||
\begin{lstlisting}
|
||||
def forward(self, x): # x: [B, T, D]
|
||||
B, T, _ = x.shape
|
||||
H, r = self.num_heads, self.kv_up.in_features
|
||||
|
||||
# Step 1: latent + query
|
||||
c = self.kv_norm(self.kv_down(x)) # [B, T, r]
|
||||
q = self.q_up(self.q_norm(self.q_down(x))) # [B, T, H*d_q]
|
||||
q = q.view(B, T, H, self.qk_nope_head_dim) # [B, T, H, d_q]
|
||||
|
||||
# Split W_UK, W_UV from kv_up.weight
|
||||
w = self.kv_up.weight # [H*(d_q+d_v), r]
|
||||
w_uk = w[:H*d_q].view(H, d_q, r) # [H, d_q, r]
|
||||
w_uv = w[H*d_q:].view(H, d_v, r) # [H, d_v, r]
|
||||
|
||||
# Step 2: absorb W_UK, score
|
||||
q_absorb = einsum('bthd,hdj->bthj', q, w_uk) # [B,T,H,r]
|
||||
scores = einsum('bthj,bsj->bhts', q_absorb, c) # [B,H,T,T]
|
||||
scores = scores.masked_fill(causal_mask, -inf)
|
||||
attn = softmax(scores, dim=-1) # [B,H,T,T]
|
||||
|
||||
# Step 3: latent-space weighted sum, then W_UV
|
||||
latent_out = einsum('bhts,bsj->bhtj', attn, c) # [B,H,T,r]
|
||||
o_heads = einsum('bhtj,hvj->bhtv', latent_out, w_uv) # [B,H,T,d_v]
|
||||
|
||||
# Step 4: output gate
|
||||
o_heads = o_heads.transpose(1,2).reshape(B,T, H*d_v)
|
||||
gate = sigmoid(self.gate(x)) # [B,T,H*d_v]
|
||||
return self.o_proj(gate * o_heads) # [B,T,D]
|
||||
\end{lstlisting}
|
||||
\end{codemathtop}
|
||||
|
||||
\begin{warningbox}{为什么可以先加权再乘 $W_{UV}$?}
|
||||
标准做法:$o = \mathrm{attn} \cdot V = \mathrm{attn} \cdot (c \cdot W_{UV}^T)$
|
||||
|
||||
交换顺序:$o = (\mathrm{attn} \cdot c) \cdot W_{UV}^T$
|
||||
|
||||
这能成立是因为矩阵乘法的结合律:$A(BC) = (AB)C$。
|
||||
$\mathrm{attn} \cdot c$ 先在 latent 空间 $r$ 维上加权求和,
|
||||
得到的 \shape{B,H,T,r} 再乘 $W_{UV}^T$ 还原到 $d_v$ 维。
|
||||
全程不需要显式构造 $H \cdot T$ 大小的 $V$ 矩阵。
|
||||
\end{warningbox}
|
||||
|
||||
\subsection{形状与参数对比}
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{lll}
|
||||
\toprule
|
||||
参数 & 形状 & 说明 \\
|
||||
\midrule
|
||||
\texttt{kv\_down.weight} & \shape{r, D} & KV latent 压缩 \\
|
||||
\texttt{kv\_up.weight} & \shape{H \cdot (d_q+d_v), r} & 包含 $W_{UK}$ 和 $W_{UV}$ \\
|
||||
\texttt{q\_down.weight} & \shape{r_q, D} & Q 低秩 \\
|
||||
\texttt{q\_up.weight} & \shape{H \cdot d_q, r_q} & Q 解压 \\
|
||||
\texttt{gate.weight} & \shape{H \cdot d_v, D} & 输出门 \\
|
||||
\texttt{o\_proj.weight} & \shape{D, H \cdot d_v} & 输出投影 \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\noindent KV cache 大小对比(推理时):
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{ll}
|
||||
\toprule
|
||||
方法 & Cache 大小 per token \\
|
||||
\midrule
|
||||
标准 MHA & $2 \times H \times d = 2 H d$ \\
|
||||
MLA (latent) & $r$(只存 $c$) \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\subsection{本章小结}
|
||||
|
||||
Gated MLA 把 KV 压缩到低秩 latent $c$ \shape{B,T,r},通过矩阵吸收
|
||||
($q_{\mathrm{abs}} = q \cdot W_{UK}$)直接在 latent 空间打分和加权,
|
||||
永不解压 K/V。输出通过 sigmoid 门控。NoPE:不使用 RoPE,位置感交给夹层 KDA 的 decay/gate。
|
||||
@@ -0,0 +1,158 @@
|
||||
% teach:
|
||||
% gap: 读者知道 MoE 的 top-k 路由但不知道 LatentMoE 的 latent 接口和 SiTU-GLU
|
||||
% takeaway: LatentMoE 通过 latent 接口把 routed 专家限制在 ℓ=d/2 上算, SiTU-GLU 用软上限防溢出
|
||||
% jump: 为什么 routed 专家在 latent 空间而 shared 在全宽?省参数
|
||||
% omit: load balancing loss
|
||||
|
||||
\section{SiTU-GLU 与 Stable LatentMoE}
|
||||
\splabel{C5}
|
||||
|
||||
\subsection{SiTU-GLU:带软上限的激活}
|
||||
|
||||
SwiGLU 在低精度(fp16/bf16)训练时可能溢出:$\mathrm{silu}(x) \cdot x$ 没有上限。
|
||||
SiTU-GLU 用 $\tanh$ 给门控和上投影加软上限:
|
||||
|
||||
\[
|
||||
\mathrm{SiTU}(x) = W_o \big[\underbrace{\beta_1 \tanh\!\left(\frac{W_g x}{\beta_1}\right) \cdot \sigma(W_g x)}_{\text{gate}} \;\cdot\; \underbrace{\beta_2 \tanh\!\left(\frac{W_u x}{\beta_2}\right)}_{\text{up}}\big]
|
||||
\]
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{lp{8cm}}
|
||||
\toprule
|
||||
性质 & 说明 \\
|
||||
\midrule
|
||||
输出上限 & $\|\mathrm{SiTU}\|_\infty \leq \beta_1 \cdot \beta_2 = 4 \times 25 = 100$ \\
|
||||
原点附近 & $\tanh(x/\beta) \approx x/\beta$,所以 $\beta \cdot \tanh(x/\beta) \approx x$,退化为 SwiGLU \\
|
||||
远端 & 软饱和,防 fp16 溢出 \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\begin{codemathtop}{layers/latent\_moe.py — SiTU}
|
||||
\begin{lstlisting}
|
||||
class SiTU(nn.Module):
|
||||
def __init__(self, dim_in, dim_ff, beta1=4.0, beta2=25.0):
|
||||
self.w_g = nn.Linear(dim_in, dim_ff, bias=False)
|
||||
self.w_u = nn.Linear(dim_in, dim_ff, bias=False)
|
||||
self.w_o = nn.Linear(dim_ff, dim_in, bias=False)
|
||||
|
||||
def forward(self, x): # [*, dim_in]
|
||||
wg = self.w_g(x)
|
||||
g = self.beta1 * tanh(wg / self.beta1) * sigmoid(wg) # gate
|
||||
u = self.beta2 * tanh(self.w_u(x) / self.beta2) # up
|
||||
return self.w_o(g * u) # [*, dim_in]
|
||||
\end{lstlisting}
|
||||
\end{codemathtop}
|
||||
|
||||
\subsection{LatentMoE 架构}
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{rl}
|
||||
\toprule
|
||||
组件 & 说明 \\
|
||||
\midrule
|
||||
\textbf{Shared 专家} & $n_{\mathrm{shared}}$ 个 SiTU,全宽 $d \to d$,所有 token 都经过 \\
|
||||
\textbf{Routed 专家} & $n_{\mathrm{routed}}$ 个 SiTU,半宽 $\ell \to \ell$($\ell = d/2$) \\
|
||||
\textbf{Latent 接口} & $W_\downarrow: d \to \ell$, $W_\uparrow: \ell \to d$(压缩/还原) \\
|
||||
\textbf{Router} & $W_r: d \to n_{\mathrm{routed}}$,Top-k 选择 + softmax 归一化 \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\subsection{计算流(五步)}
|
||||
|
||||
\begin{enumerate}[leftmargin=2em]
|
||||
\item \textbf{Latent 压缩}:
|
||||
\[
|
||||
z = W_\downarrow \cdot x \qquad \shape{B, T, \ell}
|
||||
\]
|
||||
|
||||
\item \textbf{Routing}:
|
||||
\[
|
||||
\mathrm{logits} = W_r \cdot x \qquad \shape{B, T, n_{\mathrm{routed}}}
|
||||
\]
|
||||
\[
|
||||
\mathrm{ids}, \mathrm{probs} = \mathrm{TopK}(\mathrm{logits}, k)
|
||||
\qquad \mathrm{ids}: \shape{B, T, k}, \;\; \mathrm{probs}: \shape{B, T, k}
|
||||
\]
|
||||
|
||||
\item \textbf{Routed 专家}(在 latent 空间 $\ell$ 上):
|
||||
\[
|
||||
u = \sum_{i \in \mathrm{Top\text{-}k}} p_i \cdot E_i^{\mathrm{rt}}(z)
|
||||
\qquad \shape{B, T, \ell}
|
||||
\]
|
||||
|
||||
\item \textbf{Shared 专家}(全宽 $d$):
|
||||
\[
|
||||
s = \sum_j E_j^{\mathrm{sh}}(x) \qquad \shape{B, T, d}
|
||||
\]
|
||||
|
||||
\item \textbf{合并}:
|
||||
\[
|
||||
y = s + W_\uparrow \cdot \mathrm{RMSNorm}(u) \qquad \shape{B, T, d}
|
||||
\]
|
||||
\end{enumerate}
|
||||
|
||||
\begin{importantbox}{如果你只记一件事}
|
||||
Routed 专家只在 $\ell = d/2$ 的 latent 空间操作,
|
||||
参数量是全宽专家的 $1/4$($\ell^2$ vs $d^2$)。
|
||||
Shared 专家保持全宽 $d$,提供基础表达能力。
|
||||
\end{importantbox}
|
||||
|
||||
\subsection{代码对照}
|
||||
|
||||
\begin{codemathtop}{layers/latent\_moe.py — LatentMoE.forward}
|
||||
\begin{lstlisting}
|
||||
def forward(self, x): # [B, T, d]
|
||||
z = self.down(x) # [B, T, ell]
|
||||
|
||||
logits = self.router(x) # [B, T, n_routed]
|
||||
topk = torch.topk(logits, self.top_k, dim=-1)
|
||||
ids = topk.indices # [B, T, k]
|
||||
probs = F.softmax(topk.values, dim=-1) # [B, T, k]
|
||||
|
||||
# All expert outputs (vectorized)
|
||||
all_out = stack([e(z) for e in self.experts]) # [R, B, T, ell]
|
||||
# Gather top-k and weighted sum
|
||||
u = zeros(B, T, ell)
|
||||
for i in range(self.top_k):
|
||||
idx = ids[:,:,i].reshape(B*T)
|
||||
sel = all_out[arange, idx]
|
||||
u += probs[:,:,i:i+1] * sel.reshape(B, T, ell)
|
||||
|
||||
shared_out = stack([e(x) for e in self.shared]).sum(0) # [B, T, d]
|
||||
return shared_out + self.up(self.norm(u)) # [B, T, d]
|
||||
\end{lstlisting}
|
||||
\end{codemathtop}
|
||||
|
||||
\begin{warningbox}{为什么 router 用 $x$(全宽)而不是 $z$(latent)?}
|
||||
路由需要看到 token 的完整表示才能做好选择。
|
||||
如果用 $z$ 路由,压缩过程可能丢失路由需要的信息。
|
||||
K3 论文里也是用全宽 $x$ 做 Top-k,然后在 $\ell$ 空间计算。
|
||||
\end{warningbox}
|
||||
|
||||
\subsection{形状总览}
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{llll}
|
||||
\toprule
|
||||
变量 & 形状 & 说明 \\
|
||||
\midrule
|
||||
$x$ & \shape{B, T, d} & 输入 \\
|
||||
$z$ & \shape{B, T, \ell} & latent($\ell = d/2$)\\
|
||||
logits & \shape{B, T, n_r} & router 输出 \\
|
||||
ids & \shape{B, T, k} & Top-k 专家索引 \\
|
||||
probs & \shape{B, T, k} & Top-k softmax 权重 \\
|
||||
\texttt{all\_out} & \shape{n_r, B, T, \ell} & 所有 routed 专家输出 \\
|
||||
$u$ & \shape{B, T, \ell} & 加权求和后的 routed 输出 \\
|
||||
\texttt{shared\_out} & \shape{B, T, d} & shared 专家求和 \\
|
||||
$y$ & \shape{B, T, d} & 最终输出 \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\subsection{本章小结}
|
||||
|
||||
LatentMoE 把 routed 专家限制在 $\ell = d/2$ 的 latent 空间,省参数。
|
||||
SiTU-GLU 给 gate 和 up 加 $\tanh$ 软上限($\beta_1=4, \beta_2=25$),
|
||||
防止低精度溢出。Shared 专家全宽,提供基础能力;routed 专家通过 Top-k 路由提供专业化能力。
|
||||
@@ -0,0 +1,162 @@
|
||||
% teach:
|
||||
% gap: 读者已知各组件但不知道怎么组装成完整模型
|
||||
% takeaway: K3 = Hybrid(3 KDA + 1 MLA) × DecoderBlock(attn + MoE), 末层强制 MLA
|
||||
% jump: 为什么每 4 层才放一次 MLA?位置感知只需要周期性提供
|
||||
% omit: 0.5b preset 的训练超参
|
||||
|
||||
\section{K3 混合架构}
|
||||
|
||||
\subsection{整体结构}
|
||||
|
||||
\begin{center}
|
||||
\texttt{Embedding} $\to$ \texttt{DecoderBlock} $\times L$ $\to$ \texttt{RMSNorm} $\to$ \texttt{LM Head}
|
||||
\end{center}
|
||||
|
||||
每个 \texttt{DecoderBlock} 是 Pre-Norm 残差:
|
||||
|
||||
\begin{lstlisting}
|
||||
def forward(self, x):
|
||||
x = x + self.attn(self.attn_norm(x)) # mixing
|
||||
return x + self.ffn(self.ffn_norm(x)) # channel
|
||||
\end{lstlisting}
|
||||
|
||||
\subsection{Hybrid Attention Pattern}
|
||||
|
||||
K3 用两种 attention 层交替:
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{ccccccccc}
|
||||
\toprule
|
||||
层 & 0 & 1 & 2 & 3 & 4 & 5 & 6 & 7 \\
|
||||
\midrule
|
||||
Attn & KDA & KDA & KDA & \textbf{MLA} & KDA & KDA & KDA & \textbf{MLA} \\
|
||||
FFN & MoE & MoE & MoE & MoE & MoE & MoE & MoE & MoE \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\noindent 规则:每 4 层放 1 次 MLA(0-based 层 3, 7, 11, ...),\textbf{末层强制 MLA}。
|
||||
|
||||
\begin{codemathtop}{models/k3\_config.py — layer\_types}
|
||||
\begin{lstlisting}
|
||||
def layer_types(self) -> list[str]:
|
||||
"""Hybrid: 3 KDA + 1 MLA per group, last always MLA."""
|
||||
types = ["kda"] * self.num_hidden_layers
|
||||
for i in range(self.num_hidden_layers):
|
||||
if i % 4 == 3:
|
||||
types[i] = "mla"
|
||||
types[-1] = "mla" # last layer forced
|
||||
return types
|
||||
|
||||
def layer_specs(self):
|
||||
return [(kind, "moe") for kind in self.layer_types()]
|
||||
\end{lstlisting}
|
||||
\end{codemathtop}
|
||||
|
||||
\begin{knowledgebox}{为什么 KDA 不需要 RoPE?}
|
||||
KDA 的 gate/decay 机制天然提供位置感知:
|
||||
远的 token 衰减更多,近的保留更多。
|
||||
但 softmax attention(MLA)没有这个机制,所以真实 K3 用 NoPE
|
||||
(本复现的 MLA 也是 NoPE)。
|
||||
位置感知从 KDA 层``渗透''到 MLA 层——3:1 的比例足够了。
|
||||
\end{knowledgebox}
|
||||
|
||||
\subsection{CausalLM 完整数据流}
|
||||
|
||||
\begin{codemathtop}{models/causal\_lm.py — CausalLM}
|
||||
\begin{lstlisting}
|
||||
class CausalLM(nn.Module):
|
||||
def __init__(self, config):
|
||||
self.embedding = nn.Embedding(vocab_size, D) # [V, D]
|
||||
self.blocks = ModuleList([
|
||||
DecoderBlock.from_spec(config, attn, ffn)
|
||||
for attn, ffn in config.layer_specs()
|
||||
])
|
||||
self.norm = RMSNorm(D)
|
||||
self.lm_head = nn.Linear(D, vocab_size) # [V, D]
|
||||
if config.tie_word_embeddings:
|
||||
self.lm_head.weight = self.embedding.weight
|
||||
|
||||
def forward(self, input_ids, labels=None):
|
||||
x = self.embedding(input_ids) # [B,T] -> [B,T,D]
|
||||
if self.mixer is None: # attnres="off"
|
||||
for block in self.blocks:
|
||||
x = block(x) # [B,T,D] -> [B,T,D]
|
||||
else:
|
||||
x = self.mixer(x) # AttnRes 深度残差, 见 §9
|
||||
logits = self.lm_head(self.norm(x)) # [B,T,D] -> [B,T,V]
|
||||
if labels is None:
|
||||
return logits
|
||||
# Shifted CE: predict next token
|
||||
return cross_entropy(logits[:,:-1], labels[:,1:])
|
||||
\end{lstlisting}
|
||||
\end{codemathtop}
|
||||
|
||||
\begin{knowledgebox}{残差流是可替换的}
|
||||
上面的 \texttt{DecoderBlock} 逐层堆叠(\texttt{x = x + sublayer(norm(x))})
|
||||
是 \texttt{config.attnres="off"} 时的默认路径。
|
||||
置为 \texttt{"full"} / \texttt{"block"} 时,\texttt{CausalLM} 会把每个 block
|
||||
拆成 attn / ffn 两个原子子层交给 \texttt{mixer},用\textbf{深度维注意力}
|
||||
代替等权残差加法——见 \S9。真实 K3 用的是 \texttt{block} 模式。
|
||||
\end{knowledgebox}
|
||||
|
||||
\subsection{两种配置}
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{lll}
|
||||
\toprule
|
||||
& \textbf{KDAConfig}(纯 KDA) & \textbf{K3Config}(混合) \\
|
||||
\midrule
|
||||
Attn & KDA only & 3 KDA + 1 MLA \\
|
||||
FFN & SwiGLU & LatentMoE \\
|
||||
典型规模 & \textasciitilde8M (toy) & \textasciitilde8M (toy) / \textasciitilde500M (0.5b) \\
|
||||
\texttt{layer\_specs()} & \texttt{[("kda","swiglu")] * L} & \texttt{[(kind,"moe") for kind in ...]} \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\subsection{K3 toy 尺寸}
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{llll}
|
||||
\toprule
|
||||
参数 & 真实 K3 & toy 复现 & 缩比 \\
|
||||
\midrule
|
||||
$D$ & 7168 & 256 & 28$\times$ \\
|
||||
$L$ & 93 & 4 & 23$\times$ \\
|
||||
$H = H_V$ & 96 & 8 & 12$\times$ \\
|
||||
$K = V$ & 128 & 16 & 8$\times$ \\
|
||||
kv\_lora\_rank & 512 & 32 & 16$\times$ \\
|
||||
q\_lora\_rank & 1536 & 64 & 24$\times$ \\
|
||||
$\ell$ (MoE latent) & 3584 & 128 & 28$\times$ \\
|
||||
$n_{\mathrm{routed}}$ / Top-$k$ & 896/16 & 16/2 & 56$\times$ / 8$\times$ \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\subsection{DecoderBlock 构建}
|
||||
|
||||
\begin{codemathtop}{layers/block.py — build\_attn / build\_ffn}
|
||||
\begin{lstlisting}
|
||||
def build_attn(config, kind: str) -> nn.Module:
|
||||
if kind == "kda": return KDAAttention.from_config(config)
|
||||
if kind == "mla": return GatedMLA.from_config(config)
|
||||
|
||||
def build_ffn(config, kind: str) -> nn.Module:
|
||||
if kind == "swiglu": return SwiGLUMLP.from_config(config)
|
||||
if kind == "moe": return LatentMoE.from_config(config)
|
||||
|
||||
class DecoderBlock(nn.Module):
|
||||
def forward(self, x):
|
||||
x = x + self.attn(self.attn_norm(x))
|
||||
return x + self.ffn(self.ffn_norm(x))
|
||||
\end{lstlisting}
|
||||
\end{codemathtop}
|
||||
|
||||
\subsection{本章小结}
|
||||
|
||||
K3 架构 = Hybrid Attention(3 KDA + 1 MLA,末层强制 MLA)+ LatentMoE。
|
||||
KDA 层提供线性复杂度的序列混合和位置感知(通过 decay),
|
||||
MLA 层提供全局 softmax attention(NoPE,利用 KDA 渗透的位置信息)。
|
||||
每层默认是 Pre-Norm 残差 DecoderBlock;\texttt{config.attnres} 可以把这条
|
||||
等权残差流换成 AttnRes 深度注意力(\S9)。
|
||||
@@ -0,0 +1,379 @@
|
||||
% teach:
|
||||
% gap: 读者知道 Pre-Norm 残差是"无条件等权累加", 但不知道怎么把它换成"按内容选择读哪一层"
|
||||
% takeaway: AttnRes = 深度维 softmax 注意力残差; Block 版把 O(N^2) 源数压到 O(N/S); 两阶段 = inter 批量 + intra online-softmax 合并
|
||||
% jump: 为什么打分用 RMSNorm 后的 v, 加权和却用原始 v
|
||||
% omit: 论文里的 kernel 级调度与 pipeline 重叠
|
||||
|
||||
\section{Attention Residual 深度残差}
|
||||
|
||||
\subsection{从"等权累加"到"按内容选择"}
|
||||
|
||||
标准 Pre-Norm 残差把每层输出\textbf{无条件加}进残差流:
|
||||
|
||||
\[
|
||||
x_l = x_{l-1} + f_l(x_{l-1}),
|
||||
\qquad
|
||||
x_N = x_0 + \sum_{l=1}^{N} f_l(x_{l-1})
|
||||
\]
|
||||
|
||||
\noindent 展开后每一项权重恒为 1:第 3 层的输出和第 80 层的输出对最终表示的
|
||||
"名义"贡献一样大,深层无法表达"我这一步应该主要读第 12 层的结果"。
|
||||
|
||||
AttnRes(\texttt{arXiv:2603.15031})把这个加法换成\textbf{深度维上的 softmax 注意力}:
|
||||
第 $l$ 层持有一个可学习 query 向量 $w_l \in \mathbb{R}^D$,
|
||||
把此前所有层的输出当成"可读的记忆":
|
||||
|
||||
\begin{align}
|
||||
s_{l,i} &= w_l^{\top}\,\mathrm{RMS}(v_i),
|
||||
& i = 0,1,\dots,l-1 \tag{A1} \\
|
||||
\alpha_{l,i} &= \frac{\exp(s_{l,i})}{\sum_{j} \exp(s_{l,j})}
|
||||
& \shape{n, B, T} \tag{A2} \\
|
||||
h_l &= \sum_{i} \alpha_{l,i}\, v_i
|
||||
& \shape{B, T, D} \tag{A3} \\
|
||||
v_l &= f_l(h_l) \tag{A4}
|
||||
\end{align}
|
||||
|
||||
\noindent 其中 $v_0 = x$(embedding 输出),$f_l$ 是已经含 Pre-Norm 的原子子层,
|
||||
$\mathrm{RMS}(\cdot)$ 是不带 gain 的 RMS 归一化。
|
||||
最后(\texttt{is\_final\_aggregate=True})再用一个独立 query 聚合所有源得到 $y$。
|
||||
|
||||
\begin{importantbox}{注意力权重是逐 token 的}
|
||||
$s_{l,i}$ 的形状是 \shape{n, B, T}——每个 batch、每个位置 $t$ 都有自己的一套深度权重。
|
||||
所以同一个位置在不同深度可以读不同的层,但\textbf{不跨时间混合},
|
||||
因果性完全不受影响(\texttt{test\_attnres\_is\_still\_causal})。
|
||||
\end{importantbox}
|
||||
|
||||
\subsection{DepthResidual:三个实现细节}
|
||||
|
||||
\begin{codemathtop}{layers/attn\_res.py — DepthResidual}
|
||||
\begin{lstlisting}
|
||||
class DepthResidual(nn.Module):
|
||||
def __init__(self, dim, eps=1e-8, zero_init=True):
|
||||
self.query = nn.Parameter(torch.zeros(dim)) # [D]
|
||||
self.norm = RMSNorm(dim, eps=eps) # gain gamma
|
||||
|
||||
def effective_query(self):
|
||||
return (self.query * self.norm.weight).float() # 折叠 gain
|
||||
|
||||
def forward(self, sources):
|
||||
sources = stack_layers(sources) # [n,B,T,D]
|
||||
q = self.effective_query() # [D]
|
||||
k = rms(sources.float(), self.norm.eps) # 只用于打分
|
||||
logits = einsum('d, n b t d -> n b t', q, k)
|
||||
w = logits.softmax(dim=0) # 在深度维 softmax
|
||||
out = einsum('n b t, n b t d -> b t d', w, sources.float())
|
||||
return out.to(sources.dtype)
|
||||
\end{lstlisting}
|
||||
\end{codemathtop}
|
||||
|
||||
\paragraph{(1) gain 折叠}
|
||||
RMSNorm 的可学习 gain $\gamma$ 本该作用在 key 上,但
|
||||
$w^{\top}(\gamma \odot \mathrm{RMS}(v)) = (w \odot \gamma)^{\top}\mathrm{RMS}(v)$,
|
||||
所以直接把 $\gamma$ 折进 query:$\tilde{w}_l = w_l \odot \gamma_l$。
|
||||
少一次 \shape{n,B,T,D} 的逐元素乘法,两阶段算法里也只需要传一个向量。
|
||||
|
||||
\paragraph{(2) 打分用归一化的 $v$,加权和用原始 $v$}
|
||||
注意 \texttt{logits} 用 \texttt{k = rms(sources)},而 \texttt{out} 用的是
|
||||
\texttt{sources} 本身。
|
||||
|
||||
\begin{knowledgebox}{为什么这样不对称?}
|
||||
打分要的是\textbf{方向}:$\mathrm{RMS}$ 之后 $s_{l,i}$ 与 $\|v_i\|$ 无关,
|
||||
一层输出幅度大不会自动抢到高权重,softmax 只按"内容像不像我要读的东西"分配。\\
|
||||
加权和要的是\textbf{原始信息}:如果对归一化后的 $v$ 求和,每层输出的模长
|
||||
(承载着"这层贡献多大"的信息)就被抹掉了,深层的小幅修正会被放大到和主干同量级。
|
||||
\end{knowledgebox}
|
||||
|
||||
\paragraph{(3) zero-init query}
|
||||
\texttt{query} 默认初始化为 $0$ $\Rightarrow$ 所有 logits 为 $0$
|
||||
$\Rightarrow$ softmax 均匀 $\Rightarrow$
|
||||
\[
|
||||
h_l = \frac{1}{l}\sum_{i=0}^{l-1} v_i
|
||||
\]
|
||||
训练起步就是\textbf{等权深度平均}(已实测:零初始化时 \texttt{forward} 输出与
|
||||
\texttt{sources.mean(0)} 逐位相同),行为接近标准残差但自带 $1/l$ 缩放,
|
||||
之后由梯度慢慢学出偏好。设 \texttt{zero\_init\_queries=False} 则用 $\mathcal{N}(0, 0.02^2)$。
|
||||
|
||||
\subsection{Full 与 Block:源数量的差别}
|
||||
|
||||
两种堆叠方式的区别只在\textbf{谁有资格进入源列表}:
|
||||
|
||||
\begin{itemize}[nosep, leftmargin=2em]
|
||||
\item \texttt{FullAttnResStack}\\
|
||||
保留\textbf{每一个原子层}的输出作为源,第 $l$ 层在 $l+1$ 个源上做注意力。
|
||||
\item \texttt{BlockAttnResStack}\\
|
||||
把 $N$ 个原子层切成大小为 $S$ 的块,\textbf{块内退化成普通求和}
|
||||
(running partial $p \leftarrow p + v$),
|
||||
只有\textbf{块的输出} $b_j$ 才进入源列表。
|
||||
\end{itemize}
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{lccc}
|
||||
\toprule
|
||||
& \textbf{Full} & \textbf{Block ($S$)} & 标准残差 \\
|
||||
\midrule
|
||||
注意力源数 & 最多 $N+1$ & 最多 $N/S + 2$ & 1 \\
|
||||
需保留的 \shape{B,T,D} 激活 & $O(N)$ & $O(N/S)$ & $O(1)$ \\
|
||||
深度注意力 FLOPs & $O(N^2 BTD)$ & $O(N^2 BTD / S)$ & 0 \\
|
||||
新增参数 & $2(N{+}1)D$ & $2(N{+}1)D$ & 0 \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\noindent 参数量不变(每个原子层都有自己的 query),变的是\textbf{显存与带宽}。
|
||||
真实 K3($L=93$,$N=186$)取 $S=24$ 个原子层(12 个 DecoderBlock),
|
||||
源数从 187 降到 $\le 10$。
|
||||
|
||||
\begin{codemathtop}{layers/attn\_res.py — BlockAttnResStack.forward\_naive(语义参考实现)}
|
||||
\begin{lstlisting}
|
||||
blocks = [x] # b_0 = embedding
|
||||
partial = None
|
||||
for layer_idx, (layer, residual) in enumerate(zip(self.layers, self.residuals), 1):
|
||||
sources = blocks if partial is None else blocks + [partial]
|
||||
h = residual(sources) # 深度注意力
|
||||
out = layer(h)
|
||||
partial = out if partial is None else (partial + out) # 块内: 普通累加
|
||||
if (layer_idx % self.block_size == 0) or (layer_idx == len(self.layers)):
|
||||
blocks.append(partial) # 块边界: 定型成一个新源
|
||||
partial = None
|
||||
return self.final_residual(blocks)
|
||||
\end{lstlisting}
|
||||
\end{codemathtop}
|
||||
|
||||
\subsection{两阶段算法(inter / intra)}
|
||||
|
||||
块内逐层跑上面的 naive 版本有个浪费:块内每一层的 query 面对的
|
||||
\textbf{块间源 $b_0 \dots b_{j-1}$ 是完全相同且固定的},
|
||||
唯一在变的只有 running partial $p$。于是拆成两个阶段:
|
||||
|
||||
\begin{enumerate}[nosep]
|
||||
\item \textbf{inter(批量)}:把块内 $S$ 个 query 堆成 \shape{S, D},
|
||||
对固定源做\textbf{一次}批量 einsum,拿到每个 query 的 online-softmax 三元组
|
||||
$(m,\ \text{numer},\ \text{denom})$;
|
||||
\item \textbf{intra(串行)}:逐层把新出现的 $p$ 作为\textbf{单个源}合并进去,
|
||||
用 online softmax 的 merge 规则更新三元组,再 \texttt{normalized()} 出 $h$。
|
||||
\end{enumerate}
|
||||
|
||||
\noindent online softmax 的三元组定义与合并规则(和 FlashAttention 同构,
|
||||
只是"序列维"换成了"深度维"):
|
||||
|
||||
\begin{align}
|
||||
m = \max_i s_i,
|
||||
\quad
|
||||
n = \sum_i e^{s_i - m} v_i,
|
||||
\quad
|
||||
d = \sum_i e^{s_i - m},
|
||||
\quad
|
||||
h = n / d
|
||||
\tag{OS1}
|
||||
\end{align}
|
||||
|
||||
\noindent 合并两组统计量 $(m_a, n_a, d_a)$ 与 $(m_b, n_b, d_b)$,
|
||||
令 $m = \max(m_a, m_b)$、$w_a = e^{m_a - m}$、$w_b = e^{m_b - m}$:
|
||||
|
||||
\begin{align}
|
||||
n = w_a\, n_a + w_b\, n_b,
|
||||
\quad
|
||||
d = w_a\, d_a + w_b\, d_b
|
||||
\tag{OS2}
|
||||
\end{align}
|
||||
|
||||
\noindent 单个源 $p$ 的三元组是 $(\,m = s_p,\ \text{numer} = p,\ \text{denom} = 1\,)$
|
||||
——因为 $e^{s_p - m} = 1$,不需要真的算指数(\texttt{single\_source\_stats})。
|
||||
|
||||
\begin{codemathtop}{layers/attn\_res.py — \_run\_block\_two\_phase}
|
||||
\begin{lstlisting}
|
||||
queries = torch.stack([self.residuals[i].effective_query()
|
||||
for i in range(start, end)], dim=0) # [S, D]
|
||||
inter = attn_with_stats(queries, stack_layers(blocks), self.eps) # phase 1: 一次算完
|
||||
|
||||
partial = None
|
||||
for local_idx, layer_idx in enumerate(range(start, end)): # phase 2: 串行
|
||||
stats = inter.select(local_idx)
|
||||
if partial is not None:
|
||||
intra = single_source_stats(queries[local_idx], partial, self.eps)
|
||||
stats = merge_attn_stats(stats, intra) # online softmax merge
|
||||
h = stats.normalized()
|
||||
out = self.layers[layer_idx](h)
|
||||
partial = out if partial is None else (partial + out)
|
||||
return partial
|
||||
\end{lstlisting}
|
||||
\end{codemathtop}
|
||||
|
||||
\begin{importantbox}{等价性是被测出来的,不是假设的}
|
||||
\texttt{test\_block\_two\_phase\_matches\_naive} 直接对拍
|
||||
\texttt{mixer.forward\_naive(emb)} 与 \texttt{mixer(emb)},
|
||||
\texttt{atol=rtol=1e-5} 通过。Full 版同理:不传 \texttt{schedule\_block\_size}
|
||||
走 naive,传了走两阶段,两者一致。
|
||||
\end{importantbox}
|
||||
|
||||
\subsection{接入 CausalLM}
|
||||
|
||||
\subsubsection*{原子层 = 半个 DecoderBlock}
|
||||
|
||||
深度注意力的粒度是\textbf{原子层}而不是 DecoderBlock:
|
||||
每个 block 拆成"norm + attn"和"norm + ffn"两个 Pre-Norm 原子层,
|
||||
所以原子层数 $N = 2L$。
|
||||
|
||||
\begin{codemathtop}{models/causal\_lm.py — \_build\_mixer}
|
||||
\begin{lstlisting}
|
||||
atomics = []
|
||||
for block in blocks:
|
||||
atomics.append(BorrowedSubLayer(block.attn_norm, block.attn))
|
||||
atomics.append(BorrowedSubLayer(block.ffn_norm, block.ffn))
|
||||
|
||||
if mode == "full":
|
||||
return FullAttnResStack(D, atomics, eps=..., zero_init_queries=..., ...)
|
||||
if mode == "block":
|
||||
return BlockAttnResStack(D, atomics,
|
||||
block_size=atomic_block_size(config.num_hidden_layers,
|
||||
config.attnres_block_size), ...)
|
||||
\end{lstlisting}
|
||||
\end{codemathtop}
|
||||
|
||||
\subsubsection*{BorrowedSubLayer:借用而不注册}
|
||||
|
||||
\begin{codemathtop}{layers/attn\_res.py — BorrowedSubLayer}
|
||||
\begin{lstlisting}
|
||||
class BorrowedSubLayer(nn.Module):
|
||||
def __init__(self, norm, fn):
|
||||
self._borrowed = (norm, fn) # 普通 tuple, 不是 self.norm = norm
|
||||
|
||||
def forward(self, x):
|
||||
norm, fn = self._borrowed
|
||||
return fn(norm(x))
|
||||
\end{lstlisting}
|
||||
\end{codemathtop}
|
||||
|
||||
\begin{warningbox}{为什么必须用 tuple 藏起来?}
|
||||
如果写成 \texttt{self.norm = norm},\texttt{nn.Module} 会把它\textbf{注册成子模块},
|
||||
于是同一份权重同时挂在 \texttt{blocks.0.attn.*} 和 \texttt{mixer.layers.0.fn.*} 下:
|
||||
\begin{itemize}[nosep]
|
||||
\item \texttt{model.parameters()} 出现重复 $\Rightarrow$ 优化器对同一参数更新两次
|
||||
\item \texttt{state\_dict()} 多出一份镜像键 $\Rightarrow$ 旧 checkpoint 加载不上
|
||||
\end{itemize}
|
||||
放进普通 tuple 后 \texttt{blocks.*} 仍是唯一属主,
|
||||
\texttt{mixer} 下只多出 depth query 与 gain(\texttt{test\_no\_duplicate\_parameter\_ids} 守这条)。
|
||||
\end{warningbox}
|
||||
|
||||
\subsubsection*{forward:mixer 接管整条残差流}
|
||||
|
||||
\begin{codemathtop}{models/causal\_lm.py — CausalLM.forward}
|
||||
\begin{lstlisting}
|
||||
x = self.embedding(input_ids)
|
||||
if self.mixer is None:
|
||||
for block in self.blocks: # attnres="off": 老路径
|
||||
x = block(x)
|
||||
else:
|
||||
x = self.mixer(x) # full / block: DecoderBlock.forward 被完全绕过
|
||||
logits = self.lm_head(self.norm(x))
|
||||
\end{lstlisting}
|
||||
\end{codemathtop}
|
||||
|
||||
\noindent 注意 \texttt{mixer} 打开后 \texttt{DecoderBlock.forward}
|
||||
(\S8 里的 \texttt{x = x + attn(...)})\textbf{一次都不会被调用}——
|
||||
残差加法整个交给深度注意力,DecoderBlock 退化成"两个子层的容器"。
|
||||
|
||||
\subsubsection*{新增参数量:可忽略}
|
||||
|
||||
每个 DepthResidual 只有 query \shape{D} 和 gain \shape{D},共 $N+1$ 个:
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{lrrr}
|
||||
\toprule
|
||||
配置 & $D$ / $L$ & 原子层 $N$ & 新增参数 \\
|
||||
\midrule
|
||||
toy (K3Config) & 256 / 4 & 8 & 4{,}608 \\
|
||||
0.5b preset & 768 / 24 & 48 & 75{,}264 \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\subsection{配置与命令行}
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{lll}
|
||||
\toprule
|
||||
字段 & 默认 & 含义 \\
|
||||
\midrule
|
||||
\texttt{attnres} & \texttt{"off"} & \texttt{off} / \texttt{full} / \texttt{block} \\
|
||||
\texttt{attnres\_block\_size} & \texttt{None} & 每块几个 \textbf{DecoderBlock};\texttt{None} $\to \lceil L/8 \rceil$ \\
|
||||
\texttt{attnres\_zero\_init\_queries} & \texttt{True} & query 零初始化(等权起步)\\
|
||||
\texttt{attnres\_final\_aggregate} & \texttt{True} & 末尾再做一次全源聚合 \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\begin{codemathtop}{layers/attn\_res.py — atomic\_block\_size}
|
||||
\begin{lstlisting}
|
||||
def atomic_block_size(num_hidden_layers, attnres_block_size):
|
||||
"""DecoderBlock 数 -> 原子层数。None 时目标约 8 块。"""
|
||||
layers_per_block = (attnres_block_size if attnres_block_size is not None
|
||||
else max(1, (num_hidden_layers + 7) // 8))
|
||||
return layers_per_block * 2 # 每个 DecoderBlock = attn|ffn 两个原子层
|
||||
\end{lstlisting}
|
||||
\end{codemathtop}
|
||||
|
||||
\noindent 单位换算是最容易踩的一处:
|
||||
配置字段的单位是 \textbf{DecoderBlock 数},
|
||||
而堆叠类收到的 \texttt{block\_size} 是\textbf{原子层数}($\times 2$)。
|
||||
例如 $L = 24$、块大小留 \texttt{None}:
|
||||
\[
|
||||
\lceil 24/8 \rceil = 3 \text{ 个 DecoderBlock}
|
||||
\;\to\; S = 6 \text{ 个原子层}
|
||||
\;\to\; N/S = 48/6 = 8 \text{ 块}
|
||||
\]
|
||||
|
||||
\begin{lstlisting}
|
||||
uv run python train_k3.py --preset toy --attnres block --attnres-block-size 2
|
||||
uv run python train_k3.py --preset 0.5b --attnres block # 块大小自动 ~L/8
|
||||
\end{lstlisting}
|
||||
|
||||
\noindent \texttt{KDAConfig} 与 \texttt{K3Config} 都在
|
||||
\texttt{\_\_post\_init\_\_} 里调 \texttt{validate\_attnres},
|
||||
非法模式 / 块大小在构造时就报错。
|
||||
旧 checkpoint 的 config 里没有这几个字段,加载时回落到
|
||||
\texttt{off}(见 \S 9.7 验证清单最后两行)。
|
||||
|
||||
\subsection{验证清单}
|
||||
|
||||
\texttt{tests/integration/test\_attn\_res.py},14 项全过:
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{ll}
|
||||
\toprule
|
||||
测试 & 守住的性质 \\
|
||||
\midrule
|
||||
\texttt{default\_attnres\_is\_off} & 默认不改变任何既有行为,\texttt{mixer is None} \\
|
||||
\texttt{invalid\_attnres\_rejected} & 非法 mode / \texttt{block\_size=0} 构造期报错 \\
|
||||
\texttt{mixer\_kind\_and\_atomic\_count} & 原子层数 $= 2L$,块大小 $\times 2$ 换算 \\
|
||||
\texttt{auto\_block\_size\_targets\_eight\_blocks} & $L=93 \to 24$(K3 $S=12$ 个 block)\\
|
||||
\texttt{no\_duplicate\_parameter\_ids} & 借用不注册,参数 id / 名字均无重复 \\
|
||||
\texttt{off\_and\_block\_differ\_at\_same\_seed} & 同种子下确实换了计算图 \\
|
||||
\texttt{block\_two\_phase\_matches\_naive} & 两阶段 $\equiv$ naive,\texttt{atol 1e-5} \\
|
||||
\texttt{attnres\_is\_still\_causal} & 改末位 token 不影响前缀 logits \\
|
||||
\texttt{kda\_config\_block\_runs} & 纯 KDA 配置也能开 \\
|
||||
\texttt{attnres\_ckpt\_roundtrip} & 存取后逐位一致,config 字段保真 \\
|
||||
\texttt{old\_ckpt\_without\_attnres\_stays\_off} & 向后兼容 \\
|
||||
\texttt{attnres\_block\_overfits\_single\_batch} & 200 步 loss $< 0.5$,能训 \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\begin{warningbox}{混合精度}
|
||||
\texttt{DepthResidual.forward}(naive 路径)显式 \texttt{.float()} 后再算 softmax
|
||||
与加权和,最后 cast 回原 dtype;两阶段路径的 query 是 fp32、源保持原 dtype,
|
||||
靠 einsum 的类型提升处理。bf16 autocast 下前向实测正常。
|
||||
和 \S1 的结论一致:\textbf{指数/累和一律不要放进 fp16}。
|
||||
\end{warningbox}
|
||||
|
||||
\subsection{本章小结}
|
||||
|
||||
AttnRes 把残差流从"等权累加"升级成"深度维 softmax 注意力":
|
||||
每层用自己的 query 决定读此前哪些层的输出,打分在 RMS 归一化后做(方向)、
|
||||
加权和在原始张量上做(保留模长),query 零初始化让训练从等权平均起步。
|
||||
Full 版源数随深度线性增长,Block 版把块内退化成普通求和、只让块输出进入源列表,
|
||||
把源数压到 $O(N/S)$;两阶段算法进一步把块间注意力批量化,
|
||||
块内用 online softmax 增量合并,与 naive 实现数值等价。
|
||||
接入 \texttt{CausalLM} 时每个 DecoderBlock 拆成 attn / ffn 两个原子层,
|
||||
\texttt{BorrowedSubLayer} 用普通 tuple 借用权重以免重复注册,
|
||||
\texttt{attnres="off"} 保持旧路径不变。
|
||||
@@ -0,0 +1,157 @@
|
||||
% teach:
|
||||
% gap: 读者会用 autograd 但不知道手写 KDA backward 的具体展开
|
||||
% takeaway: backward = 逆序遍历时间步, 每步求 dq/dk/dv/dg/dbeta + 累积 dS; GVA 反传 = view+sum
|
||||
% jump: 为什么 dS_dec 要加 dS_acc 和 -k⊗dr 两项
|
||||
% omit: Triton backward 优化
|
||||
|
||||
\section{反向传播推导}
|
||||
|
||||
\subsection{Forward 回顾}
|
||||
|
||||
逐步写下 forward(省略 batch/head 下标):
|
||||
|
||||
\begin{align}
|
||||
S_{\mathrm{dec}} &= \exp(g_t) \odot S_{t-1} \tag{F1} \\
|
||||
p_t &= k_t^T S_{\mathrm{dec}} \tag{F2} \\
|
||||
r_t &= v_t - p_t \tag{F3} \\
|
||||
a_t &= \beta_t \cdot k_t \tag{F4} \\
|
||||
S_t &= S_{\mathrm{dec}} + a_t \otimes r_t \tag{F5} \\
|
||||
o_t &= q_t^T S_t \tag{F6}
|
||||
\end{align}
|
||||
|
||||
\subsection{反传公式(BPTT,$T \to 0$)}
|
||||
|
||||
设 $dS_{\mathrm{acc}}$ 是从时间步 $t$ 开始累积到 $S_t$ 上的梯度。逆序遍历:
|
||||
|
||||
\paragraph{Step 1: $o_t = q_t^T S_t$}
|
||||
|
||||
\begin{align}
|
||||
dS_{\mathrm{acc}} &\mathrel{+}= q_t \otimes do_t
|
||||
& \xrightarrow{\texttt{einsum('bhk,bhv->bhkv')}}
|
||||
& \quad \shape{B, HV, K, V} \\
|
||||
dq_t &= do_t^T S_t
|
||||
& \xrightarrow{\texttt{einsum('bhv,bhkv->bhk')}}
|
||||
& \quad \shape{B, HV, K}
|
||||
\end{align}
|
||||
|
||||
\paragraph{Step 2: $S_t = S_{\mathrm{dec}} + a_t \otimes r_t$}
|
||||
|
||||
外积的反传:$d(a \otimes r) = (\cdot)$,分解为:
|
||||
|
||||
\begin{align}
|
||||
da_t &= \sum_v r_{t,v} \cdot dS_{\mathrm{acc},\cdot,v}
|
||||
& \xrightarrow{\texttt{einsum('bhv,bhkv->bhk')}}
|
||||
& \quad \shape{B, HV, K} \\
|
||||
dr_t &= \sum_k a_{t,k} \cdot dS_{\mathrm{acc},k,\cdot}
|
||||
& \xrightarrow{\texttt{einsum('bhk,bhkv->bhv')}}
|
||||
& \quad \shape{B, HV, V}
|
||||
\end{align}
|
||||
|
||||
\paragraph{Step 3: $a_t = \beta_t \cdot k_t$}
|
||||
|
||||
\begin{align}
|
||||
d\beta_t &= \sum_k k_{t,k} \cdot da_{t,k}
|
||||
& \xrightarrow{\texttt{einsum('bhk,bhk->bh')}}
|
||||
& \quad \shape{B, HV} \\
|
||||
dk_t^{(a)} &= \beta_t \cdot da_t
|
||||
& & \quad \shape{B, HV, K}
|
||||
\end{align}
|
||||
|
||||
\paragraph{Step 4: $r_t = v_t - p_t = v_t - k_t^T S_{\mathrm{dec}}$}
|
||||
|
||||
\begin{align}
|
||||
dv_t &= dr_t & & \shape{B, HV, V} \\
|
||||
dk_t^{(r)} &= -S_{\mathrm{dec}}^T \cdot dr_t
|
||||
& \xrightarrow{\texttt{einsum('bhv,bhkv->bhk')}}
|
||||
& \quad \shape{B, HV, K} \\
|
||||
dS_{\mathrm{dec}}^{(r)} &= -k_t \otimes dr_t
|
||||
& \xrightarrow{\texttt{einsum('bhv,bhk->bhkv')}}
|
||||
& \quad \shape{B, HV, K, V}
|
||||
\end{align}
|
||||
|
||||
\paragraph{Step 5: 合并 $dS_{\mathrm{dec}}$ 并传递 $dg_t$, $dS_{t-1}$}
|
||||
|
||||
\[
|
||||
dS_{\mathrm{dec}}^{\mathrm{total}} = dS_{\mathrm{acc}} + dS_{\mathrm{dec}}^{(r)}
|
||||
= dS_{\mathrm{acc}} - k_t \otimes dr_t
|
||||
\]
|
||||
|
||||
因为 $S_{\mathrm{dec}} = \exp(g_t) \odot S_{t-1}$:
|
||||
|
||||
\begin{align}
|
||||
dg_t &= S_{\mathrm{dec}} \odot dS_{\mathrm{dec}}^{\mathrm{total}}
|
||||
& \xrightarrow{\texttt{einsum('bhkv,bhkv->bhk')}}
|
||||
& \quad \shape{B, HV, K} \\
|
||||
dS_{t-1} &= \exp(g_t) \odot dS_{\mathrm{dec}}^{\mathrm{total}}
|
||||
& & \quad \shape{B, HV, K, V}
|
||||
\end{align}
|
||||
|
||||
\paragraph{Step 6: 合并 $dk_t$ 和 GVA 归约}
|
||||
|
||||
\[
|
||||
dk_t = dk_t^{(a)} + dk_t^{(r)}
|
||||
= \beta_t \cdot da_t - S_{\mathrm{dec}}^T \cdot dr_t
|
||||
\]
|
||||
|
||||
GVA 反传($H_V \to H$):
|
||||
\[
|
||||
dq_H = dq_{H_V}.\texttt{view}(B, T, H, G, K).\texttt{sum}(\text{dim}=3) \cdot \mathrm{scale}
|
||||
\]
|
||||
\[
|
||||
dk_H = dk_{H_V}.\texttt{view}(B, T, H, G, K).\texttt{sum}(\text{dim}=3)
|
||||
\]
|
||||
|
||||
\subsection{代码对照}
|
||||
|
||||
\begin{codemathtop}{ops/reference/recurrent.py — KDAFunction.backward}
|
||||
\begin{lstlisting}
|
||||
for t in range(T - 1, -1, -1):
|
||||
q_t, k_t, b_t = q_ts[:,t], k_ts[:,t], b_ts[:,t]
|
||||
S_dec, r_t, a_t = S_decs[:,t], r_ts[:,t], a_ts[:,t]
|
||||
exp_g_t, do_t = exp_g_ts[:,t], do[:,t]
|
||||
|
||||
# Step 1: o_t = q_t . S_t
|
||||
S_t = S_dec + einsum('bhk,bhv->bhkv', a_t, r_t)
|
||||
dS_acc += einsum('bhk,bhv->bhkv', q_t, do_t)
|
||||
dq_e[:,t] = einsum('bhv,bhkv->bhk', do_t, S_t)
|
||||
|
||||
# Step 2: outer product grads
|
||||
da_t = einsum('bhv,bhkv->bhk', r_t, dS_acc)
|
||||
dr_t = einsum('bhk,bhkv->bhv', a_t, dS_acc)
|
||||
|
||||
# Step 3: a_t = beta_t * k_t
|
||||
dbeta[:,t] = einsum('bhk,bhk->bh', k_t, da_t)
|
||||
dk_t_a = b_t.unsqueeze(-1) * da_t
|
||||
|
||||
# Step 4: r_t = v_t - k_t . S_dec
|
||||
dv[:,t] = dr_t
|
||||
dS_dec_from_r = -einsum('bhv,bhk->bhkv', dr_t, k_t)
|
||||
dk_t_r = -einsum('bhv,bhkv->bhk', dr_t, S_dec)
|
||||
|
||||
# Step 5: S_dec = exp(g) * S_{t-1}
|
||||
dS_dec_total = dS_acc + dS_dec_from_r
|
||||
dk_e[:,t] = dk_t_a + dk_t_r
|
||||
dg[:,t] = einsum('bhkv,bhkv->bhk', S_dec, dS_dec_total)
|
||||
dS_acc = exp_g_t.unsqueeze(-1) * dS_dec_total
|
||||
|
||||
# Step 6: GVA reduce
|
||||
dq_H = dq_e.view(B,T,H,G,K).sum(dim=3) * scale
|
||||
dk_H = dk_e.view(B,T,H,G,K).sum(dim=3)
|
||||
\end{lstlisting}
|
||||
\end{codemathtop}
|
||||
|
||||
\begin{warningbox}{$dS_{\mathrm{dec}}^{\mathrm{total}}$ 为什么包含两项?}
|
||||
$S_t = S_{\mathrm{dec}} + a_t \otimes r_t$,$S_{\mathrm{dec}}$ 同时参与了:
|
||||
\begin{enumerate}[nosep]
|
||||
\item 直接传递到 $dS_{\mathrm{acc}}$(作为 $S_t$ 的一部分被读出)
|
||||
\item 通过 $r_t = v_t - k_t \cdot S_{\mathrm{dec}}$ 参与 delta rule
|
||||
\end{enumerate}
|
||||
所以 $dS_{\mathrm{dec}}^{\mathrm{total}} = dS_{\mathrm{acc}} + dS_{\mathrm{dec}}^{(r)}$,
|
||||
两条路径的梯度要\textbf{加}起来(chain rule 分叉处求和)。
|
||||
\end{warningbox}
|
||||
|
||||
\subsection{本章小结}
|
||||
|
||||
KDA backward 是 BPTT 展开:逆序遍历时间步,每步 6 个 einsum + 一次 $dS$ 累积更新。
|
||||
GVA 反传在最后做 \texttt{view+sum}。手写 backward 的关键是正确处理
|
||||
$dS_{\mathrm{dec}}$ 的两条梯度路径(直接传递 + 通过 $r_t$ 的 delta rule 路径)。
|
||||
@@ -0,0 +1,183 @@
|
||||
% teach:
|
||||
% gap: none — this is a reference appendix
|
||||
% takeaway: 一表查所有符号
|
||||
% jump: none
|
||||
% omit: none
|
||||
|
||||
\section{符号表}
|
||||
|
||||
\subsection{形状参数}
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{lll}
|
||||
\toprule
|
||||
符号 & 含义 & 典型值 (toy) \\
|
||||
\midrule
|
||||
$B$ & batch size & 2--4 \\
|
||||
$T$ & 序列长度 & 128--2048 \\
|
||||
$D$ & hidden\_size & 64 / 256 \\
|
||||
$H$ & query/key 头数 & 4 / 8 \\
|
||||
$H_V$ & value 头数(GVA) & $G \cdot H$ \\
|
||||
$G$ & GVA 组数 & $H_V / H$ \\
|
||||
$K$ & key/query 头维度 & 16 \\
|
||||
$V$ & value 头维度($= K$) & 16 \\
|
||||
$C$ & chunk\_size & 16 / 64 \\
|
||||
$r$ & KV latent rank (MLA) & 32 \\
|
||||
$d_q$ & MLA query head dim & 16 \\
|
||||
$d_v$ & MLA value head dim & 16 \\
|
||||
$\ell$ & MoE latent width ($= D/2$) & 128 \\
|
||||
$n_r$ & routed 专家数 & 16 \\
|
||||
$k$ & Top-$k$ & 2 \\
|
||||
$n_s$ & shared 专家数 & 2 \\
|
||||
$d_{\mathrm{ff}}$ & 专家中间维度 & 96 \\
|
||||
$N$ & AttnRes 原子层数 ($= 2L$) & 8 \\
|
||||
$S$ & AttnRes 块大小(原子层) & 2--24 \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\subsection{KDA 变量}
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{llp{7cm}}
|
||||
\toprule
|
||||
符号 & 形状 & 含义 \\
|
||||
\midrule
|
||||
$q_t$ & \shape{B, HV, K} & query(已 GVA 展开 + scale) \\
|
||||
$k_t$ & \shape{B, HV, K} & key(已 GVA 展开) \\
|
||||
$v_t$ & \shape{B, HV, V} & value \\
|
||||
$g_t$ & \shape{B, HV, K} & gate(log-space 衰减,逐维逐头) \\
|
||||
$\beta_t$ & \shape{B, HV} & 写入强度 \\
|
||||
$S_t$ & \shape{B, HV, K, V} & KV 状态矩阵 \\
|
||||
$S_{\mathrm{dec}}$ & \shape{B, HV, K, V} & 衰减后的状态 \\
|
||||
$p_t$ & \shape{B, HV, V} & 旧状态对 $k_t$ 的预测 \\
|
||||
$r_t$ & \shape{B, HV, V} & delta rule 残差 = $v_t - p_t$ \\
|
||||
$a_t$ & \shape{B, HV, K} & 写入向量 = $\beta_t \cdot k_t$ \\
|
||||
$o_t$ & \shape{B, HV, V} & 读出 = $q_t \cdot S_t$ \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\subsection{Gate 变量}
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{llp{6cm}}
|
||||
\toprule
|
||||
符号 & 形状 & 含义 \\
|
||||
\midrule
|
||||
$g_{\mathrm{raw}}$ & \shape{B, T, HV, K} & gate 投影原始输出 \\
|
||||
$A_{\log}$ & \shape{HV} & head-wise 衰减参数 (log-space) \\
|
||||
$\Delta_b$ & \shape{HV, K} & per-dim gate bias \\
|
||||
$\mathrm{rate}$ & \shape{HV, 1} & $\exp(A_{\log})$ \\
|
||||
$\mathrm{input}$ & \shape{B, T, HV, K} & $g_{\mathrm{raw}} + \Delta_b$ \\
|
||||
$L$ & 标量 & lower\_bound ($-5.0$) \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\subsection{MLA 变量}
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{llp{6cm}}
|
||||
\toprule
|
||||
符号 & 形状 & 含义 \\
|
||||
\midrule
|
||||
$c$ & \shape{B, T, r} & KV latent(推理时缓存这个) \\
|
||||
$q$ & \shape{B, T, H, d_q} & query(低秩路径输出) \\
|
||||
$W_{UK}$ & \shape{H, d_q, r} & key 解压矩阵(吸收进 $q$) \\
|
||||
$W_{UV}$ & \shape{H, d_v, r} & value 解压矩阵 \\
|
||||
$q_{\mathrm{abs}}$ & \shape{B, T, H, r} & 吸收后的 query \\
|
||||
score & \shape{B, H, T, T} & $q_{\mathrm{abs}} \cdot c^T$ \\
|
||||
attn & \shape{B, H, T, T} & causal softmax \\
|
||||
$\tilde{o}_{\mathrm{lat}}$ & \shape{B, H, T, r} & latent 加权输出 \\
|
||||
$\tilde{o}$ & \shape{B, H, T, d_v} & 解压后的输出 \\
|
||||
gate & \shape{B, T, H \cdot d_v} & $\sigma(W_g x)$ \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\subsection{LatentMoE 变量}
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{llp{6cm}}
|
||||
\toprule
|
||||
符号 & 形状 & 含义 \\
|
||||
\midrule
|
||||
$x$ & \shape{B, T, D} & 输入 \\
|
||||
$z$ & \shape{B, T, \ell} & latent ($\ell = D/2$) \\
|
||||
logits & \shape{B, T, n_r} & router logits \\
|
||||
ids & \shape{B, T, k} & Top-$k$ 专家索引 \\
|
||||
probs & \shape{B, T, k} & softmax 权重 \\
|
||||
$u$ & \shape{B, T, \ell} & routed 加权输出 \\
|
||||
$s$ & \shape{B, T, D} & shared 专家求和 \\
|
||||
$y$ & \shape{B, T, D} & $s + W_\uparrow \mathrm{RMSNorm}(u)$ \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\subsection{AttnRes 变量}
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{llp{6.4cm}}
|
||||
\toprule
|
||||
符号 & 形状 & 含义 \\
|
||||
\midrule
|
||||
$v_i$ & \shape{B, T, D} & 第 $i$ 个源($v_0 =$ embedding 输出) \\
|
||||
$w_l$ & \shape{D} & 第 $l$ 层的 depth query(零初始化) \\
|
||||
$\gamma_l$ & \shape{D} & DepthResidual 的 RMSNorm gain \\
|
||||
$\tilde{w}_l$ & \shape{D} & 折叠后的 query $= w_l \odot \gamma_l$ \\
|
||||
$s_{l,i}$ & \shape{n, B, T} & 深度打分 $= \tilde{w}_l^{\top}\mathrm{RMS}(v_i)$ \\
|
||||
$\alpha_{l,i}$ & \shape{n, B, T} & 深度维 softmax 权重 \\
|
||||
$h_l$ & \shape{B, T, D} & 第 $l$ 层的输入 $= \sum_i \alpha_{l,i} v_i$ \\
|
||||
$b_j$ & \shape{B, T, D} & 第 $j$ 个块的输出(Block 版的源) \\
|
||||
$p$ & \shape{B, T, D} & 块内 running partial \\
|
||||
$m, n, d$ & \shape{B, T} / \shape{B,T,D} / \shape{B,T} & online softmax 三元组 \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\subsection{Einsum 速查}
|
||||
|
||||
\begin{center}
|
||||
\small
|
||||
\begin{tabular}{p{6cm}lp{3.5cm}}
|
||||
\toprule
|
||||
操作 & einsum & 结果形状 \\
|
||||
\midrule
|
||||
key 查状态 & \texttt{'bhk,bhkv->bhv'} & $p_t$ \shape{B,HV,V} \\
|
||||
外积写入 & \texttt{'bhk,bhv->bhkv'} & $a_t \otimes r_t$ \shape{B,HV,K,V} \\
|
||||
读出 & \texttt{'bhk,bhkv->bhv'} & $o_t$ \shape{B,HV,V} \\
|
||||
MLA 吸收 $W_{UK}$ & \texttt{'bthd,hdj->bthj'} & $q_{\mathrm{abs}}$ \shape{B,T,H,r} \\
|
||||
MLA 打分 & \texttt{'bthj,bsj->bhts'} & score \shape{B,H,T,T} \\
|
||||
MLA latent 加权 & \texttt{'bhts,bsj->bhtj'} & $\tilde{o}_{\mathrm{lat}}$ \shape{B,H,T,r} \\
|
||||
MLA 解压 & \texttt{'bhtj,hvj->bhtv'} & $\tilde{o}$ \shape{B,H,T,d_v} \\
|
||||
AttnRes 深度打分 & \texttt{'d,nbtd->nbt'} & $s_{l,i}$ \shape{n,B,T} \\
|
||||
AttnRes 深度加权和 & \texttt{'nbt,nbtd->btd'} & $h_l$ \shape{B,T,D} \\
|
||||
AttnRes 批量打分(inter) & \texttt{'qd,nbtd->qnbt'} & logits \shape{S,n,B,T} \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\subsection{总结与延伸}
|
||||
|
||||
\subsubsection*{核心要点回顾}
|
||||
|
||||
\begin{enumerate}[nosep]
|
||||
\item \textbf{KDA} = delta rule 状态更新 + gate 衰减,线性复杂度
|
||||
\item \textbf{分块} = chunk 内下三角解 + chunk 间状态递推,等价于 naive recurrent
|
||||
\item \textbf{GVA} = $H_V = G \cdot H$,forward repeat\_interleave / backward view+sum
|
||||
\item \textbf{MLA} = 低秩 latent + 矩阵吸收,KV cache 从 $2Hd$ 降到 $r$
|
||||
\item \textbf{LatentMoE} = shared 全宽 + routed 半宽 latent + SiTU-GLU 防溢出
|
||||
\item \textbf{K3 Hybrid} = 3 KDA + 1 MLA,KDA 提供位置感知
|
||||
\item \textbf{AttnRes} = 深度维 softmax 残差,Block 版把源数压到 $O(N/S)$,
|
||||
两阶段 = inter 批量 + intra online-softmax 合并
|
||||
\end{enumerate}
|
||||
|
||||
\subsubsection*{未完成项}
|
||||
|
||||
\begin{itemize}[nosep]
|
||||
\item L5 — 项目内自研 fused gate Triton kernel
|
||||
\item L6 — recurrent decode cache(推理加速)
|
||||
\item AttnRes 与 recurrent decode 的组合(增量解码时的深度源缓存)
|
||||
\item AttnRes 开 / 关的收敛质量对比实验(目前只验证了等价性与可训练性)
|
||||
\end{itemize}
|
||||
@@ -0,0 +1,106 @@
|
||||
"""验证 reference `_decayed_dot` 的 g_ref 因子在真实训练门控量级下是否溢出。
|
||||
|
||||
_decayed_dot 把 exp(g_i-g_j) 拆成 (x*exp(g_i-g_ref)) @ (k*exp(g_ref-g_j)),
|
||||
g_ref = g_cumsum[..., 0, :]。因 g<0,最大中间因子为
|
||||
exp(g_ref - g_{C-1}) = exp(sum_{t=1}^{C-1} |g_t|)
|
||||
溢出阈值: sum|g| > ln(3.39e38)=88.7 (fp32/bf16) ; > ln(65504)=11.09 (fp16)
|
||||
|
||||
门控: safe_gate, g = lower_bound * sigmoid(exp(A_log) * (g_raw + dt_bias))
|
||||
lower_bound=-5 => |g| ∈ (0, 5) 逐步硬上界
|
||||
"""
|
||||
import math
|
||||
import sys
|
||||
|
||||
import torch
|
||||
|
||||
import kda.ops.api as kapi
|
||||
from kda.models.causal_lm import CausalLM
|
||||
from kda.models.config import KDAConfig
|
||||
from kda.models.k3_config import K3Config
|
||||
|
||||
dev = "cuda"
|
||||
LIM = {"fp32/bf16": math.log(3.39e38), "fp16": math.log(65504.0)}
|
||||
|
||||
# ---------- 1) 解析上界: 每个 chunk_size 需要的平均 |g| ----------
|
||||
print("=== 解析: 溢出所需的 chunk 内平均 |g|/step (硬上界 |g|<5) ===")
|
||||
print(f"{'C':>4} | {'fp32/bf16 阈值':>15} | {'可达?':>6} | {'fp16 阈值':>10} | {'可达?':>6}")
|
||||
for C in (16, 32, 64, 128):
|
||||
a, b = LIM["fp32/bf16"] / (C - 1), LIM["fp16"] / (C - 1)
|
||||
print(f"{C:4d} | {a:15.2f} | {'YES' if a < 5 else 'no':>6} | "
|
||||
f"{b:10.2f} | {'YES' if b < 5 else 'no':>6}")
|
||||
|
||||
# ---------- 2) 实测: 真实 checkpoint 上的 chunk 内 sum|g| ----------
|
||||
_orig = kapi.naive_chunk_kda # api.py 在 import 时绑定, 必须 patch 这里
|
||||
stats = []
|
||||
|
||||
|
||||
def max_sum_abs_g(g, C):
|
||||
"""chunk 内最大 sum_{t=1..C-1}|g_t| (= max exp(g_ref-g_j) 的指数)。"""
|
||||
T = g.shape[1]
|
||||
gg = g[:, : T - T % C].reshape(g.shape[0], -1, C, *g.shape[2:])
|
||||
return gg[:, :, 1:].abs().sum(dim=2).max().item()
|
||||
|
||||
|
||||
def patched(q, k, v, g, beta, **kw):
|
||||
stats.append(g.detach().float())
|
||||
return _orig(q, k, v, g, beta, **kw)
|
||||
|
||||
|
||||
kapi.naive_chunk_kda = patched
|
||||
|
||||
|
||||
def build(cfg_dict):
|
||||
"""checkpoint 的 config 是 dict, 按字段集合判断是 K3 还是纯 KDA。"""
|
||||
cls = K3Config if "moe_latent_size" in cfg_dict else KDAConfig
|
||||
return cls(**{k: v for k, v in cfg_dict.items()
|
||||
if k in cls.__dataclass_fields__})
|
||||
|
||||
|
||||
for path in ("ckpts/k3_wiki.pt", "ckpts/kda_toy.pt"):
|
||||
ck = torch.load(path, map_location="cpu", weights_only=False)
|
||||
cfg = build(ck["config"])
|
||||
model = CausalLM(cfg).to(dev).eval()
|
||||
# 旧 checkpoint 用 moe/moe_norm 命名, 现已重命名为 ffn/ffn_norm
|
||||
sd = {k.replace(".moe_norm.", ".ffn_norm.").replace(".moe.", ".ffn.")
|
||||
.replace(".mlp_norm.", ".ffn_norm.").replace(".mlp.", ".ffn."): v
|
||||
for k, v in ck["model_state"].items()}
|
||||
missing, unexpected = model.load_state_dict(sd, strict=False)
|
||||
assert not missing and not unexpected, (path, missing[:5], unexpected[:5])
|
||||
|
||||
stats.clear()
|
||||
ids = torch.randint(0, cfg.vocab_size, (2, 512), device=dev)
|
||||
with torch.no_grad():
|
||||
model(ids)
|
||||
if not stats:
|
||||
print(f"\n{path}: 未走 reference 路径 (backend={cfg.kda_backend})")
|
||||
continue
|
||||
gs = list(stats)
|
||||
print(f"\n=== {path} (训练用 C={cfg.chunk_size}, KDA 层数 {len(gs)}) ===")
|
||||
print(f" |g| mean/step: {sum(g.abs().mean().item() for g in gs)/len(gs):.4f} "
|
||||
f"|g| max/step: {max(g.abs().max().item() for g in gs):.4f} "
|
||||
f"(硬上界 {abs(cfg.lower_bound)})")
|
||||
print(f" {'C':>4} | {'max chunk sum|g|':>16} | {'max exp 因子':>12} | "
|
||||
f"{'fp32/bf16':>18} | {'fp16':>12}")
|
||||
for C in (16, 32, 64, 128):
|
||||
mg = max(max_sum_abs_g(g, C) for g in gs)
|
||||
fac = math.exp(mg) if mg < 709 else float("inf")
|
||||
f32 = f"{LIM['fp32/bf16']/mg:.2f}x OK" if mg < LIM["fp32/bf16"] else "OVERFLOW"
|
||||
f16 = f"{LIM['fp16']/mg:.2f}x OK" if mg < LIM["fp16"] else "OVERFLOW"
|
||||
mark = " <- 训练配置" if C == cfg.chunk_size else ""
|
||||
print(f" {C:4d} | {mg:16.2f} | {fac:12.3e} | {f32:>18} | {f16:>12}{mark}")
|
||||
|
||||
# ---------- 3) 随机初始化模型 (未训练) 同样测一遍 ----------
|
||||
for preset in ("toy", "0.5b"):
|
||||
cfg = K3Config.preset(preset)
|
||||
cfg.vocab_size = 2048
|
||||
cfg.kda_backend = "reference"
|
||||
model = CausalLM(cfg).to(dev).eval()
|
||||
stats.clear()
|
||||
with torch.no_grad():
|
||||
model(torch.randint(0, cfg.vocab_size, (2, 512), device=dev))
|
||||
if stats:
|
||||
gs = list(stats)
|
||||
mg = max(max_sum_abs_g(g, cfg.chunk_size) for g in gs)
|
||||
print(f"\n=== init preset={preset} (C={cfg.chunk_size}) ===")
|
||||
print(f" |g| mean/step {sum(g.abs().mean().item() for g in gs)/len(gs):.4f} "
|
||||
f"max chunk sum|g| {mg:.3f} fp32 余量 {LIM['fp32/bf16']/max(mg,1e-9):.0f}x")
|
||||
@@ -0,0 +1,83 @@
|
||||
"""验证: 不做 L2-norm 时 KDA 的爆炸是"数学上真实"还是"浮点误差"。
|
||||
|
||||
判据: 逐步递推 (naive_kda_fwd, 无三角求解) 在 float64 下的输出。
|
||||
- 若 fp64 递推也 ~1e32 => 爆炸是 KDA 递推本身的数学性质
|
||||
- 若 fp64 递推 O(1) 而 chunkwise 爆炸 => 是 solve 的浮点失效
|
||||
"""
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from kda.ops.reference.chunkwise import naive_chunk_kda
|
||||
from kda.ops.reference.recurrent import naive_kda_fwd
|
||||
|
||||
torch.manual_seed(0)
|
||||
dev = "cuda"
|
||||
B, T, H, K, V = 1, 64, 1, 16, 32
|
||||
C = 64
|
||||
|
||||
|
||||
def make(norm: bool, dtype):
|
||||
g0 = torch.Generator(device=dev).manual_seed(0)
|
||||
q = torch.randn(B, T, H, K, device=dev, dtype=dtype, generator=g0)
|
||||
k = torch.randn(B, T, H, K, device=dev, dtype=dtype, generator=g0)
|
||||
v = torch.randn(B, T, H, V, device=dev, dtype=dtype, generator=g0)
|
||||
# 训练里 g = -A.exp()*softplus(...) < 0,量级温和
|
||||
g = -F.softplus(torch.randn(B, T, H, K, device=dev, dtype=dtype, generator=g0)) * 0.1
|
||||
beta = torch.rand(B, T, H, device=dev, dtype=dtype, generator=g0)
|
||||
if norm:
|
||||
q, k = F.normalize(q, dim=-1), F.normalize(k, dim=-1)
|
||||
return q, k, v, g, beta
|
||||
|
||||
|
||||
def amax(x):
|
||||
return x.abs().max().item()
|
||||
|
||||
|
||||
print(f"{'norm':>5} | {'‖k‖':>6} | {'rec fp64':>10} | {'rec fp32':>10} | "
|
||||
f"{'chunk fp64':>10} | {'chunk fp32':>10} | {'rel(chunk64,rec64)':>18}")
|
||||
print("-" * 100)
|
||||
|
||||
for norm in (True, False):
|
||||
row = {}
|
||||
for dt in (torch.float64, torch.float32):
|
||||
q, k, v, g, beta = make(norm, dt)
|
||||
o_rec, _ = naive_kda_fwd(q, k, v, g, beta)
|
||||
o_chk, _ = naive_chunk_kda(q, k, v, g, beta, chunk_size=C)
|
||||
row[dt] = (amax(o_rec), amax(o_chk), o_chk, o_rec)
|
||||
knorm = make(norm, torch.float64)[1].norm(dim=-1).mean().item()
|
||||
r64, c64, oc64, or64 = row[torch.float64]
|
||||
r32, c32, oc32, _ = row[torch.float32]
|
||||
rel = ((oc64 - or64).abs().max() / (or64.abs().max() + 1e-30)).item()
|
||||
print(f"{str(norm):>5} | {knorm:6.2f} | {r64:10.3e} | {r32:10.3e} | "
|
||||
f"{c64:10.3e} | {c32:10.3e} | {rel:18.3e}")
|
||||
|
||||
# ---- M 的结构诊断 ----
|
||||
print("\n=== M = I + tril(A_kk*beta, -1) 诊断 (float64) ===")
|
||||
print(f"{'norm':>5} | {'max|N|':>9} | {'‖M⁻¹‖∞':>10} | {'cond2(M)':>10} | {'|S_final|max':>12}")
|
||||
for norm in (True, False):
|
||||
q, k, v, g, beta = make(norm, torch.float64)
|
||||
gc = g.cumsum(dim=1)[0, :, 0] # [T,K]
|
||||
kk = k[0, :, 0] # [T,K]
|
||||
gref = gc[:1]
|
||||
A = (kk * (gc - gref).exp()) @ (kk * (gref - gc).exp()).T
|
||||
N = (A * beta[0, :, 0][None, :]).tril(-1)
|
||||
M = torch.eye(T, dtype=torch.float64, device=dev) + N
|
||||
Minv = torch.linalg.inv(M)
|
||||
_, Sf = naive_kda_fwd(q, k, v, g, beta, output_final_state=True)
|
||||
print(f"{str(norm):>5} | {amax(N):9.3e} | {Minv.abs().sum(1).max():10.3e} | "
|
||||
f"{torch.linalg.cond(M).item():10.3e} | {amax(Sf):12.3e}")
|
||||
|
||||
# ---- ‖M⁻¹‖ 随 chunk 长度的增长 ----
|
||||
print("\n=== ‖M⁻¹‖∞ vs chunk 长度 C (float64) ===")
|
||||
for norm in (True, False):
|
||||
q, k, v, g, beta = make(norm, torch.float64)
|
||||
gc = g.cumsum(dim=1)[0, :, 0]
|
||||
kk = k[0, :, 0]
|
||||
out = []
|
||||
for C_ in (4, 8, 16, 32, 64):
|
||||
gs, ks, bs = gc[:C_], kk[:C_], beta[0, :C_, 0]
|
||||
gref = gs[:1]
|
||||
A = (ks * (gs - gref).exp()) @ (ks * (gref - gs).exp()).T
|
||||
M = torch.eye(C_, dtype=torch.float64, device=dev) + (A * bs[None, :]).tril(-1)
|
||||
out.append(f"C={C_:2d}:{torch.linalg.inv(M).abs().sum(1).max():.2e}")
|
||||
print(f" norm={str(norm):>5} " + " ".join(out))
|
||||
@@ -0,0 +1,72 @@
|
||||
"""严重度扫描: 把"数学爆炸"和"两条路径分道扬镳"分开测。
|
||||
|
||||
对每个 k 缩放系数 s:
|
||||
rec64 = 逐步递推 float64 (无三角求解) -> 数学真值
|
||||
chk64 = chunkwise float64 (全局 solve)
|
||||
chk32 = chunkwise float32
|
||||
tri32 = FLA/vendored triton 16x16 分块路径 (float32 in/out)
|
||||
"""
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from kda.ops.reference.chunkwise import naive_chunk_kda
|
||||
from kda.ops.reference.recurrent import naive_kda_fwd
|
||||
|
||||
dev = "cuda"
|
||||
B, T, H, K, V = 1, 64, 1, 16, 32
|
||||
|
||||
try:
|
||||
from kda.ops.triton.chunk import chunk_kda as triton_chunk_kda
|
||||
except Exception as e: # pragma: no cover
|
||||
triton_chunk_kda = None
|
||||
print("triton backend unavailable:", e)
|
||||
|
||||
|
||||
def make(scale_k, dtype):
|
||||
gen = torch.Generator(device=dev).manual_seed(0)
|
||||
q = torch.randn(B, T, H, K, device=dev, dtype=dtype, generator=gen)
|
||||
k = torch.randn(B, T, H, K, device=dev, dtype=dtype, generator=gen)
|
||||
v = torch.randn(B, T, H, V, device=dev, dtype=dtype, generator=gen)
|
||||
g = -F.softplus(torch.randn(B, T, H, K, device=dev, dtype=dtype, generator=gen)) * 0.1
|
||||
beta = torch.rand(B, T, H, device=dev, dtype=dtype, generator=gen)
|
||||
return q, F.normalize(k, dim=-1) * scale_k, v, g, beta
|
||||
|
||||
|
||||
def amax(x):
|
||||
return x.abs().max().item()
|
||||
|
||||
|
||||
hdr = (f"{'‖k‖':>6} | {'max|N|':>9} | {'‖M⁻¹‖∞':>10} | {'rec fp64':>10} | "
|
||||
f"{'chk fp64':>10} | {'chk fp32':>10} | {'tri fp32':>10} | "
|
||||
f"{'rel(chk64/rec64)':>16} | {'rel(chk32/chk64)':>16} | {'rel(tri32/chk32)':>16}")
|
||||
print(hdr)
|
||||
print("-" * len(hdr))
|
||||
|
||||
for s in (1.0, 1.5, 2.0, 2.5, 3.0, 4.0, 5.0):
|
||||
q, k, v, g, beta = make(s, torch.float64)
|
||||
o_rec, _ = naive_kda_fwd(q, k, v, g, beta)
|
||||
o_c64, _ = naive_chunk_kda(q, k, v, g, beta, chunk_size=64)
|
||||
|
||||
q3, k3, v3, g3, b3 = [x.float() for x in (q, k, v, g, beta)]
|
||||
o_c32, _ = naive_chunk_kda(q3, k3, v3, g3, b3, chunk_size=64)
|
||||
if triton_chunk_kda is not None:
|
||||
o_t32, _ = triton_chunk_kda(q3, k3, v3, g3, b3, chunk_size=64)
|
||||
o_t32 = o_t32.double()
|
||||
else:
|
||||
o_t32 = torch.full_like(o_c64, float("nan"))
|
||||
|
||||
gc = g.cumsum(dim=1)[0, :, 0]
|
||||
kk = k[0, :, 0]
|
||||
gref = gc[:1]
|
||||
A = (kk * (gc - gref).exp()) @ (kk * (gref - gc).exp()).T
|
||||
M = torch.eye(T, dtype=torch.float64, device=dev) + (A * beta[0, :, 0][None, :]).tril(-1)
|
||||
minv = torch.linalg.inv(M).abs().sum(1).max().item()
|
||||
|
||||
def rel(a, b):
|
||||
return ((a - b).abs().max() / (b.abs().max() + 1e-300)).item()
|
||||
|
||||
print(f"{k.norm(dim=-1).mean():6.2f} | {amax((A * beta[0, :, 0][None, :]).tril(-1)):9.2e} | "
|
||||
f"{minv:10.2e} | {amax(o_rec):10.2e} | {amax(o_c64):10.2e} | "
|
||||
f"{amax(o_c32):10.2e} | {amax(o_t32):10.2e} | "
|
||||
f"{rel(o_c64, o_rec):16.2e} | {rel(o_c32.double(), o_c64):16.2e} | "
|
||||
f"{rel(o_t32, o_c32.double()):16.2e}")
|
||||
Reference in New Issue
Block a user