Initial K3 snapshot: 0.5B KDA/MLA/MoE train path
Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
This commit is contained in:
@@ -0,0 +1,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
@@ -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
@@ -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"]
|
||||||
@@ -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
|
||||||
|
```
|
||||||
@@ -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:
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
# Mount as /data/eval: one sentence per line, e.g. zh2en.src.txt + zh2en.ref.txt
|
||||||
@@ -0,0 +1,20 @@
|
|||||||
|
今天天气很好。
|
||||||
|
请把窗户打开。
|
||||||
|
猫坐在垫子上。
|
||||||
|
这本书值得一读。
|
||||||
|
他昨天去了北京。
|
||||||
|
我们需要更多的训练数据。
|
||||||
|
太阳从东边升起。
|
||||||
|
她正在学习线性代数。
|
||||||
|
不要把评测集拿去训练。
|
||||||
|
河对面有一座旧桥。
|
||||||
|
科学是对自然的系统探索。
|
||||||
|
他们在公园里散步。
|
||||||
|
这台电脑的内存是十六吉字节。
|
||||||
|
翻译时不要照抄原文。
|
||||||
|
春天的风很温和。
|
||||||
|
我把钥匙放在桌子上了。
|
||||||
|
火车中午到达。
|
||||||
|
水在一百摄氏度沸腾。
|
||||||
|
小模型仍然可以学会狭窄的任务。
|
||||||
|
冻结测试文件保持只读。
|
||||||
@@ -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.
|
||||||
@@ -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.
|
||||||
@@ -0,0 +1,20 @@
|
|||||||
|
今天天气很好。
|
||||||
|
人工智能的发展改变了世界。
|
||||||
|
请把窗户打开。
|
||||||
|
猫坐在垫子上。
|
||||||
|
这本书值得一读。
|
||||||
|
会议将在下午三点开始。
|
||||||
|
他昨天去了北京。
|
||||||
|
我们需要更多的训练数据。
|
||||||
|
太阳从东边升起。
|
||||||
|
这个问题还没有答案。
|
||||||
|
她正在学习线性代数。
|
||||||
|
不要把评测集拿去训练。
|
||||||
|
河对面有一座旧桥。
|
||||||
|
请稍等,我马上回来。
|
||||||
|
科学是对自然的系统探索。
|
||||||
|
他们在公园里散步。
|
||||||
|
这台电脑的内存是十六吉字节。
|
||||||
|
翻译时不要照抄原文。
|
||||||
|
春天的风很温和。
|
||||||
|
我把钥匙放在桌子上了。
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
# Mount bilingual pretrain shards here (not baked into the image).
|
||||||
@@ -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.
|
||||||
@@ -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"}
|
||||||
@@ -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()
|
||||||
@@ -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"
|
||||||
@@ -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.
|
||||||
@@ -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
|
||||||
@@ -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"
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
# Vendored FLA modules used by KDA (l2norm).
|
||||||
@@ -0,0 +1,3 @@
|
|||||||
|
from kda._fla.ops.backends import BackendRegistry, BaseBackend, dispatch
|
||||||
|
|
||||||
|
__all__ = ["BackendRegistry", "BaseBackend", "dispatch"]
|
||||||
@@ -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)
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
# Vendored FLA ops subset.
|
||||||
@@ -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"]
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
# Vendored FLA common kernels used by KDA.
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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)
|
||||||
@@ -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"]
|
||||||
@@ -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
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
# Vendored GLA chunk output kernel used by KDA.
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,7 @@
|
|||||||
|
from .chunk import chunk_kda
|
||||||
|
from .fused_recurrent import fused_recurrent_kda
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"chunk_kda",
|
||||||
|
"fused_recurrent_kda",
|
||||||
|
]
|
||||||
@@ -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,
|
||||||
|
)
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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",
|
||||||
|
]
|
||||||
@@ -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."
|
||||||
|
)
|
||||||
@@ -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
|
||||||
@@ -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)",
|
||||||
|
)
|
||||||
@@ -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)
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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()
|
||||||
@@ -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
|
||||||
@@ -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",
|
||||||
|
]
|
||||||
@@ -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",
|
||||||
|
]
|
||||||
@@ -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))
|
||||||
@@ -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))
|
||||||
@@ -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)
|
||||||
@@ -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]
|
||||||
@@ -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
|
||||||
@@ -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))
|
||||||
@@ -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"]
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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()]
|
||||||
@@ -0,0 +1,5 @@
|
|||||||
|
"""KDA operator API and implementation backends."""
|
||||||
|
|
||||||
|
from .api import chunk_kda
|
||||||
|
|
||||||
|
__all__ = ["chunk_kda"]
|
||||||
+167
@@ -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),
|
||||||
|
)
|
||||||
@@ -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"]
|
||||||
@@ -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"]
|
||||||
@@ -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",
|
||||||
|
]
|
||||||
@@ -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
|
||||||
@@ -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"]
|
||||||
@@ -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)
|
||||||
@@ -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"]
|
||||||
@@ -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"]
|
||||||
@@ -0,0 +1,5 @@
|
|||||||
|
"""Vendored FLA chunk KDA backward."""
|
||||||
|
|
||||||
|
from kda._fla.ops.kda.chunk_bwd import chunk_kda_bwd
|
||||||
|
|
||||||
|
__all__ = ["chunk_kda_bwd"]
|
||||||
@@ -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,
|
||||||
|
)
|
||||||
@@ -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",
|
||||||
|
]
|
||||||
@@ -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"]
|
||||||
@@ -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"]
|
||||||
@@ -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
|
||||||
@@ -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()
|
||||||
@@ -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}"
|
||||||
@@ -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)
|
||||||
@@ -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
|
||||||
@@ -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()
|
||||||
@@ -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()
|
||||||
@@ -0,0 +1,187 @@
|
|||||||
|
schema: superpaper.ledger/v1
|
||||||
|
retired_ids: []
|
||||||
|
paper:
|
||||||
|
id: "kda-project"
|
||||||
|
title: "KDA 训练→推理 手写实现 — 完整笔记"
|
||||||
|
authors: ["dela"]
|
||||||
|
notes_language: zh
|
||||||
|
source:
|
||||||
|
kind: markdown
|
||||||
|
|
||||||
|
coverage:
|
||||||
|
mode: full
|
||||||
|
sections_in:
|
||||||
|
- "KDA 递归核心"
|
||||||
|
- "Gate 激活"
|
||||||
|
- "分块并行计算"
|
||||||
|
- "GVA 分组值注意力"
|
||||||
|
- "KDAAttention 层"
|
||||||
|
- "Gated MLA 矩阵吸收版"
|
||||||
|
- "SiTU-GLU 与 Stable LatentMoE"
|
||||||
|
- "K3 混合架构"
|
||||||
|
- "Attention Residual 深度残差"
|
||||||
|
- "反向传播推导"
|
||||||
|
sections_skipped:
|
||||||
|
- "Triton kernel 细节"
|
||||||
|
- "Docker 部署"
|
||||||
|
- "AttnRes 论文的 kernel 级调度与 pipeline 重叠"
|
||||||
|
|
||||||
|
questions:
|
||||||
|
- id: Q1
|
||||||
|
text: "KDA 的状态更新如何避免 softmax、实现线性复杂度?"
|
||||||
|
- id: Q2
|
||||||
|
text: "safe gate 与 standard gate 的区别是什么?"
|
||||||
|
- id: Q3
|
||||||
|
text: "分块并行如何在保持递归等价的同时利用 GPU 并行?"
|
||||||
|
- id: Q4
|
||||||
|
text: "GVA 的 repeat_interleave + sum 反向是怎么回事?"
|
||||||
|
- id: Q5
|
||||||
|
text: "MLA 矩阵吸收如何避免解压 K/V?"
|
||||||
|
- id: Q6
|
||||||
|
text: "SiTU-GLU 为什么比 SwiGLU 更稳定?"
|
||||||
|
- id: Q7
|
||||||
|
text: "AttnRes 如何把残差流从等权累加换成按内容选择?"
|
||||||
|
- id: Q8
|
||||||
|
text: "Block AttnRes 的两阶段算法为什么和 naive 逐层实现数值等价?"
|
||||||
|
- id: Q9
|
||||||
|
text: "深度残差接入 CausalLM 时怎样避免参数被重复注册?"
|
||||||
|
|
||||||
|
claims:
|
||||||
|
- id: C1
|
||||||
|
text: "KDA 用 delta rule 更新 KV 状态矩阵,不需要 softmax,复杂度 O(T·K·V)"
|
||||||
|
kind: methodological
|
||||||
|
status: core
|
||||||
|
- id: C2
|
||||||
|
text: "safe gate = lower_bound · σ(rate · input),保证 gate 值在 [lower_bound, 0] 范围内"
|
||||||
|
kind: methodological
|
||||||
|
status: core
|
||||||
|
- id: C3
|
||||||
|
text: "分块计算:chunk 内用下三角解,chunk 间用状态递推,数值等价于 naive recurrent"
|
||||||
|
kind: methodological
|
||||||
|
status: core
|
||||||
|
- id: C4
|
||||||
|
text: "MLA 矩阵吸收:q 吸收 W_UK 后直接与 latent c 内积,永不解压 K/V"
|
||||||
|
kind: methodological
|
||||||
|
status: core
|
||||||
|
- id: C5
|
||||||
|
text: "LatentMoE 通过 latent 接口把 routed 专家限制在半宽空间 ℓ=d/2"
|
||||||
|
kind: methodological
|
||||||
|
status: core
|
||||||
|
- id: C6
|
||||||
|
text: "AttnRes 用逐 token 的深度维 softmax 代替等权残差累加:打分在 RMS 归一化后做,加权和在原始张量上做"
|
||||||
|
kind: methodological
|
||||||
|
status: core
|
||||||
|
- id: C7
|
||||||
|
text: "Block AttnRes 块内退化为普通求和、只让块输出进入源列表,源数从 O(N) 降到 O(N/S)"
|
||||||
|
kind: methodological
|
||||||
|
status: core
|
||||||
|
- id: C8
|
||||||
|
text: "两阶段算法 = inter 块间批量 einsum + intra online-softmax 增量合并,与 naive 逐层实现数值等价 (atol 1e-5)"
|
||||||
|
kind: methodological
|
||||||
|
status: core
|
||||||
|
- id: C9
|
||||||
|
text: "BorrowedSubLayer 用普通 tuple 持有 norm/fn,不注册为子模块,保证参数与 state_dict 键不重复"
|
||||||
|
kind: methodological
|
||||||
|
status: supporting
|
||||||
|
|
||||||
|
symbols:
|
||||||
|
- {name: B, latex: "B", meaning: "batch size", kind: "shape parameter"}
|
||||||
|
- {name: T, latex: "T", meaning: "序列长度", kind: "shape parameter"}
|
||||||
|
- {name: H, latex: "H", meaning: "query/key 头数", kind: "shape parameter"}
|
||||||
|
- {name: HV, latex: "H_V", meaning: "value 头数 (GVA)", kind: "shape parameter"}
|
||||||
|
- {name: G, latex: "G", meaning: "GVA 组数 = HV/H", kind: "shape parameter"}
|
||||||
|
- {name: K, latex: "K", meaning: "key/query 头维度", kind: "shape parameter"}
|
||||||
|
- {name: V, latex: "V", meaning: "value 头维度 (= K)", kind: "shape parameter"}
|
||||||
|
- {name: D, latex: "D", meaning: "hidden_size", kind: "shape parameter"}
|
||||||
|
- {name: C, latex: "C", meaning: "chunk_size", kind: "shape parameter"}
|
||||||
|
- {name: r, latex: "r", meaning: "KV latent rank (kv_lora_rank)", kind: "shape parameter"}
|
||||||
|
- {name: ell, latex: "\\ell", meaning: "MoE latent 宽度 = d/2", kind: "shape parameter"}
|
||||||
|
- {name: S, latex: "S", meaning: "KV 状态矩阵", domain: "[B, HV, K, V]", kind: value}
|
||||||
|
- {name: q, latex: "q", meaning: "query", domain: "[B, T, H, K]", kind: value}
|
||||||
|
- {name: k, latex: "k", meaning: "key", domain: "[B, T, H, K]", kind: value}
|
||||||
|
- {name: v, latex: "v", meaning: "value", domain: "[B, T, HV, V]", kind: value}
|
||||||
|
- {name: g, latex: "g", meaning: "gate (log-space decay)", domain: "[B, T, HV, K]", kind: value}
|
||||||
|
- {name: beta, latex: "\\beta", meaning: "学习率/写入强度", domain: "[B, T, HV]", kind: value}
|
||||||
|
- {name: A_log, latex: "A_{\\log}", meaning: "head-wise 衰减参数 (log-space)", domain: "[HV]", kind: value}
|
||||||
|
- {name: dt_bias, latex: "\\Delta_b", meaning: "per-dim gate bias", domain: "[HV, K]", kind: value}
|
||||||
|
- {name: c, latex: "c", meaning: "KV latent 向量", domain: "[B, T, r]", kind: value}
|
||||||
|
- {name: W_UK, latex: "W_{UK}", meaning: "Key 解压矩阵 (MLA)", domain: "[H, d_q, r]", kind: value}
|
||||||
|
- {name: W_UV, latex: "W_{UV}", meaning: "Value 解压矩阵 (MLA)", domain: "[H, d_v, r]", kind: value}
|
||||||
|
- {name: N, latex: "N", meaning: "AttnRes 原子层数 = 2L", kind: "shape parameter"}
|
||||||
|
- {name: S, latex: "S", meaning: "AttnRes 块大小(原子层)", kind: "shape parameter"}
|
||||||
|
- {name: v_i, latex: "v_i", meaning: "AttnRes 第 i 个源(v_0 = embedding 输出)", domain: "[B, T, D]", kind: value}
|
||||||
|
- {name: w_l, latex: "w_l", meaning: "第 l 层 depth query(零初始化)", domain: "[D]", kind: value}
|
||||||
|
- {name: alpha, latex: "\\alpha_{l,i}", meaning: "深度维 softmax 权重", domain: "[n, B, T]", kind: value}
|
||||||
|
- {name: h_l, latex: "h_l", meaning: "深度注意力聚合出的层输入", domain: "[B, T, D]", kind: value}
|
||||||
|
- {name: b_j, latex: "b_j", meaning: "Block AttnRes 第 j 块的输出", domain: "[B, T, D]", kind: value}
|
||||||
|
- {name: p, latex: "p", meaning: "块内 running partial", domain: "[B, T, D]", kind: value}
|
||||||
|
|
||||||
|
terms:
|
||||||
|
- {canonical: "KDA", aliases: ["Key-Decayed Attention", "键衰减注意力"]}
|
||||||
|
- {canonical: "GVA", aliases: ["Grouped Value Attention", "分组值注意力"]}
|
||||||
|
- {canonical: "MLA", aliases: ["Multi-head Latent Attention", "多头隐变量注意力"]}
|
||||||
|
- {canonical: "MoE", aliases: ["Mixture of Experts", "混合专家"]}
|
||||||
|
- {canonical: "SiTU-GLU", aliases: ["Sigmoid Tanh Unit GLU"]}
|
||||||
|
- {canonical: "delta rule", aliases: ["δ 规则"]}
|
||||||
|
- {canonical: "safe gate", aliases: ["安全门控"]}
|
||||||
|
- {canonical: "matrix absorption", aliases: ["矩阵吸收"]}
|
||||||
|
- {canonical: "AttnRes", aliases: ["Attention Residual", "注意力残差", "深度残差"]}
|
||||||
|
- {canonical: "depth residual", aliases: ["DepthResidual", "深度维残差"]}
|
||||||
|
- {canonical: "online softmax", aliases: ["在线 softmax", "增量 softmax"]}
|
||||||
|
- {canonical: "atomic layer", aliases: ["原子层", "atomic sublayer"]}
|
||||||
|
|
||||||
|
derivations:
|
||||||
|
- id: DER1
|
||||||
|
claim: C1
|
||||||
|
title: "KDA 递归状态更新推导"
|
||||||
|
expand: true
|
||||||
|
figure: null
|
||||||
|
steps:
|
||||||
|
- {id: "1", from: "S_{t-1}", to: "S_{\\mathrm{dec}} = \\exp(g_t) \\odot S_{t-1}", rule: scale}
|
||||||
|
- {id: "2", from: "S_{\\mathrm{dec}}", to: "r_t = v_t - k_t \\cdot S_{\\mathrm{dec}}", rule: definition}
|
||||||
|
- {id: "3", from: "r_t", to: "S_t = S_{\\mathrm{dec}} + (\\beta_t k_t) \\otimes r_t", rule: definition}
|
||||||
|
- {id: "4", from: "S_t", to: "o_t = (q_t \\cdot \\text{scale}) \\cdot S_t", rule: definition}
|
||||||
|
- id: DER2
|
||||||
|
claim: C4
|
||||||
|
title: "MLA 矩阵吸收推导"
|
||||||
|
expand: true
|
||||||
|
figure: null
|
||||||
|
steps:
|
||||||
|
- {id: "1", from: "q \\in [B,T,H,d_q]", to: "q_{\\mathrm{abs}} = q \\cdot W_{UK} \\in [B,T,H,r]", rule: substitute}
|
||||||
|
- {id: "2", from: "q_{\\mathrm{abs}}, c", to: "\\text{score} = q_{\\mathrm{abs}} \\cdot c^T \\in [B,H,T,T]", rule: definition}
|
||||||
|
- {id: "3", from: "\\text{attn}, c", to: "\\tilde{o}_{\\mathrm{lat}} = \\text{attn} \\cdot c \\in [B,H,T,r]", rule: definition}
|
||||||
|
- {id: "4", from: "\\tilde{o}_{\\mathrm{lat}}", to: "\\tilde{o} = \\tilde{o}_{\\mathrm{lat}} \\cdot W_{UV}^T \\in [B,H,T,d_v]", rule: substitute}
|
||||||
|
- id: DER3
|
||||||
|
claim: C8
|
||||||
|
title: "AttnRes 两阶段 online softmax 合并推导"
|
||||||
|
expand: true
|
||||||
|
figure: null
|
||||||
|
steps:
|
||||||
|
- {id: "1", from: "s_{l,i} = \\tilde{w}_l^T \\mathrm{RMS}(v_i)", to: "(m, n, d) = (\\max_i s_i, \\sum_i e^{s_i - m} v_i, \\sum_i e^{s_i - m})", rule: definition}
|
||||||
|
- {id: "2", from: "inter sources b_0..b_{j-1} 固定", to: "一次批量 einsum 'q d, n b t d -> q n b t' 得块内全部 query 的 (m,n,d)", rule: substitute}
|
||||||
|
- {id: "3", from: "单源 partial p", to: "(m, n, d) = (s_p, p, 1),因为 e^{s_p - m} = 1", rule: definition}
|
||||||
|
- {id: "4", from: "(m_a,n_a,d_a), (m_b,n_b,d_b)", to: "m = \\max(m_a,m_b);\\ n = e^{m_a-m} n_a + e^{m_b-m} n_b;\\ d = e^{m_a-m} d_a + e^{m_b-m} d_b", rule: scale}
|
||||||
|
- {id: "5", from: "(m, n, d)", to: "h_l = n / d,与 forward_naive 逐位一致", rule: definition}
|
||||||
|
|
||||||
|
figures:
|
||||||
|
- id: F1
|
||||||
|
claim: C1
|
||||||
|
title: "KDA 递归状态更新张量图"
|
||||||
|
grammar: tensor-face
|
||||||
|
toolkit: supertensor
|
||||||
|
signals: [shape, contraction, broadcast]
|
||||||
|
status: planned
|
||||||
|
- id: F2
|
||||||
|
claim: C4
|
||||||
|
title: "MLA 矩阵吸收计算流"
|
||||||
|
grammar: tensor-face
|
||||||
|
toolkit: supertensor
|
||||||
|
signals: [shape, contraction, transpose]
|
||||||
|
status: planned
|
||||||
|
- id: F3
|
||||||
|
claim: C7
|
||||||
|
title: "Full vs Block AttnRes 的源列表增长"
|
||||||
|
grammar: tensor-face
|
||||||
|
toolkit: supertensor
|
||||||
|
signals: [shape, contraction]
|
||||||
|
status: planned
|
||||||
@@ -0,0 +1,94 @@
|
|||||||
|
% KDA 笔记 preamble — 基于 superpaper/assets/notes-macros.tex
|
||||||
|
\usepackage[fontset=fandol]{ctex}
|
||||||
|
\usepackage{amsmath,amssymb}
|
||||||
|
\usepackage{graphicx}
|
||||||
|
\usepackage[margin=2.2cm]{geometry}
|
||||||
|
\usepackage[most]{tcolorbox}
|
||||||
|
\usepackage{etoolbox}
|
||||||
|
\usepackage{listings}
|
||||||
|
\usepackage{booktabs}
|
||||||
|
\usepackage{subcaption}
|
||||||
|
\usepackage{float}
|
||||||
|
\usepackage{tikz}
|
||||||
|
\usepackage{hyperref}
|
||||||
|
\usepackage{xcolor}
|
||||||
|
\usepackage{multicol}
|
||||||
|
\usepackage{tabularx}
|
||||||
|
\usepackage{array}
|
||||||
|
\usepackage{enumitem}
|
||||||
|
|
||||||
|
% ---------- 颜色 ----------
|
||||||
|
\definecolor{codebg}{HTML}{F7F7F7}
|
||||||
|
\definecolor{codeframe}{HTML}{CCCCCC}
|
||||||
|
\definecolor{shapecolor}{HTML}{2E86C1}
|
||||||
|
\definecolor{einsumcolor}{HTML}{884EA0}
|
||||||
|
\definecolor{notegreen}{HTML}{27AE60}
|
||||||
|
\definecolor{warnorange}{HTML}{E67E22}
|
||||||
|
|
||||||
|
% ---------- 代码样式 ----------
|
||||||
|
\lstdefinestyle{pycode}{
|
||||||
|
language=Python,
|
||||||
|
backgroundcolor=\color{codebg},
|
||||||
|
frame=single,
|
||||||
|
rulecolor=\color{codeframe},
|
||||||
|
basicstyle=\ttfamily\small,
|
||||||
|
keywordstyle=\color{blue!70!black}\bfseries,
|
||||||
|
commentstyle=\color{gray},
|
||||||
|
stringstyle=\color{red!60!black},
|
||||||
|
showstringspaces=false,
|
||||||
|
breaklines=true,
|
||||||
|
tabsize=4,
|
||||||
|
columns=flexible,
|
||||||
|
xleftmargin=4pt,
|
||||||
|
xrightmargin=4pt,
|
||||||
|
aboveskip=6pt,
|
||||||
|
belowskip=6pt,
|
||||||
|
}
|
||||||
|
\lstset{style=pycode}
|
||||||
|
|
||||||
|
% ---------- 形状标注命令 ----------
|
||||||
|
\newcommand{\shape}[1]{{\color{shapecolor}\ensuremath{[#1]}}}
|
||||||
|
\newcommand{\einsum}[1]{{\color{einsumcolor}\texttt{einsum(#1)}}}
|
||||||
|
\newcommand{\note}[1]{{\color{notegreen}\textit{#1}}}
|
||||||
|
|
||||||
|
% ---------- 盒子 ----------
|
||||||
|
\newtcolorbox{knowledgebox}[1]{
|
||||||
|
enhanced, colback=blue!5!white, colframe=blue!75!black, colbacktitle=blue!75!black,
|
||||||
|
coltitle=white, fonttitle=\bfseries, title=#1,
|
||||||
|
attach boxed title to top left={yshift=-2mm, xshift=2mm},
|
||||||
|
boxrule=1pt, sharp corners
|
||||||
|
}
|
||||||
|
\newtcolorbox{importantbox}[1]{
|
||||||
|
enhanced, colback=yellow!10!white, colframe=yellow!80!black, colbacktitle=yellow!80!black,
|
||||||
|
coltitle=black, fonttitle=\bfseries, title=#1, sharp corners
|
||||||
|
}
|
||||||
|
\newtcolorbox{warningbox}[1]{
|
||||||
|
enhanced, colback=red!5!white, colframe=red!75!black, colbacktitle=red!75!black,
|
||||||
|
coltitle=white, fonttitle=\bfseries, title=#1, sharp corners
|
||||||
|
}
|
||||||
|
\newtcolorbox{tensorbox}[1]{
|
||||||
|
enhanced, colback=blue!3!white, colframe=blue!40!black, colbacktitle=blue!50!black,
|
||||||
|
coltitle=white, fonttitle=\bfseries, title=#1, sharp corners,
|
||||||
|
boxrule=0.8pt
|
||||||
|
}
|
||||||
|
|
||||||
|
% 代码-公式并行盒子
|
||||||
|
\newtcolorbox{codemathtop}[1]{
|
||||||
|
enhanced, colback=codebg, colframe=codeframe,
|
||||||
|
fonttitle=\bfseries\ttfamily\small, title=#1,
|
||||||
|
sharp corners, boxrule=0.6pt,
|
||||||
|
left=4pt, right=4pt, top=2pt, bottom=2pt,
|
||||||
|
}
|
||||||
|
\newtcolorbox{codemathbot}{
|
||||||
|
enhanced, colback=white, colframe=blue!30!black,
|
||||||
|
sharp corners, boxrule=0.6pt,
|
||||||
|
left=4pt, right=4pt, top=4pt, bottom=4pt,
|
||||||
|
}
|
||||||
|
|
||||||
|
% ---------- 文档元信息 ----------
|
||||||
|
\newcommand{\notetitle}{KDA 训练→推理 完整笔记}
|
||||||
|
\newcommand{\notesubtitle}{递归 · 分块 · Gate · MLA · MoE · AttnRes · K3 架构}
|
||||||
|
\newcommand{\notedate}{\today}
|
||||||
|
|
||||||
|
\newcommand{\splabel}[1]{\hypertarget{sp:#1}{}\label{sp:#1}}
|
||||||
|
\newcommand{\spref}[1]{\hyperlink{sp:#1}{\texttt{#1}}}
|
||||||
Binary file not shown.
@@ -0,0 +1,39 @@
|
|||||||
|
\documentclass[a4paper,11pt]{article}
|
||||||
|
\input{notes-macros}
|
||||||
|
|
||||||
|
\begin{document}
|
||||||
|
|
||||||
|
% ---------- 封面 ----------
|
||||||
|
\begin{titlepage}
|
||||||
|
\centering
|
||||||
|
\vspace{2cm}
|
||||||
|
{\huge\bfseries \notetitle\par}
|
||||||
|
\vspace{0.6cm}
|
||||||
|
{\Large \notesubtitle\par}
|
||||||
|
\vspace{0.8cm}
|
||||||
|
{\large \notedate\par}
|
||||||
|
\vspace{1.5cm}
|
||||||
|
\begin{tcolorbox}[width=0.88\textwidth, colback=black!2!white, colframe=black!60, sharp corners]
|
||||||
|
\textbf{项目}:\texttt{projects/kda/} — KDA 手写实现(naive recurrent → chunked → Triton)\\
|
||||||
|
\textbf{架构}:KDA + Gated MLA + Stable LatentMoE + AttnRes 深度残差 (K3-like)\\
|
||||||
|
\textbf{参考}:KDA arXiv:2510.26692; AttnRes arXiv:2603.15031; Kimi K3 architecture notes\\
|
||||||
|
\textbf{代码}:\texttt{kda/ops/}, \texttt{kda/layers/}, \texttt{kda/models/}
|
||||||
|
\end{tcolorbox}
|
||||||
|
\end{titlepage}
|
||||||
|
|
||||||
|
\tableofcontents
|
||||||
|
\newpage
|
||||||
|
|
||||||
|
\input{sections/sec-01} % KDA 递归核心
|
||||||
|
\input{sections/sec-02} % Gate 激活
|
||||||
|
\input{sections/sec-03} % 分块并行计算
|
||||||
|
\input{sections/sec-04} % GVA 分组值注意力
|
||||||
|
\input{sections/sec-05} % KDAAttention 层
|
||||||
|
\input{sections/sec-06} % Gated MLA 矩阵吸收版
|
||||||
|
\input{sections/sec-07} % SiTU-GLU 与 Stable LatentMoE
|
||||||
|
\input{sections/sec-08} % K3 混合架构
|
||||||
|
\input{sections/sec-09} % Attention Residual 深度残差
|
||||||
|
\input{sections/sec-10} % 反向传播推导
|
||||||
|
\input{sections/sec-11} % 符号表
|
||||||
|
|
||||||
|
\end{document}
|
||||||
@@ -0,0 +1,127 @@
|
|||||||
|
% teach:
|
||||||
|
% gap: 读者知道 softmax attention 但不知道线性注意力怎么维护状态矩阵
|
||||||
|
% takeaway: KDA 用 delta rule 逐步更新 [K,V] 状态矩阵, 写入=擦旧写新, 每步 O(KV)
|
||||||
|
% jump: 为什么 r_t = v - k·S 而不是直接用 v?delta rule 的"先擦再写"
|
||||||
|
% omit: KDA 论文的 related work、实验细节
|
||||||
|
|
||||||
|
\section{KDA 递归核心}
|
||||||
|
\splabel{C1}
|
||||||
|
|
||||||
|
\subsection{动机:从 softmax 到状态矩阵}
|
||||||
|
|
||||||
|
标准 attention 每个 token 都要回看所有历史,复杂度 $O(T^2)$。
|
||||||
|
线性注意力换掉 softmax,把 $\sum_j v_j k_j^T$ 压成一个 $K \times V$ 的状态矩阵 $S$,
|
||||||
|
每步只做 $o_t = q_t \cdot S$,复杂度降到 $O(T \cdot K \cdot V)$。
|
||||||
|
|
||||||
|
但裸线性注意力的问题是:$S$ 只能加,不能改。写进去的信息永远在那里。
|
||||||
|
KDA 的核心想法是给 $S$ 加两个操作:\textbf{衰减}(逐渐忘记旧信息)和
|
||||||
|
\textbf{delta rule}(先擦旧的,再写新的)。
|
||||||
|
|
||||||
|
\begin{importantbox}{如果你只记一件事}
|
||||||
|
KDA 的状态更新 = 衰减旧状态 + $\beta_t k_t$ 写入 $(v_t - k_t \cdot S_{\mathrm{dec}})$。
|
||||||
|
减去 $k_t \cdot S_{\mathrm{dec}}$ 就是"先把 $k_t$ 方向的旧预测擦掉"。
|
||||||
|
\end{importantbox}
|
||||||
|
|
||||||
|
\subsection{逐步公式}
|
||||||
|
|
||||||
|
\noindent\textbf{输入张量:}
|
||||||
|
|
||||||
|
\begin{center}
|
||||||
|
\begin{tabular}{lll}
|
||||||
|
\toprule
|
||||||
|
符号 & 形状 & 含义 \\
|
||||||
|
\midrule
|
||||||
|
$q_t$ & \shape{B, HV, K} & query(已经 repeat\_interleave 到 HV) \\
|
||||||
|
$k_t$ & \shape{B, HV, K} & key(同上) \\
|
||||||
|
$v_t$ & \shape{B, HV, V} & value \\
|
||||||
|
$g_t$ & \shape{B, HV, K} & gate(log-space 衰减率,逐维) \\
|
||||||
|
$\beta_t$ & \shape{B, HV} & 写入强度标量 \\
|
||||||
|
$S_{t-1}$ & \shape{B, HV, K, V} & 上一步的 KV 状态 \\
|
||||||
|
\bottomrule
|
||||||
|
\end{tabular}
|
||||||
|
\end{center}
|
||||||
|
|
||||||
|
\noindent\textbf{四步更新:}
|
||||||
|
|
||||||
|
\begin{enumerate}[leftmargin=2em]
|
||||||
|
\item \textbf{衰减旧状态}(逐元素,$g_t$ 是 log-space 所以取 exp):
|
||||||
|
\[
|
||||||
|
S_{\mathrm{dec}} = \exp(g_t) \odot S_{t-1}
|
||||||
|
\qquad \shape{B, HV, K, V}
|
||||||
|
\]
|
||||||
|
|
||||||
|
\item \textbf{计算残差}(先用 $k_t$ 查旧状态,得到"旧预测",再减掉):
|
||||||
|
\[
|
||||||
|
p_t = \sum_k k_{t,k} \cdot S_{\mathrm{dec},k,\cdot}
|
||||||
|
= \texttt{einsum('bhk, bhkv -> bhv')}
|
||||||
|
\qquad \shape{B, HV, V}
|
||||||
|
\]
|
||||||
|
\[
|
||||||
|
r_t = v_t - p_t \qquad \shape{B, HV, V}
|
||||||
|
\]
|
||||||
|
|
||||||
|
\item \textbf{写入状态}(外积 rank-1 更新):
|
||||||
|
\[
|
||||||
|
a_t = \beta_t \cdot k_t \qquad \shape{B, HV, K}
|
||||||
|
\]
|
||||||
|
\[
|
||||||
|
S_t = S_{\mathrm{dec}} + a_t \otimes r_t
|
||||||
|
= S_{\mathrm{dec}} + \texttt{einsum('bhk, bhv -> bhkv')}
|
||||||
|
\qquad \shape{B, HV, K, V}
|
||||||
|
\]
|
||||||
|
|
||||||
|
\item \textbf{读出}:
|
||||||
|
\[
|
||||||
|
o_t = \frac{1}{\sqrt{K}} \cdot q_t \cdot S_t
|
||||||
|
= \texttt{einsum('bhk, bhkv -> bhv')}
|
||||||
|
\qquad \shape{B, HV, V}
|
||||||
|
\]
|
||||||
|
\end{enumerate}
|
||||||
|
|
||||||
|
\subsection{代码对照}
|
||||||
|
|
||||||
|
\begin{codemathtop}{ops/reference/recurrent.py — naive\_kda\_fwd (核心循环)}
|
||||||
|
\begin{lstlisting}
|
||||||
|
for t in range(T):
|
||||||
|
q_t = qe[:, t] # [B, HV, K]
|
||||||
|
k_t = ke[:, t] # [B, HV, K]
|
||||||
|
v_t = v[:, t] # [B, HV, V]
|
||||||
|
g_t = g[:, t] # [B, HV, K]
|
||||||
|
b_t = beta[:, t] # [B, HV]
|
||||||
|
|
||||||
|
# Step 1: decay
|
||||||
|
S_dec = S * g_t.exp().unsqueeze(-1) # [B,HV,K,V]
|
||||||
|
|
||||||
|
# Step 2: residual (delta rule)
|
||||||
|
p_t = einsum('bhk, bhkv -> bhv', k_t, S_dec)
|
||||||
|
r_t = v_t - p_t # [B,HV,V]
|
||||||
|
|
||||||
|
# Step 3: write (rank-1 update)
|
||||||
|
a_t = b_t.unsqueeze(-1) * k_t # [B,HV,K]
|
||||||
|
S = S_dec + einsum('bhk, bhv -> bhkv', a_t, r_t)
|
||||||
|
|
||||||
|
# Step 4: read
|
||||||
|
o[:, t] = einsum('bhk, bhkv -> bhv', q_t, S)
|
||||||
|
\end{lstlisting}
|
||||||
|
\end{codemathtop}
|
||||||
|
|
||||||
|
\begin{warningbox}{为什么 exp(g\_t) 要 unsqueeze(-1)?}
|
||||||
|
$g_t$ 的形状是 \shape{B, HV, K},而 $S$ 是 \shape{B, HV, K, V}。
|
||||||
|
衰减是在 $K$ 维上逐元素(同一 $k$ 索引的所有 $v$ 维度共享同一个衰减率),
|
||||||
|
所以 \texttt{exp(g\_t).unsqueeze(-1)} 把 K 维 broadcast 到 $K \times V$。
|
||||||
|
\end{warningbox}
|
||||||
|
|
||||||
|
\subsection{Delta rule 的直觉}
|
||||||
|
|
||||||
|
\begin{knowledgebox}{为什么减去 $k_t \cdot S_{\mathrm{dec}}$?}
|
||||||
|
把 $S$ 想象成一个 $K \to V$ 的线性映射。用 $k_t$ 去查它,得到的 $p_t = k_t^T S$
|
||||||
|
就是``旧状态对 $k_t$ 方向的预测''。如果 $p_t$ 已经很接近 $v_t$,说明这个方向的信息
|
||||||
|
已经写好了,不需要再写。$r_t = v_t - p_t$ 就是``需要修正的量''。
|
||||||
|
|
||||||
|
这就是 Widrow-Hoff delta rule:不是盲目地加,而是只修正误差。
|
||||||
|
\end{knowledgebox}
|
||||||
|
|
||||||
|
\subsection{本章小结}
|
||||||
|
|
||||||
|
KDA 的状态更新 = 衰减旧状态 + $\beta_t k_t$ 写入残差 $(v_t - k_t \cdot S_{\mathrm{dec}})$。
|
||||||
|
每步复杂度 $O(K \cdot V)$(两次矩阵-向量乘 + 一次外积),不需要 softmax。
|
||||||
@@ -0,0 +1,111 @@
|
|||||||
|
% teach:
|
||||||
|
% gap: 读者知道 g_t 是 gate 但不知道它怎么从 raw projection 变成一个负的 log-space 衰减
|
||||||
|
% takeaway: safe gate 用 sigmoid 把值夹在 [lower_bound, 0], standard gate 用 -softplus 保证负
|
||||||
|
% jump: 论文没解释为什么需要 A_log 和 dt_bias 两层
|
||||||
|
% omit: Triton gate kernel 的 fused 实现细节
|
||||||
|
|
||||||
|
\section{Gate 激活}
|
||||||
|
\splabel{C2}
|
||||||
|
|
||||||
|
\subsection{Gate 的角色}
|
||||||
|
|
||||||
|
回顾 §1:$S_{\mathrm{dec}} = \exp(g_t) \odot S_{t-1}$。$g_t$ 必须 $\leq 0$
|
||||||
|
才是衰减($\exp(g_t) \leq 1$),否则状态会指数增长爆炸。
|
||||||
|
|
||||||
|
\texttt{g\_raw} 是从 \texttt{g\_proj(x)} 出来的 raw 值,没有约束。
|
||||||
|
Gate 激活函数的任务是:把 raw 值映射到一个保证 $\leq 0$ 的范围。
|
||||||
|
|
||||||
|
\subsection{两种 Gate}
|
||||||
|
|
||||||
|
\begin{center}
|
||||||
|
\begin{tabular}{p{3cm}p{5.5cm}p{5cm}}
|
||||||
|
\toprule
|
||||||
|
& \textbf{Standard gate} & \textbf{Safe gate} \\
|
||||||
|
\midrule
|
||||||
|
公式 &
|
||||||
|
$g = -\mathrm{rate} \cdot \mathrm{softplus}(\mathrm{input})$ &
|
||||||
|
$g = L \cdot \sigma(\mathrm{rate} \cdot \mathrm{input})$ \\
|
||||||
|
值域 &
|
||||||
|
$(-\infty, 0]$ &
|
||||||
|
$[L, 0]$($L$ 是 lower\_bound,如 $-5$) \\
|
||||||
|
衰减范围 &
|
||||||
|
$\exp(g) \in (0, 1]$ &
|
||||||
|
$\exp(g) \in [\exp(L), 1]$ \\
|
||||||
|
稳定性 &
|
||||||
|
衰减可以任意快 &
|
||||||
|
衰减有下限,不会瞬间清零 \\
|
||||||
|
\bottomrule
|
||||||
|
\end{tabular}
|
||||||
|
\end{center}
|
||||||
|
|
||||||
|
\noindent 其中:
|
||||||
|
\begin{itemize}[nosep]
|
||||||
|
\item $\mathrm{input} = g_{\mathrm{raw}} + \Delta_b$ \quad($\Delta_b$
|
||||||
|
是 \texttt{dt\_bias} \shape{HV, K})
|
||||||
|
\item $\mathrm{rate} = \exp(A_{\log})$ \quad($A_{\log}$ 是
|
||||||
|
\texttt{A\_log} \shape{HV},head-wise 可学习)
|
||||||
|
\end{itemize}
|
||||||
|
|
||||||
|
\begin{importantbox}{如果你只记一件事}
|
||||||
|
Safe gate = $L \cdot \sigma(\mathrm{rate} \cdot \mathrm{input})$,
|
||||||
|
$L=-5$ 时 $\exp(g) \geq \exp(-5) \approx 0.0067$,
|
||||||
|
状态永远不会被``一次性清零''。
|
||||||
|
\end{importantbox}
|
||||||
|
|
||||||
|
\subsection{代码对照}
|
||||||
|
|
||||||
|
\begin{codemathtop}{ops/reference/gate.py — kda\_gate\_reference}
|
||||||
|
\begin{lstlisting}
|
||||||
|
def kda_gate_reference(g, A_log, dt_bias=None, *,
|
||||||
|
safe_gate=False, lower_bound=None):
|
||||||
|
HV, K = g.shape[-2:]
|
||||||
|
gate_input = g if dt_bias is None else g + dt_bias.view(HV, K)
|
||||||
|
rate = A_log.view(HV, 1).exp()
|
||||||
|
|
||||||
|
if safe_gate:
|
||||||
|
# safe: g in [lower_bound, 0]
|
||||||
|
return lower_bound * torch.sigmoid(rate * gate_input)
|
||||||
|
# standard: g in (-inf, 0]
|
||||||
|
return -rate * F.softplus(gate_input)
|
||||||
|
\end{lstlisting}
|
||||||
|
\end{codemathtop}
|
||||||
|
|
||||||
|
\subsection{初始化与默认值}
|
||||||
|
|
||||||
|
\begin{center}
|
||||||
|
\begin{tabular}{llp{7cm}}
|
||||||
|
\toprule
|
||||||
|
参数 & 初始值 & 效果 \\
|
||||||
|
\midrule
|
||||||
|
\texttt{A\_log} & $\mathbf{0}$ \shape{HV} & $\mathrm{rate} = \exp(0) = 1$,不缩放 \\
|
||||||
|
\texttt{dt\_bias} & $-4.0$ \shape{HV, K} & 初始时 $\mathrm{input} \approx g_{\mathrm{raw}} - 4$,
|
||||||
|
配合 safe gate ($L=-5$) 得到 $g \approx -5 \cdot \sigma(-4) \approx -0.09$,
|
||||||
|
即 $\exp(g) \approx 0.91$(约 91\% 状态保留) \\
|
||||||
|
\texttt{lower\_bound} & $-5.0$ & safe gate 的下限 \\
|
||||||
|
\bottomrule
|
||||||
|
\end{tabular}
|
||||||
|
\end{center}
|
||||||
|
|
||||||
|
\subsection{Gate 在 API 中的位置}
|
||||||
|
|
||||||
|
Gate 激活在 \texttt{ops/api.py} 的 \texttt{chunk\_kda} 中调用,
|
||||||
|
在进入 chunkwise 或 recurrent 核心之前完成。
|
||||||
|
当 \texttt{use\_gate\_in\_kernel=True} 时,\texttt{g\_raw} 进入 API,
|
||||||
|
API 内部完成 gate 激活;否则调用方自己完成。
|
||||||
|
|
||||||
|
\begin{lstlisting}
|
||||||
|
# ops/api.py (simplified)
|
||||||
|
if use_gate_in_kernel:
|
||||||
|
gate_input = g + dt_bias.view(g.shape[-2:])
|
||||||
|
rate = A_log.exp().view(1, 1, -1, 1)
|
||||||
|
if safe_gate:
|
||||||
|
g = lower_bound * torch.sigmoid(rate * gate_input)
|
||||||
|
else:
|
||||||
|
g = -rate * F.softplus(gate_input)
|
||||||
|
\end{lstlisting}
|
||||||
|
|
||||||
|
\subsection{本章小结}
|
||||||
|
|
||||||
|
Gate 把 raw projection 映射到 $\leq 0$ 的 log-space 衰减率。
|
||||||
|
Safe gate 用 sigmoid 限制在 $[L, 0]$,防止瞬间清零;
|
||||||
|
standard gate 用 softplus 不限制下限。默认配置下初始状态保留约 91\%。
|
||||||
@@ -0,0 +1,148 @@
|
|||||||
|
% teach:
|
||||||
|
% gap: 读者知道递归形式但不知道怎么在 GPU 上并行, 以为只能一步步算
|
||||||
|
% takeaway: chunk 内用下三角线性系统并行求解, chunk 间递推状态, 数值等价于 naive recurrent
|
||||||
|
% jump: 论文直接写了 triangular solve 但没解释为什么要 solve 而不是直接矩阵乘
|
||||||
|
% omit: Triton 实现细节
|
||||||
|
|
||||||
|
\section{分块并行计算(Chunkwise)}
|
||||||
|
\splabel{C3}
|
||||||
|
|
||||||
|
\subsection{为什么需要分块?}
|
||||||
|
|
||||||
|
Naive recurrent 一步一步算,$T$ 步串行,GPU 利用率低。
|
||||||
|
分块的想法是把序列切成 $T/C$ 个长度为 $C$ 的 chunk:
|
||||||
|
\begin{itemize}[nosep]
|
||||||
|
\item \textbf{chunk 内}:$C$ 个 token 之间的依赖可以用矩阵运算并行处理
|
||||||
|
\item \textbf{chunk 间}:状态 $S$ 从上一个 chunk 传到下一个,仍然是递推
|
||||||
|
\end{itemize}
|
||||||
|
|
||||||
|
\begin{importantbox}{如果你只记一件事}
|
||||||
|
Chunkwise = chunk 内并行 + chunk 间递推。数值结果与 naive recurrent 逐位一致。
|
||||||
|
\end{importantbox}
|
||||||
|
|
||||||
|
\subsection{chunk 内的 cumsum 与下三角解}
|
||||||
|
|
||||||
|
在每个 chunk 内,先对 $g$ 做 cumsum(前缀和),这样衰减就变成了相对距离的函数:
|
||||||
|
\[
|
||||||
|
g_{\mathrm{cum},i} = \sum_{j=0}^{i} g_j, \qquad
|
||||||
|
\text{token } i \text{ 对 token } j \text{ 的衰减} = \exp(g_{\mathrm{cum},i} - g_{\mathrm{cum},j})
|
||||||
|
\]
|
||||||
|
|
||||||
|
定义 \textbf{decayed dot} 矩阵(chunk 内 $C \times C$):
|
||||||
|
\[
|
||||||
|
A_{ij} = \langle x_i, \exp(g_{\mathrm{cum},i} - g_{\mathrm{cum},j}) \cdot k_j \rangle
|
||||||
|
\qquad \shape{..., C, C}
|
||||||
|
\]
|
||||||
|
|
||||||
|
这个矩阵的构造是 chunk 内计算的核心。用它可以构造一个下三角线性系统:
|
||||||
|
\[
|
||||||
|
M = I + \mathrm{tril}(A_{kk} \cdot \beta, \text{diagonal}=-1)
|
||||||
|
\qquad \shape{..., C, C}
|
||||||
|
\]
|
||||||
|
\[
|
||||||
|
M \cdot W = \exp(g_{\mathrm{cum}}) \cdot k \qquad \Rightarrow \qquad
|
||||||
|
W = M^{-1} (\exp(g_{\mathrm{cum}}) \cdot k)
|
||||||
|
\]
|
||||||
|
\[
|
||||||
|
M \cdot U = v \qquad \Rightarrow \qquad U = M^{-1} v
|
||||||
|
\]
|
||||||
|
|
||||||
|
\subsection{代码对照}
|
||||||
|
|
||||||
|
\begin{codemathtop}{ops/reference/chunkwise.py — naive\_chunk\_kda (核心)}
|
||||||
|
\begin{lstlisting}
|
||||||
|
# Rearrange: [B,T,H,K] -> [B,H,N,C,K] where N=T/C
|
||||||
|
q, k = [rearrange(x, 'b (n c) h d -> b h n c d', c=C)
|
||||||
|
.repeat_interleave(HV//H, dim=1) for x in (q, k)]
|
||||||
|
v, g = [rearrange(x, 'b (n c) h d -> b h n c d', c=C)
|
||||||
|
for x in (v, g)]
|
||||||
|
beta = rearrange(beta, 'b (n c) h -> b h n c', c=C)
|
||||||
|
q = q * scale
|
||||||
|
g = g.cumsum(dim=-2) # chunk 内 cumsum
|
||||||
|
|
||||||
|
# Construct triangular system
|
||||||
|
A_kk = _decayed_dot(k, k, g) # [B,HV,N,C,C]
|
||||||
|
M = eye + (A_kk * beta[...,None,:]).masked_fill(mask_upper, 0)
|
||||||
|
W = solve_triangular(M, g.exp() * k, upper=False)
|
||||||
|
U = solve_triangular(M, v, upper=False)
|
||||||
|
|
||||||
|
# A_qk: query 对 key 的 decayed dot (含对角线)
|
||||||
|
A_qk = (_decayed_dot(q, k, g) * beta[...,None,:])
|
||||||
|
.masked_fill(mask_strict_upper, 0)
|
||||||
|
\end{lstlisting}
|
||||||
|
\end{codemathtop}
|
||||||
|
|
||||||
|
\subsection{chunk 间递推}
|
||||||
|
|
||||||
|
每个 chunk 内算完后,用 $W$ 和 $U$ 来处理跨 chunk 的状态:
|
||||||
|
|
||||||
|
\begin{codemathtop}{ops/reference/chunkwise.py — chunk 间循环}
|
||||||
|
\begin{lstlisting}
|
||||||
|
S = zeros(B, HV, K, V) # inter-chunk state
|
||||||
|
for n in range(T // C):
|
||||||
|
# r = "local residual, adjusted by cross-chunk state"
|
||||||
|
r = U[:,:,n] - W[:,:,n] @ S # [B,HV,C,V]
|
||||||
|
|
||||||
|
# output: cross-chunk part + intra-chunk part
|
||||||
|
o[:,:,n] = (q_n * g_n.exp()) @ S + A_qk[:,:,n] @ r
|
||||||
|
|
||||||
|
# update cross-chunk state
|
||||||
|
decay = (g_n[:,:,-1:,:] - g_n).exp() # decay to chunk end
|
||||||
|
S = S * g_n[:,:,-1,:,None].exp() # decay old state
|
||||||
|
S = S + (decay * k_n).T @ (r * beta_n) # write new
|
||||||
|
\end{lstlisting}
|
||||||
|
\end{codemathtop}
|
||||||
|
|
||||||
|
\subsection{形状流水线}
|
||||||
|
|
||||||
|
\begin{center}
|
||||||
|
\begin{tabular}{lll}
|
||||||
|
\toprule
|
||||||
|
变量 & 形状 & 说明 \\
|
||||||
|
\midrule
|
||||||
|
\texttt{q, k} (chunked) & \shape{B, HV, N, C, K} & $N = T/C$ 个 chunk \\
|
||||||
|
\texttt{v, g} (chunked) & \shape{B, HV, N, C, V/K} & \\
|
||||||
|
\texttt{beta} (chunked) & \shape{B, HV, N, C} & \\
|
||||||
|
\texttt{A\_kk} & \shape{B, HV, N, C, C} & key-key decayed dot \\
|
||||||
|
\texttt{M} & \shape{B, HV, N, C, C} & 下三角系统 \\
|
||||||
|
\texttt{W} & \shape{B, HV, N, C, K} & $M^{-1}(\exp(g) \cdot k)$ \\
|
||||||
|
\texttt{U} & \shape{B, HV, N, C, V} & $M^{-1} v$ \\
|
||||||
|
\texttt{A\_qk} & \shape{B, HV, N, C, C} & query-key decayed dot \\
|
||||||
|
\texttt{S} & \shape{B, HV, K, V} & 跨 chunk 状态 \\
|
||||||
|
\texttt{r} & \shape{B, HV, C, V} & 调整后的残差 \\
|
||||||
|
\texttt{o (chunk n)} & \shape{B, HV, C, V} & 本 chunk 输出 \\
|
||||||
|
\bottomrule
|
||||||
|
\end{tabular}
|
||||||
|
\end{center}
|
||||||
|
|
||||||
|
\begin{warningbox}{为什么用 triangular solve 而不是直接矩阵乘?}
|
||||||
|
Delta rule 的"先擦再写"引入了 chunk 内 token 之间的递归依赖:
|
||||||
|
token $i$ 的写入依赖 token $j < i$ 的写入结果。
|
||||||
|
这个依赖关系恰好形成一个下三角线性系统 $M \cdot x = b$,
|
||||||
|
用 \texttt{solve\_triangular} 可以在 $O(C^2)$ 内并行求解,
|
||||||
|
而展开递归需要 $O(C)$ 步串行。
|
||||||
|
\end{warningbox}
|
||||||
|
|
||||||
|
\subsection{Decayed dot 函数}
|
||||||
|
|
||||||
|
\begin{codemathtop}{ops/reference/chunkwise.py — \_decayed\_dot}
|
||||||
|
\begin{lstlisting}
|
||||||
|
def _decayed_dot(x, k, g):
|
||||||
|
"""A[..., i, j] = <x_i, exp(g_i - g_j) * k_j>"""
|
||||||
|
C = x.shape[-2]
|
||||||
|
A = empty(*x.shape[:-2], C, C)
|
||||||
|
for i in range(C):
|
||||||
|
decay = (g[..., i:i+1, :] - g).exp() # [.., 1, K] - [.., C, K]
|
||||||
|
A[..., i, :] = einsum('...jk,...jk->...j',
|
||||||
|
x[..., i, None, :] * decay, k)
|
||||||
|
return A
|
||||||
|
\end{lstlisting}
|
||||||
|
\end{codemathtop}
|
||||||
|
|
||||||
|
\noindent 这是一个 $C \times C$ 的矩阵,每个元素 $(i,j)$ 是
|
||||||
|
$x_i$ 和 $\exp(g_i - g_j) \cdot k_j$ 的内积。Triton 实现会把这个双循环融合成一个 kernel。
|
||||||
|
|
||||||
|
\subsection{本章小结}
|
||||||
|
|
||||||
|
分块把 $T$ 步串行拆成 $T/C$ 个 chunk,chunk 内用下三角 solve 并行处理 delta rule 依赖,
|
||||||
|
chunk 间递推状态 $S$。最终输出与 naive recurrent 逐位相同。
|
||||||
@@ -0,0 +1,114 @@
|
|||||||
|
% teach:
|
||||||
|
% gap: 读者不知道 q/k 和 v 为什么可以有不同的头数, 以及 repeat_interleave 的反传怎么做
|
||||||
|
% takeaway: GVA 让 G 组 value heads 共享一组 q/k, forward repeat_interleave, backward sum
|
||||||
|
% jump: 论文没解释为什么反传是 sum 而不是 mean
|
||||||
|
% omit: GQA 的历史
|
||||||
|
|
||||||
|
\section{GVA(分组值注意力)}
|
||||||
|
\splabel{GVA}
|
||||||
|
|
||||||
|
\subsection{为什么头数不一样?}
|
||||||
|
|
||||||
|
标准 MHA 里 $H_q = H_k = H_v$。GQA(Grouped Query Attention)让多组 q/k 共享同一组 v/k,
|
||||||
|
减少 KV cache。KDA 反过来做:$H$ 组 q/k 对应 $H_V = G \cdot H$ 组 value heads。
|
||||||
|
|
||||||
|
直觉:value 维度决定表达能力,多一点 value head 增加容量;
|
||||||
|
q/k 主要负责路由(``看哪里''),可以共享。
|
||||||
|
|
||||||
|
\begin{center}
|
||||||
|
\begin{tabular}{lll}
|
||||||
|
\toprule
|
||||||
|
& 标准头数 & GVA \\
|
||||||
|
\midrule
|
||||||
|
$q, k$ & \shape{B, T, H, K} & \shape{B, T, H, K}(不变)\\
|
||||||
|
$v$ & \shape{B, T, H, V} & \shape{B, T, HV, V}($H_V = G \cdot H$) \\
|
||||||
|
$g, \beta$ & \shape{B, T, H, K/1} & \shape{B, T, HV, K/1} \\
|
||||||
|
$S$ & \shape{B, H, K, V} & \shape{B, HV, K, V} \\
|
||||||
|
\bottomrule
|
||||||
|
\end{tabular}
|
||||||
|
\end{center}
|
||||||
|
|
||||||
|
\subsection{Forward: repeat\_interleave}
|
||||||
|
|
||||||
|
进入 KDA 核心前,$q$ 和 $k$ 从 $H$ 维复制到 $H_V$ 维:
|
||||||
|
|
||||||
|
\begin{lstlisting}
|
||||||
|
G = HV // H
|
||||||
|
qe = q.repeat_interleave(G, dim=2) * scale # [B,T,H,K] -> [B,T,HV,K]
|
||||||
|
ke = k.repeat_interleave(G, dim=2) # [B,T,H,K] -> [B,T,HV,K]
|
||||||
|
\end{lstlisting}
|
||||||
|
|
||||||
|
\noindent 例如 $H=4, G=2, H_V=8$:head 0 的 q/k 复制到 value head 0 和 1,
|
||||||
|
head 1 复制到 value head 2 和 3,依此类推。
|
||||||
|
|
||||||
|
\subsection{Backward: view + sum}
|
||||||
|
|
||||||
|
反传时,$dq_e$ 和 $dk_e$ 的形状是 \shape{B, T, HV, K}(在 $H_V$ 维上计算的梯度)。
|
||||||
|
因为 forward 是复制,反传就是求和:
|
||||||
|
|
||||||
|
\begin{lstlisting}
|
||||||
|
# Backward: HV -> H
|
||||||
|
dq_H = dq_e.view(B, T, H, G, K).sum(dim=3) # [B,T,HV,K] -> [B,T,H,K]
|
||||||
|
dk_H = dk_e.view(B, T, H, G, K).sum(dim=3)
|
||||||
|
\end{lstlisting}
|
||||||
|
|
||||||
|
\begin{warningbox}{为什么是 sum 不是 mean?}
|
||||||
|
\texttt{repeat\_interleave} 是\textbf{复制}:$y_0 = x_0, y_1 = x_0, y_2 = x_1, \ldots$
|
||||||
|
|
||||||
|
对 $x_0$ 的梯度 = $\frac{\partial L}{\partial y_0} + \frac{\partial L}{\partial y_1}$
|
||||||
|
= \textbf{sum}(不是 mean)。
|
||||||
|
|
||||||
|
这和 \texttt{.expand()} 的反传一样:复制的反传是求和。
|
||||||
|
\end{warningbox}
|
||||||
|
|
||||||
|
\subsection{scale 的处理}
|
||||||
|
|
||||||
|
$q$ 在 repeat\_interleave 之后乘了 \texttt{scale = $1/\sqrt{K}$}。
|
||||||
|
反传时 chain rule 要求 $dq_{\mathrm{orig}} = dq_e \cdot \texttt{scale}$:
|
||||||
|
|
||||||
|
\begin{lstlisting}
|
||||||
|
# q 在 forward 内被乘过 scale, chain rule:
|
||||||
|
dq_H = dq_H * scale
|
||||||
|
\end{lstlisting}
|
||||||
|
|
||||||
|
\subsection{KDAAttention 层中的投影}
|
||||||
|
|
||||||
|
\begin{codemathtop}{layers/kda\_attn.py — forward}
|
||||||
|
\begin{lstlisting}
|
||||||
|
def forward(self, x): # x: [B, T, D]
|
||||||
|
B, T, _ = x.shape
|
||||||
|
H, HV, K, V = self.num_heads, self.num_value_heads, ...
|
||||||
|
|
||||||
|
q = self.q_proj(x).view(B, T, H, K) # [B,T,D] -> [B,T,H*K] -> [B,T,H,K]
|
||||||
|
k = self.k_proj(x).view(B, T, H, K) # 同上
|
||||||
|
v = self.v_proj(x).view(B, T, HV, V) # [B,T,D] -> [B,T,HV*V] -> [B,T,HV,V]
|
||||||
|
g_raw = self.g_proj(x).view(B, T, HV, K)
|
||||||
|
beta_raw = self.beta_proj(x).view(B, T, HV)
|
||||||
|
|
||||||
|
o, _ = chunk_kda(q, k, v, g_raw, beta_raw, ...)
|
||||||
|
return self.o_proj(o.reshape(B, T, HV * V)) # [B,T,HV,V] -> [B,T,D]
|
||||||
|
\end{lstlisting}
|
||||||
|
\end{codemathtop}
|
||||||
|
|
||||||
|
\subsection{投影矩阵形状总览}
|
||||||
|
|
||||||
|
\begin{center}
|
||||||
|
\begin{tabular}{llll}
|
||||||
|
\toprule
|
||||||
|
投影 & 权重形状 & 输入 & 输出 \\
|
||||||
|
\midrule
|
||||||
|
\texttt{q\_proj} & \shape{H \cdot K, D} & \shape{B,T,D} & \shape{B,T,H,K} \\
|
||||||
|
\texttt{k\_proj} & \shape{H \cdot K, D} & \shape{B,T,D} & \shape{B,T,H,K} \\
|
||||||
|
\texttt{v\_proj} & \shape{HV \cdot V, D} & \shape{B,T,D} & \shape{B,T,HV,V} \\
|
||||||
|
\texttt{g\_proj} & \shape{HV \cdot K, D} & \shape{B,T,D} & \shape{B,T,HV,K} \\
|
||||||
|
\texttt{beta\_proj} & \shape{HV, D} & \shape{B,T,D} & \shape{B,T,HV} \\
|
||||||
|
\texttt{o\_proj} & \shape{D, HV \cdot V} & \shape{B,T,HV \cdot V} & \shape{B,T,D} \\
|
||||||
|
\bottomrule
|
||||||
|
\end{tabular}
|
||||||
|
\end{center}
|
||||||
|
|
||||||
|
\subsection{本章小结}
|
||||||
|
|
||||||
|
GVA 让 $H_V = G \cdot H$ 组 value heads 共享 $H$ 组 q/k。
|
||||||
|
Forward 用 \texttt{repeat\_interleave} 复制,backward 用 \texttt{view+sum} 归约。
|
||||||
|
$v, g, \beta$ 直接在 $H_V$ 维投影,q/k 在 $H$ 维投影。
|
||||||
@@ -0,0 +1,75 @@
|
|||||||
|
% teach:
|
||||||
|
% gap: 读者已知各组件, 但不清楚它们怎么黏在一起成为一个层
|
||||||
|
% takeaway: KDAAttention = 投影 → gate+norm → chunk_kda → output 投影, 整个层就是 x → y [B,T,D]
|
||||||
|
% jump: none
|
||||||
|
% omit: from_config 工厂方法细节
|
||||||
|
|
||||||
|
\section{KDAAttention 层}
|
||||||
|
|
||||||
|
\subsection{完整数据流}
|
||||||
|
|
||||||
|
\texttt{KDAAttention} 把投影、gate 激活、KDA 核心计算和输出投影封装成一个
|
||||||
|
\texttt{[B,T,D] $\to$ [B,T,D]} 的模块。
|
||||||
|
|
||||||
|
\begin{center}
|
||||||
|
\begin{tabular}{rlll}
|
||||||
|
\toprule
|
||||||
|
步骤 & 操作 & 输入形状 & 输出形状 \\
|
||||||
|
\midrule
|
||||||
|
1 & \texttt{q\_proj(x)} & \shape{B,T,D} & \shape{B,T,H,K} \\
|
||||||
|
2 & \texttt{k\_proj(x)} & \shape{B,T,D} & \shape{B,T,H,K} \\
|
||||||
|
3 & \texttt{v\_proj(x)} & \shape{B,T,D} & \shape{B,T,HV,V} \\
|
||||||
|
4 & \texttt{g\_proj(x)} & \shape{B,T,D} & \shape{B,T,HV,K} \\
|
||||||
|
5 & \texttt{beta\_proj(x)} & \shape{B,T,D} & \shape{B,T,HV} \\
|
||||||
|
6 & \texttt{chunk\_kda(...)} & 上述 5 项 + 参数 & \shape{B,T,HV,V} \\
|
||||||
|
7 & \texttt{o.reshape(...)} & \shape{B,T,HV,V} & \shape{B,T,HV \cdot V} \\
|
||||||
|
8 & \texttt{o\_proj(...)} & \shape{B,T,HV \cdot V} & \shape{B,T,D} \\
|
||||||
|
\bottomrule
|
||||||
|
\end{tabular}
|
||||||
|
\end{center}
|
||||||
|
|
||||||
|
\subsection{chunk\_kda 内部做了什么}
|
||||||
|
|
||||||
|
\texttt{chunk\_kda}(\texttt{ops/api.py})是统一入口,按 \texttt{backend} 分发:
|
||||||
|
|
||||||
|
\begin{enumerate}[nosep]
|
||||||
|
\item 如果 \texttt{use\_qk\_l2norm\_in\_kernel}:$q, k \leftarrow \text{L2-normalize}(q), \text{L2-normalize}(k)$
|
||||||
|
\item 如果 \texttt{use\_beta\_sigmoid\_in\_kernel}:$\beta \leftarrow \sigma(\beta_{\mathrm{raw}})$
|
||||||
|
\item 如果 \texttt{use\_gate\_in\_kernel}:应用 gate 激活(§2)
|
||||||
|
\item 调用 \texttt{naive\_chunk\_kda}(或 triton/fla 版本)
|
||||||
|
\end{enumerate}
|
||||||
|
|
||||||
|
\begin{knowledgebox}{三个 ``in\_kernel'' 开关}
|
||||||
|
\begin{itemize}[nosep]
|
||||||
|
\item \texttt{use\_qk\_l2norm}:L2-norm 让 $\langle q, k \rangle$ 变成余弦相似度,
|
||||||
|
稳定训练
|
||||||
|
\item \texttt{use\_beta\_sigmoid}:sigmoid 把 $\beta$ 限制在 $(0,1)$,
|
||||||
|
控制写入强度
|
||||||
|
\item \texttt{use\_gate\_in\_kernel}:gate 激活在 API 内部完成(vs 调用方自己做)
|
||||||
|
\end{itemize}
|
||||||
|
默认三个都是 \texttt{True}。
|
||||||
|
\end{knowledgebox}
|
||||||
|
|
||||||
|
\subsection{可学习参数清单}
|
||||||
|
|
||||||
|
\begin{center}
|
||||||
|
\begin{tabular}{lll}
|
||||||
|
\toprule
|
||||||
|
参数 & 形状 & 说明 \\
|
||||||
|
\midrule
|
||||||
|
\texttt{q\_proj.weight} & \shape{H \cdot K, D} & query 投影 \\
|
||||||
|
\texttt{k\_proj.weight} & \shape{H \cdot K, D} & key 投影 \\
|
||||||
|
\texttt{v\_proj.weight} & \shape{HV \cdot V, D} & value 投影 \\
|
||||||
|
\texttt{g\_proj.weight} & \shape{HV \cdot K, D} & gate 投影 \\
|
||||||
|
\texttt{beta\_proj.weight} & \shape{HV, D} & beta 投影 \\
|
||||||
|
\texttt{o\_proj.weight} & \shape{D, HV \cdot V} & 输出投影 \\
|
||||||
|
\texttt{A\_log} & \shape{HV} & head-wise 衰减率(log-space)\\
|
||||||
|
\texttt{dt\_bias} & \shape{HV, K} & per-dim gate bias \\
|
||||||
|
\bottomrule
|
||||||
|
\end{tabular}
|
||||||
|
\end{center}
|
||||||
|
|
||||||
|
\subsection{本章小结}
|
||||||
|
|
||||||
|
KDAAttention 是一个完整的 mixing 模块:5 个线性投影 + gate 激活 + KDA 核心 + 输出投影。
|
||||||
|
三个 ``in\_kernel'' 开关控制 L2-norm、sigmoid、gate 是否在 API 内部完成。
|
||||||
@@ -0,0 +1,173 @@
|
|||||||
|
% teach:
|
||||||
|
% gap: 读者知道标准 MHA 但不知道 MLA 怎么压缩 KV、矩阵吸收怎么避免解压
|
||||||
|
% takeaway: MLA 把 KV 压成低秩 latent c, 通过吸收 W_UK 进 q 直接在 latent 空间算 attention
|
||||||
|
% jump: 为什么可以先在 latent 加权再乘 W_UV?因为矩阵乘和加权求和可交换
|
||||||
|
% omit: RoPE (K3 用 NoPE)
|
||||||
|
|
||||||
|
\section{Gated MLA(矩阵吸收版)}
|
||||||
|
\splabel{C4}
|
||||||
|
|
||||||
|
\subsection{标准 MHA 的 KV cache 问题}
|
||||||
|
|
||||||
|
标准 MHA 推理时需要缓存所有历史 token 的 $K, V$,cache 大小 $\propto T \cdot H \cdot d$。
|
||||||
|
MLA 的想法:把 $K, V$ 压缩成一个低秩 latent $c$,cache 大小 $\propto T \cdot r$,
|
||||||
|
其中 $r \ll H \cdot d$。
|
||||||
|
|
||||||
|
\subsection{低秩压缩}
|
||||||
|
|
||||||
|
\[
|
||||||
|
c = \mathrm{RMSNorm}(W_{\downarrow} \cdot x) \qquad \shape{B, T, r}
|
||||||
|
\]
|
||||||
|
|
||||||
|
推理时只缓存 $c$,不缓存解压后的 $K, V$。
|
||||||
|
解压矩阵 $W_{\mathrm{KV}\uparrow}$ 包含两部分:
|
||||||
|
\[
|
||||||
|
W_{\mathrm{KV}\uparrow} = \begin{bmatrix} W_{UK} \\ W_{UV} \end{bmatrix}
|
||||||
|
\qquad \shape{H \cdot (d_q + d_v), r}
|
||||||
|
\]
|
||||||
|
拆开:$W_{UK} \in \mathbb{R}^{H \times d_q \times r}$(key 解压),
|
||||||
|
$W_{UV} \in \mathbb{R}^{H \times d_v \times r}$(value 解压)。
|
||||||
|
|
||||||
|
\subsection{矩阵吸收的核心思路}
|
||||||
|
|
||||||
|
\textbf{不解压} $K$ 和 $V$。标准做法会先解压再算 attention:
|
||||||
|
|
||||||
|
\begin{center}
|
||||||
|
\textit{标准}:$k_h = c \cdot W_{UK,h}^T$ \shape{B,T,d_q},
|
||||||
|
$\mathrm{score} = q_h \cdot k_h^T$
|
||||||
|
\end{center}
|
||||||
|
|
||||||
|
矩阵吸收反过来:把 $W_{UK}$ 吸收进 $q$:
|
||||||
|
|
||||||
|
\begin{center}
|
||||||
|
\textit{吸收}:$q_{\mathrm{abs},h} = q_h \cdot W_{UK,h}$ \shape{B,T,r},
|
||||||
|
$\mathrm{score} = q_{\mathrm{abs},h} \cdot c^T$
|
||||||
|
\end{center}
|
||||||
|
|
||||||
|
\begin{importantbox}{如果你只记一件事}
|
||||||
|
$(q \cdot W_{UK}^T) \cdot c^T = q \cdot (W_{UK}^T \cdot c^T) = q_{\mathrm{abs}} \cdot c^T$
|
||||||
|
|
||||||
|
吸收后,attention 直接在 latent 空间 $r$ 维上算,永不解压到 $H \cdot d_q$ 维。
|
||||||
|
\end{importantbox}
|
||||||
|
|
||||||
|
\subsection{完整计算流(四步)}
|
||||||
|
|
||||||
|
\begin{enumerate}[leftmargin=2em]
|
||||||
|
\item \textbf{Q 低秩路径}(NoPE,只有 nope 段):
|
||||||
|
\[
|
||||||
|
q = W_{q\uparrow} \cdot \mathrm{RMSNorm}(W_{q\downarrow} \cdot x)
|
||||||
|
\qquad \shape{B, T, H, d_q}
|
||||||
|
\]
|
||||||
|
|
||||||
|
\item \textbf{吸收 $W_{UK}$ + 打分}:
|
||||||
|
\[
|
||||||
|
q_{\mathrm{abs}} = q \cdot W_{UK} \quad
|
||||||
|
\xrightarrow{\texttt{einsum('bthd,hdj->bthj')}} \quad \shape{B, T, H, r}
|
||||||
|
\]
|
||||||
|
\[
|
||||||
|
\mathrm{score} = q_{\mathrm{abs}} \cdot c^T \quad
|
||||||
|
\xrightarrow{\texttt{einsum('bthj,bsj->bhts')}} \quad \shape{B, H, T, T}
|
||||||
|
\]
|
||||||
|
\[
|
||||||
|
\mathrm{attn} = \mathrm{softmax}(\mathrm{causal\_mask}(\mathrm{score}))
|
||||||
|
\qquad \shape{B, H, T, T}
|
||||||
|
\]
|
||||||
|
|
||||||
|
\item \textbf{先在 latent 加权,再乘 $W_{UV}^T$}:
|
||||||
|
\[
|
||||||
|
\tilde{o}_{\mathrm{lat}} = \mathrm{attn} \cdot c \quad
|
||||||
|
\xrightarrow{\texttt{einsum('bhts,bsj->bhtj')}} \quad \shape{B, H, T, r}
|
||||||
|
\]
|
||||||
|
\[
|
||||||
|
\tilde{o} = \tilde{o}_{\mathrm{lat}} \cdot W_{UV}^T \quad
|
||||||
|
\xrightarrow{\texttt{einsum('bhtj,hvj->bhtv')}} \quad \shape{B, H, T, d_v}
|
||||||
|
\]
|
||||||
|
|
||||||
|
\item \textbf{输出门 + 投影}:
|
||||||
|
\[
|
||||||
|
y = W_o \big[ \sigma(W_g \cdot x) \odot \tilde{o}_{\mathrm{flat}} \big]
|
||||||
|
\qquad \shape{B, T, D}
|
||||||
|
\]
|
||||||
|
\end{enumerate}
|
||||||
|
|
||||||
|
\subsection{代码对照}
|
||||||
|
|
||||||
|
\begin{codemathtop}{layers/mla.py — GatedMLA.forward}
|
||||||
|
\begin{lstlisting}
|
||||||
|
def forward(self, x): # x: [B, T, D]
|
||||||
|
B, T, _ = x.shape
|
||||||
|
H, r = self.num_heads, self.kv_up.in_features
|
||||||
|
|
||||||
|
# Step 1: latent + query
|
||||||
|
c = self.kv_norm(self.kv_down(x)) # [B, T, r]
|
||||||
|
q = self.q_up(self.q_norm(self.q_down(x))) # [B, T, H*d_q]
|
||||||
|
q = q.view(B, T, H, self.qk_nope_head_dim) # [B, T, H, d_q]
|
||||||
|
|
||||||
|
# Split W_UK, W_UV from kv_up.weight
|
||||||
|
w = self.kv_up.weight # [H*(d_q+d_v), r]
|
||||||
|
w_uk = w[:H*d_q].view(H, d_q, r) # [H, d_q, r]
|
||||||
|
w_uv = w[H*d_q:].view(H, d_v, r) # [H, d_v, r]
|
||||||
|
|
||||||
|
# Step 2: absorb W_UK, score
|
||||||
|
q_absorb = einsum('bthd,hdj->bthj', q, w_uk) # [B,T,H,r]
|
||||||
|
scores = einsum('bthj,bsj->bhts', q_absorb, c) # [B,H,T,T]
|
||||||
|
scores = scores.masked_fill(causal_mask, -inf)
|
||||||
|
attn = softmax(scores, dim=-1) # [B,H,T,T]
|
||||||
|
|
||||||
|
# Step 3: latent-space weighted sum, then W_UV
|
||||||
|
latent_out = einsum('bhts,bsj->bhtj', attn, c) # [B,H,T,r]
|
||||||
|
o_heads = einsum('bhtj,hvj->bhtv', latent_out, w_uv) # [B,H,T,d_v]
|
||||||
|
|
||||||
|
# Step 4: output gate
|
||||||
|
o_heads = o_heads.transpose(1,2).reshape(B,T, H*d_v)
|
||||||
|
gate = sigmoid(self.gate(x)) # [B,T,H*d_v]
|
||||||
|
return self.o_proj(gate * o_heads) # [B,T,D]
|
||||||
|
\end{lstlisting}
|
||||||
|
\end{codemathtop}
|
||||||
|
|
||||||
|
\begin{warningbox}{为什么可以先加权再乘 $W_{UV}$?}
|
||||||
|
标准做法:$o = \mathrm{attn} \cdot V = \mathrm{attn} \cdot (c \cdot W_{UV}^T)$
|
||||||
|
|
||||||
|
交换顺序:$o = (\mathrm{attn} \cdot c) \cdot W_{UV}^T$
|
||||||
|
|
||||||
|
这能成立是因为矩阵乘法的结合律:$A(BC) = (AB)C$。
|
||||||
|
$\mathrm{attn} \cdot c$ 先在 latent 空间 $r$ 维上加权求和,
|
||||||
|
得到的 \shape{B,H,T,r} 再乘 $W_{UV}^T$ 还原到 $d_v$ 维。
|
||||||
|
全程不需要显式构造 $H \cdot T$ 大小的 $V$ 矩阵。
|
||||||
|
\end{warningbox}
|
||||||
|
|
||||||
|
\subsection{形状与参数对比}
|
||||||
|
|
||||||
|
\begin{center}
|
||||||
|
\begin{tabular}{lll}
|
||||||
|
\toprule
|
||||||
|
参数 & 形状 & 说明 \\
|
||||||
|
\midrule
|
||||||
|
\texttt{kv\_down.weight} & \shape{r, D} & KV latent 压缩 \\
|
||||||
|
\texttt{kv\_up.weight} & \shape{H \cdot (d_q+d_v), r} & 包含 $W_{UK}$ 和 $W_{UV}$ \\
|
||||||
|
\texttt{q\_down.weight} & \shape{r_q, D} & Q 低秩 \\
|
||||||
|
\texttt{q\_up.weight} & \shape{H \cdot d_q, r_q} & Q 解压 \\
|
||||||
|
\texttt{gate.weight} & \shape{H \cdot d_v, D} & 输出门 \\
|
||||||
|
\texttt{o\_proj.weight} & \shape{D, H \cdot d_v} & 输出投影 \\
|
||||||
|
\bottomrule
|
||||||
|
\end{tabular}
|
||||||
|
\end{center}
|
||||||
|
|
||||||
|
\noindent KV cache 大小对比(推理时):
|
||||||
|
|
||||||
|
\begin{center}
|
||||||
|
\begin{tabular}{ll}
|
||||||
|
\toprule
|
||||||
|
方法 & Cache 大小 per token \\
|
||||||
|
\midrule
|
||||||
|
标准 MHA & $2 \times H \times d = 2 H d$ \\
|
||||||
|
MLA (latent) & $r$(只存 $c$) \\
|
||||||
|
\bottomrule
|
||||||
|
\end{tabular}
|
||||||
|
\end{center}
|
||||||
|
|
||||||
|
\subsection{本章小结}
|
||||||
|
|
||||||
|
Gated MLA 把 KV 压缩到低秩 latent $c$ \shape{B,T,r},通过矩阵吸收
|
||||||
|
($q_{\mathrm{abs}} = q \cdot W_{UK}$)直接在 latent 空间打分和加权,
|
||||||
|
永不解压 K/V。输出通过 sigmoid 门控。NoPE:不使用 RoPE,位置感交给夹层 KDA 的 decay/gate。
|
||||||
@@ -0,0 +1,158 @@
|
|||||||
|
% teach:
|
||||||
|
% gap: 读者知道 MoE 的 top-k 路由但不知道 LatentMoE 的 latent 接口和 SiTU-GLU
|
||||||
|
% takeaway: LatentMoE 通过 latent 接口把 routed 专家限制在 ℓ=d/2 上算, SiTU-GLU 用软上限防溢出
|
||||||
|
% jump: 为什么 routed 专家在 latent 空间而 shared 在全宽?省参数
|
||||||
|
% omit: load balancing loss
|
||||||
|
|
||||||
|
\section{SiTU-GLU 与 Stable LatentMoE}
|
||||||
|
\splabel{C5}
|
||||||
|
|
||||||
|
\subsection{SiTU-GLU:带软上限的激活}
|
||||||
|
|
||||||
|
SwiGLU 在低精度(fp16/bf16)训练时可能溢出:$\mathrm{silu}(x) \cdot x$ 没有上限。
|
||||||
|
SiTU-GLU 用 $\tanh$ 给门控和上投影加软上限:
|
||||||
|
|
||||||
|
\[
|
||||||
|
\mathrm{SiTU}(x) = W_o \big[\underbrace{\beta_1 \tanh\!\left(\frac{W_g x}{\beta_1}\right) \cdot \sigma(W_g x)}_{\text{gate}} \;\cdot\; \underbrace{\beta_2 \tanh\!\left(\frac{W_u x}{\beta_2}\right)}_{\text{up}}\big]
|
||||||
|
\]
|
||||||
|
|
||||||
|
\begin{center}
|
||||||
|
\begin{tabular}{lp{8cm}}
|
||||||
|
\toprule
|
||||||
|
性质 & 说明 \\
|
||||||
|
\midrule
|
||||||
|
输出上限 & $\|\mathrm{SiTU}\|_\infty \leq \beta_1 \cdot \beta_2 = 4 \times 25 = 100$ \\
|
||||||
|
原点附近 & $\tanh(x/\beta) \approx x/\beta$,所以 $\beta \cdot \tanh(x/\beta) \approx x$,退化为 SwiGLU \\
|
||||||
|
远端 & 软饱和,防 fp16 溢出 \\
|
||||||
|
\bottomrule
|
||||||
|
\end{tabular}
|
||||||
|
\end{center}
|
||||||
|
|
||||||
|
\begin{codemathtop}{layers/latent\_moe.py — SiTU}
|
||||||
|
\begin{lstlisting}
|
||||||
|
class SiTU(nn.Module):
|
||||||
|
def __init__(self, dim_in, dim_ff, beta1=4.0, beta2=25.0):
|
||||||
|
self.w_g = nn.Linear(dim_in, dim_ff, bias=False)
|
||||||
|
self.w_u = nn.Linear(dim_in, dim_ff, bias=False)
|
||||||
|
self.w_o = nn.Linear(dim_ff, dim_in, bias=False)
|
||||||
|
|
||||||
|
def forward(self, x): # [*, dim_in]
|
||||||
|
wg = self.w_g(x)
|
||||||
|
g = self.beta1 * tanh(wg / self.beta1) * sigmoid(wg) # gate
|
||||||
|
u = self.beta2 * tanh(self.w_u(x) / self.beta2) # up
|
||||||
|
return self.w_o(g * u) # [*, dim_in]
|
||||||
|
\end{lstlisting}
|
||||||
|
\end{codemathtop}
|
||||||
|
|
||||||
|
\subsection{LatentMoE 架构}
|
||||||
|
|
||||||
|
\begin{center}
|
||||||
|
\begin{tabular}{rl}
|
||||||
|
\toprule
|
||||||
|
组件 & 说明 \\
|
||||||
|
\midrule
|
||||||
|
\textbf{Shared 专家} & $n_{\mathrm{shared}}$ 个 SiTU,全宽 $d \to d$,所有 token 都经过 \\
|
||||||
|
\textbf{Routed 专家} & $n_{\mathrm{routed}}$ 个 SiTU,半宽 $\ell \to \ell$($\ell = d/2$) \\
|
||||||
|
\textbf{Latent 接口} & $W_\downarrow: d \to \ell$, $W_\uparrow: \ell \to d$(压缩/还原) \\
|
||||||
|
\textbf{Router} & $W_r: d \to n_{\mathrm{routed}}$,Top-k 选择 + softmax 归一化 \\
|
||||||
|
\bottomrule
|
||||||
|
\end{tabular}
|
||||||
|
\end{center}
|
||||||
|
|
||||||
|
\subsection{计算流(五步)}
|
||||||
|
|
||||||
|
\begin{enumerate}[leftmargin=2em]
|
||||||
|
\item \textbf{Latent 压缩}:
|
||||||
|
\[
|
||||||
|
z = W_\downarrow \cdot x \qquad \shape{B, T, \ell}
|
||||||
|
\]
|
||||||
|
|
||||||
|
\item \textbf{Routing}:
|
||||||
|
\[
|
||||||
|
\mathrm{logits} = W_r \cdot x \qquad \shape{B, T, n_{\mathrm{routed}}}
|
||||||
|
\]
|
||||||
|
\[
|
||||||
|
\mathrm{ids}, \mathrm{probs} = \mathrm{TopK}(\mathrm{logits}, k)
|
||||||
|
\qquad \mathrm{ids}: \shape{B, T, k}, \;\; \mathrm{probs}: \shape{B, T, k}
|
||||||
|
\]
|
||||||
|
|
||||||
|
\item \textbf{Routed 专家}(在 latent 空间 $\ell$ 上):
|
||||||
|
\[
|
||||||
|
u = \sum_{i \in \mathrm{Top\text{-}k}} p_i \cdot E_i^{\mathrm{rt}}(z)
|
||||||
|
\qquad \shape{B, T, \ell}
|
||||||
|
\]
|
||||||
|
|
||||||
|
\item \textbf{Shared 专家}(全宽 $d$):
|
||||||
|
\[
|
||||||
|
s = \sum_j E_j^{\mathrm{sh}}(x) \qquad \shape{B, T, d}
|
||||||
|
\]
|
||||||
|
|
||||||
|
\item \textbf{合并}:
|
||||||
|
\[
|
||||||
|
y = s + W_\uparrow \cdot \mathrm{RMSNorm}(u) \qquad \shape{B, T, d}
|
||||||
|
\]
|
||||||
|
\end{enumerate}
|
||||||
|
|
||||||
|
\begin{importantbox}{如果你只记一件事}
|
||||||
|
Routed 专家只在 $\ell = d/2$ 的 latent 空间操作,
|
||||||
|
参数量是全宽专家的 $1/4$($\ell^2$ vs $d^2$)。
|
||||||
|
Shared 专家保持全宽 $d$,提供基础表达能力。
|
||||||
|
\end{importantbox}
|
||||||
|
|
||||||
|
\subsection{代码对照}
|
||||||
|
|
||||||
|
\begin{codemathtop}{layers/latent\_moe.py — LatentMoE.forward}
|
||||||
|
\begin{lstlisting}
|
||||||
|
def forward(self, x): # [B, T, d]
|
||||||
|
z = self.down(x) # [B, T, ell]
|
||||||
|
|
||||||
|
logits = self.router(x) # [B, T, n_routed]
|
||||||
|
topk = torch.topk(logits, self.top_k, dim=-1)
|
||||||
|
ids = topk.indices # [B, T, k]
|
||||||
|
probs = F.softmax(topk.values, dim=-1) # [B, T, k]
|
||||||
|
|
||||||
|
# All expert outputs (vectorized)
|
||||||
|
all_out = stack([e(z) for e in self.experts]) # [R, B, T, ell]
|
||||||
|
# Gather top-k and weighted sum
|
||||||
|
u = zeros(B, T, ell)
|
||||||
|
for i in range(self.top_k):
|
||||||
|
idx = ids[:,:,i].reshape(B*T)
|
||||||
|
sel = all_out[arange, idx]
|
||||||
|
u += probs[:,:,i:i+1] * sel.reshape(B, T, ell)
|
||||||
|
|
||||||
|
shared_out = stack([e(x) for e in self.shared]).sum(0) # [B, T, d]
|
||||||
|
return shared_out + self.up(self.norm(u)) # [B, T, d]
|
||||||
|
\end{lstlisting}
|
||||||
|
\end{codemathtop}
|
||||||
|
|
||||||
|
\begin{warningbox}{为什么 router 用 $x$(全宽)而不是 $z$(latent)?}
|
||||||
|
路由需要看到 token 的完整表示才能做好选择。
|
||||||
|
如果用 $z$ 路由,压缩过程可能丢失路由需要的信息。
|
||||||
|
K3 论文里也是用全宽 $x$ 做 Top-k,然后在 $\ell$ 空间计算。
|
||||||
|
\end{warningbox}
|
||||||
|
|
||||||
|
\subsection{形状总览}
|
||||||
|
|
||||||
|
\begin{center}
|
||||||
|
\begin{tabular}{llll}
|
||||||
|
\toprule
|
||||||
|
变量 & 形状 & 说明 \\
|
||||||
|
\midrule
|
||||||
|
$x$ & \shape{B, T, d} & 输入 \\
|
||||||
|
$z$ & \shape{B, T, \ell} & latent($\ell = d/2$)\\
|
||||||
|
logits & \shape{B, T, n_r} & router 输出 \\
|
||||||
|
ids & \shape{B, T, k} & Top-k 专家索引 \\
|
||||||
|
probs & \shape{B, T, k} & Top-k softmax 权重 \\
|
||||||
|
\texttt{all\_out} & \shape{n_r, B, T, \ell} & 所有 routed 专家输出 \\
|
||||||
|
$u$ & \shape{B, T, \ell} & 加权求和后的 routed 输出 \\
|
||||||
|
\texttt{shared\_out} & \shape{B, T, d} & shared 专家求和 \\
|
||||||
|
$y$ & \shape{B, T, d} & 最终输出 \\
|
||||||
|
\bottomrule
|
||||||
|
\end{tabular}
|
||||||
|
\end{center}
|
||||||
|
|
||||||
|
\subsection{本章小结}
|
||||||
|
|
||||||
|
LatentMoE 把 routed 专家限制在 $\ell = d/2$ 的 latent 空间,省参数。
|
||||||
|
SiTU-GLU 给 gate 和 up 加 $\tanh$ 软上限($\beta_1=4, \beta_2=25$),
|
||||||
|
防止低精度溢出。Shared 专家全宽,提供基础能力;routed 专家通过 Top-k 路由提供专业化能力。
|
||||||
@@ -0,0 +1,162 @@
|
|||||||
|
% teach:
|
||||||
|
% gap: 读者已知各组件但不知道怎么组装成完整模型
|
||||||
|
% takeaway: K3 = Hybrid(3 KDA + 1 MLA) × DecoderBlock(attn + MoE), 末层强制 MLA
|
||||||
|
% jump: 为什么每 4 层才放一次 MLA?位置感知只需要周期性提供
|
||||||
|
% omit: 0.5b preset 的训练超参
|
||||||
|
|
||||||
|
\section{K3 混合架构}
|
||||||
|
|
||||||
|
\subsection{整体结构}
|
||||||
|
|
||||||
|
\begin{center}
|
||||||
|
\texttt{Embedding} $\to$ \texttt{DecoderBlock} $\times L$ $\to$ \texttt{RMSNorm} $\to$ \texttt{LM Head}
|
||||||
|
\end{center}
|
||||||
|
|
||||||
|
每个 \texttt{DecoderBlock} 是 Pre-Norm 残差:
|
||||||
|
|
||||||
|
\begin{lstlisting}
|
||||||
|
def forward(self, x):
|
||||||
|
x = x + self.attn(self.attn_norm(x)) # mixing
|
||||||
|
return x + self.ffn(self.ffn_norm(x)) # channel
|
||||||
|
\end{lstlisting}
|
||||||
|
|
||||||
|
\subsection{Hybrid Attention Pattern}
|
||||||
|
|
||||||
|
K3 用两种 attention 层交替:
|
||||||
|
|
||||||
|
\begin{center}
|
||||||
|
\begin{tabular}{ccccccccc}
|
||||||
|
\toprule
|
||||||
|
层 & 0 & 1 & 2 & 3 & 4 & 5 & 6 & 7 \\
|
||||||
|
\midrule
|
||||||
|
Attn & KDA & KDA & KDA & \textbf{MLA} & KDA & KDA & KDA & \textbf{MLA} \\
|
||||||
|
FFN & MoE & MoE & MoE & MoE & MoE & MoE & MoE & MoE \\
|
||||||
|
\bottomrule
|
||||||
|
\end{tabular}
|
||||||
|
\end{center}
|
||||||
|
|
||||||
|
\noindent 规则:每 4 层放 1 次 MLA(0-based 层 3, 7, 11, ...),\textbf{末层强制 MLA}。
|
||||||
|
|
||||||
|
\begin{codemathtop}{models/k3\_config.py — layer\_types}
|
||||||
|
\begin{lstlisting}
|
||||||
|
def layer_types(self) -> list[str]:
|
||||||
|
"""Hybrid: 3 KDA + 1 MLA per group, last always MLA."""
|
||||||
|
types = ["kda"] * self.num_hidden_layers
|
||||||
|
for i in range(self.num_hidden_layers):
|
||||||
|
if i % 4 == 3:
|
||||||
|
types[i] = "mla"
|
||||||
|
types[-1] = "mla" # last layer forced
|
||||||
|
return types
|
||||||
|
|
||||||
|
def layer_specs(self):
|
||||||
|
return [(kind, "moe") for kind in self.layer_types()]
|
||||||
|
\end{lstlisting}
|
||||||
|
\end{codemathtop}
|
||||||
|
|
||||||
|
\begin{knowledgebox}{为什么 KDA 不需要 RoPE?}
|
||||||
|
KDA 的 gate/decay 机制天然提供位置感知:
|
||||||
|
远的 token 衰减更多,近的保留更多。
|
||||||
|
但 softmax attention(MLA)没有这个机制,所以真实 K3 用 NoPE
|
||||||
|
(本复现的 MLA 也是 NoPE)。
|
||||||
|
位置感知从 KDA 层``渗透''到 MLA 层——3:1 的比例足够了。
|
||||||
|
\end{knowledgebox}
|
||||||
|
|
||||||
|
\subsection{CausalLM 完整数据流}
|
||||||
|
|
||||||
|
\begin{codemathtop}{models/causal\_lm.py — CausalLM}
|
||||||
|
\begin{lstlisting}
|
||||||
|
class CausalLM(nn.Module):
|
||||||
|
def __init__(self, config):
|
||||||
|
self.embedding = nn.Embedding(vocab_size, D) # [V, D]
|
||||||
|
self.blocks = ModuleList([
|
||||||
|
DecoderBlock.from_spec(config, attn, ffn)
|
||||||
|
for attn, ffn in config.layer_specs()
|
||||||
|
])
|
||||||
|
self.norm = RMSNorm(D)
|
||||||
|
self.lm_head = nn.Linear(D, vocab_size) # [V, D]
|
||||||
|
if config.tie_word_embeddings:
|
||||||
|
self.lm_head.weight = self.embedding.weight
|
||||||
|
|
||||||
|
def forward(self, input_ids, labels=None):
|
||||||
|
x = self.embedding(input_ids) # [B,T] -> [B,T,D]
|
||||||
|
if self.mixer is None: # attnres="off"
|
||||||
|
for block in self.blocks:
|
||||||
|
x = block(x) # [B,T,D] -> [B,T,D]
|
||||||
|
else:
|
||||||
|
x = self.mixer(x) # AttnRes 深度残差, 见 §9
|
||||||
|
logits = self.lm_head(self.norm(x)) # [B,T,D] -> [B,T,V]
|
||||||
|
if labels is None:
|
||||||
|
return logits
|
||||||
|
# Shifted CE: predict next token
|
||||||
|
return cross_entropy(logits[:,:-1], labels[:,1:])
|
||||||
|
\end{lstlisting}
|
||||||
|
\end{codemathtop}
|
||||||
|
|
||||||
|
\begin{knowledgebox}{残差流是可替换的}
|
||||||
|
上面的 \texttt{DecoderBlock} 逐层堆叠(\texttt{x = x + sublayer(norm(x))})
|
||||||
|
是 \texttt{config.attnres="off"} 时的默认路径。
|
||||||
|
置为 \texttt{"full"} / \texttt{"block"} 时,\texttt{CausalLM} 会把每个 block
|
||||||
|
拆成 attn / ffn 两个原子子层交给 \texttt{mixer},用\textbf{深度维注意力}
|
||||||
|
代替等权残差加法——见 \S9。真实 K3 用的是 \texttt{block} 模式。
|
||||||
|
\end{knowledgebox}
|
||||||
|
|
||||||
|
\subsection{两种配置}
|
||||||
|
|
||||||
|
\begin{center}
|
||||||
|
\begin{tabular}{lll}
|
||||||
|
\toprule
|
||||||
|
& \textbf{KDAConfig}(纯 KDA) & \textbf{K3Config}(混合) \\
|
||||||
|
\midrule
|
||||||
|
Attn & KDA only & 3 KDA + 1 MLA \\
|
||||||
|
FFN & SwiGLU & LatentMoE \\
|
||||||
|
典型规模 & \textasciitilde8M (toy) & \textasciitilde8M (toy) / \textasciitilde500M (0.5b) \\
|
||||||
|
\texttt{layer\_specs()} & \texttt{[("kda","swiglu")] * L} & \texttt{[(kind,"moe") for kind in ...]} \\
|
||||||
|
\bottomrule
|
||||||
|
\end{tabular}
|
||||||
|
\end{center}
|
||||||
|
|
||||||
|
\subsection{K3 toy 尺寸}
|
||||||
|
|
||||||
|
\begin{center}
|
||||||
|
\begin{tabular}{llll}
|
||||||
|
\toprule
|
||||||
|
参数 & 真实 K3 & toy 复现 & 缩比 \\
|
||||||
|
\midrule
|
||||||
|
$D$ & 7168 & 256 & 28$\times$ \\
|
||||||
|
$L$ & 93 & 4 & 23$\times$ \\
|
||||||
|
$H = H_V$ & 96 & 8 & 12$\times$ \\
|
||||||
|
$K = V$ & 128 & 16 & 8$\times$ \\
|
||||||
|
kv\_lora\_rank & 512 & 32 & 16$\times$ \\
|
||||||
|
q\_lora\_rank & 1536 & 64 & 24$\times$ \\
|
||||||
|
$\ell$ (MoE latent) & 3584 & 128 & 28$\times$ \\
|
||||||
|
$n_{\mathrm{routed}}$ / Top-$k$ & 896/16 & 16/2 & 56$\times$ / 8$\times$ \\
|
||||||
|
\bottomrule
|
||||||
|
\end{tabular}
|
||||||
|
\end{center}
|
||||||
|
|
||||||
|
\subsection{DecoderBlock 构建}
|
||||||
|
|
||||||
|
\begin{codemathtop}{layers/block.py — build\_attn / build\_ffn}
|
||||||
|
\begin{lstlisting}
|
||||||
|
def build_attn(config, kind: str) -> nn.Module:
|
||||||
|
if kind == "kda": return KDAAttention.from_config(config)
|
||||||
|
if kind == "mla": return GatedMLA.from_config(config)
|
||||||
|
|
||||||
|
def build_ffn(config, kind: str) -> nn.Module:
|
||||||
|
if kind == "swiglu": return SwiGLUMLP.from_config(config)
|
||||||
|
if kind == "moe": return LatentMoE.from_config(config)
|
||||||
|
|
||||||
|
class DecoderBlock(nn.Module):
|
||||||
|
def forward(self, x):
|
||||||
|
x = x + self.attn(self.attn_norm(x))
|
||||||
|
return x + self.ffn(self.ffn_norm(x))
|
||||||
|
\end{lstlisting}
|
||||||
|
\end{codemathtop}
|
||||||
|
|
||||||
|
\subsection{本章小结}
|
||||||
|
|
||||||
|
K3 架构 = Hybrid Attention(3 KDA + 1 MLA,末层强制 MLA)+ LatentMoE。
|
||||||
|
KDA 层提供线性复杂度的序列混合和位置感知(通过 decay),
|
||||||
|
MLA 层提供全局 softmax attention(NoPE,利用 KDA 渗透的位置信息)。
|
||||||
|
每层默认是 Pre-Norm 残差 DecoderBlock;\texttt{config.attnres} 可以把这条
|
||||||
|
等权残差流换成 AttnRes 深度注意力(\S9)。
|
||||||
@@ -0,0 +1,379 @@
|
|||||||
|
% teach:
|
||||||
|
% gap: 读者知道 Pre-Norm 残差是"无条件等权累加", 但不知道怎么把它换成"按内容选择读哪一层"
|
||||||
|
% takeaway: AttnRes = 深度维 softmax 注意力残差; Block 版把 O(N^2) 源数压到 O(N/S); 两阶段 = inter 批量 + intra online-softmax 合并
|
||||||
|
% jump: 为什么打分用 RMSNorm 后的 v, 加权和却用原始 v
|
||||||
|
% omit: 论文里的 kernel 级调度与 pipeline 重叠
|
||||||
|
|
||||||
|
\section{Attention Residual 深度残差}
|
||||||
|
|
||||||
|
\subsection{从"等权累加"到"按内容选择"}
|
||||||
|
|
||||||
|
标准 Pre-Norm 残差把每层输出\textbf{无条件加}进残差流:
|
||||||
|
|
||||||
|
\[
|
||||||
|
x_l = x_{l-1} + f_l(x_{l-1}),
|
||||||
|
\qquad
|
||||||
|
x_N = x_0 + \sum_{l=1}^{N} f_l(x_{l-1})
|
||||||
|
\]
|
||||||
|
|
||||||
|
\noindent 展开后每一项权重恒为 1:第 3 层的输出和第 80 层的输出对最终表示的
|
||||||
|
"名义"贡献一样大,深层无法表达"我这一步应该主要读第 12 层的结果"。
|
||||||
|
|
||||||
|
AttnRes(\texttt{arXiv:2603.15031})把这个加法换成\textbf{深度维上的 softmax 注意力}:
|
||||||
|
第 $l$ 层持有一个可学习 query 向量 $w_l \in \mathbb{R}^D$,
|
||||||
|
把此前所有层的输出当成"可读的记忆":
|
||||||
|
|
||||||
|
\begin{align}
|
||||||
|
s_{l,i} &= w_l^{\top}\,\mathrm{RMS}(v_i),
|
||||||
|
& i = 0,1,\dots,l-1 \tag{A1} \\
|
||||||
|
\alpha_{l,i} &= \frac{\exp(s_{l,i})}{\sum_{j} \exp(s_{l,j})}
|
||||||
|
& \shape{n, B, T} \tag{A2} \\
|
||||||
|
h_l &= \sum_{i} \alpha_{l,i}\, v_i
|
||||||
|
& \shape{B, T, D} \tag{A3} \\
|
||||||
|
v_l &= f_l(h_l) \tag{A4}
|
||||||
|
\end{align}
|
||||||
|
|
||||||
|
\noindent 其中 $v_0 = x$(embedding 输出),$f_l$ 是已经含 Pre-Norm 的原子子层,
|
||||||
|
$\mathrm{RMS}(\cdot)$ 是不带 gain 的 RMS 归一化。
|
||||||
|
最后(\texttt{is\_final\_aggregate=True})再用一个独立 query 聚合所有源得到 $y$。
|
||||||
|
|
||||||
|
\begin{importantbox}{注意力权重是逐 token 的}
|
||||||
|
$s_{l,i}$ 的形状是 \shape{n, B, T}——每个 batch、每个位置 $t$ 都有自己的一套深度权重。
|
||||||
|
所以同一个位置在不同深度可以读不同的层,但\textbf{不跨时间混合},
|
||||||
|
因果性完全不受影响(\texttt{test\_attnres\_is\_still\_causal})。
|
||||||
|
\end{importantbox}
|
||||||
|
|
||||||
|
\subsection{DepthResidual:三个实现细节}
|
||||||
|
|
||||||
|
\begin{codemathtop}{layers/attn\_res.py — DepthResidual}
|
||||||
|
\begin{lstlisting}
|
||||||
|
class DepthResidual(nn.Module):
|
||||||
|
def __init__(self, dim, eps=1e-8, zero_init=True):
|
||||||
|
self.query = nn.Parameter(torch.zeros(dim)) # [D]
|
||||||
|
self.norm = RMSNorm(dim, eps=eps) # gain gamma
|
||||||
|
|
||||||
|
def effective_query(self):
|
||||||
|
return (self.query * self.norm.weight).float() # 折叠 gain
|
||||||
|
|
||||||
|
def forward(self, sources):
|
||||||
|
sources = stack_layers(sources) # [n,B,T,D]
|
||||||
|
q = self.effective_query() # [D]
|
||||||
|
k = rms(sources.float(), self.norm.eps) # 只用于打分
|
||||||
|
logits = einsum('d, n b t d -> n b t', q, k)
|
||||||
|
w = logits.softmax(dim=0) # 在深度维 softmax
|
||||||
|
out = einsum('n b t, n b t d -> b t d', w, sources.float())
|
||||||
|
return out.to(sources.dtype)
|
||||||
|
\end{lstlisting}
|
||||||
|
\end{codemathtop}
|
||||||
|
|
||||||
|
\paragraph{(1) gain 折叠}
|
||||||
|
RMSNorm 的可学习 gain $\gamma$ 本该作用在 key 上,但
|
||||||
|
$w^{\top}(\gamma \odot \mathrm{RMS}(v)) = (w \odot \gamma)^{\top}\mathrm{RMS}(v)$,
|
||||||
|
所以直接把 $\gamma$ 折进 query:$\tilde{w}_l = w_l \odot \gamma_l$。
|
||||||
|
少一次 \shape{n,B,T,D} 的逐元素乘法,两阶段算法里也只需要传一个向量。
|
||||||
|
|
||||||
|
\paragraph{(2) 打分用归一化的 $v$,加权和用原始 $v$}
|
||||||
|
注意 \texttt{logits} 用 \texttt{k = rms(sources)},而 \texttt{out} 用的是
|
||||||
|
\texttt{sources} 本身。
|
||||||
|
|
||||||
|
\begin{knowledgebox}{为什么这样不对称?}
|
||||||
|
打分要的是\textbf{方向}:$\mathrm{RMS}$ 之后 $s_{l,i}$ 与 $\|v_i\|$ 无关,
|
||||||
|
一层输出幅度大不会自动抢到高权重,softmax 只按"内容像不像我要读的东西"分配。\\
|
||||||
|
加权和要的是\textbf{原始信息}:如果对归一化后的 $v$ 求和,每层输出的模长
|
||||||
|
(承载着"这层贡献多大"的信息)就被抹掉了,深层的小幅修正会被放大到和主干同量级。
|
||||||
|
\end{knowledgebox}
|
||||||
|
|
||||||
|
\paragraph{(3) zero-init query}
|
||||||
|
\texttt{query} 默认初始化为 $0$ $\Rightarrow$ 所有 logits 为 $0$
|
||||||
|
$\Rightarrow$ softmax 均匀 $\Rightarrow$
|
||||||
|
\[
|
||||||
|
h_l = \frac{1}{l}\sum_{i=0}^{l-1} v_i
|
||||||
|
\]
|
||||||
|
训练起步就是\textbf{等权深度平均}(已实测:零初始化时 \texttt{forward} 输出与
|
||||||
|
\texttt{sources.mean(0)} 逐位相同),行为接近标准残差但自带 $1/l$ 缩放,
|
||||||
|
之后由梯度慢慢学出偏好。设 \texttt{zero\_init\_queries=False} 则用 $\mathcal{N}(0, 0.02^2)$。
|
||||||
|
|
||||||
|
\subsection{Full 与 Block:源数量的差别}
|
||||||
|
|
||||||
|
两种堆叠方式的区别只在\textbf{谁有资格进入源列表}:
|
||||||
|
|
||||||
|
\begin{itemize}[nosep, leftmargin=2em]
|
||||||
|
\item \texttt{FullAttnResStack}\\
|
||||||
|
保留\textbf{每一个原子层}的输出作为源,第 $l$ 层在 $l+1$ 个源上做注意力。
|
||||||
|
\item \texttt{BlockAttnResStack}\\
|
||||||
|
把 $N$ 个原子层切成大小为 $S$ 的块,\textbf{块内退化成普通求和}
|
||||||
|
(running partial $p \leftarrow p + v$),
|
||||||
|
只有\textbf{块的输出} $b_j$ 才进入源列表。
|
||||||
|
\end{itemize}
|
||||||
|
|
||||||
|
\begin{center}
|
||||||
|
\begin{tabular}{lccc}
|
||||||
|
\toprule
|
||||||
|
& \textbf{Full} & \textbf{Block ($S$)} & 标准残差 \\
|
||||||
|
\midrule
|
||||||
|
注意力源数 & 最多 $N+1$ & 最多 $N/S + 2$ & 1 \\
|
||||||
|
需保留的 \shape{B,T,D} 激活 & $O(N)$ & $O(N/S)$ & $O(1)$ \\
|
||||||
|
深度注意力 FLOPs & $O(N^2 BTD)$ & $O(N^2 BTD / S)$ & 0 \\
|
||||||
|
新增参数 & $2(N{+}1)D$ & $2(N{+}1)D$ & 0 \\
|
||||||
|
\bottomrule
|
||||||
|
\end{tabular}
|
||||||
|
\end{center}
|
||||||
|
|
||||||
|
\noindent 参数量不变(每个原子层都有自己的 query),变的是\textbf{显存与带宽}。
|
||||||
|
真实 K3($L=93$,$N=186$)取 $S=24$ 个原子层(12 个 DecoderBlock),
|
||||||
|
源数从 187 降到 $\le 10$。
|
||||||
|
|
||||||
|
\begin{codemathtop}{layers/attn\_res.py — BlockAttnResStack.forward\_naive(语义参考实现)}
|
||||||
|
\begin{lstlisting}
|
||||||
|
blocks = [x] # b_0 = embedding
|
||||||
|
partial = None
|
||||||
|
for layer_idx, (layer, residual) in enumerate(zip(self.layers, self.residuals), 1):
|
||||||
|
sources = blocks if partial is None else blocks + [partial]
|
||||||
|
h = residual(sources) # 深度注意力
|
||||||
|
out = layer(h)
|
||||||
|
partial = out if partial is None else (partial + out) # 块内: 普通累加
|
||||||
|
if (layer_idx % self.block_size == 0) or (layer_idx == len(self.layers)):
|
||||||
|
blocks.append(partial) # 块边界: 定型成一个新源
|
||||||
|
partial = None
|
||||||
|
return self.final_residual(blocks)
|
||||||
|
\end{lstlisting}
|
||||||
|
\end{codemathtop}
|
||||||
|
|
||||||
|
\subsection{两阶段算法(inter / intra)}
|
||||||
|
|
||||||
|
块内逐层跑上面的 naive 版本有个浪费:块内每一层的 query 面对的
|
||||||
|
\textbf{块间源 $b_0 \dots b_{j-1}$ 是完全相同且固定的},
|
||||||
|
唯一在变的只有 running partial $p$。于是拆成两个阶段:
|
||||||
|
|
||||||
|
\begin{enumerate}[nosep]
|
||||||
|
\item \textbf{inter(批量)}:把块内 $S$ 个 query 堆成 \shape{S, D},
|
||||||
|
对固定源做\textbf{一次}批量 einsum,拿到每个 query 的 online-softmax 三元组
|
||||||
|
$(m,\ \text{numer},\ \text{denom})$;
|
||||||
|
\item \textbf{intra(串行)}:逐层把新出现的 $p$ 作为\textbf{单个源}合并进去,
|
||||||
|
用 online softmax 的 merge 规则更新三元组,再 \texttt{normalized()} 出 $h$。
|
||||||
|
\end{enumerate}
|
||||||
|
|
||||||
|
\noindent online softmax 的三元组定义与合并规则(和 FlashAttention 同构,
|
||||||
|
只是"序列维"换成了"深度维"):
|
||||||
|
|
||||||
|
\begin{align}
|
||||||
|
m = \max_i s_i,
|
||||||
|
\quad
|
||||||
|
n = \sum_i e^{s_i - m} v_i,
|
||||||
|
\quad
|
||||||
|
d = \sum_i e^{s_i - m},
|
||||||
|
\quad
|
||||||
|
h = n / d
|
||||||
|
\tag{OS1}
|
||||||
|
\end{align}
|
||||||
|
|
||||||
|
\noindent 合并两组统计量 $(m_a, n_a, d_a)$ 与 $(m_b, n_b, d_b)$,
|
||||||
|
令 $m = \max(m_a, m_b)$、$w_a = e^{m_a - m}$、$w_b = e^{m_b - m}$:
|
||||||
|
|
||||||
|
\begin{align}
|
||||||
|
n = w_a\, n_a + w_b\, n_b,
|
||||||
|
\quad
|
||||||
|
d = w_a\, d_a + w_b\, d_b
|
||||||
|
\tag{OS2}
|
||||||
|
\end{align}
|
||||||
|
|
||||||
|
\noindent 单个源 $p$ 的三元组是 $(\,m = s_p,\ \text{numer} = p,\ \text{denom} = 1\,)$
|
||||||
|
——因为 $e^{s_p - m} = 1$,不需要真的算指数(\texttt{single\_source\_stats})。
|
||||||
|
|
||||||
|
\begin{codemathtop}{layers/attn\_res.py — \_run\_block\_two\_phase}
|
||||||
|
\begin{lstlisting}
|
||||||
|
queries = torch.stack([self.residuals[i].effective_query()
|
||||||
|
for i in range(start, end)], dim=0) # [S, D]
|
||||||
|
inter = attn_with_stats(queries, stack_layers(blocks), self.eps) # phase 1: 一次算完
|
||||||
|
|
||||||
|
partial = None
|
||||||
|
for local_idx, layer_idx in enumerate(range(start, end)): # phase 2: 串行
|
||||||
|
stats = inter.select(local_idx)
|
||||||
|
if partial is not None:
|
||||||
|
intra = single_source_stats(queries[local_idx], partial, self.eps)
|
||||||
|
stats = merge_attn_stats(stats, intra) # online softmax merge
|
||||||
|
h = stats.normalized()
|
||||||
|
out = self.layers[layer_idx](h)
|
||||||
|
partial = out if partial is None else (partial + out)
|
||||||
|
return partial
|
||||||
|
\end{lstlisting}
|
||||||
|
\end{codemathtop}
|
||||||
|
|
||||||
|
\begin{importantbox}{等价性是被测出来的,不是假设的}
|
||||||
|
\texttt{test\_block\_two\_phase\_matches\_naive} 直接对拍
|
||||||
|
\texttt{mixer.forward\_naive(emb)} 与 \texttt{mixer(emb)},
|
||||||
|
\texttt{atol=rtol=1e-5} 通过。Full 版同理:不传 \texttt{schedule\_block\_size}
|
||||||
|
走 naive,传了走两阶段,两者一致。
|
||||||
|
\end{importantbox}
|
||||||
|
|
||||||
|
\subsection{接入 CausalLM}
|
||||||
|
|
||||||
|
\subsubsection*{原子层 = 半个 DecoderBlock}
|
||||||
|
|
||||||
|
深度注意力的粒度是\textbf{原子层}而不是 DecoderBlock:
|
||||||
|
每个 block 拆成"norm + attn"和"norm + ffn"两个 Pre-Norm 原子层,
|
||||||
|
所以原子层数 $N = 2L$。
|
||||||
|
|
||||||
|
\begin{codemathtop}{models/causal\_lm.py — \_build\_mixer}
|
||||||
|
\begin{lstlisting}
|
||||||
|
atomics = []
|
||||||
|
for block in blocks:
|
||||||
|
atomics.append(BorrowedSubLayer(block.attn_norm, block.attn))
|
||||||
|
atomics.append(BorrowedSubLayer(block.ffn_norm, block.ffn))
|
||||||
|
|
||||||
|
if mode == "full":
|
||||||
|
return FullAttnResStack(D, atomics, eps=..., zero_init_queries=..., ...)
|
||||||
|
if mode == "block":
|
||||||
|
return BlockAttnResStack(D, atomics,
|
||||||
|
block_size=atomic_block_size(config.num_hidden_layers,
|
||||||
|
config.attnres_block_size), ...)
|
||||||
|
\end{lstlisting}
|
||||||
|
\end{codemathtop}
|
||||||
|
|
||||||
|
\subsubsection*{BorrowedSubLayer:借用而不注册}
|
||||||
|
|
||||||
|
\begin{codemathtop}{layers/attn\_res.py — BorrowedSubLayer}
|
||||||
|
\begin{lstlisting}
|
||||||
|
class BorrowedSubLayer(nn.Module):
|
||||||
|
def __init__(self, norm, fn):
|
||||||
|
self._borrowed = (norm, fn) # 普通 tuple, 不是 self.norm = norm
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
norm, fn = self._borrowed
|
||||||
|
return fn(norm(x))
|
||||||
|
\end{lstlisting}
|
||||||
|
\end{codemathtop}
|
||||||
|
|
||||||
|
\begin{warningbox}{为什么必须用 tuple 藏起来?}
|
||||||
|
如果写成 \texttt{self.norm = norm},\texttt{nn.Module} 会把它\textbf{注册成子模块},
|
||||||
|
于是同一份权重同时挂在 \texttt{blocks.0.attn.*} 和 \texttt{mixer.layers.0.fn.*} 下:
|
||||||
|
\begin{itemize}[nosep]
|
||||||
|
\item \texttt{model.parameters()} 出现重复 $\Rightarrow$ 优化器对同一参数更新两次
|
||||||
|
\item \texttt{state\_dict()} 多出一份镜像键 $\Rightarrow$ 旧 checkpoint 加载不上
|
||||||
|
\end{itemize}
|
||||||
|
放进普通 tuple 后 \texttt{blocks.*} 仍是唯一属主,
|
||||||
|
\texttt{mixer} 下只多出 depth query 与 gain(\texttt{test\_no\_duplicate\_parameter\_ids} 守这条)。
|
||||||
|
\end{warningbox}
|
||||||
|
|
||||||
|
\subsubsection*{forward:mixer 接管整条残差流}
|
||||||
|
|
||||||
|
\begin{codemathtop}{models/causal\_lm.py — CausalLM.forward}
|
||||||
|
\begin{lstlisting}
|
||||||
|
x = self.embedding(input_ids)
|
||||||
|
if self.mixer is None:
|
||||||
|
for block in self.blocks: # attnres="off": 老路径
|
||||||
|
x = block(x)
|
||||||
|
else:
|
||||||
|
x = self.mixer(x) # full / block: DecoderBlock.forward 被完全绕过
|
||||||
|
logits = self.lm_head(self.norm(x))
|
||||||
|
\end{lstlisting}
|
||||||
|
\end{codemathtop}
|
||||||
|
|
||||||
|
\noindent 注意 \texttt{mixer} 打开后 \texttt{DecoderBlock.forward}
|
||||||
|
(\S8 里的 \texttt{x = x + attn(...)})\textbf{一次都不会被调用}——
|
||||||
|
残差加法整个交给深度注意力,DecoderBlock 退化成"两个子层的容器"。
|
||||||
|
|
||||||
|
\subsubsection*{新增参数量:可忽略}
|
||||||
|
|
||||||
|
每个 DepthResidual 只有 query \shape{D} 和 gain \shape{D},共 $N+1$ 个:
|
||||||
|
|
||||||
|
\begin{center}
|
||||||
|
\begin{tabular}{lrrr}
|
||||||
|
\toprule
|
||||||
|
配置 & $D$ / $L$ & 原子层 $N$ & 新增参数 \\
|
||||||
|
\midrule
|
||||||
|
toy (K3Config) & 256 / 4 & 8 & 4{,}608 \\
|
||||||
|
0.5b preset & 768 / 24 & 48 & 75{,}264 \\
|
||||||
|
\bottomrule
|
||||||
|
\end{tabular}
|
||||||
|
\end{center}
|
||||||
|
|
||||||
|
\subsection{配置与命令行}
|
||||||
|
|
||||||
|
\begin{center}
|
||||||
|
\begin{tabular}{lll}
|
||||||
|
\toprule
|
||||||
|
字段 & 默认 & 含义 \\
|
||||||
|
\midrule
|
||||||
|
\texttt{attnres} & \texttt{"off"} & \texttt{off} / \texttt{full} / \texttt{block} \\
|
||||||
|
\texttt{attnres\_block\_size} & \texttt{None} & 每块几个 \textbf{DecoderBlock};\texttt{None} $\to \lceil L/8 \rceil$ \\
|
||||||
|
\texttt{attnres\_zero\_init\_queries} & \texttt{True} & query 零初始化(等权起步)\\
|
||||||
|
\texttt{attnres\_final\_aggregate} & \texttt{True} & 末尾再做一次全源聚合 \\
|
||||||
|
\bottomrule
|
||||||
|
\end{tabular}
|
||||||
|
\end{center}
|
||||||
|
|
||||||
|
\begin{codemathtop}{layers/attn\_res.py — atomic\_block\_size}
|
||||||
|
\begin{lstlisting}
|
||||||
|
def atomic_block_size(num_hidden_layers, attnres_block_size):
|
||||||
|
"""DecoderBlock 数 -> 原子层数。None 时目标约 8 块。"""
|
||||||
|
layers_per_block = (attnres_block_size if attnres_block_size is not None
|
||||||
|
else max(1, (num_hidden_layers + 7) // 8))
|
||||||
|
return layers_per_block * 2 # 每个 DecoderBlock = attn|ffn 两个原子层
|
||||||
|
\end{lstlisting}
|
||||||
|
\end{codemathtop}
|
||||||
|
|
||||||
|
\noindent 单位换算是最容易踩的一处:
|
||||||
|
配置字段的单位是 \textbf{DecoderBlock 数},
|
||||||
|
而堆叠类收到的 \texttt{block\_size} 是\textbf{原子层数}($\times 2$)。
|
||||||
|
例如 $L = 24$、块大小留 \texttt{None}:
|
||||||
|
\[
|
||||||
|
\lceil 24/8 \rceil = 3 \text{ 个 DecoderBlock}
|
||||||
|
\;\to\; S = 6 \text{ 个原子层}
|
||||||
|
\;\to\; N/S = 48/6 = 8 \text{ 块}
|
||||||
|
\]
|
||||||
|
|
||||||
|
\begin{lstlisting}
|
||||||
|
uv run python train_k3.py --preset toy --attnres block --attnres-block-size 2
|
||||||
|
uv run python train_k3.py --preset 0.5b --attnres block # 块大小自动 ~L/8
|
||||||
|
\end{lstlisting}
|
||||||
|
|
||||||
|
\noindent \texttt{KDAConfig} 与 \texttt{K3Config} 都在
|
||||||
|
\texttt{\_\_post\_init\_\_} 里调 \texttt{validate\_attnres},
|
||||||
|
非法模式 / 块大小在构造时就报错。
|
||||||
|
旧 checkpoint 的 config 里没有这几个字段,加载时回落到
|
||||||
|
\texttt{off}(见 \S 9.7 验证清单最后两行)。
|
||||||
|
|
||||||
|
\subsection{验证清单}
|
||||||
|
|
||||||
|
\texttt{tests/integration/test\_attn\_res.py},14 项全过:
|
||||||
|
|
||||||
|
\begin{center}
|
||||||
|
\begin{tabular}{ll}
|
||||||
|
\toprule
|
||||||
|
测试 & 守住的性质 \\
|
||||||
|
\midrule
|
||||||
|
\texttt{default\_attnres\_is\_off} & 默认不改变任何既有行为,\texttt{mixer is None} \\
|
||||||
|
\texttt{invalid\_attnres\_rejected} & 非法 mode / \texttt{block\_size=0} 构造期报错 \\
|
||||||
|
\texttt{mixer\_kind\_and\_atomic\_count} & 原子层数 $= 2L$,块大小 $\times 2$ 换算 \\
|
||||||
|
\texttt{auto\_block\_size\_targets\_eight\_blocks} & $L=93 \to 24$(K3 $S=12$ 个 block)\\
|
||||||
|
\texttt{no\_duplicate\_parameter\_ids} & 借用不注册,参数 id / 名字均无重复 \\
|
||||||
|
\texttt{off\_and\_block\_differ\_at\_same\_seed} & 同种子下确实换了计算图 \\
|
||||||
|
\texttt{block\_two\_phase\_matches\_naive} & 两阶段 $\equiv$ naive,\texttt{atol 1e-5} \\
|
||||||
|
\texttt{attnres\_is\_still\_causal} & 改末位 token 不影响前缀 logits \\
|
||||||
|
\texttt{kda\_config\_block\_runs} & 纯 KDA 配置也能开 \\
|
||||||
|
\texttt{attnres\_ckpt\_roundtrip} & 存取后逐位一致,config 字段保真 \\
|
||||||
|
\texttt{old\_ckpt\_without\_attnres\_stays\_off} & 向后兼容 \\
|
||||||
|
\texttt{attnres\_block\_overfits\_single\_batch} & 200 步 loss $< 0.5$,能训 \\
|
||||||
|
\bottomrule
|
||||||
|
\end{tabular}
|
||||||
|
\end{center}
|
||||||
|
|
||||||
|
\begin{warningbox}{混合精度}
|
||||||
|
\texttt{DepthResidual.forward}(naive 路径)显式 \texttt{.float()} 后再算 softmax
|
||||||
|
与加权和,最后 cast 回原 dtype;两阶段路径的 query 是 fp32、源保持原 dtype,
|
||||||
|
靠 einsum 的类型提升处理。bf16 autocast 下前向实测正常。
|
||||||
|
和 \S1 的结论一致:\textbf{指数/累和一律不要放进 fp16}。
|
||||||
|
\end{warningbox}
|
||||||
|
|
||||||
|
\subsection{本章小结}
|
||||||
|
|
||||||
|
AttnRes 把残差流从"等权累加"升级成"深度维 softmax 注意力":
|
||||||
|
每层用自己的 query 决定读此前哪些层的输出,打分在 RMS 归一化后做(方向)、
|
||||||
|
加权和在原始张量上做(保留模长),query 零初始化让训练从等权平均起步。
|
||||||
|
Full 版源数随深度线性增长,Block 版把块内退化成普通求和、只让块输出进入源列表,
|
||||||
|
把源数压到 $O(N/S)$;两阶段算法进一步把块间注意力批量化,
|
||||||
|
块内用 online softmax 增量合并,与 naive 实现数值等价。
|
||||||
|
接入 \texttt{CausalLM} 时每个 DecoderBlock 拆成 attn / ffn 两个原子层,
|
||||||
|
\texttt{BorrowedSubLayer} 用普通 tuple 借用权重以免重复注册,
|
||||||
|
\texttt{attnres="off"} 保持旧路径不变。
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user