"""Decode a saved H3 audio latent to WAV.""" import argparse import subprocess from pathlib import Path import torch from h3_blackwell_runtime.audio_vae_decoder import MiniMaxH3AudioVAE parser = argparse.ArgumentParser() parser.add_argument("--latent", type=Path, required=True) parser.add_argument("--output", type=Path, required=True) parser.add_argument("--ffmpeg-loglevel", default="error") parser.add_argument("--quiet", action="store_true") args = parser.parse_args() state = torch.load(args.latent, map_location="cuda", weights_only=False) if isinstance(state, dict): latent = state.get("audio_latent", state.get("final_audio", state.get("latent"))) if latent is None: raise ValueError("saved state does not contain 'audio_latent', 'final_audio', or 'latent'") else: latent = state latent = latent.to("cuda") vae = MiniMaxH3AudioVAE.from_safetensors("/vae/minimax_h3_audio_vae_fp32.safetensors", device="cuda").eval() with torch.inference_mode(): waveform = vae.decode(latent.to(next(vae.parameters()).dtype)).clamp(-1, 1).cpu()[0] args.output.parent.mkdir(parents=True, exist_ok=True) raw = args.output.with_suffix(".f32le") waveform.transpose(0, 1).contiguous().numpy().tofile(raw) subprocess.run([ "ffmpeg", "-hide_banner", "-loglevel", args.ffmpeg_loglevel, "-y", "-f", "f32le", "-ar", "32000", "-ac", "2", "-i", str(raw), str(args.output), ], check=True) raw.unlink() if not args.quiet: print({"output": str(args.output), "sample_rate": 32000, "shape": tuple(waveform.shape)})