Llama 1B baseline (no QK-norm), 6B tokens

A dense ~1B-parameter Llama 3-style decoder (4 KV heads, versus 6 in the later runs) pretrained from scratch on 6B tokens of FineWeb. This is our first baseline run, trained without QK normalization. Its gradient norm stayed high and wavy for the whole run, so we reran it with QK-norm on; that QK-norm baseline is the reference model for the rest of our ablations. We release this checkpoint so the effect of QK-norm can be compared directly.

This is a base model. It is not instruction-tuned or safety-tuned.

Results

Training

Metric Value vs. QK-norm baseline
Final train loss (step 3,052) 2.6211 +0.0512
Final eval loss (step 3,000) 2.5981 +0.0074
Final grad norm (step 3,052) 0.1260 +0.0806
Peak grad norm after step 200 1.344 +0.805
Mean grad norm over the run 0.58 +0.38
Tokens / steps 6B / 3,052 same

Zero-shot benchmarks

Scores at the final step, from the run's own evaluation. Shared-9 is the unweighted mean of the nine tasks.

Benchmark Metric Score
HellaSwag acc_norm 39.01
WinoGrande acc 50.43
ARC-Easy acc_norm 39.94
ARC-Challenge acc_norm 24.91
PIQA acc_norm 68.01
OpenBookQA acc_norm 30.40
CommonsenseQA acc 19.82
SciQ acc_norm 66.30
LAMBADA acc 26.66
Shared-9 average 40.61

These were measured earlier in the project than the other four checkpoints' scores, so small differences against them (LAMBADA in particular) may come from the evaluation setup as well as the model.

Usage

This is a standard LlamaForCausalLM checkpoint, so it needs no custom code. Tested with transformers==5.8.0.

import torch
from transformers import AutoModelForCausalLM, AutoTokenizer

repo = "Mercity/pretrain-baseline-non-qknorm"

tokenizer = AutoTokenizer.from_pretrained(repo)
model = AutoModelForCausalLM.from_pretrained(
    repo,
    torch_dtype=torch.bfloat16,
    device_map="auto",
)

inputs = tokenizer("The capital of France is", return_tensors="pt").to(model.device)
output = model.generate(**inputs, max_new_tokens=32, do_sample=False)
print(tokenizer.decode(output[0], skip_special_tokens=True))

To score text instead of generating it, pass labels and read the loss:

batch = tokenizer("FineWeb is a large web-text dataset.", return_tensors="pt").to(model.device)
with torch.no_grad():
    loss = model(**batch, labels=batch["input_ids"]).loss
print(f"loss={loss.item():.3f}  ppl={loss.exp().item():.1f}")

Model details

Setting Value
Architecture LlamaForCausalLM (Llama 3-style dense decoder)
Layers 32
Hidden size 1,536
Intermediate size (SwiGLU) 5,120
Total parameters 1.006B
Attention heads / KV heads 12 / 4 (GQA)
QK normalization Off
Attention bias None
Max sequence length 8,192
Tokenizer Llama 2, 32,000 tokens
Embeddings Tied input and output

Training

Setting Value
Data FineWeb sample-10BT, packed 8,192-token sequences
Tokens / steps 6B / 3,052
Batch 12 per device × 20 gradient accumulation (~1.97M tokens per step)
Gradient clipping 1.0
Optimizer Muon (LR 0.02, momentum 0.95, 5 Newton-Schulz steps, WD 0.1) + AdamW (LR 3e-4, β 0.9/0.95, WD 0.1)
Schedule Cosine, 100 warmup steps, min LR factor 0.1
Hardware 1 × NVIDIA B200, ~14.5 hours
Stack TorchTitan, FlashAttention 4, Liger kernels

Related checkpoints

Model Change from this model Shared-9
Baseline (QK-norm) QK normalization on; the reference for the ablations 41.23
KDA QK-norm, and 8 of 32 attention layers replaced with Kimi Delta Attention 41.37
N-gram 25% QK-norm, and ~25% of parameters moved into LongCat n-gram tables 40.56
N-gram 50% QK-norm, and ~48% of parameters moved into LongCat n-gram tables 39.54

Limitations

Trained on 6B English web tokens only, a small budget for a 1B model. Training was less stable than the QK-norm rerun, and benchmark scores are single-seed. The model will repeat or make up facts and has had no alignment training.

Downloads last month
172
Safetensors
Model size
1B params
Tensor type
BF16
·
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Dataset used to train Mercity/pretrain-baseline-non-qknorm

Collection including Mercity/pretrain-baseline-non-qknorm