dela e7185cbf49 Pull OPUS-100 en-zh for SFT instead of a checked-in jsonl
train_sft --data opus-100 streams Helsinki-NLP/opus-100, writes both
directions, and skips frozen eval sentences. Runtime cache stays under
data/sft/ (gitignored).
2026-08-25 21:39:58 +08:00

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

当前进度

  • 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") = 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。

# 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。安装并配置:

# 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
S
Description
No description provided
Readme
1.8 MiB
Languages
Python 86.1%
TeX 13.6%
Dockerfile 0.3%