h3-blackwell-runtime/tools/decode_audio_latent.py
2026-08-14 14:11:36 +07:00

42 lines
1.5 KiB
Python

"""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("latent"))
if latent is None:
raise ValueError("saved state does not contain 'audio_latent' 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)})