Blogsadvanced

Train Your LLM from Scratch

A comprehensive guide to training a language model from scratch � from data preparation and tokenization through pretraining, instruction tuning, and reasoning with RLHF/DPO.

2 hoursLLM Training, PyTorch, Pretraining, SFT, LoRA, RLHF, DPOSeries: Train Your LLM from Scratch

prerequisites

3
  • Python proficiency
  • PyTorch basics
  • Understanding of transformer architecture

Stage 0: Preparation

Before writing a single training loop, you need three things: an environment that won't crash mid-run, a tokenizer that can represent your data, and a data pipeline that feeds batches efficiently. Skip any of these and you'll waste GPU hours debugging.

Environment Setup

Bash
# Create a clean environment
conda create -n llm-train python=3.11 -y
conda activate llm-train

# Core dependencies
pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121
pip install transformers datasets tokenizers wandb accelerate
pip install flash-attn --no-build-isolation  # FlashAttention-2

# Verify GPU
python -c "import torch; print(f'CUDA: {torch.cuda.is_available()}, GPUs: {torch.cuda.device_count()}')"

Building a Tokenizer

We'll train a BPE tokenizer from scratch using HuggingFace `tokenizers`. This is the same approach used by Llama, GPT, and Mistral.

Python
from tokenizers import Tokenizer, models, trainers, pre_tokenizers, decoders

def train_tokenizer(
    data_files: list[str],
    vocab_size: int = 32_000,
    save_path: str = "tokenizer.json"
):
    tokenizer = Tokenizer(models.BPE())
    tokenizer.pre_tokenizer = pre_tokenizers.ByteLevel(add_prefix_space=False)
    tokenizer.decoder = decoders.ByteLevel()

    trainer = trainers.BpeTrainer(
        vocab_size=vocab_size,
        special_tokens=["<|pad|>", "<|eos|>", "<|bos|>", "<|unk|>"],
        min_frequency=2,
        show_progress=True,
    )

    tokenizer.train(data_files, trainer)
    tokenizer.save(save_path)
    print(f"Trained tokenizer: {tokenizer.get_vocab_size()} tokens")
    return tokenizer

# Train on your corpus
tokenizer = train_tokenizer(["corpus_part1.txt", "corpus_part2.txt"])

# Test it
encoded = tokenizer.encode("The transformer architecture revolutionized NLP.")
print(f"Tokens: {encoded.tokens}")
print(f"IDs:    {encoded.ids}")

Key decisions:

  • Vocab size: 32K is a good default. Larger (64K) improves multilingual; smaller (16K) reduces embedding params
  • Special tokens: At minimum you need pad, eos, bos, unk. Add `<|im_start|>` and `<|im_end|>` if you plan to instruction-tune later

Data Pipeline

Efficient data loading is critical. We pack multiple documents into fixed-length sequences to maximize GPU utilization:

Python
import torch
from torch.utils.data import Dataset, DataLoader
from pathlib import Path
import numpy as np

