#!/usr/bin/env python3 """ latency_plot_trained_model.py — Load trained checkpoints, evaluate accuracy on Sudoku test data, measure inference latency, and generate combined plots. Usage: source venv/bin/activate python latency_plot_trained_model.py \ --baseline "checkpoints/Sudoku-extreme-1k-aug-1000 ACT-torch/HierarchicalReasoningModel_ACTV1 belligerent-squirrel/step_52080" \ --tiered "checkpoints/Sudoku-extreme-1k-aug-1000 ACT-torch/HRM_Tiered realistic-dalmatian/step_52080" """ import argparse import json import os import sys import yaml # Disable torch.compile — avoids 10+ min compilation during eval # and prevents inference_mode/compile conflicts os.environ["DISABLE_COMPILE"] = "1" import torch import numpy as np import matplotlib matplotlib.use('Agg') import matplotlib.pyplot as plt from matplotlib.gridspec import GridSpec sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) from pretrain import PretrainConfig, init_train_state, evaluate, create_dataloader # ═══════════════════════════════════════════════════════════ # Load a trained checkpoint # ═══════════════════════════════════════════════════════════ def load_trained_model(ckpt_path, device="cuda"): """Load checkpoint, return (train_state, config, eval_loader, latency_loader, eval_metadata). Model hierarchy: torch.compile → ACTLossHead → ACTV1/HRM_Tiered → _Inner Returns TWO eval loaders: one for accuracy (consumed by evaluate()), one for latency. """ ckpt_dir = os.path.dirname(ckpt_path) config_path = os.path.join(ckpt_dir, "all_config.yaml") with open(config_path, "r") as f: content = f.read() if "!!python/object" not in content: raw = yaml.safe_load(content) else: # Fallback for the irreparably mangled tiered config dump print(" [Warning] Tiered config YAML is mangled, using robust fallback.") raw = { "arch": { "name": "hrm.hrm_tiered@HRM_Tiered", "hidden_size": 512, "num_heads": 8, "puzzle_emb_ndim": 512, "pos_encodings": "rope", "H_layers": 4, "H_cycles": 2, "L_layers": 4, "L_cycles": 2, "expansion": 4, "halt_max_steps": 16, "halt_exploration_prob": 0.1, "memory_tier": {"sram_capacity_mb": 48, "enable_tracking": True}, "loss": {"loss_type": "stablemax_cross_entropy", "name": "losses@ACTLossHead"} }, "global_batch_size": 384, "skip_eval": False, "eval_save_outputs": [], "checkpoint_path": ckpt_dir, "epochs": 20000, "lr": 7.0e-05, "lr_min_ratio": 1.0, "lr_warmup_steps": 2000, "weight_decay": 1.0, "beta1": 0.9, "beta2": 0.95, "puzzle_emb_lr": 7.0e-05, "puzzle_emb_weight_decay": 1.0, "eval_interval": 2000, "data_path": "data/sudoku-extreme-1k-aug-1000", "project_name": "Sudoku-extreme-1k-aug-1000 ACT-torch", "run_name": "HRM_Tiered realistic-dalmatian", "checkpoint_every_eval": True } config = PretrainConfig(**raw) config.checkpoint_path = ckpt_dir # Build dataloaders — need TWO because evaluate() consumes its loader _, train_metadata = create_dataloader( config, "train", test_set_mode=False, epochs_per_iter=1, global_batch_size=config.global_batch_size, rank=0, world_size=1, ) eval_loader, eval_metadata = create_dataloader( config, "test", test_set_mode=True, epochs_per_iter=1, global_batch_size=config.global_batch_size, rank=0, world_size=1, ) latency_loader, _ = create_dataloader( config, "test", test_set_mode=True, epochs_per_iter=1, global_batch_size=config.global_batch_size, rank=0, world_size=1, ) # Build model (torch.compile → ACTLossHead → model) and load weights train_state = init_train_state(config, train_metadata, world_size=1) try: train_state.model.load_state_dict( torch.load(ckpt_path, map_location=device, weights_only=True), assign=True ) except Exception: state = torch.load(ckpt_path, map_location=device, weights_only=True) train_state.model.load_state_dict( {k.removeprefix("_orig_mod."): v for k, v in state.items()}, assign=True ) ckpt_name = os.path.basename(ckpt_path) if ckpt_name.startswith("step_"): train_state.step = int(ckpt_name.removeprefix("step_")) train_state.model.eval() return train_state, config, eval_loader, latency_loader, eval_metadata def unwrap_model(compiled_model): """Unwrap torch.compile + ACTLossHead to get the ACTV1/HRM_Tiered wrapper. Hierarchy: OptimizedModule._orig_mod = ACTLossHead.model = ACTV1/HRM_Tiered """ model = compiled_model # Unwrap torch.compile if hasattr(model, '_orig_mod'): model = model._orig_mod # Unwrap ACTLossHead to get to the ACTV1/HRM_Tiered wrapper if hasattr(model, 'model'): model = model.model return model # ═══════════════════════════════════════════════════════════ # Evaluate accuracy on real Sudoku test set # ═══════════════════════════════════════════════════════════ def eval_accuracy(config, train_state, eval_loader, eval_metadata, limit_batches=20): """Run the real evaluation on a subset of batches and return metrics dict.""" import itertools class LimitedLoader: def __init__(self, loader, limit): self.loader = loader self.limit = limit def __iter__(self): return itertools.islice(self.loader, self.limit) limited_eval_loader = LimitedLoader(eval_loader, limit_batches) metrics = evaluate(config, train_state, limited_eval_loader, eval_metadata, rank=0, world_size=1) if metrics is None: return {} # Flatten and convert to floats (skip nested dicts / non-numeric) result = {} for k, v in metrics.items(): if isinstance(v, torch.Tensor): result[k] = v.item() elif isinstance(v, (int, float)): result[k] = float(v) elif isinstance(v, dict): for kk, vv in v.items(): if isinstance(vv, torch.Tensor): result[f"{k}/{kk}"] = vv.item() elif isinstance(vv, (int, float)): result[f"{k}/{kk}"] = float(vv) return result # ═══════════════════════════════════════════════════════════ # Measure inference latency on real data # ═══════════════════════════════════════════════════════════ @torch.no_grad() def measure_latency(compiled_model, eval_loader, device, warmup=3, iterations=20): """Time forward pass using the unwrapped model wrapper (ACTV1/HRM_Tiered). The unwrapped model has: - initial_carry(batch) → carry - forward(carry, batch) → (new_carry, outputs) """ # Unwrap torch.compile + ACTLossHead model = unwrap_model(compiled_model) model.eval() # Collect batches from the eval loader batches = [] for set_name, batch, global_bs in eval_loader: batch = {k: v.to(device) if isinstance(v, torch.Tensor) else v for k, v in batch.items()} batches.append(batch) if len(batches) >= warmup + iterations: break if not batches: return {"latency_ms": 0, "latency_std": 0, "throughput": 0} # Helper: create carry and move all tensors to device def make_carry(batch): carry = model.initial_carry(batch) carry.inner_carry.z_H = carry.inner_carry.z_H.to(device) carry.inner_carry.z_L = carry.inner_carry.z_L.to(device) carry.steps = carry.steps.to(device) carry.halted = carry.halted.to(device) carry.current_data = {k: v.to(device) for k, v in carry.current_data.items()} return carry # Warmup for i in range(min(warmup, len(batches))): batch = batches[i] carry = make_carry(batch) model(carry, batch) torch.cuda.synchronize() # Timed runs latencies = [] bs_total = 0 n_iters = min(iterations, max(1, len(batches) - warmup)) for i in range(n_iters): batch = batches[(warmup + i) % len(batches)] carry = make_carry(batch) start = torch.cuda.Event(enable_timing=True) end = torch.cuda.Event(enable_timing=True) start.record() model(carry, batch) end.record() torch.cuda.synchronize() latencies.append(start.elapsed_time(end)) bs_total += batch["inputs"].shape[0] lat = np.array(latencies) avg_bs = bs_total / len(latencies) if latencies else 1 return { "latency_ms": float(np.mean(lat)), "latency_std": float(np.std(lat)), "throughput": float(avg_bs / (np.mean(lat) / 1000)) if lat.mean() > 0 else 0, } # ═══════════════════════════════════════════════════════════ # Generate combined plots # ═══════════════════════════════════════════════════════════ def create_combined_plot(base_data, tier_data, output_dir): os.makedirs(output_dir, exist_ok=True) c_base, c_tier = "#4A90D9", "#E85D75" bg, text, grid = "#1a1a2e", "#e0e0e0", "#333355" plt.rcParams.update({ "figure.facecolor": bg, "axes.facecolor": "#16213e", "axes.edgecolor": grid, "axes.labelcolor": text, "text.color": text, "xtick.color": text, "ytick.color": text, "grid.color": grid, "grid.alpha": 0.3, "font.family": "sans-serif", "font.size": 11, }) fig = plt.figure(figsize=(18, 10)) fig.suptitle("HRM Trained Model Comparison: Baseline vs Tiered", fontsize=18, fontweight="bold", y=0.98) gs = GridSpec(2, 3, figure=fig, hspace=0.35, wspace=0.35) labels = ["Baseline", "Tiered"] def bar_ax(ax, title, ylabel, vals, fmt=".2f"): bars = ax.bar(labels, vals, color=[c_base, c_tier], edgecolor="white", linewidth=0.5, width=0.5) ax.set_title(title, fontweight="bold") ax.set_ylabel(ylabel) for b, v in zip(bars, vals): ax.text(b.get_x() + b.get_width()/2, b.get_height() * 1.02, f"{v:{fmt}}", ha="center", fontsize=11, color=text) ax.grid(axis="y") # Extract metrics with safe defaults def get_acc(data, key, normalize=True): count = data["accuracy"].get("eval/count", 1) val = data["accuracy"].get(key, 0) if normalize and count > 0: return val / count * 100 return val # 1. Exact Accuracy bar_ax(fig.add_subplot(gs[0, 0]), "Exact Accuracy (Sudoku)", "%", [get_acc(base_data, "eval/exact_accuracy"), get_acc(tier_data, "eval/exact_accuracy")]) # 2. Cell Accuracy bar_ax(fig.add_subplot(gs[0, 1]), "Cell-level Accuracy", "%", [get_acc(base_data, "eval/accuracy"), get_acc(tier_data, "eval/accuracy")]) # 3. Avg Reasoning Steps bar_ax(fig.add_subplot(gs[0, 2]), "Avg Reasoning Steps (ACT)", "steps", [get_acc(base_data, "eval/steps"), get_acc(tier_data, "eval/steps")], fmt=".1f") # 4. Inference Latency bar_ax(fig.add_subplot(gs[1, 0]), "Inference Latency", "ms", [base_data["latency"]["latency_ms"], tier_data["latency"]["latency_ms"]]) # 5. Throughput bar_ax(fig.add_subplot(gs[1, 1]), "Throughput", "samples/sec", [base_data["latency"]["throughput"], tier_data["latency"]["throughput"]], fmt=".0f") # 6. Summary ax6 = fig.add_subplot(gs[1, 2]) speedup = (base_data["latency"]["latency_ms"] / tier_data["latency"]["latency_ms"] if tier_data["latency"]["latency_ms"] > 0 else 0) base_exact = get_acc(base_data, "eval/exact_accuracy") tier_exact = get_acc(tier_data, "eval/exact_accuracy") summary = ( f"Exact Accuracy:\n" f" Baseline: {base_exact:.1f}%\n" f" Tiered: {tier_exact:.1f}%\n\n" f"Speedup: {speedup:.2f}x\n" f"Throughput:\n" f" {tier_data['latency']['throughput']:.0f} vs " f"{base_data['latency']['throughput']:.0f}/s" ) ax6.text(0.5, 0.5, summary, transform=ax6.transAxes, ha="center", va="center", fontsize=13, fontfamily="monospace", bbox=dict(boxstyle="round,pad=0.5", facecolor="#0f3460", alpha=0.8)) ax6.set_title("Summary", fontweight="bold") ax6.axis("off") path = os.path.join(output_dir, "trained_model_comparison.png") fig.savefig(path, dpi=150, bbox_inches="tight") plt.close() print(f" Plot saved → {path}") return path # ═══════════════════════════════════════════════════════════ # Main # ═══════════════════════════════════════════════════════════ def main(): parser = argparse.ArgumentParser(description="Evaluate trained Baseline vs Tiered HRM") parser.add_argument("--baseline", type=str, required=True, help="Baseline checkpoint path") parser.add_argument("--tiered", type=str, required=True, help="Tiered checkpoint path") parser.add_argument("--latency-iters", type=int, default=20) parser.add_argument("--output-dir", type=str, default="benchmark_results") args = parser.parse_args() device = "cuda" print("=" * 64) print(" Trained Model Comparison: Baseline vs Tiered") print(f" Device: {torch.cuda.get_device_name(0)}") print("=" * 64) results = {} # ── Baseline ── print("\n [1/4] Loading Baseline checkpoint...") base_state, base_cfg, base_eval_loader, base_lat_loader, base_eval_meta = load_trained_model(args.baseline, device) n_params = sum(p.numel() for p in base_state.model.parameters()) / 1e6 print(f" Step: {base_state.step}, Params: {n_params:.1f}M") print(" [2/4] Evaluating Baseline accuracy + latency...") base_acc = eval_accuracy(base_cfg, base_state, base_eval_loader, base_eval_meta) print(f" Accuracy metrics: {base_acc}") base_lat = measure_latency(base_state.model, base_lat_loader, device, iterations=args.latency_iters) print(f" Latency: {base_lat['latency_ms']:.2f} ms ± {base_lat['latency_std']:.2f}") results["baseline"] = {"accuracy": base_acc, "latency": base_lat} # Free memory del base_state, base_eval_loader, base_lat_loader torch.cuda.empty_cache() # ── Tiered ── print("\n [3/4] Loading Tiered checkpoint...") tier_state, tier_cfg, tier_eval_loader, tier_lat_loader, tier_eval_meta = load_trained_model(args.tiered, device) n_params = sum(p.numel() for p in tier_state.model.parameters()) / 1e6 print(f" Step: {tier_state.step}, Params: {n_params:.1f}M") print(" [4/4] Evaluating Tiered accuracy + latency...") tier_acc = eval_accuracy(tier_cfg, tier_state, tier_eval_loader, tier_eval_meta) print(f" Accuracy metrics: {tier_acc}") tier_lat = measure_latency(tier_state.model, tier_lat_loader, device, iterations=args.latency_iters) print(f" Latency: {tier_lat['latency_ms']:.2f} ms ± {tier_lat['latency_std']:.2f}") results["tiered"] = {"accuracy": tier_acc, "latency": tier_lat} del tier_state, tier_eval_loader, tier_lat_loader torch.cuda.empty_cache() # ── Plots ── print("\n Generating comparison plots...") create_combined_plot(results["baseline"], results["tiered"], args.output_dir) # ── Save JSON ── json_path = os.path.join(args.output_dir, "trained_model_results.json") os.makedirs(args.output_dir, exist_ok=True) with open(json_path, "w") as f: json.dump(results, f, indent=2, default=str) print(f" Results saved → {json_path}") print("\n" + "=" * 64) print(" Done!") print("=" * 64) if __name__ == "__main__": main()