252 lines
13 KiB
Python
252 lines
13 KiB
Python
"""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()
|