Files
dela 24c9d56b72 Keep attnres on resume, fix final chunk_index, default 0.5b to Yi-6B
CLI default attnres=off was overwriting block checkpoints on resume
so later loads hit Unexpected key(s). Only apply flags the user
passed. Track next_chunk so a budget-exit save does not skip the
untrained yield. 0.5b now uses 01-ai/Yi-6B (64k); refuse resume
when the ckpt tokenizer does not match.
2026-08-26 14:23:29 +08:00

296 lines
15 KiB
Markdown
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# 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")` ≈ 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
```