Yi-6B sets model_max_length=4096. encode() would clip long Wikipedia pages before we pack seq_len chunks. Raise the cap so only our chunker limits context.
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 |
建议执行顺序
- L1 → L2:先纯 PyTorch 把 fwd+bwd 写对,gradcheck 双保险
- L3 → L4:写 Triton kernel 时,L2 当 reference,逐项 diff
- L5:和 L3/L4 解耦测试(用 PyTorch 等价 naive gate 当参考)
- L6:复用 L1 单步公式,独立测试
- 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
当前进度
- L1 — recurrent PyTorch reference + gradcheck
- L2 — chunked PyTorch reference + recurrent 对拍 + gradcheck
- L7 — 可训练 CausalLM、toy overfit、checkpoint round-trip
- 统一
ops.api.chunk_kda接口:reference/triton/fla显式选择 - K3-like — Gated MLA + LatentMoE + Hybrid 3:1,
train_k3.py双语 wiki 预训练 - L3/L4 — vendored FLA Triton chunk fwd/bwd(
kda/_fla,不静默 import 上游fla) - L5 — vendored FLA fused gate
- L6 — vendored FLA fused recurrent decode
- AttnRes — 深度残差 mixer 接入
CausalLM(config.attnres = off | full | block) - 翻译训练环 —
--max-tokens/--resume/ held-out / cosine、train_sft.py、冻结data/eval
当前可直接运行小模型训练:
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。
# 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。安装并配置:
# 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
验证:
docker run --rm --gpus all nvidia/cuda:12.8.0-base-ubuntu24.04 nvidia-smi
# 应能看到 RTX 3060 与 driver 610.57
构建
docker build -t kda:latest .
--network=host必要:本机 docker daemon 配置了代理http://127.0.0.1:7890, 构建容器内的127.0.0.1指向容器自身而非宿主,不走 host 网络时 pip 下载会失败。
运行
无命令时只打印用法(退出码 2)。正式跑必须写入口。
# 用法
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)。
镜像内验证
docker run --rm kda:latest python -c "import torch, einops, swanlab, sacrebleu; print(torch.__version__, torch.version.cuda)"
# 预期: 2.9.0+cu128 12.8