class PretrainingDataset(Dataset):
    """Concatenates all documents and chunks into fixed-length sequences."""

    def __init__(
        self,
        data_dir: str,
        tokenizer,
        seq_length: int = 2048,
        stride: int = 2048,  # No overlap by default
    ):
        self.seq_length = seq_length
        self.stride = stride

        # Tokenize and concatenate all files
        all_ids = []
        for file in sorted(Path(data_dir).glob("*.txt")):
            text = file.read_text(encoding="utf-8")
            encoded = tokenizer.encode(text)
            all_ids.extend(encoded.ids)
            all_ids.append(tokenizer.token_to_id("<|eos|>"))

        self.data = np.array(all_ids, dtype=np.uint16)
        self.n_chunks = max(1, (len(self.data) - seq_length) // stride)
        print(f"Dataset: {len(self.data):,} tokens → {self.n_chunks:,} chunks")

    def __len__(self):
        return self.n_chunks

    def __getitem__(self, idx):
        start = idx * self.stride
        end = start + self.seq_length + 1  # +1 for target shift

        chunk = torch.tensor(self.data[start:end], dtype=torch.long)
        x = chunk[:-1]   # Input
        y = chunk[1:]    # Target (shifted by 1)
        return x, y

def create_dataloaders(
    train_dir: str,
    val_dir: str,
    tokenizer,
    batch_size: int = 8,
    seq_length: int = 2048,
):
    train_ds = PretrainingDataset(train_dir, tokenizer, seq_length)
    val_ds = PretrainingDataset(val_dir, tokenizer, seq_length)

    train_loader = DataLoader(
        train_ds, batch_size=batch_size, shuffle=True,
        num_workers=4, pin_memory=True, drop_last=True,
    )
    val_loader = DataLoader(
        val_ds, batch_size=batch_size, shuffle=False,
        num_workers=2, pin_memory=True,
    )
    return train_loader, val_loader

Data Preparation Checklist

StepActionWhy
1Deduplicate documentsPrevents memorization of repeated text
2Filter low-quality textRemoves boilerplate, ads, HTML artifacts
3Shuffle at document levelPrevents domain clustering in batches
4Split train/val (99/1)Val set should be representative but small
5Tokenize and save as binaryAvoids re-tokenizing every training run

Hyperparameter Cheat Sheet

For a ~125M parameter model (good for learning):

Python
config = {
    "vocab_size": 32_000,
    "d_model": 768,
    "n_heads": 12,
    "n_layers": 12,
    "d_ff": 3072,          # 4 * d_model
    "seq_length": 2048,
    "dropout": 0.1,
    "batch_size": 8,       # Per GPU
    "gradient_accumulation": 4,
    "lr": 3e-4,
    "warmup_steps": 1000,
    "total_steps": 100_000,
    "weight_decay": 0.1,
}

Next: Stage 1 — Pretraining, where we build the model and write the training loop.


Stage 1: Pretraining

Pretraining is where a language model learns to predict the next token. This is the most compute-intensive phase — you're teaching the model the statistical patterns of language from raw text.

Model Architecture

We'll build a decoder-only transformer (GPT-style). This is the same architecture used by GPT, Llama, and Mistral:

Python
import torch
import torch.nn as nn
import torch.nn.functional as F
import math

class RMSNorm(nn.Module):
    """RMSNorm (used by Llama, Mistral instead of LayerNorm)."""
    def __init__(self, dim: int, eps: float = 1e-6):
        super().__init__()
        self.eps = eps
        self.weight = nn.Parameter(torch.ones(dim))

    def forward(self, x):
        norm = torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
        return x * norm * self.weight


class CausalSelfAttention(nn.Module):
    def __init__(self, d_model: int, n_heads: int, max_seq_len: int = 2048):
        super().__init__()
        assert d_model % n_heads == 0
        self.n_heads = n_heads
        self.d_k = d_model // n_heads

        self.qkv = nn.Linear(d_model, 3 * d_model, bias=False)
        self.proj = nn.Linear(d_model, d_model, bias=False)

        # Causal mask
        mask = torch.triu(torch.ones(max_seq_len, max_seq_len), diagonal=1).bool()
        self.register_buffer("mask", mask)

    def forward(self, x):
        B, T, C = x.shape

        # Project Q, K, V in one shot
        qkv = self.qkv(x).reshape(B, T, 3, self.n_heads, self.d_k)
        q, k, v = qkv.unbind(dim=2)
        q, k, v = [t.transpose(1, 2) for t in (q, k, v)]  # [B, heads, T, d_k]

        # Scaled dot-product attention with causal mask
        scores = (q @ k.transpose(-2, -1)) / math.sqrt(self.d_k)
        scores = scores.masked_fill(self.mask[:T, :T], float("-inf"))
        attn = F.softmax(scores, dim=-1)

        out = (attn @ v).transpose(1, 2).reshape(B, T, C)
        return self.proj(out)


class FeedForward(nn.Module):
    """SwiGLU feedforward (used by Llama, Mistral)."""
    def __init__(self, d_model: int, d_ff: int):
        super().__init__()
        self.w1 = nn.Linear(d_model, d_ff, bias=False)
        self.w2 = nn.Linear(d_ff, d_model, bias=False)
        self.w3 = nn.Linear(d_model, d_ff, bias=False)

    def forward(self, x):
        return self.w2(F.silu(self.w1(x)) * self.w3(x))


class TransformerBlock(nn.Module):
    def __init__(self, d_model, n_heads, d_ff):
        super().__init__()
        self.norm1 = RMSNorm(d_model)
        self.attn = CausalSelfAttention(d_model, n_heads)
        self.norm2 = RMSNorm(d_model)
        self.ff = FeedForward(d_model, d_ff)

    def forward(self, x):
        x = x + self.attn(self.norm1(x))
        x = x + self.ff(self.norm2(x))
        return x


class GPT(nn.Module):
    def __init__(self, vocab_size, d_model, n_heads, n_layers, d_ff, max_seq_len):
        super().__init__()
        self.tok_emb = nn.Embedding(vocab_size, d_model)
        self.pos_emb = nn.Embedding(max_seq_len, d_model)
        self.blocks = nn.ModuleList([
            TransformerBlock(d_model, n_heads, d_ff) for _ in range(n_layers)
        ])
        self.norm = RMSNorm(d_model)
        self.head = nn.Linear(d_model, vocab_size, bias=False)

        # Weight tying
        self.head.weight = self.tok_emb.weight

        n_params = sum(p.numel() for p in self.parameters())
        print(f"Model parameters: {n_params / 1e6:.1f}M")

    def forward(self, idx, targets=None):
        B, T = idx.shape
        pos = torch.arange(T, device=idx.device)

        x = self.tok_emb(idx) + self.pos_emb(pos)
        for block in self.blocks:
            x = block(x)
        x = self.norm(x)
        logits = self.head(x)

        loss = None
        if targets is not None:
            loss = F.cross_entropy(
                logits.view(-1, logits.size(-1)),
                targets.view(-1),
                ignore_index=-1,
            )
        return logits, loss

Training Loop

Here's the complete training loop with mixed precision, gradient accumulation, and cosine LR schedule:

Python
from torch.cuda.amp import autocast, GradScaler
import wandb

def train(model, train_loader, val_loader, config):
    device = torch.device("cuda")
    model = model.to(device)

    optimizer = torch.optim.AdamW(
        model.parameters(),
        lr=config["lr"],
        betas=(0.9, 0.95),
        weight_decay=config["weight_decay"],
    )

    # Cosine LR with warmup
    def lr_schedule(step):
        if step < config["warmup_steps"]:
            return step / config["warmup_steps"]
        progress = (step - config["warmup_steps"]) / (config["total_steps"] - config["warmup_steps"])
        return 0.1 + 0.9 * 0.5 * (1 + math.cos(math.pi * progress))

    scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_schedule)
    scaler = GradScaler()  # For mixed precision

    wandb.init(project="llm-from-scratch", config=config)

    global_step = 0
    for epoch in range(config.get("epochs", 1)):
        model.train()
        for batch_idx, (x, y) in enumerate(train_loader):
            x, y = x.to(device), y.to(device)

            with autocast(dtype=torch.bfloat16):
                _, loss = model(x, y)
                loss = loss / config["gradient_accumulation"]

            scaler.scale(loss).backward()

            if (batch_idx + 1) % config["gradient_accumulation"] == 0:
                scaler.unscale_(optimizer)
                torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
                scaler.step(optimizer)
                scaler.update()
                optimizer.zero_grad()
                scheduler.step()
                global_step += 1

                if global_step % 100 == 0:
                    wandb.log({
                        "train/loss": loss.item() * config["gradient_accumulation"],
                        "train/lr": scheduler.get_last_lr()[0],
                        "train/step": global_step,
                    })

                if global_step % 1000 == 0:
                    val_loss = validate(model, val_loader, device)
                    wandb.log({"val/loss": val_loss, "val/perplexity": math.exp(val_loss)})
                    print(f"Step {global_step} | Val loss: {val_loss:.4f} | PPL: {math.exp(val_loss):.2f}")
                    save_checkpoint(model, optimizer, global_step)

                if global_step >= config["total_steps"]:
                    return

