DM-MambaGQA-1.7B-MP
A 1.67 B-parameter Mamba2 : GQA hybrid pretrained from scratch on 30 B tokens with standard BF16/FP32 mixed-precision attention (bf16 cuDNN 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-MambaGQA-1.7B-MP | clean | bf16 cuDNN SDPA (standard BF16/FP32 mixed precision) | FFN |
| DM-MambaGQA-1.7B-SD | stale | FlashMatch FP8 attention, stale delta (naive FP8) | FFN + attention q/k/v/o |
| DM-MambaGQA-1.7B-DM | match | FlashMatch FP8 attention, DeltaMatching (matched delta) | FFN + attention q/k/v/o |
Model
| layers | 24 = [Mamba2, Mamba2, Mamba2, GQA] × 6 |
| width | d_model 2048, SwiGLU FFN 7,104, tied embeddings |
| Mamba2 | head_dim 64, d_state 128, expand 2, 1 group, chunk 256 |
| 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,666,495,872 (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).
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.4144 | 51.01 | 0.5907 | 0.6711 | 0.1593 | 0.3057 | 0.0959 |
| SD (stale) | 1.4215 | 56.01 | 0.5857 | 0.6785 | 0.1571 | 0.2967 | 0.0857 |
| DM (match) | 1.4147 | 52.00 | 0.5905 | 0.6940 | 0.1588 | 0.3005 | 0.1249 |
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; across them no arm separates from clean in this cell, and stale's RULER lead comes from one bimodal
subtask (niah_multikey_2). 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, mamba-ssm and causal-conv1d (tested with torch 2.12,
transformers 5.9.0, mamba-ssm 2.3.2, causal-conv1d 1.6.2).
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
repo = "tturing/DM-MambaGQA-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 attention layers' K/V and the Mamba2 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 Mamba2 layers have no padding mask. - Fine-tuning.
model(input_ids=ids, labels=ids).lossis the next-token cross-entropy, so the model trains with thetransformersTrainerin 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 standard SDPA attention. - Fidelity. The shipped code is the study's evaluation path: its logits match the original loader bit for bit.
- Downloads last month
- 11