HRM_sudoku / compare_models.py
Code2aum's picture
Upload folder using huggingface_hub
5dc80b3 verified
Raw
History Blame Contribute Delete
12.3 kB
#!/usr/bin/env python3
"""
Compare Baseline vs Tiered HRM β€” Multi-GPU Benchmark + Comparison Plots.
Usage:
source venv/bin/activate
python compare_models.py # quick compare
python compare_models.py --sweep # batch-size sweep + plots
python compare_models.py --iterations 50 --sweep # thorough
"""
import argparse
import json
import os
import sys
from dataclasses import dataclass, asdict
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 models.memory_tier import MemoryTierManager
from models.hrm.hrm_tiered import HRM_Tiered
from models.hrm.hrm_act_v1 import HierarchicalReasoningModel_ACTV1
# ═══════════════════════════════════════════════════════════
# Helpers
# ═══════════════════════════════════════════════════════════
def make_config(batch_size, seq_len, hidden_size, num_heads):
return {
"batch_size": batch_size, "seq_len": seq_len,
"puzzle_emb_ndim": 0, "num_puzzle_identifiers": batch_size,
"vocab_size": 32,
"H_cycles": 2, "L_cycles": 2, "H_layers": 4, "L_layers": 4,
"hidden_size": hidden_size, "expansion": 4.0,
"num_heads": num_heads, "pos_encodings": "rope",
"halt_max_steps": 1, "halt_exploration_prob": 0.0,
}
def make_batch(batch_size, seq_len, device):
return {
"inputs": torch.randint(0, 31, (batch_size, seq_len), device=device),
"labels": torch.randint(0, 31, (batch_size, seq_len), device=device),
"puzzle_identifiers": torch.arange(batch_size, device=device),
}
def build_model(arch, config_dict, device):
"""Build using the wrapper classes (same pattern as eval_dummy.py)."""
if arch == "tiered":
mm = MemoryTierManager(device=device, enable_tracking=True)
model = HRM_Tiered(config_dict, memory_manager=mm).to(device)
else:
model = HierarchicalReasoningModel_ACTV1(config_dict).to(device)
model.eval()
return model
@torch.no_grad()
def benchmark(model, batch, device, warmup=5, iterations=20):
"""Time forward pass, return dict of metrics."""
bs = batch["inputs"].shape[0]
# Warmup
for _ in range(warmup):
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()}
model(carry, batch)
torch.cuda.reset_peak_memory_stats(device)
torch.cuda.synchronize()
latencies = []
for _ in range(iterations):
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()}
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))
lat = np.array(latencies)
return {
"latency_ms": float(np.mean(lat)),
"latency_std": float(np.std(lat)),
"throughput": float(bs / (np.mean(lat) / 1000)),
"peak_gpu_mb": float(torch.cuda.max_memory_allocated(device) / 1e6),
"params_m": sum(p.numel() for p in model.parameters()) / 1e6,
}
# ═══════════════════════════════════════════════════════════
# Plotting
# ═══════════════════════════════════════════════════════════
def create_plots(base_res, tier_res, sweep_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,
})
n_plots = 6 if sweep_data else 5
fig = plt.figure(figsize=(16, 10))
fig.suptitle("HRM Baseline vs Tiered (SRAM/DRAM) Comparison",
fontsize=18, fontweight="bold", y=0.98)
gs = GridSpec(2, 3, figure=fig, hspace=0.35, wspace=0.35)
labels = ["Baseline", "Tiered"]
def bar_plot(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=10, color=text)
ax.grid(axis="y")
# 1. Latency
bar_plot(fig.add_subplot(gs[0, 0]), "Inference Latency", "ms",
[base_res["latency_ms"], tier_res["latency_ms"]])
# 2. Throughput
bar_plot(fig.add_subplot(gs[0, 1]), "Throughput", "samples/sec",
[base_res["throughput"], tier_res["throughput"]], fmt=".0f")
# 3. GPU Memory
bar_plot(fig.add_subplot(gs[0, 2]), "Peak GPU Memory", "MB",
[base_res["peak_gpu_mb"], tier_res["peak_gpu_mb"]], fmt=".0f")
# 4. Parameters
bar_plot(fig.add_subplot(gs[1, 0]), "Model Parameters", "Millions",
[base_res["params_m"], tier_res["params_m"]], fmt=".1f")
# 5. Summary text
ax5 = fig.add_subplot(gs[1, 1])
speedup = base_res["latency_ms"] / tier_res["latency_ms"]
mem_diff = tier_res["peak_gpu_mb"] - base_res["peak_gpu_mb"]
tp_gain = (tier_res["throughput"] / base_res["throughput"] - 1) * 100
summary = (
f"Speedup: {speedup:.2f}x\n"
f"Throughput: {tp_gain:+.1f}%\n"
f"Memory Ξ”: {mem_diff:+.0f} MB\n"
f"Params: identical"
)
ax5.text(0.5, 0.5, summary, transform=ax5.transAxes,
ha="center", va="center", fontsize=14, fontfamily="monospace",
bbox=dict(boxstyle="round,pad=0.5", facecolor="#0f3460", alpha=0.8))
ax5.set_title("Summary", fontweight="bold")
ax5.axis("off")
# 6. Sweep plot
ax6 = fig.add_subplot(gs[1, 2])
if sweep_data:
bs_list = [s["batch_size"] for s in sweep_data["baseline"]]
ax6.plot(bs_list, [s["latency_ms"] for s in sweep_data["baseline"]],
"o-", color=c_base, label="Baseline", linewidth=2, markersize=6)
ax6.plot(bs_list, [s["latency_ms"] for s in sweep_data["tiered"]],
"s-", color=c_tier, label="Tiered", linewidth=2, markersize=6)
ax6.set_xlabel("Batch Size")
ax6.set_ylabel("Latency (ms)")
ax6.set_title("Latency vs Batch Size", fontweight="bold")
ax6.legend(facecolor="#16213e", edgecolor=grid)
ax6.grid(True)
else:
ax6.text(0.5, 0.5, "Run with --sweep\nfor batch size\ncomparison",
transform=ax6.transAxes, ha="center", va="center", fontsize=12)
ax6.set_title("Latency vs Batch Size", fontweight="bold")
ax6.axis("off")
path = os.path.join(output_dir, "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="Compare Baseline vs Tiered HRM")
parser.add_argument("--batch-size", type=int, default=32)
parser.add_argument("--seq-len", type=int, default=81)
parser.add_argument("--hidden-size", type=int, default=512)
parser.add_argument("--num-heads", type=int, default=8)
parser.add_argument("--warmup", type=int, default=5)
parser.add_argument("--iterations", type=int, default=20)
parser.add_argument("--sweep", action="store_true", help="Batch-size sweep")
parser.add_argument("--output-dir", type=str, default="benchmark_results")
args = parser.parse_args()
device = torch.device("cuda")
cfg = make_config(args.batch_size, args.seq_len, args.hidden_size, args.num_heads)
print("=" * 64)
print(" HRM Model Comparison: Baseline vs Tiered (SRAM/DRAM)")
print(f" Device: {torch.cuda.get_device_name(0)}")
print(f" Config: bs={args.batch_size}, seq={args.seq_len}, hidden={args.hidden_size}")
print("=" * 64)
# ── Build ──
print("\n Building Baseline...")
base_model = build_model("baseline", cfg, device)
batch = make_batch(args.batch_size, args.seq_len, device)
print(" Building Tiered...")
tier_model = build_model("tiered", cfg, device)
# ── Benchmark ──
print(f"\n Benchmarking Baseline ({args.iterations} iters)...")
base_res = benchmark(base_model, batch, device, args.warmup, args.iterations)
print(f" β†’ {base_res['latency_ms']:.2f} ms | {base_res['throughput']:.0f} samp/s | {base_res['peak_gpu_mb']:.0f} MB")
batch_t = make_batch(args.batch_size, args.seq_len, device)
print(f" Benchmarking Tiered ({args.iterations} iters)...")
tier_res = benchmark(tier_model, batch_t, device, args.warmup, args.iterations)
print(f" β†’ {tier_res['latency_ms']:.2f} ms | {tier_res['throughput']:.0f} samp/s | {tier_res['peak_gpu_mb']:.0f} MB")
speedup = base_res["latency_ms"] / tier_res["latency_ms"]
print(f"\n Speedup: {speedup:.2f}x")
# ── Sweep ──
sweep_data = None
if args.sweep:
print("\n Running batch-size sweep...")
sweep_data = {"baseline": [], "tiered": []}
del base_model, tier_model
torch.cuda.empty_cache()
for bs in [1, 4, 8, 16, 32, 64]:
print(f" bs={bs}...", end=" ", flush=True)
c = make_config(bs, args.seq_len, args.hidden_size, args.num_heads)
b = make_batch(bs, args.seq_len, device)
bm = build_model("baseline", c, device)
br = benchmark(bm, b, device, warmup=3, iterations=10)
br["batch_size"] = bs
sweep_data["baseline"].append(br)
del bm
tm = build_model("tiered", c, device)
tr = benchmark(tm, b, device, warmup=3, iterations=10)
tr["batch_size"] = bs
sweep_data["tiered"].append(tr)
del tm
torch.cuda.empty_cache()
print(f"base={br['latency_ms']:.2f}ms, tier={tr['latency_ms']:.2f}ms")
# ── Plots ──
print("\n Generating plots...")
create_plots(base_res, tier_res, sweep_data, args.output_dir)
# ── Save JSON ──
results = {"baseline": base_res, "tiered": tier_res, "speedup": speedup}
if sweep_data:
results["sweep"] = sweep_data
json_path = os.path.join(args.output_dir, "comparison_results.json")
os.makedirs(args.output_dir, exist_ok=True)
with open(json_path, "w") as f:
json.dump(results, f, indent=2)
print(f" Results saved β†’ {json_path}")
print("\n" + "=" * 64)
print(" Done!")
print("=" * 64)
if __name__ == "__main__":
main()