HRM_sudoku / train_and_benchmark.py
Code2aum's picture
Upload folder using huggingface_hub
5dc80b3 verified
Raw
History Blame Contribute Delete
8.85 kB
#!/usr/bin/env python3
"""
train_and_benchmark.py β€” Train Baseline vs Tiered HRM, then benchmark both.
Usage:
source venv/bin/activate
# Quick smoke test (tiny dataset, 2 epochs)
python train_and_benchmark.py --data-path data/sudoku-1k --epochs 2 --batch-size 384
# Full run
python train_and_benchmark.py --data-path data/sudoku-1k --epochs 1000 --batch-size 384
# Benchmark only (skip training)
python train_and_benchmark.py --benchmark-only
Outputs:
- Console: live training metrics + benchmark table
- benchmark_results/results.json
- benchmark_results/benchmark_comparison.png
- benchmark_results/memory_analysis.png
"""
import argparse
import json
import os
import sys
import time
import torch
# ── Project imports ──────────────────────────────────────────────────
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from omegaconf import OmegaConf
from pretrain import PretrainConfig, init_train_state, train_batch, evaluate, create_dataloader, save_train_state
from benchmark import (
benchmark_tiered_model,
benchmark_baseline_model,
compare_models,
print_results_table,
generate_plots,
)
# =====================================================================
# PART 1: TRAINING
# =====================================================================
def run_training(arch_name: str, args, device):
"""Train one architecture variant and return final metrics."""
print(f"\n{'='*60}")
print(f" Training: {arch_name}")
print(f" Data: {args.data_path}")
print(f" Epochs: {args.epochs} | Batch: {args.batch_size}")
print(f"{'='*60}\n")
# Build Hydra-style config dict
cfg = OmegaConf.create({
"arch": arch_name,
"data_path": args.data_path,
"global_batch_size": args.batch_size,
"epochs": args.epochs,
"lr": 7e-5,
"lr_min_ratio": 0.0,
"lr_warmup_steps": 100,
"weight_decay": 1.0,
"beta1": 0.9,
"beta2": 0.95,
"puzzle_emb_lr": 7e-5,
"puzzle_emb_weight_decay": 1.0,
"seed": 0,
"skip_eval": False,
"eval_interval": max(1, args.epochs // 5), # eval 5 times
})
# Merge with the arch yaml
from hydra import compose, initialize_config_dir
from hydra.core.global_hydra import GlobalHydra
GlobalHydra.instance().clear()
with initialize_config_dir(config_dir=os.path.join(os.path.abspath("."), "config"), version_base=None):
hydra_cfg = compose(config_name="cfg_pretrain", overrides=[
f"arch={arch_name}",
f"data_path={args.data_path}",
f"global_batch_size={args.batch_size}",
f"epochs={args.epochs}",
])
config = PretrainConfig(**OmegaConf.to_container(hydra_cfg, resolve=True))
config.eval_interval = max(1, args.epochs // 5)
# Dataloaders
train_loader, train_meta = create_dataloader(
config, "train", 0, 1,
test_set_mode=False, epochs_per_iter=1,
global_batch_size=config.global_batch_size,
)
eval_loader, eval_meta = create_dataloader(
config, "test", 0, 1,
test_set_mode=True, epochs_per_iter=1,
global_batch_size=config.global_batch_size,
)
# Model
train_state = init_train_state(config, train_meta, world_size=1)
param_count = sum(p.numel() for p in train_state.model.parameters())
print(f" Parameters: {param_count:,} ({param_count/1e6:.1f}M)")
# Training loop
best_acc = 0.0
epoch_times = []
for epoch in range(1, config.epochs + 1):
train_state.model.train()
t0 = time.perf_counter()
last_metrics = None
for set_name, batch, gbs in train_loader:
metrics = train_batch(config, train_state, batch, gbs, rank=0, world_size=1)
if metrics:
last_metrics = metrics
dt = time.perf_counter() - t0
epoch_times.append(dt)
# Log
if last_metrics and epoch % max(1, args.epochs // 20) == 0:
loss = last_metrics.get("train/total_loss", 0)
acc = last_metrics.get("train/exact_accuracy", 0)
print(f" Epoch {epoch:>5}/{config.epochs} | Loss: {loss:.4f} | Acc: {acc:.2%} | {dt:.1f}s")
# Eval
if config.eval_interval and epoch % config.eval_interval == 0:
train_state.model.eval()
eval_results = evaluate(config, train_state, eval_loader, eval_meta, rank=0, world_size=1)
if eval_results:
for s_name, s_metrics in eval_results.items():
ea = s_metrics.get("exact_accuracy", 0)
print(f" eval/{s_name}: exact_accuracy={ea:.2%}")
best_acc = max(best_acc, ea)
# Summary
avg_epoch = sum(epoch_times) / len(epoch_times) if epoch_times else 0
print(f"\n ── {arch_name} Done ──")
print(f" Best Eval Accuracy: {best_acc:.2%}")
print(f" Avg Epoch Time: {avg_epoch:.2f}s")
print(f" GPU Peak Memory: {torch.cuda.max_memory_allocated()/1e9:.2f} GB")
return {"arch": arch_name, "best_acc": best_acc, "avg_epoch_s": avg_epoch}
# =====================================================================
# PART 2: HARDWARE BENCHMARK
# =====================================================================
def run_benchmark(args):
"""Compare Tiered vs Baseline on synthetic data."""
print(f"\n{'='*60}")
print(f" Hardware Benchmark: Tiered vs Baseline")
print(f" Batch sizes: {args.bench_batch_sizes}")
print(f" Seq lengths: {args.bench_seq_lens}")
print(f"{'='*60}")
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
bs_list = [int(x) for x in args.bench_batch_sizes.split(",")]
sl_list = [int(x) for x in args.bench_seq_lens.split(",")]
results = compare_models(
batch_sizes=bs_list, seq_lens=sl_list,
hidden_size=512, warmup=5, iterations=args.bench_iters,
device=device,
)
# Save
os.makedirs(args.output_dir, exist_ok=True)
out_path = os.path.join(args.output_dir, "results.json")
with open(out_path, "w") as f:
json.dump(results, f, indent=2, default=str)
print(f"\n Results saved β†’ {out_path}")
# Plots
try:
generate_plots(results, output_dir=args.output_dir)
except Exception as e:
print(f" (Plots skipped: {e})")
return results
# =====================================================================
# MAIN
# =====================================================================
def main():
parser = argparse.ArgumentParser(description="HRM: Train + Benchmark")
# Training
parser.add_argument("--data-path", type=str, default="data/sudoku-1k")
parser.add_argument("--epochs", type=int, default=100)
parser.add_argument("--batch-size", type=int, default=384)
# Benchmark
parser.add_argument("--bench-batch-sizes", type=str, default="1,8,32")
parser.add_argument("--bench-seq-lens", type=str, default="64,128")
parser.add_argument("--bench-iters", type=int, default=20)
parser.add_argument("--output-dir", type=str, default="benchmark_results")
# Flow control
parser.add_argument("--benchmark-only", action="store_true", help="Skip training, only run benchmark")
parser.add_argument("--train-only", action="store_true", help="Skip benchmark, only run training")
args = parser.parse_args()
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"\n Device: {device}")
if torch.cuda.is_available():
print(f" GPU: {torch.cuda.get_device_name(0)}")
# ── Training ──
train_summary = {}
if not args.benchmark_only:
for arch in ["hrm_v1", "hrm_tiered"]:
try:
result = run_training(arch, args, device)
train_summary[arch] = result
except Exception as e:
print(f"\n ⚠ Training failed for {arch}: {e}")
train_summary[arch] = {"error": str(e)}
torch.cuda.empty_cache()
# Print training comparison
print(f"\n{'='*60}")
print(f" Training Comparison")
print(f"{'='*60}")
print(f" {'Arch':<15} {'Best Acc':>10} {'Avg Epoch':>12}")
print(f" {'-'*37}")
for arch, r in train_summary.items():
if "error" in r:
print(f" {arch:<15} {'FAILED':>10} {'':>12}")
else:
print(f" {arch:<15} {r['best_acc']:>10.2%} {r['avg_epoch_s']:>10.2f}s")
# ── Benchmark ──
if not args.train_only:
run_benchmark(args)
print(f"\n All done!\n")
if __name__ == "__main__":
main()