|
Download README.md from tturing/DM-KDAGQA-1.7B-SD: direct link, hf CLI and curl.
- Browser
- Download file 5.07 kB
-
https://huggingface.co/tturing/DM-KDAGQA-1.7B-SD/resolve/main/README.md
- Command line
-
hf download hf://tturing/DM-KDAGQA-1.7B-SD/README.md
-
curl -L -o README.md https://huggingface.co/tturing/DM-KDAGQA-1.7B-SD/resolve/main/README.md
5.07 kB
| tags: | |
| - kimi-delta-attention | |
| - kda | |
| - hybrid | |
| - fp8 | |
| - deltamatching | |
| # DM-KDAGQA-1.7B-SD | |
| A 1.66 B-parameter KDA : GQA hybrid (Kimi Delta Attention recurrent layers, grouped-query attention) pretrained from | |
| scratch on 30 B tokens with naive FP8 attention: FlashMatch FP8 attention with the stale delta (`sb_mode=fp8base`). It is the **SD** (stale) arm of a study in which three models were trained | |
| identically except for the attention's training precision: | |
| | repo | arm | attention in training | FP8 GEMMs | | |
| |---|---|---|---| | |
| | [DM-KDAGQA-1.7B-MP](https://huggingface.co/tturing/DM-KDAGQA-1.7B-MP) | clean | bf16 cuDNN SDPA (standard BF16/FP32 mixed precision) | FFN | | |
| | [DM-KDAGQA-1.7B-SD](https://huggingface.co/tturing/DM-KDAGQA-1.7B-SD) | stale | FlashMatch FP8 attention, stale delta (naive FP8) | FFN + attention q/k/v/o | | |
| | [DM-KDAGQA-1.7B-DM](https://huggingface.co/tturing/DM-KDAGQA-1.7B-DM) | match | FlashMatch FP8 attention, DeltaMatching (matched delta) | FFN + attention q/k/v/o | | |
| ## Model | |
| | | | | |
| |---|---| | |
| | layers | 24 = [KDA, KDA, KDA, GQA] × 6 | | |
| | width | d_model 2048, SwiGLU FFN 8,064, tied embeddings | | |
| | KDA | Kimi Delta Attention (flash-linear-attention, chunk mode), 16 heads × 128, value width = key width, short convolution, per-channel decay gate, gated output | | |
| | attention | GQA, 16 query / 4 KV heads × 128, qk-norm, partial RoPE 0.25 (θ = 1e7), gated output | | |
| | vocabulary, context | Llama-2 32k tokenizer, 8,192 tokens | | |
| | parameters | 1,664,776,224 (bf16 safetensors) | | |
| The export runs attention through bf16 SDPA in every arm, so the three repos share one architecture and differ only in | |
| their weights. `config.train.json` records the training-time layer config (the FP8 arms train their attention layers | |
| through the `sagebwd` FlashMatch mixer). The KDA layers are bf16 in all three arms. | |
| ## Training | |
| | | | | |
| |---|---| | |
| | data | Nemotron-CC (`nemotron_cc_v2d1_hq_dqa`), packed 8,192-token sequences | | |
| | budget | 28,610 steps × global batch 128 × 8,192 tokens = 30.0 B tokens | | |
| | optimizer | AdamW, β (0.9, 0.95), ε 1e-8, weight decay 0.1, gradient clip 1.0, z-loss 1e-4 | | |
| | schedule | WSD: 3.33 % warmup to 2.4e-3, constant, linear decay over the last 20 % | | |
| | seed | 11785 (initialization only; every run in the study reads the same data in the same order) | | |
| | hardware | 8 × H200 | | |
| ## Evaluation (seed 11785) | |
| | model | val CE | RULER | CSense-9 | Extract | TriviaQA | MMLU | MQAR | | |
| |---|---|---|---|---|---|---|---| | |
| | MP (clean) | 1.3990 | 58.66 | 0.5994 | 0.6829 | 0.1764 | 0.3472 | 0.1242 | | |
| | **SD (stale)** | **1.7124** | **34.58** | **0.5072** | **0.5546** | **0.0444** | **0.2646** | **0.1337** | | |
| | DM (match) | 1.3989 | 55.61 | 0.5993 | 0.6556 | 0.1851 | 0.3374 | 0.1369 | | |
| val CE: nats/token on 2,000 held-out 8,192-token windows (lower is better). RULER: 13 tasks at 4k and 8k, mean of the | |
| two lengths. CSense-9: LAMBADA, HellaSwag, PIQA, ARC-e, ARC-c, SciQ, OpenBookQA, WinoGrande, COPA. Extract: SWDE, FDA, | |
| SQuAD completion. TriviaQA: 5-shot exact match. MQAR: synthetic multi-query key-value recall. Each arm was trained | |
| with two seeds. In this cell the stale arm drifted away from clean during training (two-seed means: val CE +0.42, | |
| RULER 24.5 against 57.8), while the match arm equals clean on val CE (−0.0006) and stays within 0.0011 of it in train | |
| CE from step 5,000 on; its RULER −2.1 comes from one subtask (`niah_multikey_2`). Per-seed and per-task results and the | |
| full configuration: `report.md` in [tturing/n8t-train-curve](https://huggingface.co/datasets/tturing/n8t-train-curve). | |
| ## Usage | |
| The model code ships with this repo (`modeling_bqalm.py`, `configuration_bqalm.py`), so no other codebase is needed. | |
| It needs a CUDA GPU with `torch`, `transformers` >= 5 and `flash-linear-attention` (tested with torch 2.12, | |
| transformers 5.9.0, flash-linear-attention 0.5.0). | |
| ```python | |
| import torch | |
| from transformers import AutoModelForCausalLM, AutoTokenizer | |
| repo = "tturing/DM-KDAGQA-1.7B-SD" | |
| model = AutoModelForCausalLM.from_pretrained(repo, dtype=torch.bfloat16, trust_remote_code=True).cuda() | |
| tokenizer = AutoTokenizer.from_pretrained(repo, trust_remote_code=True) | |
| inputs = tokenizer("The capital of France is", return_tensors="pt").to("cuda") | |
| print(tokenizer.decode(model.generate(**inputs, max_new_tokens=32)[0], skip_special_tokens=True)) | |
| ``` | |
| - **Inference.** `generate()` keeps a cache (the GQA layers' K/V and the KDA recurrent and conv states), so each new | |
| token costs one step. Greedy decoding and sampling are supported (`num_beams=1`); prompts batched together must | |
| share one length, since the KDA layers have no padding mask. | |
| - **Fine-tuning.** `model(input_ids=ids, labels=ids).loss` is the next-token cross-entropy, so the model trains with | |
| the `transformers` `Trainer` in bf16. The FP8 FlashMatch attention the SD and DM arms were trained with needs a | |
| compiled CUDA kernel that is not shipped; the exported weights are bf16 and run bf16 SDPA attention. | |
| - **Fidelity.** The shipped code is the study's evaluation path: its logits match the original loader bit for bit. | |