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:
dela
2026-08-25 14:43:17 +08:00
commit 584f7e9e73
140 changed files with 21592 additions and 0 deletions
+15
View File
@@ -0,0 +1,15 @@
# Build context is projects/kda/. Keep the image a runtime, not a data dump.
.venv/
**/__pycache__/
**/*.pyc
**/*.pyo
**/.pytest_cache/
**/.ruff_cache/
ckpts/
*.pt
notes/
.git/
.gitignore
uv.lock
pyrightconfig.json
inspect_tensors.py
+23
View File
@@ -0,0 +1,23 @@
.venv/
swanlog/
.pytest_cache/
__pycache__/
*.py[cod]
ckpts/
data/spm_4k.*
data/pretrain/*
!data/pretrain/.gitkeep
data/sft/*
!data/sft/.gitkeep
!data/sft/toy.jsonl
data/eval/flores*
# LaTeX build output
*.aux
*.log
*.out
*.toc
*.fls
*.fdb_latexmk
*.xdv
wiki jsonl
+48
View File
@@ -0,0 +1,48 @@
# syntax=docker/dockerfile:1
#
# Single GPU runtime for train + SwanLab client + eval.
# Build: docker build -t kda:<tag> .
#
# Does not bake data, checkpoints, or API keys. Mount them at run time.
# Host: NVIDIA driver >= 570, nvidia-container-toolkit. See README.
FROM pytorch/pytorch:2.9.0-cuda12.8-cudnn9-devel
ENV DEBIAN_FRONTEND=noninteractive \
PIP_NO_CACHE_DIR=1 \
PYTHONUNBUFFERED=1 \
PYTHONPATH=/workspace/kda \
HF_HOME=/cache/huggingface \
HUGGINGFACE_HUB_CACHE=/cache/huggingface \
HF_HUB_DISABLE_TELEMETRY=1
RUN apt-get update && apt-get install -y --no-install-recommends \
git \
ca-certificates \
&& rm -rf /var/lib/apt/lists/*
WORKDIR /workspace/kda
# Layer cache: install deps from pyproject before the rest of the tree.
COPY pyproject.toml ./
COPY kda ./kda
# Base image already has torch/cuda/triton; do not let pip re-resolve torch.
RUN pip install --no-cache-dir --no-deps -e . && \
pip install --no-cache-dir \
"einops>=0.7.0" \
"packaging>=23.0" \
"sentencepiece>=0.2.0" \
"datasets>=3.0.0" \
"transformers>=4.51.0" \
"swanlab>=0.6.0" \
"sacrebleu>=2.4.0" \
"langdetect>=1.0.9" \
"pytest>=7.0"
COPY . /workspace/kda
RUN pip install --no-cache-dir --no-deps -e . && \
mkdir -p /cache/huggingface /workspace/kda/ckpts /workspace/kda/swanlog \
/data/pretrain /data/eval /data/sft
# Require an explicit entry (train / eval / pytest / swanlab ping).
CMD ["python", "scripts/container_help.py"]
+295
View File
@@ -0,0 +1,295 @@
# KDA 训练 → 推理 手写实现
从 naive recurrent 到 fused Triton kernel + 训练 + 推理,逐层手写实现。
对拍时复用上游 `flash-linear-attention` 的 `naive_recurrent_kda` / `naive_chunk_kda` 作为参考。
## 目录结构
```
kda/
├── pyproject.toml
├── README.md
├── train.py # KDA-only toy overfit
├── train_k3.py # K3-like 双语 wiki 预训练
├── train_sft.py # zh↔en 指令 SFT(模板与 eval_mt 相同)
├── Dockerfile # 训练 / SwanLab 客户端 / 评估 同一 GPU 镜像
├── compose.yaml
├── scripts/container_help.py # 镜像默认入口(拒绝无命令启动)
├── kda/ # 可安装包 `from kda import ...`
│ ├── _fla/ # FLA NVIDIA-Triton KDA 子集(不 import 上游 fla 包)
│ ├── ops/
│ │ ├── api.py # 显式选择 reference / triton / fla
│ │ ├── reference/ # L1/L2 + gate 的 PyTorch 正确性实现
│ │ ├── triton/ # vendored FLA chunk/gate/wy wrappers
│ │ └── recurrent/ # L6 fused recurrent decode + state cache
│ ├── layers/ # kda_attn / mla / swiglu / latent_moe / block
│ │ # attn_res.py = 深度残差 mixer(config.attnres)
│ ├── models/ # CausalLM + KDAConfig / K3Config
│ └── training/ # toy、双语 wiki、SFT、eval_mt / success
└── tests/
├── correctness/ # recurrent、chunkwise、gate reference
├── kernels/ # Triton / FLA 对拍
├── inference/ # recurrent decode
└── integration/ # Causal LM toy overfit、K3 架构
```
## 实现路线表
| # | 层级 | 文件 | 核心公式 / 关键操作 | 需手写的梯度 | 验证方法 | atol |
|---|------|------|---------------------|--------------|----------|------|
| L1 | Naive recurrent fwd+bwd | `ops/reference/recurrent.py` | `S_t = exp(g_t)⊙S + (β_t⊙k_t)⊗(v_t − k_t·S)`<br>`o_t = (q_t·scale)·S` | `dq,dk,dv,dg,dβ` | `torch.autograd.gradcheck` (float64) | 1e-4 |
| L2 | Naive chunked fwd+bwd | `ops/reference/chunkwise.py` | ① g→cumsum in chunk<br>② 构造下三角系统<br>③ triangular solve<br>④ chunk 间状态传递 | PyTorch autograd | recurrent 对拍 + gradcheck | 1e-4 |
| L3 | Triton fused fwd | `ops/triton/chunk_fwd.py` | chunk 内与 chunk 间并行计算 | — | Triton out vs L2 | 1e-4 |
| L4 | Triton fused bwd | `ops/triton/chunk_bwd.py` | dAv、状态反传、dqkg、intra 修正 | 端到端梯度 | gradcheck + L2 bwd 对拍 | 1e-3 |
| L5 | Gate reference / fusion | `ops/reference/gate.py`, `ops/triton/gate.py` | standard gate 与 safe gate 使用官方语义 | `dA_log,d_dt_bias,dg` | 公式对拍 + gradcheck | 1e-4 |
| L6 | recurrent decode | `ops/recurrent/fused.py` | `KDAState [B,HV,K,V]` | — | 逐 token vs L1 | 1e-5 |
| L7 | Model + training | `layers/`, `models/`, `training/` | KDA Causal LM + shifted CE + toy overfit | AdamW autograd | loss < 0.1 | — |
## 关键约束 / 风险点
| 项 | 约定 |
|---|---|
| **GVA (G=HV/H)** | `repeat_interleave(G, dim=2)` 把 q/k 从 H 扩到 HV;bwd 的 `dq/dk` 在 HV 维算完后 **必须 `view(B,T,H,G,K).sum(dim=3)` 回到 H 维** |
| **scale** | `1/√K`,在 forward 入口乘入 q(不要散布到 kernel 内部) |
| **gradcheck** | 强制 `dtype=torch.float64`;`eps=1e-6, atol=1e-4`;inputs 显式 `requires_grad_(True)` |
| **chunk_size** | 默认 64;必须 `T % BT == 0`;H100 可上 128 |
| **exp 下溢** | g 应用前 clamp;gate 路径 `lower_bound=−5.0`;显式 exp 路径用 `−softplus(−g)` 防大量 0 |
| **L4 dxs 落位** | `dq, dk` 在 HV 维算完后归约到 H;`dv/dg/dβ` 在 HV 维直接输出 |
| **naive 测试 shape** | T=16, B=2, H=2, HV=4, K=V=8 |
| **fused_whole 测试 shape** | T=128/512, B=2/4, H=4, HV=8, K=V=64 |
## 建议执行顺序
1. **L1 → L2**:先纯 PyTorch 把 fwd+bwd 写对,gradcheck 双保险
2. **L3 → L4**:写 Triton kernel 时,L2 当 reference,逐项 diff
3. **L5**:和 L3/L4 **解耦测试**(用 PyTorch 等价 naive gate 当参考)
4. **L6**:复用 L1 单步公式,独立测试
5. **L7**:把 L3+L4+L5 黏到 layer 里,配 SwiGLU + RMSNorm + CE 训练
## 测试状态
```
| Layer | Test | Status | max_diff | atol | Notes |
|-------|----------------------|----------|--------------|--------|--------------|
| L1/L2 | reference correctness | PASSED | | | recurrent + chunkwise |
| L3/L4 | vendored FLA Triton chunk | PASSED | | | CUDA, chunk_size 32/64, q/k L2-norm |
| L5 reference | gate formula + gradcheck | PASSED | | | standard + safe gate |
| L5 Triton | fused gate | PASSED | | | vendored FLA kda_gate_* |
| L6 | recurrent decode | PASSED | | | vendored FLA fused_recurrent_kda |
| L7 | train overfit | PASSED | | | reference backend |
| L7 | TorchLens trace/extract | PASSED | | | 计算图展开 + 激活提取 |
| L7 | TensorLens viewer API | PASSED | | | trace/normalize/Flask 端点 |
| K3 | hybrid + MLA + MoE | PASSED | | | `test_k3_arch.py` 6 项 |
| AttnRes | 深度残差 mixer | PASSED | | | `test_attn_res.py` 14 项 |
```
## 参考资源
- 上游 naive 实现 (`对拍用`): `(optional sibling) flash-linear-attention/fla/ops/kda/naive.py`
- 上游 fused kernel (`读源码用`): `(optional sibling) flash-linear-attention/fla/ops/kda/chunk_{fwd,bwd,intra}.py`
- KDA 笔记: `(optional sibling) flash-linear-attention/KDA_学习笔记.md`
- KDA paper: https://arxiv.org/abs/2510.26692
- AttnRes paper: https://arxiv.org/abs/2603.15031
- 本项目完整笔记(LaTeX): `notes/notes.tex` → `notes/notes.pdf`
## 当前进度
- [x] L1 — recurrent PyTorch reference + gradcheck
- [x] L2 — chunked PyTorch reference + recurrent 对拍 + gradcheck
- [x] L7 — 可训练 CausalLM、toy overfit、checkpoint round-trip
- [x] 统一 `ops.api.chunk_kda` 接口:`reference` / `triton` / `fla` 显式选择
- [x] K3-like — Gated MLA + LatentMoE + Hybrid 3:1,`train_k3.py` 双语 wiki 预训练
- [x] L3/L4 — vendored FLA Triton chunk fwd/bwd(`kda/_fla`,不静默 import 上游 `fla`)
- [x] L5 — vendored FLA fused gate
- [x] L6 — vendored FLA fused recurrent decode
- [x] AttnRes — 深度残差 mixer 接入 `CausalLM`(`config.attnres = off | full | block`)
- [x] 翻译训练环 — `--max-tokens` / `--resume` / held-out / cosine、`train_sft.py`、冻结 `data/eval`
当前可直接运行小模型训练:
```bash
PYTHONPATH=. python train.py
```
默认 `reference` 始终使用本仓库 PyTorch 实现。`backend="triton"` 走 `kda/_fla`
里搬运的 FLA NVIDIA Triton kernel(CUDA,`chunk_size` 为 32 或 64)。`backend="fla"`
仍要求完整安装上游 `flash-linear-attention`,只用于显式对拍。
## Kimi K3 小规模复现(KDA + Gated MLA + Stable LatentMoE)
按 `learning/kimi-k3-notes` 的架构笔记,复现 K3 的**核心三件套**(Hybrid Attention
三选一 + MoE),规模受 6GB 显存限制缩小约 30 倍:
| K3 组件 | 真实 K3 | 本复现 (toy) | 实现 |
|---|---|---|---|
| KDA | safe gate, NoPE, H=HV=96 | H=HV=8, K=V=16 | `kda/ops/`(已验证)|
| Hybrid | 每 4 层 1×Gated MLA,末层 MLA | 同 pattern (L=4) | `kda/layers/mla.py` |
| Gated MLA | kv_lora 512, NoPE, 矩阵吸收 | kv_lora 32, 吸收版 | 同上 |
| Stable LatentMoE | ℓ=d/2=3584, 896/16, shared 2 | ℓ=d/2=128, 16/2, shared 2 | `kda/layers/latent_moe.py` |
| SiTU-GLU | β1=4, β2=25 | 同 | 同上 |
| AttnRes | Block, S≈L/8 个 DecoderBlock | 同(`--attnres block`,默认 off)| `kda/layers/attn_res.py` |
验证(`tests/integration/test_k3_arch.py`,6 项全过):
- MLA 矩阵吸收版 **fwd+bwd 均与解压版逐位一致**
- LatentMoE Top-k 路由权重正确、shared 全宽贡献
- Hybrid pattern(每 4 层 1 MLA + 末层强制)
- K3 模型因果性 + 小模型单 batch 收敛
### 复现训练
目标是 **~0.5B zh↔en 指令翻译模型**(`K3Config.preset("0.5b")` = 482M,tied Qwen3 词表)。本机 RTX 3060 6GB 只跑 8M 全流程孪生;0.5B 预训练需要 **32–40GB Ampere bf16**。
成功标准是冻结集上的 `translation_success()`,**不是** wiki train loss。wiki 预训练没见过 `Translate to English:\n...`,预训练阶段 `eval_mt` 的 success_rate 预期 ≈0。
```bash
# 8M 孪生(自训 8k SentencePiece;默认 zh+en wiki,缓存 data/pretrain/)
uv run python kda/training/train_tokenizer.py --out data/spm_4k --vocab-size 8192 --limit 20000
uv run python train_k3.py --preset toy --limit 8000 --steps 800 --batch 4 --seq-len 256
uv run python train_sft.py --ckpt ckpts/k3_wiki.pt --data data/sft/toy.jsonl \
--src data/eval/zh2en.src.txt --ref data/eval/zh2en.ref.txt --target-lang en
# AttnRes 对照(默认 off,块大小不给则自动 ≈ L/8 个 DecoderBlock)
uv run python train_k3.py --preset toy --attnres block --attnres-block-size 2
# ~0.5B 冒烟(默认 --steps 2000 ≈ 8.2M token,不是语言模型)
uv run python train_k3.py --preset 0.5b --attnres block
# 32–40GB Ampere:1B-token 双语预训练(可 --resume ckpts/k3_0.5b_last.pt 续到 2–5B)
uv run python train_k3.py --preset 0.5b --attnres block \
--max-tokens 1000000000 --warmup 2000
```
`0.5b` 预设:`d=768`,`L=24`(6×3 KDA + 1 MLA),`H=12`,`head=64`,`chunk=64`,LatentMoE `ℓ=384` / 16 routed / Top-2 / shared 2,tied Qwen3 embedding,activation checkpoint 默认开。训练默认 seq 2048、micro-batch 2、grad-acc 8、lr 3e-4、warmup **64 optimizer steps**。checkpoint:`ckpts/k3_0.5b.pt`,另写 `_last` / `_best`(best 按 held-out CE)。
Token 会计:`step` 仍是 micro-batch;有效 token = `batch × seq_len × micro_steps`。默认 0.5b 冒烟是 **8.2M token ≈ 0.017 tok/param**。翻译前置 LM 的最低有意义预算是 **1B token**(`--max-tokens`),不是 2000 step。
训练指标参考(800 步 @ 8M 参数,RTX 3060,bf16 + grad clip 1.0):
- loss:9.0(ln 8192 均匀起步)→ best 3.81,无 nan
- 生成:乱码 → 高频中文片段 → 短句雏形
- fp16 AMP 会在 KDA `exp(g)/cumsum` 上发散,必须 bf16 或 fp32
入口:`train.py` 用 `KDAConfig`;`train_k3.py` / `train_sft.py` 用 checkpoint 里的 `K3Config`(或 KDA)。SFT 指令模板与 `eval_mt` 相同,见 `kda/training/prompts.py`。冻结评测句在 `data/eval/`(不要拿去训练)。FLORES 导出:`uv run python scripts/export_flores.py --out data/eval`。
## Docker 训练镜像
一个 GPU 运行时覆盖 **训练、SwanLab 客户端监控、冻结集评估**。构建上下文 = 本目录。
默认 `CMD` 只打印用法并退出码 2,必须显式传入口,避免误跑 toy overfit。
语料、checkpoint、`SWANLAB_API_KEY`、HF 词表缓存 **不打进镜像**,运行时挂载。
### 版本 / CUDA 矩阵(与本地开发环境一致,已实测)
| 组件 | 版本 | 来源 |
|------|------|------|
| 基础镜像 | `pytorch/pytorch:2.9.0-cuda12.8-cudnn9-devel` | Docker Hub(~9.6GB)|
| Python | 3.12 | 基础镜像自带 |
| torch | 2.9.0+cu128 | 基础镜像自带 |
| CUDA | 12.8 | 基础镜像 runtime |
| cuDNN | 9.x | 基础镜像自带 |
| triton | 3.5.0 | torch 附带 |
| numpy | 2.3.x | torch 附带 |
| 训练依赖 | einops / sentencepiece / datasets / transformers | `pyproject.toml` |
| 监控 / 评估 | swanlab / sacrebleu / langdetect | extra `train` |
宿主要求:
- **NVIDIA driver >= 570.00**(CUDA 12.8 最低要求)。本机 driver 610.57 ✅
- **nvidia-container-toolkit**(容器内用 GPU 的前置,见下)
- GPU 架构:cu128 完整支持 Ampere 及以后(RTX 3060 / sm_86 ✅)
- 显存注意:RTX 3060 只有 6GB,toy 训练无压力;大 batch / `chunk_size=128`
可能 OOM,按需调小
### 宿主机一次性配置:nvidia-container-toolkit
当前机器**未安装**,`docker run --gpus all` 会报
`failed to discover GPU vendor from CDI`。安装并配置:
```bash
# Arch Linux
sudo pacman -S nvidia-container-toolkit
sudo nvidia-ctk runtime configure --runtime=docker
sudo systemctl restart docker
# Ubuntu / Debian(若换机器)
curl -fsSL https://nvidia.github.io/libnvidia-container/gpgkey \
| sudo gpg --dearmor -o /usr/share/keyrings/nvidia-container-toolkit-keyring.gpg
curl -s -L https://nvidia.github.io/libnvidia-container/stable/deb/nvidia-container-toolkit.list \
| sed 's#deb https://#deb [signed-by=/usr/share/keyrings/nvidia-container-toolkit-keyring.gpg] https://#g' \
| sudo tee /etc/apt/sources.list.d/nvidia-container-toolkit.list
sudo apt-get update && sudo apt-get install -y nvidia-container-toolkit
sudo nvidia-ctk runtime configure --runtime=docker
sudo systemctl restart docker
```
验证:
```bash
docker run --rm --gpus all nvidia/cuda:12.8.0-base-ubuntu24.04 nvidia-smi
# 应能看到 RTX 3060 与 driver 610.57
```
### 构建
```bash
docker build -t kda:latest .
```
> `--network=host` 必要:本机 docker daemon 配置了代理 `http://127.0.0.1:7890`,
> 构建容器内的 `127.0.0.1` 指向容器自身而非宿主,不走 host 网络时 pip 下载会失败。
### 运行
无命令时只打印用法(退出码 2)。正式跑必须写入口。
```bash
# 用法
docker run --rm kda:latest
# 监控连通(训练机出网到 api.swanlab.cn)
docker run --rm --gpus all -e SWANLAB_API_KEY kda:latest swanlab ping
# toy 训练 + 云端监控;ckpt 落在宿主
mkdir -p ckpts
docker run --rm --gpus all \
-e SWANLAB_API_KEY \
-v "$PWD/ckpts:/workspace/kda/ckpts" \
-v "$PWD/data:/workspace/kda/data" \
kda:latest python train_k3.py --preset toy --attnres off
# 0.5b 冒烟(词表走 HF 缓存卷)。真正的 1B-token 预训练加 --max-tokens 1000000000
docker run --rm --gpus all \
-e SWANLAB_API_KEY \
-v "$PWD/ckpts:/workspace/kda/ckpts" \
-v "$PWD/data:/workspace/kda/data" \
-v hf-cache:/cache/huggingface \
kda:latest python train_k3.py --preset 0.5b --attnres block
# 评估(src/ref 一行一句,挂冻结集)
docker run --rm --gpus all \
-v "$PWD/ckpts:/workspace/kda/ckpts" \
-v "$PWD/data/eval:/data/eval:ro" \
kda:latest python -m kda.training.eval_mt \
--ckpt /workspace/kda/ckpts/k3_wiki.pt \
--src /data/eval/zh2en.src.txt --ref /data/eval/zh2en.ref.txt \
--target-lang en
# 镜像内测试
docker run --rm --gpus all kda:latest python -m pytest -q \
tests/correctness tests/integration/test_causal_lm.py \
tests/integration/test_k3_arch.py tests/integration/test_attn_res.py \
tests/integration/test_eval_mt.py
# 交互
docker run --rm -it --gpus all --entrypoint bash kda:latest
```
也可用 `docker compose run --rm train python train_k3.py --preset toy`(`compose.yaml`)。
### 镜像内验证
```bash
docker run --rm kda:latest python -c "import torch, einops, swanlab, sacrebleu; print(torch.__version__, torch.version.cuda)"
# 预期: 2.9.0+cu128 12.8
```
+29
View File
@@ -0,0 +1,29 @@
# GPU train / eval against the image in Dockerfile.
# docker compose run --rm train python train_k3.py --preset toy
# docker compose run --rm train swanlab ping
# docker compose run --rm train python -m kda.training.eval_mt --ckpt ...
#
# Secrets and corpora stay on the host.
services:
train:
build: .
image: kda:latest
gpus: all
ipc: host
shm_size: "2gb"
working_dir: /workspace/kda
environment:
SWANLAB_API_KEY: ${SWANLAB_API_KEY:-}
HF_HOME: /cache/huggingface
HUGGINGFACE_HUB_CACHE: /cache/huggingface
volumes:
- ./ckpts:/workspace/kda/ckpts
- ./data:/workspace/kda/data
- ${PRETRAIN_DATA:-./data/pretrain}:/data/pretrain:ro
- ${EVAL_DATA:-./data/eval}:/data/eval:ro
- ${SFT_DATA:-./data/sft}:/data/sft:ro
- hf-cache:/cache/huggingface
volumes:
hf-cache:
+1
View File
@@ -0,0 +1 @@
# Mount as /data/eval: one sentence per line, e.g. zh2en.src.txt + zh2en.ref.txt
+20
View File
@@ -0,0 +1,20 @@
今天天气很好。
请把窗户打开。
猫坐在垫子上。
这本书值得一读。
他昨天去了北京。
我们需要更多的训练数据。
太阳从东边升起。
她正在学习线性代数。
不要把评测集拿去训练。
河对面有一座旧桥。
科学是对自然的系统探索。
他们在公园里散步。
这台电脑的内存是十六吉字节。
翻译时不要照抄原文。
春天的风很温和。
我把钥匙放在桌子上了。
火车中午到达。
水在一百摄氏度沸腾。
小模型仍然可以学会狭窄的任务。
冻结测试文件保持只读。
+20
View File
@@ -0,0 +1,20 @@
The weather is very nice today.
Please open the window.
The cat is sitting on the mat.
This book is worth reading.
He went to Beijing yesterday.
We need more training data.
The sun rises in the east.
She is studying linear algebra.
Do not train on the evaluation set.
There is an old bridge across the river.
Science is the systematic exploration of nature.
They are walking in the park.
This computer has sixteen gigabytes of memory.
Do not copy the source text when translating.
The spring wind is gentle.
I left the keys on the table.
The train arrives at noon.
Water boils at one hundred degrees Celsius.
A small model can still learn a narrow task.
Keep the frozen test files read-only.
+20
View File
@@ -0,0 +1,20 @@
The weather is very nice today.
The development of artificial intelligence has changed the world.
Please open the window.
The cat is sitting on the mat.
This book is worth reading.
The meeting will start at three in the afternoon.
He went to Beijing yesterday.
We need more training data.
The sun rises in the east.
This question still has no answer.
She is studying linear algebra.
Do not train on the evaluation set.
There is an old bridge across the river.
Please wait a moment, I will be right back.
Science is the systematic exploration of nature.
They are walking in the park.
This computer has sixteen gigabytes of memory.
Do not copy the source text when translating.
The spring wind is gentle.
I left the keys on the table.
+20
View File
@@ -0,0 +1,20 @@
今天天气很好。
人工智能的发展改变了世界。
请把窗户打开。
猫坐在垫子上。
这本书值得一读。
会议将在下午三点开始。
他昨天去了北京。
我们需要更多的训练数据。
太阳从东边升起。
这个问题还没有答案。
她正在学习线性代数。
不要把评测集拿去训练。
河对面有一座旧桥。
请稍等,我马上回来。
科学是对自然的系统探索。
他们在公园里散步。
这台电脑的内存是十六吉字节。
翻译时不要照抄原文。
春天的风很温和。
我把钥匙放在桌子上了。
+1
View File
@@ -0,0 +1 @@
# Mount bilingual pretrain shards here (not baked into the image).
+2
View File
@@ -0,0 +1,2 @@
SFT bitext (jsonl or tsv) is mounted here. toy.jsonl is a tiny in-repo twin set.
OPUS / WMT dumps stay out of git.
+32
View File
@@ -0,0 +1,32 @@
{"src": "你好。", "tgt": "Hello.", "target_lang": "en"}
{"src": "Hello.", "tgt": "你好。", "target_lang": "zh"}
{"src": "谢谢。", "tgt": "Thank you.", "target_lang": "en"}
{"src": "Thank you.", "tgt": "谢谢。", "target_lang": "zh"}
{"src": "我爱北京。", "tgt": "I love Beijing.", "target_lang": "en"}
{"src": "I love Beijing.", "tgt": "我爱北京。", "target_lang": "zh"}
{"src": "现在几点?", "tgt": "What time is it now?", "target_lang": "en"}
{"src": "What time is it now?", "tgt": "现在几点?", "target_lang": "zh"}
{"src": "水是透明的。", "tgt": "Water is transparent.", "target_lang": "en"}
{"src": "Water is transparent.", "tgt": "水是透明的。", "target_lang": "zh"}
{"src": "他是一名教师。", "tgt": "He is a teacher.", "target_lang": "en"}
{"src": "He is a teacher.", "tgt": "他是一名教师。", "target_lang": "zh"}
{"src": "明天会下雨。", "tgt": "It will rain tomorrow.", "target_lang": "en"}
{"src": "It will rain tomorrow.", "tgt": "明天会下雨。", "target_lang": "zh"}
{"src": "请坐。", "tgt": "Please sit down.", "target_lang": "en"}
{"src": "Please sit down.", "tgt": "请坐。", "target_lang": "zh"}
{"src": "这是一只狗。", "tgt": "This is a dog.", "target_lang": "en"}
{"src": "This is a dog.", "tgt": "这是一只狗。", "target_lang": "zh"}
{"src": "大门在左边。", "tgt": "The gate is on the left.", "target_lang": "en"}
{"src": "The gate is on the left.", "tgt": "大门在左边。", "target_lang": "zh"}
{"src": "我们走吧。", "tgt": "Let's go.", "target_lang": "en"}
{"src": "Let's go.", "tgt": "我们走吧。", "target_lang": "zh"}
{"src": "夜空中有星星。", "tgt": "There are stars in the night sky.", "target_lang": "en"}
{"src": "There are stars in the night sky.", "tgt": "夜空中有星星。", "target_lang": "zh"}
{"src": "面包放在厨房里。", "tgt": "The bread is in the kitchen.", "target_lang": "en"}
{"src": "The bread is in the kitchen.", "tgt": "面包放在厨房里。", "target_lang": "zh"}
{"src": "孩子在睡觉。", "tgt": "The child is sleeping.", "target_lang": "en"}
{"src": "The child is sleeping.", "tgt": "孩子在睡觉。", "target_lang": "zh"}
{"src": "这座山很高。", "tgt": "This mountain is very high.", "target_lang": "en"}
{"src": "This mountain is very high.", "tgt": "这座山很高。", "target_lang": "zh"}
{"src": "请关上门。", "tgt": "Please close the door.", "target_lang": "en"}
{"src": "Please close the door.", "tgt": "请关上门。", "target_lang": "zh"}
+124
View File
@@ -0,0 +1,124 @@
"""Inspect KDA activations: print stats, TorchLens extract, or TensorLens web UI.
uv run python inspect_tensors.py # shape / min / max / nan
uv run python inspect_tensors.py --lens torch # named activations
uv run python inspect_tensors.py --lens web # http://127.0.0.1:8000
"""
from __future__ import annotations
import argparse
import torch
from kda.models.causal_lm import CausalLM
from kda.models.config import KDAConfig
MODULES = (
"embedding",
"blocks.0.attn.q_proj",
"blocks.0.attn.k_proj",
"blocks.0.attn.v_proj",
"blocks.0.attn",
"blocks.0.ffn",
"blocks.0",
"norm",
)
def _model() -> CausalLM:
torch.manual_seed(51)
return CausalLM(
KDAConfig(
hidden_size=16,
num_hidden_layers=1,
num_heads=2,
num_value_heads=2,
head_dim=4,
chunk_size=4,
vocab_size=32,
intermediate_size=32,
kda_backend="reference",
)
).eval()
def _tokens() -> torch.Tensor:
return torch.tensor([[1, 2, 3, 4]])
def _named_modules(model: torch.nn.Module) -> dict[str, torch.nn.Module]:
return dict(model.named_modules())
def capture_activations(model: CausalLM, x: torch.Tensor) -> dict[str, torch.Tensor]:
captured: dict[str, torch.Tensor] = {}
hooks = []
modules = _named_modules(model)
for name in MODULES:
module = modules[name]
def _hook(_module, _inp, out, key=name):
captured[key] = out.detach()
hooks.append(module.register_forward_hook(_hook))
with torch.no_grad():
captured["logits"] = model(x).detach()
for hook in hooks:
hook.remove()
return captured
def print_stats(tensors: dict[str, torch.Tensor]) -> None:
print(f"{'name':28} {'shape':18} {'dtype':10} {'min':>10} {'max':>10} {'mean':>10} nan/inf")
for name, tensor in tensors.items():
finite = torch.isfinite(tensor)
n_bad = int((~finite).sum())
stats = tensor.float() if tensor.is_floating_point() else tensor
print(
f"{name:28} {str(tuple(tensor.shape)):18} {str(tensor.dtype):10} "
f"{stats.min().item():10.4f} {stats.max().item():10.4f} "
f"{stats.float().mean().item():10.4f} {n_bad}"
)
def inspect_torchlens(model: CausalLM, x: torch.Tensor) -> None:
import torchlens as tl
names = [*MODULES, "output"]
with torch.no_grad():
acts = tl.extract(model, x, names)
print_stats(acts)
def inspect_web(model: CausalLM, x: torch.Tensor, host: str, port: int) -> None:
from tensorlens.tensorlens import trace
from tensorlens.web.server import app
acts = capture_activations(model, x)
for name, tensor in acts.items():
trace(name, tensor.cpu().float().numpy(), normalization="minmax")
trace("lm_head.weight", model.lm_head.weight.detach().cpu().float().numpy(), normalization="minmax")
print(f"TensorLens: http://{host}:{port} (Ctrl-C to stop)")
app.run(host=host, port=port, debug=False, use_reloader=False)
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--lens", choices=("print", "torch", "web"), default="print")
parser.add_argument("--host", default="127.0.0.1")
parser.add_argument("--port", type=int, default=8000)
args = parser.parse_args()
model = _model()
x = _tokens()
if args.lens == "torch":
inspect_torchlens(model, x)
return
if args.lens == "web":
inspect_web(model, x, args.host, args.port)
return
print_stats(capture_activations(model, x))
if __name__ == "__main__":
main()
+14
View File
@@ -0,0 +1,14 @@
"""KDA operators, composable layers, and CausalLM."""
from .models.causal_lm import CausalLM
from .models.config import KDAConfig
from .models.k3_config import K3Config
from .ops import chunk_kda
__all__ = [
"CausalLM",
"KDAConfig",
"K3Config",
"chunk_kda",
]
__version__ = "0.0.1"
+21
View File
@@ -0,0 +1,21 @@
MIT License
Copyright (c) 2023-2026 Songlin Yang, Yu Zhang, Zhiyuan Li
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
+9
View File
@@ -0,0 +1,9 @@
# Vendored FLA KDA kernels
Subset of [flash-linear-attention](https://github.com/fla-org/flash-linear-attention)
used by `kda.ops` `backend="triton"`.
- License: MIT (see `LICENSE`)
- Upstream version tag in `__init__.py`
- Import path is `kda._fla.*`, not `fla.*`
- Not included: context parallel, Ascend, TileLang, `flash_kda`, non-KDA ops
+7
View File
@@ -0,0 +1,7 @@
"""Vendored NVIDIA-Triton KDA path from flash-linear-attention (MIT).
This package is imported as ``kda._fla``, never as the upstream ``fla``
distribution. Context-parallel, Ascend, and TileLang backends are omitted.
"""
__version__ = "0.5.2"
+1
View File
@@ -0,0 +1 @@
# Vendored FLA modules used by KDA (l2norm).
+3
View File
@@ -0,0 +1,3 @@
from kda._fla.ops.backends import BackendRegistry, BaseBackend, dispatch
__all__ = ["BackendRegistry", "BaseBackend", "dispatch"]
+299
View File
@@ -0,0 +1,299 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import torch
import torch.nn as nn
import triton
import triton.language as tl
from kda._fla.modules.backends import dispatch
from kda._fla.ops.utils.cache import fla_cache_autotune
from kda._fla.utils import IS_AMD, autotune_cache_kwargs, input_guard
BT_LIST = [8, 16, 32, 64, 128]
NUM_WARPS_AUTOTUNE = [1, 2, 4, 8, 16] if IS_AMD else [1, 2, 4, 8, 16, 32]
@triton.autotune(
configs=[triton.Config({}, num_warps=num_warps) for num_warps in NUM_WARPS_AUTOTUNE],
key=["D"],
**autotune_cache_kwargs,
)
@triton.jit
def l2norm_fwd_kernel1(
x,
y,
rstd,
eps,
D,
BD: tl.constexpr,
):
i_t = tl.program_id(0).to(tl.int64)
x += i_t * D
y += i_t * D
# Compute mean and variance
cols = tl.arange(0, BD)
mask = cols < D
b_x = tl.load(x + cols, mask=mask, other=0.0).to(tl.float32)
b_rstd = 1 / tl.sqrt(tl.sum(b_x * b_x) + eps)
b_y = b_x * b_rstd
tl.store(y + cols, b_y, mask=mask)
tl.store(rstd + i_t, b_rstd)
@triton.autotune(
configs=[triton.Config({}, num_warps=num_warps) for num_warps in NUM_WARPS_AUTOTUNE],
key=["D"],
**autotune_cache_kwargs,
)
@triton.jit
def l2norm_bwd_kernel1(
y,
rstd,
dy,
dx,
eps,
D,
BD: tl.constexpr,
):
i_t = tl.program_id(0).to(tl.int64)
y += i_t * D
dx += i_t * D
dy += i_t * D
cols = tl.arange(0, BD)
mask = cols < D
b_y = tl.load(y + cols, mask=mask, other=0.0).to(tl.float32)
b_rstd = tl.load(rstd + i_t).to(tl.float32)
b_dy = tl.load(dy + cols, mask=mask, other=0.0).to(tl.float32)
b_dx = b_dy * b_rstd - tl.sum(b_dy * b_y) * b_y * b_rstd
tl.store(dx + cols, b_dx, mask=mask)
@fla_cache_autotune(
configs=[triton.Config({"BT": BT}, num_warps=num_warps) for num_warps in [1, 2, 4, 8, 16] for BT in BT_LIST],
key=["D", "NB"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=["T"])
def l2norm_fwd_kernel(
x,
y,
rstd,
eps,
T,
D: tl.constexpr,
BD: tl.constexpr,
NB: tl.constexpr,
BT: tl.constexpr,
):
i_t = tl.program_id(0).to(tl.int64)
o_t = i_t * BT + tl.arange(0, BT)
o_d = tl.arange(0, BD)
m_t = o_t < T
m_x = m_t[:, None] & (o_d[None, :] < D)
p_x = x + o_t[:, None] * D + o_d[None, :]
p_y = y + o_t[:, None] * D + o_d[None, :]
p_rstd = rstd + o_t
b_x = tl.load(p_x, mask=m_x, other=0.0).to(tl.float32)
b_rstd = 1 / tl.sqrt(tl.sum(b_x * b_x, 1) + eps)
b_y = b_x * b_rstd[:, None]
tl.store(p_y, b_y.to(p_y.dtype.element_ty), mask=m_x)
tl.store(p_rstd, b_rstd.to(p_rstd.dtype.element_ty), mask=m_t)
@fla_cache_autotune(
configs=[triton.Config({"BT": BT}, num_warps=num_warps) for num_warps in [1, 2, 4, 8, 16] for BT in BT_LIST],
key=["D", "NB"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=["T"])
def l2norm_bwd_kernel(
y,
rstd,
dy,
dx,
eps,
T,
D: tl.constexpr,
BD: tl.constexpr,
NB: tl.constexpr,
BT: tl.constexpr,
):
i_t = tl.program_id(0).to(tl.int64)
o_t = i_t * BT + tl.arange(0, BT)
o_d = tl.arange(0, BD)
m_t = o_t < T
m_x = m_t[:, None] & (o_d[None, :] < D)
p_y = y + o_t[:, None] * D + o_d[None, :]
p_rstd = rstd + o_t
p_dy = dy + o_t[:, None] * D + o_d[None, :]
p_dx = dx + o_t[:, None] * D + o_d[None, :]
b_y = tl.load(p_y, mask=m_x, other=0.0).to(tl.float32)
b_rstd = tl.load(p_rstd, mask=m_t, other=0.0).to(tl.float32)
b_dy = tl.load(p_dy, mask=m_x, other=0.0).to(tl.float32)
b_dx = b_dy * b_rstd[:, None] - tl.sum(b_dy * b_y, 1)[:, None] * b_y * b_rstd[:, None]
tl.store(p_dx, b_dx.to(p_dx.dtype.element_ty), mask=m_x)
@dispatch('modules')
def l2norm_fwd(
x: torch.Tensor,
eps: float = 1e-6,
output_dtype: torch.dtype | None = None,
):
x_shape_og = x.shape
x = x.view(-1, x.shape[-1])
# allocate output
if output_dtype is None:
y = torch.empty_like(x)
else:
y = torch.empty_like(x, dtype=output_dtype)
assert y.stride(-1) == 1
T, D = x.shape[0], x.shape[-1]
# Less than 64KB per feature: enqueue fused kernel
MAX_FUSED_SIZE = 65536 // x.element_size()
BD = min(MAX_FUSED_SIZE, triton.next_power_of_2(D))
if D > BD:
raise RuntimeError("This layer doesn't support feature dim >= 64KB.")
rstd = torch.empty((T,), dtype=torch.float32, device=x.device)
if D <= 512:
# NOTE(tylerr): Avoid excessive recompilation and autotuning by tolerating a larger range
# of T before recompiling the kernel.
# NB = triton.cdiv(T, 2048)
NB = triton.cdiv(T, 2048 * 32)
def grid(meta):
return (triton.cdiv(T, meta["BT"]),)
l2norm_fwd_kernel[grid](
x=x,
y=y,
rstd=rstd,
eps=eps,
T=T,
D=D,
BD=BD,
NB=NB,
)
else:
l2norm_fwd_kernel1[(T,)](
x=x,
y=y,
rstd=rstd,
eps=eps,
D=D,
BD=BD,
)
return y.view(x_shape_og), rstd.view(x_shape_og[:-1])
@dispatch('modules')
def l2norm_bwd(
y: torch.Tensor,
rstd: torch.Tensor,
dy: torch.Tensor,
eps: float = 1e-6,
):
y_shape_og = y.shape
y = y.view(-1, dy.shape[-1])
dy = dy.view(-1, dy.shape[-1])
assert dy.shape == y.shape
# allocate output
dx = torch.empty_like(y)
T, D = y.shape[0], y.shape[-1]
# Less than 64KB per feature: enqueue fused kernel
MAX_FUSED_SIZE = 65536 // y.element_size()
BD = min(MAX_FUSED_SIZE, triton.next_power_of_2(D))
if D > BD:
raise RuntimeError("This layer norm doesn't support feature dim >= 64KB.")
if D <= 512:
# NOTE(tylerr): Avoid excessive recompilation and autotuning by tolerating a larger range
# of T before recompiling the kernel.
# NB = triton.cdiv(T, 2048)
NB = triton.cdiv(T, 2048 * 32)
def grid(meta):
return (triton.cdiv(T, meta["BT"]),)
l2norm_bwd_kernel[grid](
y=y,
rstd=rstd,
dy=dy,
dx=dx,
eps=eps,
T=T,
D=D,
BD=BD,
NB=NB,
)
else:
l2norm_bwd_kernel1[(T,)](
y=y,
rstd=rstd,
dy=dy,
dx=dx,
eps=eps,
D=D,
BD=BD,
)
return dx.view(y_shape_og)
class L2NormFunction(torch.autograd.Function):
@staticmethod
@input_guard
def forward(
ctx,
x,
eps=1e-6,
output_dtype=None,
):
y, rstd = l2norm_fwd(x, eps, output_dtype)
ctx.eps = eps
ctx.x_dtype = x.dtype
ctx.save_for_backward(y, rstd)
return y
@staticmethod
@input_guard
def backward(ctx, dy):
y, rstd = ctx.saved_tensors
dx = l2norm_bwd(y, rstd, dy, ctx.eps)
return dx, None, None
def l2norm(
x: torch.Tensor,
eps: float = 1e-6,
output_dtype: torch.dtype | None = None,
) -> torch.Tensor:
return L2NormFunction.apply(x, eps, output_dtype)
l2_norm = l2norm
class L2Norm(nn.Module):
def __init__(
self,
eps: float = 1e-6,
output_dtype: torch.dtype | None = None,
):
super().__init__()
self.eps = eps
self.output_dtype = output_dtype
def forward(self, x: torch.Tensor) -> torch.Tensor:
return l2norm(x, self.eps, self.output_dtype)
+1
View File
@@ -0,0 +1 @@
# Vendored FLA ops subset.
+34
View File
@@ -0,0 +1,34 @@
"""Identity dispatch: keep the NVIDIA Triton implementation in this tree."""
from __future__ import annotations
from collections.abc import Callable
from typing import TypeVar
F = TypeVar("F", bound=Callable)
def dispatch(operation: str):
def decorator(func: F) -> F:
return func
return decorator
class BaseBackend:
backend_type = "triton"
def is_available(self) -> bool:
return True
def is_enabled(self) -> bool:
return True
class BackendRegistry:
@classmethod
def ensure_initialized(cls, operation: str) -> None:
return None
__all__ = ["BackendRegistry", "BaseBackend", "dispatch"]
+1
View File
@@ -0,0 +1 @@
# Vendored FLA common kernels used by KDA.
+806
View File
@@ -0,0 +1,806 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import torch
import triton
import triton.language as tl
from kda._fla.ops.backends import dispatch
from kda._fla.ops.utils import prepare_chunk_indices, prepare_chunk_offsets
from kda._fla.ops.utils.cache import fla_cache_autotune
from kda._fla.ops.utils.op import exp2
from kda._fla.utils import (
IS_INTEL,
IS_NVIDIA_BLACKWELL,
IS_NVIDIA_HOPPER,
autotune_cache_kwargs,
check_shared_mem,
)
NUM_WARPS = [2, 4] if IS_NVIDIA_HOPPER else [2, 4, 8, 16]
# TODO: Triton mainline fixes a Blackwell tl.dot recurrence race.
# Keep this kernel on num_warps=2 for Blackwell until Triton 3.8 is released
# and we re-validate the wider config space.
# Intel needs more warps than NVIDIA here: 8 warps is ~1.5x faster than the best
# config reachable under the [2, 4] cap.
if IS_NVIDIA_BLACKWELL:
GATED_DELTA_RULE_FWD_H_NUM_WARPS = [2]
elif IS_INTEL:
GATED_DELTA_RULE_FWD_H_NUM_WARPS = [2, 4, 8, 16]
else:
GATED_DELTA_RULE_FWD_H_NUM_WARPS = [2, 4]
@triton.heuristics({
'USE_G': lambda args: args['g'] is not None,
'USE_GK': lambda args: args['gk'] is not None,
'USE_INITIAL_STATE': lambda args: args['h0'] is not None,
'STORE_FINAL_STATE': lambda args: args['ht'] is not None,
'SAVE_NEW_VALUE': lambda args: args['v_new'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({'BV': BV}, num_warps=num_warps, num_stages=num_stages)
for num_warps in GATED_DELTA_RULE_FWD_H_NUM_WARPS
for num_stages in ([2, 3, 4] if check_shared_mem('ampere') else [2, 1])
for BV in ([32, 64] if check_shared_mem('ada') else [32])
],
key=['H', 'HV', 'K', 'V', 'BT', 'STATE_V_FIRST'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_gated_delta_rule_fwd_kernel_h_blockdim64(
k,
v,
w,
v_new,
g,
gk,
h,
h0,
ht,
cu_seqlens,
chunk_offsets,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BT: tl.constexpr,
BV: tl.constexpr,
USE_G: tl.constexpr,
USE_GK: tl.constexpr,
USE_INITIAL_STATE: tl.constexpr,
STORE_FINAL_STATE: tl.constexpr,
SAVE_NEW_VALUE: tl.constexpr,
STATE_V_FIRST: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
pid = tl.program_id(0)
NV = tl.cdiv(V, BV)
i_v, i_nh = pid % NV, (pid // NV).to(tl.int64)
i_n, i_h = i_nh // HV, i_nh % HV
if IS_VARLEN:
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
NT = tl.cdiv(T, BT)
boh = tl.load(chunk_offsets + i_n).to(tl.int64)
else:
bos, eos = i_n * T, i_n * T + T
NT = tl.cdiv(T, BT)
boh = i_n * NT
if STATE_V_FIRST:
b_h1 = tl.zeros([BV, 64], dtype=tl.float32)
if K > 64:
b_h2 = tl.zeros([BV, 64], dtype=tl.float32)
if K > 128:
b_h3 = tl.zeros([BV, 64], dtype=tl.float32)
if K > 192:
b_h4 = tl.zeros([BV, 64], dtype=tl.float32)
else:
b_h1 = tl.zeros([64, BV], dtype=tl.float32)
if K > 64:
b_h2 = tl.zeros([64, BV], dtype=tl.float32)
if K > 128:
b_h3 = tl.zeros([64, BV], dtype=tl.float32)
if K > 192:
b_h4 = tl.zeros([64, BV], dtype=tl.float32)
# calculate offset
h += (boh * HV + i_h).to(tl.int64) * K*V
v += (bos * HV + i_h).to(tl.int64) * V
k += (bos * H + i_h // (HV // H)).to(tl.int64) * K
w += (bos * HV + i_h).to(tl.int64) * K
if SAVE_NEW_VALUE:
v_new += (bos * HV + i_h).to(tl.int64) * V
if USE_INITIAL_STATE:
h0 = h0 + i_nh * K*V
if STORE_FINAL_STATE:
ht = ht + i_nh * K*V
# load initial state
o_v = i_v * BV + tl.arange(0, BV)
m_v = o_v < V
o_k1 = tl.arange(0, 64)
m_k1 = o_k1 < K
o_k2 = 64 + o_k1
m_k2 = o_k2 < K
o_k3 = 128 + o_k1
m_k3 = o_k3 < K
o_k4 = 192 + o_k1
m_k4 = o_k4 < K
if USE_INITIAL_STATE:
if STATE_V_FIRST:
p_h0_1 = h0 + o_v[:, None] * K + o_k1[None, :]
m_h0_1 = m_v[:, None] & m_k1[None, :]
else:
p_h0_1 = h0 + o_k1[:, None] * V + o_v[None, :]
m_h0_1 = m_k1[:, None] & m_v[None, :]
b_h1 += tl.load(p_h0_1, mask=m_h0_1, other=0.0).to(tl.float32)
if K > 64:
if STATE_V_FIRST:
p_h0_2 = h0 + o_v[:, None] * K + o_k2[None, :]
m_h0_2 = m_v[:, None] & m_k2[None, :]
else:
p_h0_2 = h0 + o_k2[:, None] * V + o_v[None, :]
m_h0_2 = m_k2[:, None] & m_v[None, :]
b_h2 += tl.load(p_h0_2, mask=m_h0_2, other=0.0).to(tl.float32)
if K > 128:
if STATE_V_FIRST:
p_h0_3 = h0 + o_v[:, None] * K + o_k3[None, :]
m_h0_3 = m_v[:, None] & m_k3[None, :]
else:
p_h0_3 = h0 + o_k3[:, None] * V + o_v[None, :]
m_h0_3 = m_k3[:, None] & m_v[None, :]
b_h3 += tl.load(p_h0_3, mask=m_h0_3, other=0.0).to(tl.float32)
if K > 192:
if STATE_V_FIRST:
p_h0_4 = h0 + o_v[:, None] * K + o_k4[None, :]
m_h0_4 = m_v[:, None] & m_k4[None, :]
else:
p_h0_4 = h0 + o_k4[:, None] * V + o_v[None, :]
m_h0_4 = m_k4[:, None] & m_v[None, :]
b_h4 += tl.load(p_h0_4, mask=m_h0_4, other=0.0).to(tl.float32)
# main recurrence
for i_t in range(NT):
i_t_int64 = i_t.to(tl.int64)
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
if STATE_V_FIRST:
p_h1 = h + i_t_int64 * HV*K*V + o_v[:, None] * K + o_k1[None, :]
m_h1 = m_v[:, None] & m_k1[None, :]
else:
p_h1 = h + i_t_int64 * HV*K*V + o_k1[:, None] * V + o_v[None, :]
m_h1 = m_k1[:, None] & m_v[None, :]
tl.store(p_h1, b_h1.to(p_h1.dtype.element_ty), mask=m_h1)
if K > 64:
if STATE_V_FIRST:
p_h2 = h + i_t_int64 * HV*K*V + o_v[:, None] * K + o_k2[None, :]
m_h2 = m_v[:, None] & m_k2[None, :]
else:
p_h2 = h + i_t_int64 * HV*K*V + o_k2[:, None] * V + o_v[None, :]
m_h2 = m_k2[:, None] & m_v[None, :]
tl.store(p_h2, b_h2.to(p_h2.dtype.element_ty), mask=m_h2)
if K > 128:
if STATE_V_FIRST:
p_h3 = h + i_t_int64 * HV*K*V + o_v[:, None] * K + o_k3[None, :]
m_h3 = m_v[:, None] & m_k3[None, :]
else:
p_h3 = h + i_t_int64 * HV*K*V + o_k3[:, None] * V + o_v[None, :]
m_h3 = m_k3[:, None] & m_v[None, :]
tl.store(p_h3, b_h3.to(p_h3.dtype.element_ty), mask=m_h3)
if K > 192:
if STATE_V_FIRST:
p_h4 = h + i_t_int64 * HV*K*V + o_v[:, None] * K + o_k4[None, :]
m_h4 = m_v[:, None] & m_k4[None, :]
else:
p_h4 = h + i_t_int64 * HV*K*V + o_k4[:, None] * V + o_v[None, :]
m_h4 = m_k4[:, None] & m_v[None, :]
tl.store(p_h4, b_h4.to(p_h4.dtype.element_ty), mask=m_h4)
p_w = w + o_t[:, None] * (HV*K) + o_k1[None, :]
b_w = tl.load(p_w, mask=m_t[:, None] & m_k1[None, :], other=0.0)
if STATE_V_FIRST:
b_v = tl.dot(b_w, tl.trans(b_h1).to(b_w.dtype))
else:
b_v = tl.dot(b_w, b_h1.to(b_w.dtype))
if K > 64:
p_w = w + o_t[:, None] * (HV*K) + o_k2[None, :]
b_w = tl.load(p_w, mask=m_t[:, None] & m_k2[None, :], other=0.0)
if STATE_V_FIRST:
b_v += tl.dot(b_w, tl.trans(b_h2).to(b_w.dtype))
else:
b_v += tl.dot(b_w, b_h2.to(b_w.dtype))
if K > 128:
p_w = w + o_t[:, None] * (HV*K) + o_k3[None, :]
b_w = tl.load(p_w, mask=m_t[:, None] & m_k3[None, :], other=0.0)
if STATE_V_FIRST:
b_v += tl.dot(b_w, tl.trans(b_h3).to(b_w.dtype))
else:
b_v += tl.dot(b_w, b_h3.to(b_w.dtype))
if K > 192:
p_w = w + o_t[:, None] * (HV*K) + o_k4[None, :]
b_w = tl.load(p_w, mask=m_t[:, None] & m_k4[None, :], other=0.0)
if STATE_V_FIRST:
b_v += tl.dot(b_w, tl.trans(b_h4).to(b_w.dtype))
else:
b_v += tl.dot(b_w, b_h4.to(b_w.dtype))
p_v = v + o_t[:, None] * (HV*V) + o_v[None, :]
b_v = tl.load(p_v, mask=m_t[:, None] & m_v[None, :], other=0.0) - b_v
if SAVE_NEW_VALUE:
p_v = v_new + o_t[:, None] * (HV*V) + o_v[None, :]
tl.store(p_v, b_v.to(p_v.dtype.element_ty), mask=m_t[:, None] & m_v[None, :])
last_idx = min((i_t + 1) * BT, T) - 1
if USE_G:
b_g_last = tl.load(g + (bos * HV + last_idx * HV + i_h).to(tl.int64)).to(tl.float32)
p_g = g + (bos * HV + i_h).to(tl.int64) + o_t * HV
b_g = tl.load(p_g, mask=m_t, other=0.0).to(tl.float32)
b_v = b_v * tl.where(m_t, exp2(b_g_last - b_g), 0)[:, None]
b_g_last = exp2(b_g_last)
b_h1 *= b_g_last
if K > 64:
b_h2 *= b_g_last
if K > 128:
b_h3 *= b_g_last
if K > 192:
b_h4 *= b_g_last
if USE_GK:
o_k1 = tl.arange(0, 64)
b_gk_last1 = tl.load(gk + (bos + last_idx) * HV*K + i_h * K + o_k1, mask=(o_k1 < K), other=0.).to(tl.float32)
if STATE_V_FIRST:
b_h1 *= exp2(b_gk_last1)[None, :]
else:
b_h1 *= exp2(b_gk_last1)[:, None]
if K > 64:
o_k2 = 64 + o_k1
b_gk_last2 = tl.load(gk + (bos + last_idx) * HV*K + i_h * K + o_k2, mask=(o_k2 < K), other=0.).to(tl.float32)
if STATE_V_FIRST:
b_h2 *= exp2(b_gk_last2)[None, :]
else:
b_h2 *= exp2(b_gk_last2)[:, None]
if K > 128:
o_k3 = 128 + o_k1
b_gk_last3 = tl.load(gk + (bos + last_idx) * HV*K + i_h * K + o_k3, mask=(o_k3 < K), other=0.).to(tl.float32)
if STATE_V_FIRST:
b_h3 *= exp2(b_gk_last3)[None, :]
else:
b_h3 *= exp2(b_gk_last3)[:, None]
if K > 192:
o_k4 = 192 + o_k1
b_gk_last4 = tl.load(gk + (bos + last_idx) * HV*K + i_h * K + o_k4, mask=(o_k4 < K), other=0.).to(tl.float32)
if STATE_V_FIRST:
b_h4 *= exp2(b_gk_last4)[None, :]
else:
b_h4 *= exp2(b_gk_last4)[:, None]
b_v = b_v.to(k.dtype.element_ty)
p_k = k + o_k1[:, None] + o_t[None, :] * (H*K)
b_k = tl.load(p_k, mask=m_k1[:, None] & m_t[None, :], other=0.0)
if STATE_V_FIRST:
b_h1 += tl.trans(tl.dot(b_k, b_v))
else:
b_h1 += tl.dot(b_k, b_v)
if K > 64:
p_k = k + o_k2[:, None] + o_t[None, :] * (H*K)
b_k = tl.load(p_k, mask=m_k2[:, None] & m_t[None, :], other=0.0)
if STATE_V_FIRST:
b_h2 += tl.trans(tl.dot(b_k, b_v))
else:
b_h2 += tl.dot(b_k, b_v)
if K > 128:
p_k = k + o_k3[:, None] + o_t[None, :] * (H*K)
b_k = tl.load(p_k, mask=m_k3[:, None] & m_t[None, :], other=0.0)
if STATE_V_FIRST:
b_h3 += tl.trans(tl.dot(b_k, b_v))
else:
b_h3 += tl.dot(b_k, b_v)
if K > 192:
p_k = k + o_k4[:, None] + o_t[None, :] * (H*K)
b_k = tl.load(p_k, mask=m_k4[:, None] & m_t[None, :], other=0.0)
if STATE_V_FIRST:
b_h4 += tl.trans(tl.dot(b_k, b_v))
else:
b_h4 += tl.dot(b_k, b_v)
if STORE_FINAL_STATE:
if STATE_V_FIRST:
p_ht = ht + o_v[:, None] * K + o_k1[None, :]
m_ht = m_v[:, None] & m_k1[None, :]
else:
p_ht = ht + o_k1[:, None] * V + o_v[None, :]
m_ht = m_k1[:, None] & m_v[None, :]
tl.store(p_ht, b_h1.to(p_ht.dtype.element_ty), mask=m_ht)
if K > 64:
if STATE_V_FIRST:
p_ht = ht + o_v[:, None] * K + o_k2[None, :]
m_ht = m_v[:, None] & m_k2[None, :]
else:
p_ht = ht + o_k2[:, None] * V + o_v[None, :]
m_ht = m_k2[:, None] & m_v[None, :]
tl.store(p_ht, b_h2.to(p_ht.dtype.element_ty), mask=m_ht)
if K > 128:
if STATE_V_FIRST:
p_ht = ht + o_v[:, None] * K + o_k3[None, :]
m_ht = m_v[:, None] & m_k3[None, :]
else:
p_ht = ht + o_k3[:, None] * V + o_v[None, :]
m_ht = m_k3[:, None] & m_v[None, :]
tl.store(p_ht, b_h3.to(p_ht.dtype.element_ty), mask=m_ht)
if K > 192:
if STATE_V_FIRST:
p_ht = ht + o_v[:, None] * K + o_k4[None, :]
m_ht = m_v[:, None] & m_k4[None, :]
else:
p_ht = ht + o_k4[:, None] * V + o_v[None, :]
m_ht = m_k4[:, None] & m_v[None, :]
tl.store(p_ht, b_h4.to(p_ht.dtype.element_ty), mask=m_ht)
@triton.heuristics({
'USE_G': lambda args: args['g'] is not None,
'USE_GK': lambda args: args['gk'] is not None,
'USE_INITIAL_STATE': lambda args: args['dh0'] is not None,
'USE_FINAL_STATE_GRADIENT': lambda args: args['dht'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({'BV': BV}, num_warps=num_warps, num_stages=num_stages)
for num_warps in [2, 4]
for num_stages in ([2, 3, 4] if check_shared_mem('ampere') else [1])
for BV in ([32, 64] if check_shared_mem('ada') else [32])
],
key=['H', 'HV', 'K', 'V', 'BT', 'BV', 'USE_G', 'STATE_V_FIRST'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_gated_delta_rule_bwd_kernel_dhu_blockdim64(
q,
k,
w,
g,
gk,
dht,
dh0,
do,
dh,
dv,
dv2,
cu_seqlens,
chunk_offsets,
scale,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BT: tl.constexpr,
BV: tl.constexpr,
USE_G: tl.constexpr,
USE_GK: tl.constexpr,
USE_INITIAL_STATE: tl.constexpr,
USE_FINAL_STATE_GRADIENT: tl.constexpr,
STATE_V_FIRST: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
pid = tl.program_id(0)
NV = tl.cdiv(V, BV)
i_v, i_nh = pid % NV, (pid // NV).to(tl.int64)
i_n, i_h = i_nh // HV, i_nh % HV
if IS_VARLEN:
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
NT = tl.cdiv(T, BT)
boh = tl.load(chunk_offsets + i_n).to(tl.int64)
else:
bos, eos = i_n * T, i_n * T + T
NT = tl.cdiv(T, BT)
boh = i_n * NT
if STATE_V_FIRST:
b_dh1 = tl.zeros([BV, 64], dtype=tl.float32)
if K > 64:
b_dh2 = tl.zeros([BV, 64], dtype=tl.float32)
if K > 128:
b_dh3 = tl.zeros([BV, 64], dtype=tl.float32)
if K > 192:
b_dh4 = tl.zeros([BV, 64], dtype=tl.float32)
else:
b_dh1 = tl.zeros([64, BV], dtype=tl.float32)
if K > 64:
b_dh2 = tl.zeros([64, BV], dtype=tl.float32)
if K > 128:
b_dh3 = tl.zeros([64, BV], dtype=tl.float32)
if K > 192:
b_dh4 = tl.zeros([64, BV], dtype=tl.float32)
# calculate offset
q += (bos * H + i_h // (HV // H)).to(tl.int64) * K
k += (bos * H + i_h // (HV // H)).to(tl.int64) * K
w += (bos * HV + i_h).to(tl.int64) * K
do += (bos * HV + i_h).to(tl.int64) * V
dv += (bos * HV + i_h).to(tl.int64) * V
dv2 += (bos * HV + i_h).to(tl.int64) * V
dh += (boh * HV + i_h).to(tl.int64) * K*V
if USE_GK:
gk += (bos * HV + i_h).to(tl.int64) * K
if USE_INITIAL_STATE:
dh0 += i_nh * K*V
if USE_FINAL_STATE_GRADIENT:
dht += i_nh * K*V
o_v = i_v * BV + tl.arange(0, BV)
m_v = o_v < V
o_k1 = tl.arange(0, 64)
m_k1 = o_k1 < K
o_k2 = 64 + o_k1
m_k2 = o_k2 < K
o_k3 = 128 + o_k1
m_k3 = o_k3 < K
o_k4 = 192 + o_k1
m_k4 = o_k4 < K
if USE_FINAL_STATE_GRADIENT:
if STATE_V_FIRST:
p_dht1 = dht + o_v[:, None] * K + o_k1[None, :]
m_dht1 = m_v[:, None] & m_k1[None, :]
else:
p_dht1 = dht + o_k1[:, None] * V + o_v[None, :]
m_dht1 = m_k1[:, None] & m_v[None, :]
b_dh1 += tl.load(p_dht1, mask=m_dht1, other=0.0)
if K > 64:
if STATE_V_FIRST:
p_dht2 = dht + o_v[:, None] * K + o_k2[None, :]
m_dht2 = m_v[:, None] & m_k2[None, :]
else:
p_dht2 = dht + o_k2[:, None] * V + o_v[None, :]
m_dht2 = m_k2[:, None] & m_v[None, :]
b_dh2 += tl.load(p_dht2, mask=m_dht2, other=0.0)
if K > 128:
if STATE_V_FIRST:
p_dht3 = dht + o_v[:, None] * K + o_k3[None, :]
m_dht3 = m_v[:, None] & m_k3[None, :]
else:
p_dht3 = dht + o_k3[:, None] * V + o_v[None, :]
m_dht3 = m_k3[:, None] & m_v[None, :]
b_dh3 += tl.load(p_dht3, mask=m_dht3, other=0.0)
if K > 192:
if STATE_V_FIRST:
p_dht4 = dht + o_v[:, None] * K + o_k4[None, :]
m_dht4 = m_v[:, None] & m_k4[None, :]
else:
p_dht4 = dht + o_k4[:, None] * V + o_v[None, :]
m_dht4 = m_k4[:, None] & m_v[None, :]
b_dh4 += tl.load(p_dht4, mask=m_dht4, other=0.0)
for i_t in range(NT - 1, -1, -1):
i_t_int64 = i_t.to(tl.int64)
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
if STATE_V_FIRST:
p_dh1 = dh + i_t_int64*HV*K*V + o_v[:, None] * K + o_k1[None, :]
m_dh1 = m_v[:, None] & m_k1[None, :]
else:
p_dh1 = dh + i_t_int64*HV*K*V + o_k1[:, None] * V + o_v[None, :]
m_dh1 = m_k1[:, None] & m_v[None, :]
tl.store(p_dh1, b_dh1.to(p_dh1.dtype.element_ty), mask=m_dh1)
if K > 64:
if STATE_V_FIRST:
p_dh2 = dh + i_t_int64*HV*K*V + o_v[:, None] * K + o_k2[None, :]
m_dh2 = m_v[:, None] & m_k2[None, :]
else:
p_dh2 = dh + i_t_int64*HV*K*V + o_k2[:, None] * V + o_v[None, :]
m_dh2 = m_k2[:, None] & m_v[None, :]
tl.store(p_dh2, b_dh2.to(p_dh2.dtype.element_ty), mask=m_dh2)
if K > 128:
if STATE_V_FIRST:
p_dh3 = dh + i_t_int64*HV*K*V + o_v[:, None] * K + o_k3[None, :]
m_dh3 = m_v[:, None] & m_k3[None, :]
else:
p_dh3 = dh + i_t_int64*HV*K*V + o_k3[:, None] * V + o_v[None, :]
m_dh3 = m_k3[:, None] & m_v[None, :]
tl.store(p_dh3, b_dh3.to(p_dh3.dtype.element_ty), mask=m_dh3)
if K > 192:
if STATE_V_FIRST:
p_dh4 = dh + i_t_int64*HV*K*V + o_v[:, None] * K + o_k4[None, :]
m_dh4 = m_v[:, None] & m_k4[None, :]
else:
p_dh4 = dh + i_t_int64*HV*K*V + o_k4[:, None] * V + o_v[None, :]
m_dh4 = m_k4[:, None] & m_v[None, :]
tl.store(p_dh4, b_dh4.to(p_dh4.dtype.element_ty), mask=m_dh4)
last_idx = min((i_t + 1) * BT, T) - 1
if USE_G:
bg_last = tl.load(g + (bos + last_idx) * HV + i_h).to(tl.float32)
p_g = g + bos * HV + i_h + o_t * HV
b_g = tl.load(p_g, mask=m_t, other=0.0).to(tl.float32)
bg_last_exp = exp2(bg_last)
b_g_exp = exp2(b_g)
p_dv = dv + o_t[:, None] * (HV*V) + o_v[None, :]
p_dv2 = dv2 + o_t[:, None] * (HV*V) + o_v[None, :]
p_do = do + o_t[:, None] * (HV*V) + o_v[None, :]
b_do = tl.load(p_do, mask=m_t[:, None] & m_v[None, :], other=0.0)
# Update dv
p_k = k + o_t[:, None] * (H*K) + o_k1[None, :]
b_k = tl.load(p_k, mask=m_t[:, None] & m_k1[None, :], other=0.0)
if USE_GK:
o_k1 = tl.arange(0, 64)
b_gk_last1 = tl.load(gk + last_idx * HV*K + o_k1, mask=(o_k1 < K), other=0.).to(tl.float32)
if STATE_V_FIRST:
b_dv = tl.dot(b_k, tl.trans(b_dh1).to(b_k.dtype))
else:
b_dv = tl.dot(b_k, b_dh1.to(b_k.dtype))
if K > 64:
p_k = k + o_t[:, None] * (H*K) + o_k2[None, :]
b_k = tl.load(p_k, mask=m_t[:, None] & m_k2[None, :], other=0.0)
if USE_GK:
b_gk_last2 = tl.load(gk + last_idx * HV*K + o_k2, mask=(o_k2 < K), other=0.).to(tl.float32)
if STATE_V_FIRST:
b_dv += tl.dot(b_k, tl.trans(b_dh2).to(b_k.dtype))
else:
b_dv += tl.dot(b_k, b_dh2.to(b_k.dtype))
if K > 128:
p_k = k + o_t[:, None] * (H*K) + o_k3[None, :]
b_k = tl.load(p_k, mask=m_t[:, None] & m_k3[None, :], other=0.0)
if USE_GK:
b_gk_last3 = tl.load(gk + last_idx * HV*K + o_k3, mask=(o_k3 < K), other=0.).to(tl.float32)
if STATE_V_FIRST:
b_dv += tl.dot(b_k, tl.trans(b_dh3).to(b_k.dtype))
else:
b_dv += tl.dot(b_k, b_dh3.to(b_k.dtype))
if K > 192:
p_k = k + o_t[:, None] * (H*K) + o_k4[None, :]
b_k = tl.load(p_k, mask=m_t[:, None] & m_k4[None, :], other=0.0)
if USE_GK:
b_gk_last4 = tl.load(gk + last_idx * HV*K + o_k4, mask=(o_k4 < K), other=0.).to(tl.float32)
if STATE_V_FIRST:
b_dv += tl.dot(b_k, tl.trans(b_dh4).to(b_k.dtype))
else:
b_dv += tl.dot(b_k, b_dh4.to(b_k.dtype))
if USE_G:
b_dv *= tl.where(m_t, exp2(bg_last - b_g), 0)[:, None]
b_dv += tl.load(p_dv, mask=m_t[:, None] & m_v[None, :], other=0.0)
tl.store(p_dv2, b_dv.to(p_dv.dtype.element_ty), mask=m_t[:, None] & m_v[None, :])
# Update dh
p_w = w + o_k1[:, None] + o_t[None, :] * (HV*K)
p_q = q + o_k1[:, None] + o_t[None, :] * (H*K)
b_w = tl.load(p_w, mask=m_k1[:, None] & m_t[None, :], other=0.0)
b_q = tl.load(p_q, mask=m_k1[:, None] & m_t[None, :], other=0.0)
if USE_G:
b_dh1 *= bg_last_exp
b_q = b_q * b_g_exp[None, :]
if USE_GK:
if STATE_V_FIRST:
b_dh1 *= exp2(b_gk_last1)[None, :]
else:
b_dh1 *= exp2(b_gk_last1[:, None])
if STATE_V_FIRST:
b_dh1 += tl.trans(tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype)))
else:
b_dh1 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype))
if K > 64:
p_q = q + o_k2[:, None] + o_t[None, :] * (H*K)
p_w = w + o_k2[:, None] + o_t[None, :] * (HV*K)
b_q = tl.load(p_q, mask=m_k2[:, None] & m_t[None, :], other=0.0)
b_w = tl.load(p_w, mask=m_k2[:, None] & m_t[None, :], other=0.0)
if USE_G:
b_dh2 *= bg_last_exp
b_q = b_q * b_g_exp[None, :]
if USE_GK:
if STATE_V_FIRST:
b_dh2 *= exp2(b_gk_last2)[None, :]
else:
b_dh2 *= exp2(b_gk_last2[:, None])
if STATE_V_FIRST:
b_dh2 += tl.trans(tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype)))
else:
b_dh2 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype))
if K > 128:
p_q = q + o_k3[:, None] + o_t[None, :] * (H*K)
p_w = w + o_k3[:, None] + o_t[None, :] * (HV*K)
b_q = tl.load(p_q, mask=m_k3[:, None] & m_t[None, :], other=0.0)
b_w = tl.load(p_w, mask=m_k3[:, None] & m_t[None, :], other=0.0)
if USE_G:
b_dh3 *= bg_last_exp
b_q = b_q * b_g_exp[None, :]
if USE_GK:
if STATE_V_FIRST:
b_dh3 *= exp2(b_gk_last3)[None, :]
else:
b_dh3 *= exp2(b_gk_last3[:, None])
if STATE_V_FIRST:
b_dh3 += tl.trans(tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype)))
else:
b_dh3 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype))
if K > 192:
p_q = q + o_k4[:, None] + o_t[None, :] * (H*K)
p_w = w + o_k4[:, None] + o_t[None, :] * (HV*K)
b_q = tl.load(p_q, mask=m_k4[:, None] & m_t[None, :], other=0.0)
b_w = tl.load(p_w, mask=m_k4[:, None] & m_t[None, :], other=0.0)
if USE_G:
b_dh4 *= bg_last_exp
b_q = b_q * b_g_exp[None, :]
if USE_GK:
if STATE_V_FIRST:
b_dh4 *= exp2(b_gk_last4)[None, :]
else:
b_dh4 *= exp2(b_gk_last4[:, None])
if STATE_V_FIRST:
b_dh4 += tl.trans(tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype)))
else:
b_dh4 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype))
if USE_INITIAL_STATE:
if STATE_V_FIRST:
p_dh0 = dh0 + o_v[:, None] * K + o_k1[None, :]
m_dh0 = m_v[:, None] & m_k1[None, :]
else:
p_dh0 = dh0 + o_k1[:, None] * V + o_v[None, :]
m_dh0 = m_k1[:, None] & m_v[None, :]
tl.store(p_dh0, b_dh1.to(p_dh0.dtype.element_ty), mask=m_dh0)
if K > 64:
if STATE_V_FIRST:
p_dh1 = dh0 + o_v[:, None] * K + o_k2[None, :]
m_dh1 = m_v[:, None] & m_k2[None, :]
else:
p_dh1 = dh0 + o_k2[:, None] * V + o_v[None, :]
m_dh1 = m_k2[:, None] & m_v[None, :]
tl.store(p_dh1, b_dh2.to(p_dh1.dtype.element_ty), mask=m_dh1)
if K > 128:
if STATE_V_FIRST:
p_dh2 = dh0 + o_v[:, None] * K + o_k3[None, :]
m_dh2 = m_v[:, None] & m_k3[None, :]
else:
p_dh2 = dh0 + o_k3[:, None] * V + o_v[None, :]
m_dh2 = m_k3[:, None] & m_v[None, :]
tl.store(p_dh2, b_dh3.to(p_dh2.dtype.element_ty), mask=m_dh2)
if K > 192:
if STATE_V_FIRST:
p_dh3 = dh0 + o_v[:, None] * K + o_k4[None, :]
m_dh3 = m_v[:, None] & m_k4[None, :]
else:
p_dh3 = dh0 + o_k4[:, None] * V + o_v[None, :]
m_dh3 = m_k4[:, None] & m_v[None, :]
tl.store(p_dh3, b_dh4.to(p_dh3.dtype.element_ty), mask=m_dh3)
@dispatch('common')
def chunk_gated_delta_rule_fwd_h(
k: torch.Tensor,
w: torch.Tensor,
u: torch.Tensor,
g: torch.Tensor | None = None,
gk: torch.Tensor | None = None,
initial_state: torch.Tensor | None = None,
output_final_state: bool = False,
chunk_size: int = 64,
save_new_value: bool = True,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
cu_seqlens_cpu: torch.LongTensor | None = None,
chunk_indices: torch.LongTensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None]:
B, T, H, K, V, HV = *k.shape, u.shape[-1], u.shape[2]
BT = chunk_size
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, chunk_size)
# N: the actual number of sequences in the batch with either equal or variable lengths
if cu_seqlens is None:
N, NT, chunk_offsets = B, triton.cdiv(T, BT), None
else:
N, NT, chunk_offsets = len(cu_seqlens) - 1, len(chunk_indices), prepare_chunk_offsets(cu_seqlens, BT)
assert K <= 256, "current kernel does not support head dimension larger than 256."
if state_v_first:
h = k.new_empty(B, NT, HV, V, K)
final_state = k.new_zeros(N, HV, V, K, dtype=torch.float32) if output_final_state else None
else:
h = k.new_empty(B, NT, HV, K, V)
final_state = k.new_zeros(N, HV, K, V, dtype=torch.float32) if output_final_state else None
v_new = torch.empty_like(u) if save_new_value else None
def grid(meta): return (triton.cdiv(V, meta['BV']) * N * HV, )
chunk_gated_delta_rule_fwd_kernel_h_blockdim64[grid](
k=k,
v=u,
w=w,
v_new=v_new,
g=g,
gk=gk,
h=h,
h0=initial_state,
ht=final_state,
cu_seqlens=cu_seqlens,
chunk_offsets=chunk_offsets,
T=T,
H=H,
HV=HV,
K=K,
V=V,
BT=BT,
STATE_V_FIRST=state_v_first,
)
return h, v_new, final_state
@dispatch('common')
def chunk_gated_delta_rule_bwd_dhu(
q: torch.Tensor,
k: torch.Tensor,
w: torch.Tensor,
do: torch.Tensor,
dv: torch.Tensor,
g: torch.Tensor | None = None,
gk: torch.Tensor | None = None,
h0: torch.Tensor | None = None,
dht: torch.Tensor | None = None,
scale: float | None = None,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
chunk_size: int = 64,
chunk_indices: torch.LongTensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
B, T, H, K, V, HV = *q.shape, do.shape[-1], do.shape[2]
# N: the actual number of sequences in the batch with either equal or variable lengths
BT = chunk_size
assert K <= 256, "current kernel does not support head dimension being larger than 256."
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, chunk_size)
if cu_seqlens is None:
N, NT, chunk_offsets = B, triton.cdiv(T, BT), None
else:
N, NT, chunk_offsets = len(cu_seqlens) - 1, len(chunk_indices), prepare_chunk_offsets(cu_seqlens, BT)
if state_v_first:
dh = q.new_empty(B, NT, HV, V, K)
else:
dh = q.new_empty(B, NT, HV, K, V)
dh0 = torch.empty_like(h0, dtype=torch.float32) if h0 is not None else None
dv2 = torch.empty_like(dv)
def grid(meta): return (triton.cdiv(V, meta['BV']) * N * HV, )
chunk_gated_delta_rule_bwd_kernel_dhu_blockdim64[grid](
q=q,
k=k,
w=w,
g=g,
gk=gk,
dht=dht,
dh0=dh0,
do=do,
dh=dh,
dv=dv,
dv2=dv2,
cu_seqlens=cu_seqlens,
chunk_offsets=chunk_offsets,
scale=scale,
T=T,
H=H,
HV=HV,
K=K,
V=V,
BT=BT,
STATE_V_FIRST=state_v_first,
)
return dh, dh0, dv2
+432
View File
@@ -0,0 +1,432 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import torch
import triton
import triton.language as tl
from kda._fla.ops.backends import dispatch
from kda._fla.ops.utils import prepare_chunk_offsets
from kda._fla.ops.utils.op import exp2
from kda._fla.utils import autotune_cache_kwargs, check_shared_mem
BKV_LIST = [32, 64] if check_shared_mem() else [16, 32]
@triton.heuristics({
'USE_INITIAL_STATE': lambda args: args['h0'] is not None,
'STORE_FINAL_STATE': lambda args: args['ht'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@triton.autotune(
configs=[
triton.Config({'BK': BK, 'BV': BV}, num_warps=num_warps, num_stages=num_stages)
for BK in BKV_LIST
for BV in BKV_LIST
for num_warps in [1, 2, 4, 8]
for num_stages in [2, 3, 4]
],
key=['BT', 'USE_G', 'USE_GK', 'USE_GV', 'STATE_V_FIRST'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_fwd_kernel_h(
k,
v,
h,
g,
g_gamma,
gk,
gv,
h0,
ht,
cu_seqlens,
split_offsets,
T,
H: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BT: tl.constexpr,
BS: tl.constexpr,
BK: tl.constexpr,
BV: tl.constexpr,
USE_G: tl.constexpr,
USE_G_GAMMA: tl.constexpr,
USE_GK: tl.constexpr,
USE_GV: tl.constexpr,
USE_INITIAL_STATE: tl.constexpr,
STORE_FINAL_STATE: tl.constexpr,
IS_VARLEN: tl.constexpr,
STATE_V_FIRST: tl.constexpr,
):
i_k, i_v, i_nh = tl.program_id(0), tl.program_id(1), tl.program_id(2).to(tl.int64)
i_n, i_h = i_nh // H, i_nh % H
if IS_VARLEN:
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
NT, NS = tl.cdiv(T, BT), tl.cdiv(T, BS)
boh = tl.load(split_offsets + i_n).to(tl.int64)
else:
bos, eos = i_n * T, i_n * T + T
NT, NS = tl.cdiv(T, BT), tl.cdiv(T, BS)
boh = i_n * NS
NTS = BS // BT
if USE_G_GAMMA:
# decay rate given the head index
b_gamma = tl.load(g_gamma + i_h)
b_g = b_gamma * (tl.arange(0, BT) + 1)
# [BK, BV] accumulator; STATE_V_FIRST only flips the stored state's HBM layout to [V, K], applied at the load/store below.
b_h = tl.zeros([BK, BV], dtype=tl.float32)
o_k = i_k * BK + tl.arange(0, BK)
o_v = i_v * BV + tl.arange(0, BV)
if USE_INITIAL_STATE:
if STATE_V_FIRST:
p_h0 = h0 + i_nh * K*V + o_v[:, None] * K + o_k[None, :]
b_h = tl.trans(tl.load(p_h0, mask=(o_v[:, None] < V) & (o_k[None, :] < K), other=0.0)).to(tl.float32)
else:
p_h0 = h0 + i_nh * K*V + o_k[:, None] * V + o_v[None, :]
b_h = tl.load(p_h0, mask=(o_k[:, None] < K) & (o_v[None, :] < V), other=0.0).to(tl.float32)
for i_t in range(NT):
i_s = i_t // NTS
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
p_k = k + (bos*H + i_h) * K + o_k[:, None] + o_t[None, :] * (H*K)
p_v = v + (bos*H + i_h) * V + o_t[:, None] * (H*V) + o_v[None, :]
o_h = ((boh + i_s) * H + i_h).to(tl.int64) * K*V
if STATE_V_FIRST:
p_h = h + o_h + o_v[:, None] * K + o_k[None, :]
m_h = (o_v[:, None] < V) & (o_k[None, :] < K)
else:
p_h = h + o_h + o_k[:, None] * V + o_v[None, :]
m_h = (o_k[:, None] < K) & (o_v[None, :] < V)
if i_t % NTS == 0:
tl.store(p_h, (tl.trans(b_h) if STATE_V_FIRST else b_h).to(p_h.dtype.element_ty), mask=m_h)
# [BK, BT]
b_k = tl.load(p_k, mask=(o_k[:, None] < K) & m_t[None, :], other=0.0)
# [BT, BV]
b_v = tl.load(p_v, mask=m_t[:, None] & (o_v < V)[None, :], other=0.0)
last_idx = min((i_t + 1) * BT, T) - 1
# scalar decay
if USE_G:
b_g_last = tl.load(g + bos * H + last_idx * H + i_h)
p_g = g + bos*H + (i_t * BT + tl.arange(0, BT)) * H + i_h
b_g = tl.load(p_g, mask=(i_t * BT + tl.arange(0, BT) < T), other=0.)
b_h *= exp2(b_g_last)
b_v = (b_v * exp2(b_g_last - b_g)[:, None]).to(b_v.dtype)
if USE_G_GAMMA:
b_g_last = b_gamma * min(BT, T - i_t * BT)
b_h *= exp2(b_g_last)
b_v = (b_v * exp2(b_g_last - b_g)[:, None]).to(b_v.dtype)
# vector decay, h = Diag(gk) @ h
if USE_GK:
p_gk = gk + (bos*H + i_h) * K + o_k[:, None] + o_t[None, :] * (H*K)
p_gk_last = gk + (bos + last_idx) * H*K + i_h * K + i_k * BK + tl.arange(0, BK)
b_gk_last = tl.load(p_gk_last, mask=(i_k * BK + tl.arange(0, BK) < K), other=0.)
b_gk = tl.load(p_gk, mask=(o_k[:, None] < K) & m_t[None, :], other=0.0)
b_h *= exp2(b_gk_last)[:, None]
b_k = (b_k * exp2(b_gk_last[:, None] - b_gk)).to(b_k.dtype)
# vector decay, h = h @ Diag(gv)
if USE_GV:
p_gv = gv + (bos*H + i_h) * V + o_t[:, None] * (H*V) + o_v[None, :]
p_gv_last = gv + (bos + last_idx) * H*V + i_h * V + i_v * BV + tl.arange(0, BV)
b_gv_last = tl.load(p_gv_last, mask=(i_v * BV + tl.arange(0, BV) < V), other=0.)
b_gv = tl.load(p_gv, mask=m_t[:, None] & (o_v < V)[None, :], other=0.0)
b_h *= exp2(b_gv_last)[None, :]
b_v = (b_v * exp2(b_gv_last[None, :] - b_gv)).to(b_v.dtype)
b_h += tl.dot(b_k, b_v)
if STORE_FINAL_STATE:
if STATE_V_FIRST:
p_ht = ht + i_nh * K*V + o_v[:, None] * K + o_k[None, :]
tl.store(p_ht, tl.trans(b_h).to(p_ht.dtype.element_ty), mask=(o_v[:, None] < V) & (o_k[None, :] < K))
else:
p_ht = ht + i_nh * K*V + o_k[:, None] * V + o_v[None, :]
tl.store(p_ht, b_h.to(p_ht.dtype.element_ty), mask=(o_k[:, None] < K) & (o_v[None, :] < V))
@triton.heuristics({
'STORE_INITIAL_STATE_GRADIENT': lambda args: args['dh0'] is not None,
'USE_FINAL_STATE_GRADIENT': lambda args: args['dht'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@triton.autotune(
configs=[
triton.Config({'BK': BK, 'BV': BV}, num_warps=num_warps, num_stages=num_stages)
for BK in BKV_LIST
for BV in BKV_LIST
for num_warps in [1, 2, 4, 8]
for num_stages in [2, 3, 4]
],
key=['BT', 'USE_G', 'USE_GK', 'USE_GV', 'STATE_V_FIRST'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_bwd_kernel_dh(
q,
g,
g_gamma,
gk,
gv,
do,
dh,
dht,
dh0,
cu_seqlens,
split_offsets,
scale,
T,
HQ: tl.constexpr,
H: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BT: tl.constexpr,
BS: tl.constexpr,
BK: tl.constexpr,
BV: tl.constexpr,
NG: tl.constexpr,
USE_G: tl.constexpr,
USE_G_GAMMA: tl.constexpr,
USE_GK: tl.constexpr,
USE_GV: tl.constexpr,
STORE_INITIAL_STATE_GRADIENT: tl.constexpr,
USE_FINAL_STATE_GRADIENT: tl.constexpr,
IS_VARLEN: tl.constexpr,
STATE_V_FIRST: tl.constexpr,
):
i_k, i_v, i_nh = tl.program_id(0), tl.program_id(1), tl.program_id(2).to(tl.int64)
i_n, i_hq = i_nh // HQ, i_nh % HQ
i_h = i_hq // NG
if IS_VARLEN:
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
NT = tl.cdiv(T, BT)
NS = tl.cdiv(T, BS)
boh = tl.load(split_offsets + i_n).to(tl.int64)
else:
bos, eos = i_n * T, i_n * T + T
NT = tl.cdiv(T, BT)
NS = tl.cdiv(T, BS)
boh = i_n * NS
if USE_G_GAMMA:
b_gamma = tl.load(g_gamma + i_h)
b_g = b_gamma * (tl.arange(0, BT) + 1)
# [BK, BV] accumulator; STATE_V_FIRST only flips the stored state's HBM layout to [V, K], applied at the load/store below.
b_dh = tl.zeros([BK, BV], dtype=tl.float32)
o_k = i_k * BK + tl.arange(0, BK)
o_v = i_v * BV + tl.arange(0, BV)
if USE_FINAL_STATE_GRADIENT:
if STATE_V_FIRST:
p_dht = dht + i_nh * K*V + o_v[:, None] * K + o_k[None, :]
b_dh += tl.trans(tl.load(p_dht, mask=(o_v[:, None] < V) & (o_k[None, :] < K), other=0.0)).to(tl.float32)
else:
p_dht = dht + i_nh * K*V + o_k[:, None] * V + o_v[None, :]
b_dh += tl.load(p_dht, mask=(o_k[:, None] < K) & (o_v[None, :] < V), other=0.0).to(tl.float32)
for i_t in range(NT - 1, -1, -1):
i_s = i_t // (BS // BT)
o_dh = ((boh + i_s) * H + i_h).to(tl.int64) * K*V
if STATE_V_FIRST:
p_dh = dh + o_dh + o_v[:, None] * K + o_k[None, :]
m_dh = (o_v[:, None] < V) & (o_k[None, :] < K)
else:
p_dh = dh + o_dh + o_k[:, None] * V + o_v[None, :]
m_dh = (o_k[:, None] < K) & (o_v[None, :] < V)
if i_t % (BS // BT) == 0:
tl.store(p_dh, (tl.trans(b_dh) if STATE_V_FIRST else b_dh).to(p_dh.dtype.element_ty), mask=m_dh)
last_idx = min(i_t * BT + BT, T) - 1
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
# [BK, BT]
p_q = q + (bos*HQ + i_hq) * K + o_k[:, None] + o_t[None, :] * (HQ*K)
p_do = do + (bos*HQ + i_hq) * V + o_t[:, None] * (HQ*V) + o_v[None, :]
b_q = tl.load(p_q, mask=(o_k[:, None] < K) & m_t[None, :], other=0.0)
b_q = (b_q * scale).to(b_q.dtype)
# [BT, BV]
b_do = tl.load(p_do, mask=m_t[:, None] & (o_v < V)[None, :], other=0.0)
if USE_G:
p_g = g + (bos + i_t * BT + tl.arange(0, BT)) * H + i_h
b_g_last = tl.load(g + (bos + last_idx) * H + i_h)
b_g = tl.load(p_g, mask=(i_t * BT + tl.arange(0, BT) < T), other=0.)
b_q = (b_q * exp2(b_g)[None, :]).to(b_q.dtype)
b_dh *= exp2(b_g_last)
if USE_G_GAMMA:
b_g_last = b_gamma * min(BT, T - i_t * BT)
b_q = (b_q * exp2(b_g)[None, :]).to(b_q.dtype)
b_dh *= exp2(b_g_last)
if USE_GK:
p_gk = gk + (bos*H + i_h) * K + o_k[:, None] + o_t[None, :] * (H*K)
p_gk_last = gk + (bos + last_idx) * H*K + i_h * K + i_k * BK + tl.arange(0, BK)
b_gk = tl.load(p_gk, mask=(o_k[:, None] < K) & m_t[None, :], other=0.0)
b_gk_last = tl.load(p_gk_last, mask=(i_k * BK + tl.arange(0, BK) < K), other=0.)
b_q = (b_q * exp2(b_gk)).to(b_q.dtype)
b_dh *= exp2(b_gk_last)[:, None]
if USE_GV:
p_gv = gv + (bos*H + i_h) * V + o_t[:, None] * (H*V) + o_v[None, :]
p_gv_last = gv + (bos + last_idx) * H*V + i_h * V + i_v * BV + tl.arange(0, BV)
b_gv = tl.load(p_gv, mask=m_t[:, None] & (o_v < V)[None, :], other=0.0)
b_gv_last = tl.load(p_gv_last, mask=(i_v * BV + tl.arange(0, BV) < V), other=0.)
b_do = (b_do * exp2(b_gv))
b_dh *= exp2(b_gv_last)[None, :]
b_dh += tl.dot(b_q, b_do.to(b_q.dtype))
if STORE_INITIAL_STATE_GRADIENT:
if STATE_V_FIRST:
p_dh0 = dh0 + i_nh * K*V + o_v[:, None] * K + o_k[None, :]
tl.store(p_dh0, tl.trans(b_dh).to(p_dh0.dtype.element_ty), mask=(o_v[:, None] < V) & (o_k[None, :] < K))
else:
p_dh0 = dh0 + i_nh * K*V + o_k[:, None] * V + o_v[None, :]
tl.store(p_dh0, b_dh.to(p_dh0.dtype.element_ty), mask=(o_k[:, None] < K) & (o_v[None, :] < V))
@dispatch('common')
def chunk_fwd_h(
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor | None = None,
g_gamma: torch.Tensor | None = None,
gk: torch.Tensor | None = None,
gv: torch.Tensor | None = None,
h0: torch.Tensor | None = None,
output_final_state: bool = False,
state_v_first: bool = False,
cu_seqlens: torch.Tensor | None = None,
chunk_size: int = 64,
split_size: int | None = None,
states_in_fp32: bool = False,
) -> tuple[torch.Tensor, torch.Tensor]:
B, T, H, K, V = *k.shape, v.shape[-1]
BT = chunk_size
BS = BT if split_size is None else split_size
assert BS % BT == 0, f"The `split_size` (got {BS}) must be a multiple of `chunk_size` {BT}"
# N: the actual number of sequences in the batch with either equal or variable lengths
if cu_seqlens is None:
N, NS, split_offsets = B, triton.cdiv(T, BS), None
else:
split_offsets = prepare_chunk_offsets(cu_seqlens, BS)
N, NS = len(cu_seqlens) - 1, split_offsets[-1].item()
# `state_v_first` stores the states in V-first `[V, K]` layout instead of `[K, V]`
state_shape = (V, K) if state_v_first else (K, V)
h = k.new_empty(B, NS, H, *state_shape, dtype=k.dtype if not states_in_fp32 else torch.float)
ht = k.new_empty(N, H, *state_shape, dtype=torch.float) if output_final_state else None
def grid(meta): return (triton.cdiv(K, meta['BK']), triton.cdiv(V, meta['BV']), N * H)
chunk_fwd_kernel_h[grid](
k=k,
v=v,
h=h,
g=g,
g_gamma=g_gamma,
gk=gk,
gv=gv,
h0=h0,
ht=ht,
cu_seqlens=cu_seqlens,
split_offsets=split_offsets,
T=T,
H=H,
K=K,
V=V,
BT=BT,
BS=BS,
USE_G=g is not None,
USE_G_GAMMA=g_gamma is not None,
USE_GK=gk is not None,
USE_GV=gv is not None,
STATE_V_FIRST=state_v_first,
)
return h, ht
@dispatch('common')
def chunk_bwd_dh(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
do: torch.Tensor,
h0: torch.Tensor,
dht: torch.Tensor,
scale: float,
g: torch.Tensor | None = None,
g_gamma: torch.Tensor | None = None,
gk: torch.Tensor | None = None,
gv: torch.Tensor | None = None,
state_v_first: bool = False,
cu_seqlens: torch.Tensor | None = None,
chunk_size: int = 64,
split_size: int | None = None,
states_in_fp32: bool = False,
) -> tuple[torch.Tensor, torch.Tensor]:
B, T, H, K, V = *k.shape, v.shape[-1]
HQ = q.shape[2]
BT = chunk_size
BS = BT if split_size is None else split_size
assert BS % BT == 0, f"The `split_size` (got {BS}) must be a multiple of `chunk_size` {BT}"
# N: the actual number of sequences in the batch with either equal or variable lengths
# NG: number of groups in GQA
if cu_seqlens is None:
N, NS, split_offsets = B, triton.cdiv(T, BS), None
else:
split_offsets = prepare_chunk_offsets(cu_seqlens, BS)
N, NS = len(cu_seqlens) - 1, split_offsets[-1].item()
NG = HQ // H
# `state_v_first` stores the states in V-first `[V, K]` layout instead of `[K, V]`
state_shape = (V, K) if state_v_first else (K, V)
dh = k.new_empty(B, NS, HQ, *state_shape, dtype=k.dtype if not states_in_fp32 else torch.float)
dh0 = torch.empty_like(h0, dtype=torch.float) if h0 is not None else None
def grid(meta): return (triton.cdiv(K, meta['BK']), triton.cdiv(V, meta['BV']), N * H)
chunk_bwd_kernel_dh[grid](
q=q,
g=g,
g_gamma=g_gamma,
gk=gk,
gv=gv,
do=do,
dh=dh,
dht=dht,
dh0=dh0,
cu_seqlens=cu_seqlens,
split_offsets=split_offsets,
scale=scale,
T=T,
HQ=HQ,
H=H,
K=K,
V=V,
BT=BT,
BS=BS,
NG=NG,
USE_G=g is not None,
USE_G_GAMMA=g_gamma is not None,
USE_GK=gk is not None,
USE_GV=gv is not None,
STATE_V_FIRST=state_v_first,
)
return dh, dh0
+110
View File
@@ -0,0 +1,110 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
# Shared gate helpers reused across delta-rule family ops (KDA, GDN, ...).
import torch
import triton
import triton.language as tl
from kda._fla.ops.backends import dispatch
from kda._fla.utils import autocast_custom_bwd, autocast_custom_fwd, input_guard
@triton.jit
def fused_beta_sigmoid_fwd_kernel(
x,
y,
scale,
n_elements,
BLOCK_SIZE: tl.constexpr,
):
pid = tl.program_id(0).to(tl.int64)
offs = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE).to(tl.int64)
mask = offs < n_elements
b_x = tl.load(x + offs, mask=mask, other=0).to(tl.float32)
b_y = scale * tl.sigmoid(b_x)
tl.store(y + offs, b_y.to(y.dtype.element_ty), mask=mask)
@triton.jit
def fused_beta_sigmoid_bwd_kernel(
x,
dy,
dx,
scale,
n_elements,
BLOCK_SIZE: tl.constexpr,
):
pid = tl.program_id(0).to(tl.int64)
offs = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE).to(tl.int64)
mask = offs < n_elements
b_x = tl.load(x + offs, mask=mask, other=0).to(tl.float32)
b_dy = tl.load(dy + offs, mask=mask, other=0).to(tl.float32)
b_y = tl.sigmoid(b_x)
b_dx = b_dy * scale * b_y * (1.0 - b_y)
tl.store(dx + offs, b_dx.to(dx.dtype.element_ty), mask=mask)
_BETA_SIGMOID_BLOCK_SIZE = 2048
_BETA_SIGMOID_NUM_WARPS = 8
@dispatch('common')
def fused_beta_sigmoid_fwd(x: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
y = torch.empty_like(x, dtype=torch.float32)
n_elements = x.numel()
grid = (triton.cdiv(n_elements, _BETA_SIGMOID_BLOCK_SIZE),)
fused_beta_sigmoid_fwd_kernel[grid](
x,
y,
scale,
n_elements,
BLOCK_SIZE=_BETA_SIGMOID_BLOCK_SIZE,
num_warps=_BETA_SIGMOID_NUM_WARPS,
)
return y
@dispatch('common')
def fused_beta_sigmoid_bwd(x: torch.Tensor, dy: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
dx = torch.empty_like(x)
n_elements = x.numel()
grid = (triton.cdiv(n_elements, _BETA_SIGMOID_BLOCK_SIZE),)
fused_beta_sigmoid_bwd_kernel[grid](
x,
dy,
dx,
scale,
n_elements,
BLOCK_SIZE=_BETA_SIGMOID_BLOCK_SIZE,
num_warps=_BETA_SIGMOID_NUM_WARPS,
)
return dx
class BetaSigmoidFunction(torch.autograd.Function):
@staticmethod
@input_guard
@autocast_custom_fwd
def forward(ctx, x: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
y = fused_beta_sigmoid_fwd(x, scale)
ctx.save_for_backward(x)
ctx.scale = scale
return y
@staticmethod
@input_guard
@autocast_custom_bwd
def backward(ctx, dy: torch.Tensor):
(x,) = ctx.saved_tensors
dx = fused_beta_sigmoid_bwd(x, dy, ctx.scale)
return dx.type_as(x), None
def fused_beta_sigmoid(x: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
return BetaSigmoidFunction.apply(x, scale)
+13
View File
@@ -0,0 +1,13 @@
"""Context-parallel stubs. Pass ``cp_context=None`` (the default)."""
class FLACPContext:
cu_seqlens = None
cu_seqlens_cpu = None
def build_cp_context(*args, **kwargs):
raise RuntimeError("Context parallel is not included in the vendored KDA kernels")
__all__ = ["FLACPContext", "build_cp_context"]
+11
View File
@@ -0,0 +1,11 @@
"""Context-parallel hooks referenced by KDA fwd/bwd. Not implemented here."""
def _cp_unsupported(*args, **kwargs):
raise RuntimeError("Context parallel is not included in the vendored KDA kernels")
chunk_gated_delta_rule_fwd_h_pre_process = _cp_unsupported
compress_h0 = _cp_unsupported
chunk_gated_delta_rule_bwd_dhu_pre_process = _cp_unsupported
expand_h0 = _cp_unsupported
+1
View File
@@ -0,0 +1 @@
# Vendored GLA chunk output kernel used by KDA.
File diff suppressed because it is too large Load Diff
+7
View File
@@ -0,0 +1,7 @@
from .chunk import chunk_kda
from .fused_recurrent import fused_recurrent_kda
__all__ = [
"chunk_kda",
"fused_recurrent_kda",
]
+443
View File
@@ -0,0 +1,443 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
# Related files are modified and supported by the Moonshot AI Team
import warnings
import torch
from kda._fla.modules.l2norm import l2norm_bwd, l2norm_fwd
from kda._fla.ops.backends import dispatch
from kda._fla.ops.common.gate import fused_beta_sigmoid, fused_beta_sigmoid_bwd
from kda._fla.ops.cp import FLACPContext
from kda._fla.ops.kda.chunk_bwd import chunk_kda_bwd
from kda._fla.ops.kda.chunk_fwd import chunk_kda_fwd
from kda._fla.ops.utils.index import prepare_chunk_indices
from kda._fla.utils import autocast_custom_bwd, autocast_custom_fwd, input_guard
class ChunkKDAFunction(torch.autograd.Function):
@staticmethod
@input_guard
@autocast_custom_fwd
def forward(
ctx,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
A_log: torch.Tensor,
dt_bias: torch.Tensor,
scale: float,
initial_state: torch.Tensor,
output_final_state: bool = False,
use_qk_l2norm_in_kernel: bool = False,
use_gate_in_kernel: bool = False,
use_beta_sigmoid_in_kernel: bool = False,
allow_neg_eigval: bool = False,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
cu_seqlens_cpu: torch.LongTensor | None = None,
safe_gate: bool = False,
lower_bound: float | None = None,
chunk_size: int = 64,
disable_recompute: bool = False,
return_intermediate_states: bool = False,
cp_context: FLACPContext | None = None,
):
# Apply l2norm
q_rstd, k_rstd = None, None
if use_qk_l2norm_in_kernel:
q, q_rstd = l2norm_fwd(q)
k, k_rstd = l2norm_fwd(k)
beta_raw = beta
if use_beta_sigmoid_in_kernel:
beta = fused_beta_sigmoid(beta_raw, scale=2.0 if allow_neg_eigval else 1.0)
chunk_indices = None
if cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(
cu_seqlens,
chunk_size,
cu_seqlens_cpu=cu_seqlens_cpu,
)
g_input = g
(o, final_state, g_cumsum, Aqk, Akk, w, u, qg, kg, v_new, h, initial_state) = chunk_kda_fwd(
q=q,
k=k,
v=v,
g=g_input,
beta=beta,
scale=scale,
initial_state=initial_state,
output_final_state=output_final_state,
cu_seqlens=cu_seqlens,
cu_seqlens_cpu=cu_seqlens_cpu,
chunk_indices=chunk_indices,
safe_gate=safe_gate,
lower_bound=lower_bound,
use_gate_in_kernel=use_gate_in_kernel,
A_log=A_log,
dt_bias=dt_bias,
chunk_size=chunk_size,
disable_recompute=disable_recompute,
return_intermediate_states=return_intermediate_states,
cp_context=cp_context,
state_v_first=state_v_first,
)
if return_intermediate_states:
assert torch.is_inference_mode_enabled(), "return_intermediate_states is only allowed in inference mode"
assert disable_recompute is False, "return_intermediate_states must be used with disable_recompute=False"
return o.type_as(q), final_state, h
ctx.save_for_backward(
q, q_rstd, k, k_rstd, v, g_cumsum, g_input, beta_raw, beta, A_log, dt_bias, Aqk, Akk,
w, u, qg, kg, v_new, h,
initial_state, cu_seqlens, chunk_indices
)
ctx.chunk_size = chunk_size
ctx.safe_gate = safe_gate
ctx.scale = scale
ctx.lower_bound = lower_bound
ctx.use_qk_l2norm_in_kernel = use_qk_l2norm_in_kernel
ctx.use_gate_in_kernel = use_gate_in_kernel
ctx.use_beta_sigmoid_in_kernel = use_beta_sigmoid_in_kernel
ctx.allow_neg_eigval = allow_neg_eigval
ctx.disable_recompute = disable_recompute
ctx.cp_context = cp_context
ctx.state_v_first = state_v_first
return o.type_as(q), final_state
@staticmethod
@input_guard
@autocast_custom_bwd
def backward(
ctx,
do: torch.Tensor,
dht: torch.Tensor,
):
(q, q_rstd, k, k_rstd, v, g_cumsum, g_input, beta_raw, beta, A_log, dt_bias, Aqk, Akk,
w, u, qg, kg, v_new, h,
initial_state, cu_seqlens, chunk_indices) = (
ctx.saved_tensors
)
dq, dk, dv, db, dg, dh0, dA, dbias = chunk_kda_bwd(
q=q,
k=k,
v=v,
beta=beta,
Aqk=Aqk,
Akk=Akk,
scale=ctx.scale,
initial_state=initial_state,
do=do,
dht=dht,
g=g_cumsum,
g_org=g_input if ctx.use_gate_in_kernel else None,
state_v_first=ctx.state_v_first,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
chunk_size=ctx.chunk_size,
safe_gate=ctx.safe_gate,
lower_bound=ctx.lower_bound,
use_gate_in_kernel=ctx.use_gate_in_kernel,
A_log=A_log,
dt_bias=dt_bias,
disable_recompute=ctx.disable_recompute,
cp_context=ctx.cp_context,
w=w,
u=u,
qg=qg,
kg=kg,
v_new=v_new,
h=h,
)
if ctx.use_qk_l2norm_in_kernel:
dq = l2norm_bwd(q, q_rstd, dq)
dk = l2norm_bwd(k, k_rstd, dk)
if ctx.use_beta_sigmoid_in_kernel:
db = fused_beta_sigmoid_bwd(beta_raw, db, scale=2.0 if ctx.allow_neg_eigval else 1.0)
return (dq.to(q), dk.to(k), dv.to(v), dg.to(g_input), db.to(beta_raw), dA, dbias, None, dh0,
None, None, None, None, None, None, None, None, None, None, None, None, None, None)
@dispatch('kda')
@torch.compiler.disable
def chunk_kda(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
scale: float | None = None,
initial_state: torch.Tensor | None = None,
output_final_state: bool = False,
use_qk_l2norm_in_kernel: bool = False,
use_gate_in_kernel: bool = False,
use_beta_sigmoid_in_kernel: bool = False,
allow_neg_eigval: bool = False,
safe_gate: bool = False,
lower_bound: float | None = None,
disable_recompute: bool = False,
return_intermediate_states: bool = False,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
cu_seqlens_cpu: torch.LongTensor | None = None,
cp_context: FLACPContext = None,
**kwargs,
):
r"""
Args:
q (torch.Tensor):
queries of shape ``[B, T, H, K]``.
k (torch.Tensor):
keys of shape ``[B, T, H, K]``.
v (torch.Tensor):
values of shape ``[B, T, HV, V]``.
GVA (Grouped Value Attention) is applied if ``HV > H``, where ``HV`` must be divisible by ``H``.
g (torch.Tensor):
(forget) gating tensor (in log space!) of shape ``[B, T, HV, K]``.
When ``use_gate_in_kernel=False`` (default), ``g`` should be the pre-computed decay value.
When ``use_gate_in_kernel=True``, ``g`` is the raw input before gate activation;
the kernel fuses ``-exp(A_log) * softplus(g + dt_bias)`` + chunk cumsum internally.
beta (torch.Tensor):
betas of shape ``[B, T, HV]``.
scale (Optional[float]):
Scale factor for the KDA attention scores.
If not provided, it will default to ``1 / sqrt(K)``. Default: ``None``.
initial_state (Optional[torch.Tensor]):
Initial state of shape ``[N, HV, K, V]`` for ``N`` input sequences.
For equal-length input sequences, ``N`` equals the batch size ``B``.
Default: ``None``.
output_final_state (Optional[bool]):
Whether to output the final state of shape ``[N, HV, K, V]``. Default: ``False``.
use_qk_l2norm_in_kernel (bool):
Whether to apply L2norm to the q,k tensor internally. Default: ``False``.
use_gate_in_kernel (bool):
Whether to compute the log-space KDA decay internally.
- If ``True``:
The passed ``g`` acts as the raw input for ``-exp(A_log) * softplus(g + dt_bias.view(HV, K))``.
Note that as part of the input arguments,
``A_log`` (shape ``[HV]``) and the optional ``dt_bias`` (shape ``[HV * K]``) should be provided.
When ``lower_bound`` is set, ``A_log`` may be ``None``,
in which case the gate is ``lower_bound * sigmoid(g + dt_bias)``.
- If ``False``, ``g`` is expected to be the pre-computed decay value.
Default: ``False``.
use_beta_sigmoid_in_kernel (bool):
Whether to apply ``torch.sigmoid(beta)`` before launching the chunk kernel.
- If ``True``, the passed ``beta`` acts as the raw beta logits.
- If ``False``, ``beta`` is expected to already be in post-sigmoid space.
Default: ``False``.
allow_neg_eigval (bool):
Whether to allow negative eigenvalues by scaling ``beta`` to ``[0, 2)``.
Only takes effect together with ``use_beta_sigmoid_in_kernel=True``, in which case
the kernel computes ``2 * sigmoid(beta)`` instead of ``sigmoid(beta)``.
Default: ``False``.
safe_gate (bool):
Whether to clamp the gate to ``[lower_bound, 0)`` and enable M=16 TensorCore
acceleration for higher throughput. Requires ``lower_bound`` to be set.
Default: ``False``.
lower_bound (Optional[float]):
Lower bound for the forget gate (in log space). When set together with
``safe_gate=True``, changes the gate activation from
``-exp(A_log) * softplus(g + dt_bias)`` to
``lower_bound * sigmoid(exp(A_log) * (g + dt_bias))``,
which naturally clamps the output to ``[lower_bound, 0)``.
Recommended value: ``-5`` (i.e., per-step decay ``exp(-5) ≈ 0.0067``).
Default: ``None``.
disable_recompute (bool):
Whether to disable gradient recomputation in the kernel. When ``True``, the kernel
will save all intermediate activations for backward pass, which is beneficial
for training small models at the cost of increased memory usage. Default: ``False``.
return_intermediate_states (bool):
If True, returns intermediate state ``h`` for inference scenarios (e.g., vLLM).
Must be used within ``torch.inference_mode()`` and will return a 3-tuple instead of 2-tuple.
This is not intended for training as it bypasses autograd. Default: ``False``.
state_v_first (Optional[bool]):
Store the recurrent state in V-first ``[V, K]`` layout instead of the default ``[K, V]``. Default: ``False``.
cu_seqlens (torch.LongTensor):
Cumulative sequence lengths of shape ``[N+1]`` used for variable-length training,
consistent with the FlashAttention API.
cu_seqlens_cpu (torch.LongTensor):
Cumulative sequence lengths of shape ``[N+1]`` used for variable-length training,
consistent with the FlashAttention API.
cp_context (Optional[FLACPContext]):
Context parallel context for distributed training across multiple devices.
When provided, ``initial_state`` and ``output_final_state`` are not supported,
and ``cu_seqlens`` will be overridden by the context. Default: ``None``.
Returns:
- Normal mode (return_intermediate_states=False): A tuple (o, final_state)
o (torch.Tensor):
Outputs of shape ``[B, T, HV, V]``.
final_state (torch.Tensor):
Final state of shape ``[N, HV, K, V]`` if ``output_final_state=True`` else ``None``.
- Inference mode (return_intermediate_states=True): A tuple (o, final_state, h)
o (torch.Tensor):
Outputs of shape ``[B, T, HV, V]``.
final_state (torch.Tensor):
Final state of shape ``[N, HV, K, V]`` if ``output_final_state=True`` else ``None``.
h (torch.Tensor):
Intermediate states of shape ``[B, NT, HV, K, V]`` and dtype ``bfloat16``.
- For equal-length sequences: ``NT = ceil(T / chunk_size)``
- For variable-length sequences (cu_seqlens): B is always 1 (flattened),
NT is the total number of chunks across all sequences.
Examples::
>>> import torch
>>> import torch.nn.functional as F
>>> from einops import rearrange
>>> from fla.ops.kda import chunk_kda
# inputs with equal lengths (no GVA, HV == H)
>>> B, T, H, K, V = 4, 2048, 4, 512, 512
>>> q = torch.randn(B, T, H, K, dtype=torch.bfloat16, device='cuda')
>>> k = torch.randn(B, T, H, K, dtype=torch.bfloat16, device='cuda')
>>> v = torch.randn(B, T, H, V, dtype=torch.bfloat16, device='cuda')
>>> beta = torch.rand(B, T, H, dtype=torch.bfloat16, device='cuda')
>>> g = torch.rand(B, T, H, K, dtype=torch.bfloat16, device='cuda')
>>> h0 = torch.randn(B, H, K, V, dtype=torch.bfloat16, device='cuda')
>>> A_log = torch.randn(H, dtype=torch.float32, device='cuda')
>>> dt_bias = torch.randn(H * K, dtype=torch.float32, device='cuda')
>>> o, ht = chunk_kda(
q, k, v, g, beta,
A_log=A_log,
dt_bias=dt_bias,
use_qk_l2norm_in_kernel=True,
use_gate_in_kernel=True,
initial_state=h0,
output_final_state=True
)
# GVA mode (HV > H)
>>> HV = 8 # 2x more value heads than qk heads
>>> v = torch.randn(B, T, HV, V, dtype=torch.bfloat16, device='cuda')
>>> g = torch.rand(B, T, HV, K, dtype=torch.bfloat16, device='cuda')
>>> beta = torch.rand(B, T, HV, dtype=torch.bfloat16, device='cuda')
>>> h0 = torch.randn(B, HV, K, V, dtype=torch.bfloat16, device='cuda')
>>> A_log = torch.randn(HV, dtype=torch.float32, device='cuda')
>>> dt_bias = torch.randn(HV * K, dtype=torch.float32, device='cuda')
>>> o, ht = chunk_kda(
q, k, v, g, beta,
A_log=A_log,
dt_bias=dt_bias,
use_qk_l2norm_in_kernel=True,
use_gate_in_kernel=True,
initial_state=h0,
output_final_state=True
)
# for variable-length inputs, the batch size `B` is expected to be 1 and `cu_seqlens` is required
>>> q, k, v, beta, g = map(lambda x: rearrange(x, 'b t ... -> 1 (b t) ...'), (q, k, v, beta, g))
# for a batch with 4 sequences, `cu_seqlens` with 5 start/end positions are expected
>>> cu_seqlens = q.new_tensor([0, 2048, 4096, 6144, 8192], dtype=torch.long)
>>> o, ht = chunk_kda(
q, k, v, g, beta,
A_log=A_log,
dt_bias=dt_bias,
use_qk_l2norm_in_kernel=True,
use_gate_in_kernel=True,
initial_state=h0,
output_final_state=True,
cu_seqlens=cu_seqlens
)
"""
if 'transpose_state_layout' in kwargs:
if state_v_first:
raise ValueError("Cannot pass both `state_v_first` and the deprecated `transpose_state_layout`.")
warnings.warn(
"`transpose_state_layout` is deprecated and renamed to `state_v_first`.",
DeprecationWarning,
stacklevel=2,
)
state_v_first = kwargs.pop('transpose_state_layout')
if cp_context is not None:
assert initial_state is None, "Initial state is not supported for CP"
assert output_final_state is False, "Output final state is not supported for CP"
assert cp_context.cu_seqlens is not None, "cu_seqlens is required for CP"
# Override cu_seqlens and cu_seqlens_cpu with the ones from the context
cu_seqlens = cp_context.cu_seqlens
if cp_context.cu_seqlens_cpu is not None:
cu_seqlens_cpu = cp_context.cu_seqlens_cpu
if cu_seqlens is not None:
if q.shape[0] != 1:
raise ValueError(
f"The batch size is expected to be 1 rather than {q.shape[0]} when using `cu_seqlens`."
f"Please flatten variable-length inputs before processing.",
)
if initial_state is not None and initial_state.shape[0] != len(cu_seqlens) - 1:
raise ValueError(
f"The number of initial states is expected to be equal to the number of input sequences, "
f"i.e., {len(cu_seqlens) - 1} rather than {initial_state.shape[0]}.",
)
if initial_state is not None:
assert initial_state.dtype == torch.float32, "initial_state must be in float32."
A_log, dt_bias = None, None
if use_gate_in_kernel:
A_log, dt_bias = kwargs.get("A_log"), kwargs.get("dt_bias")
if A_log is None and lower_bound is None:
raise ValueError("`A_log` must be provided when `use_gate_in_kernel=True` and `lower_bound` is not set.")
chunk_size = kwargs.pop("chunk_size", 64)
if chunk_size not in (32, 64):
raise ValueError(f"`chunk_size` must be either 32 or 64 for KDA, got {chunk_size}.")
if safe_gate and use_gate_in_kernel:
if lower_bound is None:
raise ValueError("`lower_bound` must be specified when `safe_gate=True` and `use_gate_in_kernel=True`.")
if not (-5 <= lower_bound < 0):
raise ValueError(f"`lower_bound` must be in the safe range [-5, 0), got {lower_bound}.")
if allow_neg_eigval and not use_beta_sigmoid_in_kernel:
raise ValueError("`allow_neg_eigval=True` requires `use_beta_sigmoid_in_kernel=True`.")
# Validate head dimensions for GVA
B, T, H, K, HV = *q.shape, v.shape[2]
assert q.shape == k.shape, f"q and k must have the same shape, got q={q.shape} vs k={k.shape}"
assert K <= 256, f"Currently we only support key headdim <=256 for KDA, got {K}."
assert HV % H == 0, (
f"For GVA, num_v_heads (HV={HV}) must be evenly divisible by num_qk_heads (H={H}), "
f"but got HV % H = {HV % H}"
)
assert g.shape == (B, T, HV, K), f"g must have shape [B, T, HV, K]={[B, T, HV, K]}, got {list(g.shape)}"
assert beta.shape == (B, T, HV), f"beta must have shape [B, T, HV]={[B, T, HV]}, got {list(beta.shape)}"
if scale is None:
scale = K ** -0.5
return ChunkKDAFunction.apply(
q,
k,
v,
g,
beta,
A_log,
dt_bias,
scale,
initial_state,
output_final_state,
use_qk_l2norm_in_kernel,
use_gate_in_kernel,
use_beta_sigmoid_in_kernel,
allow_neg_eigval,
state_v_first,
cu_seqlens,
cu_seqlens_cpu,
safe_gate,
lower_bound,
chunk_size,
disable_recompute,
return_intermediate_states,
cp_context,
)
+651
View File
@@ -0,0 +1,651 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import torch
import triton
import triton.language as tl
from kda._fla.ops.backends import dispatch
from kda._fla.ops.common.chunk_delta_h import (
chunk_gated_delta_rule_bwd_dhu,
chunk_gated_delta_rule_fwd_h,
)
from kda._fla.ops.cp import FLACPContext
from kda._fla.ops.cp.chunk_delta_h import (
chunk_gated_delta_rule_bwd_dhu_pre_process,
expand_h0,
)
from kda._fla.ops.kda.chunk_intra import chunk_kda_bwd_intra
from kda._fla.ops.kda.gate import kda_gate_bwd, kda_gate_chunk_cumsum
from kda._fla.ops.kda.wy_fast import recompute_w_u_fwd
from kda._fla.ops.utils import chunk_local_cumsum, prepare_chunk_indices
from kda._fla.ops.utils.cache import fla_cache_autotune
from kda._fla.ops.utils.constant import RCP_LN2
from kda._fla.ops.utils.op import exp2
from kda._fla.utils import (
IS_NVIDIA_HOPPER,
IS_NVIDIA_SM100,
autotune_cache_kwargs,
check_shared_mem,
)
BK_LIST = [32, 64] if check_shared_mem() else [16, 32]
BV_LIST = [64, 128] if check_shared_mem("ampere") else [16, 32]
NUM_WARPS = [2, 4] if IS_NVIDIA_HOPPER else [2, 4, 8]
@triton.heuristics(
{
"IS_VARLEN": lambda args: args["cu_seqlens"] is not None,
}
)
@fla_cache_autotune(
configs=[
triton.Config({}, num_warps=num_warps, num_stages=num_stages)
for num_warps in NUM_WARPS
for num_stages in [2, 3, 4]
],
key=["H", "HV", "K", "V", "BT", "BK", "BV"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=["T"])
def chunk_kda_bwd_kernel_dAv(
q,
k,
v,
A,
do,
dv,
dA,
cu_seqlens,
chunk_indices,
scale,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BT: tl.constexpr,
BK: tl.constexpr,
BV: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_t, i_bh = tl.program_id(0).to(tl.int64), tl.program_id(1).to(tl.int64)
i_b, i_hv = i_bh // HV, i_bh % HV
i_h = i_hv // (HV // H)
if IS_VARLEN:
i_n, i_t = (
tl.load(chunk_indices + i_t * 2).to(tl.int32),
tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64),
)
bos, eos = (
tl.load(cu_seqlens + i_n).to(tl.int64),
tl.load(cu_seqlens + i_n + 1).to(tl.int64),
)
T = eos - bos
else:
bos, eos = i_b * T, i_b * T + T
# offset calculation
q += (bos * H + i_h) * K
k += (bos * H + i_h) * K
v += (bos * HV + i_hv) * V
do += (bos * HV + i_hv) * V
dv += (bos * HV + i_hv) * V
dA += (bos * HV + i_hv) * BT
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
o_A = tl.arange(0, BT)
m_AT = (o_A[:, None] < BT) & m_t[None, :]
p_A = A + (bos * HV + i_hv) * BT + o_A[:, None] + o_t[None, :] * (HV * BT)
b_A = tl.load(p_A, mask=m_AT, other=0.0)
m_A = (o_t[:, None] <= o_t[None, :]) & (m_t[:, None] & m_t)
b_A = tl.where(m_A, b_A, 0).to(do.dtype.element_ty)
b_dA = tl.zeros([BT, BT], dtype=tl.float32)
for i_v in range(tl.cdiv(V, BV)):
o_v = i_v * BV + tl.arange(0, BV)
m_v = o_v < V
m_vT = m_v[:, None] & m_t[None, :]
m_tv = m_t[:, None] & m_v[None, :]
p_v = v + o_v[:, None] + o_t[None, :] * (HV * V)
p_do = do + o_t[:, None] * (HV * V) + o_v[None, :]
p_dv = dv + o_t[:, None] * (HV * V) + o_v[None, :]
# [BV, BT]
b_v = tl.load(p_v, mask=m_vT, other=0.0)
# [BT, BV]
b_do = tl.load(p_do, mask=m_tv, other=0.0)
# [BT, BT]
b_dA += tl.dot(b_do, b_v)
# [BT, BV]
b_dv = tl.dot(b_A.to(b_do.dtype), b_do)
tl.store(p_dv, b_dv.to(p_dv.dtype.element_ty), mask=m_tv)
m_dA = m_t[:, None] & (o_A[None, :] < BT)
p_dA = dA + o_t[:, None] * (HV * BT) + o_A[None, :]
b_dA = tl.where(o_t[:, None] >= o_t, b_dA * scale, 0.0)
tl.store(p_dA, b_dA.to(p_dA.dtype.element_ty), mask=m_dA)
@triton.heuristics(
{
"IS_VARLEN": lambda args: args["cu_seqlens"] is not None,
}
)
@fla_cache_autotune(
configs=[
triton.Config({"BK": BK, "BV": BV}, num_warps=num_warps, num_stages=num_stages)
for BK in BK_LIST
for BV in BV_LIST
for num_warps in NUM_WARPS
for num_stages in [2, 3, 4]
if not (IS_NVIDIA_HOPPER and BK == 32 and num_warps == 4)
if not (IS_NVIDIA_SM100 and BK == 32 and num_warps != 2)
],
key=["BT", "HV", "STATE_V_FIRST"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=["T"])
def chunk_kda_bwd_kernel_wy_dqkg_fused(
q,
k,
v,
v_new,
g,
beta,
A,
h,
do,
dh,
dq,
dk,
dv,
dv2,
dg,
db,
dA,
cu_seqlens,
chunk_indices,
scale,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BT: tl.constexpr,
BK: tl.constexpr,
BV: tl.constexpr,
STATE_V_FIRST: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_t, i_bh = tl.program_id(0).to(tl.int64), tl.program_id(1)
i_b, i_hv = i_bh // HV, i_bh % HV
i_h = i_hv // (HV // H)
if IS_VARLEN:
i_tg = i_t.to(tl.int64)
i_n, i_t = (
tl.load(chunk_indices + i_t * 2).to(tl.int32),
tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64),
)
bos, eos = (
tl.load(cu_seqlens + i_n).to(tl.int64),
tl.load(cu_seqlens + i_n + 1).to(tl.int64),
)
T = (eos - bos).to(tl.int32)
NT = tl.cdiv(T, BT)
else:
NT = tl.cdiv(T, BT)
i_tg = (i_b * NT + i_t).to(tl.int64)
bos, eos = (i_b * T).to(tl.int64), (i_b * T + T).to(tl.int64)
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
m_last = o_t == min(T, i_t * BT + BT) - 1
q += (bos * H + i_h) * K
k += (bos * H + i_h) * K
v += (bos * HV + i_hv) * V
v_new += (bos * HV + i_hv) * V
g += (bos * HV + i_hv) * K
beta += bos * HV + i_hv
A += (bos * HV + i_hv) * BT
h += (i_tg * HV + i_hv) * K * V
do += (bos * HV + i_hv) * V
dh += (i_tg * HV + i_hv) * K * V
dq += (bos * HV + i_hv) * K
dk += (bos * HV + i_hv) * K
dv += (bos * HV + i_hv) * V
dv2 += (bos * HV + i_hv) * V
dg += (bos * HV + i_hv) * K
db += bos * HV + i_hv
dA += (bos * HV + i_hv) * BT
p_beta = beta + o_t * HV
b_beta = tl.load(p_beta, mask=m_t, other=0.0)
o_A = tl.arange(0, BT)
m_AT = (o_A[:, None] < BT) & m_t[None, :]
p_A = A + o_A[:, None] + o_t[None, :] * (HV * BT)
b_A = tl.load(p_A, mask=m_AT, other=0.0)
b_dA = tl.zeros([BT, BT], dtype=tl.float32)
b_db = tl.zeros([BT], dtype=tl.float32)
for i_k in range(tl.cdiv(K, BK)):
o_k = i_k * BK + tl.arange(0, BK)
m_k = o_k < K
m_tk = m_t[:, None] & m_k[None, :]
p_k = k + o_t[:, None] * (H * K) + o_k[None, :]
p_g = g + o_t[:, None] * (HV * K) + o_k[None, :]
b_k = tl.load(p_k, mask=m_tk, other=0.0)
b_g = tl.load(p_g, mask=m_tk, other=0.0).to(tl.float32)
p_gn = g + (min(T, i_t * BT + BT) - 1).to(tl.int64) * HV * K + o_k
b_gn = tl.load(p_gn, mask=m_k, other=0).to(tl.float32)
b_dq = tl.zeros([BT, BK], dtype=tl.float32)
b_dk = tl.zeros([BT, BK], dtype=tl.float32)
b_dw = tl.zeros([BT, BK], dtype=tl.float32)
b_dgk = tl.zeros([BK], dtype=tl.float32)
for i_v in range(tl.cdiv(V, BV)):
o_v = i_v * BV + tl.arange(0, BV)
m_tv = m_t[:, None] & (o_v[None, :] < V)
m_h = (o_v[:, None] < V) & m_k[None, :]
p_v_new = v_new + o_t[:, None] * (HV * V) + o_v[None, :]
p_do = do + o_t[:, None] * (HV * V) + o_v[None, :]
if STATE_V_FIRST:
p_h = h + o_v[:, None] * K + o_k[None, :]
p_dh = dh + o_v[:, None] * K + o_k[None, :]
else:
p_h = h + o_v[:, None] + o_k[None, :] * V
p_dh = dh + o_v[:, None] + o_k[None, :] * V
p_dv = dv + o_t[:, None] * (HV * V) + o_v[None, :]
# [BT, BV]
b_v_new = tl.load(p_v_new, mask=m_tv, other=0.0)
b_do = tl.load(p_do, mask=m_tv, other=0.0)
# [BV, BK]
b_h = tl.load(p_h, mask=m_h, other=0.0)
b_dh = tl.load(p_dh, mask=m_h, other=0.0)
# [BT, BV]
b_dv = tl.load(p_dv, mask=m_tv, other=0.0)
b_dgk += tl.sum(b_h * b_dh, axis=0)
b_dq += tl.dot(b_do, b_h.to(b_do.dtype))
b_dk += tl.dot(b_v_new, b_dh.to(b_v_new.dtype))
b_dw += tl.dot(b_dv.to(b_v_new.dtype), b_h.to(b_v_new.dtype))
tl.debug_barrier() # DO NOT REMOVE THIS LINE!
if i_k == 0:
p_v = v + o_t[:, None] * (HV * V) + o_v[None, :]
p_dv2 = dv2 + o_t[:, None] * (HV * V) + o_v[None, :]
b_v = tl.load(p_v, mask=m_tv, other=0.0)
b_dA += tl.dot(b_dv, tl.trans(b_v))
b_dvb = tl.dot(b_A, b_dv)
b_dv2 = b_dvb * b_beta[:, None]
b_db += tl.sum(b_dvb * b_v, 1)
tl.store(p_dv2, b_dv2.to(p_dv2.dtype.element_ty), mask=m_tv)
b_gk_exp = exp2(b_g)
b_gb = b_gk_exp * b_beta[:, None]
b_dgk *= exp2(b_gn)
b_dq = b_dq * b_gk_exp * scale
b_dk = b_dk * tl.where(m_t[:, None], exp2(b_gn[None, :] - b_g), 0)
b_kg = b_k * b_gk_exp
b_dw = -b_dw.to(b_A.dtype)
b_dA += tl.dot(b_dw, tl.trans(b_kg.to(b_A.dtype)))
b_dkgb = tl.dot(b_A, b_dw)
b_db += tl.sum(b_dkgb * b_kg, 1)
p_q = q + o_t[:, None] * (H * K) + o_k[None, :]
b_q = tl.load(p_q, mask=m_tk, other=0.0)
b_kdk = b_k * b_dk
b_dgk += tl.sum(b_kdk, axis=0)
b_dg = (
b_q * b_dq
- b_kdk
+ m_last[:, None] * b_dgk
+ b_kg * b_dkgb * b_beta[:, None]
)
b_dk = b_dk + b_dkgb * b_gb
p_dq = dq + o_t[:, None] * (HV * K) + o_k[None, :]
p_dk = dk + o_t[:, None] * (HV * K) + o_k[None, :]
p_dg = dg + o_t[:, None] * (HV * K) + o_k[None, :]
tl.store(p_dq, b_dq.to(p_dq.dtype.element_ty), mask=m_tk)
tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), mask=m_tk)
tl.store(p_dg, b_dg.to(p_dg.dtype.element_ty), mask=m_tk)
m_A = (o_t[:, None] > o_t[None, :]) & (m_t[:, None] & m_t)
b_dA = tl.where(m_A, b_dA * b_beta[None, :], 0)
b_dA = tl.dot(b_dA.to(b_A.dtype), b_A)
b_dA = tl.dot(b_A, b_dA.to(b_A.dtype))
b_dA = tl.where(m_A, -b_dA, 0)
m_dA = m_t[:, None] & (o_A[None, :] < BT)
p_dA = dA + o_t[:, None] * (HV * BT) + o_A[None, :]
p_db = db + o_t * HV
tl.store(p_dA, b_dA.to(p_dA.dtype.element_ty), mask=m_dA)
tl.store(p_db, b_db.to(p_db.dtype.element_ty), mask=m_t)
@dispatch("kda")
def chunk_kda_bwd_dAv(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
do: torch.Tensor,
A: torch.Tensor | None = None,
scale: float = None,
cu_seqlens: torch.LongTensor | None = None,
chunk_size: int = 64,
chunk_indices: torch.LongTensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
B, T, H, K, HV, V = *k.shape, do.shape[2], do.shape[-1]
BT = chunk_size
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
# H100 can have larger block size
if check_shared_mem("hopper", k.device.index):
CONST_TILING = 128
elif check_shared_mem:
CONST_TILING = 64
else:
CONST_TILING = 32
BK = min(max(triton.next_power_of_2(K), 16), CONST_TILING)
BV = min(max(triton.next_power_of_2(V), 16), CONST_TILING)
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
dA = v.new_empty(B, T, HV, BT, dtype=torch.float)
dv = torch.empty_like(do)
grid = (NT, B * HV)
chunk_kda_bwd_kernel_dAv[grid](
q=q,
k=k,
v=v,
A=A,
do=do,
dv=dv,
dA=dA,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
scale=scale,
T=T,
H=H,
HV=HV,
K=K,
V=V,
BT=BT,
BK=BK,
BV=BV,
)
return dA, dv
@dispatch("kda")
def chunk_kda_bwd_wy_dqkg_fused(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
v_new: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
A: torch.Tensor,
h: torch.Tensor,
do: torch.Tensor,
dh: torch.Tensor,
dv: torch.Tensor,
scale: float | None = None,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
chunk_size: int = 64,
chunk_indices: torch.LongTensor | None = None,
):
B, T, H, K, HV, V = *k.shape, v.shape[2], v.shape[-1]
BT = chunk_size
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
# dq, dk are allocated at HV dimension; caller reduces to H if GVA
dq = g.new_empty(B, T, HV, K, dtype=torch.float)
dk = g.new_empty(B, T, HV, K, dtype=torch.float)
dv2 = torch.empty_like(v)
dg = torch.empty_like(g, dtype=torch.float)
db = torch.empty_like(beta, dtype=torch.float)
dA = torch.empty_like(A, dtype=torch.float)
grid = (NT, B * HV)
chunk_kda_bwd_kernel_wy_dqkg_fused[grid](
q=q,
k=k,
v=v,
v_new=v_new,
g=g,
beta=beta,
A=A,
h=h,
do=do,
dh=dh,
dq=dq,
dk=dk,
dv=dv,
dv2=dv2,
dg=dg,
db=db,
dA=dA,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
scale=scale,
T=T,
H=H,
HV=HV,
K=K,
V=V,
BT=BT,
STATE_V_FIRST=state_v_first,
)
dv = dv2
return dq, dk, dv, db, dg, dA
def chunk_kda_bwd(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
beta: torch.Tensor,
Aqk: torch.Tensor,
Akk: torch.Tensor,
scale: float,
initial_state: torch.Tensor,
do: torch.Tensor,
dht: torch.Tensor,
g: torch.Tensor | None = None,
g_org: torch.Tensor | None = None,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
chunk_indices: torch.LongTensor | None = None,
chunk_size: int = 64,
safe_gate: bool = False,
lower_bound: float | None = None,
use_gate_in_kernel: bool = False,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
disable_recompute: bool = False,
cp_context: FLACPContext | None = None,
**kwargs,
):
H, HV = q.shape[2], v.shape[2]
G = HV // H
if disable_recompute is False:
if use_gate_in_kernel:
g = kda_gate_chunk_cumsum(
g=g_org,
A_log=A_log,
dt_bias=dt_bias,
scale=RCP_LN2,
chunk_size=chunk_size,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
lower_bound=lower_bound,
)
w, u, qg, kg = recompute_w_u_fwd(
q=q,
k=k,
v=v,
beta=beta,
A=Akk,
gk=g,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
)
if cp_context is not None:
# Restore the full initial_state tensor from the compressed version.
# Only the first sequence's state is non-zero as it's the only one that could be cross-rank.
initial_state = expand_h0(initial_state, context=cp_context)
h, v_new, _ = chunk_gated_delta_rule_fwd_h(
k=kg,
w=w,
u=u,
gk=g,
initial_state=initial_state,
output_final_state=False,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
chunk_size=chunk_size,
state_v_first=state_v_first,
)
else:
w, u, qg, kg, v_new, h = (
kwargs["w"],
kwargs["u"],
kwargs["qg"],
kwargs["kg"],
kwargs["v_new"],
kwargs["h"],
)
if cp_context is not None:
# Restore the full initial_state tensor from the compressed version.
# Only the first sequence's state is non-zero as it's the only one that could be cross-rank.
initial_state = expand_h0(initial_state, context=cp_context)
# dAqk = do @ v.T
# dv = A @ do
dAqk, dv = chunk_kda_bwd_dAv(
q=q,
k=k,
v=v_new,
do=do,
A=Aqk,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_size=chunk_size,
chunk_indices=chunk_indices,
)
if cp_context is not None:
# initial_state is None in the CP mode
# We only need to compute dht of current rank and pass it to the backward kernel
dht, initial_state = chunk_gated_delta_rule_bwd_dhu_pre_process(
q=qg,
k=kg,
w=w,
do=do,
dv=dv,
gk=g,
scale=scale,
cu_seqlens=cu_seqlens,
dht=dht,
initial_state=initial_state,
context=cp_context,
chunk_size=chunk_size,
state_v_first=state_v_first,
)
dh, dh0, dv = chunk_gated_delta_rule_bwd_dhu(
q=qg,
k=kg,
w=w,
gk=g,
h0=initial_state,
dht=dht,
do=do,
dv=dv,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_size=chunk_size,
chunk_indices=chunk_indices,
state_v_first=state_v_first,
)
dq, dk, dv, db, dg, dAkk = chunk_kda_bwd_wy_dqkg_fused(
q=q,
k=k,
v=v,
v_new=v_new,
g=g,
beta=beta,
A=Akk,
h=h,
do=do,
dh=dh,
dv=dv,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_size=chunk_size,
chunk_indices=chunk_indices,
state_v_first=state_v_first,
)
dq, dk, db, dg = chunk_kda_bwd_intra(
q=q,
k=k,
g=g,
beta=beta,
dAqk=dAqk,
dAkk=dAkk,
dq=dq,
dk=dk,
db=db,
dg=dg,
cu_seqlens=cu_seqlens,
chunk_size=chunk_size,
chunk_indices=chunk_indices,
safe_gate=safe_gate,
)
# For GVA, reduce dq and dk from [B, T, HV, K] back to [B, T, H, K]
if HV > H:
dq = dq.view(*dq.shape[:2], H, G, dq.shape[-1]).sum(dim=3)
dk = dk.view(*dk.shape[:2], H, G, dk.shape[-1]).sum(dim=3)
dA, dbias = None, None
dg = chunk_local_cumsum(
dg,
chunk_size=chunk_size,
reverse=True,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
)
if use_gate_in_kernel:
dg, dA, dbias = kda_gate_bwd(
g=g_org, A_log=A_log, dt_bias=dt_bias, dyg=dg, lower_bound=lower_bound
)
return dq, dk, dv, db, dg, dh0, dA, dbias
+134
View File
@@ -0,0 +1,134 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import torch
from kda._fla.ops.common.chunk_delta_h import chunk_gated_delta_rule_fwd_h
from kda._fla.ops.cp import FLACPContext
from kda._fla.ops.cp.chunk_delta_h import chunk_gated_delta_rule_fwd_h_pre_process, compress_h0
from kda._fla.ops.gla.chunk import chunk_gla_fwd_o_gk
from kda._fla.ops.kda.chunk_intra import chunk_kda_fwd_intra
from kda._fla.ops.kda.gate import kda_gate_chunk_cumsum
from kda._fla.ops.utils import chunk_local_cumsum
from kda._fla.ops.utils.constant import RCP_LN2
def chunk_kda_fwd(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
scale: float,
initial_state: torch.Tensor,
output_final_state: bool,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
cu_seqlens_cpu: torch.LongTensor | None = None,
chunk_indices: torch.LongTensor | None = None,
chunk_size: int = 64,
safe_gate: bool = False,
lower_bound: float | None = None,
use_gate_in_kernel: bool = False,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
disable_recompute: bool = False,
return_intermediate_states: bool = False,
cp_context: FLACPContext | None = None,
):
# Apply gate activation
g_org = None
if use_gate_in_kernel:
g_org = g
g = kda_gate_chunk_cumsum(
g=g_org,
A_log=A_log,
dt_bias=dt_bias,
scale=RCP_LN2,
chunk_size=chunk_size,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
lower_bound=lower_bound,
)
else:
g = chunk_local_cumsum(
g=g,
scale=RCP_LN2,
chunk_size=chunk_size,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices
)
# qg = None if disable_recompute is False
w, u, qg, kg, Aqk, Akk = chunk_kda_fwd_intra(
q=q,
k=k,
v=v,
gk=g,
beta=beta,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_size=chunk_size,
chunk_indices=chunk_indices,
safe_gate=safe_gate,
disable_recompute=disable_recompute
)
if cp_context is not None:
initial_state = chunk_gated_delta_rule_fwd_h_pre_process(
k=kg,
w=w,
u=u,
gk=g,
cu_seqlens=cu_seqlens,
initial_state=initial_state,
context=cp_context,
chunk_size=chunk_size,
state_v_first=state_v_first,
)
h, v_new, final_state = chunk_gated_delta_rule_fwd_h(
k=kg,
w=w,
u=u,
gk=g,
initial_state=initial_state,
output_final_state=output_final_state,
cu_seqlens=cu_seqlens,
cu_seqlens_cpu=cu_seqlens_cpu,
chunk_indices=chunk_indices,
chunk_size=chunk_size,
state_v_first=state_v_first,
)
if cp_context is not None:
# In Context Parallel (CP) mode, global initial states are not supported at the entry point.
# The `initial_state` here is computed internally via inter-rank communication.
# Since only the first sequence in the local batch can be a continuation of a cross-rank sequence,
# only the first state in the tensor is relevant. We compress it to optimize memory for `save_for_backward`.
initial_state = compress_h0(initial_state, context=cp_context)
o = chunk_gla_fwd_o_gk(
q=q,
v=v_new,
g=g,
A=Aqk,
h=h,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_size=chunk_size,
chunk_indices=chunk_indices,
state_v_first=state_v_first,
)
if disable_recompute is False:
# Delete to save memory
w, u, qg, kg, v_new = None, None, None, None, None
if not return_intermediate_states:
h = None
if use_gate_in_kernel:
g = None
return o, final_state, g, Aqk, Akk, w, u, qg, kg, v_new, h, initial_state
+962
View File
@@ -0,0 +1,962 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import torch
import triton
import triton.language as tl
from kda._fla.ops.backends import dispatch
from kda._fla.ops.kda.chunk_intra_token_parallel import chunk_kda_fwd_intra_token_parallel
from kda._fla.ops.kda.wy_fast import recompute_w_u_fwd
from kda._fla.ops.utils import prepare_chunk_indices
from kda._fla.ops.utils.cache import fla_cache_autotune
from kda._fla.ops.utils.op import exp2, gather
from kda._fla.utils import IS_GATHER_SUPPORTED, IS_TF32_SUPPORTED, autotune_cache_kwargs
if IS_TF32_SUPPORTED:
SOLVE_TRIL_DOT_PRECISION = tl.constexpr('tf32')
else:
SOLVE_TRIL_DOT_PRECISION = tl.constexpr('ieee')
################################################################################
# Fused inter + solve_tril kernel: compute off-diagonal Akk and solve in one pass
################################################################################
@triton.heuristics({
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({'BK': BK}, num_warps=num_warps)
for BK in [32, 64]
for num_warps in [1, 2, 4]
],
key=["H", "HV", "K", "BT", "BC", "NC"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_kda_fwd_kernel_inter_solve_fused(
q,
k,
g,
beta,
Aqk,
Akkd,
Akk,
scale,
cu_seqlens,
chunk_indices,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
BT: tl.constexpr,
BC: tl.constexpr,
NC: tl.constexpr,
BK: tl.constexpr,
IS_VARLEN: tl.constexpr,
USE_SAFE_GATE: tl.constexpr,
):
"""
Fused kernel: compute inter-subchunk Akk + solve_tril in one pass.
Prerequisite: token_parallel has already computed diagonal Akk blocks in Akkd.
This kernel:
1. Computes off-diagonal Aqk blocks -> writes to global
2. Computes off-diagonal Akk blocks -> keeps in registers
3. Loads diagonal Akk blocks from Akkd (fp32)
4. Does forward substitution on diagonals
5. Computes merged Akk_inv
6. Writes Akk_inv to Akk
"""
i_t, i_bh = tl.program_id(0).to(tl.int64), tl.program_id(1).to(tl.int64)
i_b, i_hv = i_bh // HV, i_bh % HV
i_h = i_hv // (HV // H)
if IS_VARLEN:
i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64)
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
else:
bos, eos = i_b * T, i_b * T + T
if i_t * BT >= T:
return
i_tc0 = i_t * BT
i_tc1 = i_t * BT + BC
i_tc2 = i_t * BT + 2 * BC
i_tc3 = i_t * BT + 3 * BC
q += (bos * H + i_h) * K
k += (bos * H + i_h) * K
g += (bos * HV + i_hv) * K
Aqk += (bos * HV + i_hv) * BT
Akk += (bos * HV + i_hv) * BT
Akkd += (bos * HV + i_hv) * BC
o_i = tl.arange(0, BC)
m_tc1 = (i_tc1 + o_i) < T
m_tc2 = (i_tc2 + o_i) < T
m_tc3 = (i_tc3 + o_i) < T
o_c0 = i_tc0 + o_i
o_c1 = i_tc1 + o_i
o_c2 = i_tc2 + o_i
o_c3 = i_tc3 + o_i
m_tc0 = o_c0 < T
m_A0 = m_tc0[:, None] & (o_i[None, :] < BT)
m_A1 = m_tc1[:, None] & (o_i[None, :] < BT)
m_A2 = m_tc2[:, None] & (o_i[None, :] < BT)
m_A3 = m_tc3[:, None] & (o_i[None, :] < BT)
b_Aqk10 = tl.zeros([BC, BC], dtype=tl.float32)
b_Akk10 = tl.zeros([BC, BC], dtype=tl.float32)
b_Aqk20 = tl.zeros([BC, BC], dtype=tl.float32)
b_Akk20 = tl.zeros([BC, BC], dtype=tl.float32)
b_Aqk21 = tl.zeros([BC, BC], dtype=tl.float32)
b_Akk21 = tl.zeros([BC, BC], dtype=tl.float32)
b_Aqk30 = tl.zeros([BC, BC], dtype=tl.float32)
b_Akk30 = tl.zeros([BC, BC], dtype=tl.float32)
b_Aqk31 = tl.zeros([BC, BC], dtype=tl.float32)
b_Akk31 = tl.zeros([BC, BC], dtype=tl.float32)
b_Aqk32 = tl.zeros([BC, BC], dtype=tl.float32)
b_Akk32 = tl.zeros([BC, BC], dtype=tl.float32)
################################################################################
# off-diagonal blocks
################################################################################
for i_k in range(tl.cdiv(K, BK)):
o_k = i_k * BK + tl.arange(0, BK)
m_k = o_k < K
m_ck0 = m_tc0[:, None] & m_k[None, :]
p_k0 = k + o_c0[:, None] * (H*K) + o_k[None, :]
p_g0 = g + o_c0[:, None] * (HV*K) + o_k[None, :]
b_k0 = tl.load(p_k0, mask=m_ck0, other=0.0).to(tl.float32)
b_g0 = tl.load(p_g0, mask=m_ck0, other=0.0).to(tl.float32)
if i_tc1 < T:
m_ck1 = m_tc1[:, None] & m_k[None, :]
p_q1 = q + o_c1[:, None] * (H*K) + o_k[None, :]
p_k1 = k + o_c1[:, None] * (H*K) + o_k[None, :]
p_g1 = g + o_c1[:, None] * (HV*K) + o_k[None, :]
# [BC, BK]
b_q1 = tl.load(p_q1, mask=m_ck1, other=0.0).to(tl.float32)
b_k1 = tl.load(p_k1, mask=m_ck1, other=0.0).to(tl.float32)
b_g1 = tl.load(p_g1, mask=m_ck1, other=0.0).to(tl.float32)
# [BK]
b_gn1 = tl.load(g + i_tc1 * HV*K + o_k, mask=m_k, other=0).to(tl.float32)
# [BC, BK]
b_gqn = tl.where(m_tc1[:, None], exp2(b_g1 - b_gn1[None, :]), 0)
# [BK, BC]
b_kgt = tl.trans(b_k0 * exp2(b_gn1[None, :] - b_g0))
# [BC, BC]
b_Aqk10 += tl.dot(b_q1 * b_gqn, b_kgt)
b_Akk10 += tl.dot(b_k1 * b_gqn, b_kgt)
if NC >= 3 and i_tc2 < T:
m_ck2 = m_tc2[:, None] & m_k[None, :]
p_q2 = q + o_c2[:, None] * (H*K) + o_k[None, :]
p_k2 = k + o_c2[:, None] * (H*K) + o_k[None, :]
p_g2 = g + o_c2[:, None] * (HV*K) + o_k[None, :]
# [BC, BK]
b_q2 = tl.load(p_q2, mask=m_ck2, other=0.0).to(tl.float32)
b_k2 = tl.load(p_k2, mask=m_ck2, other=0.0).to(tl.float32)
b_g2 = tl.load(p_g2, mask=m_ck2, other=0.0).to(tl.float32)
# [BK]
b_gn2 = tl.load(g + i_tc2 * HV*K + o_k, mask=m_k, other=0).to(tl.float32)
# [BC, BK]
b_gqn2 = tl.where(m_tc2[:, None], exp2(b_g2 - b_gn2[None, :]), 0)
b_qg2 = b_q2 * b_gqn2
b_kg2 = b_k2 * b_gqn2
# [BK, BC]
b_kgt = tl.trans(b_k0 * exp2(b_gn2[None, :] - b_g0))
b_Aqk20 += tl.dot(b_qg2, b_kgt)
b_Akk20 += tl.dot(b_kg2, b_kgt)
# [BC, BC]
b_kgt = tl.trans(b_k1 * exp2(b_gn2[None, :] - b_g1))
# [BC, BC]
b_Aqk21 += tl.dot(b_qg2, b_kgt)
b_Akk21 += tl.dot(b_kg2, b_kgt)
if NC >= 4 and i_tc3 < T:
m_ck3 = m_tc3[:, None] & m_k[None, :]
p_q3 = q + o_c3[:, None] * (H*K) + o_k[None, :]
p_k3 = k + o_c3[:, None] * (H*K) + o_k[None, :]
p_g3 = g + o_c3[:, None] * (HV*K) + o_k[None, :]
# [BC, BK]
b_q3 = tl.load(p_q3, mask=m_ck3, other=0.0).to(tl.float32)
b_k3 = tl.load(p_k3, mask=m_ck3, other=0.0).to(tl.float32)
b_g3 = tl.load(p_g3, mask=m_ck3, other=0.0).to(tl.float32)
# [BK]
b_gn3 = tl.load(g + i_tc3 * HV*K + o_k, mask=m_k, other=0).to(tl.float32)
# [BC, BK]
b_gqn3 = tl.where(m_tc3[:, None], exp2(b_g3 - b_gn3[None, :]), 0)
b_qg3 = b_q3 * b_gqn3
b_kg3 = b_k3 * b_gqn3
# [BK, BC]
b_kgt = tl.trans(b_k0 * exp2(b_gn3[None, :] - b_g0))
# [BC, BC]
b_Aqk30 += tl.dot(b_qg3, b_kgt)
b_Akk30 += tl.dot(b_kg3, b_kgt)
# [BK, BC]
b_kgt = tl.trans(b_k1 * exp2(b_gn3[None, :] - b_g1))
# [BC, BC]
b_Aqk31 += tl.dot(b_qg3, b_kgt)
b_Akk31 += tl.dot(b_kg3, b_kgt)
# [BK, BC]
b_kgt = tl.trans(b_k2 * exp2(b_gn3[None, :] - b_g2))
# [BC, BC]
b_Aqk32 += tl.dot(b_qg3, b_kgt)
b_Akk32 += tl.dot(b_kg3, b_kgt)
################################################################################
# save off-diagonal Aqk blocks and prepare Akk
################################################################################
if i_tc1 < T:
p_Aqk10 = Aqk + o_c1[:, None] * (HV*BT) + o_i[None, :]
tl.store(p_Aqk10, (b_Aqk10 * scale).to(Aqk.dtype.element_ty), mask=m_A1)
p_b1 = beta + bos * HV + i_hv + o_c1 * HV
b_b1 = tl.load(p_b1, mask=m_tc1, other=0.0).to(tl.float32)
b_Akk10 = b_Akk10 * b_b1[:, None]
if NC >= 3 and i_tc2 < T:
p_Aqk20 = Aqk + o_c2[:, None] * (HV*BT) + o_i[None, :]
p_Aqk21 = Aqk + o_c2[:, None] * (HV*BT) + (o_i + BC)[None, :]
tl.store(p_Aqk20, (b_Aqk20 * scale).to(Aqk.dtype.element_ty), mask=m_A2)
tl.store(p_Aqk21, (b_Aqk21 * scale).to(Aqk.dtype.element_ty), mask=m_A2)
p_b2 = beta + bos * HV + i_hv + o_c2 * HV
b_b2 = tl.load(p_b2, mask=m_tc2, other=0.0).to(tl.float32)
b_Akk20 = b_Akk20 * b_b2[:, None]
b_Akk21 = b_Akk21 * b_b2[:, None]
if NC >= 4 and i_tc3 < T:
p_Aqk30 = Aqk + o_c3[:, None] * (HV*BT) + o_i[None, :]
p_Aqk31 = Aqk + o_c3[:, None] * (HV*BT) + (o_i + BC)[None, :]
p_Aqk32 = Aqk + o_c3[:, None] * (HV*BT) + (o_i + 2*BC)[None, :]
tl.store(p_Aqk30, (b_Aqk30 * scale).to(Aqk.dtype.element_ty), mask=m_A3)
tl.store(p_Aqk31, (b_Aqk31 * scale).to(Aqk.dtype.element_ty), mask=m_A3)
tl.store(p_Aqk32, (b_Aqk32 * scale).to(Aqk.dtype.element_ty), mask=m_A3)
p_b3 = beta + bos * HV + i_hv + o_c3 * HV
b_b3 = tl.load(p_b3, mask=m_tc3, other=0.0).to(tl.float32)
b_Akk30 = b_Akk30 * b_b3[:, None]
b_Akk31 = b_Akk31 * b_b3[:, None]
b_Akk32 = b_Akk32 * b_b3[:, None]
p_Akk00 = Akkd + o_c0[:, None] * (HV*BC) + o_i[None, :]
p_Akk11 = Akkd + o_c1[:, None] * (HV*BC) + o_i[None, :]
b_Ai00 = tl.load(p_Akk00, mask=m_A0, other=0.0).to(tl.float32)
b_Ai11 = tl.load(p_Akk11, mask=m_A1, other=0.0).to(tl.float32)
if NC >= 3:
p_Akk22 = Akkd + o_c2[:, None] * (HV*BC) + o_i[None, :]
b_Ai22 = tl.load(p_Akk22, mask=m_A2, other=0.0).to(tl.float32)
if NC >= 4:
p_Akk33 = Akkd + o_c3[:, None] * (HV*BC) + o_i[None, :]
b_Ai33 = tl.load(p_Akk33, mask=m_A3, other=0.0).to(tl.float32)
################################################################################
# forward substitution on diagonals
################################################################################
if not USE_SAFE_GATE:
m_A = o_i[:, None] > o_i[None, :]
m_I = o_i[:, None] == o_i[None, :]
b_Ai00 = -tl.where(m_A, b_Ai00, 0)
b_Ai11 = -tl.where(m_A, b_Ai11, 0)
if NC >= 3:
b_Ai22 = -tl.where(m_A, b_Ai22, 0)
if NC >= 4:
b_Ai33 = -tl.where(m_A, b_Ai33, 0)
for i in range(2, min(BC, T - i_tc0)):
b_a00 = -tl.load(Akkd + (i_tc0 + i) * HV*BC + o_i)
b_a00 = tl.where(o_i < i, b_a00, 0.)
b_a00 += tl.sum(b_a00[:, None] * b_Ai00, 0)
b_Ai00 = tl.where((o_i == i)[:, None], b_a00, b_Ai00)
for i in range(BC + 2, min(2*BC, T - i_tc0)):
b_a11 = -tl.load(Akkd + (i_tc0 + i) * HV*BC + o_i)
b_a11 = tl.where(o_i < i - BC, b_a11, 0.)
b_a11 += tl.sum(b_a11[:, None] * b_Ai11, 0)
b_Ai11 = tl.where((o_i == i - BC)[:, None], b_a11, b_Ai11)
if NC >= 3:
for i in range(2*BC + 2, min(3*BC, T - i_tc0)):
b_a22 = -tl.load(Akkd + (i_tc0 + i) * HV*BC + o_i)
b_a22 = tl.where(o_i < i - 2*BC, b_a22, 0.)
b_a22 += tl.sum(b_a22[:, None] * b_Ai22, 0)
b_Ai22 = tl.where((o_i == i - 2*BC)[:, None], b_a22, b_Ai22)
if NC >= 4:
for i in range(3*BC + 2, min(4*BC, T - i_tc0)):
b_a33 = -tl.load(Akkd + (i_tc0 + i) * HV*BC + o_i)
b_a33 = tl.where(o_i < i - 3*BC, b_a33, 0.)
b_a33 += tl.sum(b_a33[:, None] * b_Ai33, 0)
b_Ai33 = tl.where((o_i == i - 3*BC)[:, None], b_a33, b_Ai33)
b_Ai00 += m_I
b_Ai11 += m_I
if NC >= 3:
b_Ai22 += m_I
if NC >= 4:
b_Ai33 += m_I
################################################################################
# compute merged inverse using off-diagonals
################################################################################
# we used tf32 to maintain matrix inverse's precision whenever possible.
b_Ai10 = -tl.dot(
tl.dot(b_Ai11, b_Akk10, input_precision=SOLVE_TRIL_DOT_PRECISION),
b_Ai00,
input_precision=SOLVE_TRIL_DOT_PRECISION
)
if NC >= 3:
b_Ai21 = -tl.dot(
tl.dot(b_Ai22, b_Akk21, input_precision=SOLVE_TRIL_DOT_PRECISION),
b_Ai11,
input_precision=SOLVE_TRIL_DOT_PRECISION
)
b_Ai20 = -tl.dot(
b_Ai22,
tl.dot(b_Akk20, b_Ai00, input_precision=SOLVE_TRIL_DOT_PRECISION) +
tl.dot(b_Akk21, b_Ai10, input_precision=SOLVE_TRIL_DOT_PRECISION),
input_precision=SOLVE_TRIL_DOT_PRECISION
)
if NC >= 4:
b_Ai32 = -tl.dot(
tl.dot(b_Ai33, b_Akk32, input_precision=SOLVE_TRIL_DOT_PRECISION),
b_Ai22,
input_precision=SOLVE_TRIL_DOT_PRECISION
)
b_Ai31 = -tl.dot(
b_Ai33,
tl.dot(b_Akk31, b_Ai11, input_precision=SOLVE_TRIL_DOT_PRECISION) +
tl.dot(b_Akk32, b_Ai21, input_precision=SOLVE_TRIL_DOT_PRECISION),
input_precision=SOLVE_TRIL_DOT_PRECISION
)
b_Ai30 = -tl.dot(
b_Ai33,
tl.dot(b_Akk30, b_Ai00, input_precision=SOLVE_TRIL_DOT_PRECISION) +
tl.dot(b_Akk31, b_Ai10, input_precision=SOLVE_TRIL_DOT_PRECISION) +
tl.dot(b_Akk32, b_Ai20, input_precision=SOLVE_TRIL_DOT_PRECISION),
input_precision=SOLVE_TRIL_DOT_PRECISION
)
################################################################################
# store full Akk_inv to Akk
################################################################################
p_Akk00 = Akk + o_c0[:, None] * (HV*BT) + o_i[None, :]
p_Akk10 = Akk + o_c1[:, None] * (HV*BT) + o_i[None, :]
p_Akk11 = Akk + o_c1[:, None] * (HV*BT) + (o_i + BC)[None, :]
tl.store(p_Akk00, b_Ai00.to(Akk.dtype.element_ty), mask=m_A0)
tl.store(p_Akk10, b_Ai10.to(Akk.dtype.element_ty), mask=m_A1)
tl.store(p_Akk11, b_Ai11.to(Akk.dtype.element_ty), mask=m_A1)
if NC >= 3:
p_Akk20 = Akk + o_c2[:, None] * (HV*BT) + o_i[None, :]
p_Akk21 = Akk + o_c2[:, None] * (HV*BT) + (o_i + BC)[None, :]
p_Akk22 = Akk + o_c2[:, None] * (HV*BT) + (o_i + 2*BC)[None, :]
tl.store(p_Akk20, b_Ai20.to(Akk.dtype.element_ty), mask=m_A2)
tl.store(p_Akk21, b_Ai21.to(Akk.dtype.element_ty), mask=m_A2)
tl.store(p_Akk22, b_Ai22.to(Akk.dtype.element_ty), mask=m_A2)
if NC >= 4:
p_Akk30 = Akk + o_c3[:, None] * (HV*BT) + o_i[None, :]
p_Akk31 = Akk + o_c3[:, None] * (HV*BT) + (o_i + BC)[None, :]
p_Akk32 = Akk + o_c3[:, None] * (HV*BT) + (o_i + 2*BC)[None, :]
p_Akk33 = Akk + o_c3[:, None] * (HV*BT) + (o_i + 3*BC)[None, :]
tl.store(p_Akk30, b_Ai30.to(Akk.dtype.element_ty), mask=m_A3)
tl.store(p_Akk31, b_Ai31.to(Akk.dtype.element_ty), mask=m_A3)
tl.store(p_Akk32, b_Ai32.to(Akk.dtype.element_ty), mask=m_A3)
tl.store(p_Akk33, b_Ai33.to(Akk.dtype.element_ty), mask=m_A3)
@triton.heuristics({
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({}, num_warps=num_warps, num_stages=num_stages)
for num_warps in [1, 2, 4, 8]
for num_stages in [2, 3, 4]
],
key=['BK', 'NC', 'BT', 'HV'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['B', 'T'])
def chunk_kda_bwd_kernel_intra(
q,
k,
g,
beta,
dAqk,
dAkk,
dq,
dq2,
dk,
dk2,
dg,
dg2,
db,
cu_seqlens,
chunk_indices,
B,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
BT: tl.constexpr,
BC: tl.constexpr,
BK: tl.constexpr,
NC: tl.constexpr,
IS_VARLEN: tl.constexpr,
SAFE_GATE: tl.constexpr,
USE_GATHER: tl.constexpr,
):
i_kc, i_t, i_bh = tl.program_id(0), tl.program_id(1).to(tl.int64), tl.program_id(2).to(tl.int64)
i_b, i_hv = i_bh // HV, i_bh % HV
i_h = i_hv // (HV // H)
i_k, i_i = i_kc // NC, i_kc % NC
all = B * T
if IS_VARLEN:
i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64)
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
else:
bos, eos = i_b * T, i_b * T + T
T = eos - bos
i_ti = i_t * BT + i_i * BC
if i_ti >= T:
return
o_k = i_k * BK + tl.arange(0, BK)
m_k = o_k < K
q += (bos * H + i_h) * K
k += (bos * H + i_h) * K
g += (bos * HV + i_hv) * K
beta += bos * HV + i_hv
dAqk += (bos * HV + i_hv) * BT
dAkk += (bos * HV + i_hv) * BT
dq += (bos * HV + i_hv) * K
dq2 += (bos * HV + i_hv) * K
dk += (bos * HV + i_hv) * K
dk2 += (bos * HV + i_hv) * K
dg += (bos * HV + i_hv) * K
dg2 += (bos * HV + i_hv) * K
db += (i_k * all + bos) * HV + i_hv
o_i = tl.arange(0, BC)
o_c = i_ti + o_i
m_c = o_c < T
m_ck = m_c[:, None] & m_k[None, :]
m_dAf = m_c[:, None] & (o_i[None, :] < BT)
m_dAt = (o_i[:, None] < BT) & m_c[None, :]
p_g = g + o_c[:, None] * (HV*K) + o_k[None, :]
b_g = tl.load(p_g, mask=m_ck, other=0.0).to(tl.float32)
p_b = beta + o_c * HV
b_b = tl.load(p_b, mask=m_c, other=0.0)
b_dq2 = tl.zeros([BC, BK], dtype=tl.float32)
b_dk2 = tl.zeros([BC, BK], dtype=tl.float32)
if i_i > 0:
p_gn = g + i_ti * HV*K + o_k
# [BK,]
b_gn = tl.load(p_gn, mask=m_k, other=0).to(tl.float32)[None, :]
for i_j in range(0, i_i):
o_j = i_t * BT + i_j * BC + o_i
m_jk = (o_j < T)[:, None] & m_k[None, :]
p_k = k + o_j[:, None] * (H*K) + o_k[None, :]
p_gk = g + o_j[:, None] * (HV*K) + o_k[None, :]
p_dAqk = dAqk + o_c[:, None] * (HV*BT) + (i_j * BC + o_i)[None, :]
p_dAkk = dAkk + o_c[:, None] * (HV*BT) + (i_j * BC + o_i)[None, :]
# [BC, BK]
b_k = tl.load(p_k, mask=m_jk, other=0.0)
b_gk = tl.load(p_gk, mask=m_jk, other=0.0)
b_kg = b_k * exp2(b_gn - b_gk)
# [BC, BC]
b_dAqk = tl.load(p_dAqk, mask=m_dAf, other=0.0)
b_dAkk = tl.load(p_dAkk, mask=m_dAf, other=0.0)
# [BC, BK]
b_dq2 += tl.dot(b_dAqk, b_kg)
b_dk2 += tl.dot(b_dAkk, b_kg)
b_gqn = exp2(b_g - b_gn)
b_dq2 *= b_gqn
b_dk2 *= b_gqn
o_i = tl.arange(0, BC)
m_dA = (i_ti + o_i) < T
o_dA = (i_ti + o_i) * HV*BT + i_i * BC
p_kj = k + i_ti * H*K + o_k
p_gkj = g + i_ti * HV*K + o_k
p_q = q + o_c[:, None] * (H*K) + o_k[None, :]
p_k = k + o_c[:, None] * (H*K) + o_k[None, :]
b_q = tl.load(p_q, mask=m_ck, other=0.0)
b_k = tl.load(p_k, mask=m_ck, other=0.0)
if SAFE_GATE:
if USE_GATHER:
b_gn = gather(b_g, tl.full([1, BK], min(BC//2, T - i_ti - 1), dtype=tl.int16), axis=0)
else:
p_gn = g + (i_ti + min(BC // 2, T - i_ti - 1)) * HV*K + o_k
b_gn = tl.load(p_gn, mask=m_k, other=0)[None, :]
p_dAqk = dAqk + o_c[:, None] * (HV*BT) + (i_i * BC + o_i)[None, :]
p_dAkk = dAkk + o_c[:, None] * (HV*BT) + (i_i * BC + o_i)[None, :]
b_dAqk_diag_qk = tl.load(p_dAqk, mask=m_dAf, other=0.0).to(tl.float32)
b_dAkk_diag_qk = tl.load(p_dAkk, mask=m_dAf, other=0.0).to(tl.float32)
m_i_diag_qk = (o_i[:, None] >= o_i[None, :]) & ((i_ti + o_i[:, None]) < T) & ((i_ti + o_i[None, :]) < T)
m_j_diag_qk = (i_ti + o_i[:, None]) < T
b_dAqk_diag_qk = tl.where(m_i_diag_qk, b_dAqk_diag_qk, 0.)
b_dAkk_diag_qk = tl.where(m_i_diag_qk, b_dAkk_diag_qk, 0.)
b_g_diag_qk = tl.where(m_j_diag_qk, b_g - b_gn, 0.)
exp_b_g_diag_qk = tl.where(m_j_diag_qk, exp2(b_g_diag_qk), 0.)
exp_neg_b_g_diag_qk = tl.where(m_j_diag_qk, exp2(-b_g_diag_qk), 0.)
b_k_exp_diag_qk = b_k * exp_neg_b_g_diag_qk
b_dq2 += tl.dot(b_dAqk_diag_qk, b_k_exp_diag_qk) * exp_b_g_diag_qk
b_dk2 += tl.dot(b_dAkk_diag_qk, b_k_exp_diag_qk) * exp_b_g_diag_qk
else:
for j in range(0, min(BC, T - i_t * BT - i_i * BC)):
# [BC]
b_dAqk = tl.load(dAqk + o_dA + j, mask=m_dA, other=0)
b_dAkk = tl.load(dAkk + o_dA + j, mask=m_dA, other=0)
# [BK]
b_kj = tl.load(p_kj, mask=m_k, other=0).to(tl.float32)
b_gkj = tl.load(p_gkj, mask=m_k, other=0).to(tl.float32)
# [BC, BK]
m_i = o_i[:, None] >= j
# [BC, BK]
b_gqk = exp2(b_g - b_gkj[None, :])
b_dq2 += tl.where(m_i, b_dAqk[:, None] * b_kj[None, :] * b_gqk, 0.)
b_dk2 += tl.where(m_i, b_dAkk[:, None] * b_kj[None, :] * b_gqk, 0.)
p_kj += H*K
p_gkj += HV*K
b_db = tl.sum(b_dk2 * b_k, 1)
b_dk2 *= b_b[:, None]
p_dq = dq + o_c[:, None] * (HV*K) + o_k[None, :]
p_dq2 = dq2 + o_c[:, None] * (HV*K) + o_k[None, :]
p_db = db + o_c * HV
b_dg2 = b_q * b_dq2
b_dq2 = b_dq2 + tl.load(p_dq, mask=m_ck, other=0.0)
tl.store(p_dq2, b_dq2.to(p_dq2.dtype.element_ty), mask=m_ck)
tl.store(p_db, b_db.to(p_db.dtype.element_ty), mask=m_c)
tl.debug_barrier()
b_dkt = tl.zeros([BC, BK], dtype=tl.float32)
NC = min(NC, tl.cdiv(T - i_t * BT, BC))
if i_i < NC - 1:
p_gn = g + (min(i_ti + BC, T) - 1) * HV*K + o_k
# [BK,]
b_gn = tl.load(p_gn, mask=m_k, other=0).to(tl.float32)[None, :]
for i_j in range(i_i + 1, NC):
o_j = i_t * BT + i_j * BC + o_i
m_j = o_j < T
m_jk = m_j[:, None] & m_k[None, :]
m_dAj = (o_i[:, None] < BT) & m_j[None, :]
p_q = q + o_j[:, None] * (H*K) + o_k[None, :]
p_k = k + o_j[:, None] * (H*K) + o_k[None, :]
p_gk = g + o_j[:, None] * (HV*K) + o_k[None, :]
p_b = beta + o_j * HV
p_dAqk = dAqk + (i_i * BC + o_i)[:, None] + o_j[None, :] * (HV*BT)
p_dAkk = dAkk + (i_i * BC + o_i)[:, None] + o_j[None, :] * (HV*BT)
# [BC]
b_b = tl.load(p_b, mask=m_j, other=0.0)
# [BC, BK]
b_q = tl.load(p_q, mask=m_jk, other=0.0)
b_kb = tl.load(p_k, mask=m_jk, other=0.0) * b_b[:, None]
b_gk = tl.load(p_gk, mask=m_jk, other=0.0).to(tl.float32)
# [BC, BC]
b_dAqk = tl.load(p_dAqk, mask=m_dAj, other=0.0)
b_dAkk = tl.load(p_dAkk, mask=m_dAj, other=0.0)
# [BC, BK]
b_gkn = exp2(b_gk - b_gn)
b_qg = b_q * tl.where(m_j[:, None], b_gkn, 0)
b_kbg = b_kb * tl.where(m_j[:, None], b_gkn, 0)
# [BC, BK]
# (SY 09/17) important to not use bf16 here to have a good precision.
b_dkt += tl.dot(b_dAqk, b_qg)
b_dkt += tl.dot(b_dAkk, b_kbg)
b_dkt *= exp2(b_gn - b_g)
o_dA = i_ti * HV*BT + i_i * BC + o_i
p_qj = q + i_ti * H*K + o_k
p_kj = k + i_ti * H*K + o_k
p_gkj = g + i_ti * HV*K + o_k
p_bj = beta + i_ti * HV
if SAFE_GATE:
if USE_GATHER:
b_gn = gather(b_g, tl.full([1, BK], min(BC//2, T - i_ti - 1), dtype=tl.int16), axis=0)
else:
p_gn = g + (i_ti + min(BC // 2, T - i_ti - 1)) * HV*K + o_k
b_gn = tl.load(p_gn, mask=m_k, other=0).to(tl.float32)[None, :]
p_q = q + o_c[:, None] * (H*K) + o_k[None, :]
b_q = tl.load(p_q, mask=m_ck, other=0.0)
p_b = beta + o_c * HV
b_b = tl.load(p_b, mask=m_c, other=0.0)
p_dAqk = dAqk + (i_i * BC + o_i)[:, None] + o_c[None, :] * (HV*BT)
p_dAkk = dAkk + (i_i * BC + o_i)[:, None] + o_c[None, :] * (HV*BT)
b_dAqk_diag_kk = tl.load(p_dAqk, mask=m_dAt, other=0.0).to(tl.float32)
b_dAkk_diag_kk = tl.load(p_dAkk, mask=m_dAt, other=0.0).to(tl.float32)
m_i_diag_kk = (o_i[:, None] <= o_i[None, :]) & ((i_ti + o_i[:, None]) < T) & ((i_ti + o_i[None, :]) < T)
m_j_diag_kk = (i_ti + o_i[:, None]) < T
b_dAqk_diag_kk = tl.where(m_i_diag_kk, b_dAqk_diag_kk, 0.)
b_dAkk_diag_kk = tl.where(m_i_diag_kk, b_dAkk_diag_kk, 0.)
# ensure numerical stability
b_g_diag_kk = tl.where(m_j_diag_kk, b_g - b_gn, 0.)
exp_b_g_diag_kk = tl.where(m_j_diag_kk, exp2(b_g_diag_kk), 0.)
exp_neg_b_g_diag_kk = tl.where(m_j_diag_kk, exp2(-b_g_diag_kk), 0.)
b_q_exp = b_q * exp_b_g_diag_kk
b_kb_exp = b_k * b_b[:, None] * exp_b_g_diag_kk
b_dkt += tl.dot(b_dAqk_diag_kk, b_q_exp) * exp_neg_b_g_diag_kk
b_dkt += tl.dot(b_dAkk_diag_kk, b_kb_exp) * exp_neg_b_g_diag_kk
else:
for j in range(0, min(BC, T - i_t * BT - i_i * BC)):
# [BC,]
b_dAqk = tl.load(dAqk + o_dA + j * HV*BT)
b_dAkk = tl.load(dAkk + o_dA + j * HV*BT)
# [BK,]
b_qj = tl.load(p_qj, mask=m_k, other=0).to(tl.float32)
b_kbj = tl.load(p_kj, mask=m_k, other=0).to(tl.float32) * tl.load(p_bj)
b_gkj = tl.load(p_gkj, mask=m_k, other=0).to(tl.float32)
# [BC, BK]
m_i = o_i[:, None] <= j
b_gkq = exp2(b_gkj[None, :] - b_g)
b_dkt += tl.where(m_i, b_dAqk[:, None] * b_qj[None, :] * b_gkq, 0.)
b_dkt += tl.where(m_i, b_dAkk[:, None] * b_kbj[None, :] * b_gkq, 0.)
p_qj += H*K
p_kj += H*K
p_gkj += HV*K
p_bj += HV
p_dk = dk + o_c[:, None] * (HV*K) + o_k[None, :]
p_dk2 = dk2 + o_c[:, None] * (HV*K) + o_k[None, :]
p_dg = dg + o_c[:, None] * (HV*K) + o_k[None, :]
p_dg2 = dg2 + o_c[:, None] * (HV*K) + o_k[None, :]
b_dg2 += (b_dk2 - b_dkt) * b_k + tl.load(p_dg, mask=m_ck, other=0.0)
b_dk2 += tl.load(p_dk, mask=m_ck, other=0.0)
b_dk2 += b_dkt
tl.store(p_dk2, b_dk2.to(p_dk2.dtype.element_ty), mask=m_ck)
tl.store(p_dg2, b_dg2.to(p_dg2.dtype.element_ty), mask=m_ck)
@triton.heuristics({
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({}, num_warps=num_warps, num_stages=num_stages)
for num_warps in [1, 2, 4, 8]
for num_stages in [2, 3, 4]
],
key=["BT", "BC", "HV"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_kda_fwd_kernel_intra_sub_chunk(
q,
k,
g,
beta,
Aqk,
Akk,
scale,
cu_seqlens,
chunk_indices,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
BT: tl.constexpr,
BC: tl.constexpr,
BK: tl.constexpr,
IS_VARLEN: tl.constexpr,
USE_GATHER: tl.constexpr,
):
i_t, i_i, i_bh = tl.program_id(0).to(tl.int64), tl.program_id(1), tl.program_id(2).to(tl.int64)
i_b, i_hv = i_bh // HV, i_bh % HV
i_h = i_hv // (HV // H)
if IS_VARLEN:
i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64)
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
else:
bos, eos = i_b * T, i_b * T + T
i_ti = i_t * BT + i_i * BC
if i_ti >= T:
return
o_c = i_ti + tl.arange(0, BC)
m_c = o_c < T
q = q + (bos * H + i_h) * K
k = k + (bos * H + i_h) * K
g = g + (bos * HV + i_hv) * K
beta = beta + bos * HV + i_hv
Aqk = Aqk + (bos * HV + i_hv) * BT
Akk = Akk + (bos * HV + i_hv) * BC
o_k = tl.arange(0, BK)
m_k = o_k < K
m_ck = m_c[:, None] & m_k[None, :]
p_q = q + o_c[:, None] * (H*K) + o_k[None, :]
p_k = k + o_c[:, None] * (H*K) + o_k[None, :]
p_g = g + o_c[:, None] * (HV*K) + o_k[None, :]
p_beta = beta + o_c * HV
b_q = tl.load(p_q, mask=m_ck, other=0.0)
b_k = tl.load(p_k, mask=m_ck, other=0.0)
b_g = tl.load(p_g, mask=m_ck, other=0.0)
b_beta = tl.load(p_beta, mask=m_c, other=0.0)
if USE_GATHER:
b_gn = gather(b_g, tl.full([1, BK], min(BC//2, T - i_ti - 1), dtype=tl.int16), axis=0)
else:
# caculate offset
p_gn = g + (i_ti + min(BC // 2, T - i_ti - 1)) * HV*K + tl.arange(0, BK)
b_gn = tl.load(p_gn, mask=tl.arange(0, BK) < K, other=0.0)
b_gn = b_gn[None, :]
# current block, keep numerical stability by subtracting the left boundary
# less than 85 to avoid overflow in exp2
b_gm = (b_g - b_gn).to(tl.float32)
b_gq = tl.where(m_c[:, None], exp2(b_gm), 0.)
b_gk = tl.where(m_c[:, None], exp2(-b_gm), 0.)
b_kgt = tl.trans(b_k * b_gk)
b_Aqk = tl.dot(b_q * b_gq, b_kgt) * scale
b_Akk = tl.dot(b_k * b_gq, b_kgt) * b_beta[:, None]
o_i = tl.arange(0, BC)
m_Aqk = o_i[:, None] >= o_i[None, :]
m_Akk = o_i[:, None] > o_i[None, :]
m_I = o_i[:, None] == o_i[None, :]
b_Aqk = tl.where(m_Aqk, b_Aqk, 0.0)
b_Akk = tl.where(m_Akk, b_Akk, 0.0)
m_Aqk_st = m_c[:, None] & (o_i[None, :] < BT)
m_Akk_st = m_c[:, None] & (o_i[None, :] < BC)
p_Aqk = Aqk + o_c[:, None] * (HV*BT) + (i_i * BC + o_i)[None, :]
p_Akk = Akk + o_c[:, None] * (HV*BC) + o_i[None, :]
tl.store(p_Aqk, b_Aqk.to(Aqk.dtype.element_ty), mask=m_Aqk_st)
tl.store(p_Akk, b_Akk.to(Akk.dtype.element_ty), mask=m_Akk_st)
tl.debug_barrier()
################################################################################
# forward substitution
################################################################################
b_Ai = -b_Akk
for i in range(2, min(BC, T - i_ti)):
b_a = -tl.load(Akk + (i_ti + i) * HV*BC + o_i)
b_a = tl.where(o_i < i, b_a, 0.)
b_a += tl.sum(b_a[:, None] * b_Ai, 0)
b_Ai = tl.where((o_i == i)[:, None], b_a, b_Ai)
b_Ai += m_I
tl.store(p_Akk, b_Ai.to(Akk.dtype.element_ty), mask=m_Akk_st)
@dispatch('kda')
def chunk_kda_fwd_intra(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
gk: torch.Tensor | None = None,
beta: torch.Tensor | None = None,
scale: float | None = None,
cu_seqlens: torch.LongTensor | None = None,
chunk_size: int = 64,
chunk_indices: torch.LongTensor | None = None,
safe_gate: bool = False,
disable_recompute: bool = False,
):
B, T, H, K, HV = *k.shape, gk.shape[2]
BT = chunk_size
if BT not in (32, 64):
raise ValueError(f"KDA intra chunk kernel only supports chunk_size 32 or 64, got {BT}.")
BC = 16
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
NC = triton.cdiv(BT, BC)
Aqk = torch.empty(B, T, HV, BT, device=k.device, dtype=k.dtype)
# Akk must be zero-initialized - kernel only writes lower triangular
Akk = torch.zeros(B, T, HV, BT, device=k.device, dtype=k.dtype)
# Separate fp32 buffer for diagonal 16x16 blocks (for precision in solve_tril)
Akkd = torch.empty(B, T, HV, BC, device=k.device, dtype=torch.float32)
# Step 1: Run token_parallel first to compute diagonal blocks into Akkd (fp32)
# Step 1: compute diagonal blocks into Akk_diag (fp32)
if safe_gate:
grid = (NT, NC, B * HV)
BK = triton.next_power_of_2(K)
chunk_kda_fwd_kernel_intra_sub_chunk[grid](
q=q,
k=k,
g=gk,
beta=beta,
Aqk=Aqk,
Akk=Akkd,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
T=T,
H=H,
HV=HV,
K=K,
BT=BT,
BC=BC,
BK=BK,
USE_GATHER=IS_GATHER_SUPPORTED,
)
else:
Aqk, Akkd = chunk_kda_fwd_intra_token_parallel(
q=q,
k=k,
gk=gk,
beta=beta,
Aqk=Aqk,
Akk=Akkd,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_size=BT,
sub_chunk_size=BC,
)
# Step 2: Fused inter + solve_tril (works for both fixed-len and varlen)
grid = (NT, B * HV)
chunk_kda_fwd_kernel_inter_solve_fused[grid](
q=q,
k=k,
g=gk,
beta=beta,
Aqk=Aqk,
Akkd=Akkd,
Akk=Akk,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
T=T,
H=H,
HV=HV,
K=K,
BT=BT,
BC=BC,
NC=NC,
USE_SAFE_GATE=safe_gate,
)
w, u, qg, kg = recompute_w_u_fwd(
k=k,
v=v,
beta=beta,
A=Akk,
q=q if disable_recompute else None,
gk=gk,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
)
return w, u, qg, kg, Aqk, Akk
@dispatch('kda')
def chunk_kda_bwd_intra(
q: torch.Tensor,
k: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
dAqk: torch.Tensor,
dAkk: torch.Tensor,
dq: torch.Tensor,
dk: torch.Tensor,
db: torch.Tensor,
dg: torch.Tensor,
cu_seqlens: torch.LongTensor | None = None,
chunk_indices: torch.LongTensor | None = None,
chunk_size: int = 64,
safe_gate: bool = False,
):
B, T, H, K, HV = *k.shape, g.shape[2]
BT = chunk_size
BC = min(16, BT)
BK = min(32, triton.next_power_of_2(K))
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
NC = triton.cdiv(BT, BC)
NK = triton.cdiv(K, BK)
dq2 = torch.empty_like(dq)
dk2 = torch.empty_like(dk)
db2 = beta.new_empty(NK, *beta.shape, dtype=torch.float)
dg2 = torch.empty_like(dg, dtype=torch.float)
grid = (NK * NC, NT, B * HV)
chunk_kda_bwd_kernel_intra[grid](
q=q,
k=k,
g=g,
beta=beta,
dAqk=dAqk,
dAkk=dAkk,
dq=dq,
dq2=dq2,
dk=dk,
dk2=dk2,
dg=dg,
dg2=dg2,
db=db2,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
B=B,
T=T,
H=H,
HV=HV,
K=K,
BT=BT,
BC=BC,
BK=BK,
NC=NC,
SAFE_GATE=safe_gate,
USE_GATHER=IS_GATHER_SUPPORTED,
)
dq = dq2
dk = dk2
db = db2.sum(0).add_(db)
dg = dg2
return dq, dk, db, dg
@@ -0,0 +1,182 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
# Token-parallel implementation of KDA intra chunk kernel
import torch
import triton
import triton.language as tl
from kda._fla.ops.backends import dispatch
from kda._fla.ops.utils.cache import fla_cache_autotune
from kda._fla.ops.utils.op import exp2
from kda._fla.utils import autotune_cache_kwargs
@triton.heuristics({
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({'BH': BH}, num_warps=num_warps)
for BH in [1, 2, 4, 8]
for num_warps in [1, 2, 4, 8]
],
key=["K", "H", "HV"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T', 'N'])
def chunk_kda_fwd_kernel_intra_token_parallel(
q,
k,
g,
beta,
Aqk,
Akk,
scale,
cu_seqlens,
N,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
BK: tl.constexpr,
BT: tl.constexpr,
BC: tl.constexpr,
BH: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_tg, i_hg = tl.program_id(0).to(tl.int64), tl.program_id(1)
if IS_VARLEN:
i_n = 0
left, right = 0, N
# Unrolled binary search (max B=2^32)
# We can limit iterations based on expected max batch size if needed
# 20 iterations covers B=1M, usually enough
for _ in range(20):
if left < right:
mid = (left + right) // 2
if i_tg < tl.load(cu_seqlens + mid + 1).to(tl.int32):
right = mid
else:
left = mid + 1
i_n = left
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
i_t = i_tg - bos
else:
bos = (i_tg // T) * T
i_t = i_tg % T
if i_t >= T:
return
i_c = i_t // BT
i_s = (i_t % BT) // BC
i_tc = i_c * BT
i_ts = i_tc + i_s * BC
G: tl.constexpr = HV // H
q += bos * H*K
k += bos * H*K
g += bos * HV*K
Aqk += bos * HV*BT
Akk += bos * HV*BC
beta += bos * HV
o_hv = i_hg * BH + tl.arange(0, BH)
o_h = o_hv // G
o_k = tl.arange(0, BK)
m_hv = o_hv < HV
m_k = o_k < K
m_hk = m_hv[:, None] & m_k[None, :]
# q/k: [B, T, H, K], manual load via mapped qk head index
p_qk = o_h[:, None] * K + o_k[None, :]
b_q = tl.load(q + i_t * H * K + p_qk, mask=m_hk, other=0).to(tl.float32)
b_k = tl.load(k + i_t * H * K + p_qk, mask=m_hk, other=0).to(tl.float32)
# g: [B, T, HV, K], beta: [B, T, HV]
p_g = g + i_t * HV * K + o_hv[:, None] * K + o_k[None, :]
p_beta = beta + i_t * HV + o_hv
b_g = tl.load(p_g, mask=m_hk, other=0.0).to(tl.float32)
b_k = b_k * tl.load(p_beta, mask=m_hv, other=0.0).to(tl.float32)[:, None]
for j in range(i_ts, min(i_t + 1, min(T, i_ts + BC))):
b_kj = tl.load(k + j * H * K + p_qk, mask=m_hk, other=0).to(tl.float32)
p_gj = g + j * HV * K + o_hv[:, None] * K + o_k[None, :]
b_gj = tl.load(p_gj, mask=m_hk, other=0.0).to(tl.float32)
b_kgj = tl.where(m_k[None, :], b_kj * exp2(b_g - b_gj), 0.0)
b_Aqk = tl.sum(b_q * b_kgj, axis=1) * scale
b_Akk = tl.sum(b_k * b_kgj, axis=1) * tl.where(j < i_t, 1.0, 0.0)
tl.store(Aqk + i_t * HV * BT + o_hv * BT + j % BT, b_Aqk.to(Aqk.dtype.element_ty), mask=m_hv)
tl.store(Akk + i_t * HV * BC + o_hv * BC + j - i_ts, b_Akk.to(Akk.dtype.element_ty), mask=m_hv)
@dispatch('kda')
def chunk_kda_fwd_intra_token_parallel(
q: torch.Tensor,
k: torch.Tensor,
gk: torch.Tensor,
beta: torch.Tensor,
Aqk: torch.Tensor,
Akk: torch.Tensor,
scale: float,
cu_seqlens: torch.LongTensor | None = None,
chunk_size: int = 64,
sub_chunk_size: int = 16,
) -> None:
"""
Token-parallel implementation: each token gets its own thread block.
Supports both fixed-length and variable-length sequences.
Reduces wasted computation on padding.
Writes directly to Aqk and Akk tensors (in-place).
Args:
q: [B, T, H, K]
k: [B, T, H, K]
gk: [B, T, HV, K] cumsum of gates (HV >= H for GVA)
beta: [B, T, HV]
Aqk: [B, T, HV, BT] output tensor to write to
Akk: [B, T, HV, BC] output tensor for diagonal blocks (fp32)
scale: attention scale
chunk_size: BT (default 64)
sub_chunk_size: BC (default 16)
"""
B, T, H, K, HV = *q.shape, gk.shape[2]
N = len(cu_seqlens) - 1 if cu_seqlens is not None else B
BT = chunk_size
BC = sub_chunk_size
BK = triton.next_power_of_2(K)
def grid(meta): return (B * T, triton.cdiv(HV, meta['BH']))
chunk_kda_fwd_kernel_intra_token_parallel[grid](
q=q,
k=k,
g=gk,
beta=beta,
Aqk=Aqk,
Akk=Akk,
scale=scale,
cu_seqlens=cu_seqlens,
N=N,
T=T,
H=H,
HV=HV,
K=K,
BK=BK,
BT=BT,
BC=BC,
)
return Aqk, Akk
+491
View File
@@ -0,0 +1,491 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
# This kernel is modified from the Decode kernel of the vllm gdn/kda model.
import warnings
import torch
import triton
import triton.language as tl
from kda._fla.ops.backends import dispatch
from kda._fla.ops.utils.op import exp
from kda._fla.ops.utils.softplus import softplus
from kda._fla.utils import input_guard
@triton.heuristics(
{
"USE_INITIAL_STATE": lambda args: args["h0"] is not None,
"STORE_FINAL_STATE": lambda args: args["ht"] is not None,
"IS_VARLEN": lambda args: args["cu_seqlens"] is not None,
"IS_CONTINUOUS_BATCHING": lambda args: args["ssm_state_indices"] is not None,
"IS_SPEC_DECODING": lambda args: args["num_accepted_tokens"] is not None,
"HAS_A": lambda args: args["A_log"] is not None,
"HAS_BIAS": lambda args: args["dt_bias"] is not None,
"USE_LOWER_BOUND": lambda args: args["lower_bound"] is not None,
}
)
@triton.jit(do_not_specialize=["N", "T"])
def fused_recurrent_kda_fwd_kernel(
q,
k,
v,
g,
beta,
A_log,
dt_bias,
o,
h0,
ht,
cu_seqlens,
ssm_state_indices,
num_accepted_tokens,
lower_bound,
scale: tl.constexpr,
N: tl.int64, # num of sequences
T: tl.int64, # num of tokens
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BK: tl.constexpr,
BV: tl.constexpr,
stride_init_state_token: tl.constexpr,
stride_final_state_token: tl.constexpr,
stride_indices_seq: tl.constexpr,
stride_indices_tok: tl.constexpr,
USE_INITIAL_STATE: tl.constexpr, # whether to use initial state
INPLACE_FINAL_STATE: tl.constexpr, # whether to store final state inplace
IS_BETA_HEADWISE: tl.constexpr, # whether beta is headwise vector or scalar,
USE_QK_L2NORM_IN_KERNEL: tl.constexpr,
IS_VARLEN: tl.constexpr,
IS_CONTINUOUS_BATCHING: tl.constexpr,
IS_SPEC_DECODING: tl.constexpr,
STORE_FINAL_STATE: tl.constexpr,
HAS_A: tl.constexpr,
HAS_BIAS: tl.constexpr,
USE_GATE_IN_KERNEL: tl.constexpr,
USE_LOWER_BOUND: tl.constexpr,
APPLY_BETA_SIGMOID: tl.constexpr,
ALLOW_NEG_EIGVAL: tl.constexpr,
STATE_V_FIRST: tl.constexpr,
num_stages: tl.constexpr,
):
pid = tl.program_id(0).to(tl.int64)
NV = tl.cdiv(V, BV)
NK = tl.cdiv(K, BK)
i_k = pid % NK
pid_rest = pid // NK
i_v = pid_rest % NV
i_nh = pid_rest // NV
i_n, i_hv = i_nh // HV, i_nh % HV
i_h = i_hv // (HV // H)
if IS_VARLEN:
bos, eos = (
tl.load(cu_seqlens + i_n).to(tl.int64),
tl.load(cu_seqlens + i_n + 1).to(tl.int64),
)
T = eos - bos
else:
bos, eos = i_n * T, i_n * T + T
if T == 0:
# no tokens to process for this sequence
return
o_k = i_k * BK + tl.arange(0, BK)
o_v = i_v * BV + tl.arange(0, BV)
p_q = q + (bos * H + i_h) * K + o_k
p_k = k + (bos * H + i_h) * K + o_k
p_v = v + (bos * HV + i_hv) * V + o_v
if IS_BETA_HEADWISE:
p_beta = beta + (bos * HV + i_hv) * V + o_v
else:
p_beta = beta + bos * HV + i_hv
p_g = g + (bos * HV + i_hv) * K + o_k
p_o = o + (bos * HV + i_hv) * V + o_v
mask_k = o_k < K
mask_v = o_v < V
if STATE_V_FIRST:
mask_h = mask_v[:, None] & mask_k[None, :]
else:
mask_h = mask_k[:, None] & mask_v[None, :]
if STATE_V_FIRST:
b_h = tl.zeros([BV, BK], dtype=tl.float32)
else:
b_h = tl.zeros([BK, BV], dtype=tl.float32)
if USE_INITIAL_STATE:
if IS_CONTINUOUS_BATCHING:
if IS_SPEC_DECODING:
i_t = tl.load(num_accepted_tokens + i_n).to(tl.int64) - 1
else:
i_t = 0
p_h0 = (
h0
+ tl.load(ssm_state_indices + i_n * stride_indices_seq + i_t).to(
tl.int64
)
* stride_init_state_token
)
if STATE_V_FIRST:
p_h0 = p_h0 + i_hv * K * V + o_v[:, None] * K + o_k[None, :]
else:
p_h0 = p_h0 + i_hv * K * V + o_k[:, None] * V + o_v[None, :]
else:
if STATE_V_FIRST:
p_h0 = h0 + (i_n * HV + i_hv) * K * V + o_v[:, None] * K + o_k[None, :]
else:
p_h0 = h0 + (i_n * HV + i_hv) * K * V + o_k[:, None] * V + o_v[None, :]
b_h += tl.load(p_h0, mask=mask_h, other=0).to(tl.float32)
for i_t in tl.range(0, T, num_stages=num_stages):
b_q = tl.load(p_q, mask=mask_k, other=0, eviction_policy='evict_last').to(tl.float32)
b_k = tl.load(p_k, mask=mask_k, other=0, eviction_policy='evict_last').to(tl.float32)
b_v = tl.load(p_v, mask=mask_v, other=0, eviction_policy='evict_first').to(tl.float32)
if USE_QK_L2NORM_IN_KERNEL:
b_q = b_q / tl.sqrt(tl.sum(b_q * b_q) + 1e-6)
b_k = b_k / tl.sqrt(tl.sum(b_k * b_k) + 1e-6)
b_q = b_q * scale
b_g = tl.load(p_g, mask=mask_k, other=0, eviction_policy='evict_last').to(tl.float32)
if USE_GATE_IN_KERNEL:
b_A = tl.load(A_log + i_hv).to(tl.float32) if HAS_A else 1.0
if HAS_BIAS:
b_bias = tl.load(dt_bias + i_hv * K + o_k, mask=mask_k, other=0).to(tl.float32)
b_g = b_g + b_bias
if USE_LOWER_BOUND:
b_gk = lower_bound * tl.sigmoid((exp(b_A) if HAS_A else b_A) * b_g)
else:
b_gk = -exp(b_A) * softplus(b_g)
else:
b_gk = b_g
if STATE_V_FIRST:
b_h *= exp(b_gk[None, :])
else:
b_h *= exp(b_gk[:, None])
if STATE_V_FIRST:
b_v -= tl.sum(b_h * b_k[None, :], 1)
else:
b_v -= tl.sum(b_h * b_k[:, None], 0)
if IS_BETA_HEADWISE:
b_beta = tl.load(p_beta, mask=mask_v, other=0, eviction_policy='evict_first').to(tl.float32)
else:
b_beta = tl.load(p_beta, eviction_policy='evict_last').to(tl.float32)
if APPLY_BETA_SIGMOID:
b_beta = tl.sigmoid(b_beta)
if ALLOW_NEG_EIGVAL:
b_beta = b_beta * 2
b_v *= b_beta
if STATE_V_FIRST:
b_h += b_v[:, None] * b_k[None, :]
b_o = tl.sum(b_h * b_q[None, :], 1)
else:
b_h += b_k[:, None] * b_v[None, :]
b_o = tl.sum(b_h * b_q[:, None], 0)
tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=mask_v, eviction_policy='evict_first')
if IS_CONTINUOUS_BATCHING:
if INPLACE_FINAL_STATE:
p_ht = (
ht
+ tl.load(ssm_state_indices + i_n * stride_indices_seq + i_t).to(
tl.int64
)
* stride_final_state_token
)
else:
p_ht = ht + (bos + i_t) * stride_final_state_token
if STATE_V_FIRST:
p_ht = p_ht + i_hv * K * V + o_v[:, None] * K + o_k[None, :]
else:
p_ht = p_ht + i_hv * K * V + o_k[:, None] * V + o_v[None, :]
tl.store(p_ht, b_h.to(p_ht.dtype.element_ty), mask=mask_h)
p_q += H * K
p_k += H * K
p_o += HV * V
p_v += HV * V
p_g += HV * K
p_beta += HV * (V if IS_BETA_HEADWISE else 1)
if not IS_CONTINUOUS_BATCHING:
if STORE_FINAL_STATE:
if STATE_V_FIRST:
p_ht = ht + (i_n * HV + i_hv) * K * V + o_v[:, None] * K + o_k[None, :]
else:
p_ht = ht + (i_n * HV + i_hv) * K * V + o_k[:, None] * V + o_v[None, :]
tl.store(p_ht, b_h.to(p_ht.dtype.element_ty), mask=mask_h)
@dispatch("kda")
def fused_recurrent_kda_fwd(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
initial_state: torch.Tensor | None = None,
scale: float | None = None,
output_final_state: bool = False,
inplace_final_state: bool = True,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
ssm_state_indices: torch.Tensor | None = None,
num_accepted_tokens: torch.Tensor | None = None,
use_qk_l2norm_in_kernel: bool = False,
use_gate_in_kernel: bool = False,
use_beta_sigmoid_in_kernel: bool = False,
allow_neg_eigval: bool = False,
lower_bound: float | None = None,
out: torch.Tensor | None = None,
**kwargs,
) -> tuple[torch.Tensor, torch.Tensor]:
if scale is None:
scale = k.shape[-1] ** -0.5
B, T, H, K, V = *k.shape, v.shape[-1]
HV = v.shape[2]
N = B if cu_seqlens is None else len(cu_seqlens) - 1
BK = triton.next_power_of_2(K)
BV = 32
if out is None:
out = torch.zeros_like(v)
else:
assert out.shape == v.shape
if inplace_final_state:
assert initial_state is not None
final_state = initial_state
elif output_final_state:
if state_v_first:
final_state = q.new_empty(N, HV, V, K, dtype=torch.float32)
else:
final_state = q.new_empty(N, HV, K, V, dtype=torch.float32)
else:
final_state = None
stride_init_state_token = initial_state.stride(0) if initial_state is not None else 1
stride_final_state_token = final_state.stride(0) if final_state is not None else 1
if ssm_state_indices is None:
stride_indices_seq, stride_indices_tok = 1, 1
elif ssm_state_indices.ndim == 1:
stride_indices_seq, stride_indices_tok = ssm_state_indices.stride(0), 1
else:
stride_indices_seq, stride_indices_tok = ssm_state_indices.stride()
grid = (triton.cdiv(V, BV) * N * HV, )
fused_recurrent_kda_fwd_kernel[grid](
q=q,
k=k,
v=v,
g=g,
beta=beta,
A_log=A_log,
dt_bias=dt_bias,
o=out,
h0=initial_state,
ht=final_state,
cu_seqlens=cu_seqlens,
ssm_state_indices=ssm_state_indices,
num_accepted_tokens=num_accepted_tokens,
lower_bound=lower_bound,
scale=scale,
N=N,
T=T,
H=H,
HV=HV,
K=K,
V=V,
BK=BK,
BV=BV,
stride_init_state_token=stride_init_state_token,
stride_final_state_token=stride_final_state_token,
stride_indices_seq=stride_indices_seq,
stride_indices_tok=stride_indices_tok,
IS_BETA_HEADWISE=beta.ndim == v.ndim,
USE_QK_L2NORM_IN_KERNEL=use_qk_l2norm_in_kernel,
INPLACE_FINAL_STATE=inplace_final_state,
USE_GATE_IN_KERNEL=use_gate_in_kernel,
APPLY_BETA_SIGMOID=use_beta_sigmoid_in_kernel,
ALLOW_NEG_EIGVAL=allow_neg_eigval,
STATE_V_FIRST=state_v_first,
num_warps=4,
num_stages=2,
)
return out, final_state
@input_guard
def fused_recurrent_kda(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
scale: float | None = None,
initial_state: torch.Tensor = None,
output_final_state: bool = False,
use_qk_l2norm_in_kernel: bool = False,
use_gate_in_kernel: bool = False,
use_beta_sigmoid_in_kernel: bool = False,
allow_neg_eigval: bool = False,
lower_bound: float | None = None,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
**kwargs,
) -> tuple[torch.Tensor, torch.Tensor]:
r"""
Args:
q (torch.Tensor):
queries of shape `[B, T, H, K]`.
k (torch.Tensor):
keys of shape `[B, T, H, K]`.
v (torch.Tensor):
values of shape `[B, T, HV, V]`.
GVA is applied if `HV > H`.
g (torch.Tensor):
g (decays) of shape `[B, T, HV, K]`.
beta (torch.Tensor):
betas of shape `[B, T, HV]`.
A_log (Optional[torch.Tensor]):
Decay parameter of shape `[HV]`.
When `use_gate_in_kernel=True` together with `lower_bound`,
may be `None` to use `lower_bound * sigmoid(g + dt_bias)`.
dt_bias (Optional[torch.Tensor]):
Bias added to `g` before activation, of shape `[HV]`. Only used when `use_gate_in_kernel=True`.
scale (Optional[float]):
Scale factor for the RetNet attention scores.
If not provided, it will default to `1 / sqrt(K)`. Default: `None`.
initial_state (Optional[torch.Tensor]):
Initial state of shape `[N, HV, K, V]` for `N` input sequences.
For equal-length input sequences, `N` equals the batch size `B`.
Default: `None`.
output_final_state (Optional[bool]):
Whether to output the final state of shape `[N, HV, K, V]`. Default: `False`.
use_qk_l2norm_in_kernel (Optional[bool]):
Whether to use L2 normalization in the kernel. Default: `False`.
use_gate_in_kernel (Optional[bool]):
Whether to compute the log-space KDA decay internally.
When `True`, `g` is the raw input and the kernel fuses gate activation into the recurrence.
Default: `False`.
use_beta_sigmoid_in_kernel (Optional[bool]):
Whether to apply `torch.sigmoid(beta)` inside the kernel.
- If `True`, the passed `beta` acts as the raw beta logits.
- If `False`, `beta` is expected to already be in post-sigmoid space.
Default: `False`.
allow_neg_eigval (Optional[bool]):
Whether to allow negative eigenvalues by scaling `beta` to `[0, 2)`.
Only takes effect together with `use_beta_sigmoid_in_kernel=True`, in which case
the kernel computes `2 * sigmoid(beta)` instead of `sigmoid(beta)`. Default: `False`.
lower_bound (Optional[float]):
Lower bound for the forget gate (in log space). Only used when `use_gate_in_kernel=True`. Default: `None`.
state_v_first (Optional[bool]):
Store the recurrent state in V-first ``[V, K]`` layout instead of the default ``[K, V]``. Default: ``False``.
cu_seqlens (torch.LongTensor):
Cumulative sequence lengths of shape `[N+1]` used for variable-length training,
consistent with the FlashAttention API.
Returns:
o (torch.Tensor):
Outputs of shape `[B, T, HV, V]`.
final_state (torch.Tensor):
Final state of shape `[N, HV, K, V]` if `output_final_state=True` else `None`.
Examples::
>>> import torch
>>> import torch.nn.functional as F
>>> from einops import rearrange
>>> from fla.ops.kda import fused_recurrent_kda
# inputs with equal lengths
>>> B, T, H, HV, K, V = 4, 2048, 4, 8, 512, 512
>>> q = torch.randn(B, T, H, K, device='cuda')
>>> k = F.normalize(torch.randn(B, T, H, K, device='cuda'), p=2, dim=-1)
>>> v = torch.randn(B, T, HV, V, device='cuda')
>>> g = F.logsigmoid(torch.rand(B, T, HV, K, device='cuda'))
>>> beta = torch.rand(B, T, HV, device='cuda').sigmoid()
>>> h0 = torch.randn(B, HV, K, V, device='cuda')
>>> o, ht = fused_recurrent_kda(
q, k, v, g, beta,
initial_state=h0,
output_final_state=True
)
# for variable-length inputs, the batch size `B` is expected to be 1 and `cu_seqlens` is required
>>> q, k, v, g, beta = map(lambda x: rearrange(x, 'b t ... -> 1 (b t) ...'), (q, k, v, g, beta))
# for a batch with 4 sequences, `cu_seqlens` with 5 start/end positions are expected
>>> cu_seqlens = q.new_tensor([0, 2048, 4096, 6144, 8192], dtype=torch.long)
>>> o_var, ht_var = fused_recurrent_kda(
q, k, v, g, beta,
initial_state=h0,
output_final_state=True,
cu_seqlens=cu_seqlens
)
"""
if 'transpose_state_layout' in kwargs:
if state_v_first:
raise ValueError("Cannot pass both `state_v_first` and the deprecated `transpose_state_layout`.")
warnings.warn(
"`transpose_state_layout` is deprecated and renamed to `state_v_first`.",
DeprecationWarning,
stacklevel=2,
)
state_v_first = kwargs.pop('transpose_state_layout')
if cu_seqlens is not None:
if q.shape[0] != 1:
raise ValueError(
f"The batch size is expected to be 1 rather than {q.shape[0]} when using `cu_seqlens`."
f"Please flatten variable-length inputs before processing.",
)
if initial_state is not None and initial_state.shape[0] != len(cu_seqlens) - 1:
raise ValueError(
f"The number of initial states is expected to be equal to the number of input sequences, "
f"i.e., {len(cu_seqlens) - 1} rather than {initial_state.shape[0]}.",
)
if scale is None:
scale = k.shape[-1] ** -0.5
if allow_neg_eigval and not use_beta_sigmoid_in_kernel:
raise ValueError("`allow_neg_eigval=True` requires `use_beta_sigmoid_in_kernel=True`.")
o, final_state = fused_recurrent_kda_fwd(
q=q,
k=k,
v=v,
g=g,
beta=beta,
A_log=A_log,
dt_bias=dt_bias,
scale=scale,
initial_state=initial_state,
inplace_final_state=False,
output_final_state=output_final_state,
use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel,
use_gate_in_kernel=use_gate_in_kernel,
use_beta_sigmoid_in_kernel=use_beta_sigmoid_in_kernel,
allow_neg_eigval=allow_neg_eigval,
lower_bound=lower_bound,
cu_seqlens=cu_seqlens,
state_v_first=state_v_first,
)
return o, final_state
+514
View File
@@ -0,0 +1,514 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
# This file is modified and supported by the Moonshot AI Team
import torch
import torch.nn.functional as F
import triton
import triton.language as tl
from kda._fla.ops.backends import dispatch
from kda._fla.ops.utils.cache import fla_cache_autotune
from kda._fla.ops.utils.index import prepare_chunk_indices
from kda._fla.ops.utils.op import exp
from kda._fla.ops.utils.softplus import softplus
from kda._fla.utils import IS_AMD, autocast_custom_bwd, autocast_custom_fwd, autotune_cache_kwargs, check_shared_mem, input_guard
BS_LIST = [32, 64] if check_shared_mem() else [16, 32]
BT_LIST_AUTOTUNE = [32, 64, 128]
NUM_WARPS_AUTOTUNE = [2, 4, 8, 16] if IS_AMD else [4, 8, 16, 32]
def naive_kda_gate(
g: torch.Tensor,
A_log: torch.Tensor,
dt_bias: torch.Tensor | None = None,
output_dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
"""
Torch reference implementation for KDA gate computation.
Computes: g = -A_log.exp().unsqueeze(-1) * softplus(g + dt_bias.view(g.shape[-2:]))
Args:
g (torch.Tensor):
Input tensor of shape `[..., H, K]`.
A_log (torch.Tensor):
Parameter tensor with `H` elements.
dt_bias (torch.Tensor | None):
Optional bias tensor added to `g` before activation, shape `[H * K]`.
Returns:
Output tensor of shape `[..., H, K]` .
"""
H, _ = g.shape[-2:]
g = g.float()
if dt_bias is not None:
g = g + dt_bias.view(H, -1)
g = (-A_log.view(H, 1).float().exp() * F.softplus(g.float())).to(output_dtype)
return g
def naive_kda_lowerbound_gate(
g: torch.Tensor,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
lower_bound: float = -5.0,
output_dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
"""
Torch reference implementation for KDA lowerbound gate computation.
Computes: ``g = lower_bound * sigmoid(exp(A_log) * (g + dt_bias))``.
When ``A_log`` is ``None``: ``g = lower_bound * sigmoid(g + dt_bias)``.
Args:
g (torch.Tensor):
Input tensor of shape `[..., H, K]`.
A_log (torch.Tensor | None):
Optional parameter tensor with `H` elements.
dt_bias (torch.Tensor | None):
Optional bias tensor added to `g` before activation, shape `[H * K]`.
lower_bound (float):
Lower bound for the gate output. Default: `-5.0`.
output_dtype (torch.dtype):
The dtype of the output tensor. Default: `torch.float32`.
Returns:
Output tensor of shape `[..., H, K]`.
"""
H, _ = g.shape[-2:]
g = g.float()
if dt_bias is not None:
g = g + dt_bias.view(H, -1)
if A_log is not None:
g = A_log.view(H, 1).float().exp() * g
g = lower_bound * F.sigmoid(g)
return g.to(output_dtype)
@triton.heuristics({
"HAS_A": lambda args: args["A_log"] is not None,
"HAS_BIAS": lambda args: args["dt_bias"] is not None,
"HAS_BETA": lambda args: args["beta"] is not None,
'USE_LOWER_BOUND': lambda args: args['lower_bound'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({"BT": BT}, num_warps=num_warps, num_stages=num_stages)
for BT in BT_LIST_AUTOTUNE
for num_warps in NUM_WARPS_AUTOTUNE
for num_stages in [2, 3]
],
key=["H", "D"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def kda_gate_fwd_kernel(
g,
A_log,
dt_bias,
beta,
yg,
yb,
lower_bound,
T,
H: tl.constexpr,
D: tl.constexpr,
BT: tl.constexpr,
BD: tl.constexpr,
HAS_A: tl.constexpr,
HAS_BIAS: tl.constexpr,
HAS_BETA: tl.constexpr,
USE_LOWER_BOUND: tl.constexpr,
):
i_t, i_h = tl.program_id(0).to(tl.int64), tl.program_id(1)
b_A = tl.load(A_log + i_h).to(tl.float32) if HAS_A else 1.0
o_t = i_t * BT + tl.arange(0, BT)
o_d = tl.arange(0, BD)
m_t = o_t < T
m_g = m_t[:, None] & (o_d[None, :] < D)
p_g = g + i_h * D + o_t[:, None] * (H * D) + o_d[None, :]
p_yg = yg + i_h * D + o_t[:, None] * (H * D) + o_d[None, :]
# [BT, BD]
b_g = tl.load(p_g, mask=m_g, other=0.0).to(tl.float32)
if HAS_BIAS:
o_b = i_h * D + tl.arange(0, BD)
b_g = b_g + tl.load(dt_bias + o_b, mask=o_b < H * D, other=0.0).to(tl.float32)
if not USE_LOWER_BOUND:
b_yg = -exp(b_A) * softplus(b_g)
else:
b_yg = lower_bound * tl.sigmoid((exp(b_A) if HAS_A else b_A) * b_g)
tl.store(p_yg, b_yg.to(p_yg.dtype.element_ty), mask=m_g)
if HAS_BETA:
p_b = beta + i_h + o_t * H
p_yb = yb + i_h + o_t * H
b_yb = tl.sigmoid(tl.load(p_b, mask=m_t, other=0.0).to(tl.float32))
tl.store(p_yb, b_yb.to(p_yb.dtype.element_ty), mask=m_t)
@triton.heuristics({
"HAS_A": lambda args: args["A_log"] is not None,
"HAS_BIAS": lambda args: args["dt_bias"] is not None,
"HAS_BETA": lambda args: args["beta"] is not None,
'USE_LOWER_BOUND': lambda args: args['lower_bound'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({}, num_warps=num_warps, num_stages=num_stages)
for num_warps in NUM_WARPS_AUTOTUNE
for num_stages in [2, 3]
],
key=["H", "D"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def kda_gate_bwd_kernel(
g,
A_log,
dt_bias,
beta,
dyg,
dyb,
dg,
dA,
dbeta,
lower_bound,
T,
H: tl.constexpr,
D: tl.constexpr,
BT: tl.constexpr,
BD: tl.constexpr,
HAS_A: tl.constexpr,
HAS_BIAS: tl.constexpr,
HAS_BETA: tl.constexpr,
USE_LOWER_BOUND: tl.constexpr,
):
i_t, i_h = tl.program_id(0).to(tl.int64), tl.program_id(1)
b_A = tl.load(A_log + i_h).to(tl.float32) if HAS_A else 1.0
o_t = i_t * BT + tl.arange(0, BT)
o_d = tl.arange(0, BD)
m_t = o_t < T
m_g = m_t[:, None] & (o_d[None, :] < D)
p_g = g + i_h * D + o_t[:, None] * (H * D) + o_d[None, :]
p_dg = dg + i_h * D + o_t[:, None] * (H * D) + o_d[None, :]
p_dyg = dyg + i_h * D + o_t[:, None] * (H * D) + o_d[None, :]
# [BT, BD]
b_g = tl.load(p_g, mask=m_g, other=0.0).to(tl.float32)
b_dyg = tl.load(p_dyg, mask=m_g, other=0.0).to(tl.float32)
if HAS_BIAS:
o_b = i_h * D + tl.arange(0, BD)
b_g = b_g + tl.load(dt_bias + o_b, mask=o_b < H * D, other=0.0).to(tl.float32)
# [BT, BD]
if not USE_LOWER_BOUND:
b_A = -exp(b_A)
b_yg = b_A * softplus(b_g)
b_dg = b_A * (b_dyg * tl.sigmoid(b_g))
b_dA = tl.sum(tl.sum(b_dyg * b_yg, 1), 0)
else:
b_A = exp(b_A) if HAS_A else b_A
b_inner = b_A * b_g
b_sig = tl.sigmoid(b_inner)
b_dsig = b_sig * (1.0 - b_sig)
# Common term: dy * (LB * dsig)
b_d_inner_term = b_dyg * (lower_bound * b_dsig)
# dg = d_inner_term * A
b_dg = b_d_inner_term * b_A
b_dA = tl.sum(tl.sum(b_dg * b_g, 1), 0) if HAS_A else 0.0
tl.store(p_dg, b_dg.to(p_dg.dtype.element_ty), mask=m_g)
if HAS_A:
tl.store(dA + i_t * H + i_h, b_dA)
if HAS_BETA:
p_b = beta + i_h + o_t * H
p_db = dbeta + i_h + o_t * H
p_dyb = dyb + i_h + o_t * H
b_b = tl.load(p_b, mask=m_t, other=0.0).to(tl.float32)
b_db = tl.load(p_dyb, mask=m_t, other=0.0).to(tl.float32) * b_b * (1.0 - b_b)
tl.store(p_db, b_db.to(p_db.dtype.element_ty), mask=m_t)
@dispatch('kda')
def kda_gate_fwd(
g: torch.Tensor,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
lower_bound: float | None = None,
output_dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
H, K = g.shape[-2:]
T = g.numel() // (H * K)
yg = torch.empty_like(g, dtype=output_dtype)
def grid(meta):
return (triton.cdiv(T, meta["BT"]), H)
kda_gate_fwd_kernel[grid](
g=g,
A_log=A_log,
dt_bias=dt_bias,
beta=None,
yg=yg,
yb=None,
T=T,
H=H,
D=K,
BD=triton.next_power_of_2(K),
lower_bound=lower_bound,
)
return yg
@dispatch('kda')
def kda_gate_bwd(
g: torch.Tensor,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
dyg: torch.Tensor | None = None,
lower_bound: float | None = None,
) -> tuple[torch.Tensor, torch.Tensor | None, torch.Tensor | None]:
H, K = g.shape[-2:]
T = g.numel() // (H * K)
BT = 32
NT = triton.cdiv(T, BT)
dg = torch.empty_like(g, dtype=torch.float32)
dA = g.new_empty(NT, H, dtype=torch.float32) if A_log is not None else None
grid = (triton.cdiv(T, BT), H)
kda_gate_bwd_kernel[grid](
g=g,
A_log=A_log,
dt_bias=dt_bias,
beta=None,
dyg=dyg,
dyb=None,
dg=dg,
dA=dA,
dbeta=None,
T=T,
H=H,
D=K,
BT=BT,
BD=triton.next_power_of_2(K),
lower_bound=lower_bound,
)
dg = dg.view_as(g).type_as(g)
dA = dA.sum(0).view_as(A_log).type_as(A_log) if A_log is not None else None
# dt_bias is [HV, K] in KDAAttention and [HV*K] in some FLA call sites.
dbias = (
dg.view(-1, H * K).sum(0).reshape_as(dt_bias).type_as(dt_bias)
if dt_bias is not None
else None
)
return dg, dA, dbias
class KDAGateFunction(torch.autograd.Function):
@staticmethod
@input_guard
@autocast_custom_fwd
def forward(
ctx,
g: torch.Tensor,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
lower_bound: float | None = None,
output_dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
yg = kda_gate_fwd(
g=g,
A_log=A_log,
dt_bias=dt_bias,
lower_bound=lower_bound,
output_dtype=output_dtype
)
ctx.save_for_backward(g, A_log, dt_bias)
ctx.lower_bound = lower_bound
return yg
@staticmethod
@input_guard
@autocast_custom_bwd
def backward(ctx, dyg: torch.Tensor):
g, A_log, dt_bias = ctx.saved_tensors
dg, dA, dbias = kda_gate_bwd(
g=g,
A_log=A_log,
dt_bias=dt_bias,
dyg=dyg,
lower_bound=ctx.lower_bound
)
return dg, dA, dbias, None, None
@dispatch('kda')
@torch.compiler.disable
def fused_kda_gate(
g: torch.Tensor,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
lower_bound: float | None = None,
output_dtype: torch.dtype = torch.float32,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
"""
Fused KDA gate computation with autograd support.
Computes: g = -A_log.exp().unsqueeze(-1) * softplus(g + dt_bias.view(g.shape[-2:]))
When ``lower_bound`` is set: g = lower_bound * sigmoid(exp(A_log) * (g + dt_bias)).
When ``A_log`` is ``None`` (requires ``lower_bound``): g = lower_bound * sigmoid(g + dt_bias).
Args:
g (torch.Tensor):
Input tensor of shape `[..., H, K]`.
A_log (torch.Tensor | None):
Optional parameter tensor with `H` elements.
When ``None``, the gate reduces to ``lower_bound * sigmoid(g + dt_bias)`` (requires ``lower_bound``).
dt_bias (torch.Tensor | None):
Optional bias tensor added to `g` before activation, shape `[H * K]`.
Returns:
Output tensor of shape `[..., H, K]`.
"""
return KDAGateFunction.apply(g, A_log, dt_bias, lower_bound, output_dtype)
@triton.heuristics({
"HAS_A": lambda args: args["A_log"] is not None,
"HAS_BIAS": lambda args: args["dt_bias"] is not None,
'HAS_SCALE': lambda args: args['scale'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
'USE_LOWER_BOUND': lambda args: args['lower_bound'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({'BS': BS}, num_warps=num_warps)
for BS in BS_LIST
for num_warps in [2, 4, 8]
],
key=['H', 'S', 'BT', 'IS_VARLEN', 'REVERSE'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def kda_gate_chunk_cumsum_vector_kernel(
s,
A_log,
dt_bias,
o,
scale,
cu_seqlens,
chunk_indices,
lower_bound,
T,
H: tl.constexpr,
S: tl.constexpr,
BT: tl.constexpr,
BS: tl.constexpr,
REVERSE: tl.constexpr,
HAS_A: tl.constexpr,
HAS_BIAS: tl.constexpr,
HAS_SCALE: tl.constexpr,
IS_VARLEN: tl.constexpr,
USE_LOWER_BOUND: tl.constexpr,
):
i_s, i_t, i_bh = tl.program_id(0), tl.program_id(1).to(tl.int64), tl.program_id(2).to(tl.int64)
i_b, i_h = i_bh // H, i_bh % H
if IS_VARLEN:
i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64)
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
else:
bos, eos = i_b * T, i_b * T + T
o_t = i_t * BT + tl.arange(0, BT)
o_s = i_s * BS + tl.arange(0, BS)
m_s = (o_t[:, None] < T) & (o_s[None, :] < S)
p_s = s + (bos * H + i_h) * S + o_t[:, None] * (H*S) + o_s[None, :]
p_o = o + (bos * H + i_h) * S + o_t[:, None] * (H*S) + o_s[None, :]
# [BT, BS]
b_s = tl.load(p_s, mask=m_s, other=0.0).to(tl.float32)
# Apply dt_bias if exists
if HAS_BIAS:
b_bias = tl.load(dt_bias + i_h * S + o_s, mask=o_s < S, other=0.0).to(tl.float32)
b_s = b_s + b_bias[None, :]
b_A = tl.load(A_log + i_h).to(tl.float32) if HAS_A else 1.0
if not USE_LOWER_BOUND:
# Apply gate: -exp(A_log) * softplus(g + bias)
b_gate = -exp(b_A) * softplus(b_s)
else:
b_gate = lower_bound * tl.sigmoid((exp(b_A) if HAS_A else b_A) * b_s)
# Apply chunk local cumsum
if REVERSE:
b_o = tl.cumsum(b_gate, axis=0, reverse=True)
else:
b_o = tl.cumsum(b_gate, axis=0)
if HAS_SCALE:
b_o *= scale
tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=m_s)
@input_guard
@dispatch('kda')
def kda_gate_chunk_cumsum(
g: torch.Tensor,
A_log: torch.Tensor | None,
chunk_size: int,
scale: float = None,
dt_bias: torch.Tensor | None = None,
cu_seqlens: torch.Tensor | None = None,
output_dtype: torch.dtype | None = torch.float,
chunk_indices: torch.LongTensor | None = None,
lower_bound: float | None = None,
**kwargs,
) -> torch.Tensor:
if cu_seqlens is not None:
assert g.shape[0] == 1, "Only batch size 1 is supported when cu_seqlens are provided"
assert len(g.shape) == 4
B, T, H, S = g.shape
BT = chunk_size
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
assert chunk_size == 2**(chunk_size.bit_length()-1), "chunk_size must be a power of 2"
g_org, g = g, torch.empty_like(g, dtype=output_dtype or g.dtype)
def grid(meta): return (triton.cdiv(meta['S'], meta['BS']), NT, B * H)
kda_gate_chunk_cumsum_vector_kernel[grid](
s=g_org,
A_log=A_log,
dt_bias=dt_bias,
o=g,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
lower_bound=lower_bound,
T=T,
H=H,
S=S,
BT=BT,
REVERSE=False,
)
return g
+369
View File
@@ -0,0 +1,369 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import torch
import triton
import triton.language as tl
from kda._fla.ops.backends import dispatch
from kda._fla.ops.utils import prepare_chunk_indices
from kda._fla.ops.utils.cache import fla_cache_autotune
from kda._fla.ops.utils.op import exp2
from kda._fla.utils import autotune_cache_kwargs, check_shared_mem
@triton.heuristics({
'STORE_QG': lambda args: args['qg'] is not None,
'STORE_KG': lambda args: args['kg'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({}, num_warps=num_warps, num_stages=num_stages)
for num_warps in [2, 4, 8]
for num_stages in [2, 3, 4]
],
key=['H', 'HV', 'K', 'V', 'BT', 'BK', 'BV', 'IS_VARLEN'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def recompute_w_u_fwd_kda_kernel(
q,
k,
qg,
kg,
v,
beta,
w,
u,
A,
gk,
cu_seqlens,
chunk_indices,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BT: tl.constexpr,
BK: tl.constexpr,
BV: tl.constexpr,
STORE_QG: tl.constexpr,
STORE_KG: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_t, i_bh = tl.program_id(0).to(tl.int64), tl.program_id(1).to(tl.int64)
i_b, i_hv = i_bh // HV, i_bh % HV
i_h = i_hv // (HV // H)
if IS_VARLEN:
i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64)
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
else:
bos, eos = i_b * T, i_b * T + T
k += (bos * H + i_h) * K
v += (bos * HV + i_hv) * V
u += (bos * HV + i_hv) * V
w += (bos * HV + i_hv) * K
gk += (bos * HV + i_hv) * K
beta += bos * HV + i_hv
A += (bos * HV + i_hv) * BT
if STORE_QG:
q += (bos * H + i_h) * K
qg += (bos * HV + i_hv) * K
if STORE_KG:
kg += (bos * HV + i_hv) * K
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
p_b = beta + o_t * HV
b_b = tl.load(p_b, mask=m_t, other=0.0)
o_A = tl.arange(0, BT)
m_A = m_t[:, None] & (o_A[None, :] < BT)
p_A = A + o_t[:, None] * (HV*BT) + o_A[None, :]
b_A = tl.load(p_A, mask=m_A, other=0.0)
for i_v in range(tl.cdiv(V, BV)):
o_v = i_v * BV + tl.arange(0, BV)
m_v = m_t[:, None] & (o_v[None, :] < V)
p_v = v + o_t[:, None] * (HV*V) + o_v[None, :]
p_u = u + o_t[:, None] * (HV*V) + o_v[None, :]
b_v = tl.load(p_v, mask=m_v, other=0.0)
b_vb = (b_v * b_b[:, None]).to(b_v.dtype)
b_u = tl.dot(b_A, b_vb)
tl.store(p_u, b_u.to(p_u.dtype.element_ty), mask=m_v)
for i_k in range(tl.cdiv(K, BK)):
o_k = i_k * BK + tl.arange(0, BK)
m_k = o_k < K
m_tk = m_t[:, None] & m_k[None, :]
p_w = w + o_t[:, None] * (HV*K) + o_k[None, :]
p_k = k + o_t[:, None] * (H*K) + o_k[None, :]
b_k = tl.load(p_k, mask=m_tk, other=0.0)
b_kb = b_k * b_b[:, None]
p_gk = gk + o_t[:, None] * (HV*K) + o_k[None, :]
b_gk = tl.load(p_gk, mask=m_tk, other=0.0).to(tl.float32)
b_kb *= exp2(b_gk)
if STORE_QG:
p_q = q + o_t[:, None] * (H*K) + o_k[None, :]
p_qg = qg + o_t[:, None] * (HV*K) + o_k[None, :]
b_q = tl.load(p_q, mask=m_tk, other=0.0)
b_qg = b_q * exp2(b_gk)
tl.store(p_qg, b_qg.to(p_qg.dtype.element_ty), mask=m_tk)
if STORE_KG:
last_idx = min(i_t * BT + BT, T) - 1
b_gn = tl.load(gk + last_idx * HV*K + o_k, mask=m_k, other=0.).to(tl.float32)
b_kg = b_k * tl.where((i_t * BT + tl.arange(0, BT) < T)[:, None], exp2(b_gn[None, :] - b_gk), 0)
p_kg = kg + o_t[:, None] * (HV*K) + o_k[None, :]
tl.store(p_kg, b_kg.to(p_kg.dtype.element_ty), mask=m_tk)
b_w = tl.dot(b_A, b_kb.to(b_k.dtype))
tl.store(p_w, b_w.to(p_w.dtype.element_ty), mask=m_tk)
@triton.heuristics({
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({}, num_warps=num_warps, num_stages=num_stages)
for num_warps in [2, 4]
for num_stages in [2, 3, 4]
],
key=['H', 'HV', 'K', 'V', 'BT', 'BK', 'BV', 'IS_VARLEN'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def prepare_wy_repr_bwd_kda_kernel(
k,
v,
beta,
gk,
A,
dA,
dw,
du,
dk,
dk2,
dv,
db,
dg,
dg2,
cu_seqlens,
chunk_indices,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BT: tl.constexpr,
BK: tl.constexpr,
BV: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_t, i_bh = tl.program_id(0).to(tl.int64), tl.program_id(1).to(tl.int64)
i_b, i_hv = i_bh // HV, i_bh % HV
i_h = i_hv // (HV // H)
if IS_VARLEN:
i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64)
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
else:
bos, eos = i_b * T, i_b * T + T
k += (bos * H + i_h) * K
v += (bos * HV + i_hv) * V
beta += bos * HV + i_hv
gk += (bos * HV + i_hv) * K
A += (bos * HV + i_hv) * BT
dA += (bos * HV + i_hv) * BT
dk += (bos * HV + i_hv) * K
dk2 += (bos * HV + i_hv) * K
dw += (bos * HV + i_hv) * K
du += (bos * HV + i_hv) * V
dv += (bos * HV + i_hv) * V
db += bos * HV + i_hv
dg += (bos * HV + i_hv) * K
dg2 += (bos * HV + i_hv) * K
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
p_b = beta + o_t * HV
p_db = db + o_t * HV
o_A = tl.arange(0, BT)
m_AT = (o_A[:, None] < BT) & m_t[None, :]
p_A = A + o_A[:, None] + o_t[None, :] * (HV*BT)
b_b = tl.load(p_b, mask=m_t, other=0.0)
b_db = tl.zeros([BT], dtype=tl.float32)
b_A = tl.load(p_A, mask=m_AT, other=0.0)
b_dA = tl.zeros([BT, BT], dtype=tl.float32)
for i_k in range(tl.cdiv(K, BK)):
o_k = i_k * BK + tl.arange(0, BK)
m_k = m_t[:, None] & (o_k[None, :] < K)
p_k = k + o_t[:, None] * (H*K) + o_k[None, :]
p_dk = dk + o_t[:, None] * (HV*K) + o_k[None, :]
p_dk2 = dk2 + o_t[:, None] * (HV*K) + o_k[None, :]
p_dw = dw + o_t[:, None] * (HV*K) + o_k[None, :]
p_dg = dg + o_t[:, None] * (HV*K) + o_k[None, :]
p_dg2 = dg2 + o_t[:, None] * (HV*K) + o_k[None, :]
# [BT, BK]
b_k = tl.load(p_k, mask=m_k, other=0.0)
p_gk = gk + o_t[:, None] * (HV*K) + o_k[None, :]
b_gk_exp = exp2(tl.load(p_gk, mask=m_k, other=0.0))
b_kbg = b_k * b_b[:, None] * b_gk_exp
b_dw = tl.load(p_dw, mask=m_k, other=0.0)
b_dA += tl.dot(b_dw, tl.trans(b_kbg).to(b_dw.dtype))
b_dkbg = tl.dot(b_A, b_dw)
b_dk = b_dkbg * b_gk_exp * b_b[:, None] + tl.load(p_dk, mask=m_k, other=0.0)
b_db += tl.sum(b_dkbg * b_k * b_gk_exp, 1)
b_dg = b_kbg * b_dkbg + tl.load(p_dg, mask=m_k, other=0.0)
tl.store(p_dk2, b_dk.to(p_dk2.dtype.element_ty), mask=m_k)
tl.store(p_dg2, b_dg.to(p_dg2.dtype.element_ty), mask=m_k)
for i_v in range(tl.cdiv(V, BV)):
o_v = i_v * BV + tl.arange(0, BV)
m_v = m_t[:, None] & (o_v[None, :] < V)
p_v = v + o_t[:, None] * (HV*V) + o_v[None, :]
p_dv = dv + o_t[:, None] * (HV*V) + o_v[None, :]
p_du = du + o_t[:, None] * (HV*V) + o_v[None, :]
b_v = tl.load(p_v, mask=m_v, other=0.0)
b_vb = (b_v * b_b[:, None]).to(b_v.dtype)
b_du = tl.load(p_du, mask=m_v, other=0.0)
b_dA += tl.dot(b_du, tl.trans(b_vb))
b_dvb = tl.dot(b_A, b_du)
b_dv = b_dvb * b_b[:, None]
b_db += tl.sum(b_dvb * b_v, 1)
tl.store(p_dv, b_dv.to(p_dv.dtype.element_ty), mask=m_v)
m_A = (o_t[:, None] > o_t[None, :]) & (m_t[:, None] & m_t)
b_dA = tl.where(m_A, b_dA, 0)
b_dA = tl.dot(b_dA.to(b_A.dtype), b_A)
b_dA = tl.dot(b_A, b_dA.to(b_A.dtype))
b_dA = tl.where(m_A, -b_dA, 0)
m_dA = m_t[:, None] & (o_A[None, :] < BT)
p_dA = dA + o_t[:, None] * (HV*BT) + o_A[None, :]
tl.store(p_dA, b_dA.to(p_dA.dtype.element_ty), mask=m_dA)
tl.store(p_db, b_db.to(p_db.dtype.element_ty), mask=m_t)
@dispatch('kda')
def recompute_w_u_fwd(
k: torch.Tensor,
v: torch.Tensor,
beta: torch.Tensor,
A: torch.Tensor,
gk: torch.Tensor,
q: torch.Tensor | None = None,
cu_seqlens: torch.LongTensor | None = None,
chunk_indices: torch.LongTensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None, torch.Tensor | None]:
B, T, H, K, V = *k.shape, v.shape[-1]
HV = v.shape[2]
BT = A.shape[-1]
BK = 64
BV = 64
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
w = torch.empty(B, T, HV, K, device=k.device, dtype=k.dtype)
u = torch.empty_like(v)
qg = torch.empty(B, T, HV, K, device=k.device, dtype=k.dtype) if q is not None else None
kg = torch.empty(B, T, HV, K, device=k.device, dtype=k.dtype)
recompute_w_u_fwd_kda_kernel[(NT, B*HV)](
q=q,
k=k,
qg=qg,
kg=kg,
v=v,
beta=beta,
w=w,
u=u,
A=A,
gk=gk,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
T=T,
H=H,
HV=HV,
K=K,
V=V,
BT=BT,
BK=BK,
BV=BV,
)
return w, u, qg, kg
def prepare_wy_repr_bwd(
k: torch.Tensor,
v: torch.Tensor,
beta: torch.Tensor,
gk: torch.Tensor,
A: torch.Tensor,
dk: torch.Tensor,
dw: torch.Tensor,
du: torch.Tensor,
dg: torch.Tensor,
cu_seqlens: torch.LongTensor | None = None,
chunk_indices: torch.LongTensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
B, T, H, K, V = *k.shape, v.shape[-1]
HV = v.shape[2]
BT = A.shape[-1]
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
CONST_TILING = 64 if check_shared_mem() else 32
BK = min(max(triton.next_power_of_2(K), 16), CONST_TILING)
BV = min(max(triton.next_power_of_2(V), 16), CONST_TILING)
dk2 = torch.empty_like(dk, dtype=torch.float)
dv = torch.empty_like(v)
dg2 = torch.empty_like(gk, dtype=torch.float)
dA = torch.empty_like(A, dtype=torch.float)
db = torch.empty_like(beta, dtype=torch.float)
prepare_wy_repr_bwd_kda_kernel[(NT, B * HV)](
k=k,
v=v,
beta=beta,
gk=gk,
A=A,
dA=dA,
dw=dw,
du=du,
dk=dk,
dk2=dk2,
dv=dv,
db=db,
dg=dg,
dg2=dg2,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
T=T,
H=H,
HV=HV,
K=K,
V=V,
BT=BT,
BK=BK,
BV=BV,
)
dk = dk2
dg = dg2
return dk, dv, db, dg, dA
+14
View File
@@ -0,0 +1,14 @@
from .cumsum import (
chunk_local_cumsum,
chunk_local_cumsum_scalar,
chunk_local_cumsum_vector,
)
from .index import prepare_chunk_indices, prepare_chunk_offsets
__all__ = [
"chunk_local_cumsum",
"chunk_local_cumsum_scalar",
"chunk_local_cumsum_vector",
"prepare_chunk_indices",
"prepare_chunk_offsets",
]
+449
View File
@@ -0,0 +1,449 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import dataclasses
import enum
import json
import logging
import os
import re
from functools import cache, lru_cache
from pathlib import Path
from typing import Any
import torch
import triton
from packaging import version
from triton.runtime.autotuner import Autotuner
TRITON_ABOVE_3_5_1 = version.parse(triton.__version__) >= version.parse("3.5.1")
TRITON_ABOVE_3_4_0 = version.parse(triton.__version__) >= version.parse("3.4.0")
class FlaCacheMode(enum.Enum):
"""Controls how FLA loads kernel configs from its config cache (FLA_CACHE_MODE env var).
DISABLED — skip all cache lookups, always fall back to Triton autotune (default when FLA_CACHE_MODE is unset)
STRICT — exact key match only; falls back to Triton autotune if no match
FUZZY — exact key match → fuzzy key match; falls back to Triton autotune if no match
FULL — exact key match → fuzzy key match → default_config fallback
DEFAULT — use only the top-level default_config field, skip key-based lookup
ALWAYS — like DEFAULT, but re-reads config files on every kernel call;
useful for debugging: edit default_config in a JSON file and the next
kernel call picks it up without restarting the process
"""
DISABLED = "disabled"
STRICT = "strict"
FUZZY = "fuzzy"
FULL = "full"
DEFAULT = "default"
ALWAYS = "always"
def uses_default_config(self) -> bool:
"""Return True for modes that may fall back to default_config (FULL, DEFAULT, ALWAYS)."""
return self in (FlaCacheMode.FULL, FlaCacheMode.DEFAULT, FlaCacheMode.ALWAYS)
@classmethod
def from_env(cls) -> "FlaCacheMode":
mode_str = os.environ.get("FLA_CACHE_MODE", cls.DISABLED.value)
try:
return cls(mode_str)
except ValueError:
valid = [m.value for m in cls]
raise ValueError(
f"Invalid FLA_CACHE_MODE={mode_str!r}. Valid values: {valid}"
) from None
FLA_CACHE_MODE: FlaCacheMode = FlaCacheMode.from_env()
logger = logging.getLogger(__name__)
def sanitize_gpu_name(gpu_name: str) -> str:
sanitized = re.sub(r"[^0-9A-Za-z]+", "_", gpu_name)
sanitized = sanitized.strip("_")
return sanitized or "unknown_gpu"
@lru_cache(maxsize=1)
def get_gpu_info():
"""Get GPU model information.
This function detects the GPU model and returns a sanitized string identifier.
It prioritizes FLA_GPU_NAME environment variable if set, then detects from
available hardware (CUDA, ROCm, Intel GPU, or CPU).
"""
# Check if GPU name is overridden via environment variable
gpu_name = None
# Check if GPU name is overridden via environment variable
if "FLA_GPU_NAME" in os.environ:
gpu_name = os.environ["FLA_GPU_NAME"]
# Try to get device name based on availability
elif torch.cuda.is_available():
# Works for both NVIDIA and AMD GPUs (ROCm)
gpu_name = torch.cuda.get_device_name(0)
elif hasattr(torch, 'xpu') and torch.xpu.is_available():
gpu_name = torch.xpu.get_device_name(0)
if gpu_name:
return sanitize_gpu_name(gpu_name)
# Default to CPU if no GPU available
return "cpu"
def get_fla_config_dir() -> Path:
"""Get FLA's configs directory.
The directory can be overridden by setting the FLA_CONFIG_DIR environment variable.
If set, configs will be loaded directly from $FLA_CONFIG_DIR/. Otherwise FLA
falls back to the default fla/configs/{GPU}/ directory in the project.
"""
# Check if custom config dir is set via environment variable
if "FLA_CONFIG_DIR" in os.environ:
return Path(os.environ["FLA_CONFIG_DIR"])
# Default: project_dir/fla/configs/{GPU}/
project_dir = Path(__file__).parent.parent.parent
return project_dir / "configs" / get_gpu_info()
@dataclasses.dataclass(frozen=True)
class AutotuneKey:
"""Autotune key with exact/fuzzy matching, serialization, and construction helpers."""
autotune_key: tuple[Any, ...]
@staticmethod
def normalize_autotune_key(value: Any) -> Any:
if isinstance(value, (list, tuple)):
return [AutotuneKey.normalize_autotune_key(v) for v in value]
if isinstance(value, dict):
return {k: AutotuneKey.normalize_autotune_key(v) for k, v in value.items()}
return value
@staticmethod
def serialize(key: Any) -> str:
return json.dumps(AutotuneKey.normalize_autotune_key(key), separators=(",", ":"), sort_keys=True)
@staticmethod
def key_hash(key: Any) -> str:
import hashlib
return hashlib.md5(AutotuneKey.serialize(key).encode()).hexdigest()
@staticmethod
def is_numeric(value: Any) -> bool:
return isinstance(value, (int, float)) and not isinstance(value, bool)
@staticmethod
def keys_fuzzy_match(cached_key: Any, requested_key: Any) -> bool:
# Fuzzy match: numeric leaves are compatible regardless of their actual numeric values
# (e.g. a config tuned for seq_len=1024 can apply to seq_len=2048).
# Structure (type, length, dict keys) must still match exactly.
if AutotuneKey.is_numeric(cached_key) and AutotuneKey.is_numeric(requested_key):
return True
if isinstance(cached_key, (list, tuple)) and isinstance(requested_key, (list, tuple)):
return len(cached_key) == len(requested_key) and all(
AutotuneKey.keys_fuzzy_match(c, r) for c, r in zip(cached_key, requested_key)
)
if isinstance(cached_key, dict) and isinstance(requested_key, dict):
return cached_key.keys() == requested_key.keys() and all(
AutotuneKey.keys_fuzzy_match(cached_key[k], requested_key[k]) for k in cached_key
)
return cached_key == requested_key
@classmethod
def build(
cls,
arg_names: list[str],
key_names: list[str],
positional_args: tuple[Any, ...],
runtime_kwargs: dict[str, Any],
) -> "AutotuneKey":
named_args = dict(zip(arg_names, positional_args))
all_args = {**named_args, **runtime_kwargs}
tracked_args = {k: v for (k, v) in all_args.items() if k in arg_names}
tuning_key = [tracked_args[name] for name in key_names if name in tracked_args]
for arg in tracked_args.values():
if hasattr(arg, "dtype"):
tuning_key.append(str(arg.dtype))
return cls(autotune_key=tuple(tuning_key))
def exact_matches(self, entry_key: Any) -> bool:
return self.serialize(self.autotune_key) == self.serialize(entry_key)
def fuzzy_matches(self, entry_key: Any) -> bool:
self_normalized = self.normalize_autotune_key(self.autotune_key)
entry_normalized = self.normalize_autotune_key(entry_key)
return (
isinstance(self_normalized, list)
and isinstance(entry_normalized, list)
and len(self_normalized) == len(entry_normalized)
and AutotuneKey.keys_fuzzy_match(self_normalized, entry_normalized)
)
@dataclasses.dataclass(frozen=True)
class KernelConfigFile:
"""Validated in-memory representation of a {kernel_name}.json config file."""
kernel_name: str | None
triton_version: str | None
autotune_entries: dict[str, dict[str, Any]] | None
default_config: dict[str, Any] | None
@classmethod
def from_dict(cls, config_file: Path, data: Any) -> "KernelConfigFile | None":
"""Parse and validate a raw JSON dict. Returns None (with a warning) if malformed."""
def fail(msg, *args):
logger.warning(msg, *args)
raise ValueError
try:
if not isinstance(data, dict):
fail("Malformed config %s: root is %s, expected dict", config_file, type(data).__name__)
raw_entries = data.get("autotune_entries")
entries: dict[str, dict[str, Any]] | None = None
if raw_entries is not None:
if not isinstance(raw_entries, dict):
fail("Malformed config %s: 'autotune_entries' is %s, expected dict",
config_file, type(raw_entries).__name__)
for h, entry in raw_entries.items():
if not isinstance(entry, dict):
fail("Malformed config %s: autotune_entries[%r] is %s, expected dict",
config_file, h, type(entry).__name__)
if not isinstance(entry.get("config"), dict):
fail("Malformed config %s: autotune_entries[%r] missing valid 'config' field", config_file, h)
entries = raw_entries
default_config = data.get("default_config")
if default_config is not None and not isinstance(default_config, dict):
fail("Malformed config %s: 'default_config' is %s, expected dict", config_file, type(default_config).__name__)
return cls(
kernel_name=data.get("kernel_name"),
triton_version=data.get("triton_version"),
autotune_entries=entries,
default_config=default_config,
)
except ValueError:
return None
@classmethod
def from_file(cls, config_file: Path) -> "KernelConfigFile | None":
"""Read and validate a config file. Returns None if the file is missing or malformed."""
config_data = read_config_file(config_file)
if config_data is None:
return None
return cls.from_dict(config_file, config_data)
def lookup_exact(self, key: AutotuneKey) -> dict[str, Any] | None:
if self.autotune_entries is None:
return None
return self.autotune_entries.get(AutotuneKey.key_hash(key.autotune_key))
def lookup_fuzzy(self, key: AutotuneKey) -> dict[str, Any] | None:
if self.autotune_entries is None:
return None
for entry in self.autotune_entries.values():
if key.fuzzy_matches(entry.get("autotune_key")):
return entry
return None
@cache
def load_config_file(config_file: Path) -> dict[str, Any] | None:
try:
with open(config_file) as f:
return json.load(f)
except Exception as e:
logger.warning("Error reading config file %s: %s", config_file, e)
return None
def read_config_file(config_file: Path) -> dict[str, Any] | None:
"""Read a config file, bypassing the in-process cache in ALWAYS mode."""
if FLA_CACHE_MODE is FlaCacheMode.ALWAYS:
return load_config_file.__wrapped__(config_file)
return load_config_file(config_file)
def load_cached_config(kernel_name: str, autotune_key: AutotuneKey | None = None) -> dict[str, Any] | None:
"""
Load cached best config for a kernel from FLA configs directory.
This function loads the cached best configuration for a given kernel name
from get_fla_config_dir()/{kernel_name}.json.
Cache files may contain multiple autotune entries keyed by Triton's
runtime tuning key plus a top-level default config.
If the config file is not found or cannot be loaded, a warning is printed
and None is returned, allowing fallback to Triton's autotune.
The lookup mode is controlled by the FLA_CACHE_MODE environment variable (see FlaCacheMode).
Args:
kernel_name: Name of the kernel (e.g., "causal_conv1d_fwd_kernel")
autotune_key: Triton autotune key for the current invocation
Returns:
Best config dictionary or None if not found or disabled
"""
if FLA_CACHE_MODE is FlaCacheMode.DISABLED:
return None
config_dir = get_fla_config_dir()
config_file = config_dir / f"{kernel_name}.json"
if not config_file.exists():
return None
config_data = read_config_file(config_file)
if config_data is None:
return None
config = KernelConfigFile.from_dict(config_file, config_data)
if config is None:
return None
if FLA_CACHE_MODE is FlaCacheMode.DEFAULT or FLA_CACHE_MODE is FlaCacheMode.ALWAYS:
return config.default_config
# STRICT mode: exact match only, no fuzzy fallback
if FLA_CACHE_MODE is FlaCacheMode.STRICT:
if autotune_key is not None:
entry = config.lookup_exact(autotune_key)
if entry is not None:
return entry["config"]
return None
# FULL and FUZZY modes: try exact key match first, then fuzzy match
if autotune_key is not None:
entry = config.lookup_exact(autotune_key) or config.lookup_fuzzy(autotune_key)
if entry is not None:
return entry["config"]
if FLA_CACHE_MODE is FlaCacheMode.FUZZY:
return None
# FULL mode: fall back to default_config, then legacy raw config (no autotune_entries)
if config.default_config is not None:
return config.default_config
if config.autotune_entries is not None:
return None
return config_data
class CachedAutotuner(Autotuner):
"""
A modified autotuner that loads best config from FLA's config directory.
This class extends Triton's Autotuner but overrides the run method to
try loading cached configuration first before falling back to autotune.
"""
def __init__(self, fn, arg_names, configs, key, reset_to_zero, restore_value, **kwargs):
super().__init__(fn, arg_names, configs, key, reset_to_zero, restore_value, **kwargs)
self.kernel_name = fn.fn.__name__ if hasattr(fn, 'fn') else fn.__name__
# None-safe pre/post hooks: Triton's defaults crash when a restore_value / reset_to_zero arg
# is None (idiomatic for optional pointers gated by a tl.constexpr flag).
# Fixed upstream in triton-lang/triton#10295 — remove this override once FLA's minimum Triton version has it.
if not self.user_defined_pre_hook and (self.reset_to_zero or self.restore_value):
def _pre_hook(kw, reset_only=False):
for n in self.reset_to_zero:
if kw[n] is not None:
kw[n].zero_()
if not reset_only:
self.restore_copies = {n: kw[n].clone() for n in self.restore_value if kw[n] is not None}
self.pre_hook = _pre_hook
if not self.user_defined_post_hook and self.restore_value:
def _post_hook(kw, exception):
for n, copy in self.restore_copies.items():
kw[n].copy_(copy)
self.restore_copies = {}
self.post_hook = _post_hook
def should_check_fla_cache(self, key: AutotuneKey) -> bool:
if FLA_CACHE_MODE is FlaCacheMode.DISABLED:
return False
if FLA_CACHE_MODE is FlaCacheMode.ALWAYS:
return True
return key.autotune_key not in self.cache
def run(self, *args, **kwargs):
key = AutotuneKey.build(self.arg_names, self.keys, args, kwargs)
if self.should_check_fla_cache(key):
self.maybe_load_cached_config(key)
return super().run(*args, **kwargs)
def maybe_load_cached_config(self, key: AutotuneKey):
best_config = load_cached_config(self.kernel_name, key)
if best_config is not None:
kw = best_config["kwargs"]
num_warps = best_config["num_warps"]
num_stages = best_config["num_stages"]
extra = {
"num_ctas": best_config["num_ctas"],
"maxnreg": best_config.get("maxnreg"),
"pre_hook": None,
"ir_override": best_config.get("ir_override"),
} if TRITON_ABOVE_3_5_1 else {}
cfg = triton.Config(kw, num_warps=num_warps, num_stages=num_stages, **extra)
self.cache[key.autotune_key] = cfg
else:
logger.debug(
"No cached config found for kernel %s and key %s; falling back to Triton autotune",
self.kernel_name,
list(key.autotune_key),
)
def fla_cache_autotune(configs, key=None, prune_configs_by=None, reset_to_zero=None, restore_value=None,
pre_hook=None, post_hook=None, warmup=None, rep=None, use_cuda_graph=False,
do_bench=None, cache_results=False):
"""
Decorator for auto-tuning a :code:`triton.jit`'d function with FLA config support.
Extends Triton's autotune to load best configurations from FLA's config directory
(default: fla/configs/{GPU}/, or FLA_CONFIG_DIR/ when overridden), keyed by kernel
name from {kernel_name}.json. Lookup behaviour is controlled by FLA_CACHE_MODE.
Falls back to normal Triton autotuning when no cached config is found.
"""
# key can be None when we want to use cache only (no fallback autotune)
if key is None:
key = []
def decorator(fn):
kwargs = {}
if TRITON_ABOVE_3_4_0:
kwargs = {"cache_results": cache_results}
return CachedAutotuner(fn, fn.arg_names, configs, key, reset_to_zero, restore_value,
pre_hook=pre_hook, post_hook=post_hook,
prune_configs_by=prune_configs_by, warmup=warmup, rep=rep,
use_cuda_graph=use_cuda_graph, do_bench=do_bench,
**kwargs,
)
return decorator
def configure_fla_cache_autotune():
triton.autotune = fla_cache_autotune
logger.info(
"configure_fla_cache_autotune() is enabling FLA fla_cache_autotune; "
"triton.autotune will be replaced with fla_cache_autotune."
)
def restore_autotune_backend():
from triton.runtime.autotuner import autotune as original_autotune
triton.autotune = original_autotune
logger.info(
"restore_autotune_backend() is restoring Triton's original autotune; "
"triton.autotune will be replaced with triton.runtime.autotuner.autotune."
)
+10
View File
@@ -0,0 +1,10 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
# Approximate value of 1/ln(2), used for log/exp base conversion
# Best FP32 approximation: 1.4426950216 (hex 0x3FB8AA3B)
RCP_LN2 = 1.4426950216
+468
View File
@@ -0,0 +1,468 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import torch
import triton
import triton.language as tl
from kda._fla.ops.backends import dispatch
from kda._fla.ops.utils.cache import fla_cache_autotune
from kda._fla.ops.utils.index import prepare_chunk_indices
from kda._fla.utils import autotune_cache_kwargs, check_shared_mem, input_guard
BS_LIST = [32, 64] if check_shared_mem() else [16, 32]
@triton.heuristics({
'HAS_SCALE': lambda args: args['scale'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({}, num_warps=num_warps)
for num_warps in [1, 2, 4, 8]
],
key=['B', 'H', 'BT', 'IS_VARLEN', 'REVERSE'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_local_cumsum_scalar_kernel(
s,
o,
scale,
cu_seqlens,
chunk_indices,
T,
B: tl.constexpr,
H: tl.constexpr,
BT: tl.constexpr,
REVERSE: tl.constexpr,
HAS_SCALE: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_t, i_bh = tl.program_id(0).to(tl.int64), tl.program_id(1).to(tl.int64)
i_b, i_h = i_bh // H, i_bh % H
if IS_VARLEN:
i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64)
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
else:
bos, eos = i_b * T, i_b * T + T
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
p_s = s + bos*H + i_h + o_t * H
p_o = o + bos*H + i_h + o_t * H
# [BT]
b_s = tl.load(p_s, mask=m_t, other=0.0).to(tl.float32)
if REVERSE:
b_o = tl.cumsum(b_s, axis=0, reverse=True)
else:
b_o = tl.cumsum(b_s, axis=0)
if HAS_SCALE:
b_o *= scale
tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=m_t)
@triton.heuristics({
'HAS_SCALE': lambda args: args['scale'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({'BS': BS}, num_warps=num_warps)
for BS in BS_LIST
for num_warps in [2, 4, 8]
],
key=['B', 'H', 'S', 'BT', 'IS_VARLEN', 'REVERSE'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_local_cumsum_vector_kernel(
s,
o,
scale,
cu_seqlens,
chunk_indices,
T,
B: tl.constexpr,
H: tl.constexpr,
S: tl.constexpr,
BT: tl.constexpr,
BS: tl.constexpr,
REVERSE: tl.constexpr,
HAS_SCALE: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_s, i_t, i_bh = tl.program_id(0), tl.program_id(1).to(tl.int64), tl.program_id(2).to(tl.int64)
i_b, i_h = i_bh // H, i_bh % H
if IS_VARLEN:
i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64)
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
else:
bos, eos = i_b * T, i_b * T + T
o_t = i_t * BT + tl.arange(0, BT)
o_s = i_s * BS + tl.arange(0, BS)
m_s = (o_t[:, None] < T) & (o_s[None, :] < S)
p_s = s + (bos * H + i_h) * S + o_t[:, None] * (H*S) + o_s[None, :]
p_o = o + (bos * H + i_h) * S + o_t[:, None] * (H*S) + o_s[None, :]
# [BT, BS]
b_s = tl.load(p_s, mask=m_s, other=0.0).to(tl.float32)
if REVERSE:
b_o = tl.cumsum(b_s, axis=0, reverse=True)
else:
b_o = tl.cumsum(b_s, axis=0)
if HAS_SCALE:
b_o *= scale
tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=m_s)
@triton.heuristics({
'HAS_SCALE': lambda args: args['scale'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@triton.autotune(
configs=[
triton.Config({'BT': BT}, num_warps=num_warps, num_stages=num_stages)
for BT in [32, 64, 128, 256]
for num_warps in [2, 4, 8]
for num_stages in [1, 2, 3, 4]
],
key=['B', 'H', 'IS_VARLEN', 'REVERSE'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_global_cumsum_scalar_kernel(
s,
o,
scale,
cu_seqlens,
T,
B: tl.constexpr,
H: tl.constexpr,
BT: tl.constexpr,
REVERSE: tl.constexpr,
HAS_SCALE: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_nh = tl.program_id(0).to(tl.int64)
i_n, i_h = i_nh // H, i_nh % H
if IS_VARLEN:
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
else:
bos, eos = i_n * T, i_n * T + T
T = eos - bos
b_z = tl.zeros([], dtype=tl.float32)
NT = tl.cdiv(T, BT)
for i_c in range(NT):
i_t = NT - 1 - i_c if REVERSE else i_c
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
p_s = s + bos*H + i_h + o_t * H
p_o = o + bos*H + i_h + o_t * H
b_s = tl.load(p_s, mask=m_t, other=0.0).to(tl.float32)
if REVERSE:
b_o = tl.cumsum(b_s, axis=0, reverse=True)
else:
b_o = tl.cumsum(b_s, axis=0)
b_ss = tl.sum(b_s, 0)
b_o += b_z
if i_c >= 0:
b_z += b_ss
if HAS_SCALE:
b_o *= scale
tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=m_t)
@triton.heuristics({
'HAS_SCALE': lambda args: args['scale'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@triton.autotune(
configs=[
triton.Config({'BT': BT}, num_warps=num_warps, num_stages=num_stages)
for BT in [16, 32, 64, 128]
for num_warps in [2, 4, 8]
for num_stages in [1, 2, 3, 4]
],
key=['B', 'H', 'S', 'IS_VARLEN', 'REVERSE'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_global_cumsum_vector_kernel(
s,
o,
scale,
cu_seqlens,
T,
B: tl.constexpr,
H: tl.constexpr,
S: tl.constexpr,
BT: tl.constexpr,
BS: tl.constexpr,
REVERSE: tl.constexpr,
HAS_SCALE: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_s, i_nh = tl.program_id(0), tl.program_id(1).to(tl.int64)
i_n, i_h = i_nh // H, i_nh % H
if IS_VARLEN:
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
else:
bos, eos = i_n * T, i_n * T + T
T = eos - bos
b_z = tl.zeros([BS], dtype=tl.float32)
NT = tl.cdiv(T, BT)
for i_c in range(NT):
i_t = NT - 1 - i_c if REVERSE else i_c
o_t = i_t * BT + tl.arange(0, BT)
o_s = i_s * BS + tl.arange(0, BS)
m_s = (o_t[:, None] < T) & (o_s[None, :] < S)
p_s = s + (bos * H + i_h) * S + o_t[:, None] * (H*S) + o_s[None, :]
p_o = o + (bos * H + i_h) * S + o_t[:, None] * (H*S) + o_s[None, :]
# [BT, BS]
b_s = tl.load(p_s, mask=m_s, other=0.0).to(tl.float32)
if REVERSE:
b_c = b_z[None, :] + tl.cumsum(b_s, axis=0, reverse=True)
else:
b_c = b_z[None, :] + tl.cumsum(b_s, axis=0)
if HAS_SCALE:
b_c *= scale
tl.store(p_o, b_c.to(p_o.dtype.element_ty), mask=m_s)
b_z += tl.sum(b_s, 0)
def chunk_local_cumsum_scalar(
g: torch.Tensor,
chunk_size: int,
reverse: bool = False,
scale: float = None,
cu_seqlens: torch.Tensor | None = None,
output_dtype: torch.dtype | None = torch.float,
chunk_indices: torch.LongTensor | None = None,
**kwargs,
) -> torch.Tensor:
if 'head_first' in kwargs:
raise DeprecationWarning(
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
)
B, T, H = g.shape
assert chunk_size == 2**(chunk_size.bit_length()-1), "chunk_size must be a power of 2"
BT = chunk_size
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
g_org, g = g, torch.empty_like(g, dtype=output_dtype or g.dtype)
grid = (NT, B * H)
chunk_local_cumsum_scalar_kernel[grid](
s=g_org,
o=g,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
T=T,
B=B,
H=H,
BT=BT,
REVERSE=reverse,
)
return g
def chunk_local_cumsum_vector(
g: torch.Tensor,
chunk_size: int,
reverse: bool = False,
scale: float = None,
cu_seqlens: torch.Tensor | None = None,
output_dtype: torch.dtype | None = torch.float,
chunk_indices: torch.LongTensor | None = None,
**kwargs,
) -> torch.Tensor:
if 'head_first' in kwargs:
raise DeprecationWarning(
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
)
B, T, H, S = g.shape
BT = chunk_size
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
assert chunk_size == 2**(chunk_size.bit_length()-1), "chunk_size must be a power of 2"
g_org, g = g, torch.empty_like(g, dtype=output_dtype or g.dtype)
def grid(meta): return (triton.cdiv(meta['S'], meta['BS']), NT, B * H)
# keep cummulative normalizer in fp32
# this kernel is equivalent to
# g = g.view(B, H, NT, BT, -1).cumsum(-2).view(B, H, T, -1)
chunk_local_cumsum_vector_kernel[grid](
s=g_org,
o=g,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
T=T,
B=B,
H=H,
S=S,
BT=BT,
REVERSE=reverse,
)
return g
@input_guard
def chunk_global_cumsum_scalar(
s: torch.Tensor,
reverse: bool = False,
cu_seqlens: torch.Tensor | None = None,
scale: float = None,
output_dtype: torch.dtype | None = torch.float,
**kwargs,
) -> torch.Tensor:
if 'head_first' in kwargs:
raise DeprecationWarning(
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
)
B, T, H = s.shape
N = len(cu_seqlens) - 1 if cu_seqlens is not None else B
z = torch.empty_like(s, dtype=output_dtype or s.dtype)
grid = (N * H,)
chunk_global_cumsum_scalar_kernel[grid](
s=s,
o=z,
scale=scale,
cu_seqlens=cu_seqlens,
T=T,
B=B,
H=H,
REVERSE=reverse,
)
return z
@input_guard
def chunk_global_cumsum_vector(
s: torch.Tensor,
reverse: bool = False,
cu_seqlens: torch.Tensor | None = None,
scale: float = None,
output_dtype: torch.dtype | None = torch.float,
**kwargs,
) -> torch.Tensor:
if 'head_first' in kwargs:
raise DeprecationWarning(
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
)
B, T, H, S = s.shape
N = len(cu_seqlens) - 1 if cu_seqlens is not None else B
BS = min(32, triton.next_power_of_2(S))
z = torch.empty_like(s, dtype=output_dtype or s.dtype)
grid = (triton.cdiv(S, BS), N * H)
chunk_global_cumsum_vector_kernel[grid](
s=s,
o=z,
scale=scale,
cu_seqlens=cu_seqlens,
T=T,
B=B,
H=H,
S=S,
BS=BS,
REVERSE=reverse,
)
return z
@input_guard
@dispatch('utils')
def chunk_global_cumsum(
s: torch.Tensor,
reverse: bool = False,
cu_seqlens: torch.Tensor | None = None,
scale: float = None,
output_dtype: torch.dtype | None = torch.float,
**kwargs,
) -> torch.Tensor:
if 'head_first' in kwargs:
raise DeprecationWarning(
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
)
if cu_seqlens is not None:
assert s.shape[0] == 1, "Only batch size 1 is supported when cu_seqlens are provided"
if len(s.shape) == 3:
return chunk_global_cumsum_scalar(
s=s,
reverse=reverse,
cu_seqlens=cu_seqlens,
scale=scale,
output_dtype=output_dtype,
)
elif len(s.shape) == 4:
return chunk_global_cumsum_vector(
s=s,
reverse=reverse,
cu_seqlens=cu_seqlens,
scale=scale,
output_dtype=output_dtype,
)
else:
raise ValueError(
f"Unsupported input shape {s.shape}, "
f"which should be [B, T, H] or [B, T, H, D]",
)
@input_guard
@dispatch('utils')
def chunk_local_cumsum(
g: torch.Tensor,
chunk_size: int,
reverse: bool = False,
scale: float = None,
cu_seqlens: torch.Tensor | None = None,
output_dtype: torch.dtype | None = torch.float,
chunk_indices: torch.LongTensor | None = None,
**kwargs,
) -> torch.Tensor:
if 'head_first' in kwargs:
raise DeprecationWarning(
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
)
if cu_seqlens is not None:
assert g.shape[0] == 1, "Only batch size 1 is supported when cu_seqlens are provided"
if len(g.shape) == 3:
return chunk_local_cumsum_scalar(
g=g,
chunk_size=chunk_size,
reverse=reverse,
scale=scale,
cu_seqlens=cu_seqlens,
output_dtype=output_dtype,
chunk_indices=chunk_indices,
)
elif len(g.shape) == 4:
return chunk_local_cumsum_vector(
g=g,
chunk_size=chunk_size,
reverse=reverse,
scale=scale,
cu_seqlens=cu_seqlens,
output_dtype=output_dtype,
chunk_indices=chunk_indices,
)
else:
raise ValueError(
f"Unsupported input shape {g.shape}, "
f"which should be (B, T, H) or (B, T, H, D)",
)
+183
View File
@@ -0,0 +1,183 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import torch
import torch.nn.functional as F
import triton
import triton.language as tl
from kda._fla.utils import autotune_cache_kwargs, tensor_cache
@triton.autotune(
configs=[
triton.Config({}, num_warps=num_warps)
for num_warps in [4, 8, 16, 32]
],
key=['B'],
**autotune_cache_kwargs,
)
@triton.jit
def prepare_position_ids_kernel(
y,
cu_seqlens,
B: tl.constexpr,
):
i_n = tl.program_id(0)
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
T = eos - bos
o = tl.arange(0, B)
for i in range(0, tl.cdiv(T, B) * B, B):
o_i = o + i
tl.store(y + bos + o_i, o_i, o_i < T)
@tensor_cache
def prepare_lens(cu_seqlens: torch.LongTensor) -> torch.LongTensor:
return torch.diff(cu_seqlens)
@tensor_cache
def prepare_lens_from_mask(mask: torch.BoolTensor) -> torch.LongTensor:
return mask.sum(dim=-1, dtype=torch.int32)
@tensor_cache
def prepare_cu_seqlens_from_lens(
lens: torch.LongTensor,
dtype: torch.dtype | None = torch.int32,
) -> torch.LongTensor:
return F.pad(lens.cumsum(dim=0, dtype=dtype), (1, 0))
@tensor_cache
def prepare_cu_seqlens_from_mask(
mask: torch.BoolTensor,
dtype: torch.dtype | None = torch.int32,
) -> torch.LongTensor:
return prepare_cu_seqlens_from_lens(prepare_lens_from_mask(mask), dtype)
@tensor_cache
def prepare_split_cu_seqlens(
batch_size: int | None = None,
seq_len: int | None = None,
split_size: int | None = None,
cu_seqlens: torch.LongTensor | None = None,
dtype: torch.dtype | None = torch.int32,
device: torch.device | None = torch.device('cpu'),
) -> torch.LongTensor:
"""Sub-split a (optionally packed) batch along the token axis.
Two calling modes:
- **Rectangular batch**: pass `batch_size` and `seq_len`, leave
`cu_seqlens=None`. Internally synthesizes `[0, L, 2L, ..., B*L]`.
- **Packed varlen**: pass `cu_seqlens`. `batch_size` and `seq_len` are
ignored (kept as optional kwargs for backward-compat with callers
that used to pass dummies).
`split_size` is always required.
The legacy positional signature `(batch_size, seq_len, split_size, ...)`
continues to work — the first two args retain their position but may now
be omitted when `cu_seqlens` is supplied.
"""
if split_size is None:
raise TypeError("prepare_split_cu_seqlens() requires `split_size`")
if cu_seqlens is None:
if batch_size is None or seq_len is None:
raise TypeError(
"prepare_split_cu_seqlens() requires either `cu_seqlens`, "
"or both `batch_size` and `seq_len`"
)
total_tokens = batch_size * seq_len
cu_seqlens = list(range(0, total_tokens, seq_len)) + [total_tokens]
else:
cu_seqlens = cu_seqlens.tolist()
return torch.tensor(
[
i
for bos, eos in zip(cu_seqlens[:-1], cu_seqlens[1:], strict=False)
for i in range(bos, eos, split_size)
] + [cu_seqlens[-1]],
dtype=dtype,
device=device,
)
def _segmented_arange(counts: torch.LongTensor) -> tuple[torch.LongTensor, torch.LongTensor]:
"""Expand per-segment counts into flat per-slot index tensors.
Given segment sizes ``counts = [c0, c1, ...]``, return two 1-D tensors of
length ``counts.sum()`` that together label every slot with its segment and
its position within that segment.
Example -- ``counts = [2, 3]`` (segment 0 spans 2 slots, segment 1 spans 3)::
seg_id = [0, 0, 1, 1, 1] # which segment each slot belongs to
intra_idx = [0, 1, 0, 1, 2] # running index within that segment
With CUDA ``counts``, ``repeat_interleave`` reads ``counts.sum()`` on the
host (one device sync). Pass host-side counts to avoid it.
"""
seg_id = torch.repeat_interleave(
torch.arange(counts.numel(), device=counts.device, dtype=counts.dtype),
counts,
)
seg_start = F.pad(counts.cumsum(0), (1, 0))[:-1]
intra_idx = torch.arange(seg_id.shape[0], device=counts.device, dtype=counts.dtype) - seg_start[seg_id]
return seg_id, intra_idx
@tensor_cache
def prepare_position_ids(cu_seqlens: torch.LongTensor, cu_seqlens_cpu: torch.LongTensor | None = None) -> torch.LongTensor:
src = cu_seqlens_cpu if cu_seqlens_cpu is not None else cu_seqlens
_, position_ids = _segmented_arange(prepare_lens(src))
return position_ids.to(cu_seqlens)
@tensor_cache
def prepare_sequence_ids(cu_seqlens: torch.LongTensor, cu_seqlens_cpu: torch.LongTensor | None = None) -> torch.LongTensor:
return prepare_position_ids(cu_seqlens, cu_seqlens_cpu).eq(0).cumsum(0) - 1
@tensor_cache
def prepare_token_indices(cu_seqlens: torch.LongTensor, cu_seqlens_cpu: torch.LongTensor | None = None) -> torch.LongTensor:
position_ids = prepare_position_ids(cu_seqlens, cu_seqlens_cpu)
return torch.stack([prepare_sequence_ids(cu_seqlens, cu_seqlens_cpu), position_ids], 1).to(cu_seqlens)
@tensor_cache
def prepare_chunk_indices(
cu_seqlens: torch.LongTensor,
chunk_size: int,
cu_seqlens_cpu: torch.LongTensor | None = None,
) -> torch.LongTensor:
src = cu_seqlens_cpu if cu_seqlens_cpu is not None else cu_seqlens
chunk_counts = (prepare_lens(src) + (chunk_size - 1)).div(chunk_size, rounding_mode='floor')
seg_id, intra_chunk_idx = _segmented_arange(chunk_counts)
return torch.stack([seg_id, intra_chunk_idx], 1).to(cu_seqlens)
@tensor_cache
def prepare_chunk_offsets(
cu_seqlens: torch.LongTensor,
chunk_size: int,
) -> torch.LongTensor:
return F.pad(triton.cdiv(prepare_lens(cu_seqlens), chunk_size), (1, 0), value=0).cumsum(-1)
@tensor_cache
def get_max_num_splits(
cu_seqlens: torch.LongTensor,
chunk_size: int,
cu_seqlens_cpu: torch.LongTensor | None = None
) -> int:
if cu_seqlens_cpu is not None:
return triton.cdiv(int(max(prepare_lens(cu_seqlens_cpu))), chunk_size)
return triton.cdiv(int(max(prepare_lens(cu_seqlens))), chunk_size)
+101
View File
@@ -0,0 +1,101 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import os
import triton
import triton.language as tl
import triton.language.extra.libdevice as tldevice
from kda._fla.utils import IS_GATHER_SUPPORTED, IS_NVIDIA_BLACKWELL
if os.environ.get('FLA_USE_FAST_OPS', '0') == '1':
@triton.jit
def exp(x): return tldevice.fast_expf(x.to(tl.float32))
@triton.jit
def exp2(x): return tldevice.exp2(x.to(tl.float32))
@triton.jit
def log(x): return tldevice.fast_logf(x.to(tl.float32))
@triton.jit
def log2(x): return tldevice.fast_log2f(x.to(tl.float32))
@triton.jit
def tanh(x): return tldevice.fast_tanhf(x.to(tl.float32))
else:
@triton.jit
def exp(x): return tl.exp(x.to(tl.float32))
@triton.jit
def exp2(x): return tl.math.exp2(x.to(tl.float32))
@triton.jit
def log(x): return tl.log(x.to(tl.float32))
@triton.jit
def log2(x): return tl.log2(x.to(tl.float32))
@triton.jit
def tanh(x): return tldevice.tanh(x.to(tl.float32))
if IS_NVIDIA_BLACKWELL:
"""
Compute tl.dot with Blackwell workaround.
On SM100 datacenter and SM120 consumer Blackwell GPUs, wraps the result in
inline assembly to prevent the TritonGPUHoistTMEMAlloc pass from incorrectly
fusing add and dot operations.
See: https://github.com/fla-org/flash-linear-attention/issues/638
TODO: Remove this workaround once the Triton compiler bug is fixed.
Track upstream issue at: https://github.com/triton-lang/triton/issues/8695
"""
@triton.jit
def safe_dot(a, b, allow_tf32: tl.constexpr = None):
return tl.inline_asm_elementwise(
asm="mov.f32 $0, $1;",
constraints="=r,r",
args=[tl.dot(a, b, allow_tf32=allow_tf32)],
dtype=tl.float32,
is_pure=True,
pack=1,
)
else:
@triton.jit
def safe_dot(a, b, allow_tf32: tl.constexpr = None):
return tl.dot(a, b, allow_tf32=allow_tf32)
if not IS_GATHER_SUPPORTED:
@triton.jit
def gather(src, index, axis, _builder=None):
"""
Gather operation that works when tl.gather is not supported.
This is a fallback implementation that returns None.
Just to make triton compiler happy.
"""
return None
else:
gather = tl.gather
if hasattr(triton.language, '_experimental_make_tensor_descriptor'):
# For Triton 3.3.x
make_tensor_descriptor = triton.language._experimental_make_tensor_descriptor
elif hasattr(triton.language, 'make_tensor_descriptor'):
# For Triton 3.4.x and later
make_tensor_descriptor = triton.language.make_tensor_descriptor
else:
"""
Fallback implementation when TMA is not supported.
Returns None to indicate TMA descriptors are unavailable.
Just make triton compiler happy.
"""
@triton.jit
def make_tensor_descriptor(
base,
shape,
strides,
block_shape,
_builder=None,
):
return None
+115
View File
@@ -0,0 +1,115 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
# REVISED FROM
# https://github.com/shawntan/stickbreaking-attention/blob/main/stickbreaking_attention/sb_varlen/softplus.py
import triton
from triton import language as tl
from kda._fla.utils import IS_NVIDIA
def _generate_softplus(num_pack):
template = """
.reg .pred p;
setp.gt.f32 p, ${in_reg}, 20.;
@p mov.f32 ${out_reg}, ${in_reg};
@!p mul.f32 ${out_reg}, ${in_reg}, 1.4426950408889634;
@!p ex2.approx.ftz.f32 ${out_reg}, ${out_reg};
@!p add.f32 ${out_reg}, ${out_reg}, 1.0;
@!p lg2.approx.ftz.f32 ${out_reg}, ${out_reg};
@!p mul.f32 ${out_reg}, ${out_reg}, 0.6931471805599453;
"""
out_str = ""
for i in range(num_pack):
inner_str = template.format(out_reg=i, in_reg=i + num_pack)
out_str += "{" + inner_str + "}\n"
# flatten out because torch.compile doesn't like newlines
out_str = " ".join(out_str.split("\n"))
return out_str
def _generate_softplus2(num_pack):
template = """
.reg .pred p;
setp.gt.f32 p, ${in_reg}, 15.;
@p mov.f32 ${out_reg}, ${in_reg};
@!p ex2.approx.ftz.f32 ${out_reg}, ${in_reg};
@!p add.f32 ${out_reg}, ${out_reg}, 1.0;
@!p lg2.approx.ftz.f32 ${out_reg}, ${out_reg};
"""
out_str = ""
for i in range(num_pack):
inner_str = template.format(out_reg=i, in_reg=i + num_pack)
out_str += "{" + inner_str + "}\n"
# flatten out because torch.compile doesn't like newlines
out_str = " ".join(out_str.split("\n"))
return out_str
def _generate_constraints(num_pack):
return ",".join("=r" for i in range(num_pack)) + "," + ",".join("r" for i in range(num_pack))
_NUM_REG = 1
s_softplus: tl.constexpr = tl.constexpr(_generate_softplus(_NUM_REG))
s_softplus2: tl.constexpr = tl.constexpr(_generate_softplus2(_NUM_REG))
s_constraints: tl.constexpr = tl.constexpr(_generate_constraints(_NUM_REG))
NUM_REG: tl.constexpr = tl.constexpr(_NUM_REG)
@triton.jit
def softplus_nv(x):
# equivalent to:
# return tl.where(x < 20.0, tl.math.log(1 + tl.math.exp(x)), x)
return tl.inline_asm_elementwise(
asm=s_softplus,
constraints=s_constraints,
pack=NUM_REG,
args=[
x,
],
dtype=tl.float32,
is_pure=True,
)
@triton.jit
def softplus_triton(x):
return tl.where(x < 20.0, tl.math.log(1 + tl.math.exp(x)), x)
@triton.jit
def softplus2_nv(x):
# equivalent to:
# return tl.where(x < 15.0, tl.math.log2(1 + tl.math.exp2(x)), x)
return tl.inline_asm_elementwise(
asm=s_softplus2,
constraints=s_constraints,
pack=NUM_REG,
args=[
x,
],
dtype=tl.float32,
is_pure=True,
)
@triton.jit
def softplus2_triton(x):
return tl.where(x < 15.0, tl.math.log2(1 + tl.math.exp2(x)), x)
if IS_NVIDIA:
softplus = softplus_nv
softplus2 = softplus2_nv
else:
softplus = softplus_triton
softplus2 = softplus2_triton
+92
View File
@@ -0,0 +1,92 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import sys
from ._compat import ( # noqa: F401
SUPPORTS_AUTOTUNE_CACHE,
TRITON_ABOVE_3_4_0,
TRITON_ABOVE_3_5_1,
TRITON_ABOVE_3_7_1,
autotune_cache_kwargs,
find_spec_cached,
has_usable_nvcc,
)
from ._config import ( # noqa: F401
FLA_CACHE_RESULTS,
FLA_CI_ENV,
FLA_DISABLE_TENSOR_CACHE,
FLA_TENSOR_CACHE_SIZE,
)
from ._decorators import ( # noqa: F401
Action,
checkpoint,
contiguous,
deprecate_kwarg,
input_guard,
require_version,
tensor_cache,
)
from ._device import ( # noqa: F401
IS_AMD,
IS_ARM,
IS_GATHER_SUPPORTED,
IS_INTEL,
IS_INTEL_ALCHEMIST,
IS_NPU,
IS_NVIDIA,
IS_NVIDIA_BLACKWELL,
IS_NVIDIA_HOPPER,
IS_NVIDIA_SM100,
IS_NVIDIA_SM120,
IS_TF32_SUPPORTED,
IS_TMA_SUPPORTED,
Backend,
autocast_custom_bwd,
autocast_custom_fwd,
check_environments,
check_pytorch_version,
check_shared_mem,
custom_device_ctx,
device,
device_name,
device_platform,
device_torch_lib,
get_all_max_shared_mem,
get_available_device,
get_device_capability,
get_device_smem_optin,
get_multiprocessor_count,
map_triton_backend_to_torch_device,
)
from ._testing import assert_close, get_abs_err, get_err_ratio # noqa: F401
def _register_aliases():
current_module = sys.modules[__name__]
for key in (
'IS_AMD',
'IS_ARM',
'IS_INTEL',
'IS_INTEL_ALCHEMIST',
'IS_NVIDIA',
'IS_NPU',
'IS_NVIDIA_BLACKWELL',
'IS_NVIDIA_HOPPER',
'IS_NVIDIA_SM100',
'IS_NVIDIA_SM120',
'IS_TF32_SUPPORTED',
'IS_GATHER_SUPPORTED',
'IS_TMA_SUPPORTED',
):
if hasattr(current_module, key):
setattr(current_module, key.lower(), getattr(current_module, key))
_register_aliases()
del _register_aliases
+65
View File
@@ -0,0 +1,65 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import functools
import importlib.metadata
import inspect
import logging
import os
import shutil
from importlib.util import find_spec
from pathlib import Path
import triton
from packaging import version as package_version
from ._config import FLA_CACHE_RESULTS
logger = logging.getLogger(__name__)
TRITON_ABOVE_3_4_0 = package_version.parse(triton.__version__) >= package_version.parse("3.4.0")
TRITON_ABOVE_3_5_1 = package_version.parse(triton.__version__) >= package_version.parse("3.5.1")
TRITON_ABOVE_3_7_1 = package_version.parse(triton.__version__) >= package_version.parse("3.7.1")
SUPPORTS_AUTOTUNE_CACHE = "cache_results" in inspect.signature(triton.autotune).parameters
autotune_cache_kwargs = {"cache_results": FLA_CACHE_RESULTS} if SUPPORTS_AUTOTUNE_CACHE else {}
@functools.cache
def find_spec_cached(name):
return find_spec(name)
@functools.cache
def has_usable_nvcc() -> bool:
"""Whether a usable nvcc compiler is available for TileLang's JIT.
Mirrors the guesses in ``tilelang.env._find_cuda_home`` (env
CUDA_HOME/CUDA_PATH, nvcc on PATH, the ``nvidia-cuda-nvcc`` wheel,
/usr/local/cuda), but verifies the nvcc binary actually exists —
only ``nvidia-cuda-nvcc`` >= 13.0 ships it, the ``-cu12`` variant
installs just ptxas.
"""
cuda_home = os.environ.get("CUDA_HOME") or os.environ.get("CUDA_PATH")
if cuda_home is not None and (Path(cuda_home) / "bin" / "nvcc").exists():
return True
if shutil.which("nvcc") is not None:
return True
try:
files = importlib.metadata.files("nvidia-cuda-nvcc") or []
except importlib.metadata.PackageNotFoundError:
files = []
if any(f.name in ("nvcc", "nvcc.exe") for f in files):
return True
if (Path("/usr/local/cuda") / "bin" / "nvcc").exists():
return True
logger.info(
"[FLA Backend] TileLang is installed but no usable nvcc compiler was found; falling back to Triton. "
"Install a CUDA toolkit or nvidia-cuda-nvcc, or set FLA_TILELANG=0 to disable TileLang explicitly."
)
return False
+17
View File
@@ -0,0 +1,17 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import os
FLA_CI_ENV = os.getenv("FLA_CI_ENV") == "1"
FLA_CACHE_RESULTS = os.getenv('FLA_CACHE_RESULTS', '1') == '1'
FLA_DISABLE_TENSOR_CACHE = os.getenv('FLA_DISABLE_TENSOR_CACHE', '0') == '1'
try:
FLA_TENSOR_CACHE_SIZE = int(os.getenv('FLA_TENSOR_CACHE_SIZE', "4"))
except ValueError:
FLA_TENSOR_CACHE_SIZE = 4
+336
View File
@@ -0,0 +1,336 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import contextlib
import functools
import inspect
import sys
import warnings
from collections import deque
from collections.abc import Callable
from enum import Enum
from typing import Any
import torch
from packaging import version as package_version
from .. import __version__
from ._config import FLA_DISABLE_TENSOR_CACHE, FLA_TENSOR_CACHE_SIZE
from ._device import custom_device_ctx
class Action(Enum):
NONE = "none"
NOTIFY = "notify"
NOTIFY_ALWAYS = "notify_always"
RAISE = "raise"
def tensor_cache(
fn: Callable[..., torch.Tensor],
) -> Callable[..., torch.Tensor]:
"""
A decorator that memoizes the most recent results of a function call by argument identity.
The decorator keeps a bounded queue of up to ``FLA_TENSOR_CACHE_SIZE`` (default 4)
recent ``(args, kwargs, result)`` triples. On each call, every cached entry is checked
in order; an entry is considered a hit when the positional arg count and kwarg key set
match and every argument is the *same object* (``is`` identity) as the cached one. On a
hit the cached result is returned and ``fn`` is skipped; on a miss ``fn`` is invoked and
the new triple is appended (evicting the oldest when the queue is full).
Caching is fully bypassed when the ``FLA_DISABLE_TENSOR_CACHE`` environment variable is
set to ``'1'``.
Args:
fn (Callable[..., torch.Tensor]):
The function to be decorated. Intended for functions whose inputs are tensors
(or other objects compared by identity) and whose output is a tensor.
Returns:
Callable[..., torch.Tensor]:
A wrapped version of ``fn`` backed by an identity-based bounded cache.
"""
cached: deque = deque(maxlen=FLA_TENSOR_CACHE_SIZE)
def cache_disabled() -> bool:
utils_module = sys.modules.get('kda._fla.utils')
return getattr(utils_module, 'FLA_DISABLE_TENSOR_CACHE', FLA_DISABLE_TENSOR_CACHE)
@functools.wraps(fn)
def wrapper(*args: Any, **kwargs: Any) -> Any:
if cache_disabled():
return fn(*args, **kwargs)
for cached_args, cached_kwargs, cached_result in cached:
if len(args) != len(cached_args) or len(kwargs) != len(cached_kwargs):
continue
if all(a is b for a, b in zip(args, cached_args, strict=False)) and \
all(k in cached_kwargs and v is cached_kwargs[k] for k, v in kwargs.items()):
return cached_result
result = fn(*args, **kwargs)
cached.append((args, kwargs, result))
return result
return wrapper
def _skip_contiguous(
no_guard_contiguous: bool | list[str] | tuple[str, ...] | set[str],
param_name: str,
skip_params: set[str],
) -> bool:
return no_guard_contiguous is True or param_name in skip_params
def _contiguous_if_needed(arg: Any, skip: bool) -> Any:
if isinstance(arg, torch.Tensor) and not skip:
return arg.contiguous()
return arg
def input_guard(
fn: Callable[..., torch.Tensor] | None = None,
*,
no_guard_contiguous: bool | list[str] | tuple[str, ...] | set[str] = False,
) -> Callable[[Callable[..., torch.Tensor]], Callable[..., torch.Tensor]] | Callable[..., torch.Tensor]:
"""
A decorator to make sure all input tensors are contiguous and set the device based on input tensors.
Args:
no_guard_contiguous (bool | list[str] | tuple[str, ...] | set[str]):
If True, skip all contiguous checks. If a list/tuple/set of parameter names, skip contiguous check for those parameters.
"""
def decorator(fn: Callable[..., torch.Tensor]) -> Callable[..., torch.Tensor]:
# Get function signature for parameter name mapping
sig = inspect.signature(fn)
param_names = list(sig.parameters.keys())
skip_params = set(no_guard_contiguous) if isinstance(no_guard_contiguous, (list, tuple, set)) else set()
@functools.wraps(fn)
def wrapper(*args, **kwargs):
# Process args with parameter name mapping
processed_args = []
for i, arg in enumerate(args):
if i < len(param_names):
param_name = param_names[i]
else:
# For *args beyond signature, use position as name
param_name = f"__arg_{i}"
processed_args.append(_contiguous_if_needed(
arg, _skip_contiguous(no_guard_contiguous, param_name, skip_params)))
# Process kwargs
processed_kwargs = {}
for k, v in kwargs.items():
processed_kwargs[k] = _contiguous_if_needed(v, _skip_contiguous(no_guard_contiguous, k, skip_params))
tensor = None
for arg in args:
if isinstance(arg, torch.Tensor):
tensor = arg
break
if tensor is None:
for value in kwargs.values():
if isinstance(value, torch.Tensor):
tensor = value
break
if tensor is not None:
ctx = custom_device_ctx(tensor.device.index)
else:
ctx = contextlib.nullcontext()
with ctx:
return fn(*processed_args, **processed_kwargs)
return wrapper
# Handle direct usage without parentheses: @input_guard
if fn is not None:
return decorator(fn)
return decorator
def contiguous(fn: Callable[..., torch.Tensor]) -> Callable[..., torch.Tensor]:
"""Alias for input_guard() without parameters."""
return input_guard(fn)
def require_version(version, hint):
"""
Perform a runtime check of the dependency versions, using the exact same syntax used by pip.
"""
def decorator(fn):
@functools.wraps(fn)
def wrapper(ctx, *args, **kwargs):
from transformers.utils.versions import require_version
require_version(version, hint)
return fn(
ctx,
*(i if not isinstance(i, torch.Tensor) else i.contiguous() for i in args),
**{k: (v if not isinstance(v, torch.Tensor) else v.contiguous()) for k, v in kwargs.items()},
)
return wrapper
return decorator
def deprecate_kwarg(
old_name: str,
version: str,
new_name: str | None = None,
warn_if_greater_or_equal_version: bool = False,
raise_if_greater_or_equal_version: bool = False,
raise_if_both_names: bool = False,
additional_message: str | None = None,
):
"""
Decorator to notify users about deprecated keyword arguments, replacing them with a new name if specified.
This decorator allows you to:
- Notify users when a keyword argument is deprecated.
- Automatically replace deprecated keyword arguments with new ones.
- Raise an error if deprecated arguments are used, depending on the specified conditions.
By default, the decorator notifies the user about the deprecated argument while the `fla.__version__` < specified `version`
in the decorator. To keep notifications with any version `warn_if_greater_or_equal_version=True` can be set.
Args:
old_name (`str`):
Name of the deprecated keyword argument.
version (`str`):
The version in which the keyword argument was (or will be) deprecated.
new_name (`Optional[str]`, *optional*):
The new name for the deprecated keyword argument.
If specified, the deprecated keyword argument will be replaced with this new name.
warn_if_greater_or_equal_version (`bool`, *optional*, defaults to `False`):
Whether to show warning if current `fla` version is greater or equal to the deprecated version.
raise_if_greater_or_equal_version (`bool`, *optional*, defaults to `False`):
Whether to raise `ValueError` if current `fla` version is greater or equal to the deprecated version.
raise_if_both_names (`bool`, *optional*, defaults to `False`):
Whether to raise `ValueError` if both deprecated and new keyword arguments are set.
additional_message (`Optional[str]`, *optional*):
An additional message to append to the default deprecation message.
Raises:
ValueError:
If `raise_if_greater_or_equal_version` is `True` and the current version >= the deprecated one,
or if `raise_if_both_names` is `True` and both old and new keyword arguments are provided.
Returns:
Callable:
A wrapped function that handles the deprecated keyword arguments according to the specified parameters.
Example usage with renaming argument:
```python
@deprecate_kwarg("reduce_labels", new_name="do_reduce_labels", version="6.0.0")
def my_function(do_reduce_labels):
print(do_reduce_labels)
my_function(reduce_labels=True) # Will show a deprecation warning and use do_reduce_labels=True
```
Example usage without renaming argument:
```python
@deprecate_kwarg("max_size", version="6.0.0")
def my_function(max_size):
print(max_size)
my_function(max_size=1333) # Will show a deprecation warning
```
"""
deprecated_version = package_version.parse(version)
current_version = package_version.parse(__version__)
is_greater_or_equal_version = current_version >= deprecated_version
if is_greater_or_equal_version:
version_message = f"and removed starting from version {version}"
else:
version_message = f"and will be removed in version {version}"
def wrapper(func):
# Required for better warning message
sig = inspect.signature(func)
function_named_args = set(sig.parameters.keys())
is_instance_method = "self" in function_named_args
is_class_method = "cls" in function_named_args
@functools.wraps(func)
def wrapped_func(*args, **kwargs):
# Get class + function name (just for better warning message)
func_name = func.__name__
if is_instance_method:
func_name = f"{args[0].__class__.__name__}.{func_name}"
elif is_class_method:
func_name = f"{args[0].__name__}.{func_name}"
minimum_action = Action.NONE
message = None
# deprecated kwarg and its new version are set for function call -> replace it with new name
if old_name in kwargs and new_name in kwargs:
minimum_action = Action.RAISE if raise_if_both_names else Action.NOTIFY_ALWAYS
message = (
f"Both `{old_name}` and `{new_name}` are set for `{func_name}`. "
f"Using `{new_name}={kwargs[new_name]}` and ignoring deprecated `{old_name}={kwargs[old_name]}`."
)
kwargs.pop(old_name)
# only deprecated kwarg is set for function call -> replace it with new name
elif old_name in kwargs and new_name is not None and new_name not in kwargs:
minimum_action = Action.NOTIFY
message = (
f"`{old_name}` is deprecated {version_message} for `{func_name}`. "
f"Use `{new_name}` instead."
)
kwargs[new_name] = kwargs.pop(old_name)
# deprecated kwarg is not set for function call and new name is not specified -> just notify
elif old_name in kwargs:
minimum_action = Action.NOTIFY
message = f"`{old_name}` is deprecated {version_message} for `{func_name}`."
if message is not None and additional_message is not None:
message = f"{message} {additional_message}"
# update minimum_action if argument is ALREADY deprecated (current version >= deprecated version)
if is_greater_or_equal_version:
# change to (NOTIFY, NOTIFY_ALWAYS) -> RAISE if specified
# in case we want to raise error for already deprecated arguments
if raise_if_greater_or_equal_version and minimum_action != Action.NONE:
minimum_action = Action.RAISE
# change to NOTIFY -> NONE if specified (NOTIFY_ALWAYS can't be changed to NONE)
elif not warn_if_greater_or_equal_version and minimum_action == Action.NOTIFY:
minimum_action = Action.NONE
# raise error or notify user
if minimum_action == Action.RAISE:
raise ValueError(message)
elif minimum_action in (Action.NOTIFY, Action.NOTIFY_ALWAYS):
# DeprecationWarning is ignored by default, so we use FutureWarning instead
warnings.warn(message, FutureWarning, stacklevel=2)
return func(*args, **kwargs)
return wrapped_func
return wrapper
def checkpoint(fn):
@functools.wraps(fn)
def wrapper(*args, **kwargs):
return torch.utils.checkpoint.checkpoint(fn, *args, **kwargs)
return wrapper
+245
View File
@@ -0,0 +1,245 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import contextlib
import functools
import logging
import os
import platform
import sys
import warnings
from enum import Enum
from functools import cache, lru_cache
import torch
import triton
from packaging import version as package_version
logger = logging.getLogger(__name__)
@lru_cache(maxsize=1)
def check_environments():
"""
Checks the current operating system, Triton version, and Python version,
issuing warnings if they don't meet recommendations.
This function's body only runs once due to lru_cache.
"""
# Check Operating System
if sys.platform == 'win32':
# Check if triton-windows is installed
try:
from importlib.metadata import PackageNotFoundError, metadata
metadata('triton-windows')
# triton-windows is installed, no warning needed
except PackageNotFoundError:
logger.warning(
"Detected Windows operating system. Consider installing triton-windows "
"(https://github.com/triton-lang/triton-windows) for better compatibility. "
"Without it, some features may not work correctly.",
)
triton_version = package_version.parse(triton.__version__)
required_triton_version = package_version.parse("3.3.0")
if triton_version < required_triton_version:
logger.warning(
f"Current Triton version {triton_version} is below the recommended 3.3.0 version. "
"Errors may occur and these issues will not be fixed. "
"Please consider upgrading Triton.",
)
# Check Python version
py_version = package_version.parse(f"{sys.version_info.major}.{sys.version_info.minor}")
required_py_version = package_version.parse("3.11")
if py_version < required_py_version:
logger.warning(
f"Current Python version {py_version} is below the recommended 3.11 version. "
"It is recommended to upgrade to Python 3.11 or higher for the best experience.",
)
return None
check_environments()
def _cpu_device_warning():
warnings.warn(('Triton is not supported on current platform, roll back to CPU.'), stacklevel=2)
@cache
def check_pytorch_version(version_s: str = '2.4') -> bool:
return package_version.parse(torch.__version__) >= package_version.parse(version_s)
@cache
def get_multiprocessor_count(tensor_idx: int = 0, *, use_aicore: bool = False) -> int:
try:
return triton.runtime.driver.active.utils.get_device_properties(tensor_idx)['multiprocessor_count']
except Exception:
# Maybe we use a NPU device.
try:
if triton.runtime.driver.active.get_current_target().backend == 'npu':
props = triton.runtime.driver.active.utils.get_device_properties(tensor_idx)
return props['num_aicore'] if use_aicore else props['num_vectorcore']
except Exception:
logger.debug('Failed to get NPU multiprocessor count, falling back to 1.', exc_info=True)
return 1
@cache
def get_device_capability(device_index: int = 0) -> tuple[int, int]:
major, minor = torch.cuda.get_device_capability(device_index)
return int(major), int(minor)
@cache
def get_device_smem_optin(device_index: int = 0) -> int:
props = torch.cuda.get_device_properties(device_index)
return int(getattr(props, 'shared_memory_per_block_optin', props.shared_memory_per_block))
@cache
def get_available_device() -> str:
try:
return triton.runtime.driver.active.get_current_target().backend
except Exception:
_cpu_device_warning()
return 'cpu'
def map_triton_backend_to_torch_device() -> str:
backend = get_available_device() # 'cuda' | 'hip' | 'xpu' | 'cpu' | ...
return {'cuda': 'cuda', 'hip': 'cuda', 'xpu': 'xpu'}.get(backend, backend)
# For AMD GPUs, the triton backend is 'hip', while for Nvidia GPUs, the triton backend is 'cuda'.
# However, the torch backend is 'cuda' for both Nvidia and AMD GPUs.
# Therefore, we need to check the triton backend to determine the actual GPU vendor.
device = get_available_device() if get_available_device() != 'hip' else 'cuda'
device_torch_lib = getattr(torch, device)
device_platform = get_available_device()
device_name = map_triton_backend_to_torch_device()
IS_AMD = (device_platform == 'hip')
IS_ARM = platform.machine().lower() in ('aarch64', 'arm64')
IS_INTEL = (device_platform == 'xpu')
IS_INTEL_ALCHEMIST = (IS_INTEL and 'Intel(R) Arc(TM) A' in torch.xpu.get_device_name(0))
IS_NPU = (device_platform == 'npu')
IS_NVIDIA = (device_platform == 'cuda')
IS_NVIDIA_HOPPER = (
IS_NVIDIA and (
'NVIDIA H' in torch.cuda.get_device_name(0)
or torch.cuda.get_device_capability()[0] == 9
)
)
IS_NVIDIA_SM100 = (IS_NVIDIA and torch.cuda.get_device_capability()[0] == 10)
# NOTE: exactly 12.0 — 12.1 (GB10) is a different target that FlashQLA rejects at import time.
IS_NVIDIA_SM120 = (IS_NVIDIA and torch.cuda.get_device_capability() == (12, 0))
IS_NVIDIA_BLACKWELL = (IS_NVIDIA and torch.cuda.get_device_capability()[0] in (10, 12))
# Nvidia Ampere or newer, haven't check AMD and intel yet.
IS_TF32_SUPPORTED = (IS_NVIDIA and torch.cuda.get_device_capability(0)[0] >= 8)
IS_GATHER_SUPPORTED = hasattr(triton.language, 'gather')
IS_TMA_SUPPORTED = (
IS_NVIDIA
and torch.cuda.get_device_capability(0)[0] >= 9
and os.environ.get('FLA_USE_TMA', '0') == '1'
and (
hasattr(triton.language, '_experimental_make_tensor_descriptor')
or hasattr(triton.language, 'make_tensor_descriptor')
)
)
if IS_NVIDIA and not IS_TF32_SUPPORTED:
# Make old card happy, since triton will use tf32 by default.
# This is a workaround for old nvidia card.
os.environ['TRITON_F32_DEFAULT'] = 'ieee'
def _default_alloc_fn(size: int, alignment: int, stream: int | None):
return torch.empty(size, device=torch.device(device_name, device_torch_lib.current_device()), dtype=torch.int8)
if IS_TMA_SUPPORTED:
logger.info('TMA is supported, using TMA by default.')
triton.set_allocator(_default_alloc_fn)
elif IS_NVIDIA_BLACKWELL:
# Blackwell (SM100 datacenter / SM120 consumer): Triton compiler may emit global_scratch for
# autotuned kernels even without TMA. Register a default allocator to
# prevent NullAllocator crashes. See triton-lang/triton#10002.
logger.info('Blackwell detected: registering default global_scratch allocator.')
triton.set_allocator(_default_alloc_fn)
def get_all_max_shared_mem():
try:
return [
triton.runtime.driver.active.utils.get_device_properties(i)['max_shared_mem']
for i in range(device_torch_lib.device_count())
]
except Exception:
_cpu_device_warning()
return [-1]
class Backend(Enum):
ADA = 101376 # RTX 4090
AMPERE = 166912 # A100
HOPPER = 232448 # H100
DEFAULT = 102400 # Default
@classmethod
def get_shared_memory(cls, arch: str) -> int:
try:
return cls[arch.upper()].value
except KeyError:
return cls.DEFAULT.value
@cache
def check_shared_mem(arch: str = "none", tensor_idx: int = 0) -> bool:
try:
device_shared_mem_list = get_all_max_shared_mem()
max_shared_memory = device_shared_mem_list[tensor_idx]
return max_shared_memory >= Backend.get_shared_memory(arch)
except Exception:
return False
if check_pytorch_version('2.4'):
if device == 'cpu':
device = 'cuda'
device_torch_lib = getattr(torch, device)
autocast_custom_fwd = functools.partial(torch.amp.custom_fwd, device_type=device)
autocast_custom_bwd = functools.partial(torch.amp.custom_bwd, device_type=device)
def custom_device_ctx(index: int):
if index is None:
return contextlib.nullcontext()
try:
return device_torch_lib.device(index)
except (AttributeError, AssertionError, RuntimeError):
return contextlib.nullcontext()
else:
assert device == 'cuda', 'Only cuda device is supported for PyTorch version < 2.4.0.'
autocast_custom_fwd = device_torch_lib.amp.custom_fwd
autocast_custom_bwd = device_torch_lib.amp.custom_bwd
def custom_device_ctx(index: int):
if index is None:
return contextlib.nullcontext()
try:
return torch.cuda.device(index)
except (AttributeError, AssertionError, RuntimeError):
return contextlib.nullcontext()
+41
View File
@@ -0,0 +1,41 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import logging
import warnings
import torch
from ._config import FLA_CI_ENV
logger = logging.getLogger(__name__)
def get_abs_err(x, y):
return (x.detach() - y.detach()).flatten().abs().max().item()
def get_err_ratio(x, y):
err = (x.detach() - y.detach()).flatten().square().mean().sqrt().item()
base = (x.detach()).flatten().square().mean().sqrt().item()
return err / (base + 1e-8)
def assert_close(prefix, ref, tri, ratio, warning=False, err_atol=1e-6):
abs_atol = get_abs_err(ref, tri)
error_rate = get_err_ratio(ref, tri)
msg = f"{prefix:>16} diff: {abs_atol:.6f} ratio: {error_rate:.6f}"
logger.info(msg)
if abs_atol <= err_atol:
return
assert not torch.isnan(ref).any(), f"{prefix}: NaN detected in ref"
assert not torch.isnan(tri).any(), f"{prefix}: NaN detected in tri"
if warning or (FLA_CI_ENV and (error_rate < 0.01 or abs_atol <= 0.3)):
if error_rate > ratio:
warnings.warn(msg)
else:
assert error_rate < ratio, msg
+23
View File
@@ -0,0 +1,23 @@
"""Composable mixing layers: attn and ffn both map [B,T,D] -> [B,T,D].
Depth mixing (AttnRes) is not a layer_specs kind. CausalLM reads
``config.attnres`` (off | full | block) and wraps DecoderBlock sublayers.
"""
from .block import DecoderBlock, build_attn, build_ffn
from .kda_attn import KDAAttention
from .latent_moe import LatentMoE
from .mla import GatedMLA
from .rmsnorm import RMSNorm
from .swiglu import SwiGLUMLP
__all__ = [
"DecoderBlock",
"GatedMLA",
"KDAAttention",
"LatentMoE",
"RMSNorm",
"SwiGLUMLP",
"build_attn",
"build_ffn",
]
+519
View File
@@ -0,0 +1,519 @@
"""
Attention Residual in one file
Reference:
Kimi Team, Guangyu Chen, Yu Zhang, Jianlin Su, Weixin Xu, Siyuan Pan,
Yaoyu Wang, Yucheng Wang, Guanduo Chen, et al.
"Attention Residuals." arXiv:2603.15031, 2026.
https://arxiv.org/abs/2603.15031
This module is a compact PyTorch reference implementation of:
- Full AttnRes
- Block AttnRes
- two-phase inter/intra-block computation from the paper
CausalLM wires Full/Block stacks when ``config.attnres`` is ``full`` or
``block``. Standard residual (``x += attn; x += ffn``) is ``attnres="off"``.
"""
import torch
import torch.nn.functional as F
from einops import rearrange
from torch import Tensor, nn
ATTNRES_MODES = ("off", "full", "block")
def exists(x):
return x is not None
def validate_attnres(mode: str, block_size: int | None) -> None:
if mode not in ATTNRES_MODES:
raise ValueError(f"attnres must be one of {ATTNRES_MODES}, got {mode!r}")
if block_size is not None and block_size < 1:
raise ValueError(f"attnres_block_size must be >= 1, got {block_size}")
def atomic_block_size(num_hidden_layers: int, attnres_block_size: int | None) -> int:
"""DecoderBlocks per AttnRes block, converted to attn|ffn atomic layers.
``None`` targets about 8 blocks: ``max(1, ceil(L / 8))`` DecoderBlocks.
"""
layers_per_block = (
attnres_block_size
if attnres_block_size is not None
else max(1, (num_hidden_layers + 7) // 8)
)
if layers_per_block < 1:
raise ValueError(f"attnres_block_size must be >= 1, got {layers_per_block}")
return layers_per_block * 2
class BorrowedSubLayer(nn.Module):
"""``fn(norm(x))`` without registering ``norm``/``fn`` (owned by DecoderBlock)."""
def __init__(self, norm: nn.Module, fn: nn.Module):
super().__init__()
self._borrowed = (norm, fn)
def forward(self, x: Tensor) -> Tensor:
norm, fn = self._borrowed
return fn(norm(x))
def rms(x: Tensor, eps: float):
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + eps)
class RMSNorm(nn.Module):
def __init__(self, dim: int, eps: float):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x: Tensor) -> Tensor:
return rms(x, self.eps) * self.weight
class DepthResidual(nn.Module):
"""
h_l = sum_i softmax_i(w_l^T RMSNorm(v_i))*v_i
Keep query and RMSNorm gain separate
Since q^T (gamma * RMS(v)) == (q * gamma)^T RMS(v),
we can fold gamma into q for scoring.
"""
def __init__(self, dim: int, eps: float = 1e-8, zero_init: bool = True):
super().__init__()
self.query = nn.Parameter(torch.zeros(dim))
self.norm = RMSNorm(dim, eps=eps)
if not zero_init:
nn.init.normal_(self.query, std=0.02)
def effective_query(self) -> Tensor:
return (self.query * self.norm.weight).float()
def logits(self, sources: Tensor | list[Tensor] | tuple[Tensor, ...]) -> Tensor:
sources = stack_layers(sources) # [n, b, t, d]
q = self.effective_query() # [d]
k = rms(sources.float(), self.norm.eps) # [n, b, t, d]
return torch.einsum("d, n b t d -> n b t", q, k)
def forward(self, sources: Tensor | list[Tensor] | tuple[Tensor, ...]) -> Tensor:
sources = stack_layers(sources)
weights = self.logits(sources).softmax(dim=0)
out = torch.einsum("n b t, n b t d -> b t d", weights, sources.float())
return out.to(sources.dtype)
class DepthResidualList(nn.Module):
def __init__(self, dim: int, depth: int, eps: float, zero_init: bool = True):
super().__init__()
# for L layers (depth), create depth residual modules
self.layers = nn.ModuleList(
[DepthResidual(dim, eps=eps, zero_init=zero_init) for _ in range(depth)]
)
def __getitem__(self, idx: int) -> DepthResidual:
return self.layers[idx]
def __iter__(self):
return iter(self.layers)
def __len__(self):
return len(self.layers)
# attnres stacks
class FullAttnResStack(nn.Module):
"""
Full AttnRes over atomic layers
eg: f_1,...,f_L
Each entry in `layers` should already be a full atomic layer fxn
x -> f_l(x)
"""
def __init__(
self,
dim: int,
layers,
*,
eps: float = 1e-8,
zero_init_queries: bool = True,
is_final_aggregate: bool = True,
):
super().__init__()
self.layers = nn.ModuleList(list(layers))
self.eps = eps
depth = len(self.layers)
self.residuals = DepthResidualList(dim, depth, eps, zero_init_queries)
self.final_residual = (
DepthResidual(dim, eps, zero_init_queries) if is_final_aggregate else None
)
def forward_naive(self, x: Tensor) -> Tensor:
sources = [x]
for layer, residual in zip(self.layers, self.residuals):
h = residual(sources)
out = layer(h)
sources.append(out)
return (
self.final_residual(sources) if exists(self.final_residual) else sources[-1]
)
def forward_two_phase(self, x: Tensor, schedule_block_size: int) -> Tensor:
assert schedule_block_size > 0
sources = [x]
depth = len(self.layers)
start = 0
while start < depth:
end = min(start + schedule_block_size, depth)
queries = torch.stack(
[self.residuals[i].effective_query() for i in range(start, end)], dim=0
)
inter_sources = stack_layers(sources)
inter_stats = attn_with_stats(queries, inter_sources, self.eps)
local_outputs = [] # outputs of intra-block
for local_idx, layer_idx in enumerate(range(start, end)):
stats = inter_stats.select(local_idx)
if len(local_outputs) > 0:
intra_sources = stack_layers(local_outputs)
intra = attn_with_stats(
queries[local_idx : local_idx + 1], intra_sources, self.eps
).select(0)
stats = merge_attn_stats(stats, intra)
h = stats.normalized()
out = self.layers[layer_idx](h)
local_outputs.append(out)
sources.append(out)
start = end
return (
self.final_residual(sources) if exists(self.final_residual) else sources[-1]
)
def forward(self, x: Tensor, schedule_block_size: int | None = None) -> Tensor:
if schedule_block_size is None:
return self.forward_naive(x)
return self.forward_two_phase(x, schedule_block_size)
class BlockAttnResStack(nn.Module):
"""
Block AttnRes over atomic layers
`block_size` is in atomic layers, not Transformer blocks.
Eg: block_size=8 -> 4 transformer blocks when layers alternate attn/MLP
The default forward path is the two-phase algorithm from the paper:
phase 1: batch inter-block attn from all queries in the block
phase 2: merge the evolving intra-block partial sum with online softmax
"""
def __init__(
self,
dim: int,
layers,
*,
block_size: int,
eps: float = 1e-8,
zero_init_queries: bool = True,
is_final_aggregate: bool = True,
):
super().__init__()
self.layers = nn.ModuleList(list(layers))
assert len(self.layers) > 0
assert block_size > 0
self.block_size = block_size
self.eps = eps
depth = len(self.layers)
self.residuals = DepthResidualList(
dim, depth, eps=eps, zero_init=zero_init_queries
)
self.final_residual = (
DepthResidual(dim, eps=eps, zero_init=zero_init_queries)
if is_final_aggregate
else None
)
def forward_naive(self, x: Tensor) -> Tensor:
blocks = [x] # b_0=embedding/input representation
partial = None
for layer_idx, (layer, residual) in enumerate(
zip(self.layers, self.residuals), start=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) if exists(self.final_residual) else blocks[-1]
)
def _run_block_two_phase(
self, blocks: list[Tensor], start: int, end: int
) -> Tensor:
queries = torch.stack(
[self.residuals[i].effective_query() for i in range(start, end)], dim=0
)
inter_sources = stack_layers(blocks)
inter = attn_with_stats(queries, inter_sources, self.eps)
partial = None
for local_idx, layer_idx in enumerate(range(start, end)):
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)
h = stats.normalized()
out = self.layers[layer_idx](h)
partial = out if partial is None else (partial + out)
return partial
def forward(self, x: Tensor) -> Tensor:
blocks = [x]
depth = len(self.layers)
start = 0
while start < depth:
end = min(start + self.block_size, depth)
blocks.append(self._run_block_two_phase(blocks, start, end))
start = end
return (
self.final_residual(blocks) if exists(self.final_residual) else blocks[-1]
)
# helpers
def stack_layers(sources: Tensor | list[Tensor] | tuple[Tensor, ...]) -> Tensor:
if isinstance(sources, Tensor):
assert sources.ndim == 4, f"expected [n, b, t, d] got {tuple(sources.shape)}"
return sources
assert len(sources) > 0, "needs at least one source"
return torch.stack(tuple(sources), dim=0)
class SingleAttnStats:
def __init__(self, numer: Tensor, denom: Tensor, max: Tensor):
self.numer = numer # [b,t,d]
self.max = max # [b,t]
self.denom = denom # [b,t]
def normalized(self) -> Tensor:
return self.numer / self.denom[..., None]
class AttnStats:
# store the numerator => e^{s_{j}-m} * v_j where m is the max score so far
# store the max m = max(s_j)
# store the denominator sum_j e^{s_{j}-m}
def __init__(self, numer: Tensor, denom: Tensor, max: Tensor):
self.numer = numer # [q,b,t,d]
self.max = max # [q,b,t]
self.denom = denom # [q,b,t]
def select(self, idx: int) -> "SingleAttnStats":
return SingleAttnStats(self.numer[idx], self.denom[idx], self.max[idx])
def attn_with_stats(queries: Tensor, sources: Tensor, eps: float = 1e-8) -> AttnStats:
"""
queries: [q, d]
sources: [n, b, t, d]
Returns the following for online softmax:
numer = sum_i exp(logit_i - m)*v_i
m = max_i logit_i
denom = sum_i exp(logit_i - m)
"""
normed = rms(sources, eps)
logits = torch.einsum("q d, n b t d -> q n b t", queries, normed)
m = logits.amax(dim=1)
weights = torch.exp(logits - m[:, None])
numer = torch.einsum("q n b t, n b t d -> q b t d", weights, sources)
denom = weights.sum(dim=1)
return AttnStats(numer, denom, m)
def single_source_stats(
query: Tensor, source: Tensor, eps: float = 1e-8
) -> SingleAttnStats:
score = torch.einsum("d, b t d -> b t", query, rms(source, eps))
denom = torch.ones_like(score)
return SingleAttnStats(source, denom, score)
def merge_attn_stats(a: SingleAttnStats, b: SingleAttnStats) -> SingleAttnStats:
m = torch.maximum(a.max, b.max)
wa = torch.exp(a.max - m)
wb = torch.exp(b.max - m)
numer = wa[..., None] * a.numer + wb[..., None] * b.numer
denom = wa * a.denom + wb * b.denom
return SingleAttnStats(numer, denom, m)
# transformer
class PreNorm(nn.Module):
def __init__(self, dim: int, fn: nn.Module, eps: float = 1e-8):
super().__init__()
self.norm = RMSNorm(dim, eps=eps)
self.fn = fn
def forward(self, x: Tensor) -> Tensor:
return self.fn(self.norm(x))
class CausalAttention(nn.Module):
def __init__(
self, dim: int, heads: int = 8, dim_head: int = 64, dropout: float = 0.0
):
super().__init__()
inner_dim = heads * dim_head
self.heads = heads
self.dim_head = dim_head
self.dropout = dropout
self.to_qkv = nn.Linear(dim, inner_dim * 3, bias=False)
self.to_out = nn.Linear(inner_dim, dim, bias=False)
def forward(self, x: Tensor) -> Tensor:
q, k, v = self.to_qkv(x).chunk(3, dim=-1)
def split_heads(y: Tensor) -> Tensor:
return rearrange(y, "b t (h d) -> b h t d", h=self.heads)
q, k, v = map(split_heads, (q, k, v))
out = F.scaled_dot_product_attention(
q, k, v, is_causal=True, dropout_p=self.dropout if self.training else 0.0
)
out = rearrange(out, "b h t d -> b t (h d)")
return self.to_out(out)
class SwiGLU(nn.Module):
def __init__(self, dim: int, mult: int = 4, dropout: float = 0.0):
# dropout not needed unless training on a smaller training data
super().__init__()
inner_dim = dim * mult
self.to_hidden = nn.Linear(dim, inner_dim * 2, bias=False)
self.to_out = nn.Linear(inner_dim, dim, bias=False)
self.dropout = nn.Dropout(dropout)
def forward(self, x: Tensor) -> Tensor:
gate, value = self.to_hidden(x).chunk(2, dim=-1)
x = F.silu(gate) * value
x = self.dropout(x)
return self.to_out(x)
class AttnResTransformer(nn.Module):
"""
Small GPT-style reference model using AttnRes
Using plain PyTorch: tok/pos embedding, alternating
causal attn, SwiGLU MLP layers, final norm, output head.
"""
def __init__(
self,
*,
num_tokens: int,
dim: int,
depth: int,
max_seq_len: int,
heads: int = 8,
dim_head: int = 64,
ff_mult: int = 4,
attn_dropout: float = 0.0,
ff_dropout: float = 0.0,
attnres: str = "block", # full or block
block_size: int = 8,
zero_init_queries: bool = True,
is_final_aggregate: bool = True,
eps: float = 1e-8,
):
super().__init__()
assert attnres in {"full", "block"}
self.max_seq_len = max_seq_len
self.attnres = attnres
self.token_emb = nn.Embedding(num_tokens, dim)
self.pos_emb = nn.Embedding(max_seq_len, dim)
atomic_layers = []
for _ in range(depth):
atomic_layers.append(
PreNorm(dim, CausalAttention(dim, heads, dim_head, attn_dropout), eps)
)
atomic_layers.append(PreNorm(dim, SwiGLU(dim, ff_mult, ff_dropout), eps))
if attnres == "full":
self.backbone = FullAttnResStack(
dim,
atomic_layers,
eps=eps,
zero_init_queries=zero_init_queries,
is_final_aggregate=is_final_aggregate,
)
else:
self.backbone = BlockAttnResStack(
dim,
atomic_layers,
block_size=block_size,
eps=eps,
zero_init_queries=zero_init_queries,
is_final_aggregate=is_final_aggregate,
)
self.final_norm = RMSNorm(dim, eps)
self.to_logits = nn.Linear(dim, num_tokens, bias=False)
def forward(self, ids: Tensor, schedule_block_size: int | None = None) -> Tensor:
b, t = ids.shape
assert t <= self.max_seq_len
pos = torch.arange(t, device=ids.device)
x = self.token_emb(ids) + self.pos_emb(pos)[None, :, :]
if self.attnres == "full":
x = self.backbone(x, schedule_block_size=schedule_block_size)
else:
x = self.backbone(x)
x = self.final_norm(x)
return self.to_logits(x)
__all__ = [
"ATTNRES_MODES",
"RMSNorm",
"DepthResidual",
"DepthResidualList",
"FullAttnResStack",
"BlockAttnResStack",
"BorrowedSubLayer",
"PreNorm",
"CausalAttention",
"SwiGLU",
"AttnResTransformer",
"atomic_block_size",
"validate_attnres",
]
+51
View File
@@ -0,0 +1,51 @@
"""Decoder block: x += attn(norm(x)); x += ffn(norm(x)).
attn/ffn are any modules with forward: [B,T,D] -> [B,T,D].
"""
from __future__ import annotations
from torch import nn
from .kda_attn import KDAAttention
from .latent_moe import LatentMoE
from .mla import GatedMLA
from .rmsnorm import RMSNorm
from .swiglu import SwiGLUMLP
def build_attn(config, kind: str) -> nn.Module:
if kind == "kda":
return KDAAttention.from_config(config)
if kind == "mla":
return GatedMLA.from_config(config)
raise ValueError(f"unknown attn kind: {kind}")
def build_ffn(config, kind: str) -> nn.Module:
if kind == "swiglu":
return SwiGLUMLP.from_config(config)
if kind == "moe":
return LatentMoE.from_config(config)
raise ValueError(f"unknown ffn kind: {kind}")
class DecoderBlock(nn.Module):
def __init__(self, hidden_size: int, norm_eps: float, attn: nn.Module, ffn: nn.Module):
super().__init__()
self.attn_norm = RMSNorm(hidden_size, norm_eps)
self.attn = attn
self.ffn_norm = RMSNorm(hidden_size, norm_eps)
self.ffn = ffn
@classmethod
def from_spec(cls, config, attn_kind: str, ffn_kind: str) -> DecoderBlock:
return cls(
config.hidden_size,
config.norm_eps,
build_attn(config, attn_kind),
build_ffn(config, ffn_kind),
)
def forward(self, x):
x = x + self.attn(self.attn_norm(x))
return x + self.ffn(self.ffn_norm(x))
+99
View File
@@ -0,0 +1,99 @@
"""KDA attention: project q/k/v/g/beta, run chunk_kda, project back to D."""
from __future__ import annotations
import torch
from torch import nn
from ..ops.api import chunk_kda
class KDAAttention(nn.Module):
"""Mixing module: x [B,T,D] -> y [B,T,D]."""
def __init__(
self,
hidden_size: int,
num_heads: int,
num_value_heads: int,
head_dim: int,
*,
chunk_size: int = 16,
initializer_range: float = 0.02,
use_gate_in_kernel: bool = True,
use_qk_l2norm_in_kernel: bool = True,
use_beta_sigmoid_in_kernel: bool = True,
lower_bound: float | None = -5.0,
kda_backend: str = "reference",
):
super().__init__()
if num_value_heads % num_heads:
raise ValueError("num_value_heads must be divisible by num_heads")
self.hidden_size = hidden_size
self.num_heads = num_heads
self.num_value_heads = num_value_heads
self.head_dim = head_dim
self.chunk_size = chunk_size
self.initializer_range = initializer_range
self.use_gate_in_kernel = use_gate_in_kernel
self.use_qk_l2norm_in_kernel = use_qk_l2norm_in_kernel
self.use_beta_sigmoid_in_kernel = use_beta_sigmoid_in_kernel
self.lower_bound = lower_bound
self.kda_backend = kda_backend
H, HV, K, V = num_heads, num_value_heads, head_dim, head_dim
self.q_proj = nn.Linear(hidden_size, H * K, bias=False)
self.k_proj = nn.Linear(hidden_size, H * K, bias=False)
self.v_proj = nn.Linear(hidden_size, HV * V, bias=False)
self.g_proj = nn.Linear(hidden_size, HV * K, bias=False)
self.beta_proj = nn.Linear(hidden_size, HV, bias=False)
self.o_proj = nn.Linear(HV * V, hidden_size, bias=False)
self.A_log = nn.Parameter(torch.zeros(HV))
# With safe_gate=-5, bias=-4 starts at g≈-0.09 (about 91% state retention).
self.dt_bias = nn.Parameter(torch.full((HV, K), -4.0))
self.apply(self._init_weights)
@classmethod
def from_config(cls, config) -> KDAAttention:
return cls(
hidden_size=config.hidden_size,
num_heads=config.num_heads,
num_value_heads=getattr(config, "num_value_heads", config.num_heads),
head_dim=config.head_dim,
chunk_size=config.chunk_size,
initializer_range=config.initializer_range,
use_gate_in_kernel=config.use_gate_in_kernel,
use_qk_l2norm_in_kernel=config.use_qk_l2norm_in_kernel,
use_beta_sigmoid_in_kernel=config.use_beta_sigmoid_in_kernel,
lower_bound=config.lower_bound,
kda_backend=config.kda_backend,
)
def _init_weights(self, module):
if isinstance(module, nn.Linear):
nn.init.normal_(module.weight, std=self.initializer_range)
def forward(self, x: torch.Tensor):
B, T, _ = x.shape
H, HV, K, V = self.num_heads, self.num_value_heads, self.head_dim, self.head_dim
q = self.q_proj(x).view(B, T, H, K)
k = self.k_proj(x).view(B, T, H, K)
v = self.v_proj(x).view(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,
A_log=self.A_log,
dt_bias=self.dt_bias,
use_qk_l2norm_in_kernel=self.use_qk_l2norm_in_kernel,
use_gate_in_kernel=self.use_gate_in_kernel,
use_beta_sigmoid_in_kernel=self.use_beta_sigmoid_in_kernel,
safe_gate=self.lower_bound is not None,
lower_bound=self.lower_bound,
chunk_size=self.chunk_size,
backend=self.kda_backend,
)
return self.o_proj(o.reshape(B, T, HV * V))
+118
View File
@@ -0,0 +1,118 @@
"""Stable LatentMoE (K3): shared 全宽 + routed 半宽专家 + SiTU-GLU + Top-k.
对照 learning/kimi-k3-notes §Stable LatentMoE:
z = W_down(x) [B, T, ℓ] ℓ = d/2 latent 接口宽
u = Σ_{i∈Top-k(x)} p_i E_i^rt(z) [B, T, ℓ] routed 专家只在 ℓ 上算
y = Σ_j E_j^sh(x) + W_up RMSNorm(u) [B, T, d] shared 全宽
SiTU-GLU: gate = β1·tanh(W_g x/β1)⊙σ(W_g x); up = β2·tanh(W_u x/β2)
||SiTU-GLU||_∞ ≤ β1·β2 (=100), 原点附近≈SwiGLU, 远端软饱和防低精度溢出.
E: R^in → R^in (内部中间维 d_ff).
Router: Top-k logits 基于全宽 x (笔记 Topk(x)); 归一化权重取 softmax(topk).
"""
from __future__ import annotations
import torch
import torch.nn.functional as F
from torch import nn
from .rmsnorm import RMSNorm
class SiTU(nn.Module):
"""SiTU-GLU expert: gate 支软上限 β1, up 支软上限 β2, 输出回到输入维."""
def __init__(self, dim_in: int, dim_ff: int, beta1: float = 4.0, beta2: float = 25.0):
super().__init__()
self.beta1, self.beta2 = beta1, beta2
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: torch.Tensor):
wg = self.w_g(x)
g = self.beta1 * torch.tanh(wg / self.beta1) * torch.sigmoid(wg)
u = self.beta2 * torch.tanh(self.w_u(x) / self.beta2)
return self.w_o(g * u)
class LatentMoE(nn.Module):
def __init__(
self,
hidden_size: int,
latent_size: int,
n_routed: int,
top_k: int,
n_shared: int,
d_ff: int,
beta1: float = 4.0,
beta2: float = 25.0,
):
super().__init__()
self.latent_size = latent_size
self.n_routed = n_routed
self.top_k = top_k
self.down = nn.Linear(hidden_size, latent_size, bias=False) # W↓
self.router = nn.Linear(hidden_size, n_routed, bias=False) # Top-k logits
self.shared = nn.ModuleList(
[SiTU(hidden_size, d_ff, beta1, beta2) for _ in range(n_shared)]
)
self.experts = nn.ModuleList(
[SiTU(latent_size, d_ff, beta1, beta2) for _ in range(n_routed)]
)
self.norm = RMSNorm(latent_size)
self.up = nn.Linear(latent_size, hidden_size, bias=False) # W↑
self.last_route_ids: torch.Tensor | None = None
@classmethod
def from_config(cls, config) -> LatentMoE:
return cls(
config.hidden_size,
config.moe_latent_size,
config.n_routed,
config.top_k,
config.n_shared,
config.moe_d_ff,
config.situ_beta1,
config.situ_beta2,
)
def forward(self, x: torch.Tensor):
B, T, _ = x.shape
z = self.down(x) # [B, T, ℓ]
logits = self.router(x) # [B, T, n_routed]
topk = torch.topk(logits, self.top_k, dim=-1)
ids = topk.indices # [B, T, k]
self.last_route_ids = ids.detach()
probs = F.softmax(topk.values, dim=-1) # [B, T, k]
# 向量化 routed: 预计算全部专家输出, 按 token 的 Top-k id 取
all_out = torch.stack([e(z) for e in self.experts]) # [R, B, T, ℓ]
all_out = all_out.permute(1, 2, 0, 3).reshape(B * T, self.n_routed, self.latent_size)
u = torch.zeros(B, T, self.latent_size, device=x.device, dtype=x.dtype)
for i in range(self.top_k):
idx = ids[:, :, i].reshape(B * T) # [B*T]
sel = all_out[torch.arange(B * T, device=x.device), idx] # [B*T, ℓ]
u += probs[:, :, i : i + 1] * sel.reshape(B, T, self.latent_size)
shared_out = torch.stack([e(x) for e in self.shared]).sum(0) # [B, T, d]
return shared_out + self.up(self.norm(u))
def moe_route_frac(model: nn.Module) -> torch.Tensor | None:
"""Mean expert occupancy over LatentMoE layers from the last forward."""
hists: list[torch.Tensor] = []
n_routed: int | None = None
for module in model.modules():
if not isinstance(module, LatentMoE) or module.last_route_ids is None:
continue
n_routed = module.n_routed
ids = module.last_route_ids.reshape(-1)
hists.append(torch.bincount(ids, minlength=n_routed).float())
if not hists or n_routed is None:
return None
stacked = torch.stack(hists).sum(0)
return stacked / stacked.sum().clamp_min(1.0)
+97
View File
@@ -0,0 +1,97 @@
"""Gated MLA (K3): NoPE, latent KV compression, matrix absorption, full-rank output gate.
K3 相对 DeepSeek MLA 的三个改动 (对照 learning/kimi-k3-notes):
1. NoPE — 不显式 RoPE; 位置感交给夹层 KDA 的 decay/gate。
2. 矩阵吸收 — 训练/推理都不解压 K/V: q 吸收 W_UK 后直接与 latent c 内积,
输出先在 latent 加权再乘 W_UV 还原 (v2 吸收版)。
3. Full-rank 输出门 — y = W_o[ σ(W_g x) ⊙ õ ]。
形状 (小规模 toy, d 为 hidden):
c = RMSNorm(kv_down(x)) [B, T, r] latent
q = q_up(RMSNorm(q_down(x))) [B, T, H, d_q] d_q = d_nope (NoPE)
W_UK = kv_up[.., :H*d_q].view(H,d_q,r) W_UV = kv_up[.., H*d_q:].view(H,d_v,r)
score = (q @ W_UK^T) @ c^T [B, H, T, T] causal
õ = (softmax(score) @ c) @ W_UV^T [B, T, H, d_v]
y = o_proj( σ(W_g x) ⊙ õ_head ) [B, T, d]
"""
from __future__ import annotations
import torch
import torch.nn.functional as F
from torch import nn
from .rmsnorm import RMSNorm
class GatedMLA(nn.Module):
def __init__(
self,
hidden_size: int,
num_heads: int,
kv_lora_rank: int,
q_lora_rank: int,
qk_nope_head_dim: int,
v_head_dim: int,
):
super().__init__()
self.hidden_size = hidden_size
self.num_heads = num_heads
self.qk_nope_head_dim = qk_nope_head_dim
self.v_head_dim = v_head_dim
# Q 低秩路径 (NoPE, 只有 nope 段)
self.q_down = nn.Linear(hidden_size, q_lora_rank, bias=False)
self.q_norm = RMSNorm(q_lora_rank)
self.q_up = nn.Linear(q_lora_rank, num_heads * qk_nope_head_dim, bias=False)
# KV latent 压缩 + 解压 (W_UK | W_UV 拼接在同一矩阵里)
self.kv_down = nn.Linear(hidden_size, kv_lora_rank, bias=False)
self.kv_norm = RMSNorm(kv_lora_rank)
self.kv_up = nn.Linear(
kv_lora_rank, num_heads * (qk_nope_head_dim + v_head_dim), bias=False
)
# Full-rank 输出门: σ(W_g x) 与 õ (H*d_v 维) 逐元素相乘
self.gate = nn.Linear(hidden_size, num_heads * v_head_dim, bias=False)
self.o_proj = nn.Linear(num_heads * v_head_dim, hidden_size, bias=False)
@classmethod
def from_config(cls, config) -> GatedMLA:
return cls(
config.hidden_size,
config.num_heads,
config.kv_lora_rank,
config.q_lora_rank,
config.qk_nope_head_dim,
config.v_head_dim,
)
def forward(self, x: torch.Tensor):
B, T, _ = x.shape
H, r = self.num_heads, self.kv_up.in_features
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]
w = self.kv_up.weight # [H*(d_q+d_v), r]
w_uk = w[: H * self.qk_nope_head_dim].view(H, self.qk_nope_head_dim, r)
w_uv = w[H * self.qk_nope_head_dim :].view(H, self.v_head_dim, r)
# 吸收 W_UK 进 query: score = (q @ W_UK^T) @ c^T
q_absorb = torch.einsum("bthd,hdj->bthj", q, w_uk) # [B, T, H, r]
scores = torch.einsum("bthj,bsj->bhts", q_absorb, c) # [B, H, T, T]
mask = torch.triu(
torch.ones(T, T, dtype=torch.bool, device=x.device), diagonal=1
)
scores = scores.masked_fill(mask, float("-inf"))
attn = F.softmax(scores, dim=-1) # [B, H, T, T]
# 先在 latent 加权, 再乘 W_UV^T 还原 v —— 永不解压
latent_out = torch.einsum("bhts,bsj->bhtj", attn, c) # [B, H, T, r]
o_heads = torch.einsum("bhtj,hvj->bhtv", latent_out, w_uv) # [B, H, T, d_v]
o_heads = o_heads.transpose(1, 2).reshape(B, T, H * self.v_head_dim)
gate = torch.sigmoid(self.gate(x)) # [B, T, H*d_v]
return self.o_proj(gate * o_heads) # [B, T, d]
+17
View File
@@ -0,0 +1,17 @@
"""RMSNorm used by attention, FFN, and the final LM stem."""
from __future__ import annotations
import torch
from torch import nn
class RMSNorm(nn.Module):
def __init__(self, dim: int, eps: float = 1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(dim))
self.eps = eps
def forward(self, x: torch.Tensor):
dtype = x.dtype
x = x.float()
return (x * torch.rsqrt(x.square().mean(-1, keepdim=True) + self.eps)).to(dtype) * self.weight
+20
View File
@@ -0,0 +1,20 @@
"""SwiGLU FFN: x [B,T,D] -> y [B,T,D]."""
from __future__ import annotations
import torch.nn.functional as F
from torch import nn
class SwiGLUMLP(nn.Module):
def __init__(self, hidden_size: int, intermediate_size: int):
super().__init__()
self.w1 = nn.Linear(hidden_size, intermediate_size, bias=False)
self.w3 = nn.Linear(hidden_size, intermediate_size, bias=False)
self.w2 = nn.Linear(intermediate_size, hidden_size, bias=False)
@classmethod
def from_config(cls, config) -> SwiGLUMLP:
return cls(config.hidden_size, config.intermediate_size)
def forward(self, x):
return self.w2(F.silu(self.w1(x)) * self.w3(x))
+7
View File
@@ -0,0 +1,7 @@
"""Configs and the single CausalLM entry."""
from .causal_lm import CausalLM
from .config import KDAConfig
from .k3_config import K3Config
__all__ = ["CausalLM", "K3Config", "KDAConfig"]
+123
View File
@@ -0,0 +1,123 @@
"""Causal LM stem: embed -> DecoderBlock* -> norm -> lm_head.
KDA-only and K3-like both use this class. Config.layer_specs() chooses
attn/ffn per layer: ("kda"|"mla", "swiglu"|"moe").
``config.attnres`` selects the depth mixer:
off — standard residual inside each DecoderBlock (default)
full — Full AttnRes over attn|ffn sublayers
block — Block AttnRes (K3); block size from ``attnres_block_size``
"""
from __future__ import annotations
import torch
import torch.nn.functional as F
from torch import nn
from torch.utils.checkpoint import checkpoint as activation_checkpoint
from ..layers.attn_res import (
BlockAttnResStack,
BorrowedSubLayer,
FullAttnResStack,
atomic_block_size,
)
from ..layers.block import DecoderBlock
from ..layers.rmsnorm import RMSNorm
def _build_mixer(config, blocks: nn.ModuleList):
mode = getattr(config, "attnres", "off")
if mode == "off":
return None
atomics = []
for block in blocks:
atomics.append(BorrowedSubLayer(block.attn_norm, block.attn))
atomics.append(BorrowedSubLayer(block.ffn_norm, block.ffn))
kwargs = dict(
eps=config.norm_eps,
zero_init_queries=getattr(config, "attnres_zero_init_queries", True),
is_final_aggregate=getattr(config, "attnres_final_aggregate", True),
)
if mode == "full":
return FullAttnResStack(config.hidden_size, atomics, **kwargs)
if mode == "block":
return BlockAttnResStack(
config.hidden_size,
atomics,
block_size=atomic_block_size(
config.num_hidden_layers, getattr(config, "attnres_block_size", None)
),
**kwargs,
)
raise ValueError(f"unknown attnres mode: {mode!r}")
class CausalLM(nn.Module):
def __init__(self, config):
super().__init__()
self.config = config
self.attnres = getattr(config, "attnres", "off")
self.embedding = nn.Embedding(config.vocab_size, config.hidden_size)
self.blocks = nn.ModuleList(
[
DecoderBlock.from_spec(config, attn, ffn)
for attn, ffn in config.layer_specs()
]
)
self.mixer = _build_mixer(config, self.blocks)
self.gradient_checkpointing = bool(
getattr(config, "gradient_checkpointing", False)
)
self.norm = RMSNorm(config.hidden_size, config.norm_eps)
self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
nn.init.normal_(self.embedding.weight, std=config.initializer_range)
nn.init.normal_(self.lm_head.weight, std=config.initializer_range)
if config.tie_word_embeddings:
self.lm_head.weight = self.embedding.weight
def forward(
self,
input_ids: torch.Tensor,
labels: torch.Tensor | None = None,
ignore_index: int = -100,
):
x = self.embedding(input_ids)
if self.mixer is None:
for block in self.blocks:
if self.gradient_checkpointing and self.training:
x = activation_checkpoint(block, x, use_reentrant=False)
else:
x = block(x)
elif self.gradient_checkpointing and self.training:
x = activation_checkpoint(self.mixer, x, use_reentrant=False)
else:
x = self.mixer(x)
logits = self.lm_head(self.norm(x))
if labels is None:
return logits
return F.cross_entropy(
logits[:, :-1].reshape(-1, logits.size(-1)),
labels[:, 1:].reshape(-1),
ignore_index=ignore_index,
)
@torch.inference_mode()
def generate(
self,
input_ids: torch.Tensor,
max_new_tokens: int,
temperature: float = 0.0,
eos_token_id: int | None = None,
):
for _ in range(max_new_tokens):
logits = self(input_ids)[:, -1]
if temperature > 0:
probs = F.softmax(logits / temperature, dim=-1)
next_token = torch.multinomial(probs, 1)
else:
next_token = logits.argmax(-1, keepdim=True)
input_ids = torch.cat((input_ids, next_token), dim=1)
if eos_token_id is not None and (next_token.squeeze(-1) == eos_token_id).all():
break
return input_ids
+62
View File
@@ -0,0 +1,62 @@
"""KDAConfig — toy Causal LM hyperparameters.
Defaults match the working reference-backend model: GVA with G=2,
safe gate (lower_bound=-5), q/k L2-norm and beta sigmoid inside the op.
"""
from __future__ import annotations
from dataclasses import dataclass
@dataclass
class KDAConfig:
hidden_size: int = 64
num_hidden_layers: int = 2
num_heads: int = 4
num_value_heads: int = 8 # G = num_value_heads // num_heads
head_dim: int = 16
chunk_size: int = 16
vocab_size: int = 256
intermediate_size: int = 128
max_position_embeddings: int = 128
initializer_range: float = 0.02
norm_eps: float = 1e-6
use_gate_in_kernel: bool = True
use_qk_l2norm_in_kernel: bool = True
use_beta_sigmoid_in_kernel: bool = True
lower_bound: float | None = -5.0
tie_word_embeddings: bool = False
kda_backend: str = "reference" # reference | triton | fla
attnres: str = "off" # off | full | block
attnres_block_size: int | None = None # DecoderBlocks / block; None ≈ L/8
attnres_zero_init_queries: bool = True
attnres_final_aggregate: bool = True
gradient_checkpointing: bool = False
@property
def H(self) -> int: return self.num_heads
@property
def G(self) -> int: return self.num_value_heads // self.num_heads
@property
def HV(self) -> int: return self.num_value_heads
@property
def K(self) -> int: return self.head_dim
@property
def V(self) -> int: return self.head_dim
def __post_init__(self):
from ..layers.attn_res import validate_attnres
if self.num_value_heads % self.num_heads:
raise ValueError("num_value_heads must be divisible by num_heads")
supported = {"reference", "triton", "fla", "torch", "auto"}
if self.kda_backend not in supported:
raise ValueError(f"kda_backend must be one of {sorted(supported)}")
validate_attnres(self.attnres, self.attnres_block_size)
def layer_specs(self) -> list[tuple[str, str]]:
return [("kda", "swiglu")] * self.num_hidden_layers
+128
View File
@@ -0,0 +1,128 @@
"""K3Config — Kimi K3 架构的小规模复现配置 (KDA + Gated MLA + Stable LatentMoE).
对照 learning/kimi-k3-notes §尺寸速查 (真实 K3 → 本 toy 缩比):
hidden 7168 → 256; L 93 → 4; H=HV 96 → 8; K=V 128 → 16;
MLA kv_lora 512 → 32, q_lora 1536 → 64, nope/v 128 → 16;
MoE ℓ=d/2=3584 → 128, 896/16 → 16/2, shared 2, d_ff 3072 → 96.
Hybrid Attention (K3): 每 4 层 1 次 Gated MLA, 末层强制 MLA.
Presets:
toy — ~8M, 自训 8k SP, 本地过拟合
0.5b — ~482M, Qwen3 词表, 32–40GB bf16;默认 step 是冒烟,翻译前置用 --max-tokens
"""
from __future__ import annotations
from dataclasses import dataclass
# Qwen3 config.json; train_k3 overrides with len(tokenizer).
QWEN3_VOCAB_SIZE = 151936
@dataclass
class K3Config:
# 主干
hidden_size: int = 256
num_hidden_layers: int = 4
vocab_size: int = 8192 # toy: data/spm_4k; 0.5b: Qwen3
initializer_range: float = 0.02
norm_eps: float = 1e-6
tie_word_embeddings: bool = False
max_position_embeddings: int = 2048 # NoPE, 仅语义保留
# KDA (K3: H = HV = 96, 无 GVA)
num_heads: int = 8
head_dim: int = 16
chunk_size: int = 16
lower_bound: float | None = -5.0
use_gate_in_kernel: bool = True
use_qk_l2norm_in_kernel: bool = True
use_beta_sigmoid_in_kernel: bool = True
# Gated MLA (NoPE)
kv_lora_rank: int = 32
q_lora_rank: int = 64
qk_nope_head_dim: int = 16
v_head_dim: int = 16
# Stable LatentMoE
moe_latent_size: int = 128 # ℓ = d/2
n_routed: int = 16
top_k: int = 2
n_shared: int = 2
moe_d_ff: int = 96
situ_beta1: float = 4.0
situ_beta2: float = 25.0
kda_backend: str = "reference"
# Depth mixer. off = DecoderBlock residual; block matches K3.
attnres: str = "off" # off | full | block
attnres_block_size: int | None = None # DecoderBlocks / AttnRes block; None ≈ L/8
attnres_zero_init_queries: bool = True
attnres_final_aggregate: bool = True
gradient_checkpointing: bool = False
def __post_init__(self):
from ..layers.attn_res import validate_attnres
validate_attnres(self.attnres, self.attnres_block_size)
@classmethod
def preset(cls, name: str) -> K3Config:
if name == "toy":
return cls()
if name in {"0.5b", "500m"}:
# H * head_dim == hidden. Routed 16: LatentMoE still runs every expert.
# ~482M with tied Qwen3 embeddings. 6×(3 KDA + 1 MLA).
return cls(
hidden_size=768,
num_hidden_layers=24,
vocab_size=QWEN3_VOCAB_SIZE,
tie_word_embeddings=True,
max_position_embeddings=2048,
num_heads=12,
head_dim=64,
chunk_size=64,
kv_lora_rank=192,
q_lora_rank=512,
qk_nope_head_dim=64,
v_head_dim=64,
moe_latent_size=384,
n_routed=16,
top_k=2,
n_shared=2,
moe_d_ff=512,
# The pure-PyTorch reference is far too slow at this size.
kda_backend="triton",
gradient_checkpointing=True,
)
raise ValueError(f"unknown preset: {name}")
@property
def H(self) -> int:
return self.num_heads
@property
def HV(self) -> int:
return self.num_heads
@property
def K(self) -> int:
return self.head_dim
@property
def V(self) -> int:
return self.head_dim
def layer_types(self) -> list[str]:
"""Hybrid pattern: 每 4 层 1 次 MLA (0-based 层 3,7,...), 末层强制 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"
return types
def layer_specs(self) -> list[tuple[str, str]]:
return [(kind, "moe") for kind in self.layer_types()]
+5
View File
@@ -0,0 +1,5 @@
"""KDA operator API and implementation backends."""
from .api import chunk_kda
__all__ = ["chunk_kda"]
+167
View File
@@ -0,0 +1,167 @@
"""Training-facing KDA operator with the same boundary as FLA's ``chunk_kda``."""
from __future__ import annotations
import warnings
from functools import lru_cache
import torch
import torch.nn.functional as F
from .reference.chunkwise import DECAY_BLOCK, _EXP_LIMIT, naive_chunk_kda
@lru_cache(maxsize=1)
def _fla_chunk_kda():
try:
from fla.ops.kda import chunk_kda
except ImportError:
return None
return chunk_kda
def _reference_chunk_size(T: int, requested: int) -> int:
size = min(T, requested)
while T % size:
size -= 1
return size
def chunk_kda(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
*,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
scale: float | None = None,
initial_state: torch.Tensor | None = None,
output_final_state: bool = False,
use_qk_l2norm_in_kernel: bool = False,
use_gate_in_kernel: bool = False,
use_beta_sigmoid_in_kernel: bool = False,
safe_gate: bool = False,
lower_bound: float | None = None,
chunk_size: int = 64,
backend: str = "reference",
):
"""Run an explicitly selected KDA implementation.
``reference`` and its legacy alias ``torch`` use this repository's
differentiable PyTorch implementation. ``triton`` uses the vendored
FLA NVIDIA Triton kernels in ``kda._fla`` (chunk_size 32 or 64,
CUDA). ``fla`` is reserved for explicit upstream parity runs.
"""
supported = {"reference", "triton", "fla", "torch", "auto"}
if backend not in supported:
raise ValueError(f"backend must be one of {sorted(supported)}")
if not use_qk_l2norm_in_kernel:
# Backend-independent: this is a property of the recurrence, not of any
# one implementation.
warnings.warn(
"use_qk_l2norm_in_kernel=False: KDA's chunkwise form assumes "
"||k||=1 so that I + tril(A_kk*beta) has a convergent Neumann "
"series. Unnormalised k makes the exact output grow like "
"||k||^chunk_size and can reach inf on any backend.",
RuntimeWarning,
stacklevel=2,
)
if backend == "auto":
warnings.warn(
"backend='auto' is deprecated and now selects the local reference backend; "
"use backend='fla' explicitly for upstream FLA",
DeprecationWarning,
stacklevel=2,
)
backend = "reference"
if backend == "torch":
backend = "reference"
if backend == "triton":
from .triton.chunk import chunk_kda as triton_chunk_kda
fla_chunk = 32 if chunk_size <= 32 else 64
return triton_chunk_kda(
q,
k,
v,
g,
beta,
scale=scale,
initial_state=initial_state,
output_final_state=output_final_state,
chunk_size=fla_chunk,
use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel,
use_gate_in_kernel=use_gate_in_kernel,
use_beta_sigmoid_in_kernel=use_beta_sigmoid_in_kernel,
A_log=A_log,
dt_bias=dt_bias,
safe_gate=safe_gate,
lower_bound=lower_bound,
)
if backend == "fla":
fused_op = _fla_chunk_kda()
if fused_op is None:
raise RuntimeError(
"backend='fla' requires a complete flash-linear-attention installation"
)
return fused_op(
q,
k,
v,
g,
beta,
A_log=A_log,
dt_bias=dt_bias,
scale=scale,
initial_state=initial_state,
output_final_state=output_final_state,
use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel,
use_gate_in_kernel=use_gate_in_kernel,
use_beta_sigmoid_in_kernel=use_beta_sigmoid_in_kernel,
safe_gate=safe_gate,
lower_bound=lower_bound,
chunk_size=32 if chunk_size <= 32 else 64,
)
if safe_gate and lower_bound is not None:
# _decayed_dot exponentiates at most DECAY_BLOCK steps of gate decay,
# and safe_gate bounds each step by |lower_bound|.
budget = DECAY_BLOCK * abs(lower_bound)
if budget > _EXP_LIMIT:
raise ValueError(
f"lower_bound={lower_bound} allows a gate span of {budget:.1f} "
f"per {DECAY_BLOCK}-row block, which overflows exp() "
f"(limit {_EXP_LIMIT:.1f}) and yields NaN. Use "
f"|lower_bound| < {_EXP_LIMIT / DECAY_BLOCK:.2f} or "
"backend='triton'."
)
if use_qk_l2norm_in_kernel:
q, k = F.normalize(q, dim=-1), F.normalize(k, dim=-1)
if use_beta_sigmoid_in_kernel:
beta = beta.sigmoid()
if use_gate_in_kernel:
if A_log is None:
raise ValueError("A_log is required when use_gate_in_kernel=True")
bias = 0 if dt_bias is None else dt_bias.view(g.shape[-2:])
gate_input = g + bias
rate = A_log.exp().view(1, 1, -1, 1)
if safe_gate:
if lower_bound is None:
raise ValueError("lower_bound is required when safe_gate=True")
g = lower_bound * torch.sigmoid(rate * gate_input)
else:
g = -rate * F.softplus(gate_input)
return naive_chunk_kda(
q,
k,
v,
g,
beta,
scale=scale,
initial_state=initial_state,
output_final_state=output_final_state,
chunk_size=_reference_chunk_size(q.shape[1], chunk_size),
)
+5
View File
@@ -0,0 +1,5 @@
"""Incremental recurrent KDA implementations and state containers."""
from .fused import KDAState, fused_recurrent_kda, fused_recurrent_kda_step
__all__ = ["KDAState", "fused_recurrent_kda", "fused_recurrent_kda_step"]
+73
View File
@@ -0,0 +1,73 @@
"""L6: FLA fused recurrent KDA decode with optional step cache."""
from __future__ import annotations
from dataclasses import dataclass
import torch
from kda._fla.ops.kda.fused_recurrent import fused_recurrent_kda as _fused_recurrent_kda
@dataclass
class KDAState:
"""Mutable recurrent state cache: ``S`` is ``[B, HV, K, V]``."""
S: torch.Tensor
pos: int = 0
def reset(self):
self.S.zero_()
self.pos = 0
def fused_recurrent_kda_step(
state: KDAState,
q_t: torch.Tensor,
k_t: torch.Tensor,
v_t: torch.Tensor,
g_t: torch.Tensor,
beta_t: torch.Tensor,
scale: float | None = None,
):
"""Single-token step. Inputs are ``[B, H|HV, ...]`` (no time dim)."""
o, ht = _fused_recurrent_kda(
q_t.unsqueeze(1),
k_t.unsqueeze(1),
v_t.unsqueeze(1),
g_t.unsqueeze(1),
beta_t.unsqueeze(1),
scale=scale,
initial_state=state.S,
output_final_state=True,
)
state.S = ht
state.pos += 1
return o.squeeze(1)
def fused_recurrent_kda(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
scale: float | None = None,
initial_state: torch.Tensor | None = None,
output_final_state: bool = False,
**kwargs,
):
return _fused_recurrent_kda(
q,
k,
v,
g,
beta,
scale=scale,
initial_state=initial_state,
output_final_state=output_final_state,
**kwargs,
)
__all__ = ["KDAState", "fused_recurrent_kda", "fused_recurrent_kda_step"]
+13
View File
@@ -0,0 +1,13 @@
"""Readable PyTorch implementations used as correctness references."""
from .chunkwise import naive_chunk_kda
from .gate import kda_gate_naive, kda_gate_reference
from .recurrent import naive_kda, naive_kda_fwd
__all__ = [
"kda_gate_naive",
"kda_gate_reference",
"naive_chunk_kda",
"naive_kda",
"naive_kda_fwd",
]
+155
View File
@@ -0,0 +1,155 @@
"""Pure-PyTorch chunked reference implementation of KDA."""
from __future__ import annotations
import math
import warnings
import torch
from einops import rearrange
#: ``exp`` overflows past this exponent in fp32 and bf16 (both top out at 3.4e38).
_EXP_LIMIT = math.log(torch.finfo(torch.float32).max)
#: Row-block size for :func:`_decayed_dot`.
#:
#: The g_ref GEMM exponentiates the gate span between the reference row and the
#: rows/columns it covers, so the block size caps that exponent at
#: ``DECAY_BLOCK * max|g|``. With the default ``lower_bound=-5`` gate that is
#: ``16 * 5 = 80 < ln(3.4e38) = 88.7``, i.e. fp32/bf16-safe for any chunk size.
#: Referencing a whole 64-row chunk instead would allow ``64 * 5 = 320`` and
#: overflow to NaN once the gate saturates.
DECAY_BLOCK = 16
def _decayed_dot(x: torch.Tensor, k: torch.Tensor, g: torch.Tensor) -> torch.Tensor:
"""Return ``A[..., i, j] = <x_i, exp(g_i-g_j) * k_j>`` (FLA g_ref GEMM).
Only the causal part (``j <= i``) is exact; callers mask the rest, which is
left at zero. Rows are processed in blocks of :data:`DECAY_BLOCK` against
the block's own first row, which is what bounds the exponent: for a row
block starting at ``r``, ``exp(g_i - g_ref)`` spans at most ``DECAY_BLOCK``
steps, and ``exp(g_ref - g_j)`` is ``<= 1`` for ``j < r`` and likewise spans
at most ``DECAY_BLOCK`` steps for ``j >= r``.
"""
C = g.shape[-2]
out = g.new_zeros(*g.shape[:-1], C)
for r in range(0, C, DECAY_BLOCK):
end = min(r + DECAY_BLOCK, C)
g_ref = g[..., r : r + 1, :]
rows = x[..., r:end, :] * (g[..., r:end, :] - g_ref).exp()
cols = k[..., :end, :] * (g_ref - g[..., :end, :]).exp()
out[..., r:end, :end] = rows @ cols.transpose(-1, -2)
return out
#: Whether :func:`naive_chunk_kda` checks the gate span against the ``exp``
#: budget. The check costs one device sync per call; set it to ``False`` if that
#: matters more than diagnosing a NaN.
CHECK_DECAY_SPAN = True
def _max_decay_span(g_cumsum: torch.Tensor) -> torch.Tensor:
"""Largest ``|g_ref - g_j|`` any row block will exponentiate."""
C = g_cumsum.shape[-2]
if C % DECAY_BLOCK == 0:
blocks = g_cumsum.unflatten(-2, (C // DECAY_BLOCK, DECAY_BLOCK))
return (blocks[..., :1, :] - blocks).abs().amax()
return torch.stack(
[
(g_cumsum[..., r : r + 1, :] - g_cumsum[..., r : r + DECAY_BLOCK, :])
.abs()
.amax()
for r in range(0, C, DECAY_BLOCK)
]
).amax()
def _warn_if_decay_span_overflows(g_cumsum: torch.Tensor) -> None:
"""Warn when a row block's gate span is about to overflow ``exp``.
``DECAY_BLOCK`` bounds this for the default ``safe_gate`` path, but an
unbounded gate (``-A.exp() * softplus(x)``) can still exceed it.
"""
span = _max_decay_span(g_cumsum).item()
if span > _EXP_LIMIT:
warnings.warn(
f"gate span within a {DECAY_BLOCK}-row block is {span:.1f} > "
f"{_EXP_LIMIT:.1f}; exp() will overflow to inf and the output will "
"be NaN. Reduce the gate magnitude (e.g. safe_gate with a smaller "
"|lower_bound|) or use backend='triton'.",
RuntimeWarning,
stacklevel=3,
)
def naive_chunk_kda(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
scale: float | None = None,
initial_state: torch.Tensor | None = None,
output_final_state: bool = False,
chunk_size: int = 64,
):
"""Chunk-parallel, inter-chunk recurrent KDA reference.
Shapes are ``q/k: [B,T,H,K]``, ``v: [B,T,HV,V]``,
``g: [B,T,HV,K]`` and ``beta: [B,T,HV]``.
"""
dtype = v.dtype
B, T, H, K = q.shape
HV, V = v.shape[2], v.shape[-1]
C = chunk_size
assert HV % H == 0, f"HV={HV} must be divisible by H={H}"
assert T % C == 0, f"T={T} must be divisible by chunk_size={C}"
scale = K**-0.5 if scale is None else scale
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)
if CHECK_DECAY_SPAN:
_warn_if_decay_span_overflows(g)
# r_i + sum_{j<i} beta_j <k_i, exp(g_i-g_j)k_j> r_j
# = v_i - <exp(g_i)k_i, S_start>.
mask_upper = torch.triu(torch.ones(C, C, dtype=torch.bool, device=q.device))
mask_strict_upper = torch.triu(mask_upper, diagonal=1)
eye = torch.eye(C, dtype=q.dtype, device=q.device)
A_kk = _decayed_dot(k, k, g)
M = eye + (A_kk * beta[..., None, :]).masked_fill(mask_upper, 0)
W = torch.linalg.solve_triangular(M, g.exp() * k, upper=False)
U = torch.linalg.solve_triangular(M, v, upper=False)
# Output includes the current token, hence the diagonal is retained.
A_qk = (_decayed_dot(q, k, g) * beta[..., None, :]).masked_fill(mask_strict_upper, 0)
S = q.new_zeros(B, HV, K, V)
if initial_state is not None:
S = S + initial_state
o = v.new_empty(B, HV, T // C, C, V)
for n in range(T // C):
q_n, k_n, g_n = q[:, :, n], k[:, :, n], g[:, :, n]
r = U[:, :, n] - W[:, :, n] @ S
o[:, :, n] = (q_n * g_n.exp()) @ S + A_qk[:, :, n] @ r
decay = (g_n[:, :, -1:, :] - g_n).exp()
S = S * g_n[:, :, -1, :, None].exp()
S = S + (decay * k_n).transpose(-1, -2) @ (r * beta[:, :, n, :, None])
if not output_final_state:
S = None
return rearrange(o, "b h n c d -> b (n c) h d").to(dtype), S
# Backward-compatible name used by earlier notes/scripts.
naive_chunk_kda_fwd = naive_chunk_kda
+47
View File
@@ -0,0 +1,47 @@
"""PyTorch references for the two KDA gate activations."""
from __future__ import annotations
import torch
import torch.nn.functional as F
def kda_gate_reference(
g: torch.Tensor,
A_log: torch.Tensor,
dt_bias: torch.Tensor | None = None,
*,
safe_gate: bool = False,
lower_bound: float | None = None,
) -> torch.Tensor:
"""Compute the official KDA gate semantics in PyTorch.
``A_log`` is head-wise with shape ``[HV]`` and ``dt_bias`` is
per-dimension with shape ``[HV, K]`` (or flattened to ``[HV*K]``).
"""
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:
if lower_bound is None:
raise ValueError("lower_bound is required when safe_gate=True")
return lower_bound * torch.sigmoid(rate * gate_input)
return -rate * F.softplus(gate_input)
def kda_gate_naive(
g: torch.Tensor,
A_log: torch.Tensor,
dt_bias: torch.Tensor | None = None,
lower_bound: float | None = None,
) -> torch.Tensor:
"""Compatibility name matching FLA's reference gate convention."""
return kda_gate_reference(
g,
A_log,
dt_bias,
safe_gate=lower_bound is not None,
lower_bound=lower_bound,
)
__all__ = ["kda_gate_naive", "kda_gate_reference"]
+298
View File
@@ -0,0 +1,298 @@
"""L1: Naive recurrent KDA fwd+bwd (torch only).
公式 (per timestep t, log-space gate; q/k 入口 H 维, 内部 repeat_interleave 到 HV):
S_t = exp(g_t) * S_{t-1} + (beta_t * k_t) outer (v_t - k_t . (exp(g_t) * S_{t-1}))
o_t = (q_t * scale) . S_t
backward (BPTT, T -> 0):
设 dS_t 为进入 t 步累积的反传梯度 (含 o_t 反传).
1. o_t = q_t . S_t -> dS_t += q_t outer do_t (i.e. dS = dS + q_t·do_t)
dq_t = do_t . S_t^T -> einsum('bhv,bhkv->bhk')
2. S_t = S_decay + a_t outer r_t, a_t = b_t k_t, r_t = v_t - k_t . S_decay
其中 S_decay = exp(g_t) * S_{t-1}
dS_{t-1} = exp(g_t) * (dS_t - r_t outer da_t - a_t outer dr_t) via residual 反传
更具体:
dS_decay = dS_t - (a_t outer dr_t) - (da_t outer r_t)
dS_{t-1} += exp(g_t) * dS_decay
其中 dr_t = -dv_t + dS_t . a_t^T (因为 r_t = v - k·S_dec, dr 来自 -dv - k·dS_decay)
da_t = -r_t outer dS_t? 让我直接推导下面.
推导 (设 G1 = S_t, 走 a = r 反向链 通过 autograd):
o_t = q_t . G1
dq_t = do_t . G1^T -> [B,HV,K]
dG1 = q_t outer do_t -> [B,HV,K,V] = dS_t (上游)
G1 = Sdec + a outer r -> Sdec = G1[...] (跳过)
dSdec = dG1
da_t = r_t outer dG1 -> [B,HV,K] (因为 a outer r 是 K-V, d(a outer r) = r outer d[...,V])
但在 einsum 表示: dA_t.grad = einsum('bhkv,bhv->bhk', dS_t, r_t)
dr_t = a_t outer dG1 -> [B,HV,V] = einsum('bhkv,bhk->bhv', dS_t, a_t)
plus: S_t = Sdec + a outer r -> a outer r - outer product 形状是 [B,HV,K,V] = einsum('bhk,bhv->bhkv')
d(a outer r) 的雅可比: let G1_m = a_t ⊗ r_t (rank-1 matrix per (b,h))
dG1_m[i,j] = da_t[i] * r_t[j] + a_t[i] * dr_t[j]
在外积形式, 即 dG1_m = a_outer r 的张量积正交分解:
da_t = sum_j r_t[j] dG1_m[i,j] = einsum('bhkv,bhv->bhk', dG1_m, r_t)
dr_t = sum_i a_t[i] dG1_m[i,j] = einsum('bhkv,bhk->bhv', dG1_m, a_t)
因为 a_t = b_t k_t -> da_t = db_t k_t + b_t dk_t (b_t 是 ...)
db_t = einsum('bhk,bhk->bh', da_t, k_t)
dk_t_a = b_t * da_t (来自 a_t 路径, 还有来自 r_t 路径和 S_dec 路径)
因为 r_t = v_t - k_t . S_dec -> 注 rk_t grad via dg,S_dec 和 dv_t
dv_t = -dr_t (实际 dr 的负梯度) 即 dv_t = -dr_t
这里 r_t = v_t - k_t · S_dec, 写作矩阵乘 r = v - einsum('bhk,bhkv->bhv', k, S_dec)
dr = -dv - einsum('bhk,bhkv->bhv', dk_from_r, S_dec) + eigengrad via S_dec
更精确的反向: r_t = v_t - k_t . S_dec
dv_t += -dr_t -> dv_t = -dr_t
dk_t_r_path = -S_dec outer dr_t (即 -dS_dec 传递来自 k_t 的部分)
具体: d(k·S) = dk·S + k·dS -> dS_dec 这层, dk 的贡献: -S_dec outer dr_t
即 dk_t_r = einsum('bhv,bhkv->bhk', -dr_t, S_dec)
dS_dec_r = -k_t outer dr_t = -einsum('bhv,bhk->bhkv', dr_t, k_t)
合并: dS_dec 合总 = dG1 + (-k_t outer dr_t)
= dS_t - k_t outer dr_t
(相加过的 dv, dk_r, dS_dec_r 都上面项)
Sdec = exp(g_t) * S_{t-1}:
dS_{t-1} = exp(g_t) ⊙ dS_dec (因为 Sdec = exp_g * S_prev, 微分后 exp_g 直接相乘)
dg_t = exp(g_t) * S_prev * dS_dec (微分时对 g_t (log-space) 求偏导数)
即 dg_t = exp(g_t) * (S_{t-1} ⊙ dS_dec) -> 沿 K 维求和
in einsum: dg_t = sum over v of (exp(g_t) * S_{t-1}) ⊙ dS_dec ...\n
= einsum('bhk, bhk, bhkv -> bhk', exp_g, S_prev, dS_dec)
更简洁: Sdec = exp_g * S_prev (per-(b,h,k)/v), 故 dSdec/dg_t = S_prev * exp_g
所以 dg_t = sum_v S_prev_sub_k_dim * exp_g * dS_dec -> [B, HV, K]
einsum: dg_t = einsum('bhkv,bhkv->bhk', Sdec, dS_dec)
(因为 Sdec = S_prev * exp_g, sum_v Sdec[:, :, :, v] * dS_dec[:, :, :, v] = sum_v Sdec_eachK * dSdec_eachK)
einsum上是 einsum('bhkv,bhkv->bhk', Sdec, dSdec)
dS_{t-1} = exp_g ⊙ dSdec (per (b,h,k,v) entrywise multiply exp_g with dSdec)
GVA 反归约:
q,k 入口 [B, T, H, K] --repeat_interleave(G, dim=2)--> [B, T, HV, K]
内部计算后, dq/dk 在 HV 维上 -> dV 拿 shape [B,T,HV,K]
bwd 通过 sum 回 H: dq_H = dq_HV.view(B,T,H,G,K).sum(dim=3) -> [B,T,H,K]
(因为 repeat_interleave 是复制, 反传是 sum 路径相同意义)
记号对照:
a_t = b_t * k_t (a = beta * k) [B, HV, K]
r_t = v_t - k_t . S_dec (residual) [B, HV, V]
S_dec = exp(g_t) * S_{t-1} [B, HV, K, V]
S_t = S_dec + a_t outer r_t [B, HV, K, V]
o_t = q_t . S_t = (q_t_eff * scale) . S_t [B, HV, V]
"""
from __future__ import annotations
import math
import torch
def naive_kda_fwd(
q: torch.Tensor, # [B, T, H, K]
k: torch.Tensor, # [B, T, H, K]
v: torch.Tensor, # [B, T, HV, V]
g: torch.Tensor, # [B, T, HV, K]
beta: torch.Tensor, # [B, T, HV]
scale: float | None = None,
initial_state: torch.Tensor | None = None, # [B, HV, K, V]
output_final_state: bool = False,
*,
force_float32: bool = False,
):
"""纯 forward, 不带 autograd. 与上游 naive_recurrent_kda 数值等价.
force_float32=True 时强制 fp32 计算 (与上游对拍时用);
默认保持输入 dtype (gradcheck 用 fp64).
"""
dtype = v.dtype
B, T, H, K = q.shape
HV, V = v.shape[2], v.shape[-1]
G = HV // H
if scale is None:
scale = 1.0 / math.sqrt(K)
# 上游强制 fp32; 本实现默认保留输入 dtype 以便 gradcheck 适用 fp64
# force_float32=True 时与上游逐位对齐
work_dtype = torch.float if force_float32 else q.dtype
q = q.to(work_dtype)
k = k.to(work_dtype)
v = v.to(work_dtype)
g = g.to(work_dtype)
beta = beta.to(work_dtype)
# GVA: expand q/k from H to HV
qe = q.repeat_interleave(G, dim=2) * scale # [B, T, HV, K]
ke = k.repeat_interleave(G, dim=2) # [B, T, HV, K]
S = torch.zeros(B, HV, K, V, dtype=work_dtype, device=q.device)
if initial_state is not None:
S = S + initial_state.to(work_dtype)
o = torch.empty(B, T, HV, V, dtype=work_dtype, device=q.device)
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]
S_dec = S * g_t.exp().unsqueeze(-1) # [B, HV, K, V]
p_t = torch.einsum('b h k, b h k v -> b h v', k_t, S_dec) # [B, HV, V]
r_t = v_t - p_t # [B, HV, V]
a_t = b_t.unsqueeze(-1) * k_t # [B, HV, K]
S = S_dec + torch.einsum('b h k, b h v -> b h k v', a_t, r_t)
o[:, t] = torch.einsum('b h k, b h k v -> b h v', q_t, S)
if not output_final_state:
S = None
return o.to(dtype), S
class KDAFunction(torch.autograd.Function):
"""autograd Function (forward + backward).
forward 入参顺序 (q, k, v, g, beta, scale, initial_state, output_final_state)
backward 必须返回一致: (dq, dk, dv, dg, dbeta, None, dinit_state, None)
"""
@staticmethod
def forward(ctx, q, k, v, g, beta, scale, initial_state, output_final_state):
dtype = v.dtype
B, T, H, K = q.shape
HV, V = v.shape[2], v.shape[-1]
G = HV // H
if scale is None:
scale = 1.0 / math.sqrt(K)
work_dtype = q.dtype
qf = q.to(work_dtype).contiguous()
kf = k.to(work_dtype).contiguous()
vf = v.to(work_dtype).contiguous()
gf = g.to(work_dtype).contiguous()
bf = beta.to(work_dtype).contiguous()
# GVA: expand q/k from H to HV
qe = qf.repeat_interleave(G, dim=2) * scale # [B, T, HV, K]
ke = kf.repeat_interleave(G, dim=2) # [B, T, HV, K]
S = torch.zeros(B, HV, K, V, dtype=work_dtype, device=q.device)
if initial_state is not None:
S = S + initial_state.to(work_dtype)
o = torch.empty(B, T, HV, V, dtype=work_dtype, device=q.device)
q_ts, k_ts, b_ts, S_decs, r_ts, a_ts, exp_g_ts = [], [], [], [], [], [], []
for t in range(T):
q_t = qe[:, t]
k_t = ke[:, t]
v_t = vf[:, t]
g_t = gf[:, t]
b_t = bf[:, t]
exp_g_t = g_t.exp()
S_dec = S * exp_g_t.unsqueeze(-1)
p_t = torch.einsum('b h k, b h k v -> b h v', k_t, S_dec)
r_t = v_t - p_t
a_t = b_t.unsqueeze(-1) * k_t
S = S_dec + torch.einsum('b h k, b h v -> b h k v', a_t, r_t)
o[:, t] = torch.einsum('b h k, b h k v -> b h v', q_t, S)
q_ts.append(q_t)
k_ts.append(k_t)
b_ts.append(b_t)
S_decs.append(S_dec)
r_ts.append(r_t)
a_ts.append(a_t)
exp_g_ts.append(exp_g_t)
ctx.save_for_backward(
torch.stack(q_ts, dim=1),
torch.stack(k_ts, dim=1),
torch.stack(b_ts, dim=1),
torch.stack(S_decs, dim=1),
torch.stack(r_ts, dim=1),
torch.stack(a_ts, dim=1),
torch.stack(exp_g_ts, dim=1),
)
ctx.G = G
ctx.H = H
ctx.HV = HV
ctx.K = K
ctx.V = V
ctx.T = T
ctx.B = B
ctx.scale = scale
ctx.dtype = dtype
ctx.has_initial_state = initial_state is not None
ctx.output_final_state = output_final_state
final_S = S if output_final_state else None
return o.to(dtype), final_S
@staticmethod
def backward(ctx, do, dS):
q_ts, k_ts, b_ts, S_decs, r_ts, a_ts, exp_g_ts = ctx.saved_tensors
B, T, H, HV, K, V, G = ctx.B, ctx.T, ctx.H, ctx.HV, ctx.K, ctx.V, ctx.G
work_dtype = q_ts.dtype
device = q_ts.device
dq_e = torch.zeros(B, T, HV, K, dtype=work_dtype, device=device)
dk_e = torch.zeros(B, T, HV, K, dtype=work_dtype, device=device)
dv = torch.zeros(B, T, HV, V, dtype=work_dtype, device=device)
dg = torch.zeros(B, T, HV, K, dtype=work_dtype, device=device)
dbeta= torch.zeros(B, T, HV, dtype=work_dtype, device=device)
if dS is None:
dS_acc = torch.zeros(B, HV, K, V, dtype=work_dtype, device=device)
else:
dS_acc = dS.to(work_dtype).clone()
for t in range(T - 1, -1, -1):
q_t = q_ts[:, t]
k_t = k_ts[:, t]
b_t = b_ts[:, t]
S_dec = S_decs[:, t]
r_t = r_ts[:, t]
a_t = a_ts[:, t]
exp_g_t = exp_g_ts[:, t]
do_t = do[:, t].to(work_dtype)
S_t = S_dec + torch.einsum('b h k, b h v -> b h k v', a_t, r_t)
dS_acc = dS_acc + torch.einsum('b h k, b h v -> b h k v', q_t, do_t)
dq_e[:, t] = torch.einsum('b h v, b h k v -> b h k', do_t, S_t)
da_t = torch.einsum('b h v, b h k v -> b h k', r_t, dS_acc)
dr_t = torch.einsum('b h k, b h k v -> b h v', a_t, dS_acc)
dbeta[:, t] = torch.einsum('b h k, b h k -> b h', k_t, da_t)
dk_t_a = b_t.unsqueeze(-1) * da_t
dv[:, t] = dr_t
dS_dec_from_r = -torch.einsum('b h v, b h k -> b h k v', dr_t, k_t)
dk_t_r = -torch.einsum('b h v, b h k v -> b h k', dr_t, S_dec)
dS_dec_total = dS_acc + dS_dec_from_r
dk_e[:, t] = dk_t_a + dk_t_r
dg[:, t] = torch.einsum('b h k v, b h k v -> b h k', S_dec, dS_dec_total)
dS_acc = exp_g_t.unsqueeze(-1) * dS_dec_total
if HV > H:
dq_H = dq_e.view(B, T, H, G, K).sum(dim=3)
dk_H = dk_e.view(B, T, H, G, K).sum(dim=3)
else:
dq_H = dq_e
dk_H = dk_e
# q 在 forward 内被乘过 scale (qe = q * scale), chain rule: dq_orig = dq_e * scale
dq_H = dq_H * ctx.scale
return (dq_H.to(ctx.dtype), dk_H.to(ctx.dtype), dv.to(ctx.dtype),
dg.to(ctx.dtype), dbeta.to(ctx.dtype), None, None, None)
def naive_kda(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
scale: float | None = None,
initial_state: torch.Tensor | None = None,
output_final_state: bool = False,
):
"""对外入口: 调 KDAFunction.apply."""
return KDAFunction.apply(q, k, v, g, beta, scale, initial_state, output_final_state)
+7
View File
@@ -0,0 +1,7 @@
"""Local Triton KDA kernels vendored from FLA chunk_{fwd,intra,bwd,wy,gate}."""
from .chunk import ChunkKDAFunction, chunk_kda
from .chunk_fwd import chunk_kda_fwd
from .gate import kda_gate_fwd
__all__ = ["ChunkKDAFunction", "chunk_kda", "chunk_kda_fwd", "kda_gate_fwd"]
+5
View File
@@ -0,0 +1,5 @@
"""FLA ``chunk_kda`` surface used by ``ops.api`` backend='triton'."""
from kda._fla.ops.kda.chunk import ChunkKDAFunction, chunk_kda
__all__ = ["ChunkKDAFunction", "chunk_kda"]
+5
View File
@@ -0,0 +1,5 @@
"""Vendored FLA chunk KDA backward."""
from kda._fla.ops.kda.chunk_bwd import chunk_kda_bwd
__all__ = ["chunk_kda_bwd"]
+37
View File
@@ -0,0 +1,37 @@
"""Vendored FLA chunk KDA forward, returning ``(o, ht)`` like the public op."""
from __future__ import annotations
import torch
from kda._fla.ops.kda.chunk import chunk_kda
from kda._fla.ops.kda.chunk_fwd import chunk_kda_fwd as fla_chunk_kda_fwd
__all__ = ["chunk_kda_fwd", "fla_chunk_kda_fwd"]
def chunk_kda_fwd(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
scale: float | None = None,
initial_state: torch.Tensor | None = None,
output_final_state: bool = False,
chunk_size: int = 64,
**kwargs,
):
"""Chunked KDA forward with FLA kernels. Returns ``(o, ht)``."""
return chunk_kda(
q,
k,
v,
g,
beta,
scale=scale,
initial_state=initial_state,
output_final_state=output_final_state,
chunk_size=chunk_size,
**kwargs,
)
+36
View File
@@ -0,0 +1,36 @@
"""Vendored FLA KDA gate fusion (standard + safe gate + chunk cumsum)."""
from __future__ import annotations
import torch
from kda._fla.ops.kda.gate import (
kda_gate_bwd,
kda_gate_chunk_cumsum,
kda_gate_fwd as _kda_gate_fwd,
)
DEFAULT_LOWER_BOUND = -5.0
def kda_gate_fwd(
g: torch.Tensor,
A_log: torch.Tensor,
dt_bias: torch.Tensor | None = None,
lower_bound: float | None = DEFAULT_LOWER_BOUND,
):
return _kda_gate_fwd(
g,
A_log=A_log,
dt_bias=dt_bias,
lower_bound=lower_bound,
output_dtype=g.dtype,
)
__all__ = [
"DEFAULT_LOWER_BOUND",
"kda_gate_bwd",
"kda_gate_chunk_cumsum",
"kda_gate_fwd",
]
+5
View File
@@ -0,0 +1,5 @@
"""Vendored FLA WY recompute used by the chunk KDA backward."""
from kda._fla.ops.kda.wy_fast import recompute_w_u_fwd
__all__ = ["recompute_w_u_fwd"]
+5
View File
@@ -0,0 +1,5 @@
"""Training and checkpoint helpers."""
from .toy import load_ckpt, make_toy_data, save_ckpt, train_one_batch
__all__ = ["load_ckpt", "make_toy_data", "save_ckpt", "train_one_batch"]
+355
View File
@@ -0,0 +1,355 @@
"""Pretrain / SFT sample construction.
Pretrain: Wikipedia parquet → tokenize → pack (B, T). Languages mix 1:1 by
token via seq_len-sized blocks so each training chunk is monolingual.
SFT: instruction-parallel rows → prompt-masked labels. Template lives in
``prompts.instruction_prompt`` (same string as eval_mt).
"""
from __future__ import annotations
import json
import os
from dataclasses import dataclass
from pathlib import Path
from typing import Iterable, Protocol
import torch
from .prompts import instruction_prompt
WIKI_SHARD_TOTAL = {"zh": 6, "en": 41}
WIKI_BASE = (
"https://huggingface.co/datasets/wikimedia/wikipedia/resolve/main/20231101.{lang}"
)
IGNORE_INDEX = -100
class Tokenizer(Protocol):
vocab_size: int
def encode(self, text: str) -> list[int]: ...
def decode(self, ids: list[int]) -> str: ...
@dataclass
class SentencePieceTokenizer:
_sp: object
@property
def vocab_size(self) -> int:
return int(self._sp.vocab_size())
def encode(self, text: str) -> list[int]:
return list(self._sp.encode(text, out_type=int))
def decode(self, ids: list[int]) -> str:
return str(self._sp.decode(ids))
@dataclass
class HuggingFaceTokenizer:
_tok: object
@property
def vocab_size(self) -> int:
return int(len(self._tok))
def encode(self, text: str) -> list[int]:
return list(self._tok.encode(text, add_special_tokens=False))
def decode(self, ids: list[int]) -> str:
return str(self._tok.decode(ids, skip_special_tokens=True))
def load_tokenizer(source: str) -> Tokenizer:
"""`.model` 走 SentencePiece, 其它当作 HuggingFace 名或本地目录."""
if source.endswith(".model"):
from sentencepiece import SentencePieceProcessor
return SentencePieceTokenizer(SentencePieceProcessor(model_file=source))
from transformers import AutoTokenizer
tok = AutoTokenizer.from_pretrained(source, trust_remote_code=True)
return HuggingFaceTokenizer(tok)
def pretrain_dir() -> Path:
for candidate in (
os.environ.get("KDA_PRETRAIN_DIR"),
"/data/pretrain",
"data/pretrain",
):
if candidate and Path(candidate).is_dir():
return Path(candidate)
return Path("data/pretrain")
def _wiki_files(lang: str, n_shards: int) -> list[str]:
if lang not in WIKI_SHARD_TOTAL:
raise ValueError(f"unsupported wiki lang {lang!r}; expected zh or en")
total = WIKI_SHARD_TOTAL[lang]
n = min(max(n_shards, 1), total)
base = WIKI_BASE.format(lang=lang)
return [f"{base}/train-{i:05d}-of-{total:05d}.parquet" for i in range(n)]
def _cache_path(cache_dir: Path, lang: str, n_shards: int, limit: int) -> Path:
return cache_dir / f"wiki-{lang}-n{n_shards}-limit{limit}.jsonl"
def fetch_wiki_texts(
limit: int,
lang: str = "zh",
n_shards: int = 2,
cache_dir: str | Path | None = None,
) -> list[str]:
"""Load up to ``limit`` article bodies, caching jsonl under pretrain_dir."""
cache = Path(cache_dir) if cache_dir is not None else pretrain_dir()
cache.mkdir(parents=True, exist_ok=True)
path = _cache_path(cache, lang, n_shards, limit)
if path.exists():
texts: list[str] = []
with path.open(encoding="utf-8") as fh:
for line in fh:
line = line.strip()
if not line:
continue
texts.append(json.loads(line)["text"])
if len(texts) >= limit:
break
if texts:
return texts
from datasets import load_dataset
files = _wiki_files(lang, n_shards)
ds = load_dataset("parquet", data_files=files, split="train", streaming=True)
texts = []
for i, row in enumerate(ds):
if i >= limit:
break
texts.append(row["text"])
tmp = path.with_suffix(path.suffix + ".tmp")
with tmp.open("w", encoding="utf-8") as fh:
for text in texts:
fh.write(json.dumps({"text": text}, ensure_ascii=False) + "\n")
tmp.replace(path)
return texts
def tokenize_corpus(texts: list[str], tok: Tokenizer) -> list[int]:
ids: list[int] = []
for text in texts:
ids.extend(tok.encode(text))
return ids
def interleave_balanced(ids_a: list[int], ids_b: list[int], block: int) -> list[int]:
"""1:1 by token: seq_len-sized monolingual blocks, drop the longer tail."""
if block < 1:
raise ValueError(f"block must be >= 1, got {block}")
n = min(len(ids_a), len(ids_b))
n = (n // block) * block
out: list[int] = []
a, b = ids_a, ids_b
for i in range(0, n, block):
out.extend(a[i : i + block])
out.extend(b[i : i + block])
return out
def chunk_ids(ids: list[int], batch: int, seq_len: int) -> torch.Tensor:
"""切成 (num_chunks, B, T); 末尾不足部分丢弃."""
n = (len(ids) // (batch * seq_len)) * (batch * seq_len)
t = torch.tensor(ids[:n], dtype=torch.long)
if n == 0:
return t.view(0, batch, seq_len)
return t.view(batch, -1, seq_len).transpose(0, 1)
def split_heldout(
chunks: torch.Tensor,
frac: float = 0.01,
min_heldout: int = 1,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Last ``frac`` of packed chunks for CE only. Empty held-out if too few."""
n = int(chunks.size(0))
if n <= 1 or frac <= 0:
return chunks, chunks[:0]
h = max(min_heldout, int(n * frac))
h = min(h, n - 1)
return chunks[:-h], chunks[-h:]
def iter_chunks(chunks: torch.Tensor):
"""逐块产出 (input_ids, labels), labels 右移 (模型内 CE shift)."""
for chunk in chunks:
yield chunk, chunk.clone()
def iter_indexed(chunks: torch.Tensor, start: int = 0):
"""Infinite cycle with a global index (for --resume)."""
n = int(chunks.size(0))
if n == 0:
raise ValueError("no training chunks")
i = start
while True:
x = chunks[i % n]
yield i, x, x.clone()
i += 1
def load_pretrain_chunks(
tok: Tokenizer,
*,
langs: Iterable[str],
limit: int,
batch: int,
seq_len: int,
heldout_frac: float = 0.01,
n_shards: int = 2,
cache_dir: str | Path | None = None,
) -> tuple[torch.Tensor, torch.Tensor, int]:
"""Fetch / cache / tokenize / pack. Returns train chunks, held-out, token count."""
lang_list = [lang.strip() for lang in langs if lang.strip()]
if not lang_list:
raise ValueError("langs must contain at least one of zh, en")
streams: list[list[int]] = []
for lang in lang_list:
print(f"loading {limit} wiki articles ({lang}) ...")
texts = fetch_wiki_texts(limit, lang=lang, n_shards=n_shards, cache_dir=cache_dir)
streams.append(tokenize_corpus(texts, tok))
print(f" {lang}: {len(streams[-1]):,} tokens from {len(texts)} articles")
if len(streams) == 1:
ids = streams[0]
else:
ids = streams[0]
for extra in streams[1:]:
ids = interleave_balanced(ids, extra, seq_len)
chunks = chunk_ids(ids, batch, seq_len)
train, held = split_heldout(chunks, heldout_frac)
return train, held, len(ids)
def pad_id(tok: Tokenizer) -> int:
inner = getattr(tok, "_tok", None)
if inner is not None:
pid = getattr(inner, "pad_token_id", None)
if pid is not None:
return int(pid)
eid = getattr(inner, "eos_token_id", None)
if eid is not None:
return int(eid)
return 0
def eos_id(tok: Tokenizer) -> int | None:
inner = getattr(tok, "_tok", None)
if inner is not None:
eid = getattr(inner, "eos_token_id", None)
if eid is not None:
return int(eid)
convert = getattr(inner, "convert_tokens_to_ids", None)
if convert is not None:
tid = convert("<|im_end|>")
if isinstance(tid, int) and tid >= 0:
return tid
return None
def encode_sft_row(
tok: Tokenizer,
src: str,
tgt: str,
target_lang: str,
max_len: int,
eos: int | None = None,
) -> tuple[list[int], list[int]]:
prompt_ids = tok.encode(instruction_prompt(src, target_lang))
tgt_ids = tok.encode(tgt)
if eos is not None:
tgt_ids = tgt_ids + [eos]
ids = prompt_ids + tgt_ids
labels = [IGNORE_INDEX] * len(prompt_ids) + list(tgt_ids)
if len(ids) > max_len:
overflow = len(ids) - max_len
cut = min(overflow, max(len(prompt_ids) - 1, 0))
ids = ids[cut:]
labels = labels[cut:]
if len(ids) > max_len:
ids = ids[:max_len]
labels = labels[:max_len]
return ids, labels
def load_sft_rows(path: str | Path) -> list[dict]:
"""jsonl ``{src,tgt,target_lang}`` or TSV ``src\\ttgt\\ttarget_lang``."""
p = Path(path)
rows: list[dict] = []
text = p.read_text(encoding="utf-8")
if p.suffix == ".jsonl" or p.suffix == ".json":
for line in text.splitlines():
line = line.strip()
if not line:
continue
obj = json.loads(line)
rows.append(
{
"src": obj["src"],
"tgt": obj["tgt"],
"target_lang": obj.get("target_lang", "en"),
}
)
return rows
for line in text.splitlines():
line = line.strip()
if not line or line.startswith("#"):
continue
parts = line.split("\t")
if len(parts) < 2:
raise ValueError(f"SFT TSV needs src, tgt [, target_lang]: {line[:80]!r}")
lang = parts[2] if len(parts) > 2 else "en"
rows.append({"src": parts[0], "tgt": parts[1], "target_lang": lang})
return rows
def collate_sft(
rows: list[dict],
tok: Tokenizer,
max_len: int,
) -> tuple[torch.Tensor, torch.Tensor]:
pad = pad_id(tok)
eos = eos_id(tok)
encoded = [
encode_sft_row(tok, r["src"], r["tgt"], r["target_lang"], max_len, eos)
for r in rows
]
width = min(max(len(ids) for ids, _ in encoded), max_len)
width = max(width, 2)
bsz = len(encoded)
input_ids = torch.full((bsz, width), pad, dtype=torch.long)
labels = torch.full((bsz, width), IGNORE_INDEX, dtype=torch.long)
for i, (ids, lab) in enumerate(encoded):
n = min(len(ids), width)
input_ids[i, :n] = torch.tensor(ids[:n], dtype=torch.long)
labels[i, :n] = torch.tensor(lab[:n], dtype=torch.long)
return input_ids, labels
def iter_sft_batches(
rows: list[dict],
tok: Tokenizer,
batch: int,
max_len: int,
start: int = 0,
):
n = len(rows)
if n == 0:
raise ValueError("no SFT rows")
i = start
while True:
sl = [rows[j % n] for j in range(i, i + batch)]
yield i, *collate_sft(sl, tok, max_len)
i += batch
+139
View File
@@ -0,0 +1,139 @@
"""Greedy translation eval on line-aligned src/ref files.
python -m kda.training.eval_mt \\
--ckpt ckpts/k3_wiki.pt --src /data/eval/zh2en.src.txt \\
--ref /data/eval/zh2en.ref.txt --target-lang en
"""
from __future__ import annotations
import argparse
import json
from pathlib import Path
import torch
from kda.training.data import eos_id, load_tokenizer
from kda.training.prompts import instruction_prompt
from kda.training.success import _chrf, _detect_lang, translation_success
from kda.training.toy import load_ckpt
def _read_lines(path: str) -> list[str]:
return [ln.strip() for ln in Path(path).read_text(encoding="utf-8").splitlines() if ln.strip()]
def _instruction(src: str, target_lang: str) -> str:
return instruction_prompt(src, target_lang)
@torch.inference_mode()
def decode_one(model, tok, prompt: str, device: str, max_new: int) -> str:
ids = tok.encode(prompt)
if not ids:
return ""
inp = torch.tensor([ids], dtype=torch.long, device=device)
out = model.generate(inp, max_new, eos_token_id=eos_id(tok))
gen = out[0, inp.size(1) :].tolist()
return tok.decode(gen).strip()
def evaluate_pairs(
model,
tok,
srcs: list[str],
refs: list[str],
*,
target_lang: str,
device: str,
max_new: int,
limit: int | None,
) -> dict:
n = len(srcs)
if limit is not None:
n = min(n, limit)
hyps: list[str] = []
wins = 0
copies = 0
lang_ok = 0
chrf_sum = 0.0
for i in range(n):
src, ref = srcs[i], refs[i]
hyp = decode_one(model, tok, _instruction(src, target_lang), device, max_new)
hyps.append(hyp)
ok = translation_success(src, hyp, ref, target_lang=target_lang)
wins += int(ok)
copies += int(_chrf(hyp, src) >= 80.0 or hyp == src)
want = "zh" if target_lang.startswith("zh") else "en"
lang_ok += int(_detect_lang(hyp) == want)
chrf_sum += _chrf(hyp, ref)
corpus = {}
try:
from sacrebleu.metrics import BLEU, CHRF
corpus["chrf"] = float(CHRF(word_order=2).corpus_score(hyps, [refs[:n]]).score)
corpus["bleu"] = float(BLEU().corpus_score(hyps, [refs[:n]]).score)
except Exception:
corpus["chrf"] = chrf_sum / max(n, 1)
corpus["bleu"] = None
return {
"n": n,
"success_rate": wins / max(n, 1),
"copy_rate": copies / max(n, 1),
"lang_ok": lang_ok / max(n, 1),
"chrf": corpus["chrf"],
"bleu": corpus["bleu"],
"hyps": hyps,
}
def main() -> None:
p = argparse.ArgumentParser(description=__doc__)
p.add_argument("--ckpt", required=True)
p.add_argument("--tokenizer", default=None, help="override ckpt tokenizer field")
p.add_argument("--src", default=None, help="one source sentence per line")
p.add_argument("--ref", default=None, help="one reference sentence per line")
p.add_argument("--target-lang", default="en", choices=["en", "zh"])
p.add_argument("--max-new", type=int, default=64)
p.add_argument("--limit", type=int, default=None)
p.add_argument("--prefix", default=None, help="single-prompt smoke decode")
p.add_argument("--device", default="auto")
args = p.parse_args()
device = args.device
if device == "auto":
device = "cuda" if torch.cuda.is_available() else "cpu"
model, _config = load_ckpt(args.ckpt)
model.to(device).eval()
payload = torch.load(args.ckpt, map_location="cpu", weights_only=False)
tok_src = args.tokenizer or payload.get("tokenizer")
if not tok_src:
raise SystemExit("need --tokenizer or a 'tokenizer' field in the checkpoint")
tok = load_tokenizer(tok_src)
if args.prefix:
print(decode_one(model, tok, args.prefix, device, args.max_new))
if args.src and args.ref:
srcs, refs = _read_lines(args.src), _read_lines(args.ref)
if len(srcs) != len(refs):
raise SystemExit(f"src/ref length mismatch: {len(srcs)} vs {len(refs)}")
out = evaluate_pairs(
model,
tok,
srcs,
refs,
target_lang=args.target_lang,
device=device,
max_new=args.max_new,
limit=args.limit,
)
printable = {k: v for k, v in out.items() if k != "hyps"}
print(json.dumps(printable, ensure_ascii=False, indent=2))
elif not args.prefix:
raise SystemExit("pass --prefix and/or --src + --ref")
if __name__ == "__main__":
main()
+7
View File
@@ -0,0 +1,7 @@
"""Instruction strings shared by SFT and eval. Do not drift."""
def instruction_prompt(src: str, target_lang: str) -> str:
if target_lang.startswith("zh"):
return f"Translate to Chinese:\n{src}"
return f"Translate to English:\n{src}"
+47
View File
@@ -0,0 +1,47 @@
"""LR scale and token-horizon helpers for train_k3 / train_sft."""
from __future__ import annotations
import math
def lr_scale(
opt_step: int,
warmup: int,
total_opt: int,
min_ratio: float = 0.1,
) -> float:
"""Linear warmup (optimizer steps) then cosine down to ``min_ratio``.
``opt_step`` is 0-indexed at the optimizer update that is about to run.
"""
if warmup > 0 and opt_step < warmup:
return (opt_step + 1) / warmup
denom = max(total_opt - warmup - 1, 1)
progress = min(max(opt_step - warmup, 0) / denom, 1.0)
cosine = 0.5 * (1.0 + math.cos(math.pi * progress))
return min_ratio + (1.0 - min_ratio) * cosine
def tokens_per_micro(batch: int, seq_len: int) -> int:
return batch * seq_len
def total_opt_steps(
*,
max_tokens: int | None,
max_micro: int | None,
batch: int,
seq_len: int,
grad_acc: int,
) -> int:
"""Optimizer-step horizon used by cosine. At least 1."""
acc = max(grad_acc, 1)
candidates: list[int] = []
if max_tokens is not None and max_tokens > 0:
tpm = max(tokens_per_micro(batch, seq_len), 1)
candidates.append(math.ceil(max_tokens / (tpm * acc)))
if max_micro is not None and max_micro > 0:
candidates.append(math.ceil(max_micro / acc))
if not candidates:
return 1
return max(min(candidates), 1)
+86
View File
@@ -0,0 +1,86 @@
"""Frozen translation success() — SFT eval and RL reward must call this."""
from __future__ import annotations
import re
CHRF_MIN = 40.0
COPY_CHRF_MAX = 80.0
_CJK = re.compile(r"[\u4e00-\u9fff]")
def _detect_lang(text: str) -> str | None:
sample = text.strip()
if not sample:
return None
try:
from langdetect import detect
tag = detect(sample)
except Exception:
if _CJK.search(sample):
return "zh"
if any(c.isascii() and c.isalpha() for c in sample):
return "en"
return None
if tag.startswith("zh"):
return "zh"
return tag[:2]
def _chrf(hyp: str, ref: str) -> float:
"""chrF++ in 0–100. Falls back to char unigram F if sacrebleu is missing."""
if not hyp or not ref:
return 0.0
try:
from sacrebleu.metrics import CHRF
return float(CHRF(word_order=2).sentence_score(hyp, [ref]).score)
except Exception:
hyp_c, ref_c = list(hyp), list(ref)
if not hyp_c:
return 0.0
ref_set = set(ref_c)
overlap = sum(1 for c in hyp_c if c in ref_set)
prec = overlap / len(hyp_c)
rec = overlap / max(len(ref_c), 1)
if prec + rec == 0:
return 0.0
return 100.0 * 2 * prec * rec / (prec + rec)
def translation_success(
src: str,
hyp: str,
ref: str | None = None,
*,
target_lang: str,
chrf_min: float = CHRF_MIN,
copy_chrf_max: float = COPY_CHRF_MAX,
) -> bool:
"""Binary task success for zh↔en instruction translation.
1. non-empty hyp, no instruction leak prefix
2. langid(hyp) matches target_lang (zh / en)
3. hyp is not a copy of src
4. if ref is given, chrF(hyp, ref) >= chrf_min
"""
hyp = hyp.strip()
src = src.strip()
if not hyp:
return False
leak = ("翻译如下", "translate to", "translation:", "译文:")
head = hyp[:40].lower()
if any(p in head or p in hyp[:20] for p in leak):
return False
want = "zh" if target_lang.startswith("zh") else "en"
got = _detect_lang(hyp)
if got != want:
return False
if src and _chrf(hyp, src) >= copy_chrf_max:
return False
if hyp == src:
return False
if ref is not None and _chrf(hyp, ref.strip()) < chrf_min:
return False
return True
+110
View File
@@ -0,0 +1,110 @@
"""L7: toy training loop — overfit 起步.
target:
端到端验证模型 + 数据流 + optimizer + ckpt + generate.
toy data:
建一份 256-token vocab 的小数据集: e.g. 1000 个长度 32 随机 token 序列
起步只取 batch=4, 看能否在 ~320 steps 内把 loss 压到 < 0.1 (overfit 单 batch).
step:
optimizer = AdamW(lr=1e-3, wd=0.01)
loss.backward(); optimizer.step(); optimizer.zero_grad()
every N steps: 打印 loss
end: 保存 ckpt to ckpts/kda_toy.pt
ckpt:
save:
torch.save({"model_state": model.state_dict(), "config": asdict(config)}, path)
load:
torch.load -> model.load_state_dict
"""
from __future__ import annotations
import os
from dataclasses import asdict, fields
import torch
from ..models.causal_lm import CausalLM
from ..models.config import KDAConfig
from ..models.k3_config import K3Config
def make_toy_data(batch: int = 4, seq_len: int = 32, vocab: int = 256, seed: int = 42):
"""单 batch overfit 数据: 同一组序列循环."""
torch.manual_seed(seed)
seq = torch.randint(0, vocab, (batch, seq_len), dtype=torch.long)
return seq # 用作 input_ids 和 labels (shift one inside forward)
def train_one_batch(model, optimizer, input_ids, labels):
optimizer.zero_grad(set_to_none=True)
loss = model(input_ids, labels=labels)
loss.backward()
optimizer.step()
return loss.detach()
def save_ckpt(model, config, path: str):
os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
torch.save({"model_state": model.state_dict(), "config": asdict(config)}, path)
#: The feed-forward submodule was named after its contents (``mlp`` in the
#: dense config, ``moe`` in K3) before both were unified under ``ffn``.
#: Checkpoints saved before that rename still carry the old prefixes.
_LEGACY_PREFIXES = {
".mlp.": ".ffn.",
".mlp_norm.": ".ffn_norm.",
".moe.": ".ffn.",
".moe_norm.": ".ffn_norm.",
}
def _rename_legacy_keys(state: dict) -> dict:
def fix(key: str) -> str:
for old, new in _LEGACY_PREFIXES.items():
if old in key:
return key.replace(old, new)
return key
return {fix(k): v for k, v in state.items()}
def _config_from(payload_config: dict) -> K3Config | KDAConfig:
"""Pick the config class the checkpoint was written with.
``moe_latent_size`` is a K3-only field, so its presence identifies the
hybrid K3 architecture; anything else is the dense KDA config.
"""
cls = K3Config if "moe_latent_size" in payload_config else KDAConfig
known = {item.name for item in fields(cls)}
return cls(**{k: v for k, v in payload_config.items() if k in known})
def load_ckpt(path: str, model: CausalLM | None = None) -> tuple[CausalLM, K3Config | KDAConfig]:
payload = torch.load(path, map_location="cpu", weights_only=False)
config = _config_from(payload["config"])
if model is None:
model = CausalLM(config)
model.load_state_dict(_rename_legacy_keys(payload["model_state"]))
return model, config
def main():
"""主入口: overfit 起步. 320 steps 期望 loss < 0.1."""
device = "cuda" if torch.cuda.is_available() else "cpu"
config = KDAConfig()
model = CausalLM(config).to(device)
tokens = make_toy_data(seq_len=32, vocab=config.vocab_size).to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=0.01)
for step in range(320):
loss = train_one_batch(model, optimizer, tokens, tokens)
if step % 64 == 0 or step == 319:
print(f"step {step:3d} loss {loss.item():.4f}")
save_ckpt(model, config, "ckpts/kda_toy.pt")
if __name__ == "__main__":
main()
+51
View File
@@ -0,0 +1,51 @@
"""Train a SentencePiece tokenizer on a Chinese Wikipedia subset.
用法:
uv run python kda/training/train_tokenizer.py \
--out data/spm_4k --vocab-size 4096 --limit 20000
产出:
data/spm_4k.model / data/spm_4k.vocab (BPE/unigram, 中文小语料)
"""
from __future__ import annotations
import argparse
import sentencepiece as spm
from .data import fetch_wiki_texts
def main() -> None:
p = argparse.ArgumentParser(description=__doc__)
p.add_argument("--out", default="data/spm_4k", help="输出前缀 (model/vocab 文件)")
p.add_argument("--vocab-size", type=int, default=8192)
p.add_argument("--limit", type=int, default=20000, help="用于训练的 wiki 文章数")
p.add_argument("--model-type", default="unigram", choices=["unigram", "bpe"])
p.add_argument("--character-coverage", type=float, default=0.9995)
args = p.parse_args()
texts = fetch_wiki_texts(args.limit)
corpus = "".join(texts)
tmp = args.out + ".corpus.txt"
with open(tmp, "w", encoding="utf-8") as f:
f.write(corpus)
print(f"corpus: {len(corpus):,} chars from {len(texts)} articles")
spm.SentencePieceTrainer.train(
input=tmp,
model_prefix=args.out,
vocab_size=args.vocab_size,
model_type=args.model_type,
character_coverage=args.character_coverage,
unk_id=0,
pad_id=1,
bos_id=-1,
eos_id=-1,
num_threads=4,
)
print(f"tokenizer saved: {args.out}.model / {args.out}.vocab")
if __name__ == "__main__":
main()
+187
View File
@@ -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
+94
View File
@@ -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}}}
BIN
View File
Binary file not shown.
+39
View File
@@ -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}
+127
View File
@@ -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。
+111
View File
@@ -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\%。
+148
View File
@@ -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 逐位相同。
+114
View File
@@ -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$ 维投影。
+75
View File
@@ -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 内部完成。
+173
View File
@@ -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。
+158
View File
@@ -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 路由提供专业化能力。
+162
View File
@@ -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)。
+379
View File
@@ -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"} 保持旧路径不变。

Some files were not shown because too many files have changed in this diff Show More