"""Replay a matched Comfy keyframe DiT capture through the direct H3 model.""" import argparse from pathlib import Path import torch from h3_blackwell_runtime.checkpoint import H3Checkpoint from h3_blackwell_runtime.denoiser import H3PackedDenoiser from h3_blackwell_runtime.packing import unpatchify_video from h3_blackwell_runtime.rope import h3_rope_rotation from h3_blackwell_runtime.sampler import _unpack_audio parser = argparse.ArgumentParser() parser.add_argument("--capture", type=Path, required=True) parser.add_argument("--model", default="/models/minimax_h3_fl2va_pruned_nvfp4.safetensors") parser.add_argument("--attention", default="sage2") args = parser.parse_args() def report(name, actual, expected): expected = expected.to(actual.device) delta = (actual.float() - expected.float()).abs() print({"stage": name, "shape": tuple(actual.shape), "mean_delta": float(delta.mean()), "max_delta": float(delta.max())}, flush=True) captured_input = torch.load(args.capture / "input.pt", map_location="cuda", weights_only=False) captured_output = torch.load(args.capture / "output.pt", map_location="cuda", weights_only=False) checkpoint = H3Checkpoint(args.model) model = H3PackedDenoiser.from_checkpoint(checkpoint, attention_backend=args.attention).eval() hidden = captured_input["hidden"].cuda() timesteps = captured_input["timesteps"].cuda() positions = captured_input["position_ids"].cuda() segments = captured_input["segments"] rotation = h3_rope_rotation(positions, model.backbone.inv_freq, hidden.dtype) with torch.inference_mode(): for index, (block, adaln) in enumerate(zip(model.backbone.blocks, model.backbone.adaln, strict=True)): hidden = block(hidden, rotation, *adaln(timesteps), segments) expected = torch.load(args.capture / "blocks" / f"{index:02d}.pt", map_location="cuda", weights_only=False) report(f"block_{index:02d}", hidden, expected) video_rows, audio_rows = model.final_layer( hidden, timesteps, tuple(captured_output["video_segment"]), tuple(captured_output["audio_segment"]), ) expected_video = captured_output["video"].cuda() expected_audio = captured_output["audio"].cuda() video = unpatchify_video(video_rows, expected_video.shape[2], expected_video.shape[3], expected_video.shape[4]) audio = _unpack_audio(audio_rows) report("video_output", video, expected_video) report("audio_output", audio, expected_audio)