160 lines
5.7 KiB
Python
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()
|