Download train_ram_18m.py from Compactbot/ram-18m-training-scripts: direct link, hf CLI and curl.
- Browser
- Download file 15.3 kB
-
https://huggingface.co/Compactbot/ram-18m-training-scripts/resolve/main/train_ram_18m.py
- Command line
-
hf download hf://Compactbot/ram-18m-training-scripts/train_ram_18m.py
-
curl -L -o train_ram_18m.py https://huggingface.co/Compactbot/ram-18m-training-scripts/resolve/main/train_ram_18m.py
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": |