Download latency_plot_trained_model.py from Code2aum/HRM_sudoku: direct link, hf CLI and curl.
- Browser
- Download file 17.3 kB
-
https://huggingface.co/Code2aum/HRM_sudoku/resolve/main/latency_plot_trained_model.py
- Command line
-
hf download hf://Code2aum/HRM_sudoku/latency_plot_trained_model.py
-
curl -L -o latency_plot_trained_model.py https://huggingface.co/Code2aum/HRM_sudoku/resolve/main/latency_plot_trained_model.py
17.3 kB
| #!/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 | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| 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() | |