76 lines
4.5 KiB
Python
76 lines
4.5 KiB
Python
"""Compare direct FL2VA H3 and sampler boundaries with captured Comfy state."""
|
|
|
|
import glob
|
|
|
|
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.sampler import _audio_sigma, res_multistep_update
|
|
|
|
|
|
root = "/artifacts/fl2va-sampler-reference"
|
|
capture = "/artifacts/capture"
|
|
initial = torch.load(f"{root}/initial.pt", map_location="cuda", weights_only=False)
|
|
steps = [torch.load(path, map_location="cuda", weights_only=False) for path in sorted(glob.glob(f"{root}/step_*.pt"))]
|
|
sigmas = initial["sigmas"].to("cuda")
|
|
model = H3PackedDenoiser.from_checkpoint(
|
|
H3Checkpoint("/models/minimax_h3_fl2va_pruned_nvfp4.safetensors"), attention_backend="sage2"
|
|
).eval()
|
|
video_shape = (1, 24, 7, 12, 20)
|
|
audio_shape = (1, 32, 2, 37)
|
|
video_count = torch.tensor(video_shape).prod().item()
|
|
audio_count = torch.tensor(audio_shape).prod().item()
|
|
old_video = old_audio = old_sigma = None
|
|
|
|
for index, reference in enumerate(steps):
|
|
h3_input = torch.load(f"{capture}/input_{index:02d}.pt", map_location="cuda", weights_only=False)
|
|
h3_output = torch.load(f"{capture}/output_{index:02d}.pt", map_location="cuda", weights_only=False)
|
|
with torch.inference_mode():
|
|
raw_video, _ = model(
|
|
h3_input["hidden"].to("cuda"),
|
|
h3_input["timesteps"].to("cuda"),
|
|
h3_input["position_ids"].to("cuda"),
|
|
h3_input["segments"],
|
|
h3_output["video_segment"],
|
|
h3_output["audio_segment"],
|
|
)
|
|
video_out = unpatchify_video(raw_video, h3_output["video"].shape[2], h3_output["video"].shape[3], h3_output["video"].shape[4])
|
|
h3_delta = (video_out.float() - h3_output["video"].to("cuda").float()).abs()
|
|
|
|
state_video = reference["x"].to("cuda").reshape(-1)[:video_count].reshape(video_shape)
|
|
state_audio = reference["x"].to("cuda").reshape(-1)[video_count:video_count + audio_count].reshape(audio_shape)
|
|
reference_denoised = reference["denoised"].to("cuda").reshape(-1)[:video_count].reshape(video_shape)
|
|
reference_audio_denoised = reference["denoised"].to("cuda").reshape(-1)[video_count:video_count + audio_count].reshape(audio_shape)
|
|
converted_denoised = state_video + sigmas[index] * h3_output["video"].to("cuda").to(torch.bfloat16).float()
|
|
denoised_delta = (converted_denoised.float() - reference_denoised.float()).abs()
|
|
sigma_audio = _audio_sigma(sigmas[index])
|
|
carry = sigma_audio / sigmas[index]
|
|
audio_model_output = (
|
|
(1.0 - 4.0) * (state_audio.to(torch.bfloat16) * carry.to(torch.bfloat16))
|
|
+ (1.0 + 3.0 * sigma_audio).to(torch.bfloat16) * (-h3_output["audio"].to("cuda").to(torch.bfloat16))
|
|
).float()
|
|
converted_audio_denoised = state_audio - sigmas[index] * audio_model_output
|
|
audio_denoised_delta = (converted_audio_denoised.float() - reference_audio_denoised.float()).abs()
|
|
|
|
if index + 1 < len(steps):
|
|
previous_sigma = sigmas[index - 1] if index else None
|
|
updated = res_multistep_update(state_video, reference_denoised, sigmas[index], sigmas[index + 1], old_video, old_sigma, previous_sigma)
|
|
updated_audio = res_multistep_update(state_audio, reference_audio_denoised, sigmas[index], sigmas[index + 1], old_audio, old_sigma, previous_sigma)
|
|
next_video = steps[index + 1]["x"].to("cuda").reshape(-1)[:video_count].reshape(video_shape)
|
|
next_audio = steps[index + 1]["x"].to("cuda").reshape(-1)[video_count:video_count + audio_count].reshape(audio_shape)
|
|
update_delta = (updated.float() - next_video.float()).abs()
|
|
update_audio_delta = (updated_audio.float() - next_audio.float()).abs()
|
|
update_text = f"update_mean={update_delta.mean().item():.6g} update_max={update_delta.max().item():.6g} audio_update_mean={update_audio_delta.mean().item():.6g} audio_update_max={update_audio_delta.max().item():.6g}"
|
|
else:
|
|
update_text = "update_mean=final-unobserved update_max=final-unobserved audio_update_mean=final-unobserved audio_update_max=final-unobserved"
|
|
|
|
print(
|
|
f"step={index:02d} "
|
|
f"h3_mean={h3_delta.mean().item():.6g} h3_max={h3_delta.max().item():.6g} "
|
|
f"denoised_mean={denoised_delta.mean().item():.6g} denoised_max={denoised_delta.max().item():.6g} "
|
|
f"audio_denoised_mean={audio_denoised_delta.mean().item():.6g} audio_denoised_max={audio_denoised_delta.max().item():.6g} "
|
|
f"{update_text}"
|
|
)
|
|
old_video, old_audio, old_sigma = reference_denoised, reference_audio_denoised, sigmas[index + 1]
|