164 lines
5.9 KiB
Diff
164 lines
5.9 KiB
Diff
|
|
diff --git a/tools/validate_cute_qkv_block.py b/tools/validate_cute_qkv_block.py
|
||
|
|
new file mode 100644
|
||
|
|
index 0000000..84811c1
|
||
|
|
--- /dev/null
|
||
|
|
+++ b/tools/validate_cute_qkv_block.py
|
||
|
|
@@ -0,0 +1,157 @@
|
||
|
|
+"""Alternate baseline and CuTe-QKV execution inside one loaded H3 block."""
|
||
|
|
+
|
||
|
|
+from __future__ import annotations
|
||
|
|
+
|
||
|
|
+import argparse
|
||
|
|
+import json
|
||
|
|
+import os
|
||
|
|
+import time
|
||
|
|
+from pathlib import Path
|
||
|
|
+
|
||
|
|
+import torch
|
||
|
|
+
|
||
|
|
+from h3_blackwell_runtime.adaln import H3CurveAdaLN
|
||
|
|
+from h3_blackwell_runtime.block import H3DiTBlock
|
||
|
|
+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
|
||
|
|
+
|
||
|
|
+
|
||
|
|
+def parse_args() -> argparse.Namespace:
|
||
|
|
+ parser = argparse.ArgumentParser(description=__doc__)
|
||
|
|
+ parser.add_argument("--output", type=Path, required=True)
|
||
|
|
+ parser.add_argument("--model-path", default="/models/minimax_h3_fl2va_pruned_nvfp4.safetensors")
|
||
|
|
+ parser.add_argument("--block-index", type=int, 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("--attention", default="sage2")
|
||
|
|
+ parser.add_argument(
|
||
|
|
+ "--feature", choices=("qkv_ring", "modulate_fusion", "swiglu_fusion"), default="qkv_ring",
|
||
|
|
+ )
|
||
|
|
+ parser.add_argument("--warmup", type=int, default=2)
|
||
|
|
+ parser.add_argument("--iterations", type=int, default=10)
|
||
|
|
+ parser.add_argument("--device", default="cuda")
|
||
|
|
+ return parser.parse_args()
|
||
|
|
+
|
||
|
|
+
|
||
|
|
+def sync() -> None:
|
||
|
|
+ torch.cuda.synchronize()
|
||
|
|
+
|
||
|
|
+
|
||
|
|
+def summarize(values: list[float]) -> dict[str, float]:
|
||
|
|
+ ordered = sorted(values)
|
||
|
|
+ middle = len(ordered) // 2
|
||
|
|
+ median = (
|
||
|
|
+ ordered[middle]
|
||
|
|
+ if len(ordered) % 2
|
||
|
|
+ else (ordered[middle - 1] + ordered[middle]) / 2
|
||
|
|
+ )
|
||
|
|
+ return {
|
||
|
|
+ "mean_s": sum(values) / len(values),
|
||
|
|
+ "p50_s": median,
|
||
|
|
+ "min_s": ordered[0],
|
||
|
|
+ "max_s": ordered[-1],
|
||
|
|
+ }
|
||
|
|
+
|
||
|
|
+
|
||
|
|
+def main() -> None:
|
||
|
|
+ args = parse_args()
|
||
|
|
+ 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, _ = random_av_latents(
|
||
|
|
+ args.width, args.height, args.frames, args.seed, device=args.device,
|
||
|
|
+ )
|
||
|
|
+ sigmas = beta_sigmas(args.steps, device=args.device)
|
||
|
|
+ sigma = sigmas[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,
|
||
|
|
+ )
|
||
|
|
+ adaln_values = tuple(value.detach() for value in adaln(timesteps))
|
||
|
|
+
|
||
|
|
+ def run(enabled: bool):
|
||
|
|
+ if args.feature == "modulate_fusion":
|
||
|
|
+ block.fused_nvfp4_modulation = enabled
|
||
|
|
+ elif args.feature == "swiglu_fusion":
|
||
|
|
+ block.mlp.fused_nvfp4_swiglu = enabled
|
||
|
|
+ else:
|
||
|
|
+ if enabled:
|
||
|
|
+ os.environ["H3_CUTE_QKV_RING"] = "1"
|
||
|
|
+ else:
|
||
|
|
+ os.environ.pop("H3_CUTE_QKV_RING", None)
|
||
|
|
+ return block(hidden, rotation, *adaln_values, segments)
|
||
|
|
+
|
||
|
|
+ with torch.inference_mode():
|
||
|
|
+ reference = run(False)
|
||
|
|
+ candidate = run(True)
|
||
|
|
+ sync()
|
||
|
|
+ delta = candidate.float() - reference.float()
|
||
|
|
+ for _ in range(args.warmup):
|
||
|
|
+ run(False)
|
||
|
|
+ run(True)
|
||
|
|
+ sync()
|
||
|
|
+ baseline_times = []
|
||
|
|
+ candidate_times = []
|
||
|
|
+ last_reference = reference
|
||
|
|
+ last_candidate = candidate
|
||
|
|
+ for _ in range(args.iterations):
|
||
|
|
+ sync()
|
||
|
|
+ started = time.perf_counter()
|
||
|
|
+ last_reference = run(False)
|
||
|
|
+ sync()
|
||
|
|
+ baseline_times.append(time.perf_counter() - started)
|
||
|
|
+
|
||
|
|
+ sync()
|
||
|
|
+ started = time.perf_counter()
|
||
|
|
+ last_candidate = run(True)
|
||
|
|
+ sync()
|
||
|
|
+ candidate_times.append(time.perf_counter() - started)
|
||
|
|
+
|
||
|
|
+ baseline = summarize(baseline_times)
|
||
|
|
+ candidate_timing = summarize(candidate_times)
|
||
|
|
+ report = {
|
||
|
|
+ "device": torch.cuda.get_device_name(),
|
||
|
|
+ "block_index": args.block_index,
|
||
|
|
+ "feature": args.feature,
|
||
|
|
+ "hidden_shape": list(hidden.shape),
|
||
|
|
+ "iterations": args.iterations,
|
||
|
|
+ "equal": torch.equal(reference, candidate),
|
||
|
|
+ "max_abs": delta.abs().max().item(),
|
||
|
|
+ "mean_abs": delta.abs().mean().item(),
|
||
|
|
+ "reference_checksum": last_reference.float().sum().item(),
|
||
|
|
+ "candidate_checksum": last_candidate.float().sum().item(),
|
||
|
|
+ "baseline": baseline,
|
||
|
|
+ "candidate": candidate_timing,
|
||
|
|
+ "p50_improvement_percent": (
|
||
|
|
+ 1.0 - candidate_timing["p50_s"] / baseline["p50_s"]
|
||
|
|
+ ) * 100.0,
|
||
|
|
+ }
|
||
|
|
+ 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 __name__ == "__main__":
|
||
|
|
+ main()
|