Files
K3/notes/sections/sec-05.tex
T
dela 584f7e9e73 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.
2026-08-25 14:43:17 +08:00

76 lines
2.9 KiB
TeX
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
% 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 内部完成。