ram-18m-training-scripts / train_ram_18m.py
Compactbot's picture
Add ram-18m training script (18,290,304 params, LLaMA-style GQA+SwiGLU) (#1)
143fec7
Raw History Blame Contribute Delete
15.3 kB
#!/usr/bin/env python3
"""ram-18m: 18.3M-param LLaMA-style language model trained from scratch.
Architecture:
d_model=384, n_heads=6, n_kv_heads=2 (GQA), n_layers=7
SwiGLU FFN (4x), RoPE, RMSNorm, vocab 8192, tied embed/head, ctx 512
Total: 18,290,304 learnable parameters
Default training:
~2B tokens (FineWeb-Edu L3), AdamW 2e-4, cosine + warmup
batch 32 (effective), seq 512, ~12,207 steps
Usage:
python3 train_ram_18m.py --stage prepare # download + tokenize data
python3 train_ram_18m.py --stage train # train the model
python3 train_ram_18m.py --stage all
python3 train_ram_18m.py --stage eval --ckpt path/to/ckpt.pt
Requirements:
pip install torch transformers datasets numpy tokenizers
"""
import os, sys, math, json, time, argparse, glob, random
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.optim import AdamW
# ============================================================================
# Architecture
# ============================================================================
VOCAB = 8192
D_MODEL = 384
N_HEADS = 6
N_KV_HEADS = 2
N_LAYERS = 7
HEAD_DIM = D_MODEL // N_HEADS # 64
KV_DIM = N_KV_HEADS * HEAD_DIM # 128
FFN_DIM = D_MODEL * 4 # 1536
SEQ_LEN = 512
ROPE_THETA = 10000.0
# Verified param count: 18,290,304
class RMSNorm(nn.Module):
def __init__(self, dim, eps=1e-6):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x):
norm = x.float().pow(2).mean(-1, keepdim=True).add(self.eps).rsqrt()
return (x.float() * norm).type_as(x) * self.weight
class RoPE(nn.Module):
def __init__(self, head_dim, theta=ROPE_THETA):
super().__init__()
freqs = 1.0 / (theta ** (torch.arange(0, head_dim, 2).float() / head_dim))
self.register_buffer("freqs", freqs, persistent=False)
def forward(self, x, pos):
freqs = self.freqs
angles = pos[:, None].float() * freqs[None, :]
cos = angles.cos()[None, None, :, :]
sin = angles.sin()[None, None, :, :]
x1 = x[..., 0::2]
x2 = x[..., 1::2]
out1 = x1 * cos - x2 * sin
out2 = x1 * sin + x2 * cos
return torch.stack([out1, out2], dim=-1).flatten(-2)
class GQAAttention(nn.Module):
def __init__(self):
super().__init__()
self.q_proj = nn.Linear(D_MODEL, N_HEADS * HEAD_DIM, bias=False)
self.k_proj = nn.Linear(D_MODEL, N_KV_HEADS * HEAD_DIM, bias=False)
self.v_proj = nn.Linear(D_MODEL, N_KV_HEADS * HEAD_DIM, bias=False)
self.o_proj = nn.Linear(N_HEADS * HEAD_DIM, D_MODEL, bias=False)
self.rope = RoPE(HEAD_DIM)
def forward(self, x, mask=None):
B, T, _ = x.shape
q = self.q_proj(x).view(B, T, N_HEADS, HEAD_DIM).transpose(1, 2)
k = self.k_proj(x).view(B, T, N_KV_HEADS, HEAD_DIM).transpose(1, 2)
v = self.v_proj(x).view(B, T, N_KV_HEADS, HEAD_DIM).transpose(1, 2)
pos = torch.arange(T, device=x.device)
q = self.rope(q, pos)
k = self.rope(k, pos)
rep = N_HEADS // N_KV_HEADS
k = k.repeat_interleave(rep, dim=1)
v = v.repeat_interleave(rep, dim=1)
scale = HEAD_DIM ** -0.5
attn = (q @ k.transpose(-2, -1)) * scale
if mask is not None:
attn = attn.masked_fill(mask[:, None, None, :] == 0, float("-inf"))
attn = F.softmax(attn, dim=-1)
out = attn @ v
out = out.transpose(1, 2).contiguous().view(B, T, N_HEADS * HEAD_DIM)
return self.o_proj(out)
class SwiGLU(nn.Module):
def __init__(self):
super().__init__()
self.gate = nn.Linear(D_MODEL, FFN_DIM, bias=False)
self.up = nn.Linear(D_MODEL, FFN_DIM, bias=False)
self.down = nn.Linear(FFN_DIM, D_MODEL, bias=False)
def forward(self, x):
return self.down(F.silu(self.gate(x)) * self.up(x))
class TransformerBlock(nn.Module):
def __init__(self):
super().__init__()
self.attn_norm = RMSNorm(D_MODEL)
self.attn = GQAAttention()
self.ffn_norm = RMSNorm(D_MODEL)
self.ffn = SwiGLU()
def forward(self, x, mask=None):
x = x + self.attn(self.attn_norm(x), mask)
x = x + self.ffn(self.ffn_norm(x))
return x
class RAM18M(nn.Module):
def __init__(self):
super().__init__()
self.tok_emb = nn.Embedding(VOCAB, D_MODEL)
self.layers = nn.ModuleList([TransformerBlock() for _ in range(N_LAYERS)])
self.norm = RMSNorm(D_MODEL)
self.head = nn.Linear(D_MODEL, VOCAB, bias=False)
self.head.weight = self.tok_emb.weight
self.apply(self._init_weights)
def _init_weights(self, module):
if isinstance(module, nn.Linear):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
if module.bias is not None:
nn.init.zeros_(module.bias)
elif isinstance(module, nn.Embedding):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
def forward(self, input_ids, targets=None):
B, T = input_ids.shape
h = self.tok_emb(input_ids)
mask = torch.tril(torch.ones(T, T, device=input_ids.device))
for layer in self.layers:
h = layer(h, mask)
h = self.norm(h)
logits = self.head(h)
loss = None
if targets is not None:
loss = F.cross_entropy(logits.view(-1, VOCAB), targets.view(-1))
return logits, loss
def count_params(self):
return sum(p.numel() for p in self.parameters() if p.requires_grad)
def get_lr(step, max_steps, warmup, base_lr, min_lr):
if step < warmup:
return base_lr * (step + 1) / warmup
if step >= max_steps:
return min_lr
progress = (step - warmup) / (max_steps - warmup)
return min_lr + 0.5 * (base_lr - min_lr) * (1 + math.cos(math.pi * progress))
# ============================================================================
# Data
# ============================================================================
DATASET = "HuggingFaceFW/fineweb-edu"
DATASET_CONFIG = "sample-100BT"
SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__))
TOK_DIR = os.path.join(SCRIPT_DIR, "tokens")
CKPT_DIR = SCRIPT_DIR
def prepare_data(target_tokens=2_000_000_000):
"""Download and tokenize FineWeb-Edu. Saves .npy files of token ids."""
from datasets import load_dataset
from tokenizers import Tokenizer
from tokenizers.models import BPE
from tokenizers.pre_tokenizers import Whitespace
from tokenizers.trainers import BpeTrainer
os.makedirs(TOK_DIR, exist_ok=True)
print("Training BPE tokenizer (vocab 8192)...")
ds = load_dataset(DATASET, DATASET_CONFIG, split="train", streaming=True)
texts = []
for i, row in enumerate(ds):
texts.append(row["text"])
if i >= 200000:
break
tokenizer = Tokenizer(BPE(unk_token="<unk>"))
tokenizer.pre_tokenizer = Whitespace()
trainer = BpeTrainer(vocab_size=VOCAB, special_tokens=["<unk>", "<pad>", "<bos>", "<eos>"])
tokenizer.train_from_iterator(texts, trainer)
tok_path = os.path.join(TOK_DIR, "tokenizer.json")
tokenizer.save(tok_path)
print(f"Tokenizer saved to {tok_path}")
print(f"Tokenizing up to {target_tokens:,} tokens...")
ds = load_dataset(DATASET, DATASET_CONFIG, split="train", streaming=True)
all_tokens = []
n_docs = 0
for row in ds:
ids = tokenizer.encode(row["text"])
if len(ids) < 10:
continue
all_tokens.extend(ids)
n_docs += 1
if len(all_tokens) >= target_tokens:
break
if n_docs % 100000 == 0:
print(f" {n_docs} docs, {len(all_tokens):,} tokens")
print(f"Total: {n_docs} docs, {len(all_tokens):,} tokens")
arr = np.array(all_tokens, dtype=np.int32)
part_size = 100_000_000
for i in range(0, len(arr), part_size):
part = arr[i:i+part_size]
path = os.path.join(TOK_DIR, f"part_{i//part_size:03d}.npy")
np.save(path, part)
print(f" Saved {path}: {len(part):,} tokens")
print("Data prep complete.")
class DataIterator:
"""Streams tokenized data from .npy parts, yielding (input, target) batches."""
def __init__(self, tok_dir, batch_size, seq_len, device="cpu"):
self.parts = sorted(glob.glob(os.path.join(tok_dir, "part_*.npy")))
if not self.parts:
raise FileNotFoundError(f"No .npy files in {tok_dir}. Run --stage prepare first.")
self.batch_size = batch_size
self.seq_len = seq_len
self.device = device
self._buf = np.array([], dtype=np.int32)
self._part_idx = 0
self._rng = np.random.default_rng(42)
def _refill(self):
need = self.batch_size * self.seq_len + self.seq_len
while len(self._buf) < need:
if self._part_idx >= len(self.parts):
self._part_idx = 0
part = np.load(self.parts[self._part_idx], mmap_mode="r")
self._buf = np.concatenate([self._buf, np.array(part)])
self._part_idx += 1
def __iter__(self):
while True:
self._refill()
max_start = len(self._buf) - self.batch_size * self.seq_len - self.seq_len
if max_start < 0:
self._refill()
continue
start = int(self._rng.integers(0, max_start))
chunk = self._buf[start:start + self.batch_size * self.seq_len + self.seq_len]
flat = chunk.reshape(self.batch_size, self.seq_len + 1)
x = torch.tensor(flat[:, :-1], dtype=torch.long, device=self.device)
y = torch.tensor(flat[:, 1:], dtype=torch.long, device=self.device)
yield x, y
# ============================================================================
# Training
# ============================================================================
def train(
steps=12207,
batch_size=32,
seq_len=SEQ_LEN,
lr=2e-4,
min_lr=2e-5,
warmup=200,
accum=1,
ckpt_every=500,
eval_every=500,
device="cuda",
resume=None,
):
"""Train ram-18m from scratch."""
torch.manual_seed(42)
model = RAM18M()
n_params = model.count_params()
print(f"Model: {n_params:,} params")
start_step = 0
if resume:
ckpt = torch.load(resume, map_location="cpu")
model.load_state_dict(ckpt["model"])
start_step = ckpt["step"]
print(f"Resumed from {resume} at step {start_step}")
if device == "cuda" and torch.cuda.is_available():
model = model.cuda()
else:
device = "cpu"
model = model.to(device)
print(f"Device: {device}")
data_iter = DataIterator(TOK_DIR, batch_size, seq_len, device)
opt = AdamW(model.parameters(), lr=lr, betas=(0.9, 0.95), weight_decay=0.0)
model.train()
t0 = time.time()
for step in range(start_step, steps):
opt.zero_grad()
loss_accum = 0.0
for _ in range(accum):
x, y = next(iter(data_iter))
_, loss = model(x, y)
loss = loss / accum
loss.backward()
loss_accum += loss.item()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
opt.step()
loss_val = loss_accum
if (step + 1) % 10 == 0:
elapsed = time.time() - t0
tok_per_sec = (step + 1 - start_step) * batch_size * seq_len / max(elapsed, 1)
lr_now = get_lr(step, steps, warmup, lr, min_lr)
print(f"step {step+1}/{steps} | loss {loss_val:.4f} | lr {lr_now:.6f} | {tok_per_sec:.0f} tok/s | {elapsed:.0f}s")
if (step + 1) % ckpt_every == 0:
path = os.path.join(CKPT_DIR, f"ckpt_step{step+1}.pt")
torch.save({"model": model.state_dict(), "opt": opt.state_dict(),
"step": step + 1, "loss": loss_val}, path)
print(f" checkpoint -> {path}")
if (step + 1) % eval_every == 0:
model.eval()
with torch.no_grad():
x, y = next(iter(data_iter))
_, eval_loss = model(x, y)
model.train()
print(f" eval_loss (1 batch): {eval_loss.item():.4f}")
path = os.path.join(CKPT_DIR, "final.pt")
torch.save({"model": model.state_dict(), "step": steps, "loss": loss_val}, path)
print(f"Training complete. Final model -> {path}")
# Print sample
model.eval()
torch.manual_seed(0)
with torch.no_grad():
prompt = torch.tensor([[3]], device=device)
for _ in range(200):
logits, _ = model(prompt)
next_tok = logits[0, -1].argmax()
prompt = torch.cat([prompt, next_tok.unsqueeze(0)], dim=1)
try:
from tokenizers import Tokenizer
tok_path = os.path.join(TOK_DIR, "tokenizer.json")
if os.path.exists(tok_path):
tok = Tokenizer.from_file(tok_path)
text = tok.decode(prompt[0].tolist())
print(f"\nSample generation:\n{text[:500]}")
except Exception:
pass
# ============================================================================
# Eval: zero-shot loglikelihood on standard benchmarks
# ============================================================================
def eval_benchmarks(ckpt_path, device="cuda", n_samples=500):
"""Run zero-shot loglikelihood eval on PIQA, ARC-Easy, ARC-Challenge, HellaSwag."""
from datasets import load_dataset
model = RAM18M()
ckpt = torch.load(ckpt_path, map_location="cpu")
model.load_state_dict(ckpt["model"])
if device == "cuda" and torch.cuda.is_available():
model = model.cuda()
else:
device = "cpu"
model.eval()
tok_path = os.path.join(TOK_DIR, "tokenizer.json")
from tokenizers import Tokenizer
tokenizer = Tokenizer.from_file(tok_path)
def encode(text):
return tokenizer.encode(text).ids
def loglikelihood(context, continuation):
full_ids = encode(context + " " + continuation)
ctx_ids = encode(context)
ctx_len = min(len(ctx_ids), len(full_ids) - 1)
if ctx_len < 1:
return -1000.0
input_ids = torch.tensor([full_ids], device=device)
with torch.no_grad():
logits, _ = model(input_ids)
log_probs = F.log_softmax(logits[0, ctx_len-1:-1, :], dim=-1)
target_ids = torch.tensor(full_ids[ctx_len:], device=device)
if len(target_ids) == 0:
return -1000.0
return log_probs.gather(1, target_ids.unsqueeze(1)).sum().item()
def accuracy(items, n=n_samples):
correct = 0
total = 0
for item in items[:n]:
ctx = item["context"]
options = item["options"]
label = item["label"]
lls = [loglikelihood(ctx, opt) for opt in options]
pred = max(range(len(lls)), key=lambda i: lls[i])
if pred == label:
correct += 1
total += 1
if total % 50 == 0:
print(f" {total}/{n} done, acc so far: {100*correct/total:.1f}%")
return 100.0 * correct / max(total, 1)
results = {}
print("Loading PIQA...")
piqa = load_dataset("ybisk/piqa", split="validation")
piqa_items = [{"context":