# 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)`
`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
② 构造下三角系统
③ triangular solve
④ 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")` ≈ 415M,tied Yi-6B 64k 词表)。本机 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 Yi-6B embedding(64k,有 EOS),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)。**不能**从 Qwen3 词表的旧 ckpt `--resume`。 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 ```