mlnomad/gelu-d12-chinchilla-261M-pytorch
GELU d=12 Chinchilla (261M) — PyTorch / HuggingFace Transformers
A 261M-parameter nanochat-architecture GPT with GELU MLP, originally trained in JAX/Flax on a TPU v6e-8 and ported to PyTorch for easy inference via the HuggingFace transformers API.
Weights are bit-exact with the Flax checkpoint (`mlnomad/gelu-d12-chinchilla-261M`) — parity validated at max |Δ logits| = 1.3e-5 on CPU/fp32.
Quick start
pip install torch transformers safetensorsimport torch
from transformers import AutoModelForCausalLM, AutoTokenizer
model = AutoModelForCausalLM.from_pretrained(
"mlnomad/gelu-d12-chinchilla-261M-pytorch",
trust_remote_code=True, # the model class ships in the repo
dtype=torch.float32,
).eval()
# The model was trained with the Mistral-7B-v0.1 tokenizer (vocab 32,768)
tokenizer = AutoTokenizer.from_pretrained("mistralai/Mistral-7B-v0.1")
prompt = "The meaning of life is"
ids = tokenizer(prompt, return_tensors="pt").input_ids
with torch.no_grad():
out = model.generate(
ids,
max_new_tokens=50,
do_sample=True,
temperature=0.8,
top_p=0.9,
use_cache=True,
pad_token_id=tokenizer.eos_token_id or 0,
)
print(tokenizer.decode(out[0], skip_special_tokens=True))Greedy completion examples from this checkpoint:
"The meaning of life is" → to live in a place where you can live. The meaning of life is to live in a "Once upon a time," → the first thing I would do was to make a list of all the things I would like to do
Model details
Architecture features
Full nanochat stack, faithfully ported to PyTorch:
- RoPE (base 100,000), split-half layout
- MHA (nhead = nkvhead = 12; the code supports GQA via nkvhead < nhead, but all d=12 models use full MHA)
- QK-norm with 1.2× scaling (after RoPE)
- Parameterless RMSNorm (no learnable gain) post-embedding and per block
- Sliding-window attention with
"SSSL"pattern - Tied embeddings (lm_head = wte.T)
- Value embeddings on alternating layers (ResFormer-style)
- Per-layer learnable residual scalars (
resid_lambdas,x0_lambdas) - Smear — learnable gate on first 24 dims of token embedding mixes in prev token
- Backout — subtract mid-layer residual from late layers
- Logit soft-cap:
15 · tanh(logits / 15) - No biases in any Linear
- MLP:
Linear(n → 4n) → F.gelu(approximate="tanh") → Linear(4n → n)
KV cache
The GeluGPTForCausalLM class implements a smear-aware KV cache for fast autoregressive generation. KV-cache parity vs full forward is validated at max |Δ| < 3e-5. Pass use_cache=True (the default for .generate()).
Files in this repo
.
├── config.json # HF config with auto_map pointing to the classes below
├── generation_config.json
├── model.safetensors # 1.04 GB, fp32 weights + persistent RoPE buffers
├── torch_gpt.py # pure PyTorch GPT module (GELU_GPT)
├── configuration_gelu_gpt.py # PretrainedConfig subclass
├── modeling_gelu_gpt.py # PreTrainedModel + GenerationMixin wrapper with KV cache
└── README.mdRelated
- `mlnomad/gelu-d12-chinchilla-261M` — original JAX/Flax Orbax checkpoint (model + AdamW optimizer state, resumable)
- flaxchat — the JAX/Flax training harness that produced the weights
Wikitext-103 evaluation
Evaluated on ~330K tokens from wikitext-103 test set (model trained on C4 only — this is a zero-shot transfer metric).
License
Apache 2.0.
