Nymbo commited on
Commit
9f153db
·
verified ·
1 Parent(s): 2050b22

ZeroGPU port: Gradio 6.29, python 3.12, torch 2.9.1, vendored minimal audiocraft EnCodec, soundfile IO, models loaded at import, literal_eval instead of eval

Browse files
README.md CHANGED
@@ -4,7 +4,8 @@ emoji: 📈
4
  colorFrom: blue
5
  colorTo: red
6
  sdk: gradio
7
- sdk_version: 4.25.0
 
8
  app_file: app.py
9
  pinned: false
10
  license: cc-by-nc-sa-4.0
 
4
  colorFrom: blue
5
  colorTo: red
6
  sdk: gradio
7
+ sdk_version: 6.29.0
8
+ python_version: "3.12"
9
  app_file: app.py
10
  pinned: false
11
  license: cc-by-nc-sa-4.0
app.py CHANGED
@@ -1,78 +1,105 @@
 
 
1
  import os
2
- # os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID"
3
- # os.environ["CUDA_VISIBLE_DEVICES"] = "1" # these are only used if developping locally
4
  import gradio as gr
 
 
5
  import torch
6
- import torchaudio
 
7
  from data.tokenizer import (
8
  AudioTokenizer,
9
  TextTokenizer,
 
10
  )
11
  from models import voicecraft
12
- import io
13
- import numpy as np
14
- import random
15
- import spaces
16
 
 
 
 
 
17
 
18
- whisper_model, voicecraft_model = None, None
19
 
20
- @spaces.GPU(duration=30)
21
- def seed_everything(seed):
22
- if seed != -1:
23
- os.environ['PYTHONHASHSEED'] = str(seed)
24
- random.seed(seed)
25
- np.random.seed(seed)
26
- torch.manual_seed(seed)
27
- torch.cuda.manual_seed(seed)
28
- torch.backends.cudnn.benchmark = False
29
- torch.backends.cudnn.deterministic = True
30
 
31
- @spaces.GPU(duration=120)
32
- def load_models(whisper_model_choice, voicecraft_model_choice):
33
- global whisper_model, voicecraft_model
34
 
35
- if whisper_model_choice is not None:
36
- import whisper
37
- from whisper.tokenizer import get_tokenizer
38
- whisper_model = {
39
- "model": whisper.load_model(whisper_model_choice),
40
- "tokenizer": get_tokenizer(multilingual=False)
41
- }
42
 
43
 
44
- device = "cuda" if torch.cuda.is_available() else "cpu"
45
-
46
- voicecraft_name = f"{voicecraft_model_choice}.pth"
47
- ckpt_fn = f"./pretrained_models/{voicecraft_name}"
48
- encodec_fn = "./pretrained_models/encodec_4cb2048_giga.th"
49
- if not os.path.exists(ckpt_fn):
50
- os.system(f"wget https://huggingface.co/pyp1/VoiceCraft/resolve/main/{voicecraft_name}\?download\=true")
51
- os.system(f"mv {voicecraft_name}\?download\=true ./pretrained_models/{voicecraft_name}")
52
- if not os.path.exists(encodec_fn):
53
- os.system(f"wget https://huggingface.co/pyp1/VoiceCraft/resolve/main/encodec_4cb2048_giga.th")
54
- os.system(f"mv encodec_4cb2048_giga.th ./pretrained_models/encodec_4cb2048_giga.th")
55
-
56
- ckpt = torch.load(ckpt_fn, map_location="cpu")
 
57
  model = voicecraft.VoiceCraft(ckpt["config"])
58
  model.load_state_dict(ckpt["model"])
59
  model.to(device)
60
  model.eval()
61
- voicecraft_model = {
 
 
 
62
  "ckpt": ckpt,
63
  "model": model,
64
  "text_tokenizer": TextTokenizer(backend="espeak"),
65
- "audio_tokenizer": AudioTokenizer(signature=encodec_fn)
66
  }
67
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
68
  return gr.Accordion()
69
 
 
70
  @spaces.GPU(duration=60)
71
  def transcribe(seed, audio_path):
72
  if whisper_model is None:
73
  raise gr.Error("Whisper model not loaded")
74
  seed_everything(seed)
75
-
 
76
  number_tokens = [
77
  i
78
  for i in range(whisper_model["tokenizer"].eot)
@@ -80,7 +107,7 @@ def transcribe(seed, audio_path):
80
  ]
81
  result = whisper_model["model"].transcribe(audio_path, suppress_tokens=[-1] + number_tokens, word_timestamps=True)
82
  words = [word_info for segment in result["segments"] for word_info in segment["words"]]
83
-
84
  transcript = result["text"]
85
  transcript_with_start_time = " ".join([f"{word['start']} {word['word']}" for word in words])
86
  transcript_with_end_time = " ".join([f"{word['word']} {word['end']}" for word in words])
@@ -98,10 +125,8 @@ def transcribe(seed, audio_path):
98
 
99
  def get_output_audio(audio_tensors, codec_audio_sr):
100
  result = torch.cat(audio_tensors, 1)
101
- buffer = io.BytesIO()
102
- torchaudio.save(buffer, result, int(codec_audio_sr), format="wav")
103
- buffer.seek(0)
104
- return buffer.read()
105
 
106
  @spaces.GPU(duration=90)
