Spaces:
Paused
Paused
| import torch | |
| from .base import AudioTokenizer | |
| class EnCodecAudioCodec(AudioTokenizer): | |
| def __init__(self, device: str = "cpu", bandwidth: float = 6.0): | |
| self.device = torch.device(device) | |
| self._bandwidth = bandwidth | |
| self._model = None | |
| self._loaded = False | |
| def _lazy_load(self): | |
| if self._loaded: | |
| return | |
| try: | |
| from encodec import EncodecModel | |
| self._model = EncodecModel.encodec_model_24khz() | |
| self._model.set_target_bandwidth(self._bandwidth) | |
| self._model.to(self.device) | |
| self._model.eval() | |
| self._loaded = True | |
| except ImportError: | |
| raise ImportError("encodec package not installed; run: pip install encodec") | |
| except Exception as e: | |
| raise RuntimeError(f"failed to load EnCodec: {e}") | |
| def encode(self, audio: torch.Tensor) -> torch.Tensor: | |
| self._lazy_load() | |
| audio = audio.to(self.device) | |
| if audio.dim() == 1: | |
| audio = audio.unsqueeze(0).unsqueeze(0) | |
| elif audio.dim() == 2: | |
| audio = audio.unsqueeze(1) | |
| from encodec.utils import convert_audio | |
| audio = convert_audio(audio, self._model.sample_rate, self._model.sample_rate, self._model.channels) | |
| frames = self._model.encode(audio) | |
| codes = torch.cat([f[0] for f in frames], dim=-1) | |
| return codes | |
| def decode(self, tokens: torch.Tensor) -> torch.Tensor: | |
| self._lazy_load() | |
| if tokens.dim() == 2: | |
| tokens = tokens.unsqueeze(0) | |
| frames = self._model.decode([(tokens, None)]) | |
| return frames[0] | |
| def sample_rate(self) -> int: | |
| return 24000 | |
| def num_codebooks(self) -> int: | |
| return 8 | |
| def name(self) -> str: | |
| return f"encodec-{self._bandwidth}kbps" | |
| class DummyAudioCodec(AudioTokenizer): | |
| def __init__(self, device: str = "cpu"): | |
| self.device = torch.device(device) | |
| def encode(self, audio: torch.Tensor) -> torch.Tensor: | |
| B = audio.shape[0] if audio.dim() > 1 else 1 | |
| return torch.randint(0, 2048, (4, B * 10), dtype=torch.int32, device=self.device) | |
| def decode(self, tokens: torch.Tensor) -> torch.Tensor: | |
| length = tokens.shape[-1] * 320 | |
| return torch.randn(1, length, device=self.device) | |
| def sample_rate(self) -> int: | |
| return 32000 | |
| def num_codebooks(self) -> int: | |
| return 4 | |
| def name(self) -> str: | |
| return "dummy" | |