"""Microbenchmark H3 attention kernels and Q/K/V layout costs from captured real block tensors.""" from __future__ import annotations import argparse import json import time import warnings from pathlib import Path warnings.filterwarnings("ignore", message="Found GPU0 NVIDIA GB10 which is of cuda capability 12.1.*", category=UserWarning) import torch from h3_blackwell_runtime.attention import AVAILABLE_BACKENDS, qkv_to_bshd, rms_rope_split_half_, run_attention, run_sol_attention_bshd from h3_blackwell_runtime.adaln import H3CurveAdaLN from h3_blackwell_runtime.block import H3DiTBlock, modulate_segments from h3_blackwell_runtime.checkpoint import H3Checkpoint from h3_blackwell_runtime.packing import H3PromptPacker from h3_blackwell_runtime.rope import h3_rope_rotation from h3_blackwell_runtime.sampler import _audio_sigma, _model_sigma, beta_sigmas from h3_blackwell_runtime.t2v import random_av_latents from h3_blackwell_runtime.attention import rms_norm def sync() -> None: if torch.cuda.is_available(): torch.cuda.synchronize() def summarize(values: list[float]) -> dict[str, float]: ordered = sorted(values) def percentile(percent: float) -> float: if len(ordered) == 1: return ordered[0] rank = (len(ordered) - 1) * percent low = int(rank) high = min(low + 1, len(ordered) - 1) weight = rank - low return ordered[low] * (1.0 - weight) + ordered[high] * weight return { "count": len(values), "mean_s": sum(values) / len(values), "p50_s": percentile(0.50), "p90_s": percentile(0.90), "p95_s": percentile(0.95), "p99_s": percentile(0.99), "min_s": ordered[0], "max_s": ordered[-1], } def timed(stats: dict[str, list[float]], name: str, fn): sync() started = time.perf_counter() value = fn() sync() stats.setdefault(name, []).append(time.perf_counter() - started) return value def prepare_qkv(block, x: torch.Tensor, rotation: torch.Tensor, segment: tuple[int, int, int] | None): attention = block.attention sequence = x.shape[0] inner = attention.heads * attention.head_dim qkv = attention.qkv_proj(x) q, k, v = qkv.split(inner, dim=-1) q = q.view(1, sequence, attention.heads, attention.head_dim) k = k.view(1, sequence, attention.heads, attention.head_dim) v = v.view(1, sequence, attention.heads, attention.head_dim) q, k = rms_rope_split_half_(q, k, rotation, attention.q_norm_weight, attention.k_norm_weight, attention.eps) full_qkv = qkv if segment is not None: start, end, _kind = segment q = q[:, start:end].contiguous() k = k[:, start:end].contiguous() v = v[:, start:end].contiguous() full_qkv = None return q, k, v, full_qkv def representative_attention_inputs(args: argparse.Namespace): torch.manual_seed(args.seed) checkpoint = H3Checkpoint(args.model_path, device=args.device) block = H3DiTBlock.from_checkpoint(checkpoint, args.block_index, attention_backend=args.attention).eval() adaln = H3CurveAdaLN.from_checkpoint(checkpoint, f"blocks.{args.block_index}.adaln_proj").eval() packer = H3PromptPacker(checkpoint) video, audio, aligned_frames = random_av_latents(args.width, args.height, args.frames, args.seed, device=args.device) sigma = beta_sigmas(args.steps, device=args.device)[args.sampler_step - 1] native_audio = audio.to(torch.bfloat16) * (_audio_sigma(sigma) / sigma) text = torch.randn(1, args.text_tokens, 5376, device=args.device, dtype=torch.bfloat16) hidden, timesteps, segments, positions, _, _ = packer(text, video, native_audio, _model_sigma(sigma)) rotation = h3_rope_rotation(positions.to(args.device), checkpoint.tensor("rope.inv_freq", dtype=torch.float32), hidden.dtype) shift_msa, scale_msa, _gate_msa, _shift_mlp, _scale_mlp, _gate_mlp = adaln(timesteps) with torch.inference_mode(): h_msa = modulate_segments(rms_norm(hidden, block.norm1_weight, block.norm_eps), shift_msa, scale_msa, segments) metadata = { "width": args.width, "height": args.height, "frames": aligned_frames, "steps": args.steps, "sampler_step": args.sampler_step, "seed": args.seed, "text_tokens": args.text_tokens, "block_index": args.block_index, "attention": args.attention, "hidden_shape": list(hidden.shape), "h_msa_shape": list(h_msa.shape), "segments": segments, } return block, h_msa, rotation, segments, metadata def run_path(q_src: torch.Tensor, k_src: torch.Tensor, v_src: torch.Tensor, backend: str, stats: dict[str, list[float]] | None = None): sequence = q_src.shape[1] inner = q_src.shape[2] * q_src.shape[3] q = timed(stats, "q_transpose_contiguous", lambda: q_src.transpose(1, 2).contiguous()) if stats is not None else q_src.transpose(1, 2).contiguous() k = timed(stats, "k_transpose_contiguous", lambda: k_src.transpose(1, 2).contiguous()) if stats is not None else k_src.transpose(1, 2).contiguous() v = timed(stats, "v_transpose_contiguous", lambda: v_src.transpose(1, 2).contiguous()) if stats is not None else v_src.transpose(1, 2).contiguous() out = timed(stats, "attention_kernel", lambda: run_attention(q, k, v, backend=backend, is_causal=False)) if stats is not None else run_attention(q, k, v, backend=backend, is_causal=False) rows = timed(stats, "output_reshape", lambda: out.transpose(1, 2).reshape(sequence, inner).contiguous()) if stats is not None else out.transpose(1, 2).reshape(sequence, inner).contiguous() return rows def run_sol_native_path(q_src: torch.Tensor, k_src: torch.Tensor, v_src: torch.Tensor, stats: dict[str, list[float]] | None = None): sequence = q_src.shape[1] inner = q_src.shape[2] * q_src.shape[3] q = timed(stats, "q_bshd_contiguous", lambda: q_src.contiguous()) if stats is not None else q_src.contiguous() k = timed(stats, "k_bshd_contiguous", lambda: k_src.contiguous()) if stats is not None else k_src.contiguous() v = timed(stats, "v_bshd_contiguous", lambda: v_src.contiguous()) if stats is not None else v_src.contiguous() out = timed(stats, "attention_kernel", lambda: run_sol_attention_bshd(q, k, v, is_causal=False)) if stats is not None else run_sol_attention_bshd(q, k, v, is_causal=False) return timed(stats, "output_reshape", lambda: out.reshape(sequence, inner).contiguous()) if stats is not None else out.reshape(sequence, inner).contiguous() def run_sol_fused_path(qkv: torch.Tensor, heads: int, head_dim: int, stats: dict[str, list[float]] | None = None): sequence = qkv.shape[0] inner = heads * head_dim q, k, v = timed(stats, "qkv_to_bshd", lambda: qkv_to_bshd(qkv, heads, head_dim)) if stats is not None else qkv_to_bshd(qkv, heads, head_dim) out = timed(stats, "attention_kernel", lambda: run_sol_attention_bshd(q, k, v, is_causal=False)) if stats is not None else run_sol_attention_bshd(q, k, v, is_causal=False) return timed(stats, "output_reshape", lambda: out.reshape(sequence, inner).contiguous()) if stats is not None else out.reshape(sequence, inner).contiguous() def layout_timing_names(layout_mode: str) -> tuple[str, ...]: if layout_mode == "sol_fused": return ("qkv_to_bshd", "attention_kernel", "output_reshape") if layout_mode == "sol_native": return ("q_bshd_contiguous", "k_bshd_contiguous", "v_bshd_contiguous", "attention_kernel", "output_reshape") return ("q_transpose_contiguous", "k_transpose_contiguous", "v_transpose_contiguous", "attention_kernel", "output_reshape") def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--model-path", default="/models/minimax_h3_fl2va_pruned_nvfp4.safetensors") parser.add_argument("--output", type=Path, default=Path("/output/h3-blackwell-runtime/benchmarks/attention-path-profile.json")) parser.add_argument("--width", type=int, default=960) parser.add_argument("--height", type=int, default=544) 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=440407) parser.add_argument("--text-tokens", type=int, default=93) parser.add_argument("--block-index", type=int, default=24) parser.add_argument("--attention", choices=AVAILABLE_BACKENDS, default="sage2", help="Backend used only while building representative upstream tensors.") parser.add_argument("--backends", nargs="+", choices=AVAILABLE_BACKENDS, default=("sol_attn", "sage2", "sage3", "sage3_mean", "kj_sage_fp8", "kj_sage_fp8pp", "sdpa")) parser.add_argument("--sol-layout", choices=("hnd", "native", "fused", "both", "all"), default="hnd", help="Compare generic HND Sol path with direct BSHD and fused QKV layout paths.") parser.add_argument("--segments", nargs="+", choices=("all", "text", "secondary", "video"), default=("all",)) parser.add_argument("--warmup", type=int, default=20) parser.add_argument("--iterations", type=int, default=100) parser.add_argument("--device", default="cuda") return parser.parse_args() def main() -> None: args = parse_args() block, x, rotation, segments, metadata = representative_attention_inputs(args) segment_map = {"all": None, "text": segments[0], "secondary": segments[1], "video": segments[2]} results = [] with torch.inference_mode(): for segment_name in args.segments: q_src, k_src, v_src, qkv_src = prepare_qkv(block, x, rotation, segment_map[segment_name]) reference = None reference_backend = None for backend in args.backends: layout_modes = ["hnd"] if backend == "sol_attn" and args.sol_layout != "hnd": layout_modes = { "native": ["sol_native"], "fused": ["sol_fused"], "both": ["hnd", "sol_native"], "all": ["hnd", "sol_native", "sol_fused"], }[args.sol_layout] for layout_mode in layout_modes: try: if layout_mode == "sol_fused" and qkv_src is None: raise ValueError("sol_fused layout currently requires the full unsegmented QKV tensor") for _ in range(args.warmup): if layout_mode == "sol_fused": run_sol_fused_path(qkv_src, block.attention.heads, block.attention.head_dim) elif layout_mode == "sol_native": run_sol_native_path(q_src, k_src, v_src) else: run_path(q_src, k_src, v_src, backend) stats: dict[str, list[float]] = {} output = None for _ in range(args.iterations): if layout_mode == "sol_fused": output = run_sol_fused_path(qkv_src, block.attention.heads, block.attention.head_dim, stats) elif layout_mode == "sol_native": output = run_sol_native_path(q_src, k_src, v_src, stats) else: output = run_path(q_src, k_src, v_src, backend, stats) if reference is None: reference = output reference_backend = f"{backend}:{layout_mode}" diff = {"max": 0.0, "mean": 0.0} else: delta = (output.float() - reference.float()).abs() diff = {"max": delta.max().item(), "mean": delta.mean().item()} summarized = {name: summarize(values) for name, values in stats.items()} total_mean = sum(summarized[name]["mean_s"] for name in layout_timing_names(layout_mode)) results.append( { "segment": segment_name, "segment_tuple": segment_map[segment_name], "backend": backend, "layout_mode": layout_mode, "q_shape": list(q_src.shape), "output_shape": list(output.shape), "timings": summarized, "layout_attention_total_mean_s": total_mean, "reference_backend": reference_backend, "reference_diff": diff, "status": "ok", } ) print(segment_name, backend, layout_mode, "attn_ms", round(summarized["attention_kernel"]["mean_s"] * 1000, 3), "total_ms", round(total_mean * 1000, 3), flush=True) except Exception as exc: results.append({"segment": segment_name, "backend": backend, "layout_mode": layout_mode, "status": "failed", "error": repr(exc)}) print(segment_name, backend, layout_mode, "FAILED", repr(exc), flush=True) output = {"metadata": metadata, "segments": segments, "warmup": args.warmup, "iterations": args.iterations, "results": results} args.output.parent.mkdir(parents=True, exist_ok=True) args.output.write_text(json.dumps(output, indent=2), encoding="utf-8") print(json.dumps(output, indent=2), flush=True) if __name__ == "__main__": main()