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.
296 lines
15 KiB
Markdown
296 lines
15 KiB
Markdown
# 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
|
||
```
|