107
  def run(seed, left_margin, right_margin, codec_audio_sr, codec_sr, top_k, top_p, temperature,
@@ -128,9 +153,11 @@ def run(seed, left_margin, right_margin, codec_audio_sr, codec_sr, top_k, top_p,
128
  else:
129
  sentences = [transcript.replace("\n", " ")]
130
 
131
- device = "cuda" if torch.cuda.is_available() else "cpu"
132
- info = torchaudio.info(audio_path)
133
- audio_dur = info.num_frames / info.sample_rate
 
 
134
 
135
  audio_tensors = []
136
  inference_transcript = ""
@@ -159,7 +186,7 @@ def run(seed, left_margin, right_margin, codec_audio_sr, codec_sr, top_k, top_p,
159
 
160
  inference_transcript += target_transcript + "\n"
161
 
162
- prompt_end_frame = int(min(audio_dur, prompt_end_time) * info.sample_rate)
163
  _, gen_audio = inference_one_sample(voicecraft_model["model"],
164
  voicecraft_model["ckpt"]["config"],
165
  voicecraft_model["ckpt"]["phn2num"],
@@ -213,8 +240,8 @@ def update_input_audio(audio_path):
213
  if audio_path is None:
214
  return 0, 0, 0
215
 
216
- info = torchaudio.info(audio_path)
217
- max_time = round(info.num_frames / info.sample_rate, 2)
218
  return [
219
  gr.Slider(maximum=max_time, value=max_time),
220
  gr.Slider(maximum=max_time, value=0),
@@ -223,7 +250,6 @@ def update_input_audio(audio_path):
223
 
224
 
225
  def change_mode(mode):
226
- tts_mode_controls, edit_mode_controls, edit_word_mode, split_text, long_tts_sentence_editor
227
  return [
228
  gr.Group(visible=mode != "Edit"),
229
  gr.Group(visible=mode == "Edit"),
@@ -377,7 +403,7 @@ with gr.Blocks() as app:
377
  with gr.Row():
378
  voicecraft_model_choice = gr.Radio(label="VoiceCraft model", value="giga830M", choices=["giga330M", "giga830M"])
379
  whisper_model_choice = gr.Radio(label="Whisper model", value="base.en",
380
- choices=[None, "tiny.en", "base.en", "small.en", "medium.en", "large"])
381
 
382
  with gr.Row():
383
  with gr.Column(scale=2):
 
1
+ import spaces # must be imported before torch on ZeroGPU
2
+
3
  import os
4
+ import random
5
+
6
  import gradio as gr
7
+ import nltk
8
+ import numpy as np
9
  import torch
10
+ from huggingface_hub import hf_hub_download
11
+
12
  from data.tokenizer import (
13
  AudioTokenizer,
14
  TextTokenizer,
15
+ audio_info,
16
  )
17
  from models import voicecraft
 
 
 
 
18
 
19
+ try:
20
+ nltk.download("punkt_tab", quiet=True) # sentence splitting in Long TTS mode
21
+ except Exception as e: # network hiccup: only "Sentence" split needs it
22
+ print(f"nltk punkt_tab download failed: {e}")
23
 
24
+ DEVICE = "cuda"
25
 
26
+ DEFAULT_WHISPER = "base.en"
27
+ DEFAULT_VOICECRAFT = "giga830M"
 
 
 
 
 
 
 
 
28
 
 
 
 
29
 
30
+ def ensure_checkpoint(filename):
31
+ path = f"./pretrained_models/{filename}"
32
+ if not os.path.exists(path):
33
+ path = hf_hub_download("pyp1/VoiceCraft", filename)
34
+ return path
 
 
35
 
36
 
37
+ def build_whisper(name, device):
38
+ if name is None:
39
+ return None
40
+ import whisper
41
+ from whisper.tokenizer import get_tokenizer
42
+ return {
43
+ "name": name,
44
+ "model": whisper.load_model(name, device=device),
45
+ "tokenizer": get_tokenizer(multilingual=False),
46
+ }
47
+
48
+
49
+ def build_voicecraft(name, device):
50
+ ckpt = torch.load(ensure_checkpoint(f"{name}.pth"), map_location="cpu", weights_only=False)
51
  model = voicecraft.VoiceCraft(ckpt["config"])
52
  model.load_state_dict(ckpt["model"])
53
  model.to(device)
54
  model.eval()
55
+ audio_tokenizer = AudioTokenizer(device=torch.device(device), signature=ensure_checkpoint("encodec_4cb2048_giga.th"))
56
+ audio_tokenizer._device = torch.device(DEVICE) # inputs go to CUDA; codec is moved there inside GPU calls
57
+ return {
58
+ "name": name,
59
  "ckpt": ckpt,
60
  "model": model,
61
  "text_tokenizer": TextTokenizer(backend="espeak"),
62
+ "audio_tokenizer": audio_tokenizer,
63
  }
64
 
65
+
66
+ # ZeroGPU: load the default models at import time and move them to CUDA here
67
+ # (globals assigned inside @spaces.GPU functions do not persist between calls).
68
+ whisper_model = build_whisper(DEFAULT_WHISPER, DEVICE)
69
+ voicecraft_model = build_voicecraft(DEFAULT_VOICECRAFT, DEVICE)
70
+ voicecraft_model["audio_tokenizer"].codec.to(DEVICE)
71
+
72
+
73
+ def seed_everything(seed):
74
+ if seed != -1:
75
+ os.environ['PYTHONHASHSEED'] = str(seed)
76
+ random.seed(seed)
77
+ np.random.seed(seed)
78
+ torch.manual_seed(seed)
79
+ torch.cuda.manual_seed(seed)
80
+ torch.backends.cudnn.benchmark = False
81
+ torch.backends.cudnn.deterministic = True
82
+
83
+
84
+ def load_models(whisper_model_choice, voicecraft_model_choice):
85
+ """Swap models on CPU in the main process; GPU functions move them to CUDA per call."""
86
+ global whisper_model, voicecraft_model
87
+ if whisper_model_choice in (None, "None"):
88
+ whisper_model = None
89
+ elif whisper_model is None or whisper_model["name"] != whisper_model_choice:
90
+ whisper_model = build_whisper(whisper_model_choice, "cpu")
91
+ if voicecraft_model["name"] != voicecraft_model_choice:
92
+ voicecraft_model = build_voicecraft(voicecraft_model_choice, "cpu")
93
  return gr.Accordion()
94
 
95
+
96
  @spaces.GPU(duration=60)
97
  def transcribe(seed, audio_path):
98
  if whisper_model is None:
99
  raise gr.Error("Whisper model not loaded")
100
  seed_everything(seed)
101
+ whisper_model["model"].to(DEVICE)
102
+
103
  number_tokens = [
104
  i
105
  for i in range(whisper_model["tokenizer"].eot)
 
107
  ]
108
  result = whisper_model["model"].transcribe(audio_path, suppress_tokens=[-1] + number_tokens, word_timestamps=True)
109
  words = [word_info for segment in result["segments"] for word_info in segment["words"]]
110
+
111
  transcript = result["text"]
112
  transcript_with_start_time = " ".join([f"{word['start']} {word['word']}" for word in words])
113
  transcript_with_end_time = " ".join([f"{word['word']} {word['end']}" for word in words])
 
125
 
126
  def get_output_audio(audio_tensors, codec_audio_sr):
127
  result = torch.cat(audio_tensors, 1)
128
+ return int(codec_audio_sr), result[0].float().numpy()
129
+
 
 
130
 
131
  @spaces.GPU(duration=90)
132
  def run(seed, left_margin, right_margin, codec_audio_sr, codec_sr, top_k, top_p, temperature,
 
153
  else:
154
  sentences = [transcript.replace("\n", " ")]
155
 
156
+ device = DEVICE
157
+ voicecraft_model["model"].to(DEVICE)
158
+ voicecraft_model["audio_tokenizer"].codec.to(DEVICE)
159
+ num_frames, sample_rate = audio_info(audio_path)
160
+ audio_dur = num_frames / sample_rate
161
 
162
  audio_tensors = []
163
  inference_transcript = ""
 
186
 
187
  inference_transcript += target_transcript + "\n"
188
 
189
+ prompt_end_frame = int(min(audio_dur, prompt_end_time) * sample_rate)
190
  _, gen_audio = inference_one_sample(voicecraft_model["model"],
191
  voicecraft_model["ckpt"]["config"],
192
  voicecraft_model["ckpt"]["phn2num"],
 
240
  if audio_path is None:
241
  return 0, 0, 0
242
 
243
+ num_frames, sample_rate = audio_info(audio_path)
244
+ max_time = round(num_frames / sample_rate, 2)
245
  return [
246
  gr.Slider(maximum=max_time, value=max_time),
247
  gr.Slider(maximum=max_time, value=0),
 
250
 
251
 
252
  def change_mode(mode):
 
253
  return [
254
  gr.Group(visible=mode != "Edit"),
255
  gr.Group(visible=mode == "Edit"),
 
403
  with gr.Row():
404
  voicecraft_model_choice = gr.Radio(label="VoiceCraft model", value="giga830M", choices=["giga330M", "giga830M"])
405
  whisper_model_choice = gr.Radio(label="Whisper model", value="base.en",
406
+ choices=["tiny.en", "base.en", "small.en", "medium.en", "large"])
407
 
408
  with gr.Row():
409
  with gr.Column(scale=2):
audiocraft/LICENSE ADDED
@@ -0,0 +1,21 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ MIT License
2
+
3
+ Copyright (c) Meta Platforms, Inc. and affiliates.
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
audiocraft/__init__.py ADDED
@@ -0,0 +1,2 @@
 
 
 
1
+ """Minimal vendored subset of facebookresearch/audiocraft (MIT, commit f83babf):
2
+ only the EnCodec model, SEANet, and RVQ needed to load VoiceCraft's codec checkpoint."""
audiocraft/models/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ from .encodec import CompressionModel, EncodecModel
audiocraft/models/encodec.py ADDED
@@ -0,0 +1,506 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+ # All rights reserved.
3
+ #
4
+ # This source code is licensed under the license found in the
5
+ # LICENSE file in the root directory of this source tree.
6
+ """Compression models or wrapper around existing models.
7
+ Also defines the main interface that a model must follow to be usable as an audio tokenizer.
8
+ """
9
+
10
+ from abc import ABC, abstractmethod
11
+ import logging
12
+ import math
13
+ from pathlib import Path
14
+ import typing as tp
15
+
16
+ from einops import rearrange
17
+ import numpy as np
18
+ import torch
19
+ from torch import nn
20
+ HFEncodecModel = tp.Any # vendored copy: the transformers EncodecModel wrapper is unused
21
+
22
+ from .. import quantization as qt
23
+
24
+
25
+ logger = logging.getLogger()
26
+
27
+
28
+ class CompressionModel(ABC, nn.Module):
29
+ """Base API for all compression model that aim at being used as audio tokenizers
30
+ with a language model.
31
+ """
32
+
33
+ @abstractmethod
34
+ def forward(self, x: torch.Tensor) -> qt.QuantizedResult:
35
+ ...
36
+
37
+ @abstractmethod
38
+ def encode(self, x: torch.Tensor) -> tp.Tuple[torch.Tensor, tp.Optional[torch.Tensor]]:
39
+ """See `EncodecModel.encode`."""
40
+ ...
41
+
42
+ @abstractmethod
43
+ def decode(self, codes: torch.Tensor, scale: tp.Optional[torch.Tensor] = None):
44
+ """See `EncodecModel.decode`."""
45
+ ...
46
+
47
+ @abstractmethod
48
+ def decode_latent(self, codes: torch.Tensor):
49
+ """Decode from the discrete codes to continuous latent space."""
50
+ ...
51
+
52
+ @property
53
+ @abstractmethod
54
+ def channels(self) -> int:
55
+ ...
56
+
57
+ @property
58
+ @abstractmethod
59
+ def frame_rate(self) -> float:
60
+ ...
61
+
62
+ @property
63
+ @abstractmethod
64
+ def sample_rate(self) -> int:
65
+ ...
66
+
67
+ @property
68
+ @abstractmethod
69
+ def cardinality(self) -> int:
70
+ ...
71
+
72
+ @property
73
+ @abstractmethod
74
+ def num_codebooks(self) -> int:
75
+ ...
76
+
77
+ @property
78
+ @abstractmethod
79
+ def total_codebooks(self) -> int:
80
+ ...
81
+
82
+ @abstractmethod
83
+ def set_num_codebooks(self, n: int):
84
+ """Set the active number of codebooks used by the quantizer."""
85
+ ...
86
+
87
+ @staticmethod
88
+ def get_pretrained(
89
+ name: str, device: tp.Union[torch.device, str] = 'cpu'
90
+ ) -> 'CompressionModel':
91
+ """Instantiate a CompressionModel from a given pretrained model.
92
+
93
+ Args:
94
+ name (Path or str): name of the pretrained model. See after.
95
+ device (torch.device or str): Device on which the model is loaded.
96
+
97
+ Pretrained models:
98
+ - dac_44khz (https://github.com/descriptinc/descript-audio-codec)
99
+ - dac_24khz (same)
100
+ - facebook/encodec_24khz (https://huggingface.co/facebook/encodec_24khz)
101
+ - facebook/encodec_32khz (https://huggingface.co/facebook/encodec_32khz)
102
+ - your own model on HugginFace. Export instructions to come...
103
+ """
104
+
105
+ from . import builders, loaders
106
+ model: CompressionModel
107
+ if name in ['dac_44khz', 'dac_24khz']:
108
+ model_type = name.split('_')[1]
109
+ logger.info("Getting pretrained compression model from DAC %s", model_type)
110
+ model = DAC(model_type)
111
+ elif name in ['debug_compression_model']:
112
+ logger.info("Getting pretrained compression model for debug")
113
+ model = builders.get_debug_compression_model()
114
+ elif Path(name).exists():
115
+ # We assume here if the paths exist that it is in fact an AC checkpoint
116
+ # that was exported using `audiocraft.utils.export` functions.
117
+ model = loaders.load_compression_model(name, device=device)
118
+ else:
119
+ logger.info("Getting pretrained compression model from HF %s", name)
120
+ hf_model = HFEncodecModel.from_pretrained(name)
121
+ model = HFEncodecCompressionModel(hf_model).to(device)
122
+ return model.to(device).eval()
123
+
124
+
125
+ class EncodecModel(CompressionModel):
126
+ """Encodec model operating on the raw waveform.
127
+
128
+ Args:
129
+ encoder (nn.Module): Encoder network.
130
+ decoder (nn.Module): Decoder network.
131
+ quantizer (qt.BaseQuantizer): Quantizer network.
132
+ frame_rate (int): Frame rate for the latent representation.
133
+ sample_rate (int): Audio sample rate.
134
+ channels (int): Number of audio channels.
135
+ causal (bool): Whether to use a causal version of the model.
136
+ renormalize (bool): Whether to renormalize the audio before running the model.
137
+ """
138
+ # we need assignment to override the property in the abstract class,
139
+ # I couldn't find a better way...
140
+ frame_rate: float = 0
141
+ sample_rate: int = 0
142
+ channels: int = 0
143
+
144
+ def __init__(self,
145
+ encoder: nn.Module,
146
+ decoder: nn.Module,
147
+ quantizer: qt.BaseQuantizer,
148
+ frame_rate: int,
149
+ sample_rate: int,
150
+ channels: int,
151
+ causal: bool = False,
152
+ renormalize: bool = False):
153
+ super().__init__()
154
+ self.encoder = encoder
155
+ self.decoder = decoder
156
+ self.quantizer = quantizer
157
+ self.frame_rate = frame_rate
158
+ self.sample_rate = sample_rate
159
+ self.channels = channels
160
+ self.renormalize = renormalize
161
+ self.causal = causal
162
+ if self.causal:
163
+ # we force disabling here to avoid handling linear overlap of segments
164
+ # as supported in original EnCodec codebase.
165
+ assert not self.renormalize, 'Causal model does not support renormalize'
166
+
167
+ @property
168
+ def total_codebooks(self):
169
+ """Total number of quantizer codebooks available."""
170
+ return self.quantizer.total_codebooks
171
+
172
+ @property
173
+ def num_codebooks(self):
174
+ """Active number of codebooks used by the quantizer."""
175
+ return self.quantizer.num_codebooks
176
+
177
+ def set_num_codebooks(self, n: int):
178
+ """Set the active number of codebooks used by the quantizer."""
179
+ self.quantizer.set_num_codebooks(n)
180
+
181
+ @property
182
+ def cardinality(self):
183
+ """Cardinality of each codebook."""
184
+ return self.quantizer.bins
185
+
186
+ def preprocess(self, x: torch.Tensor) -> tp.Tuple[torch.Tensor, tp.Optional[torch.Tensor]]:
187
+ scale: tp.Optional[torch.Tensor]
188
+ if self.renormalize:
189
+ mono = x.mean(dim=1, keepdim=True)
190
+ volume = mono.pow(2).mean(dim=2, keepdim=True).sqrt()
191
+ scale = 1e-8 + volume
192
+ x = x / scale
193
+ scale = scale.view(-1, 1)
194
+ else:
195
+ scale = None
196
+ return x, scale
197
+
198
+ def postprocess(self,
199
+ x: torch.Tensor,
200
+ scale: tp.Optional[torch.Tensor] = None) -> torch.Tensor:
201
+ if scale is not None:
202
+ assert self.renormalize
203
+ x = x * scale.view(-1, 1, 1)
204
+ return x
205
+
206
+ def forward(self, x: torch.Tensor) -> qt.QuantizedResult:
207
+ assert x.dim() == 3
208
+ length = x.shape[-1]
209
+ x, scale = self.preprocess(x)
210
+
211
+ emb = self.encoder(x)
212
+ q_res = self.quantizer(emb, self.frame_rate)
213
+ out = self.decoder(q_res.x)
214
+
215
+ # remove extra padding added by the encoder and decoder
216
+ assert out.shape[-1] >= length, (out.shape[-1], length)
217
+ out = out[..., :length]
218
+
219
+ q_res.x = self.postprocess(out, scale)
220
+
221
+ return q_res
222
+
223
+ def encode(self, x: torch.Tensor) -> tp.Tuple[torch.Tensor, tp.Optional[torch.Tensor]]:
224
+ """Encode the given input tensor to quantized representation along with scale parameter.
225
+
226
+ Args:
227
+ x (torch.Tensor): Float tensor of shape [B, C, T]
228
+
229
+ Returns:
230
+ codes, scale (tuple of torch.Tensor, torch.Tensor): Tuple composed of:
231
+ codes a float tensor of shape [B, K, T] with K the number of codebooks used and T the timestep.
232
+ scale a float tensor containing the scale for audio renormalizealization.
233
+ """
234
+ assert x.dim() == 3
235
+ x, scale = self.preprocess(x)
236
+ emb = self.encoder(x)
237
+ codes = self.quantizer.encode(emb)
238
+ return codes, scale
239
+
240
+ def decode(self, codes: torch.Tensor, scale: tp.Optional[torch.Tensor] = None):
241
+ """Decode the given codes to a reconstructed representation, using the scale to perform
242
+ audio denormalization if needed.
243
+
244
+ Args:
245
+ codes (torch.Tensor): Int tensor of shape [B, K, T]
246
+ scale (torch.Tensor, optional): Float tensor containing the scale value.
247
+
248
+ Returns:
249
+ out (torch.Tensor): Float tensor of shape [B, C, T], the reconstructed audio.
250
+ """
251
+ emb = self.decode_latent(codes)
252
+ out = self.decoder(emb)
253
+ out = self.postprocess(out, scale)
254
+ # out contains extra padding added by the encoder and decoder
255
+ return out
256
+
257
+ def decode_latent(self, codes: torch.Tensor):
258
+ """Decode from the discrete codes to continuous latent space."""
259
+ return self.quantizer.decode(codes)
260
+
261
+
262
+ class DAC(CompressionModel):
263
+ def __init__(self, model_type: str = "44khz"):
264
+ super().__init__()
265
+ try:
266
+ import dac.utils
267
+ except ImportError:
268
+ raise RuntimeError("Could not import dac, make sure it is installed, "
269
+ "please run `pip install descript-audio-codec`")
270
+ self.model = dac.utils.load_model(model_type=model_type)
271
+ self.n_quantizers = self.total_codebooks
272
+ self.model.eval()
273
+
274
+ def forward(self, x: torch.Tensor) -> qt.QuantizedResult:
275
+ # We don't support training with this.
276
+ raise NotImplementedError("Forward and training with DAC not supported.")
277
+
278
+ def encode(self, x: torch.Tensor) -> tp.Tuple[torch.Tensor, tp.Optional[torch.Tensor]]:
279
+ codes = self.model.encode(x, self.n_quantizers)[1]
280
+ return codes[:, :self.n_quantizers], None
281
+
282
+ def decode(self, codes: torch.Tensor, scale: tp.Optional[torch.Tensor] = None):
283
+ assert scale is None
284
+ z_q = self.decode_latent(codes)
285
+ return self.model.decode(z_q)
286
+
287
+ def decode_latent(self, codes: torch.Tensor):
288
+ """Decode from the discrete codes to continuous latent space."""
289
+ return self.model.quantizer.from_codes(codes)[0]
290
+
291
+ @property
292
+ def channels(self) -> int:
293
+ return 1
294
+
295
+ @property
296
+ def frame_rate(self) -> float:
297
+ return self.model.sample_rate / self.model.hop_length
298
+
299
+ @property
300
+ def sample_rate(self) -> int:
301
+ return self.model.sample_rate
302
+
303
+ @property
304
+ def cardinality(self) -> int:
305
+ return self.model.codebook_size
306
+
307
+ @property
308
+ def num_codebooks(self) -> int:
309
+ return self.n_quantizers
310
+
311
+ @property
312
+ def total_codebooks(self) -> int:
313
+ return self.model.n_codebooks
314
+
315
+ def set_num_codebooks(self, n: int):
316
+ """Set the active number of codebooks used by the quantizer.
317
+ """
318
+ assert n >= 1
319
+ assert n <= self.total_codebooks
320
+ self.n_quantizers = n
321
+
322
+
323
+ class HFEncodecCompressionModel(CompressionModel):
324
+ """Wrapper around HuggingFace Encodec.
325
+ """
326
+ def __init__(self, model: HFEncodecModel):
327
+ super().__init__()
328
+ self.model = model
329
+ bws = self.model.config.target_bandwidths
330
+ num_codebooks = [
331
+ bw * 1000 / (self.frame_rate * math.log2(self.cardinality))
332
+ for bw in bws
333
+ ]
334
+ deltas = [nc - int(nc) for nc in num_codebooks]
335
+ # Checking we didn't do some bad maths and we indeed have integers!
336
+ assert all(deltas) <= 1e-3, deltas
337
+ self.possible_num_codebooks = [int(nc) for nc in num_codebooks]
338
+ self.set_num_codebooks(max(self.possible_num_codebooks))
339
+
340
+ def forward(self, x: torch.Tensor) -> qt.QuantizedResult:
341
+ # We don't support training with this.
342
+ raise NotImplementedError("Forward and training with HF EncodecModel not supported.")
343
+
344
+ def encode(self, x: torch.Tensor) -> tp.Tuple[torch.Tensor, tp.Optional[torch.Tensor]]:
345
+ bandwidth_index = self.possible_num_codebooks.index(self.num_codebooks)
346
+ bandwidth = self.model.config.target_bandwidths[bandwidth_index]
347
+ res = self.model.encode(x, None, bandwidth)
348
+ assert len(res[0]) == 1
349
+ assert len(res[1]) == 1
350
+ return res[0][0], res[1][0]
351
+
352
+ def decode(self, codes: torch.Tensor, scale: tp.Optional[torch.Tensor] = None):
353
+ if scale is None:
354
+ scales = [None] # type: ignore
355
+ else:
356
+ scales = scale # type: ignore
357
+ res = self.model.decode(codes[None], scales)
358
+ return res[0]
359
+
360
+ def decode_latent(self, codes: torch.Tensor):
361
+ """Decode from the discrete codes to continuous latent space."""
362
+ return self.model.quantizer.decode(codes.transpose(0, 1))
363
+
364
+ @property
365
+ def channels(self) -> int:
366
+ return self.model.config.audio_channels
367
+
368
+ @property
369
+ def frame_rate(self) -> float:
370
+ hop_length = int(np.prod(self.model.config.upsampling_ratios))
371
+ return self.sample_rate / hop_length
372
+
373
+ @property
374
+ def sample_rate(self) -> int:
375
+ return self.model.config.sampling_rate
376
+
377
+ @property
378
+ def cardinality(self) -> int:
379
+ return self.model.config.codebook_size
380
+
381
+ @property
382
+ def num_codebooks(self) -> int:
383
+ return self._num_codebooks
384
+
385
+ @property
386
+ def total_codebooks(self) -> int:
387
+ return max(self.possible_num_codebooks)
388
+
389
+ def set_num_codebooks(self, n: int):
390
+ """Set the active number of codebooks used by the quantizer.
391
+ """
392
+ if n not in self.possible_num_codebooks:
393
+ raise ValueError(f"Allowed values for num codebooks: {self.possible_num_codebooks}")
394
+ self._num_codebooks = n
395
+
396
+
397
+ class InterleaveStereoCompressionModel(CompressionModel):
398
+ """Wraps a CompressionModel to support stereo inputs. The wrapped model
399
+ will be applied independently to the left and right channels, and both codebooks
400
+ will be interleaved. If the wrapped model returns a representation `[B, K ,T]` per
401
+ channel, then the output will be `[B, K * 2, T]` or `[B, K, T * 2]` depending on
402
+ `per_timestep`.
403
+
404
+ Args:
405
+ model (CompressionModel): Compression model to wrap.
406
+ per_timestep (bool): Whether to interleave on the timestep dimension
407
+ or on the codebooks dimension.
408
+ """
409
+ def __init__(self, model: CompressionModel, per_timestep: bool = False):
410
+ super().__init__()
411
+ self.model = model
412
+ self.per_timestep = per_timestep
413
+ assert self.model.channels == 1, "Wrapped model is expected to be for monophonic audio"
414
+
415
+ @property
416
+ def total_codebooks(self):
417
+ return self.model.total_codebooks
418
+
419
+ @property
420
+ def num_codebooks(self):
421
+ """Active number of codebooks used by the quantizer.
422
+
423
+ ..Warning:: this reports the number of codebooks after the interleaving
424
+ of the codebooks!
425
+ """
426
+ return self.model.num_codebooks if self.per_timestep else self.model.num_codebooks * 2
427
+
428
+ def set_num_codebooks(self, n: int):
429
+ """Set the active number of codebooks used by the quantizer.
430
+
431
+ ..Warning:: this sets the number of codebooks before the interleaving!
432
+ """
433
+ self.model.set_num_codebooks(n)
434
+
435
+ @property
436
+ def num_virtual_steps(self) -> float:
437
+ """Return the number of virtual steps, e.g. one real step
438
+ will be split into that many steps.
439
+ """
440
+ return 2 if self.per_timestep else 1
441
+
442
+ @property
443
+ def frame_rate(self) -> float:
444
+ return self.model.frame_rate * self.num_virtual_steps
445
+
446
+ @property
447
+ def sample_rate(self) -> int:
448
+ return self.model.sample_rate
449
+
450
+ @property
451
+ def channels(self) -> int:
452
+ return 2
453
+
454
+ @property
455
+ def cardinality(self):
456
+ """Cardinality of each codebook.
457
+ """
458
+ return self.model.cardinality
459
+
460
+ def forward(self, x: torch.Tensor) -> qt.QuantizedResult:
461
+ raise NotImplementedError("Not supported, use encode and decode.")
462
+
463
+ def encode(self, x: torch.Tensor) -> tp.Tuple[torch.Tensor, tp.Optional[torch.Tensor]]:
464
+ B, C, T = x.shape
465
+ assert C == self.channels, f"Expecting stereo audio but audio num channels is {C}"
466
+
467
+ indices_c0, scales_c0 = self.model.encode(x[:, 0, ...].unsqueeze(1))
468
+ indices_c1, scales_c1 = self.model.encode(x[:, 1, ...].unsqueeze(1))
469
+ indices = torch.stack([indices_c0, indices_c1], dim=0)
470
+ scales: tp.Optional[torch.Tensor] = None
471
+ if scales_c0 is not None and scales_c1 is not None:
472
+ scales = torch.stack([scales_c0, scales_c1], dim=1)
473
+
474
+ if self.per_timestep:
475
+ indices = rearrange(indices, 'c b k t -> b k (t c)', c=2)
476
+ else:
477
+ indices = rearrange(indices, 'c b k t -> b (k c) t', c=2)
478
+
479
+ return (indices, scales)
480
+
481
+ def get_left_right_codes(self, codes: torch.Tensor) -> tp.Tuple[torch.Tensor, torch.Tensor]:
482
+ if self.per_timestep:
483
+ codes = rearrange(codes, 'b k (t c) -> c b k t', c=2)
484
+ else:
485
+ codes = rearrange(codes, 'b (k c) t -> c b k t', c=2)
486
+ return codes[0], codes[1]
487
+
488
+ def decode(self, codes: torch.Tensor, scale: tp.Optional[torch.Tensor] = None):
489
+ B, K, T = codes.shape
490
+ assert T % self.num_virtual_steps == 0, "Provided codes' number of timesteps does not match"
491
+ assert K == self.num_codebooks, "Provided codes' number of codebooks does not match"
492
+
493
+ scale_c0, scale_c1 = None, None
494
+ if scale is not None:
495
+ assert scale.size(0) == B and scale.size(1) == 2, f"Scale has unexpected shape: {scale.shape}"
496
+ scale_c0 = scale[0, ...]
497
+ scale_c1 = scale[1, ...]
498
+
499
+ codes_c0, codes_c1 = self.get_left_right_codes(codes)
500
+ audio_c0 = self.model.decode(codes_c0, scale_c0)
501
+ audio_c1 = self.model.decode(codes_c1, scale_c1)
502
+ return torch.cat([audio_c0, audio_c1], dim=1)
503
+
504
+ def decode_latent(self, codes: torch.Tensor):
505
+ """Decode from the discrete codes to continuous latent space."""
506
+ raise NotImplementedError("Not supported by interleaved stereo wrapped models.")
audiocraft/modules/__init__.py ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from .conv import (
2
+ NormConv1d,
3
+ NormConv2d,
4
+ NormConvTranspose1d,
5
+ NormConvTranspose2d,
6
+ StreamableConv1d,
7
+ StreamableConvTranspose1d,
8
+ pad_for_conv1d,
9
+ pad1d,
10
+ unpad1d,
11
+ )
12
+ from .lstm import StreamableLSTM
13
+ from .seanet import SEANetEncoder, SEANetDecoder
audiocraft/modules/conv.py ADDED
@@ -0,0 +1,243 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+ # All rights reserved.
3
+ #
4
+ # This source code is licensed under the license found in the
5
+ # LICENSE file in the root directory of this source tree.
6
+
7
+ import math
8
+ import typing as tp
9
+ import warnings
10
+
11
+ import torch
12
+ from torch import nn
13
+ from torch.nn import functional as F
14
+ from torch.nn.utils import spectral_norm, weight_norm
15
+
16
+
17
+ CONV_NORMALIZATIONS = frozenset(['none', 'weight_norm', 'spectral_norm',
18
+ 'time_group_norm'])
19
+
20
+
21
+ def apply_parametrization_norm(module: nn.Module, norm: str = 'none'):
22
+ assert norm in CONV_NORMALIZATIONS
23
+ if norm == 'weight_norm':
24
+ return weight_norm(module)
25
+ elif norm == 'spectral_norm':
26
+ return spectral_norm(module)
27
+ else:
28
+ # We already check was in CONV_NORMALIZATION, so any other choice
29
+ # doesn't need reparametrization.
30
+ return module
31
+
32
+
33
+ def get_norm_module(module: nn.Module, causal: bool = False, norm: str = 'none', **norm_kwargs):
34
+ """Return the proper normalization module. If causal is True, this will ensure the returned
35
+ module is causal, or return an error if the normalization doesn't support causal evaluation.
36
+ """
37
+ assert norm in CONV_NORMALIZATIONS
38
+ if norm == 'time_group_norm':
39
+ if causal:
40
+ raise ValueError("GroupNorm doesn't support causal evaluation.")
41
+ assert isinstance(module, nn.modules.conv._ConvNd)
42
+ return nn.GroupNorm(1, module.out_channels, **norm_kwargs)
43
+ else:
44
+ return nn.Identity()
45
+
46
+
47
+ def get_extra_padding_for_conv1d(x: torch.Tensor, kernel_size: int, stride: int,
48
+ padding_total: int = 0) -> int:
49
+ """See `pad_for_conv1d`."""
50
+ length = x.shape[-1]
51
+ n_frames = (length - kernel_size + padding_total) / stride + 1
52
+ ideal_length = (math.ceil(n_frames) - 1) * stride + (kernel_size - padding_total)
53
+ return ideal_length - length
54
+
55
+
56
+ def pad_for_conv1d(x: torch.Tensor, kernel_size: int, stride: int, padding_total: int = 0):
57
+ """Pad for a convolution to make sure that the last window is full.
58
+ Extra padding is added at the end. This is required to ensure that we can rebuild
59
+ an output of the same length, as otherwise, even with padding, some time steps
60
+ might get removed.
61
+ For instance, with total padding = 4, kernel size = 4, stride = 2:
62
+ 0 0 1 2 3 4 5 0 0 # (0s are padding)
63
+ 1 2 3 # (output frames of a convolution, last 0 is never used)
64
+ 0 0 1 2 3 4 5 0 # (output of tr. conv., but pos. 5 is going to get removed as padding)
65
+ 1 2 3 4 # once you removed padding, we are missing one time step !
66
+ """
67
+ extra_padding = get_extra_padding_for_conv1d(x, kernel_size, stride, padding_total)
68
+ return F.pad(x, (0, extra_padding))
69
+
70
+
71
+ def pad1d(x: torch.Tensor, paddings: tp.Tuple[int, int], mode: str = 'constant', value: float = 0.):
72
+ """Tiny wrapper around F.pad, just to allow for reflect padding on small input.
73
+ If this is the case, we insert extra 0 padding to the right before the reflection happen.
74
+ """
75
+ length = x.shape[-1]
76
+ padding_left, padding_right = paddings
77
+ assert padding_left >= 0 and padding_right >= 0, (padding_left, padding_right)
78
+ if mode == 'reflect':
79
+ max_pad = max(padding_left, padding_right)
80
+ extra_pad = 0
81
+ if length <= max_pad:
82
+ extra_pad = max_pad - length + 1
83
+ x = F.pad(x, (0, extra_pad))
84
+ padded = F.pad(x, paddings, mode, value)
85
+ end = padded.shape[-1] - extra_pad
86
+ return padded[..., :end]
87
+ else:
88
+ return F.pad(x, paddings, mode, value)
89
+
90
+
91
+ def unpad1d(x: torch.Tensor, paddings: tp.Tuple[int, int]):
92
+ """Remove padding from x, handling properly zero padding. Only for 1d!"""
93
+ padding_left, padding_right = paddings
94
+ assert padding_left >= 0 and padding_right >= 0, (padding_left, padding_right)
95
+ assert (padding_left + padding_right) <= x.shape[-1]
96
+ end = x.shape[-1] - padding_right
97
+ return x[..., padding_left: end]
98
+
99
+
100
+ class NormConv1d(nn.Module):
101
+ """Wrapper around Conv1d and normalization applied to this conv
102
+ to provide a uniform interface across normalization approaches.
103
+ """
104
+ def __init__(self, *args, causal: bool = False, norm: str = 'none',
105
+ norm_kwargs: tp.Dict[str, tp.Any] = {}, **kwargs):
106
+ super().__init__()
107
+ self.conv = apply_parametrization_norm(nn.Conv1d(*args, **kwargs), norm)
108
+ self.norm = get_norm_module(self.conv, causal, norm, **norm_kwargs)
109
+ self.norm_type = norm
110
+
111
+ def forward(self, x):
112
+ x = self.conv(x)
113
+ x = self.norm(x)
114
+ return x
115
+
116
+
117
+ class NormConv2d(nn.Module):
118
+ """Wrapper around Conv2d and normalization applied to this conv
119
+ to provide a uniform interface across normalization approaches.
120
+ """
121
+ def __init__(self, *args, norm: str = 'none', norm_kwargs: tp.Dict[str, tp.Any] = {}, **kwargs):
122
+ super().__init__()
123
+ self.conv = apply_parametrization_norm(nn.Conv2d(*args, **kwargs), norm)
124
+ self.norm = get_norm_module(self.conv, causal=False, norm=norm, **norm_kwargs)
125
+ self.norm_type = norm
126
+
127
+ def forward(self, x):
128
+ x = self.conv(x)
129
+ x = self.norm(x)
130
+ return x
131
+
132
+
133
+ class NormConvTranspose1d(nn.Module):
134
+ """Wrapper around ConvTranspose1d and normalization applied to this conv
135
+ to provide a uniform interface across normalization approaches.
136
+ """
137
+ def __init__(self, *args, causal: bool = False, norm: str = 'none',
138
+ norm_kwargs: tp.Dict[str, tp.Any] = {}, **kwargs):
139
+ super().__init__()
140
+ self.convtr = apply_parametrization_norm(nn.ConvTranspose1d(*args, **kwargs), norm)
141
+ self.norm = get_norm_module(self.convtr, causal, norm, **norm_kwargs)
142
+ self.norm_type = norm
143
+
144
+ def forward(self, x):
145
+ x = self.convtr(x)
146
+ x = self.norm(x)
147
+ return x
148
+
149
+
150
+ class NormConvTranspose2d(nn.Module):
151
+ """Wrapper around ConvTranspose2d and normalization applied to this conv
152
+ to provide a uniform interface across normalization approaches.
153
+ """
154
+ def __init__(self, *args, norm: str = 'none', norm_kwargs: tp.Dict[str, tp.Any] = {}, **kwargs):
155
+ super().__init__()
156
+ self.convtr = apply_parametrization_norm(nn.ConvTranspose2d(*args, **kwargs), norm)
157
+ self.norm = get_norm_module(self.convtr, causal=False, norm=norm, **norm_kwargs)
158
+
159
+ def forward(self, x):
160
+ x = self.convtr(x)
161
+ x = self.norm(x)
162
+ return x
163
+
164
+
165
+ class StreamableConv1d(nn.Module):
166
+ """Conv1d with some builtin handling of asymmetric or causal padding
167
+ and normalization.
168
+ """
169
+ def __init__(self, in_channels: int, out_channels: int,
170
+ kernel_size: int, stride: int = 1, dilation: int = 1,
171
+ groups: int = 1, bias: bool = True, causal: bool = False,
172
+ norm: str = 'none', norm_kwargs: tp.Dict[str, tp.Any] = {},
173
+ pad_mode: str = 'reflect'):
174
+ super().__init__()
175
+ # warn user on unusual setup between dilation and stride
176
+ if stride > 1 and dilation > 1:
177
+ warnings.warn("StreamableConv1d has been initialized with stride > 1 and dilation > 1"
178
+ f" (kernel_size={kernel_size} stride={stride}, dilation={dilation}).")
179
+ self.conv = NormConv1d(in_channels, out_channels, kernel_size, stride,
180
+ dilation=dilation, groups=groups, bias=bias, causal=causal,
181
+ norm=norm, norm_kwargs=norm_kwargs)
182
+ self.causal = causal
183
+ self.pad_mode = pad_mode
184
+
185
+ def forward(self, x):
186
+ B, C, T = x.shape
187
+ kernel_size = self.conv.conv.kernel_size[0]
188
+ stride = self.conv.conv.stride[0]
189
+ dilation = self.conv.conv.dilation[0]
190
+ kernel_size = (kernel_size - 1) * dilation + 1 # effective kernel size with dilations
191
+ padding_total = kernel_size - stride
192
+ extra_padding = get_extra_padding_for_conv1d(x, kernel_size, stride, padding_total)
193
+ if self.causal:
194
+ # Left padding for causal
195
+ x = pad1d(x, (padding_total, extra_padding), mode=self.pad_mode)
196
+ else:
197
+ # Asymmetric padding required for odd strides
198
+ padding_right = padding_total // 2
199
+ padding_left = padding_total - padding_right
200
+ x = pad1d(x, (padding_left, padding_right + extra_padding), mode=self.pad_mode)
201
+ return self.conv(x)
202
+
203
+
204
+ class StreamableConvTranspose1d(nn.Module):
205
+ """ConvTranspose1d with some builtin handling of asymmetric or causal padding
206
+ and normalization.
207
+ """
208
+ def __init__(self, in_channels: int, out_channels: int,
209
+ kernel_size: int, stride: int = 1, causal: bool = False,
210
+ norm: str = 'none', trim_right_ratio: float = 1.,
211
+ norm_kwargs: tp.Dict[str, tp.Any] = {}):
212
+ super().__init__()
213
+ self.convtr = NormConvTranspose1d(in_channels, out_channels, kernel_size, stride,
214
+ causal=causal, norm=norm, norm_kwargs=norm_kwargs)
215
+ self.causal = causal
216
+ self.trim_right_ratio = trim_right_ratio
217
+ assert self.causal or self.trim_right_ratio == 1., \
218
+ "`trim_right_ratio` != 1.0 only makes sense for causal convolutions"
219
+ assert self.trim_right_ratio >= 0. and self.trim_right_ratio <= 1.
220
+
221
+ def forward(self, x):
222
+ kernel_size = self.convtr.convtr.kernel_size[0]
223
+ stride = self.convtr.convtr.stride[0]
224
+ padding_total = kernel_size - stride
225
+
226
+ y = self.convtr(x)
227
+
228
+ # We will only trim fixed padding. Extra padding from `pad_for_conv1d` would be
229
+ # removed at the very end, when keeping only the right length for the output,
230
+ # as removing it here would require also passing the length at the matching layer
231
+ # in the encoder.
232
+ if self.causal:
233
+ # Trim the padding on the right according to the specified ratio
234
+ # if trim_right_ratio = 1.0, trim everything from right
235
+ padding_right = math.ceil(padding_total * self.trim_right_ratio)
236
+ padding_left = padding_total - padding_right
237
+ y = unpad1d(y, (padding_left, padding_right))
238
+ else:
239
+ # Asymmetric padding required for odd strides
240
+ padding_right = padding_total // 2
241
+ padding_left = padding_total - padding_right
242
+ y = unpad1d(y, (padding_left, padding_right))
243
+ return y
audiocraft/modules/lstm.py ADDED
@@ -0,0 +1,25 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+ # All rights reserved.
3
+ #
4
+ # This source code is licensed under the license found in the
5
+ # LICENSE file in the root directory of this source tree.
6
+
7
+ from torch import nn
8
+
9
+
10
+ class StreamableLSTM(nn.Module):
11
+ """LSTM without worrying about the hidden state, nor the layout of the data.
12
+ Expects input as convolutional layout.
13
+ """
14
+ def __init__(self, dimension: int, num_layers: int = 2, skip: bool = True):
15
+ super().__init__()
16
+ self.skip = skip
17
+ self.lstm = nn.LSTM(dimension, dimension, num_layers)
18
+
19
+ def forward(self, x):
20
+ x = x.permute(2, 0, 1)
21
+ y, _ = self.lstm(x)
22
+ if self.skip:
23
+ y = y + x
24
+ y = y.permute(1, 2, 0)
25
+ return y
audiocraft/modules/seanet.py ADDED
@@ -0,0 +1,258 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+ # All rights reserved.
3
+ #
4
+ # This source code is licensed under the license found in the
5
+ # LICENSE file in the root directory of this source tree.
6
+
7
+ import typing as tp
8
+
9
+ import numpy as np
10
+ import torch.nn as nn
11
+
12
+ from .conv import StreamableConv1d, StreamableConvTranspose1d
13
+ from .lstm import StreamableLSTM
14
+
15
+
16
+ class SEANetResnetBlock(nn.Module):
17
+ """Residual block from SEANet model.
18
+
19
+ Args:
20
+ dim (int): Dimension of the input/output.
21
+ kernel_sizes (list): List of kernel sizes for the convolutions.
22
+ dilations (list): List of dilations for the convolutions.
23
+ activation (str): Activation function.
24
+ activation_params (dict): Parameters to provide to the activation function.
25
+ norm (str): Normalization method.
26
+ norm_params (dict): Parameters to provide to the underlying normalization used along with the convolution.
27
+ causal (bool): Whether to use fully causal convolution.
28
+ pad_mode (str): Padding mode for the convolutions.
29
+ compress (int): Reduced dimensionality in residual branches (from Demucs v3).
30
+ true_skip (bool): Whether to use true skip connection or a simple
31
+ (streamable) convolution as the skip connection.
32
+ """
33
+ def __init__(self, dim: int, kernel_sizes: tp.List[int] = [3, 1], dilations: tp.List[int] = [1, 1],
34
+ activation: str = 'ELU', activation_params: dict = {'alpha': 1.0},
35
+ norm: str = 'none', norm_params: tp.Dict[str, tp.Any] = {}, causal: bool = False,
36
+ pad_mode: str = 'reflect', compress: int = 2, true_skip: bool = True):
37
+ super().__init__()
38
+ assert len(kernel_sizes) == len(dilations), 'Number of kernel sizes should match number of dilations'
39
+ act = getattr(nn, activation)
40
+ hidden = dim // compress
41
+ block = []
42
+ for i, (kernel_size, dilation) in enumerate(zip(kernel_sizes, dilations)):
43
+ in_chs = dim if i == 0 else hidden
44
+ out_chs = dim if i == len(kernel_sizes) - 1 else hidden
45
+ block += [
46
+ act(**activation_params),
47
+ StreamableConv1d(in_chs, out_chs, kernel_size=kernel_size, dilation=dilation,
48
+ norm=norm, norm_kwargs=norm_params,
49
+ causal=causal, pad_mode=pad_mode),
50
+ ]
51
+ self.block = nn.Sequential(*block)
52
+ self.shortcut: nn.Module
53
+ if true_skip:
54
+ self.shortcut = nn.Identity()
55
+ else:
56
+ self.shortcut = StreamableConv1d(dim, dim, kernel_size=1, norm=norm, norm_kwargs=norm_params,
57
+ causal=causal, pad_mode=pad_mode)
58
+
59
+ def forward(self, x):
60
+ return self.shortcut(x) + self.block(x)
61
+
62
+
63
+ class SEANetEncoder(nn.Module):
64
+ """SEANet encoder.
65
+
66
+ Args:
67
+ channels (int): Audio channels.
68
+ dimension (int): Intermediate representation dimension.
69
+ n_filters (int): Base width for the model.
70
+ n_residual_layers (int): nb of residual layers.
71
+ ratios (Sequence[int]): kernel size and stride ratios. The encoder uses downsampling ratios instead of
72
+ upsampling ratios, hence it will use the ratios in the reverse order to the ones specified here
73
+ that must match the decoder order. We use the decoder order as some models may only employ the decoder.
74
+ activation (str): Activation function.
75
+ activation_params (dict): Parameters to provide to the activation function.
76
+ norm (str): Normalization method.
77
+ norm_params (dict): Parameters to provide to the underlying normalization used along with the convolution.
78
+ kernel_size (int): Kernel size for the initial convolution.
79
+ last_kernel_size (int): Kernel size for the initial convolution.
80
+ residual_kernel_size (int): Kernel size for the residual layers.
81
+ dilation_base (int): How much to increase the dilation with each layer.
82
+ causal (bool): Whether to use fully causal convolution.
83
+ pad_mode (str): Padding mode for the convolutions.
84
+ true_skip (bool): Whether to use true skip connection or a simple
85
+ (streamable) convolution as the skip connection in the residual network blocks.
86
+ compress (int): Reduced dimensionality in residual branches (from Demucs v3).
87
+ lstm (int): Number of LSTM layers at the end of the encoder.
88
+ disable_norm_outer_blocks (int): Number of blocks for which we don't apply norm.
89
+ For the encoder, it corresponds to the N first blocks.
90
+ """
91
+ def __init__(self, channels: int = 1, dimension: int = 128, n_filters: int = 32, n_residual_layers: int = 3,
92
+ ratios: tp.List[int] = [8, 5, 4, 2], activation: str = 'ELU', activation_params: dict = {'alpha': 1.0},
93
+ norm: str = 'none', norm_params: tp.Dict[str, tp.Any] = {}, kernel_size: int = 7,
94
+ last_kernel_size: int = 7, residual_kernel_size: int = 3, dilation_base: int = 2, causal: bool = False,
95
+ pad_mode: str = 'reflect', true_skip: bool = True, compress: int = 2, lstm: int = 0,
96
+ disable_norm_outer_blocks: int = 0):
97
+ super().__init__()
98
+ self.channels = channels
99
+ self.dimension = dimension
100
+ self.n_filters = n_filters
101
+ self.ratios = list(reversed(ratios))
102
+ del ratios
103
+ self.n_residual_layers = n_residual_layers
104
+ self.hop_length = np.prod(self.ratios)
105
+ self.n_blocks = len(self.ratios) + 2 # first and last conv + residual blocks
106
+ self.disable_norm_outer_blocks = disable_norm_outer_blocks
107
+ assert self.disable_norm_outer_blocks >= 0 and self.disable_norm_outer_blocks <= self.n_blocks, \
108
+ "Number of blocks for which to disable norm is invalid." \
109
+ "It should be lower or equal to the actual number of blocks in the network and greater or equal to 0."
110
+
111
+ act = getattr(nn, activation)
112
+ mult = 1
113
+ model: tp.List[nn.Module] = [
114
+ StreamableConv1d(channels, mult * n_filters, kernel_size,
115
+ norm='none' if self.disable_norm_outer_blocks >= 1 else norm,
116
+ norm_kwargs=norm_params, causal=causal, pad_mode=pad_mode)
117
+ ]
118
+ # Downsample to raw audio scale
119
+ for i, ratio in enumerate(self.ratios):
120
+ block_norm = 'none' if self.disable_norm_outer_blocks >= i + 2 else norm
121
+ # Add residual layers
122
+ for j in range(n_residual_layers):
123
+ model += [
124
+ SEANetResnetBlock(mult * n_filters, kernel_sizes=[residual_kernel_size, 1],
125
+ dilations=[dilation_base ** j, 1],
126
+ norm=block_norm, norm_params=norm_params,
127
+ activation=activation, activation_params=activation_params,
128
+ causal=causal, pad_mode=pad_mode, compress=compress, true_skip=true_skip)]
129
+
130
+ # Add downsampling layers
131
+ model += [
132
+ act(**activation_params),
133
+ StreamableConv1d(mult * n_filters, mult * n_filters * 2,
134
+ kernel_size=ratio * 2, stride=ratio,
135
+ norm=block_norm, norm_kwargs=norm_params,
136
+ causal=causal, pad_mode=pad_mode),
137
+ ]
138
+ mult *= 2
139
+
140
+ if lstm:
141
+ model += [StreamableLSTM(mult * n_filters, num_layers=lstm)]
142
+
143
+ model += [
144
+ act(**activation_params),
145
+ StreamableConv1d(mult * n_filters, dimension, last_kernel_size,
146
+ norm='none' if self.disable_norm_outer_blocks == self.n_blocks else norm,
147
+ norm_kwargs=norm_params, causal=causal, pad_mode=pad_mode)
148
+ ]
149
+
150
+ self.model = nn.Sequential(*model)
151
+
152
+ def forward(self, x):
153
+ return self.model(x)
154
+
155
+
156
+ class SEANetDecoder(nn.Module):
157
+ """SEANet decoder.
158
+
159
+ Args:
160
+ channels (int): Audio channels.
161
+ dimension (int): Intermediate representation dimension.
162
+ n_filters (int): Base width for the model.
163
+ n_residual_layers (int): nb of residual layers.
164
+ ratios (Sequence[int]): kernel size and stride ratios.
165
+ activation (str): Activation function.
166
+ activation_params (dict): Parameters to provide to the activation function.
167
+ final_activation (str): Final activation function after all convolutions.
168
+ final_activation_params (dict): Parameters to provide to the activation function.
169
+ norm (str): Normalization method.
170
+ norm_params (dict): Parameters to provide to the underlying normalization used along with the convolution.
171
+ kernel_size (int): Kernel size for the initial convolution.
172
+ last_kernel_size (int): Kernel size for the initial convolution.
173
+ residual_kernel_size (int): Kernel size for the residual layers.
174
+ dilation_base (int): How much to increase the dilation with each layer.
175
+ causal (bool): Whether to use fully causal convolution.
176
+ pad_mode (str): Padding mode for the convolutions.
177
+ true_skip (bool): Whether to use true skip connection or a simple.
178
+ (streamable) convolution as the skip connection in the residual network blocks.
179
+ compress (int): Reduced dimensionality in residual branches (from Demucs v3).
180
+ lstm (int): Number of LSTM layers at the end of the encoder.
181
+ disable_norm_outer_blocks (int): Number of blocks for which we don't apply norm.
182
+ For the decoder, it corresponds to the N last blocks.
183
+ trim_right_ratio (float): Ratio for trimming at the right of the transposed convolution under the causal setup.
184
+ If equal to 1.0, it means that all the trimming is done at the right.
185
+ """
186
+ def __init__(self, channels: int = 1, dimension: int = 128, n_filters: int = 32, n_residual_layers: int = 3,
187
+ ratios: tp.List[int] = [8, 5, 4, 2], activation: str = 'ELU', activation_params: dict = {'alpha': 1.0},
188
+ final_activation: tp.Optional[str] = None, final_activation_params: tp.Optional[dict] = None,
189
+ norm: str = 'none', norm_params: tp.Dict[str, tp.Any] = {}, kernel_size: int = 7,
190
+ last_kernel_size: int = 7, residual_kernel_size: int = 3, dilation_base: int = 2, causal: bool = False,
191
+ pad_mode: str = 'reflect', true_skip: bool = True, compress: int = 2, lstm: int = 0,
192
+ disable_norm_outer_blocks: int = 0, trim_right_ratio: float = 1.0):
193
+ super().__init__()
194
+ self.dimension = dimension
195
+ self.channels = channels
196
+ self.n_filters = n_filters
197
+ self.ratios = ratios
198
+ del ratios
199
+ self.n_residual_layers = n_residual_layers
200
+ self.hop_length = np.prod(self.ratios)
201
+ self.n_blocks = len(self.ratios) + 2 # first and last conv + residual blocks
202
+ self.disable_norm_outer_blocks = disable_norm_outer_blocks
203
+ assert self.disable_norm_outer_blocks >= 0 and self.disable_norm_outer_blocks <= self.n_blocks, \
204
+ "Number of blocks for which to disable norm is invalid." \
205
+ "It should be lower or equal to the actual number of blocks in the network and greater or equal to 0."
206
+
207
+ act = getattr(nn, activation)
208
+ mult = int(2 ** len(self.ratios))
209
+ model: tp.List[nn.Module] = [
210
+ StreamableConv1d(dimension, mult * n_filters, kernel_size,
211
+ norm='none' if self.disable_norm_outer_blocks == self.n_blocks else norm,
212
+ norm_kwargs=norm_params, causal=causal, pad_mode=pad_mode)
213
+ ]
214
+
215
+ if lstm:
216
+ model += [StreamableLSTM(mult * n_filters, num_layers=lstm)]
217
+
218
+ # Upsample to raw audio scale
219
+ for i, ratio in enumerate(self.ratios):
220
+ block_norm = 'none' if self.disable_norm_outer_blocks >= self.n_blocks - (i + 1) else norm
221
+ # Add upsampling layers
222
+ model += [
223
+ act(**activation_params),
224
+ StreamableConvTranspose1d(mult * n_filters, mult * n_filters // 2,
225
+ kernel_size=ratio * 2, stride=ratio,
226
+ norm=block_norm, norm_kwargs=norm_params,
227
+ causal=causal, trim_right_ratio=trim_right_ratio),
228
+ ]
229
+ # Add residual layers
230
+ for j in range(n_residual_layers):
231
+ model += [
232
+ SEANetResnetBlock(mult * n_filters // 2, kernel_sizes=[residual_kernel_size, 1],
233
+ dilations=[dilation_base ** j, 1],
234
+ activation=activation, activation_params=activation_params,
235
+ norm=block_norm, norm_params=norm_params, causal=causal,
236
+ pad_mode=pad_mode, compress=compress, true_skip=true_skip)]
237
+
238
+ mult //= 2
239
+
240
+ # Add final layers
241
+ model += [
242
+ act(**activation_params),
243
+ StreamableConv1d(n_filters, channels, last_kernel_size,
244
+ norm='none' if self.disable_norm_outer_blocks >= 1 else norm,
245
+ norm_kwargs=norm_params, causal=causal, pad_mode=pad_mode)
246
+ ]
247
+ # Add optional final activation to decoder (eg. tanh)
248
+ if final_activation is not None:
249
+ final_act = getattr(nn, final_activation)
250
+ final_activation_params = final_activation_params or {}
251
+ model += [
252
+ final_act(**final_activation_params)
253
+ ]
254
+ self.model = nn.Sequential(*model)
255
+
256
+ def forward(self, z):
257
+ y = self.model(z)
258
+ return y
audiocraft/quantization/__init__.py ADDED
@@ -0,0 +1,9 @@
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+ # All rights reserved.
3
+ #
4
+ # This source code is licensed under the license found in the
5
+ # LICENSE file in the root directory of this source tree.
6
+ """RVQ."""
7
+ # flake8: noqa
8
+ from .vq import ResidualVectorQuantizer
9
+ from .base import BaseQuantizer, DummyQuantizer, QuantizedResult
audiocraft/quantization/base.py ADDED
@@ -0,0 +1,99 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+ # All rights reserved.
3
+ #
4
+ # This source code is licensed under the license found in the
5
+ # LICENSE file in the root directory of this source tree.
6
+
7
+ """
8
+ Base class for all quantizers.
9
+ """
10
+
11
+ from dataclasses import dataclass, field
12
+ import typing as tp
13
+
14
+ import torch
15
+ from torch import nn
16
+
17
+
18
+ @dataclass
19
+ class QuantizedResult:
20
+ x: torch.Tensor
21
+ codes: torch.Tensor
22
+ bandwidth: torch.Tensor # bandwidth in kb/s used, per batch item.
23
+ penalty: tp.Optional[torch.Tensor] = None
24
+ metrics: dict = field(default_factory=dict)
25
+
26
+
27
+ class BaseQuantizer(nn.Module):
28
+ """Base class for quantizers.
29
+ """
30
+
31
+ def forward(self, x: torch.Tensor, frame_rate: int) -> QuantizedResult:
32
+ """
33
+ Given input tensor x, returns first the quantized (or approximately quantized)
34
+ representation along with quantized codes, bandwidth, and any penalty term for the loss.
35
+ Finally, this returns a dict of metrics to update logging etc.
36
+ Frame rate must be passed so that the bandwidth is properly computed.
37
+ """
38
+ raise NotImplementedError()
39
+
40
+ def encode(self, x: torch.Tensor) -> torch.Tensor:
41
+ """Encode a given input tensor with the specified sample rate at the given bandwidth."""
42
+ raise NotImplementedError()
43
+
44
+ def decode(self, codes: torch.Tensor) -> torch.Tensor:
45
+ """Decode the given codes to the quantized representation."""
46
+ raise NotImplementedError()
47
+
48
+ @property
49
+ def total_codebooks(self):
50
+ """Total number of codebooks."""
51
+ raise NotImplementedError()
52
+
53
+ @property
54
+ def num_codebooks(self):
55
+ """Number of active codebooks."""
56
+ raise NotImplementedError()
57
+
58
+ def set_num_codebooks(self, n: int):
59
+ """Set the number of active codebooks."""
60
+ raise NotImplementedError()
61
+
62
+
63
+ class DummyQuantizer(BaseQuantizer):
64
+ """Fake quantizer that actually does not perform any quantization.
65
+ """
66
+ def __init__(self):
67
+ super().__init__()
68
+
69
+ def forward(self, x: torch.Tensor, frame_rate: int):
70
+ q = x.unsqueeze(1)
71
+ return QuantizedResult(x, q, torch.tensor(q.numel() * 32 * frame_rate / 1000 / len(x)).to(x))
72
+
73
+ def encode(self, x: torch.Tensor) -> torch.Tensor:
74
+ """Encode a given input tensor with the specified sample rate at the given bandwidth.
75
+ In the case of the DummyQuantizer, the codes are actually identical
76
+ to the input and resulting quantized representation as no quantization is done.
77
+ """
78
+ return x.unsqueeze(1)
79
+
80
+ def decode(self, codes: torch.Tensor) -> torch.Tensor:
81
+ """Decode the given codes to the quantized representation.
82
+ In the case of the DummyQuantizer, the codes are actually identical
83
+ to the input and resulting quantized representation as no quantization is done.
84
+ """
85
+ return codes.squeeze(1)
86
+
87
+ @property
88
+ def total_codebooks(self):
89
+ """Total number of codebooks."""
90
+ return 1
91
+
92
+ @property
93
+ def num_codebooks(self):
94
+ """Total number of codebooks."""
95
+ return self.total_codebooks
96
+
97
+ def set_num_codebooks(self, n: int):
98
+ """Set the number of active codebooks."""
99
+ raise AttributeError("Cannot override the number of codebooks for the dummy quantizer")
audiocraft/quantization/core_vq.py ADDED
@@ -0,0 +1,400 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+ # All rights reserved.
3
+ #
4
+ # This source code is licensed under the license found in the
5
+ # LICENSE file in the root directory of this source tree.
6
+
7
+ import typing as tp
8
+
9
+ from einops import rearrange, repeat
10
+
11
+ import torch
12
+ from torch import nn, einsum
13
+ import torch.nn.functional as F
14
+
15
+
16
+ def exists(val: tp.Optional[tp.Any]) -> bool:
17
+ return val is not None
18
+
19
+
20
+ def default(val: tp.Any, d: tp.Any) -> tp.Any:
21
+ return val if exists(val) else d
22
+
23
+
24
+ def l2norm(t):
25
+ return F.normalize(t, p=2, dim=-1)
26
+
27
+
28
+ def ema_inplace(moving_avg, new, decay: float):
29
+ moving_avg.data.mul_(decay).add_(new, alpha=(1 - decay))
30
+
31
+
32
+ def laplace_smoothing(x, n_categories: int, epsilon: float = 1e-5):
33
+ return (x + epsilon) / (x.sum() + n_categories * epsilon)
34
+
35
+
36
+ def uniform_init(*shape: int):
37
+ t = torch.empty(shape)
38
+ nn.init.kaiming_uniform_(t)
39
+ return t
40
+
41
+
42
+ def sample_vectors(samples, num: int):
43
+ num_samples, device = samples.shape[0], samples.device
44
+
45
+ if num_samples >= num:
46
+ indices = torch.randperm(num_samples, device=device)[:num]
47
+ else:
48
+ indices = torch.randint(0, num_samples, (num,), device=device)
49
+
50
+ return samples[indices]
51
+
52
+
53
+ def kmeans(samples, num_clusters: int, num_iters: int = 10):
54
+ dim, dtype = samples.shape[-1], samples.dtype
55
+
56
+ means = sample_vectors(samples, num_clusters)
57
+
58
+ for _ in range(num_iters):
59
+ diffs = rearrange(samples, "n d -> n () d") - rearrange(
60
+ means, "c d -> () c d"
61
+ )
62
+ dists = -(diffs ** 2).sum(dim=-1)
63
+
64
+ buckets = dists.max(dim=-1).indices
65
+ bins = torch.bincount(buckets, minlength=num_clusters)
66
+ zero_mask = bins == 0
67
+ bins_min_clamped = bins.masked_fill(zero_mask, 1)
68
+
69
+ new_means = buckets.new_zeros(num_clusters, dim, dtype=dtype)
70
+ new_means.scatter_add_(0, repeat(buckets, "n -> n d", d=dim), samples)
71
+ new_means = new_means / bins_min_clamped[..., None]
72
+
73
+ means = torch.where(zero_mask[..., None], means, new_means)
74
+
75
+ return means, bins
76
+
77
+
78
+ def orthogonal_loss_fn(t):
79
+ # eq (2) from https://arxiv.org/abs/2112.00384
80
+ n = t.shape[0]
81
+ normed_codes = l2norm(t)
82
+ identity = torch.eye(n, device=t.device)
83
+ cosine_sim = einsum("i d, j d -> i j", normed_codes, normed_codes)
84
+ return ((cosine_sim - identity) ** 2).sum() / (n ** 2)
85
+
86
+
87
+ class EuclideanCodebook(nn.Module):
88
+ """Codebook with Euclidean distance.
89
+
90
+ Args:
91
+ dim (int): Dimension.
92
+ codebook_size (int): Codebook size.
93
+ kmeans_init (bool): Whether to use k-means to initialize the codebooks.
94
+ If set to true, run the k-means algorithm on the first training batch and use
95
+ the learned centroids as initialization.
96
+ kmeans_iters (int): Number of iterations used for k-means algorithm at initialization.
97
+ decay (float): Decay for exponential moving average over the codebooks.
98
+ epsilon (float): Epsilon value for numerical stability.
99
+ threshold_ema_dead_code (int): Threshold for dead code expiration. Replace any codes
100
+ that have an exponential moving average cluster size less than the specified threshold with
101
+ randomly selected vector from the current batch.
102
+ """
103
+ def __init__(
104
+ self,
105
+ dim: int,
106
+ codebook_size: int,
107
+ kmeans_init: int = False,
108
+ kmeans_iters: int = 10,
109
+ decay: float = 0.8,
110
+ epsilon: float = 1e-5,
111
+ threshold_ema_dead_code: int = 2,
112
+ ):
113
+ super().__init__()
114
+ self.decay = decay
115
+ init_fn: tp.Union[tp.Callable[..., torch.Tensor], tp.Any] = uniform_init if not kmeans_init else torch.zeros
116
+ embed = init_fn(codebook_size, dim)
117
+
118
+ self.codebook_size = codebook_size
119
+
120
+ self.kmeans_iters = kmeans_iters
121
+ self.epsilon = epsilon
122
+ self.threshold_ema_dead_code = threshold_ema_dead_code
123
+
124
+ self.register_buffer("inited", torch.Tensor([not kmeans_init]))
125
+ self.register_buffer("cluster_size", torch.zeros(codebook_size))
126
+ self.register_buffer("embed", embed)
127
+ self.register_buffer("embed_avg", embed.clone())
128
+
129
+ @torch.jit.ignore
130
+ def init_embed_(self, data):
131
+ if self.inited:
132
+ return
133
+
134
+ embed, cluster_size = kmeans(data, self.codebook_size, self.kmeans_iters)
135
+ self.embed.data.copy_(embed)
136
+ self.embed_avg.data.copy_(embed.clone())
137
+ self.cluster_size.data.copy_(cluster_size)
138
+ self.inited.data.copy_(torch.Tensor([True]))
139
+ # Make sure all buffers across workers are in sync after initialization
140
+ pass # vendored, inference-only: no distributed sync
141
+
142
+ def replace_(self, samples, mask):
143
+ modified_codebook = torch.where(
144
+ mask[..., None], sample_vectors(samples, self.codebook_size), self.embed
145
+ )
146
+ self.embed.data.copy_(modified_codebook)
147
+
148
+ def expire_codes_(self, batch_samples):
149
+ if self.threshold_ema_dead_code == 0:
150
+ return
151
+
152
+ expired_codes = self.cluster_size < self.threshold_ema_dead_code
153
+ if not torch.any(expired_codes):
154
+ return
155
+
156
+ batch_samples = rearrange(batch_samples, "... d -> (...) d")
157
+ self.replace_(batch_samples, mask=expired_codes)
158
+ pass # vendored, inference-only: no distributed sync
159
+
160
+ def preprocess(self, x):
161
+ x = rearrange(x, "... d -> (...) d")
162
+ return x
163
+
164
+ def quantize(self, x):
165
+ embed = self.embed.t()
166
+ dist = -(
167
+ x.pow(2).sum(1, keepdim=True)
168
+ - 2 * x @ embed
169
+ + embed.pow(2).sum(0, keepdim=True)
170
+ )
171
+ embed_ind = dist.max(dim=-1).indices
172
+ return embed_ind
173
+
174
+ def postprocess_emb(self, embed_ind, shape):
175
+ return embed_ind.view(*shape[:-1])
176
+
177
+ def dequantize(self, embed_ind):
178
+ quantize = F.embedding(embed_ind, self.embed)
179
+ return quantize
180
+
181
+ def encode(self, x):
182
+ shape = x.shape
183
+ # pre-process
184
+ x = self.preprocess(x)
185
+ # quantize
186
+ embed_ind = self.quantize(x)
187
+ # post-process
188
+ embed_ind = self.postprocess_emb(embed_ind, shape)
189
+ return embed_ind
190
+
191
+ def decode(self, embed_ind):
192
+ quantize = self.dequantize(embed_ind)
193
+ return quantize
194
+
195
+ def forward(self, x):
196
+ shape, dtype = x.shape, x.dtype
197
+ x = self.preprocess(x)
198
+ self.init_embed_(x)
199
+
200
+ embed_ind = self.quantize(x)
201
+ embed_onehot = F.one_hot(embed_ind, self.codebook_size).type(dtype)
202
+ embed_ind = self.postprocess_emb(embed_ind, shape)
203
+ quantize = self.dequantize(embed_ind)
204
+
205
+ if self.training:
206
+ # We do the expiry of code at that point as buffers are in sync
207
+ # and all the workers will take the same decision.
208
+ self.expire_codes_(x)
209
+ ema_inplace(self.cluster_size, embed_onehot.sum(0), self.decay)
210
+ embed_sum = x.t() @ embed_onehot
211
+ ema_inplace(self.embed_avg, embed_sum.t(), self.decay)
212
+ cluster_size = (
213
+ laplace_smoothing(self.cluster_size, self.codebook_size, self.epsilon)
214
+ * self.cluster_size.sum()
215
+ )
216
+ embed_normalized = self.embed_avg / cluster_size.unsqueeze(1)
217
+ self.embed.data.copy_(embed_normalized)
218
+
219
+ return quantize, embed_ind
220
+
221
+
222
+ class VectorQuantization(nn.Module):
223
+ """Vector quantization implementation.
224
+ Currently supports only euclidean distance.
225
+
226
+ Args:
227
+ dim (int): Dimension
228
+ codebook_size (int): Codebook size
229
+ codebook_dim (int): Codebook dimension. If not defined, uses the specified dimension in dim.
230
+ decay (float): Decay for exponential moving average over the codebooks.
231
+ epsilon (float): Epsilon value for numerical stability.
232
+ kmeans_init (bool): Whether to use kmeans to initialize the codebooks.
233
+ kmeans_iters (int): Number of iterations used for kmeans initialization.
234
+ threshold_ema_dead_code (int):
235
+ channels_last (bool): Channels are the last dimension in the input tensors.
236
+ commitment_weight (float): Weight for commitment loss.
237
+ orthogonal_reg_weight (float): Orthogonal regularization weights.
238
+ orthogonal_reg_active_codes_only (bool): Apply orthogonal regularization only on active codes.
239
+ orthogonal_reg_max_codes (optional int): Maximum number of codes to consider
240
+ for orthogonal regularization.
241
+ threshold_ema_dead_code (int): Threshold for dead code expiration. Replace any codes
242
+ that have an exponential moving average cluster size less than the specified threshold with
243
+ randomly selected vector from the current batch.
244
+ """
245
+ def __init__(
246
+ self,
247
+ dim: int,
248
+ codebook_size: int,
249
+ codebook_dim: tp.Optional[int] = None,
250
+ decay: float = 0.8,
251
+ epsilon: float = 1e-5,
252
+ kmeans_init: bool = False,
253
+ kmeans_iters: int = 10,
254
+ threshold_ema_dead_code: int = 2,
255
+ channels_last: bool = False,
256
+ commitment_weight: float = 1.,
257
+ orthogonal_reg_weight: float = 0.0,
258
+ orthogonal_reg_active_codes_only: bool = False,
259
+ orthogonal_reg_max_codes: tp.Optional[int] = None,
260
+ ):
261
+ super().__init__()
262
+ _codebook_dim: int = default(codebook_dim, dim)
263
+
264
+ requires_projection = _codebook_dim != dim
265
+ self.project_in = (nn.Linear(dim, _codebook_dim) if requires_projection else nn.Identity())
266
+ self.project_out = (nn.Linear(_codebook_dim, dim) if requires_projection else nn.Identity())
267
+
268
+ self.epsilon = epsilon
269
+ self.commitment_weight = commitment_weight
270
+
271
+ self.orthogonal_reg_weight = orthogonal_reg_weight
272
+ self.orthogonal_reg_active_codes_only = orthogonal_reg_active_codes_only
273
+ self.orthogonal_reg_max_codes = orthogonal_reg_max_codes
274
+
275
+ self._codebook = EuclideanCodebook(dim=_codebook_dim, codebook_size=codebook_size,
276
+ kmeans_init=kmeans_init, kmeans_iters=kmeans_iters,
277
+ decay=decay, epsilon=epsilon,
278
+ threshold_ema_dead_code=threshold_ema_dead_code)
279
+ self.codebook_size = codebook_size
280
+
281
+ self.channels_last = channels_last
282
+
283
+ @property
284
+ def codebook(self):
285
+ return self._codebook.embed
286
+
287
+ @property
288
+ def inited(self):
289
+ return self._codebook.inited
290
+
291
+ def _preprocess(self, x):
292
+ if not self.channels_last:
293
+ x = rearrange(x, "b d n -> b n d")
294
+ return x
295
+
296
+ def _postprocess(self, quantize):
297
+ if not self.channels_last:
298
+ quantize = rearrange(quantize, "b n d -> b d n")
299
+ return quantize
300
+
301
+ def encode(self, x):
302
+ x = self._preprocess(x)
303
+ x = self.project_in(x)
304
+ embed_in = self._codebook.encode(x)
305
+ return embed_in
306
+
307
+ def decode(self, embed_ind):
308
+ quantize = self._codebook.decode(embed_ind)
309
+ quantize = self.project_out(quantize)
310
+ quantize = self._postprocess(quantize)
311
+ return quantize
312
+
313
+ def forward(self, x):
314
+ device = x.device
315
+ x = self._preprocess(x)
316
+
317
+ x = self.project_in(x)
318
+ quantize, embed_ind = self._codebook(x)
319
+
320
+ if self.training:
321
+ quantize = x + (quantize - x).detach()
322
+
323
+ loss = torch.tensor([0.0], device=device, requires_grad=self.training)
324
+
325
+ if self.training:
326
+ if self.commitment_weight > 0:
327
+ commit_loss = F.mse_loss(quantize.detach(), x)
328
+ loss = loss + commit_loss * self.commitment_weight
329
+
330
+ if self.orthogonal_reg_weight > 0:
331
+ codebook = self.codebook
332
+
333
+ if self.orthogonal_reg_active_codes_only:
334
+ # only calculate orthogonal loss for the activated codes for this batch
335
+ unique_code_ids = torch.unique(embed_ind)
336
+ codebook = codebook[unique_code_ids]
337
+
338
+ num_codes = codebook.shape[0]
339
+ if exists(self.orthogonal_reg_max_codes) and num_codes > self.orthogonal_reg_max_codes:
340
+ rand_ids = torch.randperm(num_codes, device=device)[:self.orthogonal_reg_max_codes]
341
+ codebook = codebook[rand_ids]
342
+
343
+ orthogonal_reg_loss = orthogonal_loss_fn(codebook)
344
+ loss = loss + orthogonal_reg_loss * self.orthogonal_reg_weight
345
+
346
+ quantize = self.project_out(quantize)
347
+ quantize = self._postprocess(quantize)
348
+
349
+ return quantize, embed_ind, loss
350
+
351
+
352
+ class ResidualVectorQuantization(nn.Module):
353
+ """Residual vector quantization implementation.
354
+
355
+ Follows Algorithm 1. in https://arxiv.org/pdf/2107.03312.pdf
356
+ """
357
+ def __init__(self, *, num_quantizers, **kwargs):
358
+ super().__init__()
359
+ self.layers = nn.ModuleList(
360
+ [VectorQuantization(**kwargs) for _ in range(num_quantizers)]
361
+ )
362
+
363
+ def forward(self, x, n_q: tp.Optional[int] = None):
364
+ quantized_out = 0.0
365
+ residual = x
366
+
367
+ all_losses = []
368
+ all_indices = []
369
+
370
+ n_q = n_q or len(self.layers)
371
+
372
+ for i, layer in enumerate(self.layers[:n_q]):
373
+ quantized, indices, loss = layer(residual)
374
+ residual = residual - quantized
375
+ quantized_out = quantized_out + quantized
376
+ all_indices.append(indices)
377
+ all_losses.append(loss)
378
+
379
+ out_losses, out_indices = map(torch.stack, (all_losses, all_indices))
380
+ return quantized_out, out_indices, out_losses
381
+
382
+ def encode(self, x: torch.Tensor, n_q: tp.Optional[int] = None) -> torch.Tensor:
383
+ residual = x
384
+ all_indices = []
385
+ n_q = n_q or len(self.layers)
386
+ for layer in self.layers[:n_q]:
387
+ indices = layer.encode(residual)
388
+ quantized = layer.decode(indices)
389
+ residual = residual - quantized
390
+ all_indices.append(indices)
391
+ out_indices = torch.stack(all_indices)
392
+ return out_indices
393
+
394
+ def decode(self, q_indices: torch.Tensor) -> torch.Tensor:
395
+ quantized_out = torch.tensor(0.0, device=q_indices.device)
396
+ for i, indices in enumerate(q_indices):
397
+ layer = self.layers[i]
398
+ quantized = layer.decode(indices)
399
+ quantized_out = quantized_out + quantized
400
+ return quantized_out
audiocraft/quantization/vq.py ADDED
@@ -0,0 +1,115 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+ # All rights reserved.
3
+ #
4
+ # This source code is licensed under the license found in the
5
+ # LICENSE file in the root directory of this source tree.
6
+
7
+ import math
8
+ import typing as tp
9
+
10
+ import torch
11
+
12
+ from .base import BaseQuantizer, QuantizedResult
13
+ from .core_vq import ResidualVectorQuantization
14
+
15
+
16
+ class ResidualVectorQuantizer(BaseQuantizer):
17
+ """Residual Vector Quantizer.
18
+
19
+ Args:
20
+ dimension (int): Dimension of the codebooks.
21
+ n_q (int): Number of residual vector quantizers used.
22
+ q_dropout (bool): Random quantizer drop out at train time.
23
+ bins (int): Codebook size.
24
+ decay (float): Decay for exponential moving average over the codebooks.
25
+ kmeans_init (bool): Whether to use kmeans to initialize the codebooks.
26
+ kmeans_iters (int): Number of iterations used for kmeans initialization.
27
+ threshold_ema_dead_code (int): Threshold for dead code expiration. Replace any codes
28
+ that have an exponential moving average cluster size less than the specified threshold with
29
+ randomly selected vector from the current batch.
30
+ orthogonal_reg_weight (float): Orthogonal regularization weights.
31
+ orthogonal_reg_active_codes_only (bool): Apply orthogonal regularization only on active codes.
32
+ orthogonal_reg_max_codes (optional int): Maximum number of codes to consider.
33
+ for orthogonal regularization.
34
+ """
35
+ def __init__(
36
+ self,
37
+ dimension: int = 256,
38
+ n_q: int = 8,
39
+ q_dropout: bool = False,
40
+ bins: int = 1024,
41
+ decay: float = 0.99,
42
+ kmeans_init: bool = True,
43
+ kmeans_iters: int = 10,
44
+ threshold_ema_dead_code: int = 2,
45
+ orthogonal_reg_weight: float = 0.0,
46
+ orthogonal_reg_active_codes_only: bool = False,
47
+ orthogonal_reg_max_codes: tp.Optional[int] = None,
48
+ ):
49
+ super().__init__()
50
+ self.max_n_q = n_q
51
+ self.n_q = n_q
52
+ self.q_dropout = q_dropout
53
+ self.dimension = dimension
54
+ self.bins = bins
55
+ self.decay = decay
56
+ self.kmeans_init = kmeans_init
57
+ self.kmeans_iters = kmeans_iters
58
+ self.threshold_ema_dead_code = threshold_ema_dead_code
59
+ self.orthogonal_reg_weight = orthogonal_reg_weight
60
+ self.orthogonal_reg_active_codes_only = orthogonal_reg_active_codes_only
61
+ self.orthogonal_reg_max_codes = orthogonal_reg_max_codes
62
+ self.vq = ResidualVectorQuantization(
63
+ dim=self.dimension,
64
+ codebook_size=self.bins,
65
+ num_quantizers=self.n_q,
66
+ decay=self.decay,
67
+ kmeans_init=self.kmeans_init,
68
+ kmeans_iters=self.kmeans_iters,
69
+ threshold_ema_dead_code=self.threshold_ema_dead_code,
70
+ orthogonal_reg_weight=self.orthogonal_reg_weight,
71
+ orthogonal_reg_active_codes_only=self.orthogonal_reg_active_codes_only,
72
+ orthogonal_reg_max_codes=self.orthogonal_reg_max_codes,
73
+ channels_last=False
74
+ )
75
+
76
+ def forward(self, x: torch.Tensor, frame_rate: int):
77
+ n_q = self.n_q
78
+ if self.training and self.q_dropout:
79
+ n_q = int(torch.randint(1, self.n_q + 1, (1,)).item())
80
+ bw_per_q = math.log2(self.bins) * frame_rate / 1000
81
+ quantized, codes, commit_loss = self.vq(x, n_q=n_q)
82
+ codes = codes.transpose(0, 1)
83
+ # codes is [B, K, T], with T frames, K nb of codebooks.
84
+ bw = torch.tensor(n_q * bw_per_q).to(x)
85
+ return QuantizedResult(quantized, codes, bw, penalty=torch.mean(commit_loss))
86
+
87
+ def encode(self, x: torch.Tensor) -> torch.Tensor:
88
+ """Encode a given input tensor with the specified frame rate at the given bandwidth.
89
+ The RVQ encode method sets the appropriate number of quantizer to use
90
+ and returns indices for each quantizer.
91
+ """
92
+ n_q = self.n_q
93
+ codes = self.vq.encode(x, n_q=n_q)
94
+ codes = codes.transpose(0, 1)
95
+ # codes is [B, K, T], with T frames, K nb of codebooks.
96
+ return codes
97
+
98
+ def decode(self, codes: torch.Tensor) -> torch.Tensor:
99
+ """Decode the given codes to the quantized representation."""
100
+ # codes is [B, K, T], with T frames, K nb of codebooks, vq.decode expects [K, B, T].
101
+ codes = codes.transpose(0, 1)
102
+ quantized = self.vq.decode(codes)
103
+ return quantized
104
+
105
+ @property
106
+ def total_codebooks(self):
107
+ return self.max_n_q
108
+
109
+ @property
110
+ def num_codebooks(self):
111
+ return self.n_q
112
+
113
+ def set_num_codebooks(self, n: int):
114
+ assert n > 0 and n <= self.max_n_q
115
+ self.n_q = n
data/tokenizer.py CHANGED
@@ -98,6 +98,65 @@ def convert_audio(wav: torch.Tensor, sr: int, target_sr: int, target_channels: i
98
  wav = torchaudio.transforms.Resample(sr, target_sr)(wav)
99
  return wav
100
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
101
  class AudioTokenizer:
102
  """EnCodec audio."""
103
 
@@ -106,8 +165,7 @@ class AudioTokenizer:
106
  device: Any = None,
107
  signature = None
108
  ) -> None:
109
- from audiocraft.solvers import CompressionSolver
110
- model = CompressionSolver.model_from_checkpoint(signature)
111
  self.sample_rate = model.sample_rate
112
  self.channels = model.channels
113
 
@@ -136,10 +194,7 @@ class AudioTokenizer:
136
 
137
  def tokenize_audio(tokenizer: AudioTokenizer, audio_path: str, offset = -1, num_frames=-1):
138
  # Load and pre-process the audio waveform
139
- if offset != -1 and num_frames!=-1:
140
- wav, sr = torchaudio.load(audio_path, frame_offset=offset, num_frames=num_frames)
141
- else:
142
- wav, sr = torchaudio.load(audio_path)
143
  wav = convert_audio(wav, sr, tokenizer.sample_rate, tokenizer.channels)
144
  wav = wav.unsqueeze(0)
145
 
 
98
  wav = torchaudio.transforms.Resample(sr, target_sr)(wav)
99
  return wav
100
 
101
+ def load_encodec_checkpoint(checkpoint_path: str):
102
+ """Rebuild audiocraft's EnCodec from a CompressionSolver checkpoint.
103
+
104
+ Mirrors audiocraft.solvers.CompressionSolver.model_from_checkpoint and
105
+ models.builders.get_compression_model, using the vendored minimal audiocraft
106
+ package so the full audiocraft dependency tree is not needed.
107
+ """
108
+ from omegaconf import OmegaConf
109
+ from audiocraft import modules, quantization as qt
110
+ from audiocraft.models.encodec import EncodecModel
111
+
112
+ state = torch.load(checkpoint_path, map_location="cpu", weights_only=False)
113
+ cfg = state["xp.cfg"]
114
+ assert cfg.compression_model == "encodec", cfg.compression_model
115
+
116
+ def as_dict(node):
117
+ return OmegaConf.to_container(node, resolve=True)
118
+
119
+ kwargs = as_dict(cfg.encodec)
120
+ encoder_name = kwargs.pop("autoencoder")
121
+ quantizer_name = kwargs.pop("quantizer")
122
+ assert encoder_name == "seanet", encoder_name
123
+ seanet = as_dict(cfg.seanet)
124
+ enc_over, dec_over = seanet.pop("encoder"), seanet.pop("decoder")
125
+ encoder = modules.SEANetEncoder(**{**seanet, **enc_over})
126
+ decoder = modules.SEANetDecoder(**{**seanet, **dec_over})
127
+ if quantizer_name == "no_quant":
128
+ quantizer = qt.DummyQuantizer()
129
+ else:
130
+ q_kwargs = as_dict(getattr(cfg, quantizer_name))
131
+ q_kwargs["dimension"] = encoder.dimension
132
+ quantizer = qt.ResidualVectorQuantizer(**q_kwargs)
133
+ frame_rate = kwargs["sample_rate"] // encoder.hop_length
134
+ renormalize = kwargs.pop("renormalize", False)
135
+ kwargs.pop("renorm", None)
136
+ model = EncodecModel(encoder, decoder, quantizer, frame_rate=frame_rate,
137
+ renormalize=renormalize, **kwargs)
138
+ model.load_state_dict(state["best_state"]["model"])
139
+ model.eval()
140
+ return model
141
+
142
+
143
+ def load_audio(audio_path: str, offset: int = -1, num_frames: int = -1):
144
+ """torchaudio.load replacement (torchaudio 2.9 needs torchcodec): returns ([C, T] float tensor, sr)."""
145
+ import soundfile as sf
146
+ if offset != -1 and num_frames != -1:
147
+ data, sr = sf.read(audio_path, start=offset, frames=num_frames, dtype="float32", always_2d=True)
148
+ else:
149
+ data, sr = sf.read(audio_path, dtype="float32", always_2d=True)
150
+ return torch.from_numpy(data.T.copy()), sr
151
+
152
+
153
+ def audio_info(audio_path: str):
154
+ """torchaudio.info replacement: returns (num_frames, sample_rate)."""
155
+ import soundfile as sf
156
+ info = sf.info(audio_path)
157
+ return info.frames, info.samplerate
158
+
159
+
160
  class AudioTokenizer:
161
  """EnCodec audio."""
162
 
 
165
  device: Any = None,
166
  signature = None
167
  ) -> None:
168
+ model = load_encodec_checkpoint(signature)
 
169
  self.sample_rate = model.sample_rate
170
  self.channels = model.channels
171
 
 
194
 
195
  def tokenize_audio(tokenizer: AudioTokenizer, audio_path: str, offset = -1, num_frames=-1):
196
  # Load and pre-process the audio waveform
197
+ wav, sr = load_audio(audio_path, offset, num_frames)
 
 
 
198
  wav = convert_audio(wav, sr, tokenizer.sample_rate, tokenizer.channels)
199
  wav = wav.unsqueeze(0)
200
 
inference_speech_editing_scale.py CHANGED
@@ -1,4 +1,4 @@
1
- import argparse, pickle
2
  import logging
3
  import os, random
4
  import numpy as np
@@ -65,7 +65,7 @@ def inference_one_sample(model, model_args, phn2num, text_tokenizer, audio_token
65
  temperature=decode_config['temperature'],
66
  stop_repetition=decode_config['stop_repetition'],
67
  kvcache=decode_config['kvcache'],
68
- silence_tokens=eval(decode_config['silence_tokens']) if type(decode_config['silence_tokens']) == str else decode_config['silence_tokens'],
69
  ) # output is [1,K,T]
70
  logging.info(f"inference on one sample take: {time.time() - stime:.4f} sec.")
71
  if type(encoded_frames) == tuple:
@@ -141,7 +141,7 @@ if __name__ == "__main__":
141
  logging.basicConfig(format=formatter, level=logging.INFO)
142
  args = get_args()
143
  # args.device = 'cpu'
144
- args.allowed_repeat_tokens = eval(args.allowed_repeat_tokens)
145
  seed_everything(args.seed)
146
 
147
  # load model
 
1
+ import argparse, pickle, ast
2
  import logging
3
  import os, random
4
  import numpy as np
 
65
  temperature=decode_config['temperature'],
66
  stop_repetition=decode_config['stop_repetition'],
67
  kvcache=decode_config['kvcache'],
68
+ silence_tokens=ast.literal_eval(decode_config['silence_tokens']) if type(decode_config['silence_tokens']) == str else decode_config['silence_tokens'],
69
  ) # output is [1,K,T]
70
  logging.info(f"inference on one sample take: {time.time() - stime:.4f} sec.")
71
  if type(encoded_frames) == tuple:
 
141
  logging.basicConfig(format=formatter, level=logging.INFO)
142
  args = get_args()
143
  # args.device = 'cpu'
144
+ args.allowed_repeat_tokens = ast.literal_eval(args.allowed_repeat_tokens)
145
  seed_everything(args.seed)
146
 
147
  # load model
inference_tts_scale.py CHANGED
@@ -1,4 +1,4 @@
1
- import argparse, pickle
2
  import logging
3
  import os, random
4
  import numpy as np
@@ -69,7 +69,7 @@ def inference_one_sample(model, model_args, phn2num, text_tokenizer, audio_token
69
  temperature=decode_config['temperature'],
70
  stop_repetition=decode_config['stop_repetition'],
71
  kvcache=decode_config['kvcache'],
72
- silence_tokens=eval(decode_config['silence_tokens']) if type(decode_config['silence_tokens'])==str else decode_config['silence_tokens']
73
  ) # output is [1,K,T]
74
  else:
75
  logging.info(f"running inference with batch size {decode_config['sample_batch_size']}, i.e. return the shortest among {decode_config['sample_batch_size']} generations.")
@@ -83,7 +83,7 @@ def inference_one_sample(model, model_args, phn2num, text_tokenizer, audio_token
83
  stop_repetition=decode_config['stop_repetition'],
84
  kvcache=decode_config['kvcache'],
85
  batch_size = decode_config['sample_batch_size'],
86
- silence_tokens=eval(decode_config['silence_tokens']) if type(decode_config['silence_tokens'])==str else decode_config['silence_tokens']
87
  ) # output is [1,K,T]
88
  logging.info(f"inference on one sample take: {time.time() - stime:.4f} sec.")
89
 
 
1
+ import argparse, pickle, ast
2
  import logging
3
  import os, random
4
  import numpy as np
 
69
  temperature=decode_config['temperature'],
70
  stop_repetition=decode_config['stop_repetition'],
71
  kvcache=decode_config['kvcache'],
72
+ silence_tokens=ast.literal_eval(decode_config['silence_tokens']) if type(decode_config['silence_tokens'])==str else decode_config['silence_tokens']
73
  ) # output is [1,K,T]
74
  else:
75
  logging.info(f"running inference with batch size {decode_config['sample_batch_size']}, i.e. return the shortest among {decode_config['sample_batch_size']} generations.")
 
83
  stop_repetition=decode_config['stop_repetition'],
84
  kvcache=decode_config['kvcache'],
85
  batch_size = decode_config['sample_batch_size'],
86
+ silence_tokens=ast.literal_eval(decode_config['silence_tokens']) if type(decode_config['silence_tokens'])==str else decode_config['silence_tokens']
87
  ) # output is [1,K,T]
88
  logging.info(f"inference on one sample take: {time.time() - stime:.4f} sec.")
89
 
packages.txt CHANGED
@@ -1 +1,2 @@
1
- espeak-ng
 
 
1
+ espeak-ng
2
+ ffmpeg
requirements.txt CHANGED
@@ -1,6 +1,12 @@
1
- -e git+https://github.com/facebookresearch/audiocraft.git@f83babff6b5e97f75562127c4cc8122229c8f099#egg=audiocraft
2
- phonemizer==3.2.1
3
- gradio
4
- nltk>=3.8.1
5
- openai-whisper>=20231117
6
- spaces
 
 
 
 
 
 
 
1
+ spaces
2
+ torch==2.9.1
3
+ torchaudio==2.9.1
4
+ numpy
5
+ einops
6
+ omegaconf
7
+ torchmetrics
8
+ phonemizer>=3.2.1
9
+ nltk>=3.9
10
+ openai-whisper>=20250625
11
+ soundfile
12
+ huggingface_hub