Mandeep Sidhu
Refactor experiment pipeline and add regime paper
e7a7275
Raw
History Blame Contribute Delete
8.13 kB
"""
Derived from Andrej Karpathy's nanochat project.
MIT License
Copyright (c) 2025 Andrej Karpathy
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
"""
from __future__ import annotations
import argparse
import random
import statistics
import time
from typing import Callable
import numpy as np
import torch
from dropout_decay.experiments.artifacts import write_jsonl_row
from dropout_decay.models import DropoutGPT, GPTConfig
from dropout_decay.specs import DropoutCondition, ModelSpec
def set_seed(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
def batch_seed(
seed: int, model: ModelSpec, dropout_code: int, stage: int | None
) -> int:
return (
seed * 1_000_003
+ model.n_layer * 100_003
+ model.n_head * 10_007
+ model.n_embd * 101
+ dropout_code * 37
+ (stage or 0) * 997
)
def make_batch(
tokens: np.ndarray,
token_limit: int,
batch_size: int,
block_size: int,
rng: np.random.Generator,
device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor]:
limit = min(token_limit, len(tokens))
max_start = limit - block_size - 1
if max_start <= 0:
raise ValueError("token_limit is too small for the requested block_size")
starts = rng.integers(0, max_start, size=batch_size)
x_np = np.stack([tokens[start : start + block_size] for start in starts]).astype(
np.int64
)
y_np = np.stack(
[tokens[start + 1 : start + 1 + block_size] for start in starts]
).astype(np.int64)
return (
torch.tensor(x_np, dtype=torch.long, device=device),
torch.tensor(y_np, dtype=torch.long, device=device),
)
@torch.no_grad()
def estimate_loss(
model: DropoutGPT,
tokens: np.ndarray,
token_limit: int,
batches: int,
args: argparse.Namespace,
device: torch.device,
rng_seed: int,
) -> float:
if batches <= 0:
return float("nan")
model.eval()
rng = np.random.default_rng(rng_seed)
losses: list[float] = []
for _ in range(batches):
x, y = make_batch(
tokens, token_limit, args.batch_size, args.block_size, rng, device
)
_, loss = model(x, y)
losses.append(float(loss.item()))
model.train()
return float(statistics.fmean(losses))
def train_segment(
*,
run_mode: str,
condition: DropoutCondition,
model_spec: ModelSpec,
config: GPTConfig,
train_tokens: np.ndarray,
val_tokens: np.ndarray,
token_limit: int,
steps: int,
seed: int,
args: argparse.Namespace,
device: torch.device,
dropout_fn: Callable[[int], float],
metrics_file,
trace_file,
stage: int | None = None,
model: DropoutGPT | None = None,
optimizer: torch.optim.Optimizer | None = None,
tokens_seen_start: int = 0,
) -> tuple[DropoutGPT, torch.optim.Optimizer, int, dict]:
if model is None:
set_seed(seed)
model = DropoutGPT(config).to(device)
optimizer = torch.optim.AdamW(
model.parameters(),
lr=args.lr,
betas=(0.9, 0.95),
weight_decay=args.weight_decay,
)
else:
torch.manual_seed(seed + 10_000 + (stage or 0))
if optimizer is None:
raise ValueError("optimizer is required when reusing a model")
model.train()
dropout_code = int(round(condition.initial * 10_000))
rng = np.random.default_rng(batch_seed(seed, model_spec, dropout_code, stage))
tokens_seen = tokens_seen_start
last_loss = float("nan")
active_dropout = condition.initial
t0 = time.time()
for step in range(1, steps + 1):
active_dropout = dropout_fn(tokens_seen)
model.set_dropout(active_dropout)
x, y = make_batch(
train_tokens, token_limit, args.batch_size, args.block_size, rng, device
)
_, loss = model(x, y)
optimizer.zero_grad(set_to_none=True)
loss.backward()
if args.grad_clip > 0:
torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip)
optimizer.step()
tokens_seen += args.batch_size * args.block_size
last_loss = float(loss.item())
if args.log_every > 0 and step % args.log_every == 0:
write_jsonl_row(
trace_file,
{
"event": "train_step",
"run_mode": run_mode,
"condition": condition.name,
"model_name": model_spec.name,
"seed": seed,
"stage": stage,
"step": step,
"steps": steps,
"token_limit": int(token_limit),
"tokens_seen": int(tokens_seen),
"dropout": float(active_dropout),
"train_batch_loss": last_loss,
},
)
if args.eval_every > 0 and step % args.eval_every == 0:
train_eval = estimate_loss(
model,
train_tokens,
token_limit,
args.trace_eval_batches,
args,
device,
rng_seed=seed + 20_000 + step,
)
val_eval = estimate_loss(
model,
val_tokens,
len(val_tokens),
args.trace_eval_batches,
args,
device,
rng_seed=seed + 30_000 + step,
)
write_jsonl_row(
trace_file,
{
"event": "eval_step",
"run_mode": run_mode,
"condition": condition.name,
"model_name": model_spec.name,
"seed": seed,
"stage": stage,
"step": step,
"steps": steps,
"token_limit": int(token_limit),
"tokens_seen": int(tokens_seen),
"dropout": float(active_dropout),
"train_eval_loss": train_eval,
"val_eval_loss": val_eval,
"generalization_gap": val_eval - train_eval,
},
)
train_eval = estimate_loss(
model,
train_tokens,
token_limit,
args.train_eval_batches,
args,
device,
rng_seed=seed + 40_000 + (stage or 0),
)
val_eval = estimate_loss(
model,
val_tokens,
len(val_tokens),
args.eval_batches,
args,
device,
rng_seed=seed + 50_000 + (stage or 0),
)
row = {
"run_mode": run_mode,
"condition": condition.name,
"condition_kind": condition.kind,
"seed": seed,
"stage": stage,
"token_limit": int(token_limit),
"steps": int(steps),
"tokens_seen": int(tokens_seen),
"dropout_initial": float(condition.initial),
"dropout_final": float(condition.final),
"dropout_schedule": condition.schedule,
"dropout_active_final": float(active_dropout),
"train_loss_last": last_loss,
"train_eval_loss": train_eval,
"val_eval_loss": val_eval,
"eval_loss": val_eval,
"generalization_gap": val_eval - train_eval,
"elapsed_sec": time.time() - t0,
"parameters": model.num_parameters(),
"model_config": config.to_dict(),
**model_spec.to_dict(),
}
write_jsonl_row(metrics_file, row)
return model, optimizer, tokens_seen, row