2026-08-14 13:39:46 +07:00
|
|
|
"""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)
|
2026-08-14 14:11:36 +07:00
|
|
|
parser.add_argument("--ffmpeg-loglevel", default="error")
|
|
|
|
|
parser.add_argument("--quiet", action="store_true")
|
2026-08-14 13:39:46 +07:00
|
|
|
args = parser.parse_args()
|
|
|
|
|
|
|
|
|
|
state = torch.load(args.latent, map_location="cuda", weights_only=False)
|
|
|
|
|
if isinstance(state, dict):
|
2026-08-22 14:09:45 +07:00
|
|
|
latent = state.get("audio_latent", state.get("final_audio", state.get("latent")))
|
2026-08-14 13:39:46 +07:00
|
|
|
if latent is None:
|
2026-08-22 14:09:45 +07:00
|
|
|
raise ValueError("saved state does not contain 'audio_latent', 'final_audio', or 'latent'")
|
2026-08-14 13:39:46 +07:00
|
|
|
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([
|
2026-08-14 14:11:36 +07:00
|
|
|
"ffmpeg", "-hide_banner", "-loglevel", args.ffmpeg_loglevel,
|
|
|
|
|
"-y", "-f", "f32le", "-ar", "32000", "-ac", "2",
|
2026-08-14 13:39:46 +07:00
|
|
|
"-i", str(raw), str(args.output),
|
|
|
|
|
], check=True)
|
|
|
|
|
raw.unlink()
|
2026-08-14 14:11:36 +07:00
|
|
|
if not args.quiet:
|
|
|
|
|
print({"output": str(args.output), "sample_rate": 32000, "shape": tuple(waveform.shape)})
|