Llama 1B with Kimi Delta Attention, 6B tokens

A dense ~1B-parameter Llama 3-style decoder where one attention layer in every four is Kimi Delta Attention (KDA), a linear-attention layer from Kimi Linear. It is otherwise identical to the QK-norm baseline and was trained on the same 6B FineWeb tokens.

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

Results

Training

Metric Value vs. baseline
Final train loss (step 3,053) 2.5634 −0.0064
Final eval loss (step 3,000) 2.5848 −0.0059
Final grad norm (step 3,053) 0.0514 +0.0060
Peak grad norm after step 200 0.559 +0.020
Tokens / steps 6B / 3,053 same

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 vs. baseline
HellaSwag acc_norm 39.65 +0.67
WinoGrande acc 51.46 −0.16
ARC-Easy acc_norm 39.60 −0.55
ARC-Challenge acc_norm 24.06 +0.43
PIQA acc_norm 67.79 +1.20
OpenBookQA acc_norm 27.60 −1.40
CommonsenseQA acc 20.07 +0.25
SciQ acc_norm 63.80 +0.30
LAMBADA acc 38.29 +0.54
Shared-9 average 41.37 +0.14
Shared-9 average, 4-bit NF4 40.64 −0.18

The +0.14 Shared-9 gap is within single-seed noise, so treat KDA as matching the baseline, not beating it.

Usage

The architecture class ships with the checkpoint, so load it with trust_remote_code=True. The KDA layers need flash-linear-attention and a CUDA GPU (its kernels are written in Triton). Tested with transformers==5.8.0 and flash-linear-attention==0.5.2.

pip install "transformers==5.8.0" "flash-linear-attention==0.5.2"
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer

repo = "Mercity/pretrain-kda-1b"

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

inputs = tokenizer("The capital of France is", return_tensors="pt").to(model.device)
# use_cache=False is required: the KDA layers keep no recurrent state between
# decoding steps, so cached generation would feed them one token at a time.
output = model.generate(**inputs, max_new_tokens=32, do_sample=False, use_cache=False)
print(tokenizer.decode(output[0], skip_special_tokens=True))

Generation without the cache re-reads the full sequence at every step, so it is slow for long outputs. Scoring text in one forward pass has no such cost:

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}")

The first forward pass on a new machine spends about 90 seconds compiling the KDA kernels.

Model details

Setting Value
Architecture LlamaKDA (Llama 3-style dense decoder, hybrid KDA + softmax attention)
Total parameters 1.056B
Layers 32: 24 GQA + 8 KDA
KDA layer positions 0, 4, 8, 12, 16, 20, 24, 28
Hidden size 1,536
Intermediate size (SwiGLU) 5,120
Attention heads / KV heads 12 / 6 (GQA layers)
QK normalization On (GQA layers)
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, ~18 hours
Stack TorchTitan, FlashAttention 4, Liger kernels, flash-linear-attention

Related checkpoints

Model Change from the baseline Shared-9
Baseline (QK-norm) Reference model 41.23
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. KDA inference needs a CUDA GPU and does not support cached generation through transformers. The model will repeat or make up facts and has had no alignment training.

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

Dataset used to train Mercity/pretrain-kda-1b

Collection including Mercity/pretrain-kda-1b

Paper for Mercity/pretrain-kda-1b