h3-blackwell-runtime/research/sage2_entry_fusion/patches/0002-entry-validation.patch
2026-08-25 20:30:22 +07:00

341 lines
13 KiB
Diff

diff --git a/tools/validate_sage2_entry_fusion.py b/tools/validate_sage2_entry_fusion.py
new file mode 100644
index 0000000..b19ae7a
--- /dev/null
+++ b/tools/validate_sage2_entry_fusion.py
@@ -0,0 +1,335 @@
+"""Validate exact fused Sage2 entry preparation on randomized and real H3 tensors."""
+
+from __future__ import annotations
+
+import argparse
+import json
+import math
+from pathlib import Path
+
+import torch
+
+from h3_blackwell_runtime.adaln import H3CurveAdaLN
+from h3_blackwell_runtime.attention import rms_norm, rms_rope_split_half_
+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.sage2_entry import prepare_qk
+from h3_blackwell_runtime.sampler import _audio_sigma, _model_sigma, beta_sigmas
+from h3_blackwell_runtime.t2v import random_av_latents
+from profile_attention_path import summarize
+from profile_sage2_scheduler import run_mainloop
+
+
+def difference(actual: torch.Tensor, expected: torch.Tensor) -> dict:
+ delta = actual.float() - expected.float()
+ result = {
+ "equal": torch.equal(actual, expected),
+ "different_elements": int(torch.count_nonzero(actual != expected).item()),
+ "max_abs": delta.abs().max().item() if delta.numel() else 0.0,
+ "mean_abs": delta.abs().mean().item() if delta.numel() else 0.0,
+ }
+ if not result["equal"] and result["different_elements"] <= 8:
+ indices = torch.nonzero(actual != expected, as_tuple=False)
+ result["mismatches"] = [
+ {
+ "index": index.tolist(),
+ "actual": actual[tuple(index)].item(),
+ "expected": expected[tuple(index)].item(),
+ }
+ for index in indices
+ ]
+ return result
+
+
+def measure(fn, warmup: int, iterations: int) -> dict:
+ for _ in range(warmup):
+ fn()
+ torch.cuda.synchronize()
+ before_allocated = torch.cuda.memory_allocated()
+ before_reserved = torch.cuda.memory_reserved()
+ torch.cuda.reset_peak_memory_stats()
+ samples = []
+ for _ in range(iterations):
+ started = torch.cuda.Event(enable_timing=True)
+ finished = torch.cuda.Event(enable_timing=True)
+ started.record()
+ fn()
+ finished.record()
+ finished.synchronize()
+ samples.append(started.elapsed_time(finished) / 1000.0)
+ return {
+ "timing": summarize(samples),
+ "allocated_bytes_before": before_allocated,
+ "reserved_bytes_before": before_reserved,
+ "peak_allocated_bytes": torch.cuda.max_memory_allocated(),
+ "peak_reserved_bytes": torch.cuda.max_memory_reserved(),
+ "peak_allocated_delta_bytes": torch.cuda.max_memory_allocated() - before_allocated,
+ "peak_reserved_delta_bytes": torch.cuda.max_memory_reserved() - before_reserved,
+ }
+
+
+def views(qkv: torch.Tensor, heads: int) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
+ sequence = qkv.shape[0]
+ inner = heads * 128
+ return tuple(part.view(1, sequence, heads, 128) for part in qkv.split(inner, dim=-1))
+
+
+def q_quant(q: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
+ import sageattention.quant as sage_quant
+
+ groups = math.ceil(q.shape[1] / 128) * 4
+ output = torch.empty(q.shape, dtype=torch.int8, device=q.device)
+ scale = torch.empty((q.shape[0], q.shape[2], groups), dtype=torch.float32, device=q.device)
+ sage_quant._fused.quant_per_warp_int8_cuda(q, output, scale, 128, 32, 0)
+ return output, scale
+
+
+def k_quant(k: torch.Tensor, mean: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
+ import sageattention.quant as sage_quant
+
+ output = torch.empty(k.shape, dtype=torch.int8, device=k.device)
+ scale = torch.empty(
+ (k.shape[0], k.shape[2], math.ceil(k.shape[1] / 64)),
+ dtype=torch.float32,
+ device=k.device,
+ )
+ sage_quant._fused.quant_per_block_int8_fuse_sub_mean_cuda(
+ k, mean.squeeze(1), output, scale, 64, 0,
+ )
+ return output, scale
+
+
+def randomized_case(sequence: int, heads: int, seed: int) -> dict:
+ generator = torch.Generator(device="cuda").manual_seed(seed)
+ qkv = torch.randn(
+ (sequence, heads * 128 * 3), generator=generator, device="cuda", dtype=torch.bfloat16,
+ )
+ rotation = torch.randn(
+ (1, sequence, 1, 48, 2, 2), generator=generator, device="cuda", dtype=torch.bfloat16,
+ )
+ q_weight = torch.randn((128,), generator=generator, device="cuda", dtype=torch.bfloat16)
+ k_weight = torch.randn((128,), generator=generator, device="cuda", dtype=torch.bfloat16)
+
+ reference_storage = qkv.clone()
+ q_reference, k_reference, _ = views(reference_storage, heads)
+ rms_rope_split_half_(q_reference, k_reference, rotation, q_weight, k_weight, 1e-5)
+ q_int8_reference, q_scale_reference = q_quant(q_reference)
+
+ candidate_storage = qkv.clone()
+ q_candidate, k_candidate, _ = views(candidate_storage, heads)
+ q_int8_candidate, q_scale_candidate, q_prepared_candidate = prepare_qk(
+ q_candidate,
+ k_candidate,
+ rotation,
+ q_weight,
+ k_weight,
+ 1e-5,
+ materialize_q=True,
+ )
+ return {
+ "sequence": sequence,
+ "heads": heads,
+ "q_prepared": difference(q_prepared_candidate, q_reference),
+ "k_prepared": difference(k_candidate, k_reference),
+ "q_int8": difference(q_int8_candidate, q_int8_reference),
+ "q_scale": difference(q_scale_candidate, q_scale_reference),
+ }
+
+
+def real_case(block: H3DiTBlock, hidden: torch.Tensor, rotation: torch.Tensor, modulation, segments, warmup: int, iterations: int) -> dict:
+ shift_msa, scale_msa, *_ = modulation
+ h_msa = modulate_segments(
+ rms_norm(hidden, block.norm1_weight, block.norm_eps), shift_msa, scale_msa, segments,
+ )
+ raw_qkv = block.attention.qkv_proj(h_msa)
+ heads = block.attention.heads
+
+ reference_storage = raw_qkv.clone()
+ q_reference, k_reference, v_reference = views(reference_storage, heads)
+ rms_rope_split_half_(
+ q_reference,
+ k_reference,
+ rotation,
+ block.attention.q_norm_weight,
+ block.attention.k_norm_weight,
+ block.attention.eps,
+ )
+ q_int8_reference, q_scale_reference = q_quant(q_reference)
+ k_mean_reference = k_reference.mean(dim=1, keepdim=True)
+ k_int8_reference, k_scale_reference = k_quant(k_reference, k_mean_reference)
+
+ candidate_storage = raw_qkv.clone()
+ q_candidate, k_candidate, v_candidate = views(candidate_storage, heads)
+ q_int8_candidate, q_scale_candidate, q_prepared_candidate = prepare_qk(
+ q_candidate,
+ k_candidate,
+ rotation,
+ block.attention.q_norm_weight,
+ block.attention.k_norm_weight,
+ block.attention.eps,
+ materialize_q=True,
+ )
+ k_mean_candidate = k_candidate.mean(dim=1, keepdim=True)
+ k_int8_candidate, k_scale_candidate = k_quant(k_candidate, k_mean_candidate)
+
+ import sageattention.core as sage_core
+
+ v_fp8, v_scale, _ = sage_core.per_channel_fp8(
+ v_reference, tensor_layout="NHD", scale_max=2.25, smooth_v=False,
+ )
+ reference_output = torch.empty_like(q_reference)
+ candidate_output = torch.empty_like(q_reference)
+ run_mainloop(
+ q_int8_reference,
+ k_int8_reference,
+ v_fp8,
+ q_scale_reference,
+ k_scale_reference,
+ v_scale,
+ reference_output,
+ )
+ run_mainloop(
+ q_int8_candidate,
+ k_int8_candidate,
+ v_fp8,
+ q_scale_candidate,
+ k_scale_candidate,
+ v_scale,
+ candidate_output,
+ )
+ torch.cuda.synchronize()
+
+ baseline_timing_storage = raw_qkv.clone()
+ baseline_q, baseline_k, _ = views(baseline_timing_storage, heads)
+
+ def baseline_entry():
+ rms_rope_split_half_(
+ baseline_q,
+ baseline_k,
+ rotation,
+ block.attention.q_norm_weight,
+ block.attention.k_norm_weight,
+ block.attention.eps,
+ )
+ return q_quant(baseline_q)
+
+ candidate_timing_storage = raw_qkv.clone()
+ candidate_q, candidate_k, _ = views(candidate_timing_storage, heads)
+
+ def candidate_entry():
+ return prepare_qk(
+ candidate_q,
+ candidate_k,
+ rotation,
+ block.attention.q_norm_weight,
+ block.attention.k_norm_weight,
+ block.attention.eps,
+ materialize_q=False,
+ )
+
+ return {
+ "block": block.block_index if hasattr(block, "block_index") else None,
+ "q_shape": list(q_reference.shape),
+ "q_stride": list(q_reference.stride()),
+ "q_prepared": difference(q_prepared_candidate, q_reference),
+ "k_prepared": difference(k_candidate, k_reference),
+ "q_int8": difference(q_int8_candidate, q_int8_reference),
+ "q_scale": difference(q_scale_candidate, q_scale_reference),
+ "k_mean": difference(k_mean_candidate, k_mean_reference),
+ "k_int8": difference(k_int8_candidate, k_int8_reference),
+ "k_scale": difference(k_scale_candidate, k_scale_reference),
+ "attention_output": difference(candidate_output, reference_output),
+ "baseline_entry": measure(baseline_entry, warmup, iterations),
+ "candidate_entry": measure(candidate_entry, warmup, iterations),
+ }
+
+
+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("--blocks", nargs="+", type=int, default=(0, 24, 49))
+ parser.add_argument("--randomized-only", action="store_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("--warmup", type=int, default=3)
+ parser.add_argument("--iterations", type=int, default=10)
+ return parser.parse_args()
+
+
+def main() -> None:
+ args = parse_args()
+ randomized = [
+ randomized_case(sequence, 2, args.seed + sequence)
+ for sequence in (1, 31, 32, 33, 127, 128, 129)
+ ]
+ real = []
+ metadata = {}
+ if not args.randomized_only:
+ torch.manual_seed(args.seed)
+ checkpoint = H3Checkpoint(args.model_path, device="cuda")
+ packer = H3PromptPacker(checkpoint)
+ video, audio, aligned_frames = random_av_latents(
+ args.width, args.height, args.frames, args.seed, device="cuda",
+ )
+ sigma = beta_sigmas(args.steps, device="cuda")[args.sampler_step - 1]
+ native_audio = audio.to(torch.bfloat16) * (_audio_sigma(sigma) / sigma)
+ text = torch.randn(1, args.text_tokens, 5376, device="cuda", dtype=torch.bfloat16)
+ hidden, timesteps, segments, positions, _, _ = packer(
+ text, video, native_audio, _model_sigma(sigma),
+ )
+ rotation = h3_rope_rotation(
+ positions.to("cuda"),
+ checkpoint.tensor("rope.inv_freq", dtype=torch.float32),
+ hidden.dtype,
+ )
+ for block_index in args.blocks:
+ block = H3DiTBlock.from_checkpoint(checkpoint, block_index, attention_backend="sage2").eval()
+ block.block_index = block_index
+ adaln = H3CurveAdaLN.from_checkpoint(checkpoint, f"blocks.{block_index}.adaln_proj").eval()
+ with torch.inference_mode():
+ real.append(real_case(
+ block,
+ hidden,
+ rotation,
+ tuple(value.detach() for value in adaln(timesteps)),
+ segments,
+ args.warmup,
+ args.iterations,
+ ))
+ metadata = {
+ "resolution": [args.width, args.height],
+ "frames": aligned_frames,
+ "tokens": hidden.shape[0],
+ "seed": args.seed,
+ "blocks": args.blocks,
+ }
+ report = {
+ "device": torch.cuda.get_device_name(),
+ "torch": torch.__version__,
+ "metadata": metadata,
+ "randomized": randomized,
+ "real": real,
+ }
+ 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)
+ parity_items = randomized + real
+ if any(
+ not value["equal"]
+ for item in parity_items
+ for key, value in item.items()
+ if isinstance(value, dict) and "equal" in value
+ ):
+ raise RuntimeError("Sage2 entry fusion parity failed")
+
+
+if __name__ == "__main__":
+ main()