OxMini / src /oxmini /training.py
Shivam3002's picture
Publish trained OxMini checkpoint and measured model card
46144df verified
Raw
History Blame Contribute Delete
7.9 kB
"""Shared, CPU-first training and evaluation utilities."""
from __future__ import annotations
from dataclasses import replace
import math
import os
from pathlib import Path
import random
import shutil
from typing import Literal
import numpy as np
import torch
from .baseline_gpt import BaselineGPTForCausalLM
from .config import OxMiniConfig
from .model import OxMiniForCausalLM
Variant = Literal["baseline", "hybrid", "full"]
def set_reproducible_seed(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
def configure_cpu_threads(requested: int | None = None) -> int:
available = os.cpu_count() or 1
# More threads can hurt small matrix multiplications through scheduling
# overhead. Ten maps well to the target M4 performance-core workload while
# remaining overrideable for other CPUs.
threads = requested if requested and requested > 0 else min(10, available)
torch.set_num_threads(threads)
try:
torch.set_num_interop_threads(max(1, min(4, threads)))
except RuntimeError:
# PyTorch only permits setting this once per process.
pass
return threads
def build_model(config: OxMiniConfig, variant: Variant) -> torch.nn.Module:
# Centralizing variant construction prevents train/eval/ablation scripts
# from silently disagreeing about what "hybrid" or "full" means.
if variant == "baseline":
return BaselineGPTForCausalLM(replace(config, use_mhc=False))
if variant == "hybrid":
return OxMiniForCausalLM(replace(config, use_mhc=False))
if variant == "full":
return OxMiniForCausalLM(replace(config, use_mhc=True))
raise ValueError(f"unsupported variant: {variant}")
def parameter_count(model: torch.nn.Module) -> int:
return sum(parameter.numel() for parameter in model.parameters())
def load_split(data_dir: str | Path, split: str) -> np.memmap:
path = Path(data_dir) / f"{split}.bin"
if not path.exists():
raise FileNotFoundError(f"missing {path}; run scripts/prepare_data.py first")
# Read-only memory maps avoid loading the complete token stream into RAM and
# let the operating system cache exactly the random windows training uses.
return np.memmap(path, dtype=np.uint16, mode="r")
def sample_batch(
data: np.ndarray,
batch_size: int,
block_size: int,
rng: np.random.Generator,
device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor]:
if len(data) <= block_size:
raise ValueError(f"split has {len(data)} tokens, smaller than block_size={block_size}")
# Sample independent contiguous windows. y is x shifted one character into
# the future; converting to int64 happens only for the selected windows.
starts = rng.integers(0, len(data) - block_size, size=batch_size)
offsets = np.arange(block_size)
x_np = np.asarray(data[starts[:, None] + offsets], dtype=np.int64)
y_np = np.asarray(data[starts[:, None] + offsets + 1], dtype=np.int64)
x = torch.from_numpy(x_np).to(device=device, dtype=torch.long)
y = torch.from_numpy(y_np).to(device=device, dtype=torch.long)
return x, y
@torch.inference_mode()
def estimate_loss(
model: torch.nn.Module,
data: np.ndarray,
batch_size: int,
block_size: int,
batches: int,
seed: int,
device: torch.device,
) -> float:
was_training = model.training
model.eval()
rng = np.random.default_rng(seed)
values: list[float] = []
for _ in range(batches):
x, y = sample_batch(data, batch_size, block_size, rng, device)
loss = model(x, y).loss
if loss is None or not torch.isfinite(loss):
raise FloatingPointError("non-finite evaluation loss")
values.append(float(loss.item()))
if was_training:
model.train()
return float(np.mean(values))
@torch.inference_mode()
def estimate_loss_and_accuracy(
model: torch.nn.Module,
data: np.ndarray,
batch_size: int,
block_size: int,
batches: int,
seed: int,
device: torch.device,
) -> tuple[float, float]:
was_training = model.training
model.eval()
rng = np.random.default_rng(seed)
losses: list[float] = []
correct = 0
total = 0
for _ in range(batches):
x, y = sample_batch(data, batch_size, block_size, rng, device)
output = model(x, y)
if output.loss is None or not torch.isfinite(output.loss):
raise FloatingPointError("non-finite evaluation loss")
losses.append(float(output.loss.item()))
predictions = output.logits.argmax(dim=-1)
correct += int((predictions == y).sum().item())
total += y.numel()
if was_training:
model.train()
return float(np.mean(losses)), correct / max(total, 1)
def learning_rate_at_step(
step: int,
total_steps: int,
peak_lr: float,
warmup_steps: int,
min_lr_ratio: float = 0.1,
) -> float:
# Warm up linearly, then decay smoothly to 10% of the peak rather than zero
# so late updates remain capable of small corrections.
if warmup_steps > 0 and step < warmup_steps:
return peak_lr * (step + 1) / warmup_steps
progress = (step - warmup_steps) / max(1, total_steps - warmup_steps)
cosine = 0.5 * (1.0 + math.cos(math.pi * min(max(progress, 0.0), 1.0)))
return peak_lr * (min_lr_ratio + (1.0 - min_lr_ratio) * cosine)
def save_training_checkpoint(
directory: str | Path,
model: torch.nn.Module,
optimizer: torch.optim.Optimizer,
config: OxMiniConfig,
tokenizer_path: str | Path,
variant: Variant,
step: int,
best_val_loss: float,
rng: np.random.Generator,
metadata: dict[str, object] | None = None,
) -> Path:
directory = Path(directory)
directory.mkdir(parents=True, exist_ok=True)
# ``model.pt`` is intentionally a local resumable checkpoint, not the Hub
# artifact: optimizer and RNG state require pickle-capable torch.save. The
# publishing path extracts model weights into safetensors for inference.
payload = {
"format_version": 1,
"variant": variant,
"step": step,
"best_val_loss": best_val_loss,
"config": config.to_dict(),
"model_state": model.state_dict(),
"optimizer_state": optimizer.state_dict(),
"numpy_rng_state": rng.bit_generator.state,
"torch_rng_state": torch.get_rng_state(),
"metadata": metadata or {},
}
target = directory / "model.pt"
# Write-then-rename avoids leaving a half-written checkpoint if a run is
# interrupted while the optimizer state is being serialized.
temporary = directory / "model.pt.tmp"
torch.save(payload, temporary)
temporary.replace(target)
config.save(directory / "config.yaml")
shutil.copy2(tokenizer_path, directory / "tokenizer.json")
return target
def load_training_checkpoint(
checkpoint: str | Path,
device: torch.device,
load_optimizer: bool = False,
) -> tuple[torch.nn.Module, dict[str, object]]:
path = Path(checkpoint)
if path.is_dir():
path = path / "model.pt"
if not path.exists():
raise FileNotFoundError(path)
payload = torch.load(path, map_location=device, weights_only=False)
config = OxMiniConfig.from_dict(payload["config"])
variant: Variant = payload["variant"]
model = build_model(config, variant).to(device)
model.load_state_dict(payload["model_state"])
if not load_optimizer:
payload = {key: value for key, value in payload.items() if key != "optimizer_state"}
return model, payload
def copy_checkpoint_files(source: str | Path, destination: str | Path) -> None:
source, destination = Path(source), Path(destination)
destination.mkdir(parents=True, exist_ok=True)
for name in ("model.pt", "config.yaml", "tokenizer.json"):
shutil.copy2(source / name, destination / name)