DM-GDNMLA-1.7B-MP / README.md
tturing's picture
Add Q35MLA_clean_sdpa_ffn8_hd128_11785_fm128: GDN:MLA 1.7B, MP (clean) arm, seed 11785 (standalone, trust_remote_code)
0904b73 verified
|
Raw History Blame Contribute Delete
4.87 kB
---
tags:
- gated-deltanet
- mla
- hybrid
- fp8
- deltamatching
---
# DM-GDNMLA-1.7B-MP
A 1.67 B-parameter GatedDeltaNet : MLA hybrid pretrained from scratch on 30 B tokens with standard BF16/FP32 mixed-precision attention (bf16 SDPA). It is the **MP**
(clean) 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-GDNMLA-1.7B-MP](https://huggingface.co/tturing/DM-GDNMLA-1.7B-MP) | clean | bf16 SDPA (standard BF16/FP32 mixed precision) | FFN |
| [DM-GDNMLA-1.7B-SD](https://huggingface.co/tturing/DM-GDNMLA-1.7B-SD) | stale | FlashMatch FP8 attention, stale delta (naive FP8) | FFN + attention projections |
| [DM-GDNMLA-1.7B-DM](https://huggingface.co/tturing/DM-GDNMLA-1.7B-DM) | match | FlashMatch FP8 attention, DeltaMatching (matched delta) | FFN + attention projections |
## Model
| | |
|---|---|
| layers | 24 = [GatedDeltaNet, GatedDeltaNet, GatedDeltaNet, MLA] × 6 |
| width | d_model 2048, SwiGLU FFN 6,144, tied embeddings |
| GatedDeltaNet | 16 heads × 128, gated, short convolution (flash-linear-attention, chunk mode) |
| attention | multi-head latent attention, 16 heads × 128 (MHA on the wire), KV LoRA rank 384, qk-norm, partial RoPE 0.5 (θ = 1e7), gated output |
| vocabulary, context | Llama-2 32k tokenizer, 8,192 tokens |
| parameters | 1,665,444,672 (bf16 safetensors) |
The export runs attention through the bf16 MLA core 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 `mla_fp8` FlashMatch mixer).
## 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.4210** | **50.01** | **0.5919** | **0.6654** | **0.1596** | **0.3476** | **0.0781** |
| SD (stale) | 1.8424 | 25.97 | 0.4735 | 0.4734 | 0.0275 | 0.2666 | 0.0629 |
| DM (match) | 1.4204 | 56.24 | 0.5893 | 0.6755 | 0.1656 | 0.3484 | 0.0737 |
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 diverged during training (two-seed means: val CE +0.51, RULER −28.8 against
clean), while the match arm equals clean on val CE (+0.0003) and leads it on RULER (+4.0) and extraction (+0.008).
Per-seed and per-task results: `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-GDNMLA-1.7B-MP"
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 MLA layers' K/V and the GatedDeltaNet 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 GatedDeltaNet 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 the bf16 MLA core.
- **Fidelity.** The shipped code is the study's evaluation path: its logits match the original loader bit for bit.