#!/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="")) tokenizer.pre_tokenizer = Whitespace() trainer = BpeTrainer(vocab_size=VOCAB, special_tokens=["", "", "", ""]) 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":