h3-blackwell-runtime/tools/profile_hybrid_attention.py
2026-08-25 20:30:22 +07:00

160 lines
5.7 KiB
Python

"""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()