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()