Allow captured H3 timesteps in preview

This commit is contained in:
Daniel Maddern 2026-08-13 22:23:07 +07:00
parent 9d55aef8a6
commit d289b3c497

View file

@ -28,6 +28,7 @@ parser.add_argument("--frames", type=int, default=22)
parser.add_argument("--steps", type=int, default=12)
parser.add_argument("--seed", type=int, default=440204)
parser.add_argument("--attention", choices=("sage2", "sdpa", "sage3"), default="sage2")
parser.add_argument("--model-timesteps-capture", type=Path, help="Directory containing captured input_XX.pt H3 timesteps for strict parity checks.")
args = parser.parse_args()
checkpoint = H3Checkpoint("/models/minimax_h3_fl2va_pruned_nvfp4.safetensors")
@ -38,7 +39,13 @@ conditioner = Qwen3VLPromptConditioner(
video, audio, frames = random_av_latents(args.width, args.height, args.frames, args.seed)
model = H3PackedDenoiser.from_checkpoint(checkpoint, attention_backend=args.attention).eval()
text = H3TokenRefiner(checkpoint, attention_backend=args.attention)(conditioner(args.prompt))
latent = sample_video_res_multistep(model, H3PromptPacker(checkpoint), text, video, audio, steps=args.steps)
model_timesteps = None
if args.model_timesteps_capture is not None:
model_timesteps = [
torch.load(args.model_timesteps_capture / f"input_{index:02d}.pt", map_location="cuda", weights_only=False)["timesteps"]
for index in range(args.steps)
]
latent = sample_video_res_multistep(model, H3PromptPacker(checkpoint), text, video, audio, steps=args.steps, model_timesteps=model_timesteps)
vae = MiniMaxH3VideoVAE.from_safetensors("/vae/minimax_h3_video_vae_fp16.safetensors", device="cuda").eval()
pixels = vae.decode(latent.to(next(vae.parameters()).dtype))[:, :, :frames]
pixels = ((pixels[0].permute(1, 2, 3, 0).clamp(-1, 1) + 1) * 127.5).to(torch.uint8).cpu()