| |
| """ |
| CLI entry point for HRM SRAM/DRAM Memory Tiering Benchmarks. |
| |
| Usage: |
| # Compare tiered vs baseline HRM |
| python run_benchmark.py --mode compare --batch-sizes 1,8,32 --seq-lens 64,128 |
| |
| # Benchmark tiered model only |
| python run_benchmark.py --mode tiered --warmup 5 --iterations 50 --output results.json |
| |
| # Generate plots |
| python run_benchmark.py --mode compare --plot --output-dir benchmark_results/ |
| |
| # Quick smoke test |
| python run_benchmark.py --mode tiered --warmup 1 --iterations 3 --batch-sizes 2 --seq-lens 16 |
| """ |
|
|
| import argparse |
| import json |
| import os |
| import sys |
| from dataclasses import asdict |
|
|
| |
| sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) |
|
|
| from benchmark import ( |
| benchmark_tiered_model, |
| benchmark_baseline_model, |
| compare_models, |
| print_results_table, |
| generate_plots, |
| BenchmarkResult, |
| ) |
|
|
|
|
| def parse_int_list(s: str): |
| return [int(x.strip()) for x in s.split(',')] |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser( |
| description='HRM SRAM/DRAM Memory Tiering Benchmark Suite (Triton)', |
| formatter_class=argparse.RawDescriptionHelpFormatter, |
| epilog=__doc__, |
| ) |
|
|
| parser.add_argument( |
| '--mode', choices=['tiered', 'baseline', 'compare'], |
| default='compare', |
| help='Benchmark mode: tiered-only, baseline-only, or comparison (default: compare)', |
| ) |
| parser.add_argument( |
| '--batch-sizes', type=str, default='1,8,32', |
| help='Comma-separated batch sizes to benchmark (default: 1,8,32)', |
| ) |
| parser.add_argument( |
| '--seq-lens', type=str, default='64,128', |
| help='Comma-separated sequence lengths (default: 64,128)', |
| ) |
| parser.add_argument( |
| '--hidden-size', type=int, default=512, |
| help='Model hidden size (default: 512)', |
| ) |
| parser.add_argument( |
| '--num-heads', type=int, default=8, |
| help='Number of attention heads (default: 8)', |
| ) |
| parser.add_argument( |
| '--H-cycles', type=int, default=2, |
| help='H-level recurrence cycles (default: 2)', |
| ) |
| parser.add_argument( |
| '--L-cycles', type=int, default=2, |
| help='L-level recurrence cycles (default: 2)', |
| ) |
| parser.add_argument( |
| '--H-layers', type=int, default=4, |
| help='H-level transformer layers (default: 4)', |
| ) |
| parser.add_argument( |
| '--L-layers', type=int, default=4, |
| help='L-level transformer layers (default: 4)', |
| ) |
| parser.add_argument( |
| '--warmup', type=int, default=5, |
| help='Warmup iterations (default: 5)', |
| ) |
| parser.add_argument( |
| '--iterations', type=int, default=20, |
| help='Benchmark iterations (default: 20)', |
| ) |
| parser.add_argument( |
| '--output', type=str, default=None, |
| help='Output JSON file path for results', |
| ) |
| parser.add_argument( |
| '--output-dir', type=str, default='benchmark_results', |
| help='Directory for output files and plots (default: benchmark_results/)', |
| ) |
| parser.add_argument( |
| '--plot', action='store_true', |
| help='Generate comparison plots (requires matplotlib)', |
| ) |
| parser.add_argument( |
| '--device', type=str, default=None, |
| help='Device: cuda, cpu, or cuda:N (default: auto-detect)', |
| ) |
|
|
| args = parser.parse_args() |
|
|
| batch_sizes = parse_int_list(args.batch_sizes) |
| seq_lens = parse_int_list(args.seq_lens) |
|
|
| |
| import torch |
| if args.device: |
| device = torch.device(args.device) |
| else: |
| device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') |
|
|
| print(f"\n{'='*60}") |
| print(f" HRM SRAM/DRAM Memory Tiering Benchmark") |
| print(f" Device: {device}") |
| print(f" Mode: {args.mode}") |
| print(f" Batch sizes: {batch_sizes}") |
| print(f" Sequence lengths: {seq_lens}") |
| print(f" Hidden size: {args.hidden_size}") |
| print(f" H/L cycles: {args.H_cycles}/{args.L_cycles}") |
| print(f" H/L layers: {args.H_layers}/{args.L_layers}") |
| print(f" Warmup: {args.warmup}, Iterations: {args.iterations}") |
| print(f"{'='*60}") |
|
|
| if args.mode == 'compare': |
| results = compare_models( |
| batch_sizes=batch_sizes, |
| seq_lens=seq_lens, |
| hidden_size=args.hidden_size, |
| warmup=args.warmup, |
| iterations=args.iterations, |
| device=device, |
| ) |
|
|
| if args.plot: |
| generate_plots(results, output_dir=args.output_dir) |
|
|
| |
| output_path = args.output or os.path.join(args.output_dir, 'results.json') |
| os.makedirs(os.path.dirname(output_path) or '.', exist_ok=True) |
| with open(output_path, 'w') as f: |
| json.dump(results, f, indent=2, default=str) |
| print(f"\n Results saved: {output_path}") |
|
|
| elif args.mode == 'tiered': |
| all_results = [] |
| for bs in batch_sizes: |
| for sl in seq_lens: |
| print(f"\n Benchmarking tiered model: bs={bs}, seq={sl}") |
| r = benchmark_tiered_model( |
| batch_size=bs, seq_len=sl, |
| hidden_size=args.hidden_size, |
| num_heads=args.num_heads, |
| H_cycles=args.H_cycles, L_cycles=args.L_cycles, |
| H_layers=args.H_layers, L_layers=args.L_layers, |
| warmup=args.warmup, iterations=args.iterations, |
| device=device, |
| ) |
| all_results.append(r) |
|
|
| print_results_table(all_results) |
|
|
| if args.output: |
| with open(args.output, 'w') as f: |
| json.dump([asdict(r) for r in all_results], f, indent=2, default=str) |
| print(f" Results saved: {args.output}") |
|
|
| elif args.mode == 'baseline': |
| all_results = [] |
| for bs in batch_sizes: |
| for sl in seq_lens: |
| print(f"\n Benchmarking baseline model: bs={bs}, seq={sl}") |
| r = benchmark_baseline_model( |
| batch_size=bs, seq_len=sl, |
| hidden_size=args.hidden_size, |
| num_heads=args.num_heads, |
| H_cycles=args.H_cycles, L_cycles=args.L_cycles, |
| H_layers=args.H_layers, L_layers=args.L_layers, |
| warmup=args.warmup, iterations=args.iterations, |
| device=device, |
| ) |
| all_results.append(r) |
|
|
| print_results_table(all_results) |
|
|
| if args.output: |
| with open(args.output, 'w') as f: |
| json.dump([asdict(r) for r in all_results], f, indent=2, default=str) |
| print(f" Results saved: {args.output}") |
|
|
| print("\n Done!\n") |
|
|
|
|
| if __name__ == '__main__': |
| main() |
|
|