ggaaooppeenngg

为什么计算机科学是无限的但生命是有限的

前两篇分别实现了无遗忘的 chunked 线性注意力,以及带标量衰减的版本。两者的状态更新都只做加法:新的键值对累加进状态,旧信息靠衰减系数被动遗忘。本篇补上最后一块–删除

递推式变成:

St=St1(αt(Iβtktkt))+βtvtkt\mathbf{S}_t = \mathbf{S}_{t-1}\Big(\alpha_t\big(\mathbf{I} - \beta_t \bm{k}_t\bm{k}_t^{\intercal}\big)\Big) + \beta_t\bm{v}_t\bm{k}_t^{\intercal}

比上一篇多出的是 βt\beta_t–一个删除力度的系数:它决定在 kt\bm{k}_t 这个方向上,旧内容被擦掉多少。

上两篇见《Chunked 线性注意力》与《标量衰减》,本文沿用其符号约定与四层参考验证框架。本篇会用到 Householder 变换、WY 表示与 UT 变换这些线性代数工具,背景推导见《线性注意力的线性代数前置知识》。

阅读全文 »

上一篇实现了不带任何遗忘机制的 chunked 线性注意力,状态单调累加。本篇引入第一个衰减因子:把递推式改为 St=αSt1+vtkt\mathbf{S}_t = \alpha \mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^\intercalα\alpha 是一个标量常数。

这个衰减项来自哪里?Vanilla 线性注意力(St=St1+vtkt\mathbf{S}_t = \mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^\intercal)在语言建模上远不如 Transformer,为此 Mamba2(Dao & Gu, 2024a)引入了一个数据依赖的逐步遗忘门 αt(0,1)\alpha_t \in (0, 1) 来有选择地丢弃历史信息。本文取它的常数特例 αtα\alpha_t \equiv \alpha–这一档对应的正是 RetNet 与 Lightning-Attention,详见 §1.0。

上一篇见《TileLang 实战:KDA 从零到一–Chunked 线性注意力》,本文沿用其参考实现的验证框架。

阅读全文 »

KDA(Kimi Delta Attention)的递推式一行就写完,但直接照着它写 kernel 需要同时处理四个机制:逐通道门控、delta rule 的三角求解、log 域 cumsum、跨 chunk 状态传递。四个机制耦合在一个 kernel 里,数值出错时无法判断误差来自哪一层。 本文采用递进式实现路径,每一级只引入一个新机制、每一级都可独立运行并做数值验证。本文覆盖第一级–移除全部衰减因子与删除因子,只保留 SS+KVS \leftarrow S + K^\top V

这一级的价值不在性能,而在于将 chunkwise 分解恒等式单独隔离验证。本文同时给出一个反直觉的结论:线性注意力的全部优势是把 O(N2D)O(N^2 D) 换成 O(ND2)O(N D^2),但本文采用的逐块独立重算策略会让 FLOPs 退回 O(N2)O(N^2)–与因果 FlashAttention 同量级。§6 会说明这个退化的根源在于把序列轴放进了 grid,并对照 TileLang 官方 chunk_delta_h 的做法:序列轴不进 grid,递推就退回单 block 内的顺序循环,NN 的次数才能保住

数学背景见《KDA 的来龙去脉》,TileLang 语言基础见《TileLang 编程基本知识点》,本文复用的 tile 切分与 fragment 累加模式见《TileLang 实战:FlashAttention 前向 Kernel》。

阅读全文 »