robovoice-converter / audio_codecs.py
ChristophSchuhmann's picture
RoboVoice converter: full pipeline + CNN translator + samples
6110c17 verified
Raw
History Blame Contribute Delete
2.36 kB
"""Reusable loaders for DACVAE (target) and DramaBox/LTX audio VAE (input).
Both codecs operate on an identical grid: 25 Hz, 128-dim latents, same length.
DACVAE: 48 kHz mono. DramaBox encoder: stereo, resamples 48k->16k internally.
"""
import os, sys, torch, numpy as np
HERE = os.path.dirname(os.path.abspath(__file__))
sys.path.insert(0, os.path.join(HERE, "fast-dacvae"))
sys.path.insert(0, os.path.join(HERE, "DramaBox", "ltx2"))
sys.path.insert(0, os.path.join(HERE, "DramaBox", "src"))
DRAMABOX_CKPT = os.path.join(HERE, "weights", "dramabox-audio-components.safetensors")
def load_dacvae(device="cuda"):
from dacvae.model.dacvae import DACVAE
from huggingface_hub import hf_hub_download
sd = torch.load(hf_hub_download("facebook/dacvae-watermarked", "weights.pth"),
map_location="cpu", weights_only=True)
if "state_dict" in sd:
sd = sd["state_dict"]
m = DACVAE(encoder_dim=64, encoder_rates=[2, 8, 10, 12], decoder_rates=[12, 10, 8, 2],
latent_dim=1024, codebook_dim=128, sample_rate=48000)
miss, unexp = m.load_state_dict(sd, strict=False)
assert not miss and not unexp, (len(miss), len(unexp))
return m.to(device).eval()
class DramaBoxEncoder:
"""Wraps the LTX audio VAE encoder. encode(wav48k_mono) -> (T,128) latent."""
def __init__(self, device="cuda", dtype=torch.bfloat16):
from ltx_pipelines.utils.blocks import AudioConditioner
from ltx_core.model.audio_vae import encode_audio
from ltx_core.types import Audio
from ltx_core.components.patchifiers import AudioPatchifier
self._Audio, self._encode_audio = Audio, encode_audio
self._patch = AudioPatchifier(patch_size=1)
self.device, self.dtype = device, dtype
self._ac = AudioConditioner(checkpoint_path=DRAMABOX_CKPT, dtype=dtype, device=device, warm=True)
@torch.no_grad()
def encode(self, wav_mono):
"""wav_mono: 1-D float tensor on device, 48 kHz. Returns (T,128) float tensor."""
wt = wav_mono.view(1, 1, -1).repeat(1, 2, 1) # mono -> stereo
audio = self._Audio(waveform=wt, sampling_rate=48000)
lat = self._ac(lambda enc: self._encode_audio(audio, enc, None)) # (1,8,T,16)
pl = self._patch.patchify(lat) # (1,T,128)
return pl[0].float()