Spaces:
Runtime error
Runtime error
File size: 4,796 Bytes
982899c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 | import os
from typing import Optional
from pathlib import Path
import torch
from videox_fun.models.creator.dac_vae import DAC, DiagonalGaussianDistribution
class CreatorDACVAE(torch.nn.Module):
"""
High-level wrapper around the DAC (Descript Audio Codec) VAE in continuous mode.
Mirrors the LTXAudioVAE interface used by the inference pipeline.
"""
def __init__(self, dac_model: DAC) -> None:
super().__init__()
self.dac = dac_model
@property
def sample_rate(self) -> int:
return self.dac.sample_rate
@property
def hop_length(self) -> int:
return self.dac.hop_length
@property
def latent_dim(self) -> int:
return self.dac.latent_dim
@classmethod
def from_pretrained(
cls,
pretrained_model_path: str | os.PathLike[str],
strict: bool = False,
) -> "CreatorDACVAE":
pretrained_model_path = Path(pretrained_model_path)
if pretrained_model_path.is_dir():
dac_model = DAC.from_pretrained(pretrained_model_path)
else:
dac_model = DAC.from_pretrained(pretrained_model_path.parent)
return cls(dac_model=dac_model)
def _preprocess_waveform(self, waveform: torch.Tensor) -> torch.Tensor:
"""Normalize waveform to [B, 1, T] mono and pad to hop_length boundary."""
"""The Creator audio VAE currently supports mono audio."""
if waveform.ndim == 1:
waveform = waveform.unsqueeze(0).unsqueeze(0)
elif waveform.ndim == 2:
if waveform.size(0) > 1:
waveform = waveform.mean(dim=0, keepdim=True)
waveform = waveform.unsqueeze(0)
elif waveform.ndim == 3:
if waveform.size(1) > 1:
waveform = waveform.mean(dim=1, keepdim=True)
waveform = self.dac.preprocess(waveform, self.sample_rate)
return waveform
@torch.inference_mode()
def encode_posterior(
self,
audio: torch.Tensor,
sampling_rate: Optional[int] = None,
deterministic: bool | None = None,
) -> DiagonalGaussianDistribution:
"""Encode audio waveform and return the posterior distribution.
Parameters
----------
audio : Tensor
Raw waveform tensor. Accepts shapes [T], [C, T], or [B, C, T].
sampling_rate : int, optional
Not used directly; kept for API compatibility with LTXAudioVAE.
deterministic : bool, optional
If True, std is zeroed so sampling returns the mean.
Returns
-------
DiagonalGaussianDistribution
"""
waveform = self._preprocess_waveform(audio)
posterior, _, _, _, _ = self.dac.encode(waveform)
if deterministic is not None:
posterior.deterministic = deterministic
if deterministic:
posterior.std = torch.zeros_like(posterior.mean)
posterior.var = torch.zeros_like(posterior.mean)
return posterior
@torch.inference_mode()
def encode(
self,
audio: torch.Tensor,
sampling_rate: Optional[int] = None,
sample: bool = False,
generator: Optional[torch.Generator] = None,
) -> torch.Tensor:
"""Encode audio waveform to latent representation.
Parameters
----------
audio : Tensor
Raw waveform tensor.
sampling_rate : int, optional
Kept for API compatibility.
sample : bool
If True, sample from the posterior; otherwise return the mean.
generator : torch.Generator, optional
RNG for reproducible sampling.
Returns
-------
Tensor [B, D, T']
Continuous latent codes.
"""
posterior = self.encode_posterior(audio, sampling_rate=sampling_rate)
if sample:
return posterior.sample()
return posterior.mode()
@torch.inference_mode()
def decode(self, latent: torch.Tensor) -> torch.Tensor:
"""Decode latent codes back to waveform.
Parameters
----------
latent : Tensor [B, D, T']
Continuous latent codes.
Returns
-------
Tensor [B, 1, T]
Reconstructed waveform.
"""
return self.dac.decode(latent)
@torch.inference_mode()
def reconstruct(
self,
audio: torch.Tensor,
sampling_rate: Optional[int] = None,
sample: bool = False,
generator: Optional[torch.Generator] = None,
) -> torch.Tensor:
"""Encode then decode (round-trip reconstruction)."""
latent = self.encode(audio, sampling_rate=sampling_rate, sample=sample, generator=generator)
return self.decode(latent)
|