413 lines
15 KiB
Diff
413 lines
15 KiB
Diff
diff --git a/tools/profile_sage2_scheduler.py b/tools/profile_sage2_scheduler.py
|
|
new file mode 100644
|
|
index 0000000..03130e1
|
|
--- /dev/null
|
|
+++ b/tools/profile_sage2_scheduler.py
|
|
@@ -0,0 +1,407 @@
|
|
+"""Profile exact SageAttention2 preparation, mainloop, and tail scheduling on real H3 tensors."""
|
|
+
|
|
+from __future__ import annotations
|
|
+
|
|
+import argparse
|
|
+import hashlib
|
|
+import json
|
|
+import math
|
|
+from pathlib import Path
|
|
+
|
|
+import torch
|
|
+
|
|
+from h3_blackwell_runtime.sage2_entry import prepare_v
|
|
+from profile_attention_path import prepare_qkv, representative_attention_inputs, summarize
|
|
+
|
|
+
|
|
+CTA_Q = 128
|
|
+CTA_K = 64
|
|
+WARP_Q = 32
|
|
+WARP_K = 64
|
|
+V_SCALE_MAX = 2.25
|
|
+
|
|
+
|
|
+def tensor_sha256(value: torch.Tensor) -> str:
|
|
+ host_bytes = value.detach().contiguous().view(torch.uint8).cpu().numpy()
|
|
+ return hashlib.sha256(memoryview(host_bytes)).hexdigest()
|
|
+
|
|
+
|
|
+def event_measure(fn, *, warmup: int, iterations: int):
|
|
+ for _ in range(warmup):
|
|
+ fn()
|
|
+ torch.cuda.synchronize()
|
|
+ values = []
|
|
+ result = None
|
|
+ for _ in range(iterations):
|
|
+ started = torch.cuda.Event(enable_timing=True)
|
|
+ finished = torch.cuda.Event(enable_timing=True)
|
|
+ started.record()
|
|
+ result = fn()
|
|
+ finished.record()
|
|
+ finished.synchronize()
|
|
+ values.append(started.elapsed_time(finished) / 1000.0)
|
|
+ return summarize(values), result
|
|
+
|
|
+
|
|
+def quantize_qk(q: torch.Tensor, k: torch.Tensor, km: torch.Tensor):
|
|
+ import sageattention.core as sage_core
|
|
+
|
|
+ return sage_core.per_warp_int8_cuda(
|
|
+ q,
|
|
+ k,
|
|
+ km,
|
|
+ BLKQ=CTA_Q,
|
|
+ WARPQ=WARP_Q,
|
|
+ BLKK=CTA_K,
|
|
+ tensor_layout="NHD",
|
|
+ )
|
|
+
|
|
+
|
|
+def quantize_v(v: torch.Tensor):
|
|
+ import sageattention.core as sage_core
|
|
+
|
|
+ return sage_core.per_channel_fp8(
|
|
+ v,
|
|
+ tensor_layout="NHD",
|
|
+ scale_max=V_SCALE_MAX,
|
|
+ smooth_v=False,
|
|
+ )
|
|
+
|
|
+
|
|
+def run_mainloop(
|
|
+ q_int8: torch.Tensor,
|
|
+ k_int8: torch.Tensor,
|
|
+ v_fp8: torch.Tensor,
|
|
+ q_scale: torch.Tensor,
|
|
+ k_scale: torch.Tensor,
|
|
+ v_scale: torch.Tensor,
|
|
+ output: torch.Tensor,
|
|
+) -> torch.Tensor:
|
|
+ import sageattention.core as sage_core
|
|
+
|
|
+ sage_core.sm89_compile.qk_int8_sv_f8_accum_f16_fuse_v_scale_attn_inst_buf(
|
|
+ q_int8,
|
|
+ k_int8,
|
|
+ v_fp8,
|
|
+ output,
|
|
+ q_scale,
|
|
+ k_scale,
|
|
+ v_scale,
|
|
+ 0,
|
|
+ 0,
|
|
+ 2,
|
|
+ output.shape[-1] ** -0.5,
|
|
+ 0,
|
|
+ )
|
|
+ return output
|
|
+
|
|
+
|
|
+def prepare_quantized(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor):
|
|
+ km = k.mean(dim=1, keepdim=True)
|
|
+ q_int8, q_scale, k_int8, k_scale = quantize_qk(q, k, km)
|
|
+ v_fp8, v_scale, _ = quantize_v(v)
|
|
+ output = torch.empty(q.shape, dtype=q.dtype, device=q.device)
|
|
+ return km, q_int8, q_scale, k_int8, k_scale, v_fp8, v_scale, output
|
|
+
|
|
+
|
|
+def tail_row(
|
|
+ q: torch.Tensor,
|
|
+ k: torch.Tensor,
|
|
+ v: torch.Tensor,
|
|
+ *,
|
|
+ warmup: int,
|
|
+ iterations: int,
|
|
+) -> dict:
|
|
+ quantized = prepare_quantized(q, k, v)
|
|
+ _, q_int8, q_scale, k_int8, k_scale, v_fp8, v_scale, output = quantized
|
|
+ timing, result = event_measure(
|
|
+ lambda: run_mainloop(
|
|
+ q_int8, k_int8, v_fp8, q_scale, k_scale, v_scale, output,
|
|
+ ),
|
|
+ warmup=warmup,
|
|
+ iterations=iterations,
|
|
+ )
|
|
+ q_len = q.shape[1]
|
|
+ kv_len = k.shape[1]
|
|
+ q_ctas = math.ceil(q_len / CTA_Q)
|
|
+ k_iterations = math.ceil(kv_len / CTA_K)
|
|
+ return {
|
|
+ "q_len": q_len,
|
|
+ "kv_len": kv_len,
|
|
+ "q_ctas_per_head": q_ctas,
|
|
+ "k_iterations_per_cta": k_iterations,
|
|
+ "q_tail_rows": q_len % CTA_Q,
|
|
+ "k_tail_rows": kv_len % CTA_K,
|
|
+ "scheduled_q_rows": q_ctas * CTA_Q,
|
|
+ "q_row_efficiency": q_len / (q_ctas * CTA_Q),
|
|
+ "mainloop": timing,
|
|
+ "checksum": result.float().sum().item(),
|
|
+ }
|
|
+
|
|
+
|
|
+def difference(actual: torch.Tensor, expected: torch.Tensor) -> dict:
|
|
+ delta = actual.float() - expected.float()
|
|
+ return {
|
|
+ "equal": torch.equal(actual, expected),
|
|
+ "max_abs": delta.abs().max().item(),
|
|
+ "mean_abs": delta.abs().mean().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("--warmup", type=int, default=3)
|
|
+ parser.add_argument("--iterations", type=int, default=10)
|
|
+ parser.add_argument("--tail-iterations", type=int, default=5)
|
|
+ parser.add_argument("--skip-tail-study", action="store_true")
|
|
+ parser.add_argument("--cuda-profiler-capture", action="store_true")
|
|
+ parser.add_argument("--expected-output-sha256")
|
|
+ parser.add_argument("--device", default="cuda")
|
|
+ return parser.parse_args()
|
|
+
|
|
+
|
|
+def main() -> None:
|
|
+ args = parse_args()
|
|
+ block, hidden, rotation, _segments, metadata = representative_attention_inputs(args)
|
|
+ with torch.inference_mode():
|
|
+ q, k, v, _ = prepare_qkv(block, hidden, rotation, None)
|
|
+ km_timing, km = event_measure(
|
|
+ lambda: k.mean(dim=1, keepdim=True),
|
|
+ warmup=args.warmup,
|
|
+ iterations=args.iterations,
|
|
+ )
|
|
+ qk_timing, qk = event_measure(
|
|
+ lambda: quantize_qk(q, k, km),
|
|
+ warmup=args.warmup,
|
|
+ iterations=args.iterations,
|
|
+ )
|
|
+ q_int8, q_scale, k_int8, k_scale = qk
|
|
+ import sageattention.core as sage_core
|
|
+ import sageattention.quant as sage_quant
|
|
+
|
|
+ q_int8_probe = torch.empty(q.shape, dtype=torch.int8, device=q.device)
|
|
+ q_scale_probe = torch.empty_like(q_scale)
|
|
+ q_quant_timing, _ = event_measure(
|
|
+ lambda: sage_quant._fused.quant_per_warp_int8_cuda(
|
|
+ q, q_int8_probe, q_scale_probe, CTA_Q, WARP_Q, 0,
|
|
+ ),
|
|
+ warmup=args.warmup,
|
|
+ iterations=args.iterations,
|
|
+ )
|
|
+ k_int8_probe = torch.empty(k.shape, dtype=torch.int8, device=k.device)
|
|
+ k_scale_probe = torch.empty_like(k_scale)
|
|
+ k_quant_timing, _ = event_measure(
|
|
+ lambda: sage_quant._fused.quant_per_block_int8_fuse_sub_mean_cuda(
|
|
+ k, km.squeeze(1), k_int8_probe, k_scale_probe, CTA_K, 0,
|
|
+ ),
|
|
+ warmup=args.warmup,
|
|
+ iterations=args.iterations,
|
|
+ )
|
|
+ padded_k = math.ceil(v.shape[1] / CTA_K) * CTA_K
|
|
+ v_transposed = torch.empty(
|
|
+ (v.shape[0], v.shape[3], v.shape[2], padded_k),
|
|
+ dtype=v.dtype,
|
|
+ device=v.device,
|
|
+ )
|
|
+ v_transpose_timing, _ = event_measure(
|
|
+ lambda: sage_quant._fused.transpose_pad_permute_cuda(v, v_transposed, 0),
|
|
+ warmup=args.warmup,
|
|
+ iterations=args.iterations,
|
|
+ )
|
|
+ v_fp8_probe = torch.empty_like(v_transposed, dtype=torch.float8_e4m3fn)
|
|
+ v_scale_probe = torch.empty(
|
|
+ (v.shape[0], v.shape[2], v.shape[3]),
|
|
+ dtype=torch.float32,
|
|
+ device=v.device,
|
|
+ )
|
|
+ v_scale_quant_timing, _ = event_measure(
|
|
+ lambda: sage_quant._fused.scale_fuse_quant_cuda(
|
|
+ v_transposed, v_fp8_probe, v_scale_probe, v.shape[1], V_SCALE_MAX, 0,
|
|
+ ),
|
|
+ warmup=args.warmup,
|
|
+ iterations=args.iterations,
|
|
+ )
|
|
+ v_timing, vq = event_measure(
|
|
+ lambda: quantize_v(v),
|
|
+ warmup=args.warmup,
|
|
+ iterations=args.iterations,
|
|
+ )
|
|
+ v_fp8, v_scale, _ = vq
|
|
+ candidate_v_timing, candidate_vq = event_measure(
|
|
+ lambda: prepare_v(v, scale_max=V_SCALE_MAX),
|
|
+ warmup=args.warmup,
|
|
+ iterations=args.iterations,
|
|
+ )
|
|
+ candidate_v_fp8, candidate_v_scale = candidate_vq
|
|
+ output = torch.empty(q.shape, dtype=q.dtype, device=q.device)
|
|
+ mainloop_fn = lambda: run_mainloop(
|
|
+ q_int8, k_int8, v_fp8, q_scale, k_scale, v_scale, output,
|
|
+ )
|
|
+
|
|
+ if args.cuda_profiler_capture:
|
|
+ for _ in range(args.warmup):
|
|
+ mainloop_fn()
|
|
+ torch.cuda.synchronize()
|
|
+ torch.cuda.cudart().cudaProfilerStart()
|
|
+ captured = mainloop_fn()
|
|
+ torch.cuda.synchronize()
|
|
+ torch.cuda.cudart().cudaProfilerStop()
|
|
+ report = {
|
|
+ "metadata": metadata,
|
|
+ "capture": "one unchanged prequantized Sage2 mainloop",
|
|
+ "q_shape": list(q.shape),
|
|
+ "k_shape": list(k.shape),
|
|
+ "v_shape": list(v.shape),
|
|
+ "checksum": captured.float().sum().item(),
|
|
+ "scheduler": {
|
|
+ "cta_q": CTA_Q,
|
|
+ "cta_k": CTA_K,
|
|
+ "warp_q": WARP_Q,
|
|
+ "warp_k": WARP_K,
|
|
+ "warps_per_cta": 4,
|
|
+ "threads_per_cta": 128,
|
|
+ "dynamic_shared_memory_bytes": 32768,
|
|
+ "q_ctas_per_head": math.ceil(q.shape[1] / CTA_Q),
|
|
+ "heads": q.shape[2],
|
|
+ "grid_ctas": math.ceil(q.shape[1] / CTA_Q) * q.shape[2] * q.shape[0],
|
|
+ "k_iterations_per_cta": math.ceil(k.shape[1] / CTA_K),
|
|
+ "explicit_pipeline_stages": 2,
|
|
+ },
|
|
+ }
|
|
+ args.output.parent.mkdir(parents=True, exist_ok=True)
|
|
+ args.output.write_text(json.dumps(report, indent=2) + "\n", encoding="utf-8")
|
|
+ print(json.dumps(report, indent=2), flush=True)
|
|
+ return
|
|
+
|
|
+ mainloop_timing, manual_output = event_measure(
|
|
+ mainloop_fn,
|
|
+ warmup=args.warmup,
|
|
+ iterations=args.iterations,
|
|
+ )
|
|
+ candidate_output = torch.empty_like(output)
|
|
+ run_mainloop(
|
|
+ q_int8,
|
|
+ k_int8,
|
|
+ candidate_v_fp8,
|
|
+ q_scale,
|
|
+ k_scale,
|
|
+ candidate_v_scale,
|
|
+ candidate_output,
|
|
+ )
|
|
+ torch.cuda.synchronize()
|
|
+ reference = __import__("sageattention").sageattn(
|
|
+ q, k, v, tensor_layout="NHD", is_causal=False, smooth_k=False,
|
|
+ )
|
|
+ torch.cuda.synchronize()
|
|
+ output_sha256 = tensor_sha256(manual_output)
|
|
+ expected_output_matches = (
|
|
+ args.expected_output_sha256 is None
|
|
+ or output_sha256 == args.expected_output_sha256
|
|
+ )
|
|
+
|
|
+ tail_study = []
|
|
+ if not args.skip_tail_study:
|
|
+ q_lengths = sorted({
|
|
+ (q.shape[1] // CTA_Q) * CTA_Q,
|
|
+ (q.shape[1] // CTA_Q) * CTA_Q + 1,
|
|
+ q.shape[1],
|
|
+ })
|
|
+ kv_lengths = sorted({
|
|
+ (k.shape[1] // CTA_K) * CTA_K,
|
|
+ (k.shape[1] // CTA_K) * CTA_K + 1,
|
|
+ k.shape[1],
|
|
+ })
|
|
+ for q_len in q_lengths:
|
|
+ tail_study.append({
|
|
+ "sweep": "q_tail_fixed_kv",
|
|
+ **tail_row(
|
|
+ q[:, :q_len], k, v,
|
|
+ warmup=args.warmup,
|
|
+ iterations=args.tail_iterations,
|
|
+ ),
|
|
+ })
|
|
+ for kv_len in kv_lengths:
|
|
+ tail_study.append({
|
|
+ "sweep": "kv_tail_fixed_q",
|
|
+ **tail_row(
|
|
+ q, k[:, :kv_len], v[:, :kv_len],
|
|
+ warmup=args.warmup,
|
|
+ iterations=args.tail_iterations,
|
|
+ ),
|
|
+ })
|
|
+
|
|
+ report = {
|
|
+ "metadata": metadata,
|
|
+ "q_shape": list(q.shape),
|
|
+ "k_shape": list(k.shape),
|
|
+ "v_shape": list(v.shape),
|
|
+ "warmup": args.warmup,
|
|
+ "iterations": args.iterations,
|
|
+ "phase_timings": {
|
|
+ "k_mean_and_smoothing_preparation": km_timing,
|
|
+ "qk_int8_quantization": qk_timing,
|
|
+ "q_int8_quantization": q_quant_timing,
|
|
+ "k_int8_subtract_mean_quantization": k_quant_timing,
|
|
+ "v_fp8_transpose_scale_quantization": v_timing,
|
|
+ "v_transpose_pad_permute": v_transpose_timing,
|
|
+ "v_scale_fp8_quantization": v_scale_quant_timing,
|
|
+ "vortex_direct_v_fp8_preparation": candidate_v_timing,
|
|
+ "fused_mainloop": mainloop_timing,
|
|
+ },
|
|
+ "fused_mainloop_phases": {
|
|
+ "int8_qk": "fused inside qk_int_sv_f8_attn_kernel",
|
|
+ "scale_application": "fused inside qk_int_sv_f8_attn_kernel",
|
|
+ "online_softmax": "fused inside qk_int_sv_f8_attn_kernel",
|
|
+ "pv_accumulation": "fused inside qk_int_sv_f8_attn_kernel",
|
|
+ "final_normalization_and_output": "fused inside qk_int_sv_f8_attn_kernel",
|
|
+ "timing_policy": "Do not assign independent wall time without changing the exact kernel schedule; use source-correlated hardware counters.",
|
|
+ },
|
|
+ "scheduler": {
|
|
+ "cta_q": CTA_Q,
|
|
+ "cta_k": CTA_K,
|
|
+ "warp_q": WARP_Q,
|
|
+ "warp_k": WARP_K,
|
|
+ "warps_per_cta": 4,
|
|
+ "threads_per_cta": 128,
|
|
+ "dynamic_shared_memory_bytes": 32768,
|
|
+ "q_ctas_per_head": math.ceil(q.shape[1] / CTA_Q),
|
|
+ "heads": q.shape[2],
|
|
+ "grid_ctas": math.ceil(q.shape[1] / CTA_Q) * q.shape[2] * q.shape[0],
|
|
+ "k_iterations_per_cta": math.ceil(k.shape[1] / CTA_K),
|
|
+ "q_tail_rows": q.shape[1] % CTA_Q,
|
|
+ "k_tail_rows": k.shape[1] % CTA_K,
|
|
+ "explicit_pipeline_stages": 2,
|
|
+ },
|
|
+ "manual_decomposition_vs_public_sage2": difference(manual_output, reference),
|
|
+ "vortex_v_fp8_vs_sage2": difference(candidate_v_fp8, v_fp8),
|
|
+ "vortex_v_scale_vs_sage2": difference(candidate_v_scale, v_scale),
|
|
+ "vortex_v_mainloop_vs_sage2": difference(candidate_output, manual_output),
|
|
+ "manual_checksum": manual_output.float().sum().item(),
|
|
+ "reference_checksum": reference.float().sum().item(),
|
|
+ "output_sha256": output_sha256,
|
|
+ "expected_output_sha256": args.expected_output_sha256,
|
|
+ "expected_output_matches": expected_output_matches,
|
|
+ "tail_study": tail_study,
|
|
+ }
|
|
+ args.output.parent.mkdir(parents=True, exist_ok=True)
|
|
+ args.output.write_text(json.dumps(report, indent=2) + "\n", encoding="utf-8")
|
|
+ print(json.dumps(report, indent=2), flush=True)
|
|
+ if not expected_output_matches:
|
|
+ raise RuntimeError(
|
|
+ f"output SHA256 mismatch: expected {args.expected_output_sha256}, got {output_sha256}"
|
|
+ )
|
|
+
|
|
+
|
|
+if __name__ == "__main__":
|
|
+ main()
|