| """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 |
| |
| |
| |
| 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: |
| |
| pass |
| return threads |
|
|
|
|
| def build_model(config: OxMiniConfig, variant: Variant) -> torch.nn.Module: |
| |
| |
| 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") |
| |
| |
| 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}") |
| |
| |
| 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: |
| |
| |
| 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) |
| |
| |
| |
| 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" |
| |
| |
| 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) |
|
|