"""Profile H3 Sage2 against Sol routing with exact conditioning sinks.""" from __future__ import annotations import argparse import json import time from pathlib import Path import torch from h3_blackwell_runtime.attention import run_attention from profile_attention_path import prepare_qkv, representative_attention_inputs, summarize, sync def timed_iterations(fn, warmup: int, iterations: int) -> tuple[dict[str, float], torch.Tensor]: with torch.inference_mode(): for _ in range(warmup): fn() values = [] output = None for _ in range(iterations): sync() started = time.perf_counter() output = fn() sync() values.append(time.perf_counter() - started) return summarize(values), output def difference(actual: torch.Tensor, expected: torch.Tensor) -> dict[str, float]: delta = actual.float() - expected.float() expected_float = expected.float() return { "max_abs": delta.abs().max().item(), "mean_abs": delta.abs().mean().item(), "relative_l2": (delta.norm() / expected_float.norm().clamp_min(1e-12)).item(), } def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--model-path", default="/models/minimax_h3_fl2va_pruned_nvfp4.safetensors") parser.add_argument("--output", type=Path, required=True) parser.add_argument("--width", type=int, default=1344) parser.add_argument("--height", type=int, default=768) parser.add_argument("--frames", type=int, default=124) parser.add_argument("--steps", type=int, default=12) parser.add_argument("--sampler-step", type=int, default=1) parser.add_argument("--seed", type=int, default=440420) parser.add_argument("--text-tokens", type=int, default=100) parser.add_argument("--block-index", type=int, default=24) parser.add_argument("--attention", default="sage2", choices=("sage2",)) parser.add_argument("--taus", nargs="+", type=float, default=(0.0, 0.4, 0.8, 1.0, 1.3)) parser.add_argument("--warmup", type=int, default=3) parser.add_argument("--iterations", type=int, default=5) parser.add_argument("--device", default="cuda") parser.add_argument("--int8-qk", action="store_true") parser.add_argument("--int8-pv", action="store_true") return parser.parse_args() def main() -> None: args = parse_args() if args.int8_pv and not args.int8_qk: raise ValueError("--int8-pv requires --int8-qk") block, hidden, rotation, segments, metadata = representative_attention_inputs(args) q, k, v, _ = prepare_qkv(block, hidden, rotation, None) sequence = q.shape[1] sage_timing, sage_output_hnd = timed_iterations( lambda: run_attention( q.transpose(1, 2).contiguous(), k.transpose(1, 2).contiguous(), v.transpose(1, 2).contiguous(), backend="sage2", is_causal=False, ), args.warmup, args.iterations, ) sage_output = sage_output_hnd.transpose(1, 2).contiguous() try: from sol_kernel import sol_attn except ImportError as error: raise RuntimeError("Sol-Attn must be installed to profile the hybrid policy") from error conditioning_stop = segments[-1][0] conditioning_blocks = (conditioning_stop + 63) // 64 sink_modes = { "off": ((0, 0), (0, 0)), "exact_kv": ((0, conditioning_blocks), (0, 0)), "exact_kv_and_rows": ((0, conditioning_blocks), (0, conditioning_blocks)), } spans = { "conditioning": (0, conditioning_stop), "video": (conditioning_stop, sequence), "all": (0, sequence), } results = [] for sink_name, (sink_blocks, sink_q) in sink_modes.items(): for tau in args.taus: timing, output = timed_iterations( lambda tau=tau, sink_blocks=sink_blocks, sink_q=sink_q: sol_attn( q, k, v, tau=tau, thresh_type="diag", int8_qk=args.int8_qk, int8_pv=args.int8_pv, sink_blocks=sink_blocks, sink_q=sink_q, ), args.warmup, args.iterations, ) results.append({ "tau": tau, "sink": sink_name, "sink_blocks": list(sink_blocks), "sink_q": list(sink_q), "timing": timing, "speedup_vs_sage2_p50": sage_timing["p50_s"] / timing["p50_s"], "difference_vs_sage2": { name: difference(output[:, start:stop], sage_output[:, start:stop]) for name, (start, stop) in spans.items() }, }) print( sink_name, f"tau={tau:.2f}", f"p50_ms={timing['p50_s'] * 1000:.3f}", f"speedup={sage_timing['p50_s'] / timing['p50_s']:.3f}x", f"rel_l2={results[-1]['difference_vs_sage2']['all']['relative_l2']:.6f}", flush=True, ) report = { "metadata": metadata, "conditioning_stop": conditioning_stop, "conditioning_blocks": conditioning_blocks, "sage2": { "timing": sage_timing, "checksum": sage_output.float().sum().item(), }, "warmup": args.warmup, "iterations": args.iterations, "int8_qk": args.int8_qk, "int8_pv": args.int8_pv, "results": results, } args.output.parent.mkdir(parents=True, exist_ok=True) args.output.write_text(json.dumps(report, indent=2) + "\n", encoding="utf-8") if __name__ == "__main__": main()