yuspec-gamedev-ai / src /train.py
fovane's picture
Upload folder using huggingface_hub
295d147 verified
Raw
History Blame Contribute Delete
5.48 kB
import argparse
import math
import os
import time
from contextlib import nullcontext
import torch
import yaml
from torch.optim import AdamW
from dataset import get_batch, load_tokens
from model import GPT, GPTConfig
device = "cuda" if torch.cuda.is_available() else "cpu"
def amp_context(enabled):
if enabled:
return torch.amp.autocast("cuda", dtype=torch.float16)
return nullcontext()
def load_config(path):
with open(path, "r", encoding="utf-8") as f:
return yaml.safe_load(f)
@torch.no_grad()
def estimate_loss(model, train_data, val_data, block_size, batch_size, eval_iters, use_amp=False):
model.eval()
out = {}
for split, data in [("train", train_data), ("val", val_data)]:
losses = []
for _ in range(eval_iters):
x, y = get_batch(data, block_size, batch_size, device)
with amp_context(use_amp):
_, loss = model(x, y)
losses.append(loss.item())
out[split] = sum(losses) / len(losses)
model.train()
return out
def get_lr(it, max_lr, min_lr, warmup_iters, max_iters):
if warmup_iters > 0 and it < warmup_iters:
return max_lr * (it + 1) / warmup_iters
if it > max_iters:
return min_lr
decay_den = max(1, max_iters - warmup_iters)
decay_ratio = min(1.0, (it - warmup_iters) / decay_den)
coeff = 0.5 * (1.0 + math.cos(math.pi * decay_ratio))
return min_lr + coeff * (max_lr - min_lr)
def save_checkpoint(path, model, cfg, step, val_loss=None):
os.makedirs(os.path.dirname(path), exist_ok=True)
payload = {
"model": model.state_dict(),
"config": cfg,
"step": step,
}
if val_loss is not None:
payload["val_loss"] = val_loss
torch.save(payload, path)
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--config", default="configs/yuspec_gamedev_10m.yaml")
parser.add_argument("--max-iters", type=int)
parser.add_argument("--eval-iters", type=int)
parser.add_argument("--batch-size", type=int)
parser.add_argument("--out-dir")
parser.add_argument("--init-from", help="Optional checkpoint to initialize model weights from")
args = parser.parse_args()
cfg = load_config(args.config)
model_cfg = GPTConfig(**cfg["model"])
train_cfg = cfg["training"]
if args.max_iters is not None:
train_cfg["max_iters"] = args.max_iters
if args.eval_iters is not None:
train_cfg["eval_iters"] = args.eval_iters
if args.batch_size is not None:
train_cfg["batch_size"] = args.batch_size
if args.out_dir is not None:
train_cfg["out_dir"] = args.out_dir
train_data = load_tokens(cfg["data"]["train_bin"])
val_data = load_tokens(cfg["data"]["val_bin"])
model = GPT(model_cfg).to(device)
if args.init_from:
ckpt = torch.load(args.init_from, map_location=device)
model.load_state_dict(ckpt["model"])
print(f"initialized from {args.init_from}")
optimizer = AdamW(
model.parameters(),
lr=train_cfg["learning_rate"],
weight_decay=train_cfg["weight_decay"],
betas=(0.9, 0.95),
)
batch_size = train_cfg["batch_size"]
grad_accum = train_cfg["gradient_accumulation_steps"]
block_size = model_cfg.block_size
eval_iters = train_cfg.get("eval_iters", 50)
out_dir = train_cfg.get("out_dir", "checkpoints")
use_amp = bool(train_cfg.get("use_amp", device == "cuda")) and device == "cuda"
scaler = torch.amp.GradScaler("cuda", enabled=use_amp)
best_val_loss = float("inf")
t0 = time.time()
for step in range(train_cfg["max_iters"]):
lr = get_lr(
step,
train_cfg["learning_rate"],
train_cfg["min_lr"],
train_cfg["warmup_iters"],
train_cfg["max_iters"],
)
for param_group in optimizer.param_groups:
param_group["lr"] = lr
optimizer.zero_grad(set_to_none=True)
total_loss = 0.0
for _ in range(grad_accum):
x, y = get_batch(train_data, block_size, batch_size, device)
with amp_context(use_amp):
_, loss = model(x, y)
loss = loss / grad_accum
scaler.scale(loss).backward()
total_loss += loss.item()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(optimizer)
scaler.update()
if step % 100 == 0:
dt = time.time() - t0
print(f"step {step} | loss {total_loss:.4f} | lr {lr:.6f} | {dt:.1f}s")
t0 = time.time()
should_eval = step > 0 and step % train_cfg["eval_interval"] == 0
is_last_step = step == train_cfg["max_iters"] - 1
if should_eval or is_last_step:
losses = estimate_loss(model, train_data, val_data, block_size, batch_size, eval_iters, use_amp)
print(f"eval step {step} | train {losses['train']:.4f} | val {losses['val']:.4f}")
if losses["val"] < best_val_loss:
best_val_loss = losses["val"]
save_checkpoint(os.path.join(out_dir, "best.pt"), model, cfg, step, best_val_loss)
print(f"saved {os.path.join(out_dir, 'best.pt')}")
if step > 0 and step % train_cfg["save_interval"] == 0:
save_checkpoint(os.path.join(out_dir, f"step_{step}.pt"), model, cfg, step)
if __name__ == "__main__":
main()