141 lines
5.7 KiB
Python
141 lines
5.7 KiB
Python
|
|
"""Validate fused H3 modulation and native NVFP4 production."""
|
||
|
|
|
||
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
import argparse
|
||
|
|
import json
|
||
|
|
import time
|
||
|
|
from pathlib import Path
|
||
|
|
|
||
|
|
import torch
|
||
|
|
|
||
|
|
from h3_blackwell_runtime.adaln import H3CurveAdaLN
|
||
|
|
from h3_blackwell_runtime.attention import rms_norm
|
||
|
|
from h3_blackwell_runtime.block import H3DiTBlock, modulate_segments
|
||
|
|
from h3_blackwell_runtime.checkpoint import H3Checkpoint
|
||
|
|
from h3_blackwell_runtime.h3_fusion import segment_index
|
||
|
|
from h3_blackwell_runtime.nvfp4_quant import (
|
||
|
|
nvfp4_activation_scale,
|
||
|
|
vortex_native_quantize_modulated_nvfp4,
|
||
|
|
)
|
||
|
|
from h3_blackwell_runtime.packing import H3PromptPacker
|
||
|
|
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("--model-path", default="/models/minimax_h3_fl2va_pruned_nvfp4.safetensors")
|
||
|
|
parser.add_argument("--output", type=Path, required=True)
|
||
|
|
parser.add_argument("--block-index", type=int, default=24)
|
||
|
|
parser.add_argument("--rows", type=int, default=2048)
|
||
|
|
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)
|
||
|
|
parser.add_argument("--device", default="cuda")
|
||
|
|
return parser.parse_args()
|
||
|
|
|
||
|
|
|
||
|
|
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="sage2").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,
|
||
|
|
)
|
||
|
|
sigma = beta_sigmas(args.steps, device=args.device)[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, _, _, _ = packer(
|
||
|
|
text, video, native_audio, _model_sigma(sigma),
|
||
|
|
)
|
||
|
|
rows = min(args.rows, hidden.shape[0])
|
||
|
|
hidden = hidden[:rows].contiguous()
|
||
|
|
clipped_segments = []
|
||
|
|
for start, stop, table_row in segments:
|
||
|
|
if start >= rows:
|
||
|
|
break
|
||
|
|
clipped_segments.append((start, min(stop, rows), table_row))
|
||
|
|
shift_msa, scale_msa, *_ = (value.detach() for value in adaln(timesteps))
|
||
|
|
|
||
|
|
with torch.inference_mode():
|
||
|
|
normalized = rms_norm(hidden, block.norm1_weight, block.norm_eps)
|
||
|
|
materialized = modulate_segments(
|
||
|
|
normalized, shift_msa, scale_msa, clipped_segments,
|
||
|
|
)
|
||
|
|
reference_scale = nvfp4_activation_scale(materialized).float()
|
||
|
|
import comfy_kitchen as ck
|
||
|
|
from comfy_kitchen.tensor import TensorCoreNVFP4Layout
|
||
|
|
|
||
|
|
reference_qdata, reference_sfa = ck.quantize_nvfp4(
|
||
|
|
materialized,
|
||
|
|
reference_scale,
|
||
|
|
pad_16x=TensorCoreNVFP4Layout.get_padded_shape(tuple(materialized.shape))
|
||
|
|
!= tuple(materialized.shape),
|
||
|
|
)
|
||
|
|
fused_scale, fused_qdata, fused_sfa = vortex_native_quantize_modulated_nvfp4(
|
||
|
|
normalized,
|
||
|
|
shift_msa.contiguous(),
|
||
|
|
scale_msa.contiguous(),
|
||
|
|
segment_index(rows, clipped_segments, hidden.device),
|
||
|
|
)
|
||
|
|
torch.cuda.synchronize()
|
||
|
|
|
||
|
|
def run_fused():
|
||
|
|
return vortex_native_quantize_modulated_nvfp4(
|
||
|
|
normalized,
|
||
|
|
shift_msa.contiguous(),
|
||
|
|
scale_msa.contiguous(),
|
||
|
|
segment_index(rows, clipped_segments, hidden.device),
|
||
|
|
)
|
||
|
|
|
||
|
|
for _ in range(args.warmup):
|
||
|
|
run_fused()
|
||
|
|
fused_times = []
|
||
|
|
for _ in range(args.iterations):
|
||
|
|
torch.cuda.synchronize()
|
||
|
|
started = time.perf_counter()
|
||
|
|
run_fused()
|
||
|
|
torch.cuda.synchronize()
|
||
|
|
fused_times.append(time.perf_counter() - started)
|
||
|
|
|
||
|
|
report = {
|
||
|
|
"device": torch.cuda.get_device_name(),
|
||
|
|
"block_index": args.block_index,
|
||
|
|
"rows": rows,
|
||
|
|
"width": hidden.shape[1],
|
||
|
|
"scale_equal": torch.equal(fused_scale, reference_scale),
|
||
|
|
"scale_reference": reference_scale.item(),
|
||
|
|
"scale_fused": fused_scale.item(),
|
||
|
|
"qdata_equal": torch.equal(fused_qdata, reference_qdata),
|
||
|
|
"qdata_differences": torch.count_nonzero(fused_qdata != reference_qdata).item(),
|
||
|
|
"sfa_equal": torch.equal(fused_sfa.view(torch.uint8), reference_sfa.view(torch.uint8)),
|
||
|
|
"sfa_differences": torch.count_nonzero(
|
||
|
|
fused_sfa.view(torch.uint8) != reference_sfa.view(torch.uint8)
|
||
|
|
).item(),
|
||
|
|
"fused_producer_p50_ms": sorted(fused_times)[len(fused_times) // 2] * 1000.0,
|
||
|
|
}
|
||
|
|
report["equal"] = report["scale_equal"] and report["qdata_equal"] and report["sfa_equal"]
|
||
|
|
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 report["equal"]:
|
||
|
|
raise RuntimeError("fused modulation producer is not byte-exact")
|
||
|
|
|
||
|
|
|
||
|
|
if __name__ == "__main__":
|
||
|
|
main()
|