Validation

Perplexity is the standard metric for pretraining — lower is better. A perplexity of 20 means the model is "20-way confused" on average.

Python
@torch.no_grad()
def validate(model, val_loader, device, max_batches=50):
    model.eval()
    total_loss = 0
    n_batches = 0

    for x, y in val_loader:
        x, y = x.to(device), y.to(device)
        with autocast(dtype=torch.bfloat16):
            _, loss = model(x, y)
        total_loss += loss.item()
        n_batches += 1
        if n_batches >= max_batches:
            break

    model.train()
    return total_loss / n_batches

def save_checkpoint(model, optimizer, step, path="checkpoints"):
    Path(path).mkdir(exist_ok=True)
    torch.save({
        "model": model.state_dict(),
        "optimizer": optimizer.state_dict(),
        "step": step,
    }, f"{path}/step_{step}.pt")

What to Watch For

MetricHealthyUnhealthy
Training lossSmooth downward curveSpikes, plateaus early
Val lossTracks train loss closelyDiverges from train loss
Gradient normStable around 0.1–1.0Exploding (>10) or vanishing
Learning rateSmooth warmup → cosine decay—
PerplexitySteadily decreasingStuck above 100 after 10K steps

Generating Text (Sanity Check)

Python
@torch.no_grad()
def generate(model, tokenizer, prompt, max_new_tokens=100, temperature=0.8):
    model.eval()
    device = next(model.parameters()).device
    ids = tokenizer.encode(prompt).ids
    x = torch.tensor([ids], device=device)

    for _ in range(max_new_tokens):
        logits, _ = model(x[:, -2048:])  # Truncate to max seq len
        logits = logits[:, -1, :] / temperature
        probs = F.softmax(logits, dim=-1)
        next_id = torch.multinomial(probs, 1)
        x = torch.cat([x, next_id], dim=1)

        if next_id.item() == tokenizer.token_to_id("<|eos|>"):
            break

    return tokenizer.decode(x[0].tolist())

