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 clean bf16 cuDNN SDPA (standard BF16/FP32 mixed precision) FFN
DM-KDAGQA-1.7B-SD stale FlashMatch FP8 attention, stale delta (naive FP8) FFN + attention q/k/v/o
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.

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).

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.
Downloads last month
-
Safetensors
Model size
2B params
Tensor type
BF16
·
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support