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

A dense ~1B-parameter Llama 3-style decoder pretrained from scratch on 6B tokens of FineWeb. It is the reference model for our ablations of Kimi Delta Attention and LongCat n-gram embeddings: the other three checkpoints change one component and keep everything else identical to this one.

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

Results

Training

Metric Value
Final train loss (step 3,053) 2.5699
Final eval loss (step 3,000) 2.5907
Final grad norm (step 3,053) 0.0454
Peak grad norm after step 200 0.539
Tokens / steps 6B / 3,053

Zero-shot benchmarks

Scores from lm-eval on each task's full split. Shared-9 is the unweighted mean of the nine tasks.

Benchmark Metric Score
HellaSwag acc_norm 38.98
WinoGrande acc 51.62
ARC-Easy acc_norm 40.15
ARC-Challenge acc_norm 23.63
PIQA acc_norm 66.59
OpenBookQA acc_norm 29.00
CommonsenseQA acc 19.82
SciQ acc_norm 63.50
LAMBADA acc 37.75
Shared-9 average 41.23
Shared-9 average, 4-bit NF4 40.82

Few-NERD (LoRA fine-tuned)

Metric Score
Micro F1 0.655
Macro F1 0.595
Sentence accuracy 42.2%

Usage

The architecture class ships with the checkpoint, so load it with trust_remote_code=True. Tested with transformers==5.8.0.

import torch
from transformers import AutoModelForCausalLM, AutoTokenizer

repo = "Mercity/pretrain-baseline-qknorm"

tokenizer = AutoTokenizer.from_pretrained(repo)
model = AutoModelForCausalLM.from_pretrained(
    repo,
    trust_remote_code=True,
    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 LlamaQKNorm (Llama 3-style dense decoder with QK normalization)
Total parameters 1.031B
Layers 32
Hidden size 1,536
Intermediate size (SwiGLU) 5,120
Attention heads / KV heads 12 / 6 (GQA)
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,053
Batch 10 per device × 24 gradient accumulation (~1.97M tokens per step)
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, 150 warmup steps
Hardware 1 × NVIDIA B200, ~14.5 hours
Stack TorchTitan, FlashAttention 4, Liger kernels

Related checkpoints

Model Change from this baseline Shared-9
Baseline (no QK-norm) Same model without QK normalization; the original run 40.61
KDA 8 of 32 attention layers replaced with Kimi Delta Attention 41.37
N-gram 25% ~25% of parameters moved into LongCat n-gram tables, 23 layers 40.56
N-gram 50% ~48% of parameters moved into LongCat n-gram tables, 16 layers 39.54

Limitations

Trained on 6B English web tokens only, a small budget for a 1B model. Benchmark scores are single-seed. The model will repeat or make up facts and has had no alignment training.

Downloads last month
166
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-qknorm

Collection including Mercity/pretrain-baseline-qknorm