DM-GDNMLA-1.7B-DM

A 1.67 B-parameter GatedDeltaNet : MLA hybrid pretrained from scratch on 30 B tokens with DeltaMatching: FlashMatch FP8 attention with the matched delta (sb_mode=fp8base_md). It is the DM (match) 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 clean bf16 SDPA (standard BF16/FP32 mixed precision) FFN
DM-GDNMLA-1.7B-SD stale FlashMatch FP8 attention, stale delta (naive FP8) FFN + attention projections
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.

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