Next: Stage 2 — Instruction Tuning, where we teach the model to follow instructions.


Stage 2: Instruction Tuning

A pretrained model can complete text, but it can't follow instructions. Instruction tuning (SFT — Supervised Fine-Tuning) teaches the model to respond helpfully to user requests. This is the step that turns a "text completer" into a "chatbot."

Chat Template

First, define a chat format your model will learn:

Python
CHAT_TEMPLATE = """<|im_start|>system
{system}<|im_end|>
<|im_start|>user
{user}<|im_end|>
<|im_start|>assistant
{assistant}<|im_end|>"""

def format_conversation(system: str, user: str, assistant: str) -> str:
    return CHAT_TEMPLATE.format(
        system=system, user=user, assistant=assistant
    )

# Example
formatted = format_conversation(
    system="You are a helpful assistant.",
    user="What is gradient descent?",
    assistant="Gradient descent is an optimization algorithm..."
)

Instruction Dataset

Python
from datasets import load_dataset
from torch.utils.data import Dataset

class InstructionDataset(Dataset):
    def __init__(self, tokenizer, max_length=2048, split="train"):
        # Use a public instruction dataset
        raw = load_dataset("tatsu-lab/alpaca", split=split)
        self.samples = []
        self.tokenizer = tokenizer
        self.max_length = max_length

        for item in raw:
            instruction = item["instruction"]
            if item.get("input"):
                instruction += f"\n\nInput: {item['input']}"

            text = format_conversation(
                system="You are a helpful assistant.",
                user=instruction,
                assistant=item["output"],
            )
            ids = tokenizer.encode(text).ids
            if len(ids) <= max_length:
                self.samples.append(ids)

        print(f"Instruction dataset: {len(self.samples)} samples")

    def __len__(self):
        return len(self.samples)

    def __getitem__(self, idx):
        ids = self.samples[idx]

        # Pad to max_length
        padded = ids + [self.tokenizer.token_to_id("<|pad|>")] * (self.max_length - len(ids))
        x = torch.tensor(padded[:-1], dtype=torch.long)
        y = torch.tensor(padded[1:], dtype=torch.long)

        # Mask: only compute loss on the assistant's response
        # Find where assistant response starts
        y[:len(ids) // 2] = -1  # Simplified — mask system + user tokens
        return x, y

LoRA: Parameter-Efficient Fine-Tuning

Full fine-tuning updates all parameters. LoRA freezes the base model and adds small trainable matrices, reducing memory by ~10x:

Python
class LoRALinear(nn.Module):
    """Low-Rank Adaptation for efficient fine-tuning."""
    def __init__(self, base_layer: nn.Linear, rank: int = 16, alpha: float = 32):
        super().__init__()
        self.base = base_layer
        self.base.weight.requires_grad = False  # Freeze base

        d_in, d_out = base_layer.in_features, base_layer.out_features
        self.lora_A = nn.Parameter(torch.randn(d_in, rank) * 0.01)
        self.lora_B = nn.Parameter(torch.zeros(rank, d_out))
        self.scale = alpha / rank

    def forward(self, x):
        base_out = self.base(x)
        lora_out = (x @ self.lora_A @ self.lora_B) * self.scale
        return base_out + lora_out

def apply_lora(model, rank=16, alpha=32):
    """Apply LoRA to all attention projection layers."""
    for name, module in model.named_modules():
        if isinstance(module, CausalSelfAttention):
            module.qkv = LoRALinear(module.qkv, rank, alpha)
            module.proj = LoRALinear(module.proj, rank, alpha)

    trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
    total = sum(p.numel() for p in model.parameters())
    print(f"LoRA: {trainable/1e6:.1f}M trainable / {total/1e6:.1f}M total ({100*trainable/total:.1f}%)")

SFT Training Loop

Python
def instruction_tune(model, train_ds, val_ds, config):
    device = torch.device("cuda")
    model = model.to(device)

    # Only optimize LoRA parameters
    optimizer = torch.optim.AdamW(
        [p for p in model.parameters() if p.requires_grad],
        lr=2e-5,            # Much lower LR than pretraining
        weight_decay=0.01,
    )

    train_loader = DataLoader(train_ds, batch_size=4, shuffle=True)
    val_loader = DataLoader(val_ds, batch_size=4)

    for epoch in range(3):  # SFT typically needs only 1-3 epochs
        model.train()
        epoch_loss = 0
        for x, y in train_loader:
            x, y = x.to(device), y.to(device)

            with autocast(dtype=torch.bfloat16):
                _, loss = model(x, y)

            loss.backward()
            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
            optimizer.step()
            optimizer.zero_grad()
            epoch_loss += loss.item()

        val_loss = validate(model, val_loader, device)
        print(f"Epoch {epoch+1} | Train: {epoch_loss/len(train_loader):.4f} | Val: {val_loss:.4f}")

Evaluation

After instruction tuning, evaluate on standard benchmarks:

Python
def evaluate_instruction_following(model, tokenizer, test_prompts):
    """Manual evaluation of instruction following quality."""
    results = []
    for prompt in test_prompts:
        text = format_conversation(
            system="You are a helpful assistant.",
            user=prompt,
            assistant="",
        )
        # Remove the final <|im_end|> so the model generates the response
        text = text.rsplit("<|im_end|>", 1)[0]

        response = generate(model, tokenizer, text, max_new_tokens=256)
        results.append({"prompt": prompt, "response": response})
        print(f"Q: {prompt}")
        print(f"A: {response}\n")
    return results

test_prompts = [
    "Explain quantum computing in simple terms.",
    "Write a Python function to find the nth Fibonacci number.",
    "What are the pros and cons of remote work?",
    "Summarize the key ideas of the attention mechanism.",
]

SFT Pitfalls

IssueSymptomFix
Catastrophic forgettingModel loses general knowledgeLower LR, use LoRA, fewer epochs
OverfittingVal loss increases after epoch 1More data, higher dropout, early stopping
Template leakageModel outputs `<im_start
RepetitionModel loops the same phraseAdd repetition penalty during generation

Next: Stage 3 — Reasoning, where we teach the model to think step by step.


Stage 3: Reasoning & Alignment

An instruction-tuned model follows commands, but it doesn't think. This stage teaches the model to reason step-by-step and aligns its behavior with human preferences using RLHF or DPO.

Chain-of-Thought Training Data

The key insight: if you train on data that contains explicit reasoning steps, the model learns to reason. We create "thinking" traces:

Python
COT_TEMPLATE = """<|im_start|>system
You are a helpful assistant. Think step by step before answering.<|im_end|>
<|im_start|>user
{question}<|im_end|>
<|im_start|>assistant
<think>
{reasoning}
</think>

{answer}<|im_end|>"""

# Example training data
cot_examples = [
    {
        "question": "If a train travels 120 km in 2 hours, what is its speed in m/s?",
        "reasoning": """Step 1: Find speed in km/h.
Speed = distance / time = 120 km / 2 h = 60 km/h.

Step 2: Convert km/h to m/s.
1 km = 1000 m, 1 h = 3600 s.
60 km/h = 60 × 1000 / 3600 = 16.67 m/s.""",
        "answer": "The train's speed is approximately 16.67 m/s."
    },
    {
        "question": "A store has a 25% off sale. If an item costs $80, and there's an additional 10% member discount applied after, what's the final price?",
        "reasoning": """Step 1: Apply the 25% sale discount.
25% of $80 = $20. Price after sale: $80 - $20 = $60.

Step 2: Apply the 10% member discount on the sale price.
10% of $60 = $6. Price after member discount: $60 - $6 = $54.

Step 3: The discounts are applied sequentially, not combined.
Total discount is not 35% — it's 25% then 10% of the reduced price.""",
        "answer": "The final price is $54.00."
    },
]

Rejection Sampling for Reasoning Data

Generate multiple attempts and keep only the ones that arrive at the correct answer:

Python
def generate_cot_data(model, tokenizer, problems, n_samples=8, temperature=0.7):
    """Generate chain-of-thought data via rejection sampling."""
    good_samples = []

    for problem in problems:
        prompt = f"""<|im_start|>system
Think step by step.<|im_end|>
<|im_start|>user
{problem['question']}<|im_end|>
<|im_start|>assistant
<think>
"""
        candidates = []
        for _ in range(n_samples):
            response = generate(model, tokenizer, prompt,
                              max_new_tokens=512, temperature=temperature)
            # Check if the final answer matches
            if problem["answer"] in response:
                candidates.append(response)

        if candidates:
            # Keep the shortest correct reasoning (Occam's razor)
            best = min(candidates, key=len)
            good_samples.append({
                "question": problem["question"],
                "response": best,
            })

    print(f"Generated {len(good_samples)}/{len(problems)} valid CoT samples")
    return good_samples

DPO: Direct Preference Optimization

DPO is a simpler alternative to RLHF. Instead of training a reward model, you directly optimize on preference pairs (chosen vs rejected):

Python
class DPOTrainer:
    def __init__(self, model, ref_model, tokenizer, beta=0.1, lr=5e-7):
        self.model = model
        self.ref_model = ref_model  # Frozen copy of the SFT model
        self.tokenizer = tokenizer
        self.beta = beta

        # Freeze reference model
        for p in self.ref_model.parameters():
            p.requires_grad = False

        self.optimizer = torch.optim.AdamW(
            model.parameters(), lr=lr, weight_decay=0.01
        )

    def compute_log_probs(self, model, input_ids, labels):
        """Compute log probabilities of the target tokens."""
        logits, _ = model(input_ids)
        log_probs = F.log_softmax(logits, dim=-1)

        # Gather log probs for actual tokens
        token_log_probs = log_probs.gather(-1, labels.unsqueeze(-1)).squeeze(-1)

        # Mask padding
        mask = (labels != -1).float()
        return (token_log_probs * mask).sum(-1) / mask.sum(-1)

    def dpo_loss(self, chosen_ids, chosen_labels, rejected_ids, rejected_labels):
        """DPO loss: maximize margin between chosen and rejected."""
        # Policy log probs
        pi_chosen = self.compute_log_probs(self.model, chosen_ids, chosen_labels)
        pi_rejected = self.compute_log_probs(self.model, rejected_ids, rejected_labels)

        # Reference log probs (no gradient)
        with torch.no_grad():
            ref_chosen = self.compute_log_probs(self.ref_model, chosen_ids, chosen_labels)
            ref_rejected = self.compute_log_probs(self.ref_model, rejected_ids, rejected_labels)

        # DPO objective
        pi_logratios = pi_chosen - pi_rejected
        ref_logratios = ref_chosen - ref_rejected
        logits = self.beta * (pi_logratios - ref_logratios)

        loss = -F.logsigmoid(logits).mean()
        return loss

    def train_step(self, batch):
        loss = self.dpo_loss(
            batch["chosen_ids"], batch["chosen_labels"],
            batch["rejected_ids"], batch["rejected_labels"],
        )
        loss.backward()
        torch.nn.utils.clip_grad_norm_(self.model.parameters(), 1.0)
        self.optimizer.step()
        self.optimizer.zero_grad()
        return loss.item()

Preference Data Format

Python
preference_pairs = [
    {
        "prompt": "Explain recursion.",
        "chosen": "<think>\nRecursion is when a function calls itself...\n</think>\n\nRecursion is a programming concept where a function calls itself to solve smaller instances of the same problem...",
        "rejected": "Recursion is when something is recursive."
    },
    {
        "prompt": "Is water wet?",
        "chosen": "<think>\nThis is a nuanced question. 'Wet' typically means...\n</think>\n\nThis depends on how you define 'wet.' If wet means 'covered in water,' then water itself isn't wet — it *makes* things wet...",
        "rejected": "Yes, water is wet because it's a liquid."
    },
]

Evaluation: Reasoning Benchmarks

Python
def evaluate_reasoning(model, tokenizer, benchmark="gsm8k"):
    """Evaluate on GSM8K (grade school math)."""
    dataset = load_dataset("gsm8k", "main", split="test")
    correct = 0
    total = 0

    for item in dataset:
        prompt = f"""<|im_start|>system
Solve step by step. Put your final numerical answer after "#### ".<|im_end|>
<|im_start|>user
{item['question']}<|im_end|>
<|im_start|>assistant
<think>
"""
        response = generate(model, tokenizer, prompt, max_new_tokens=512, temperature=0.0)

        # Extract the final answer
        predicted = extract_answer(response)
        expected = item["answer"].split("####")[-1].strip()

        if predicted == expected:
            correct += 1
        total += 1

    accuracy = correct / total
    print(f"{benchmark} accuracy: {accuracy:.1%} ({correct}/{total})")
    return accuracy

Inference & Efficiency Metrics

These metrics measure how well an AI model runs on hardware and its responsiveness in production.

Throughput (Tokens Per Second)

Throughput measures the total volume of output tokens a model generates every second. High TPS is critical for high-traffic applications and batch processing.

For a given model, throughput depends on:

  • Hardware: GPU type (A100, H100), number of GPUs, interconnect bandwidth
  • Batch size: Larger batches improve throughput but increase latency
  • Quantization: INT8/INT4 quantization reduces memory and increases speed at the cost of some quality
  • Serving framework: vLLM, TensorRT-LLM, and SGLang provide optimized inference kernels
Model SizeTypical TPS (A100)Typical TPS (H100)
7B40-8080-150
13B25-5050-100
70B8-1520-40

Time to First Token (TTFT)

TTFT is the delay between a user sending a prompt and seeing the very first character of the response. Sub-200ms is the standard for a "snappy" user experience.

TTFT is dominated by the prefill phase � where the model processes all input tokens in parallel through KV-cache computation. Techniques to reduce TTFT:

  • Speculative decoding: Use a small draft model to propose tokens, verified by the large model
  • Prefix caching: Cache the KV states of common system prompts
  • Chunked prefill: Break long prompts into chunks to overlap with decode

Context Window

The context window is the "short-term memory" of the model, measured in tokens. A larger window allows the AI to process entire books or massive codebases in a single pass.

ModelContext Length~Pages of Text
GPT-4o128K~300 pages
Claude 3.5200K~500 pages
Gemini 1.5 Pro2M~5,000 pages

Key techniques for extending context:

GPU & Memory Utilization

Tracks how much hardware resources (VRAM) the model consumes. Lower utilization per query allows for more simultaneous users.

Key optimization techniques:


Quality & Intelligence Metrics

These quantify how "smart" or accurate a model's outputs are.

Perplexity

Perplexity is a mathematical measure of how "surprised" a model is by new data. Lower is better, indicating the model has a stronger internal grasp of language patterns.

Perplexity = exp(average cross-entropy loss). A perplexity of 10 means the model is, on average, "10-way uncertain" about each next token.

StageTypical Perplexity
Early pretraining100-1000+
Converged pretraining5-15
Domain-specific fine-tune3-8

Important caveat: Perplexity only measures next-token prediction quality on a held-out set. A model with great perplexity can still produce bad instruction-following results.

Accuracy & F1 Score

Standard metrics for classification and extraction tasks:

  • Accuracy: Percentage of correct predictions overall
  • Precision: Of items flagged as positive, how many actually are? (Reduces false positives)
  • Recall: Of all actual positives, how many did we find? (Reduces false negatives)
  • F1 Score: The harmonic mean of precision and recall � the "gold standard" for balancing both

For LLM benchmarks, the most commonly referenced evaluations include:

  • MMLU: 57 subjects ranging from STEM to humanities
  • HumanEval: Code generation benchmark
  • GSM8K: Grade school math reasoning
  • HellaSwag: Commonsense reasoning

Elo Rating (Human Preference)

Elo rating, borrowed from chess, is used by the LMSYS Chatbot Arena to rank models based on blind side-by-side human testing. Users see two anonymous model outputs and pick the better one.

This is arguably the most reliable quality signal because:

  • It captures holistic quality (helpfulness, safety, style, accuracy)
  • It's resistant to benchmark gaming � models can't overfit to specific test sets
  • It reflects real user preferences, not proxy metrics

Hallucination Rate

The hallucination rate measures how frequently a model generates factually incorrect or unsupported information. This is one of the biggest challenges in deploying LLMs.

Two types of hallucination:

  • Intrinsic: Contradicts the provided source material
  • Extrinsic: Generates information not supported by any source

Mitigation strategies:

Scaling Laws

Both efficiency and quality metrics improve with scale, but in predictable ways described by scaling laws:

  • Compute-optimal training (Chinchilla scaling): The optimal model size and data size grow proportionally with compute budget
  • Inference scaling: Techniques like test-time compute allow models to "think longer" on harder problems, trading latency for quality
  • Data scaling: Textbooks Are All You Need showed that high-quality data can substitute for raw scale

The Full Training Pipeline

StageWhatDataEpochsLRResult
0PreparationRaw text corpus——Tokenizer + data pipeline
1PretrainingRaw text13e-4Next-token predictor
2SFTInstruction pairs1–32e-5Instruction follower
3aCoTReasoning traces1–21e-5Step-by-step thinker
3bDPOPreference pairs15e-7Aligned reasoner

What You've Built

By completing all four stages, you've built a model that:

  1. Understands language (pretraining)
  2. Follows instructions (SFT)
  3. Thinks before answering (chain-of-thought)
  4. Prefers good answers over bad ones (DPO alignment)

This is the same pipeline used by frontier models — just at a smaller scale. The architecture, loss functions, and training stages are identical.


This completes the "Train Your LLM from Scratch" series. For production-scale training, explore DeepSpeed, FSDP, and multi-node distributed training.

Zizhao Huloading加载中