Allow captured H3 timesteps in preview
This commit is contained in:
parent
9d55aef8a6
commit
d289b3c497
1 changed files with 8 additions and 1 deletions
|
|
@ -28,6 +28,7 @@ parser.add_argument("--frames", type=int, default=22)
|
||||||
parser.add_argument("--steps", type=int, default=12)
|
parser.add_argument("--steps", type=int, default=12)
|
||||||
parser.add_argument("--seed", type=int, default=440204)
|
parser.add_argument("--seed", type=int, default=440204)
|
||||||
parser.add_argument("--attention", choices=("sage2", "sdpa", "sage3"), default="sage2")
|
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()
|
args = parser.parse_args()
|
||||||
|
|
||||||
checkpoint = H3Checkpoint("/models/minimax_h3_fl2va_pruned_nvfp4.safetensors")
|
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)
|
video, audio, frames = random_av_latents(args.width, args.height, args.frames, args.seed)
|
||||||
model = H3PackedDenoiser.from_checkpoint(checkpoint, attention_backend=args.attention).eval()
|
model = H3PackedDenoiser.from_checkpoint(checkpoint, attention_backend=args.attention).eval()
|
||||||
text = H3TokenRefiner(checkpoint, attention_backend=args.attention)(conditioner(args.prompt))
|
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()
|
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 = 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()
|
pixels = ((pixels[0].permute(1, 2, 3, 0).clamp(-1, 1) + 1) * 127.5).to(torch.uint8).cpu()
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue