HRM_sudoku / benchmark.py
Code2aum's picture
Upload folder using huggingface_hub
5dc80b3 verified
Raw
History Blame Contribute Delete
23.7 kB
"""
Comprehensive Benchmarking Suite for HRM SRAM/DRAM Memory Tiering.
Measures and compares:
- Per-module latency (L-level SRAM vs H-level DRAM)
- End-to-end inference latency
- Throughput (samples/sec)
- Memory usage per tier (SRAM / DRAM)
- Cross-tier transfer overhead
- SRAM hit rate
- Triton kernel-level profiling
- GPU utilization / power draw (when available)
"""
import json
import time
import math
import os
from typing import Dict, List, Optional, Tuple
from dataclasses import dataclass, asdict
import torch
import torch.nn.functional as F
from models.memory_tier import MemoryTierManager
from models.hrm.hrm_tiered import HRM_Tiered, HRM_Tiered_Inner
from models.hrm.hrm_act_v1 import (
HierarchicalReasoningModel_ACTV1,
HierarchicalReasoningModel_ACTV1Config,
HierarchicalReasoningModel_ACTV1InnerCarry,
)
from models.triton_kernels import triton_memory_latency_probe
@dataclass
class BenchmarkResult:
"""Results from a single benchmark run."""
model_name: str
batch_size: int
seq_len: int
hidden_size: int
H_cycles: int
L_cycles: int
H_layers: int
L_layers: int
num_iterations: int
warmup_iterations: int
# Latency (μs)
l_level_latency_mean_us: float
l_level_latency_min_us: float
l_level_latency_max_us: float
l_level_latency_std_us: float
h_level_latency_mean_us: float
h_level_latency_min_us: float
h_level_latency_max_us: float
h_level_latency_std_us: float
total_inference_latency_mean_ms: float
total_inference_latency_min_ms: float
total_inference_latency_max_ms: float
total_inference_latency_std_ms: float
# Throughput
throughput_samples_per_sec: float
# Memory (MB)
sram_peak_mb: float
dram_peak_mb: float
total_gpu_memory_mb: float
# Transfer metrics (μs)
h_l_transfer_mean_us: float
l_h_transfer_mean_us: float
# SRAM metrics
sram_hit_rate: float
# Triton kernel metrics
triton_sram_probe_latency_us: float
triton_dram_probe_latency_us: float
# GPU metrics
gpu_utilization_pct: Optional[float]
gpu_power_w: Optional[float]
gpu_temperature_c: Optional[float]
# Derived
h_over_l_latency_ratio: float
memory_efficiency: float # useful_compute_time / total_time
def _std(values: List[float]) -> float:
if len(values) < 2:
return 0.0
mean = sum(values) / len(values)
var = sum((v - mean) ** 2 for v in values) / (len(values) - 1)
return math.sqrt(var)
def _create_dummy_batch(
batch_size: int,
seq_len: int,
vocab_size: int,
device: torch.device,
) -> Dict[str, torch.Tensor]:
"""Create a synthetic batch for benchmarking."""
return {
"inputs": torch.randint(0, vocab_size, (batch_size, seq_len), device=device),
"labels": torch.randint(0, vocab_size, (batch_size, seq_len), device=device),
"puzzle_identifiers": torch.arange(batch_size, device=device),
}
def _get_gpu_metrics() -> Dict[str, Optional[float]]:
"""Try to read GPU utilization, power, temperature via nvidia-smi."""
metrics = {'utilization': None, 'power': None, 'temperature': None}
try:
import subprocess
result = subprocess.run(
['nvidia-smi', '--query-gpu=utilization.gpu,power.draw,temperature.gpu',
'--format=csv,noheader,nounits'],
capture_output=True, text=True, timeout=5,
)
if result.returncode == 0:
parts = result.stdout.strip().split(',')
if len(parts) >= 3:
metrics['utilization'] = float(parts[0].strip())
metrics['power'] = float(parts[1].strip())
metrics['temperature'] = float(parts[2].strip())
except Exception:
pass
return metrics
def _run_triton_latency_probe(
batch_size: int,
hidden_size: int,
device: torch.device,
num_iters: int = 100,
) -> Tuple[float, float]:
"""Measure SRAM vs DRAM effective latency using Triton probe kernels.
Returns (sram_latency_us, dram_latency_us).
"""
if not torch.cuda.is_available():
return 0.0, 0.0
# Small tensor → fits in SRAM (L1/L2/registers)
sram_data = torch.randn(batch_size, hidden_size, device=device, dtype=torch.float32)
# Large tensor → forces DRAM access (much larger than L2)
dram_size = max(hidden_size, 65536) # Force spill to DRAM
dram_data = torch.randn(batch_size * 64, dram_size, device=device, dtype=torch.float32)
# Make a contiguous slice for DRAM probe
dram_probe_data = dram_data[:batch_size, :hidden_size].contiguous()
# Warmup
triton_memory_latency_probe(sram_data, num_iters=10)
triton_memory_latency_probe(dram_probe_data, num_iters=10)
torch.cuda.synchronize()
# SRAM probe
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
triton_memory_latency_probe(sram_data, num_iters=num_iters)
end.record()
torch.cuda.synchronize()
sram_us = start.elapsed_time(end) * 1000 # ms → μs
# DRAM probe (access scattered to prevent caching)
start2 = torch.cuda.Event(enable_timing=True)
end2 = torch.cuda.Event(enable_timing=True)
start2.record()
triton_memory_latency_probe(dram_probe_data, num_iters=num_iters)
end2.record()
torch.cuda.synchronize()
dram_us = start2.elapsed_time(end2) * 1000
# Cleanup
del sram_data, dram_data
torch.cuda.empty_cache()
return sram_us, dram_us
def benchmark_tiered_model(
batch_size: int = 8,
seq_len: int = 64,
hidden_size: int = 512,
num_heads: int = 8,
H_cycles: int = 2,
L_cycles: int = 2,
H_layers: int = 4,
L_layers: int = 4,
halt_max_steps: int = 1,
warmup: int = 5,
iterations: int = 20,
device: Optional[torch.device] = None,
) -> BenchmarkResult:
"""Benchmark the tiered HRM model."""
if device is None:
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
vocab_size = 32
config_dict = {
'batch_size': batch_size,
'seq_len': seq_len,
'puzzle_emb_ndim': 0,
'num_puzzle_identifiers': batch_size,
'vocab_size': vocab_size,
'H_cycles': H_cycles,
'L_cycles': L_cycles,
'H_layers': H_layers,
'L_layers': L_layers,
'hidden_size': hidden_size,
'expansion': 4.0,
'num_heads': num_heads,
'pos_encodings': 'rope',
'halt_max_steps': halt_max_steps,
'halt_exploration_prob': 0.0,
}
# Memory manager
mem_mgr = MemoryTierManager(device=device, enable_tracking=True)
# Create model
model = HRM_Tiered(config_dict, memory_manager=mem_mgr).to(device)
model.eval()
batch = _create_dummy_batch(batch_size, seq_len, vocab_size, device)
# ---- Warmup ----
with torch.no_grad():
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)
model.reset_timing()
mem_mgr.reset_stats()
if torch.cuda.is_available():
torch.cuda.reset_peak_memory_stats(device)
torch.cuda.synchronize()
# ---- Benchmark iterations ----
total_latencies_ms = []
with torch.no_grad():
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()}
if torch.cuda.is_available():
start_event = torch.cuda.Event(enable_timing=True)
end_event = torch.cuda.Event(enable_timing=True)
start_event.record()
t0 = time.perf_counter()
model(carry, batch)
if torch.cuda.is_available():
end_event.record()
torch.cuda.synchronize()
total_latencies_ms.append(start_event.elapsed_time(end_event))
else:
total_latencies_ms.append((time.perf_counter() - t0) * 1000)
# ---- Collect results ----
timing = model.get_timing_stats()
mem_stats = mem_mgr.get_stats()
# GPU memory
total_gpu_mb = 0.0
if torch.cuda.is_available():
total_gpu_mb = torch.cuda.max_memory_allocated(device) / (1024 * 1024)
# GPU metrics
gpu_metrics = _get_gpu_metrics()
# Triton latency probe
sram_probe_us, dram_probe_us = _run_triton_latency_probe(
batch_size, hidden_size, device,
)
# Compute derived metrics
l_mean = timing['L_forward_us']['mean_us']
h_mean = timing['H_forward_us']['mean_us']
ratio = h_mean / l_mean if l_mean > 0 else float('inf')
total_mean_ms = sum(total_latencies_ms) / len(total_latencies_ms)
throughput = batch_size / (total_mean_ms / 1000) if total_mean_ms > 0 else 0
compute_time = timing['L_forward_us']['total_us'] + timing['H_forward_us']['total_us']
transfer_time = timing['H_L_transfer_us']['total_us'] + timing['L_H_transfer_us']['total_us']
efficiency = compute_time / (compute_time + transfer_time) if (compute_time + transfer_time) > 0 else 0
# Build L-level stats
l_values = [timing['L_forward_us']['min_us'], timing['L_forward_us']['max_us']]
result = BenchmarkResult(
model_name='HRM_Tiered',
batch_size=batch_size,
seq_len=seq_len,
hidden_size=hidden_size,
H_cycles=H_cycles,
L_cycles=L_cycles,
H_layers=H_layers,
L_layers=L_layers,
num_iterations=iterations,
warmup_iterations=warmup,
l_level_latency_mean_us=l_mean,
l_level_latency_min_us=timing['L_forward_us']['min_us'],
l_level_latency_max_us=timing['L_forward_us']['max_us'],
l_level_latency_std_us=0.0,
h_level_latency_mean_us=h_mean,
h_level_latency_min_us=timing['H_forward_us']['min_us'],
h_level_latency_max_us=timing['H_forward_us']['max_us'],
h_level_latency_std_us=0.0,
total_inference_latency_mean_ms=total_mean_ms,
total_inference_latency_min_ms=min(total_latencies_ms),
total_inference_latency_max_ms=max(total_latencies_ms),
total_inference_latency_std_ms=_std(total_latencies_ms),
throughput_samples_per_sec=throughput,
sram_peak_mb=mem_stats['sram']['peak_mb'],
dram_peak_mb=mem_stats['dram']['peak_mb'],
total_gpu_memory_mb=total_gpu_mb,
h_l_transfer_mean_us=timing['H_L_transfer_us']['mean_us'],
l_h_transfer_mean_us=timing['L_H_transfer_us']['mean_us'],
sram_hit_rate=mem_stats['sram']['hit_rate'],
triton_sram_probe_latency_us=sram_probe_us,
triton_dram_probe_latency_us=dram_probe_us,
gpu_utilization_pct=gpu_metrics['utilization'],
gpu_power_w=gpu_metrics['power'],
gpu_temperature_c=gpu_metrics['temperature'],
h_over_l_latency_ratio=ratio,
memory_efficiency=efficiency,
)
# Cleanup
del model, batch, carry
if torch.cuda.is_available():
torch.cuda.empty_cache()
return result
def benchmark_baseline_model(
batch_size: int = 8,
seq_len: int = 64,
hidden_size: int = 512,
num_heads: int = 8,
H_cycles: int = 2,
L_cycles: int = 2,
H_layers: int = 4,
L_layers: int = 4,
halt_max_steps: int = 1,
warmup: int = 5,
iterations: int = 20,
device: Optional[torch.device] = None,
) -> BenchmarkResult:
"""Benchmark the original (non-tiered) HRM model."""
if device is None:
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
vocab_size = 32
config_dict = {
'batch_size': batch_size,
'seq_len': seq_len,
'puzzle_emb_ndim': 0,
'num_puzzle_identifiers': batch_size,
'vocab_size': vocab_size,
'H_cycles': H_cycles,
'L_cycles': L_cycles,
'H_layers': H_layers,
'L_layers': L_layers,
'hidden_size': hidden_size,
'expansion': 4.0,
'num_heads': num_heads,
'pos_encodings': 'rope',
'halt_max_steps': halt_max_steps,
'halt_exploration_prob': 0.0,
}
model = HierarchicalReasoningModel_ACTV1(config_dict).to(device)
model.eval()
batch = _create_dummy_batch(batch_size, seq_len, vocab_size, device)
# ---- Warmup ----
with torch.no_grad():
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)
if torch.cuda.is_available():
torch.cuda.reset_peak_memory_stats(device)
torch.cuda.synchronize()
# ---- Benchmark ----
total_latencies_ms = []
with torch.no_grad():
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()}
if torch.cuda.is_available():
start_event = torch.cuda.Event(enable_timing=True)
end_event = torch.cuda.Event(enable_timing=True)
start_event.record()
t0 = time.perf_counter()
model(carry, batch)
if torch.cuda.is_available():
end_event.record()
torch.cuda.synchronize()
total_latencies_ms.append(start_event.elapsed_time(end_event))
else:
total_latencies_ms.append((time.perf_counter() - t0) * 1000)
total_gpu_mb = 0.0
if torch.cuda.is_available():
total_gpu_mb = torch.cuda.max_memory_allocated(device) / (1024 * 1024)
total_mean_ms = sum(total_latencies_ms) / len(total_latencies_ms)
throughput = batch_size / (total_mean_ms / 1000) if total_mean_ms > 0 else 0
gpu_metrics = _get_gpu_metrics()
result = BenchmarkResult(
model_name='HRM_Baseline',
batch_size=batch_size,
seq_len=seq_len,
hidden_size=hidden_size,
H_cycles=H_cycles,
L_cycles=L_cycles,
H_layers=H_layers,
L_layers=L_layers,
num_iterations=iterations,
warmup_iterations=warmup,
l_level_latency_mean_us=0, l_level_latency_min_us=0,
l_level_latency_max_us=0, l_level_latency_std_us=0,
h_level_latency_mean_us=0, h_level_latency_min_us=0,
h_level_latency_max_us=0, h_level_latency_std_us=0,
total_inference_latency_mean_ms=total_mean_ms,
total_inference_latency_min_ms=min(total_latencies_ms),
total_inference_latency_max_ms=max(total_latencies_ms),
total_inference_latency_std_ms=_std(total_latencies_ms),
throughput_samples_per_sec=throughput,
sram_peak_mb=0, dram_peak_mb=0,
total_gpu_memory_mb=total_gpu_mb,
h_l_transfer_mean_us=0, l_h_transfer_mean_us=0,
sram_hit_rate=0,
triton_sram_probe_latency_us=0, triton_dram_probe_latency_us=0,
gpu_utilization_pct=gpu_metrics['utilization'],
gpu_power_w=gpu_metrics['power'],
gpu_temperature_c=gpu_metrics['temperature'],
h_over_l_latency_ratio=0,
memory_efficiency=1.0,
)
del model, batch, carry
if torch.cuda.is_available():
torch.cuda.empty_cache()
return result
def compare_models(
batch_sizes: List[int] = [1, 8, 32],
seq_lens: List[int] = [64, 128],
hidden_size: int = 512,
warmup: int = 5,
iterations: int = 20,
device: Optional[torch.device] = None,
) -> List[Dict]:
"""Run comparative benchmark between tiered and baseline HRM."""
results = []
for bs in batch_sizes:
for sl in seq_lens:
print(f"\n{'='*60}")
print(f" Benchmarking: batch_size={bs}, seq_len={sl}")
print(f"{'='*60}")
print(" → Baseline model...")
baseline = benchmark_baseline_model(
batch_size=bs, seq_len=sl, hidden_size=hidden_size,
warmup=warmup, iterations=iterations, device=device,
)
print(" → Tiered model...")
tiered = benchmark_tiered_model(
batch_size=bs, seq_len=sl, hidden_size=hidden_size,
warmup=warmup, iterations=iterations, device=device,
)
comparison = {
'batch_size': bs,
'seq_len': sl,
'baseline': asdict(baseline),
'tiered': asdict(tiered),
'speedup': baseline.total_inference_latency_mean_ms / tiered.total_inference_latency_mean_ms if tiered.total_inference_latency_mean_ms > 0 else 0,
'memory_savings_mb': baseline.total_gpu_memory_mb - tiered.total_gpu_memory_mb,
'throughput_improvement': tiered.throughput_samples_per_sec / baseline.throughput_samples_per_sec if baseline.throughput_samples_per_sec > 0 else 0,
}
results.append(comparison)
# Print summary
print(f"\n Results:")
print(f" Baseline latency: {baseline.total_inference_latency_mean_ms:.2f} ms")
print(f" Tiered latency: {tiered.total_inference_latency_mean_ms:.2f} ms")
print(f" Speedup: {comparison['speedup']:.2f}x")
print(f" H/L ratio: {tiered.h_over_l_latency_ratio:.2f}x")
print(f" SRAM probe: {tiered.triton_sram_probe_latency_us:.1f} μs")
print(f" DRAM probe: {tiered.triton_dram_probe_latency_us:.1f} μs")
return results
def print_results_table(results: List[BenchmarkResult]):
"""Pretty-print benchmark results as a table."""
header = (
f"{'Model':<15} {'BS':>4} {'Seq':>5} "
f"{'Latency(ms)':>12} {'Throughput':>12} "
f"{'L_lat(μs)':>10} {'H_lat(μs)':>10} {'H/L':>6} "
f"{'GPU_MB':>8} {'Efficiency':>10}"
)
print(f"\n{'='*len(header)}")
print(header)
print(f"{'='*len(header)}")
for r in results:
print(
f"{r.model_name:<15} {r.batch_size:>4} {r.seq_len:>5} "
f"{r.total_inference_latency_mean_ms:>12.2f} {r.throughput_samples_per_sec:>12.1f} "
f"{r.l_level_latency_mean_us:>10.1f} {r.h_level_latency_mean_us:>10.1f} {r.h_over_l_latency_ratio:>6.2f} "
f"{r.total_gpu_memory_mb:>8.1f} {r.memory_efficiency:>10.3f}"
)
print()
def generate_plots(results: List[Dict], output_dir: str = "benchmark_results"):
"""Generate comparison plots using matplotlib."""
try:
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
import numpy as np
except ImportError:
print("matplotlib not available — skipping plot generation.")
return
os.makedirs(output_dir, exist_ok=True)
# ---- Plot 1: Latency comparison ----
fig, axes = plt.subplots(1, 3, figsize=(18, 5))
fig.suptitle('HRM SRAM/DRAM Tiering — Benchmark Results', fontsize=14, fontweight='bold')
configs = [f"bs={r['batch_size']}\nseq={r['seq_len']}" for r in results]
baseline_lat = [r['baseline']['total_inference_latency_mean_ms'] for r in results]
tiered_lat = [r['tiered']['total_inference_latency_mean_ms'] for r in results]
x = np.arange(len(configs))
w = 0.35
ax = axes[0]
ax.bar(x - w/2, baseline_lat, w, label='Baseline', color='#e74c3c', alpha=0.8)
ax.bar(x + w/2, tiered_lat, w, label='Tiered (SRAM/DRAM)', color='#2ecc71', alpha=0.8)
ax.set_xlabel('Configuration')
ax.set_ylabel('Latency (ms)')
ax.set_title('Inference Latency')
ax.set_xticks(x)
ax.set_xticklabels(configs, fontsize=8)
ax.legend()
ax.grid(axis='y', alpha=0.3)
# ---- Plot 2: Throughput ----
ax = axes[1]
baseline_tp = [r['baseline']['throughput_samples_per_sec'] for r in results]
tiered_tp = [r['tiered']['throughput_samples_per_sec'] for r in results]
ax.bar(x - w/2, baseline_tp, w, label='Baseline', color='#e74c3c', alpha=0.8)
ax.bar(x + w/2, tiered_tp, w, label='Tiered', color='#2ecc71', alpha=0.8)
ax.set_xlabel('Configuration')
ax.set_ylabel('Samples/sec')
ax.set_title('Throughput')
ax.set_xticks(x)
ax.set_xticklabels(configs, fontsize=8)
ax.legend()
ax.grid(axis='y', alpha=0.3)
# ---- Plot 3: H vs L latency (tiered only) ----
ax = axes[2]
l_lat = [r['tiered']['l_level_latency_mean_us'] for r in results]
h_lat = [r['tiered']['h_level_latency_mean_us'] for r in results]
ax.bar(x - w/2, l_lat, w, label='L-level (SRAM)', color='#3498db', alpha=0.8)
ax.bar(x + w/2, h_lat, w, label='H-level (DRAM)', color='#e67e22', alpha=0.8)
ax.set_xlabel('Configuration')
ax.set_ylabel('Latency (μs)')
ax.set_title('Per-Module Latency')
ax.set_xticks(x)
ax.set_xticklabels(configs, fontsize=8)
ax.legend()
ax.grid(axis='y', alpha=0.3)
plt.tight_layout()
plot_path = os.path.join(output_dir, 'benchmark_comparison.png')
plt.savefig(plot_path, dpi=150, bbox_inches='tight')
plt.close()
print(f" Plot saved: {plot_path}")
# ---- Plot 4: Memory breakdown ----
fig, axes = plt.subplots(1, 2, figsize=(12, 5))
fig.suptitle('Memory Analysis', fontsize=14, fontweight='bold')
ax = axes[0]
gpu_mem_baseline = [r['baseline']['total_gpu_memory_mb'] for r in results]
gpu_mem_tiered = [r['tiered']['total_gpu_memory_mb'] for r in results]
ax.bar(x - w/2, gpu_mem_baseline, w, label='Baseline', color='#e74c3c', alpha=0.8)
ax.bar(x + w/2, gpu_mem_tiered, w, label='Tiered', color='#2ecc71', alpha=0.8)
ax.set_xlabel('Configuration')
ax.set_ylabel('GPU Memory (MB)')
ax.set_title('Total GPU Memory')
ax.set_xticks(x)
ax.set_xticklabels(configs, fontsize=8)
ax.legend()
ax.grid(axis='y', alpha=0.3)
ax = axes[1]
triton_sram = [r['tiered']['triton_sram_probe_latency_us'] for r in results]
triton_dram = [r['tiered']['triton_dram_probe_latency_us'] for r in results]
ax.bar(x - w/2, triton_sram, w, label='SRAM Probe', color='#3498db', alpha=0.8)
ax.bar(x + w/2, triton_dram, w, label='DRAM Probe', color='#e67e22', alpha=0.8)
ax.set_xlabel('Configuration')
ax.set_ylabel('Latency (μs)')
ax.set_title('Triton Memory Probe Latency')
ax.set_xticks(x)
ax.set_xticklabels(configs, fontsize=8)
ax.legend()
ax.grid(axis='y', alpha=0.3)
plt.tight_layout()
plot_path = os.path.join(output_dir, 'memory_analysis.png')
plt.savefig(plot_path, dpi=150, bbox_inches='tight')
plt.close()
print(f" Plot saved: {plot_path}")