Compare H3 sampler pre-step states
This commit is contained in:
parent
20e18d9d0c
commit
416a7d9697
1 changed files with 10 additions and 6 deletions
|
|
@ -21,15 +21,15 @@ text = H3TokenRefiner(checkpoint)(Qwen3VLPromptConditioner("/text-encoders/qwen3
|
|||
model = H3PackedDenoiser.from_checkpoint(checkpoint, attention_backend="sage2").eval()
|
||||
packer = H3PromptPacker(checkpoint)
|
||||
|
||||
packed_initial = initial["initial_x"].to("cuda").reshape(-1)
|
||||
video_shape = (1, 24, 7, 12, 20)
|
||||
audio_shape = (1, 32, 2, 37)
|
||||
video_count = torch.tensor(video_shape).prod().item()
|
||||
video = packed_initial[:video_count].reshape(video_shape)
|
||||
audio_carried = packed_initial[video_count:].reshape(audio_shape)
|
||||
old_video = old_audio = old_sigma = None
|
||||
for index, reference in enumerate(steps):
|
||||
sigma, sigma_down = sigmas[index], sigmas[index + 1]
|
||||
packed_state = reference["x"].to("cuda").reshape(-1)
|
||||
video = packed_state[:video_count].reshape(video_shape)
|
||||
audio_carried = packed_state[video_count:].reshape(audio_shape)
|
||||
native_audio = audio_carried * (_audio_sigma(sigma) / sigma)
|
||||
hidden, times, segments, positions, video_segment, audio_segment = packer(text, video, native_audio, float(sigma))
|
||||
raw_video, raw_audio = model(hidden, times, positions, segments, video_segment, audio_segment)
|
||||
|
|
@ -44,7 +44,11 @@ for index, reference in enumerate(steps):
|
|||
previous_sigma = sigmas[index - 1] if index else None
|
||||
video = res_multistep_update(video, denoised[0], sigma, sigma_down, old_video, old_sigma, previous_sigma)
|
||||
audio_carried = res_multistep_update(audio_carried, denoised[1], sigma, sigma_down, old_audio, old_sigma, previous_sigma)
|
||||
reference_latent = reference["x"].to("cuda").reshape(-1)[:video_count].reshape(video_shape)
|
||||
latent_delta = (video.float() - reference_latent.float()).abs()
|
||||
print(f"step={index:02d} x0_video_mean={denoised_delta.mean().item():.6g} x0_video_max={denoised_delta.max().item():.6g} latent_video_mean={latent_delta.mean().item():.6g} latent_video_max={latent_delta.max().item():.6g}")
|
||||
if index + 1 < len(steps):
|
||||
reference_latent = steps[index + 1]["x"].to("cuda").reshape(-1)[:video_count].reshape(video_shape)
|
||||
latent_delta = (video.float() - reference_latent.float()).abs()
|
||||
latent_text = f"latent_video_mean={latent_delta.mean().item():.6g} latent_video_max={latent_delta.max().item():.6g}"
|
||||
else:
|
||||
latent_text = "latent_video_mean=final-unobserved latent_video_max=final-unobserved"
|
||||
print(f"step={index:02d} x0_video_mean={denoised_delta.mean().item():.6g} x0_video_max={denoised_delta.max().item():.6g} {latent_text}")
|
||||
old_video, old_audio, old_sigma = denoised[0], denoised[1], sigma_down
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue