% 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 内部完成。