Spaces:
Paused
Paused
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 +2 -1
- app.py +83 -57
- audiocraft/LICENSE +21 -0
- audiocraft/__init__.py +2 -0
- audiocraft/models/__init__.py +1 -0
- audiocraft/models/encodec.py +506 -0
- audiocraft/modules/__init__.py +13 -0
- audiocraft/modules/conv.py +243 -0
- audiocraft/modules/lstm.py +25 -0
- audiocraft/modules/seanet.py +258 -0
- audiocraft/quantization/__init__.py +9 -0
- audiocraft/quantization/base.py +99 -0
- audiocraft/quantization/core_vq.py +400 -0
- audiocraft/quantization/vq.py +115 -0
- data/tokenizer.py +61 -6
- inference_speech_editing_scale.py +3 -3
- inference_tts_scale.py +3 -3
- packages.txt +2 -1
- requirements.txt +12 -6
README.md
CHANGED
|
@@ -4,7 +4,8 @@ emoji: 📈
|
|
| 4 |
colorFrom: blue
|
| 5 |
colorTo: red
|
| 6 |
sdk: gradio
|
| 7 |
-
sdk_version:
|
|
|
|
| 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 |
-
|
| 3 |
-
|
| 4 |
import gradio as gr
|
|
|
|
|
|
|
| 5 |
import torch
|
| 6 |
-
import
|
|
|
|
| 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 |
-
|
| 19 |
|
| 20 |
-
|
| 21 |
-
|
| 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 |
-
|
| 36 |
-
|
| 37 |
-
|
| 38 |
-
|
| 39 |
-
|
| 40 |
-
"tokenizer": get_tokenizer(multilingual=False)
|
| 41 |
-
}
|
| 42 |
|
| 43 |
|
| 44 |
-
|
| 45 |
-
|
| 46 |
-
|
| 47 |
-
|
| 48 |
-
|
| 49 |
-
|
| 50 |
-
|
| 51 |
-
|
| 52 |
-
|
| 53 |
-
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
|
|
|
|
| 57 |
model = voicecraft.VoiceCraft(ckpt["config"])
|
| 58 |
model.load_state_dict(ckpt["model"])
|
| 59 |
model.to(device)
|
| 60 |
model.eval()
|
| 61 |
-
|
|
|
|
|
|
|
|
|
|
| 62 |
"ckpt": ckpt,
|
| 63 |
"model": model,
|
| 64 |
"text_tokenizer": TextTokenizer(backend="espeak"),
|
| 65 |
-
"audio_tokenizer":
|
| 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 |
-
|
| 102 |
-
|
| 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 =
|
| 132 |
-
|
| 133 |
-
|
|
|
|
|
|
|
| 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) *
|
| 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 |
-
|
| 217 |
-
max_time = round(
|
| 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=[
|
| 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 |
-
|
| 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 |
-
|
| 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=
|
| 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 =
|
| 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=
|
| 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=
|
| 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 |
-
|
| 2 |
-
|
| 3 |
-
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|