@@ -1,14 +1,15 @@
|
||||
# ComfyUI_SongGeneration
|
||||
[SongGeneration](https://github.com/tencent-ailab/SongGeneration):High-Quality Song Generation with Multi-Preference Alignment (SOTA),you can try VRAM>12G
|
||||
|
||||
# Tips
|
||||
# Update
|
||||
* 10/18 修改加载流程,支持最新的full ,new,large模型,large模型12GVram可能会OOM,修复高版本transformer 的函数错误/Modify the loading process to support the latest full, new, and large models, and fix function errors in higher versions of transformers
|
||||
* 07/29,支持bgm和人声(vocal,目前还是有bgm底噪)单独输出,选择mixed为合成全部,模型加载方式更合理,去掉诸多debug打印,新增save_separate按钮,开启则保存三个音频(bgm,vocal,mixed);
|
||||
* Test env(插件测试环境):window11,python3.11, torch2.6 ,cu124, VR12G,(transformers 4.45.1)
|
||||
|
||||
|
||||
# 1. Installation
|
||||
|
||||
In the ./ComfyUI/custom_node directory, run the following:
|
||||
In the ./ComfyUI/custom_nodes directory, run the following:
|
||||
```
|
||||
git clone https://github.com/smthemex/ComfyUI_SongGeneration.git
|
||||
```
|
||||
@@ -24,28 +25,28 @@ pip install -r requirements.txt
|
||||
```
|
||||
|
||||
# 3.Model
|
||||
* 3.1.1 download ckpt from [tencent/SongGeneration](https://huggingface.co/tencent/SongGeneration/tree/main) 国内建议魔搭[AI-ModelScope/SongGeneration](https://www.modelscope.cn/models/AI-ModelScope/SongGeneration/files)
|
||||
* 3.1.2 download htdemucs.pth [tencent/SongGeneration](https://huggingface.co/tencent/SongGeneration/tree/main/third_party/demucs/ckpt)
|
||||
* 文件结构如下:
|
||||
* 3.1.1 download ckpt from [tencent/SongGeneration](https://huggingface.co/tencent/SongGeneration/tree/main) 国内建议魔搭[AI-ModelScope/SongGeneration](https://www.modelscope.cn/models/AI-ModelScope/SongGeneration/files)
|
||||
* 3.1.2 [new base](https://huggingface.co/lglg666/SongGeneration-base-new),[large ](https://huggingface.co/lglg666/SongGeneration-large),[full](https://huggingface.co/lglg666/SongGeneration-base-full)
|
||||
* 3.1.3 new prompt,[emb](https://github.com/tencent-ailab/SongGeneration/tree/main/tools)
|
||||
* 3.1.4 download htdemucs.pth [tencent/SongGeneration](https://huggingface.co/tencent/SongGeneration/tree/main/third_party/demucs/ckpt)
|
||||
* 文件结构如下,修改了加载流程,原来的结构也能用:
|
||||
```
|
||||
-- ComfyUI/models/SongGeneration/
|
||||
|-- htdemucs.pth #150M
|
||||
|--prompt.pt # 3M
|
||||
|--new_prompt.pt # 3M
|
||||
|--model_2.safetensors
|
||||
|--model_2_fixed.safetensors
|
||||
|--new_model.pt # rename from model.pt
|
||||
|--large_model.pt # rename from model.pt
|
||||
|-- ckpt/ # 24.4G all 整个文件夹的大小
|
||||
|--encode-s12k.pt # 3.68G
|
||||
|--prompt.pt # 3M
|
||||
|--model_1rvq/
|
||||
|--all files # 全部文件
|
||||
|--model_septoken/
|
||||
|--all files # 全部文件
|
||||
|--models--lengyue233--content-vec-best/
|
||||
|--all files # 全部文件
|
||||
|--songgeneration_base/ #注意删掉了_zh notice no ‘_zh’ now
|
||||
|--all files # 全部文件
|
||||
|--vae/
|
||||
|--all files # 全部文件
|
||||
-- ComfyUI/models/vae/
|
||||
|--autoencoder_music_1320k.ckpt
|
||||
```
|
||||
# 4 Example
|
||||

|
||||

|
||||
|
||||
# 5 Citation
|
||||
```
|
||||
|
||||
@@ -21,20 +21,19 @@ from ..modules.conditioners import (
|
||||
ConditionFuser,
|
||||
)
|
||||
|
||||
|
||||
def get_audio_tokenizer_model(checkpoint_path: str, cfg: omegaconf.DictConfig):
|
||||
def get_audio_tokenizer_model(checkpoint_path: str, vae_config,vae_model,mode):
|
||||
from ..tokenizer.audio_tokenizer import AudioTokenizer
|
||||
"""Instantiate a compression model."""
|
||||
if checkpoint_path is None:
|
||||
return None
|
||||
if checkpoint_path.startswith('//pretrained/'):
|
||||
name = checkpoint_path.split('/', 3)[-1]
|
||||
return AudioTokenizer.get_pretrained(name, cfg.vae_config, cfg.vae_model, 'cpu', mode=cfg.mode)
|
||||
return AudioTokenizer.get_pretrained(name, vae_config,vae_model,'cpu', mode=mode)
|
||||
elif checkpoint_path == "":
|
||||
return None
|
||||
else:
|
||||
name = checkpoint_path
|
||||
return AudioTokenizer.get_pretrained(name, cfg.vae_config, cfg.vae_model, 'cpu', mode=cfg.mode)
|
||||
return AudioTokenizer.get_pretrained(name, vae_config,vae_model,'cpu', mode=mode)
|
||||
|
||||
def get_lm_model(cfg: omegaconf.DictConfig): #-> LMModel:
|
||||
"""Instantiate a LM."""
|
||||
|
||||
@@ -15,15 +15,14 @@ from safetensors.torch import load_file
|
||||
class Tango:
|
||||
def __init__(self, \
|
||||
model_path, \
|
||||
vae_config="",
|
||||
vae_model="",
|
||||
vae_config,
|
||||
vae_model,
|
||||
layer_num=6, \
|
||||
device="cuda:0"):
|
||||
device="cuda"):
|
||||
|
||||
self.sample_rate = 48000
|
||||
scheduler_name = "custom_nodes/ComfyUI_SongGeneration/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/configs/scheduler/stable_diffusion_2.1_largenoise_sample.json"
|
||||
self.device = device
|
||||
|
||||
self.vae = get_model(vae_config, vae_model)
|
||||
self.vae = self.vae.to(device)
|
||||
self.vae=self.vae.eval()
|
||||
@@ -43,7 +42,7 @@ class Tango:
|
||||
main_weights = torch.load(model_path, map_location=device)
|
||||
self.model.load_state_dict(main_weights, strict=False)
|
||||
#print ("Successfully loaded checkpoint from:", model_path)
|
||||
|
||||
del main_weights
|
||||
self.model.eval()
|
||||
self.model.init_device_dtype(torch.device(device), torch.float32)
|
||||
#print("scaling factor: ", self.model.normfeat.std)
|
||||
@@ -92,6 +91,7 @@ class Tango:
|
||||
# output = torch.cat([saved_samples.detach().cpu(),audio[0].detach().cpu()],0)
|
||||
# return output
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
@torch.autocast(device_type="cuda", dtype=torch.float32)
|
||||
def sound2code(self, orig_samples, batch_size=3):
|
||||
|
||||
@@ -73,15 +73,14 @@ class Tango:
|
||||
vae_model,
|
||||
layer_vocal=7,\
|
||||
layer_bgm=3,\
|
||||
device="cuda:0"):
|
||||
device="cuda"
|
||||
):
|
||||
|
||||
self.sample_rate = 48000
|
||||
scheduler_name = "configs/scheduler/stable_diffusion_2.1_largenoise_sample.json"
|
||||
self.device = device
|
||||
|
||||
self.vae = get_model(vae_config, vae_model)
|
||||
self.vae = self.vae.to(device)
|
||||
self.vae=self.vae.eval()
|
||||
self.vae=self.vae.to(device).eval()
|
||||
self.layer_vocal=layer_vocal
|
||||
self.layer_bgm=layer_bgm
|
||||
|
||||
@@ -99,7 +98,7 @@ class Tango:
|
||||
main_weights = torch.load(model_path, map_location=device)
|
||||
self.model.load_state_dict(main_weights, strict=False)
|
||||
#print ("Successfully loaded checkpoint from:", model_path)
|
||||
|
||||
del main_weights
|
||||
self.model.eval()
|
||||
self.model.init_device_dtype(torch.device(device), torch.float32)
|
||||
#print("scaling factor: ", self.model.normfeat.std)
|
||||
@@ -110,7 +109,7 @@ class Tango:
|
||||
# scheduler_name, subfolder="scheduler")
|
||||
#print("Successfully loaded inference scheduler from {}".format(scheduler_name))
|
||||
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
@torch.autocast(device_type="cuda", dtype=torch.float32)
|
||||
def sound2code(self, orig_vocal, orig_bgm, batch_size=8):
|
||||
|
||||
@@ -617,8 +617,6 @@ class PromptCondAudioDiffusion(nn.Module):
|
||||
if(scenario=='other_seg'):
|
||||
latent_masks[:,0:incontext_length] = 1
|
||||
|
||||
|
||||
|
||||
quantized_bestrq_emb = (latent_masks > 0.5).unsqueeze(-1) * quantized_bestrq_emb \
|
||||
+ (latent_masks < 0.5).unsqueeze(-1) * self.zero_cond_embedding1.reshape(1,1,1024)
|
||||
quantized_bestrq_emb_bgm = (latent_masks > 0.5).unsqueeze(-1) * quantized_bestrq_emb_bgm \
|
||||
|
||||
+1
@@ -42,6 +42,7 @@ try:
|
||||
from transformers.models.gpt2.modeling_gpt2 import GPT2SequenceSummary as SequenceSummary
|
||||
except:
|
||||
from transformers.modeling_utils import SequenceSummary
|
||||
|
||||
from transformers.pytorch_utils import Conv1D, find_pruneable_heads_and_indices, prune_conv1d_layer
|
||||
from transformers.utils import (
|
||||
ModelOutput,
|
||||
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
@@ -12,4 +12,5 @@ def get_model(model_config, path):
|
||||
state_dict = torch.load(path, map_location='cpu')
|
||||
model = create_autoencoder_from_config(model_config)
|
||||
model.load_state_dict(state_dict['state_dict'], strict=False)
|
||||
del state_dict
|
||||
return model
|
||||
|
||||
@@ -75,9 +75,9 @@ class AudioTokenizer(ABC, nn.Module):
|
||||
@staticmethod
|
||||
def get_pretrained(
|
||||
name: str,
|
||||
vae_config: str,
|
||||
vae_model: str,
|
||||
device: tp.Union[torch.device, str] = 'cpu',
|
||||
vae_config,
|
||||
vae_model,
|
||||
device: tp.Union[torch.device, str] = 'cuda',
|
||||
mode='extract'
|
||||
) -> 'AudioTokenizer':
|
||||
"""Instantiate a AudioTokenizer model from a given pretrained model.
|
||||
@@ -91,11 +91,11 @@ class AudioTokenizer(ABC, nn.Module):
|
||||
if name.split('_')[0] == 'Flow1dVAESeparate':
|
||||
model_type = name.split('_', 1)[1]
|
||||
#logger.info("Getting pretrained compression model from semantic model %s", model_type)
|
||||
model = Flow1dVAESeparate(model_type, vae_config, vae_model)
|
||||
model = Flow1dVAESeparate(model_type, vae_config,vae_model)
|
||||
elif name.split('_')[0] == 'Flow1dVAE1rvq':
|
||||
model_type = name.split('_', 1)[1]
|
||||
#logger.info("Getting pretrained compression model from semantic model %s", model_type)
|
||||
model = Flow1dVAE1rvq(model_type, vae_config, vae_model)
|
||||
model = Flow1dVAE1rvq(model_type, vae_config,vae_model)
|
||||
else:
|
||||
raise NotImplementedError("{} is not implemented in models/audio_tokenizer.py".format(
|
||||
name))
|
||||
@@ -106,19 +106,21 @@ class Flow1dVAE1rvq(AudioTokenizer):
|
||||
def __init__(
|
||||
self,
|
||||
model_type: str = "model_2_fixed.safetensors",
|
||||
vae_config: str = "",
|
||||
vae_model: str = "",
|
||||
vae_config:str = "",
|
||||
vae_model:str = "",
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
from .Flow1dVAE.generate_1rvq import Tango
|
||||
model_path = model_type
|
||||
self.model = Tango(model_path=model_path, vae_config=vae_config, vae_model=vae_model, device='cuda')
|
||||
self.model = Tango(model_path=model_path, vae_config=vae_config, vae_model=vae_model,device='cuda')
|
||||
#print ("Successfully loaded checkpoint from:", model_path)
|
||||
|
||||
|
||||
self.n_quantizers = 1
|
||||
|
||||
|
||||
|
||||
def forward(self, x: torch.Tensor) :
|
||||
# We don't support training with this.
|
||||
raise NotImplementedError("Forward and training with DAC not supported.")
|
||||
@@ -181,19 +183,21 @@ class Flow1dVAESeparate(AudioTokenizer):
|
||||
def __init__(
|
||||
self,
|
||||
model_type: str = "model_2.safetensors",
|
||||
vae_config: str = "",
|
||||
vae_model: str = "",
|
||||
vae_config:str = "",
|
||||
vae_model:str = "",
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
from .Flow1dVAE.generate_septoken import Tango
|
||||
model_path = model_type
|
||||
self.model = Tango(model_path=model_path, vae_config=vae_config, vae_model=vae_model, device='cuda')
|
||||
self.model = Tango(model_path=model_path, vae_config=vae_config, vae_model=vae_model,device='cuda')
|
||||
#print ("Successfully loaded checkpoint from:", model_path)
|
||||
|
||||
|
||||
self.n_quantizers = 1
|
||||
|
||||
|
||||
|
||||
def forward(self, x: torch.Tensor) :
|
||||
# We don't support training with this.
|
||||
raise NotImplementedError("Forward and training with DAC not supported.")
|
||||
|
||||
@@ -0,0 +1,141 @@
|
||||
# ================ Train Config ================ #
|
||||
lyric_processor:
|
||||
max_dur: 150
|
||||
min_dur: 30
|
||||
prompt_len: 10
|
||||
pad_to_max: true
|
||||
|
||||
|
||||
# ================ Audio tokenzier ================ #
|
||||
audio_tokenizer_checkpoint: Flow1dVAE1rvq_./ckpt/model_1rvq/model_2_fixed.safetensors
|
||||
audio_tokenizer_frame_rate: 25
|
||||
audio_tokenizer_code_depth: 1
|
||||
sample_rate: 48000
|
||||
|
||||
audio_tokenizer_checkpoint_sep: Flow1dVAESeparate_./ckpt/model_septoken/model_2.safetensors
|
||||
audio_tokenizer_frame_rate_sep: 25
|
||||
audio_tokenizer_code_depth_sep: 2
|
||||
sample_rate_sep: 48000
|
||||
|
||||
# ================ VAE ================ #
|
||||
vae_config: ./ckpt/vae/stable_audio_1920_vae.json
|
||||
vae_model: ./ckpt/vae/autoencoder_music_1320k.ckpt
|
||||
|
||||
# ================== LM =========================== #
|
||||
lm:
|
||||
lm_type: Llama # [Llama]
|
||||
dim: 1536
|
||||
intermediate_size: 8960
|
||||
num_heads: 12
|
||||
num_layers: 28
|
||||
num_layers_sub: 12
|
||||
code_depth: 3
|
||||
code_size: 16384
|
||||
max_position_embeddings: 8196
|
||||
max_position_embeddings_sub: 10000
|
||||
rope_theta: 100000.0
|
||||
rope_theta_sub: 500000.0
|
||||
dropout: 0.0
|
||||
use_flash_attn_2: true
|
||||
activation: gelu
|
||||
norm_first: true
|
||||
bias_ff: false
|
||||
bias_attn: false
|
||||
causal: true
|
||||
custom: false
|
||||
memory_efficient: true
|
||||
attention_as_float32: false
|
||||
layer_scale: null
|
||||
positional_embedding: sin
|
||||
xpos: false
|
||||
checkpointing: torch
|
||||
weight_init: gaussian
|
||||
depthwise_init: current
|
||||
zero_bias_init: true
|
||||
norm: layer_norm
|
||||
cross_attention: false
|
||||
qk_layer_norm: false
|
||||
qk_layer_norm_cross: false
|
||||
attention_dropout: null
|
||||
kv_repeat: 1
|
||||
|
||||
codebooks_pattern:
|
||||
modeling: delay
|
||||
delay:
|
||||
delays: [ 0, 250, 250 ]
|
||||
flatten_first: 0
|
||||
empty_initial: 0
|
||||
|
||||
# ================ Conditioners ===================== #
|
||||
classifier_free_guidance:
|
||||
# drop all conditions simultaneously
|
||||
training_dropout: 0.15
|
||||
inference_coef: 1.5
|
||||
|
||||
attribute_dropout:
|
||||
# drop each condition separately
|
||||
args:
|
||||
active_on_eval: false
|
||||
text:
|
||||
description: 0.0
|
||||
type_info: 0.5
|
||||
audio:
|
||||
prompt_audio: 0.0
|
||||
|
||||
|
||||
use_text_training: True
|
||||
fuser:
|
||||
sum: []
|
||||
prepend: [ description, prompt_audio, type_info ] # this order is the SAME with the input concatenation order
|
||||
|
||||
conditioners:
|
||||
prompt_audio:
|
||||
model: qt_embedding
|
||||
qt_embedding:
|
||||
code_size: 16384
|
||||
code_depth: 3
|
||||
max_len: ${eval:${prompt_len}*${audio_tokenizer_frame_rate}+2} # 25*10+2+1
|
||||
description:
|
||||
model: QwTokenizer
|
||||
QwTokenizer:
|
||||
token_path: third_party/Qwen2-7B
|
||||
max_len: 300
|
||||
add_token_list: ${load_yaml:conf/vocab.yaml}
|
||||
type_info:
|
||||
model: QwTextTokenizer
|
||||
QwTextTokenizer:
|
||||
token_path: third_party/Qwen2-7B
|
||||
max_len: 50
|
||||
|
||||
offload:
|
||||
audiolm:
|
||||
offload_module: self
|
||||
cpu_mem_gb: 0
|
||||
pre_copy_step: 1
|
||||
clean_cache_after_forward: false
|
||||
dtype: torch.float16
|
||||
offload_layer_dict:
|
||||
transformer: 4
|
||||
transformer2: 4
|
||||
ignore_layer_list: []
|
||||
clean_cache_wrapper:
|
||||
module: self
|
||||
method_name: _sample_next_token
|
||||
diff_mem_gb_thre: 2
|
||||
debug: false
|
||||
|
||||
wav_tokenizer_diffusion:
|
||||
offload_module: self.model.model
|
||||
pre_copy_step: 1
|
||||
clean_cache_after_forward: false
|
||||
cpu_mem_gb: -1
|
||||
dtype: null
|
||||
offload_layer_dict:
|
||||
cfm_wrapper: 5
|
||||
hubert: 4
|
||||
ignore_layer_list: []
|
||||
clean_cache_wrapper:
|
||||
module: self.model.model.cfm_wrapper.estimator
|
||||
method_name: forward
|
||||
diff_mem_gb_thre: 1
|
||||
debug: false
|
||||
@@ -0,0 +1,141 @@
|
||||
# ================ Train Config ================ #
|
||||
lyric_processor:
|
||||
max_dur: 270
|
||||
min_dur: 30
|
||||
prompt_len: 10
|
||||
pad_to_max: true
|
||||
|
||||
|
||||
# ================ Audio tokenzier ================ #
|
||||
audio_tokenizer_checkpoint: Flow1dVAE1rvq_./ckpt/model_1rvq/model_2_fixed.safetensors
|
||||
audio_tokenizer_frame_rate: 25
|
||||
audio_tokenizer_code_depth: 1
|
||||
sample_rate: 48000
|
||||
|
||||
audio_tokenizer_checkpoint_sep: Flow1dVAESeparate_./ckpt/model_septoken/model_2.safetensors
|
||||
audio_tokenizer_frame_rate_sep: 25
|
||||
audio_tokenizer_code_depth_sep: 2
|
||||
sample_rate_sep: 48000
|
||||
|
||||
# ================ VAE ================ #
|
||||
vae_config: ./ckpt/vae/stable_audio_1920_vae.json
|
||||
vae_model: ./ckpt/vae/autoencoder_music_1320k.ckpt
|
||||
|
||||
# ================== LM =========================== #
|
||||
lm:
|
||||
lm_type: Llama # [Llama]
|
||||
dim: 1536
|
||||
intermediate_size: 8960
|
||||
num_heads: 12
|
||||
num_layers: 28
|
||||
num_layers_sub: 12
|
||||
code_depth: 3
|
||||
code_size: 16384
|
||||
max_position_embeddings: 10000
|
||||
max_position_embeddings_sub: 10000
|
||||
rope_theta: 500000.0
|
||||
rope_theta_sub: 500000.0
|
||||
dropout: 0.0
|
||||
use_flash_attn_2: true
|
||||
activation: gelu
|
||||
norm_first: true
|
||||
bias_ff: false
|
||||
bias_attn: false
|
||||
causal: true
|
||||
custom: false
|
||||
memory_efficient: true
|
||||
attention_as_float32: false
|
||||
layer_scale: null
|
||||
positional_embedding: sin
|
||||
xpos: false
|
||||
checkpointing: torch
|
||||
weight_init: gaussian
|
||||
depthwise_init: current
|
||||
zero_bias_init: true
|
||||
norm: layer_norm
|
||||
cross_attention: false
|
||||
qk_layer_norm: false
|
||||
qk_layer_norm_cross: false
|
||||
attention_dropout: null
|
||||
kv_repeat: 1
|
||||
|
||||
codebooks_pattern:
|
||||
modeling: delay
|
||||
delay:
|
||||
delays: [ 0, 250, 250 ]
|
||||
flatten_first: 0
|
||||
empty_initial: 0
|
||||
|
||||
# ================ Conditioners ===================== #
|
||||
classifier_free_guidance:
|
||||
# drop all conditions simultaneously
|
||||
training_dropout: 0.15
|
||||
inference_coef: 1.5
|
||||
|
||||
attribute_dropout:
|
||||
# drop each condition separately
|
||||
args:
|
||||
active_on_eval: false
|
||||
text:
|
||||
description: 0.0
|
||||
type_info: 0.5
|
||||
audio:
|
||||
prompt_audio: 0.5
|
||||
|
||||
|
||||
use_text_training: True
|
||||
fuser:
|
||||
sum: []
|
||||
prepend: [ description, prompt_audio, type_info ] # this order is the SAME with the input concatenation order
|
||||
|
||||
conditioners:
|
||||
prompt_audio:
|
||||
model: qt_embedding
|
||||
qt_embedding:
|
||||
code_size: 16384
|
||||
code_depth: 3
|
||||
max_len: ${eval:${prompt_len}*${audio_tokenizer_frame_rate}+2} # 25*10+2+1
|
||||
description:
|
||||
model: QwTokenizer
|
||||
QwTokenizer:
|
||||
token_path: third_party/Qwen2-7B
|
||||
max_len: 600
|
||||
add_token_list: ${load_yaml:conf/vocab.yaml}
|
||||
type_info:
|
||||
model: QwTextTokenizer
|
||||
QwTextTokenizer:
|
||||
token_path: third_party/Qwen2-7B
|
||||
max_len: 100
|
||||
|
||||
offload:
|
||||
audiolm:
|
||||
offload_module: self
|
||||
cpu_mem_gb: 0
|
||||
pre_copy_step: 1
|
||||
clean_cache_after_forward: false
|
||||
dtype: torch.float16
|
||||
offload_layer_dict:
|
||||
transformer: 4
|
||||
transformer2: 4
|
||||
ignore_layer_list: []
|
||||
clean_cache_wrapper:
|
||||
module: self
|
||||
method_name: _sample_next_token
|
||||
diff_mem_gb_thre: 2
|
||||
debug: false
|
||||
|
||||
wav_tokenizer_diffusion:
|
||||
offload_module: self.model.model
|
||||
pre_copy_step: 1
|
||||
clean_cache_after_forward: false
|
||||
cpu_mem_gb: -1
|
||||
dtype: null
|
||||
offload_layer_dict:
|
||||
cfm_wrapper: 5
|
||||
hubert: 4
|
||||
ignore_layer_list: []
|
||||
clean_cache_wrapper:
|
||||
module: self.model.model.cfm_wrapper.estimator
|
||||
method_name: forward
|
||||
diff_mem_gb_thre: 1
|
||||
debug: false
|
||||
@@ -0,0 +1,141 @@
|
||||
# ================ Train Config ================ #
|
||||
lyric_processor:
|
||||
max_dur: 270
|
||||
min_dur: 30
|
||||
prompt_len: 10
|
||||
pad_to_max: true
|
||||
|
||||
|
||||
# ================ Audio tokenzier ================ #
|
||||
audio_tokenizer_checkpoint: Flow1dVAE1rvq_./ckpt/model_1rvq/model_2_fixed.safetensors
|
||||
audio_tokenizer_frame_rate: 25
|
||||
audio_tokenizer_code_depth: 1
|
||||
sample_rate: 48000
|
||||
|
||||
audio_tokenizer_checkpoint_sep: Flow1dVAESeparate_./ckpt/model_septoken/model_2.safetensors
|
||||
audio_tokenizer_frame_rate_sep: 25
|
||||
audio_tokenizer_code_depth_sep: 2
|
||||
sample_rate_sep: 48000
|
||||
|
||||
# ================ VAE ================ #
|
||||
vae_config: ./ckpt/vae/stable_audio_1920_vae.json
|
||||
vae_model: ./ckpt/vae/autoencoder_music_1320k.ckpt
|
||||
|
||||
# ================== LM =========================== #
|
||||
lm:
|
||||
lm_type: Llama # [Llama]
|
||||
dim: 2048
|
||||
intermediate_size: 11008
|
||||
num_heads: 16
|
||||
num_layers: 36
|
||||
num_layers_sub: 12
|
||||
code_depth: 3
|
||||
code_size: 16384
|
||||
max_position_embeddings: 10000
|
||||
max_position_embeddings_sub: 10000
|
||||
rope_theta: 500000.0
|
||||
rope_theta_sub: 500000.0
|
||||
dropout: 0.0
|
||||
use_flash_attn_2: true
|
||||
activation: gelu
|
||||
norm_first: true
|
||||
bias_ff: false
|
||||
bias_attn: false
|
||||
causal: true
|
||||
custom: false
|
||||
memory_efficient: true
|
||||
attention_as_float32: false
|
||||
layer_scale: null
|
||||
positional_embedding: sin
|
||||
xpos: false
|
||||
checkpointing: torch
|
||||
weight_init: gaussian
|
||||
depthwise_init: current
|
||||
zero_bias_init: true
|
||||
norm: layer_norm
|
||||
cross_attention: false
|
||||
qk_layer_norm: false
|
||||
qk_layer_norm_cross: false
|
||||
attention_dropout: null
|
||||
kv_repeat: 1
|
||||
|
||||
codebooks_pattern:
|
||||
modeling: delay
|
||||
delay:
|
||||
delays: [ 0, 250, 250 ]
|
||||
flatten_first: 0
|
||||
empty_initial: 0
|
||||
|
||||
# ================ Conditioners ===================== #
|
||||
classifier_free_guidance:
|
||||
# drop all conditions simultaneously
|
||||
training_dropout: 0.15
|
||||
inference_coef: 1.5
|
||||
|
||||
attribute_dropout:
|
||||
# drop each condition separately
|
||||
args:
|
||||
active_on_eval: false
|
||||
text:
|
||||
description: 0.0
|
||||
type_info: 0.5
|
||||
audio:
|
||||
prompt_audio: 0.5
|
||||
|
||||
|
||||
use_text_training: True
|
||||
fuser:
|
||||
sum: []
|
||||
prepend: [ description, prompt_audio, type_info ] # this order is the SAME with the input concatenation order
|
||||
|
||||
conditioners:
|
||||
prompt_audio:
|
||||
model: qt_embedding
|
||||
qt_embedding:
|
||||
code_size: 16384
|
||||
code_depth: 3
|
||||
max_len: ${eval:${prompt_len}*${audio_tokenizer_frame_rate}+2} # 25*10+2+1
|
||||
description:
|
||||
model: QwTokenizer
|
||||
QwTokenizer:
|
||||
token_path: third_party/Qwen2-7B
|
||||
max_len: 600
|
||||
add_token_list: ${load_yaml:conf/vocab.yaml}
|
||||
type_info:
|
||||
model: QwTextTokenizer
|
||||
QwTextTokenizer:
|
||||
token_path: third_party/Qwen2-7B
|
||||
max_len: 100
|
||||
|
||||
offload:
|
||||
audiolm:
|
||||
offload_module: self
|
||||
cpu_mem_gb: 0
|
||||
pre_copy_step: 1
|
||||
clean_cache_after_forward: false
|
||||
dtype: torch.float16
|
||||
offload_layer_dict:
|
||||
transformer: 4
|
||||
transformer2: 4
|
||||
ignore_layer_list: []
|
||||
clean_cache_wrapper:
|
||||
module: self
|
||||
method_name: _sample_next_token
|
||||
diff_mem_gb_thre: 5
|
||||
debug: false
|
||||
|
||||
wav_tokenizer_diffusion:
|
||||
offload_module: self.model.model
|
||||
pre_copy_step: 1
|
||||
clean_cache_after_forward: false
|
||||
cpu_mem_gb: -1
|
||||
dtype: null
|
||||
offload_layer_dict:
|
||||
cfm_wrapper: 5
|
||||
hubert: 4
|
||||
ignore_layer_list: []
|
||||
clean_cache_wrapper:
|
||||
module: self.model.model.cfm_wrapper.estimator
|
||||
method_name: forward
|
||||
diff_mem_gb_thre: 1
|
||||
debug: false
|
||||
@@ -0,0 +1,141 @@
|
||||
# ================ Train Config ================ #
|
||||
lyric_processor:
|
||||
max_dur: 150
|
||||
min_dur: 30
|
||||
prompt_len: 10
|
||||
pad_to_max: true
|
||||
|
||||
|
||||
# ================ Audio tokenzier ================ #
|
||||
audio_tokenizer_checkpoint: Flow1dVAE1rvq_./ckpt/model_1rvq/model_2_fixed.safetensors
|
||||
audio_tokenizer_frame_rate: 25
|
||||
audio_tokenizer_code_depth: 1
|
||||
sample_rate: 48000
|
||||
|
||||
audio_tokenizer_checkpoint_sep: Flow1dVAESeparate_./ckpt/model_septoken/model_2.safetensors
|
||||
audio_tokenizer_frame_rate_sep: 25
|
||||
audio_tokenizer_code_depth_sep: 2
|
||||
sample_rate_sep: 48000
|
||||
|
||||
# ================ VAE ================ #
|
||||
vae_config: ./ckpt/vae/stable_audio_1920_vae.json
|
||||
vae_model: ./ckpt/vae/autoencoder_music_1320k.ckpt
|
||||
|
||||
# ================== LM =========================== #
|
||||
lm:
|
||||
lm_type: Llama # [Llama]
|
||||
dim: 1536
|
||||
intermediate_size: 8960
|
||||
num_heads: 12
|
||||
num_layers: 28
|
||||
num_layers_sub: 12
|
||||
code_depth: 3
|
||||
code_size: 16384
|
||||
max_position_embeddings: 8196
|
||||
max_position_embeddings_sub: 10000
|
||||
rope_theta: 100000.0
|
||||
rope_theta_sub: 500000.0
|
||||
dropout: 0.0
|
||||
use_flash_attn_2: true
|
||||
activation: gelu
|
||||
norm_first: true
|
||||
bias_ff: false
|
||||
bias_attn: false
|
||||
causal: true
|
||||
custom: false
|
||||
memory_efficient: true
|
||||
attention_as_float32: false
|
||||
layer_scale: null
|
||||
positional_embedding: sin
|
||||
xpos: false
|
||||
checkpointing: torch
|
||||
weight_init: gaussian
|
||||
depthwise_init: current
|
||||
zero_bias_init: true
|
||||
norm: layer_norm
|
||||
cross_attention: false
|
||||
qk_layer_norm: false
|
||||
qk_layer_norm_cross: false
|
||||
attention_dropout: null
|
||||
kv_repeat: 1
|
||||
|
||||
codebooks_pattern:
|
||||
modeling: delay
|
||||
delay:
|
||||
delays: [ 0, 250, 250 ]
|
||||
flatten_first: 0
|
||||
empty_initial: 0
|
||||
|
||||
# ================ Conditioners ===================== #
|
||||
classifier_free_guidance:
|
||||
# drop all conditions simultaneously
|
||||
training_dropout: 0.15
|
||||
inference_coef: 1.5
|
||||
|
||||
attribute_dropout:
|
||||
# drop each condition separately
|
||||
args:
|
||||
active_on_eval: false
|
||||
text:
|
||||
description: 0.0
|
||||
type_info: 0.5
|
||||
audio:
|
||||
prompt_audio: 0.0
|
||||
|
||||
|
||||
use_text_training: True
|
||||
fuser:
|
||||
sum: []
|
||||
prepend: [ description, prompt_audio, type_info ] # this order is the SAME with the input concatenation order
|
||||
|
||||
conditioners:
|
||||
prompt_audio:
|
||||
model: qt_embedding
|
||||
qt_embedding:
|
||||
code_size: 16384
|
||||
code_depth: 3
|
||||
max_len: ${eval:${prompt_len}*${audio_tokenizer_frame_rate}+2} # 25*10+2+1
|
||||
description:
|
||||
model: QwTokenizer
|
||||
QwTokenizer:
|
||||
token_path: third_party/Qwen2-7B
|
||||
max_len: 300
|
||||
add_token_list: ${load_yaml:conf/vocab.yaml}
|
||||
type_info:
|
||||
model: QwTextTokenizer
|
||||
QwTextTokenizer:
|
||||
token_path: third_party/Qwen2-7B
|
||||
max_len: 50
|
||||
|
||||
offload:
|
||||
audiolm:
|
||||
offload_module: self
|
||||
cpu_mem_gb: 0
|
||||
pre_copy_step: 1
|
||||
clean_cache_after_forward: false
|
||||
dtype: torch.float16
|
||||
offload_layer_dict:
|
||||
transformer: 4
|
||||
transformer2: 4
|
||||
ignore_layer_list: []
|
||||
clean_cache_wrapper:
|
||||
module: self
|
||||
method_name: _sample_next_token
|
||||
diff_mem_gb_thre: 2
|
||||
debug: false
|
||||
|
||||
wav_tokenizer_diffusion:
|
||||
offload_module: self.model.model
|
||||
pre_copy_step: 1
|
||||
clean_cache_after_forward: false
|
||||
cpu_mem_gb: -1
|
||||
dtype: null
|
||||
offload_layer_dict:
|
||||
cfm_wrapper: 5
|
||||
hubert: 4
|
||||
ignore_layer_list: []
|
||||
clean_cache_wrapper:
|
||||
module: self.model.model.cfm_wrapper.estimator
|
||||
method_name: forward
|
||||
diff_mem_gb_thre: 1
|
||||
debug: false
|
||||
@@ -0,0 +1,122 @@
|
||||
{
|
||||
"model_type": "autoencoder",
|
||||
"sample_size": 403200,
|
||||
"sample_rate": 48000,
|
||||
"audio_channels": 2,
|
||||
"model": {
|
||||
"encoder": {
|
||||
"type": "oobleck",
|
||||
"config": {
|
||||
"in_channels": 2,
|
||||
"channels": 128,
|
||||
"c_mults": [1, 2, 4, 8, 16],
|
||||
"strides": [2, 4, 4, 6, 10],
|
||||
"latent_dim": 128,
|
||||
"use_snake": true
|
||||
}
|
||||
},
|
||||
"decoder": {
|
||||
"type": "oobleck",
|
||||
"config": {
|
||||
"out_channels": 2,
|
||||
"channels": 128,
|
||||
"c_mults": [1, 2, 4, 8, 16],
|
||||
"strides": [2, 4, 4, 6, 10],
|
||||
"latent_dim": 64,
|
||||
"use_snake": true,
|
||||
"final_tanh": false
|
||||
}
|
||||
},
|
||||
"bottleneck": {
|
||||
"type": "vae"
|
||||
},
|
||||
"latent_dim": 64,
|
||||
"downsampling_ratio": 1920,
|
||||
"io_channels": 2
|
||||
},
|
||||
"training": {
|
||||
"learning_rate": 1.5e-4,
|
||||
"warmup_steps": 0,
|
||||
"use_ema": true,
|
||||
"optimizer_configs": {
|
||||
"autoencoder": {
|
||||
"optimizer": {
|
||||
"type": "AdamW",
|
||||
"config": {
|
||||
"betas": [0.8, 0.99],
|
||||
"lr": 1.5e-4,
|
||||
"weight_decay": 1e-3
|
||||
}
|
||||
},
|
||||
"scheduler": {
|
||||
"type": "InverseLR",
|
||||
"config": {
|
||||
"inv_gamma": 200000,
|
||||
"power": 0.5,
|
||||
"warmup": 0.999
|
||||
}
|
||||
}
|
||||
},
|
||||
"discriminator": {
|
||||
"optimizer": {
|
||||
"type": "AdamW",
|
||||
"config": {
|
||||
"betas": [0.8, 0.99],
|
||||
"lr": 3e-4,
|
||||
"weight_decay": 1e-3
|
||||
}
|
||||
},
|
||||
"scheduler": {
|
||||
"type": "InverseLR",
|
||||
"config": {
|
||||
"inv_gamma": 200000,
|
||||
"power": 0.5,
|
||||
"warmup": 0.999
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
"loss_configs": {
|
||||
"discriminator": {
|
||||
"type": "encodec",
|
||||
"config": {
|
||||
"filters": 64,
|
||||
"n_ffts": [2048, 1024, 512, 256, 128],
|
||||
"hop_lengths": [512, 256, 128, 64, 32],
|
||||
"win_lengths": [2048, 1024, 512, 256, 128]
|
||||
},
|
||||
"weights": {
|
||||
"adversarial": 0.1,
|
||||
"feature_matching": 5.0
|
||||
}
|
||||
},
|
||||
"spectral": {
|
||||
"type": "mrstft",
|
||||
"config": {
|
||||
"fft_sizes": [2048, 1024, 512, 256, 128, 64, 32],
|
||||
"hop_sizes": [512, 256, 128, 64, 32, 16, 8],
|
||||
"win_lengths": [2048, 1024, 512, 256, 128, 64, 32],
|
||||
"perceptual_weighting": true
|
||||
},
|
||||
"weights": {
|
||||
"mrstft": 1.0
|
||||
}
|
||||
},
|
||||
"time": {
|
||||
"type": "l1",
|
||||
"weights": {
|
||||
"l1": 0.0
|
||||
}
|
||||
},
|
||||
"bottleneck": {
|
||||
"type": "kl",
|
||||
"weights": {
|
||||
"kl": 1e-4
|
||||
}
|
||||
}
|
||||
},
|
||||
"demo": {
|
||||
"demo_every": 2000
|
||||
}
|
||||
}
|
||||
}
|
||||
+16
@@ -0,0 +1,16 @@
|
||||
__version__ = "1.0.0"
|
||||
|
||||
# preserved here for legacy reasons
|
||||
__model_version__ = "latest"
|
||||
|
||||
import audiotools
|
||||
|
||||
audiotools.ml.BaseModel.INTERN += ["dac.**"]
|
||||
audiotools.ml.BaseModel.EXTERN += ["einops"]
|
||||
|
||||
|
||||
from . import nn
|
||||
from . import model
|
||||
from . import utils
|
||||
from .model import DAC
|
||||
from .model import DACFile
|
||||
+36
@@ -0,0 +1,36 @@
|
||||
import sys
|
||||
|
||||
import argbind
|
||||
|
||||
from dac.utils import download
|
||||
from dac.utils.decode import decode
|
||||
from dac.utils.encode import encode
|
||||
|
||||
STAGES = ["encode", "decode", "download"]
|
||||
|
||||
|
||||
def run(stage: str):
|
||||
"""Run stages.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
stage : str
|
||||
Stage to run
|
||||
"""
|
||||
if stage not in STAGES:
|
||||
raise ValueError(f"Unknown command: {stage}. Allowed commands are {STAGES}")
|
||||
stage_fn = globals()[stage]
|
||||
|
||||
if stage == "download":
|
||||
stage_fn()
|
||||
return
|
||||
|
||||
stage_fn()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
group = sys.argv.pop(1)
|
||||
args = argbind.parse_args(group=group)
|
||||
|
||||
with argbind.scope(args):
|
||||
run(group)
|
||||
@@ -0,0 +1 @@
|
||||
|
||||
@@ -0,0 +1,54 @@
|
||||
import torch
|
||||
from audiotools import AudioSignal
|
||||
from audiotools.ml import BaseModel
|
||||
from encodec import EncodecModel
|
||||
|
||||
|
||||
class Encodec(BaseModel):
|
||||
def __init__(self, sample_rate: int = 24000, bandwidth: float = 24.0):
|
||||
super().__init__()
|
||||
|
||||
if sample_rate == 24000:
|
||||
self.model = EncodecModel.encodec_model_24khz()
|
||||
else:
|
||||
self.model = EncodecModel.encodec_model_48khz()
|
||||
self.model.set_target_bandwidth(bandwidth)
|
||||
self.sample_rate = 44100
|
||||
|
||||
def forward(
|
||||
self,
|
||||
audio_data: torch.Tensor,
|
||||
sample_rate: int = 44100,
|
||||
n_quantizers: int = None,
|
||||
):
|
||||
signal = AudioSignal(audio_data, sample_rate)
|
||||
signal.resample(self.model.sample_rate)
|
||||
recons = self.model(signal.audio_data)
|
||||
recons = AudioSignal(recons, self.model.sample_rate)
|
||||
recons.resample(sample_rate)
|
||||
return {"audio": recons.audio_data}
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import numpy as np
|
||||
from functools import partial
|
||||
|
||||
model = Encodec()
|
||||
|
||||
for n, m in model.named_modules():
|
||||
o = m.extra_repr()
|
||||
p = sum([np.prod(p.size()) for p in m.parameters()])
|
||||
fn = lambda o, p: o + f" {p/1e6:<.3f}M params."
|
||||
setattr(m, "extra_repr", partial(fn, o=o, p=p))
|
||||
print(model)
|
||||
print("Total # of params: ", sum([np.prod(p.size()) for p in model.parameters()]))
|
||||
|
||||
length = 88200 * 2
|
||||
x = torch.randn(1, 1, length).to(model.device)
|
||||
x.requires_grad_(True)
|
||||
x.retain_grad()
|
||||
|
||||
# Make a forward pass
|
||||
out = model(x)["audio"]
|
||||
|
||||
print(x.shape, out.shape)
|
||||
@@ -0,0 +1,4 @@
|
||||
from .base import CodecMixin
|
||||
from .base import DACFile
|
||||
from .dac import DAC
|
||||
from .discriminator import Discriminator
|
||||
+294
@@ -0,0 +1,294 @@
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import tqdm
|
||||
from audiotools import AudioSignal
|
||||
from torch import nn
|
||||
|
||||
SUPPORTED_VERSIONS = ["1.0.0"]
|
||||
|
||||
|
||||
@dataclass
|
||||
class DACFile:
|
||||
codes: torch.Tensor
|
||||
|
||||
# Metadata
|
||||
chunk_length: int
|
||||
original_length: int
|
||||
input_db: float
|
||||
channels: int
|
||||
sample_rate: int
|
||||
padding: bool
|
||||
dac_version: str
|
||||
|
||||
def save(self, path):
|
||||
artifacts = {
|
||||
"codes": self.codes.numpy().astype(np.uint16),
|
||||
"metadata": {
|
||||
"input_db": self.input_db.numpy().astype(np.float32),
|
||||
"original_length": self.original_length,
|
||||
"sample_rate": self.sample_rate,
|
||||
"chunk_length": self.chunk_length,
|
||||
"channels": self.channels,
|
||||
"padding": self.padding,
|
||||
"dac_version": SUPPORTED_VERSIONS[-1],
|
||||
},
|
||||
}
|
||||
path = Path(path).with_suffix(".dac")
|
||||
with open(path, "wb") as f:
|
||||
np.save(f, artifacts)
|
||||
return path
|
||||
|
||||
@classmethod
|
||||
def load(cls, path):
|
||||
artifacts = np.load(path, allow_pickle=True)[()]
|
||||
codes = torch.from_numpy(artifacts["codes"].astype(int))
|
||||
if artifacts["metadata"].get("dac_version", None) not in SUPPORTED_VERSIONS:
|
||||
raise RuntimeError(
|
||||
f"Given file {path} can't be loaded with this version of descript-audio-codec."
|
||||
)
|
||||
return cls(codes=codes, **artifacts["metadata"])
|
||||
|
||||
|
||||
class CodecMixin:
|
||||
@property
|
||||
def padding(self):
|
||||
if not hasattr(self, "_padding"):
|
||||
self._padding = True
|
||||
return self._padding
|
||||
|
||||
@padding.setter
|
||||
def padding(self, value):
|
||||
assert isinstance(value, bool)
|
||||
|
||||
layers = [
|
||||
l for l in self.modules() if isinstance(l, (nn.Conv1d, nn.ConvTranspose1d))
|
||||
]
|
||||
|
||||
for layer in layers:
|
||||
if value:
|
||||
if hasattr(layer, "original_padding"):
|
||||
layer.padding = layer.original_padding
|
||||
else:
|
||||
layer.original_padding = layer.padding
|
||||
layer.padding = tuple(0 for _ in range(len(layer.padding)))
|
||||
|
||||
self._padding = value
|
||||
|
||||
def get_delay(self):
|
||||
# Any number works here, delay is invariant to input length
|
||||
l_out = self.get_output_length(0)
|
||||
L = l_out
|
||||
|
||||
layers = []
|
||||
for layer in self.modules():
|
||||
if isinstance(layer, (nn.Conv1d, nn.ConvTranspose1d)):
|
||||
layers.append(layer)
|
||||
|
||||
for layer in reversed(layers):
|
||||
d = layer.dilation[0]
|
||||
k = layer.kernel_size[0]
|
||||
s = layer.stride[0]
|
||||
|
||||
if isinstance(layer, nn.ConvTranspose1d):
|
||||
L = ((L - d * (k - 1) - 1) / s) + 1
|
||||
elif isinstance(layer, nn.Conv1d):
|
||||
L = (L - 1) * s + d * (k - 1) + 1
|
||||
|
||||
L = math.ceil(L)
|
||||
|
||||
l_in = L
|
||||
|
||||
return (l_in - l_out) // 2
|
||||
|
||||
def get_output_length(self, input_length):
|
||||
L = input_length
|
||||
# Calculate output length
|
||||
for layer in self.modules():
|
||||
if isinstance(layer, (nn.Conv1d, nn.ConvTranspose1d)):
|
||||
d = layer.dilation[0]
|
||||
k = layer.kernel_size[0]
|
||||
s = layer.stride[0]
|
||||
|
||||
if isinstance(layer, nn.Conv1d):
|
||||
L = ((L - d * (k - 1) - 1) / s) + 1
|
||||
elif isinstance(layer, nn.ConvTranspose1d):
|
||||
L = (L - 1) * s + d * (k - 1) + 1
|
||||
|
||||
L = math.floor(L)
|
||||
return L
|
||||
|
||||
@torch.no_grad()
|
||||
def compress(
|
||||
self,
|
||||
audio_path_or_signal: Union[str, Path, AudioSignal],
|
||||
win_duration: float = 1.0,
|
||||
verbose: bool = False,
|
||||
normalize_db: float = -16,
|
||||
n_quantizers: int = None,
|
||||
) -> DACFile:
|
||||
"""Processes an audio signal from a file or AudioSignal object into
|
||||
discrete codes. This function processes the signal in short windows,
|
||||
using constant GPU memory.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
audio_path_or_signal : Union[str, Path, AudioSignal]
|
||||
audio signal to reconstruct
|
||||
win_duration : float, optional
|
||||
window duration in seconds, by default 5.0
|
||||
verbose : bool, optional
|
||||
by default False
|
||||
normalize_db : float, optional
|
||||
normalize db, by default -16
|
||||
|
||||
Returns
|
||||
-------
|
||||
DACFile
|
||||
Object containing compressed codes and metadata
|
||||
required for decompression
|
||||
"""
|
||||
audio_signal = audio_path_or_signal
|
||||
if isinstance(audio_signal, (str, Path)):
|
||||
audio_signal = AudioSignal.load_from_file_with_ffmpeg(str(audio_signal))
|
||||
|
||||
self.eval()
|
||||
original_padding = self.padding
|
||||
original_device = audio_signal.device
|
||||
|
||||
audio_signal = audio_signal.clone()
|
||||
original_sr = audio_signal.sample_rate
|
||||
|
||||
resample_fn = audio_signal.resample
|
||||
loudness_fn = audio_signal.loudness
|
||||
|
||||
# If audio is > 10 minutes long, use the ffmpeg versions
|
||||
if audio_signal.signal_duration >= 10 * 60 * 60:
|
||||
resample_fn = audio_signal.ffmpeg_resample
|
||||
loudness_fn = audio_signal.ffmpeg_loudness
|
||||
|
||||
original_length = audio_signal.signal_length
|
||||
resample_fn(self.sample_rate)
|
||||
input_db = loudness_fn()
|
||||
|
||||
if normalize_db is not None:
|
||||
audio_signal.normalize(normalize_db)
|
||||
audio_signal.ensure_max_of_audio()
|
||||
|
||||
nb, nac, nt = audio_signal.audio_data.shape
|
||||
audio_signal.audio_data = audio_signal.audio_data.reshape(nb * nac, 1, nt)
|
||||
win_duration = (
|
||||
audio_signal.signal_duration if win_duration is None else win_duration
|
||||
)
|
||||
|
||||
if audio_signal.signal_duration <= win_duration:
|
||||
# Unchunked compression (used if signal length < win duration)
|
||||
self.padding = True
|
||||
n_samples = nt
|
||||
hop = nt
|
||||
else:
|
||||
# Chunked inference
|
||||
self.padding = False
|
||||
# Zero-pad signal on either side by the delay
|
||||
audio_signal.zero_pad(self.delay, self.delay)
|
||||
n_samples = int(win_duration * self.sample_rate)
|
||||
# Round n_samples to nearest hop length multiple
|
||||
n_samples = int(math.ceil(n_samples / self.hop_length) * self.hop_length)
|
||||
hop = self.get_output_length(n_samples)
|
||||
|
||||
codes = []
|
||||
range_fn = range if not verbose else tqdm.trange
|
||||
|
||||
for i in range_fn(0, nt, hop):
|
||||
x = audio_signal[..., i : i + n_samples]
|
||||
x = x.zero_pad(0, max(0, n_samples - x.shape[-1]))
|
||||
|
||||
audio_data = x.audio_data.to(self.device)
|
||||
audio_data = self.preprocess(audio_data, self.sample_rate)
|
||||
_, c, _, _, _ = self.encode(audio_data, n_quantizers)
|
||||
codes.append(c.to(original_device))
|
||||
chunk_length = c.shape[-1]
|
||||
|
||||
codes = torch.cat(codes, dim=-1)
|
||||
|
||||
dac_file = DACFile(
|
||||
codes=codes,
|
||||
chunk_length=chunk_length,
|
||||
original_length=original_length,
|
||||
input_db=input_db,
|
||||
channels=nac,
|
||||
sample_rate=original_sr,
|
||||
padding=self.padding,
|
||||
dac_version=SUPPORTED_VERSIONS[-1],
|
||||
)
|
||||
|
||||
if n_quantizers is not None:
|
||||
codes = codes[:, :n_quantizers, :]
|
||||
|
||||
self.padding = original_padding
|
||||
return dac_file
|
||||
|
||||
@torch.no_grad()
|
||||
def decompress(
|
||||
self,
|
||||
obj: Union[str, Path, DACFile],
|
||||
verbose: bool = False,
|
||||
) -> AudioSignal:
|
||||
"""Reconstruct audio from a given .dac file
|
||||
|
||||
Parameters
|
||||
----------
|
||||
obj : Union[str, Path, DACFile]
|
||||
.dac file location or corresponding DACFile object.
|
||||
verbose : bool, optional
|
||||
Prints progress if True, by default False
|
||||
|
||||
Returns
|
||||
-------
|
||||
AudioSignal
|
||||
Object with the reconstructed audio
|
||||
"""
|
||||
self.eval()
|
||||
if isinstance(obj, (str, Path)):
|
||||
obj = DACFile.load(obj)
|
||||
|
||||
original_padding = self.padding
|
||||
self.padding = obj.padding
|
||||
|
||||
range_fn = range if not verbose else tqdm.trange
|
||||
codes = obj.codes
|
||||
original_device = codes.device
|
||||
chunk_length = obj.chunk_length
|
||||
recons = []
|
||||
|
||||
for i in range_fn(0, codes.shape[-1], chunk_length):
|
||||
c = codes[..., i : i + chunk_length].to(self.device)
|
||||
z = self.quantizer.from_codes(c)[0]
|
||||
r = self.decode(z)
|
||||
recons.append(r.to(original_device))
|
||||
|
||||
recons = torch.cat(recons, dim=-1)
|
||||
recons = AudioSignal(recons, self.sample_rate)
|
||||
|
||||
resample_fn = recons.resample
|
||||
loudness_fn = recons.loudness
|
||||
|
||||
# If audio is > 10 minutes long, use the ffmpeg versions
|
||||
if recons.signal_duration >= 10 * 60 * 60:
|
||||
resample_fn = recons.ffmpeg_resample
|
||||
loudness_fn = recons.ffmpeg_loudness
|
||||
|
||||
recons.normalize(obj.input_db)
|
||||
resample_fn(obj.sample_rate)
|
||||
recons = recons[..., : obj.original_length]
|
||||
loudness_fn()
|
||||
recons.audio_data = recons.audio_data.reshape(
|
||||
-1, obj.channels, obj.original_length
|
||||
)
|
||||
|
||||
self.padding = original_padding
|
||||
return recons
|
||||
+364
@@ -0,0 +1,364 @@
|
||||
import math
|
||||
from typing import List
|
||||
from typing import Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from audiotools import AudioSignal
|
||||
from audiotools.ml import BaseModel
|
||||
from torch import nn
|
||||
|
||||
from .base import CodecMixin
|
||||
from ..nn.layers import Snake1d
|
||||
from ..nn.layers import WNConv1d
|
||||
from ..nn.layers import WNConvTranspose1d
|
||||
from ..nn.quantize import ResidualVectorQuantize
|
||||
|
||||
|
||||
def init_weights(m):
|
||||
if isinstance(m, nn.Conv1d):
|
||||
nn.init.trunc_normal_(m.weight, std=0.02)
|
||||
nn.init.constant_(m.bias, 0)
|
||||
|
||||
|
||||
class ResidualUnit(nn.Module):
|
||||
def __init__(self, dim: int = 16, dilation: int = 1):
|
||||
super().__init__()
|
||||
pad = ((7 - 1) * dilation) // 2
|
||||
self.block = nn.Sequential(
|
||||
Snake1d(dim),
|
||||
WNConv1d(dim, dim, kernel_size=7, dilation=dilation, padding=pad),
|
||||
Snake1d(dim),
|
||||
WNConv1d(dim, dim, kernel_size=1),
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
y = self.block(x)
|
||||
pad = (x.shape[-1] - y.shape[-1]) // 2
|
||||
if pad > 0:
|
||||
x = x[..., pad:-pad]
|
||||
return x + y
|
||||
|
||||
|
||||
class EncoderBlock(nn.Module):
|
||||
def __init__(self, dim: int = 16, stride: int = 1):
|
||||
super().__init__()
|
||||
self.block = nn.Sequential(
|
||||
ResidualUnit(dim // 2, dilation=1),
|
||||
ResidualUnit(dim // 2, dilation=3),
|
||||
ResidualUnit(dim // 2, dilation=9),
|
||||
Snake1d(dim // 2),
|
||||
WNConv1d(
|
||||
dim // 2,
|
||||
dim,
|
||||
kernel_size=2 * stride,
|
||||
stride=stride,
|
||||
padding=math.ceil(stride / 2),
|
||||
),
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
return self.block(x)
|
||||
|
||||
|
||||
class Encoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
d_model: int = 64,
|
||||
strides: list = [2, 4, 8, 8],
|
||||
d_latent: int = 64,
|
||||
):
|
||||
super().__init__()
|
||||
# Create first convolution
|
||||
self.block = [WNConv1d(1, d_model, kernel_size=7, padding=3)]
|
||||
|
||||
# Create EncoderBlocks that double channels as they downsample by `stride`
|
||||
for stride in strides:
|
||||
d_model *= 2
|
||||
self.block += [EncoderBlock(d_model, stride=stride)]
|
||||
|
||||
# Create last convolution
|
||||
self.block += [
|
||||
Snake1d(d_model),
|
||||
WNConv1d(d_model, d_latent, kernel_size=3, padding=1),
|
||||
]
|
||||
|
||||
# Wrap black into nn.Sequential
|
||||
self.block = nn.Sequential(*self.block)
|
||||
self.enc_dim = d_model
|
||||
|
||||
def forward(self, x):
|
||||
return self.block(x)
|
||||
|
||||
|
||||
class DecoderBlock(nn.Module):
|
||||
def __init__(self, input_dim: int = 16, output_dim: int = 8, stride: int = 1):
|
||||
super().__init__()
|
||||
self.block = nn.Sequential(
|
||||
Snake1d(input_dim),
|
||||
WNConvTranspose1d(
|
||||
input_dim,
|
||||
output_dim,
|
||||
kernel_size=2 * stride,
|
||||
stride=stride,
|
||||
padding=math.ceil(stride / 2),
|
||||
),
|
||||
ResidualUnit(output_dim, dilation=1),
|
||||
ResidualUnit(output_dim, dilation=3),
|
||||
ResidualUnit(output_dim, dilation=9),
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
return self.block(x)
|
||||
|
||||
|
||||
class Decoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
input_channel,
|
||||
channels,
|
||||
rates,
|
||||
d_out: int = 1,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
# Add first conv layer
|
||||
layers = [WNConv1d(input_channel, channels, kernel_size=7, padding=3)]
|
||||
|
||||
# Add upsampling + MRF blocks
|
||||
for i, stride in enumerate(rates):
|
||||
input_dim = channels // 2**i
|
||||
output_dim = channels // 2 ** (i + 1)
|
||||
layers += [DecoderBlock(input_dim, output_dim, stride)]
|
||||
|
||||
# Add final conv layer
|
||||
layers += [
|
||||
Snake1d(output_dim),
|
||||
WNConv1d(output_dim, d_out, kernel_size=7, padding=3),
|
||||
nn.Tanh(),
|
||||
]
|
||||
|
||||
self.model = nn.Sequential(*layers)
|
||||
|
||||
def forward(self, x):
|
||||
return self.model(x)
|
||||
|
||||
|
||||
class DAC(BaseModel, CodecMixin):
|
||||
def __init__(
|
||||
self,
|
||||
encoder_dim: int = 64,
|
||||
encoder_rates: List[int] = [2, 4, 8, 8],
|
||||
latent_dim: int = None,
|
||||
decoder_dim: int = 1536,
|
||||
decoder_rates: List[int] = [8, 8, 4, 2],
|
||||
n_codebooks: int = 9,
|
||||
codebook_size: int = 1024,
|
||||
codebook_dim: Union[int, list] = 8,
|
||||
quantizer_dropout: bool = False,
|
||||
sample_rate: int = 44100,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.encoder_dim = encoder_dim
|
||||
self.encoder_rates = encoder_rates
|
||||
self.decoder_dim = decoder_dim
|
||||
self.decoder_rates = decoder_rates
|
||||
self.sample_rate = sample_rate
|
||||
|
||||
if latent_dim is None:
|
||||
latent_dim = encoder_dim * (2 ** len(encoder_rates))
|
||||
|
||||
self.latent_dim = latent_dim
|
||||
|
||||
self.hop_length = np.prod(encoder_rates)
|
||||
self.encoder = Encoder(encoder_dim, encoder_rates, latent_dim)
|
||||
|
||||
self.n_codebooks = n_codebooks
|
||||
self.codebook_size = codebook_size
|
||||
self.codebook_dim = codebook_dim
|
||||
self.quantizer = ResidualVectorQuantize(
|
||||
input_dim=latent_dim,
|
||||
n_codebooks=n_codebooks,
|
||||
codebook_size=codebook_size,
|
||||
codebook_dim=codebook_dim,
|
||||
quantizer_dropout=quantizer_dropout,
|
||||
)
|
||||
|
||||
self.decoder = Decoder(
|
||||
latent_dim,
|
||||
decoder_dim,
|
||||
decoder_rates,
|
||||
)
|
||||
self.sample_rate = sample_rate
|
||||
self.apply(init_weights)
|
||||
|
||||
self.delay = self.get_delay()
|
||||
|
||||
def preprocess(self, audio_data, sample_rate):
|
||||
if sample_rate is None:
|
||||
sample_rate = self.sample_rate
|
||||
assert sample_rate == self.sample_rate
|
||||
|
||||
length = audio_data.shape[-1]
|
||||
right_pad = math.ceil(length / self.hop_length) * self.hop_length - length
|
||||
audio_data = nn.functional.pad(audio_data, (0, right_pad))
|
||||
|
||||
return audio_data
|
||||
|
||||
def encode(
|
||||
self,
|
||||
audio_data: torch.Tensor,
|
||||
n_quantizers: int = None,
|
||||
):
|
||||
"""Encode given audio data and return quantized latent codes
|
||||
|
||||
Parameters
|
||||
----------
|
||||
audio_data : Tensor[B x 1 x T]
|
||||
Audio data to encode
|
||||
n_quantizers : int, optional
|
||||
Number of quantizers to use, by default None
|
||||
If None, all quantizers are used.
|
||||
|
||||
Returns
|
||||
-------
|
||||
dict
|
||||
A dictionary with the following keys:
|
||||
"z" : Tensor[B x D x T]
|
||||
Quantized continuous representation of input
|
||||
"codes" : Tensor[B x N x T]
|
||||
Codebook indices for each codebook
|
||||
(quantized discrete representation of input)
|
||||
"latents" : Tensor[B x N*D x T]
|
||||
Projected latents (continuous representation of input before quantization)
|
||||
"vq/commitment_loss" : Tensor[1]
|
||||
Commitment loss to train encoder to predict vectors closer to codebook
|
||||
entries
|
||||
"vq/codebook_loss" : Tensor[1]
|
||||
Codebook loss to update the codebook
|
||||
"length" : int
|
||||
Number of samples in input audio
|
||||
"""
|
||||
z = self.encoder(audio_data)
|
||||
z, codes, latents, commitment_loss, codebook_loss = self.quantizer(
|
||||
z, n_quantizers
|
||||
)
|
||||
return z, codes, latents, commitment_loss, codebook_loss
|
||||
|
||||
def decode(self, z: torch.Tensor):
|
||||
"""Decode given latent codes and return audio data
|
||||
|
||||
Parameters
|
||||
----------
|
||||
z : Tensor[B x D x T]
|
||||
Quantized continuous representation of input
|
||||
length : int, optional
|
||||
Number of samples in output audio, by default None
|
||||
|
||||
Returns
|
||||
-------
|
||||
dict
|
||||
A dictionary with the following keys:
|
||||
"audio" : Tensor[B x 1 x length]
|
||||
Decoded audio data.
|
||||
"""
|
||||
return self.decoder(z)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
audio_data: torch.Tensor,
|
||||
sample_rate: int = None,
|
||||
n_quantizers: int = None,
|
||||
):
|
||||
"""Model forward pass
|
||||
|
||||
Parameters
|
||||
----------
|
||||
audio_data : Tensor[B x 1 x T]
|
||||
Audio data to encode
|
||||
sample_rate : int, optional
|
||||
Sample rate of audio data in Hz, by default None
|
||||
If None, defaults to `self.sample_rate`
|
||||
n_quantizers : int, optional
|
||||
Number of quantizers to use, by default None.
|
||||
If None, all quantizers are used.
|
||||
|
||||
Returns
|
||||
-------
|
||||
dict
|
||||
A dictionary with the following keys:
|
||||
"z" : Tensor[B x D x T]
|
||||
Quantized continuous representation of input
|
||||
"codes" : Tensor[B x N x T]
|
||||
Codebook indices for each codebook
|
||||
(quantized discrete representation of input)
|
||||
"latents" : Tensor[B x N*D x T]
|
||||
Projected latents (continuous representation of input before quantization)
|
||||
"vq/commitment_loss" : Tensor[1]
|
||||
Commitment loss to train encoder to predict vectors closer to codebook
|
||||
entries
|
||||
"vq/codebook_loss" : Tensor[1]
|
||||
Codebook loss to update the codebook
|
||||
"length" : int
|
||||
Number of samples in input audio
|
||||
"audio" : Tensor[B x 1 x length]
|
||||
Decoded audio data.
|
||||
"""
|
||||
length = audio_data.shape[-1]
|
||||
audio_data = self.preprocess(audio_data, sample_rate)
|
||||
z, codes, latents, commitment_loss, codebook_loss = self.encode(
|
||||
audio_data, n_quantizers
|
||||
)
|
||||
|
||||
x = self.decode(z)
|
||||
return {
|
||||
"audio": x[..., :length],
|
||||
"z": z,
|
||||
"codes": codes,
|
||||
"latents": latents,
|
||||
"vq/commitment_loss": commitment_loss,
|
||||
"vq/codebook_loss": codebook_loss,
|
||||
}
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import numpy as np
|
||||
from functools import partial
|
||||
|
||||
model = DAC().to("cpu")
|
||||
|
||||
for n, m in model.named_modules():
|
||||
o = m.extra_repr()
|
||||
p = sum([np.prod(p.size()) for p in m.parameters()])
|
||||
fn = lambda o, p: o + f" {p/1e6:<.3f}M params."
|
||||
setattr(m, "extra_repr", partial(fn, o=o, p=p))
|
||||
print(model)
|
||||
print("Total # of params: ", sum([np.prod(p.size()) for p in model.parameters()]))
|
||||
|
||||
length = 88200 * 2
|
||||
x = torch.randn(1, 1, length).to(model.device)
|
||||
x.requires_grad_(True)
|
||||
x.retain_grad()
|
||||
|
||||
# Make a forward pass
|
||||
out = model(x)["audio"]
|
||||
print("Input shape:", x.shape)
|
||||
print("Output shape:", out.shape)
|
||||
|
||||
# Create gradient variable
|
||||
grad = torch.zeros_like(out)
|
||||
grad[:, :, grad.shape[-1] // 2] = 1
|
||||
|
||||
# Make a backward pass
|
||||
out.backward(grad)
|
||||
|
||||
# Check non-zero values
|
||||
gradmap = x.grad.squeeze(0)
|
||||
gradmap = (gradmap != 0).sum(0) # sum across features
|
||||
rf = (gradmap != 0).sum()
|
||||
|
||||
print(f"Receptive field: {rf.item()}")
|
||||
|
||||
x = AudioSignal(torch.randn(1, 1, 44100 * 60), 44100)
|
||||
model.decompress(model.compress(x, verbose=True), verbose=True)
|
||||
@@ -0,0 +1,228 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from audiotools import AudioSignal
|
||||
from audiotools import ml
|
||||
from audiotools import STFTParams
|
||||
from einops import rearrange
|
||||
from torch.nn.utils import weight_norm
|
||||
|
||||
|
||||
def WNConv1d(*args, **kwargs):
|
||||
act = kwargs.pop("act", True)
|
||||
conv = weight_norm(nn.Conv1d(*args, **kwargs))
|
||||
if not act:
|
||||
return conv
|
||||
return nn.Sequential(conv, nn.LeakyReLU(0.1))
|
||||
|
||||
|
||||
def WNConv2d(*args, **kwargs):
|
||||
act = kwargs.pop("act", True)
|
||||
conv = weight_norm(nn.Conv2d(*args, **kwargs))
|
||||
if not act:
|
||||
return conv
|
||||
return nn.Sequential(conv, nn.LeakyReLU(0.1))
|
||||
|
||||
|
||||
class MPD(nn.Module):
|
||||
def __init__(self, period):
|
||||
super().__init__()
|
||||
self.period = period
|
||||
self.convs = nn.ModuleList(
|
||||
[
|
||||
WNConv2d(1, 32, (5, 1), (3, 1), padding=(2, 0)),
|
||||
WNConv2d(32, 128, (5, 1), (3, 1), padding=(2, 0)),
|
||||
WNConv2d(128, 512, (5, 1), (3, 1), padding=(2, 0)),
|
||||
WNConv2d(512, 1024, (5, 1), (3, 1), padding=(2, 0)),
|
||||
WNConv2d(1024, 1024, (5, 1), 1, padding=(2, 0)),
|
||||
]
|
||||
)
|
||||
self.conv_post = WNConv2d(
|
||||
1024, 1, kernel_size=(3, 1), padding=(1, 0), act=False
|
||||
)
|
||||
|
||||
def pad_to_period(self, x):
|
||||
t = x.shape[-1]
|
||||
x = F.pad(x, (0, self.period - t % self.period), mode="reflect")
|
||||
return x
|
||||
|
||||
def forward(self, x):
|
||||
fmap = []
|
||||
|
||||
x = self.pad_to_period(x)
|
||||
x = rearrange(x, "b c (l p) -> b c l p", p=self.period)
|
||||
|
||||
for layer in self.convs:
|
||||
x = layer(x)
|
||||
fmap.append(x)
|
||||
|
||||
x = self.conv_post(x)
|
||||
fmap.append(x)
|
||||
|
||||
return fmap
|
||||
|
||||
|
||||
class MSD(nn.Module):
|
||||
def __init__(self, rate: int = 1, sample_rate: int = 44100):
|
||||
super().__init__()
|
||||
self.convs = nn.ModuleList(
|
||||
[
|
||||
WNConv1d(1, 16, 15, 1, padding=7),
|
||||
WNConv1d(16, 64, 41, 4, groups=4, padding=20),
|
||||
WNConv1d(64, 256, 41, 4, groups=16, padding=20),
|
||||
WNConv1d(256, 1024, 41, 4, groups=64, padding=20),
|
||||
WNConv1d(1024, 1024, 41, 4, groups=256, padding=20),
|
||||
WNConv1d(1024, 1024, 5, 1, padding=2),
|
||||
]
|
||||
)
|
||||
self.conv_post = WNConv1d(1024, 1, 3, 1, padding=1, act=False)
|
||||
self.sample_rate = sample_rate
|
||||
self.rate = rate
|
||||
|
||||
def forward(self, x):
|
||||
x = AudioSignal(x, self.sample_rate)
|
||||
x.resample(self.sample_rate // self.rate)
|
||||
x = x.audio_data
|
||||
|
||||
fmap = []
|
||||
|
||||
for l in self.convs:
|
||||
x = l(x)
|
||||
fmap.append(x)
|
||||
x = self.conv_post(x)
|
||||
fmap.append(x)
|
||||
|
||||
return fmap
|
||||
|
||||
|
||||
BANDS = [(0.0, 0.1), (0.1, 0.25), (0.25, 0.5), (0.5, 0.75), (0.75, 1.0)]
|
||||
|
||||
|
||||
class MRD(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
window_length: int,
|
||||
hop_factor: float = 0.25,
|
||||
sample_rate: int = 44100,
|
||||
bands: list = BANDS,
|
||||
):
|
||||
"""Complex multi-band spectrogram discriminator.
|
||||
Parameters
|
||||
----------
|
||||
window_length : int
|
||||
Window length of STFT.
|
||||
hop_factor : float, optional
|
||||
Hop factor of the STFT, defaults to ``0.25 * window_length``.
|
||||
sample_rate : int, optional
|
||||
Sampling rate of audio in Hz, by default 44100
|
||||
bands : list, optional
|
||||
Bands to run discriminator over.
|
||||
"""
|
||||
super().__init__()
|
||||
|
||||
self.window_length = window_length
|
||||
self.hop_factor = hop_factor
|
||||
self.sample_rate = sample_rate
|
||||
self.stft_params = STFTParams(
|
||||
window_length=window_length,
|
||||
hop_length=int(window_length * hop_factor),
|
||||
match_stride=True,
|
||||
)
|
||||
|
||||
n_fft = window_length // 2 + 1
|
||||
bands = [(int(b[0] * n_fft), int(b[1] * n_fft)) for b in bands]
|
||||
self.bands = bands
|
||||
|
||||
ch = 32
|
||||
convs = lambda: nn.ModuleList(
|
||||
[
|
||||
WNConv2d(2, ch, (3, 9), (1, 1), padding=(1, 4)),
|
||||
WNConv2d(ch, ch, (3, 9), (1, 2), padding=(1, 4)),
|
||||
WNConv2d(ch, ch, (3, 9), (1, 2), padding=(1, 4)),
|
||||
WNConv2d(ch, ch, (3, 9), (1, 2), padding=(1, 4)),
|
||||
WNConv2d(ch, ch, (3, 3), (1, 1), padding=(1, 1)),
|
||||
]
|
||||
)
|
||||
self.band_convs = nn.ModuleList([convs() for _ in range(len(self.bands))])
|
||||
self.conv_post = WNConv2d(ch, 1, (3, 3), (1, 1), padding=(1, 1), act=False)
|
||||
|
||||
def spectrogram(self, x):
|
||||
x = AudioSignal(x, self.sample_rate, stft_params=self.stft_params)
|
||||
x = torch.view_as_real(x.stft())
|
||||
x = rearrange(x, "b 1 f t c -> (b 1) c t f")
|
||||
# Split into bands
|
||||
x_bands = [x[..., b[0] : b[1]] for b in self.bands]
|
||||
return x_bands
|
||||
|
||||
def forward(self, x):
|
||||
x_bands = self.spectrogram(x)
|
||||
fmap = []
|
||||
|
||||
x = []
|
||||
for band, stack in zip(x_bands, self.band_convs):
|
||||
for layer in stack:
|
||||
band = layer(band)
|
||||
fmap.append(band)
|
||||
x.append(band)
|
||||
|
||||
x = torch.cat(x, dim=-1)
|
||||
x = self.conv_post(x)
|
||||
fmap.append(x)
|
||||
|
||||
return fmap
|
||||
|
||||
|
||||
class Discriminator(ml.BaseModel):
|
||||
def __init__(
|
||||
self,
|
||||
rates: list = [],
|
||||
periods: list = [2, 3, 5, 7, 11],
|
||||
fft_sizes: list = [2048, 1024, 512],
|
||||
sample_rate: int = 44100,
|
||||
bands: list = BANDS,
|
||||
):
|
||||
"""Discriminator that combines multiple discriminators.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
rates : list, optional
|
||||
sampling rates (in Hz) to run MSD at, by default []
|
||||
If empty, MSD is not used.
|
||||
periods : list, optional
|
||||
periods (of samples) to run MPD at, by default [2, 3, 5, 7, 11]
|
||||
fft_sizes : list, optional
|
||||
Window sizes of the FFT to run MRD at, by default [2048, 1024, 512]
|
||||
sample_rate : int, optional
|
||||
Sampling rate of audio in Hz, by default 44100
|
||||
bands : list, optional
|
||||
Bands to run MRD at, by default `BANDS`
|
||||
"""
|
||||
super().__init__()
|
||||
discs = []
|
||||
discs += [MPD(p) for p in periods]
|
||||
discs += [MSD(r, sample_rate=sample_rate) for r in rates]
|
||||
discs += [MRD(f, sample_rate=sample_rate, bands=bands) for f in fft_sizes]
|
||||
self.discriminators = nn.ModuleList(discs)
|
||||
|
||||
def preprocess(self, y):
|
||||
# Remove DC offset
|
||||
y = y - y.mean(dim=-1, keepdims=True)
|
||||
# Peak normalize the volume of input audio
|
||||
y = 0.8 * y / (y.abs().max(dim=-1, keepdim=True)[0] + 1e-9)
|
||||
return y
|
||||
|
||||
def forward(self, x):
|
||||
x = self.preprocess(x)
|
||||
fmaps = [d(x) for d in self.discriminators]
|
||||
return fmaps
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
disc = Discriminator()
|
||||
x = torch.zeros(1, 1, 44100)
|
||||
results = disc(x)
|
||||
for i, result in enumerate(results):
|
||||
print(f"disc{i}")
|
||||
for i, r in enumerate(result):
|
||||
print(r.shape, r.mean(), r.min(), r.max())
|
||||
print()
|
||||
@@ -0,0 +1,3 @@
|
||||
from . import layers
|
||||
from . import loss
|
||||
from . import quantize
|
||||
+34
@@ -0,0 +1,34 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
from torch.nn.utils import weight_norm
|
||||
|
||||
|
||||
def WNConv1d(*args, **kwargs):
|
||||
return weight_norm(nn.Conv1d(*args, **kwargs))
|
||||
|
||||
|
||||
def WNConvTranspose1d(*args, **kwargs):
|
||||
return weight_norm(nn.ConvTranspose1d(*args, **kwargs))
|
||||
|
||||
|
||||
# Scripting this brings model speed up 1.4x
|
||||
@torch.jit.script
|
||||
def snake(x, alpha):
|
||||
shape = x.shape
|
||||
x = x.reshape(shape[0], shape[1], -1)
|
||||
x = x + (alpha + 1e-9).reciprocal() * torch.sin(alpha * x).pow(2)
|
||||
x = x.reshape(shape)
|
||||
return x
|
||||
|
||||
|
||||
class Snake1d(nn.Module):
|
||||
def __init__(self, channels):
|
||||
super().__init__()
|
||||
self.alpha = nn.Parameter(torch.ones(1, channels, 1))
|
||||
|
||||
def forward(self, x):
|
||||
return snake(x, self.alpha)
|
||||
|
||||
+368
@@ -0,0 +1,368 @@
|
||||
import typing
|
||||
from typing import List
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from audiotools import AudioSignal
|
||||
from audiotools import STFTParams
|
||||
from torch import nn
|
||||
|
||||
|
||||
class L1Loss(nn.L1Loss):
|
||||
"""L1 Loss between AudioSignals. Defaults
|
||||
to comparing ``audio_data``, but any
|
||||
attribute of an AudioSignal can be used.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
attribute : str, optional
|
||||
Attribute of signal to compare, defaults to ``audio_data``.
|
||||
weight : float, optional
|
||||
Weight of this loss, defaults to 1.0.
|
||||
|
||||
Implementation copied from: https://github.com/descriptinc/lyrebird-audiotools/blob/961786aa1a9d628cca0c0486e5885a457fe70c1a/audiotools/metrics/distance.py
|
||||
"""
|
||||
|
||||
def __init__(self, attribute: str = "audio_data", weight: float = 1.0, **kwargs):
|
||||
self.attribute = attribute
|
||||
self.weight = weight
|
||||
super().__init__(**kwargs)
|
||||
|
||||
def forward(self, x: AudioSignal, y: AudioSignal):
|
||||
"""
|
||||
Parameters
|
||||
----------
|
||||
x : AudioSignal
|
||||
Estimate AudioSignal
|
||||
y : AudioSignal
|
||||
Reference AudioSignal
|
||||
|
||||
Returns
|
||||
-------
|
||||
torch.Tensor
|
||||
L1 loss between AudioSignal attributes.
|
||||
"""
|
||||
if isinstance(x, AudioSignal):
|
||||
x = getattr(x, self.attribute)
|
||||
y = getattr(y, self.attribute)
|
||||
return super().forward(x, y)
|
||||
|
||||
|
||||
class SISDRLoss(nn.Module):
|
||||
"""
|
||||
Computes the Scale-Invariant Source-to-Distortion Ratio between a batch
|
||||
of estimated and reference audio signals or aligned features.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
scaling : int, optional
|
||||
Whether to use scale-invariant (True) or
|
||||
signal-to-noise ratio (False), by default True
|
||||
reduction : str, optional
|
||||
How to reduce across the batch (either 'mean',
|
||||
'sum', or none).], by default ' mean'
|
||||
zero_mean : int, optional
|
||||
Zero mean the references and estimates before
|
||||
computing the loss, by default True
|
||||
clip_min : int, optional
|
||||
The minimum possible loss value. Helps network
|
||||
to not focus on making already good examples better, by default None
|
||||
weight : float, optional
|
||||
Weight of this loss, defaults to 1.0.
|
||||
|
||||
Implementation copied from: https://github.com/descriptinc/lyrebird-audiotools/blob/961786aa1a9d628cca0c0486e5885a457fe70c1a/audiotools/metrics/distance.py
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
scaling: int = True,
|
||||
reduction: str = "mean",
|
||||
zero_mean: int = True,
|
||||
clip_min: int = None,
|
||||
weight: float = 1.0,
|
||||
):
|
||||
self.scaling = scaling
|
||||
self.reduction = reduction
|
||||
self.zero_mean = zero_mean
|
||||
self.clip_min = clip_min
|
||||
self.weight = weight
|
||||
super().__init__()
|
||||
|
||||
def forward(self, x: AudioSignal, y: AudioSignal):
|
||||
eps = 1e-8
|
||||
# nb, nc, nt
|
||||
if isinstance(x, AudioSignal):
|
||||
references = x.audio_data
|
||||
estimates = y.audio_data
|
||||
else:
|
||||
references = x
|
||||
estimates = y
|
||||
|
||||
nb = references.shape[0]
|
||||
references = references.reshape(nb, 1, -1).permute(0, 2, 1)
|
||||
estimates = estimates.reshape(nb, 1, -1).permute(0, 2, 1)
|
||||
|
||||
# samples now on axis 1
|
||||
if self.zero_mean:
|
||||
mean_reference = references.mean(dim=1, keepdim=True)
|
||||
mean_estimate = estimates.mean(dim=1, keepdim=True)
|
||||
else:
|
||||
mean_reference = 0
|
||||
mean_estimate = 0
|
||||
|
||||
_references = references - mean_reference
|
||||
_estimates = estimates - mean_estimate
|
||||
|
||||
references_projection = (_references**2).sum(dim=-2) + eps
|
||||
references_on_estimates = (_estimates * _references).sum(dim=-2) + eps
|
||||
|
||||
scale = (
|
||||
(references_on_estimates / references_projection).unsqueeze(1)
|
||||
if self.scaling
|
||||
else 1
|
||||
)
|
||||
|
||||
e_true = scale * _references
|
||||
e_res = _estimates - e_true
|
||||
|
||||
signal = (e_true**2).sum(dim=1)
|
||||
noise = (e_res**2).sum(dim=1)
|
||||
sdr = -10 * torch.log10(signal / noise + eps)
|
||||
|
||||
if self.clip_min is not None:
|
||||
sdr = torch.clamp(sdr, min=self.clip_min)
|
||||
|
||||
if self.reduction == "mean":
|
||||
sdr = sdr.mean()
|
||||
elif self.reduction == "sum":
|
||||
sdr = sdr.sum()
|
||||
return sdr
|
||||
|
||||
|
||||
class MultiScaleSTFTLoss(nn.Module):
|
||||
"""Computes the multi-scale STFT loss from [1].
|
||||
|
||||
Parameters
|
||||
----------
|
||||
window_lengths : List[int], optional
|
||||
Length of each window of each STFT, by default [2048, 512]
|
||||
loss_fn : typing.Callable, optional
|
||||
How to compare each loss, by default nn.L1Loss()
|
||||
clamp_eps : float, optional
|
||||
Clamp on the log magnitude, below, by default 1e-5
|
||||
mag_weight : float, optional
|
||||
Weight of raw magnitude portion of loss, by default 1.0
|
||||
log_weight : float, optional
|
||||
Weight of log magnitude portion of loss, by default 1.0
|
||||
pow : float, optional
|
||||
Power to raise magnitude to before taking log, by default 2.0
|
||||
weight : float, optional
|
||||
Weight of this loss, by default 1.0
|
||||
match_stride : bool, optional
|
||||
Whether to match the stride of convolutional layers, by default False
|
||||
|
||||
References
|
||||
----------
|
||||
|
||||
1. Engel, Jesse, Chenjie Gu, and Adam Roberts.
|
||||
"DDSP: Differentiable Digital Signal Processing."
|
||||
International Conference on Learning Representations. 2019.
|
||||
|
||||
Implementation copied from: https://github.com/descriptinc/lyrebird-audiotools/blob/961786aa1a9d628cca0c0486e5885a457fe70c1a/audiotools/metrics/spectral.py
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
window_lengths: List[int] = [2048, 512],
|
||||
loss_fn: typing.Callable = nn.L1Loss(),
|
||||
clamp_eps: float = 1e-5,
|
||||
mag_weight: float = 1.0,
|
||||
log_weight: float = 1.0,
|
||||
pow: float = 2.0,
|
||||
weight: float = 1.0,
|
||||
match_stride: bool = False,
|
||||
window_type: str = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.stft_params = [
|
||||
STFTParams(
|
||||
window_length=w,
|
||||
hop_length=w // 4,
|
||||
match_stride=match_stride,
|
||||
window_type=window_type,
|
||||
)
|
||||
for w in window_lengths
|
||||
]
|
||||
self.loss_fn = loss_fn
|
||||
self.log_weight = log_weight
|
||||
self.mag_weight = mag_weight
|
||||
self.clamp_eps = clamp_eps
|
||||
self.weight = weight
|
||||
self.pow = pow
|
||||
|
||||
def forward(self, x: AudioSignal, y: AudioSignal):
|
||||
"""Computes multi-scale STFT between an estimate and a reference
|
||||
signal.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
x : AudioSignal
|
||||
Estimate signal
|
||||
y : AudioSignal
|
||||
Reference signal
|
||||
|
||||
Returns
|
||||
-------
|
||||
torch.Tensor
|
||||
Multi-scale STFT loss.
|
||||
"""
|
||||
loss = 0.0
|
||||
for s in self.stft_params:
|
||||
x.stft(s.window_length, s.hop_length, s.window_type)
|
||||
y.stft(s.window_length, s.hop_length, s.window_type)
|
||||
loss += self.log_weight * self.loss_fn(
|
||||
x.magnitude.clamp(self.clamp_eps).pow(self.pow).log10(),
|
||||
y.magnitude.clamp(self.clamp_eps).pow(self.pow).log10(),
|
||||
)
|
||||
loss += self.mag_weight * self.loss_fn(x.magnitude, y.magnitude)
|
||||
return loss
|
||||
|
||||
|
||||
class MelSpectrogramLoss(nn.Module):
|
||||
"""Compute distance between mel spectrograms. Can be used
|
||||
in a multi-scale way.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
n_mels : List[int]
|
||||
Number of mels per STFT, by default [150, 80],
|
||||
window_lengths : List[int], optional
|
||||
Length of each window of each STFT, by default [2048, 512]
|
||||
loss_fn : typing.Callable, optional
|
||||
How to compare each loss, by default nn.L1Loss()
|
||||
clamp_eps : float, optional
|
||||
Clamp on the log magnitude, below, by default 1e-5
|
||||
mag_weight : float, optional
|
||||
Weight of raw magnitude portion of loss, by default 1.0
|
||||
log_weight : float, optional
|
||||
Weight of log magnitude portion of loss, by default 1.0
|
||||
pow : float, optional
|
||||
Power to raise magnitude to before taking log, by default 2.0
|
||||
weight : float, optional
|
||||
Weight of this loss, by default 1.0
|
||||
match_stride : bool, optional
|
||||
Whether to match the stride of convolutional layers, by default False
|
||||
|
||||
Implementation copied from: https://github.com/descriptinc/lyrebird-audiotools/blob/961786aa1a9d628cca0c0486e5885a457fe70c1a/audiotools/metrics/spectral.py
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
n_mels: List[int] = [150, 80],
|
||||
window_lengths: List[int] = [2048, 512],
|
||||
loss_fn: typing.Callable = nn.L1Loss(),
|
||||
clamp_eps: float = 1e-5,
|
||||
mag_weight: float = 1.0,
|
||||
log_weight: float = 1.0,
|
||||
pow: float = 2.0,
|
||||
weight: float = 1.0,
|
||||
match_stride: bool = False,
|
||||
mel_fmin: List[float] = [0.0, 0.0],
|
||||
mel_fmax: List[float] = [None, None],
|
||||
window_type: str = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.stft_params = [
|
||||
STFTParams(
|
||||
window_length=w,
|
||||
hop_length=w // 4,
|
||||
match_stride=match_stride,
|
||||
window_type=window_type,
|
||||
)
|
||||
for w in window_lengths
|
||||
]
|
||||
self.n_mels = n_mels
|
||||
self.loss_fn = loss_fn
|
||||
self.clamp_eps = clamp_eps
|
||||
self.log_weight = log_weight
|
||||
self.mag_weight = mag_weight
|
||||
self.weight = weight
|
||||
self.mel_fmin = mel_fmin
|
||||
self.mel_fmax = mel_fmax
|
||||
self.pow = pow
|
||||
|
||||
def forward(self, x: AudioSignal, y: AudioSignal):
|
||||
"""Computes mel loss between an estimate and a reference
|
||||
signal.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
x : AudioSignal
|
||||
Estimate signal
|
||||
y : AudioSignal
|
||||
Reference signal
|
||||
|
||||
Returns
|
||||
-------
|
||||
torch.Tensor
|
||||
Mel loss.
|
||||
"""
|
||||
loss = 0.0
|
||||
for n_mels, fmin, fmax, s in zip(
|
||||
self.n_mels, self.mel_fmin, self.mel_fmax, self.stft_params
|
||||
):
|
||||
kwargs = {
|
||||
"window_length": s.window_length,
|
||||
"hop_length": s.hop_length,
|
||||
"window_type": s.window_type,
|
||||
}
|
||||
x_mels = x.mel_spectrogram(n_mels, mel_fmin=fmin, mel_fmax=fmax, **kwargs)
|
||||
y_mels = y.mel_spectrogram(n_mels, mel_fmin=fmin, mel_fmax=fmax, **kwargs)
|
||||
|
||||
loss += self.log_weight * self.loss_fn(
|
||||
x_mels.clamp(self.clamp_eps).pow(self.pow).log10(),
|
||||
y_mels.clamp(self.clamp_eps).pow(self.pow).log10(),
|
||||
)
|
||||
loss += self.mag_weight * self.loss_fn(x_mels, y_mels)
|
||||
return loss
|
||||
|
||||
|
||||
class GANLoss(nn.Module):
|
||||
"""
|
||||
Computes a discriminator loss, given a discriminator on
|
||||
generated waveforms/spectrograms compared to ground truth
|
||||
waveforms/spectrograms. Computes the loss for both the
|
||||
discriminator and the generator in separate functions.
|
||||
"""
|
||||
|
||||
def __init__(self, discriminator):
|
||||
super().__init__()
|
||||
self.discriminator = discriminator
|
||||
|
||||
def forward(self, fake, real):
|
||||
d_fake = self.discriminator(fake.audio_data)
|
||||
d_real = self.discriminator(real.audio_data)
|
||||
return d_fake, d_real
|
||||
|
||||
def discriminator_loss(self, fake, real):
|
||||
d_fake, d_real = self.forward(fake.clone().detach(), real)
|
||||
|
||||
loss_d = 0
|
||||
for x_fake, x_real in zip(d_fake, d_real):
|
||||
loss_d += torch.mean(x_fake[-1] ** 2)
|
||||
loss_d += torch.mean((1 - x_real[-1]) ** 2)
|
||||
return loss_d
|
||||
|
||||
def generator_loss(self, fake, real):
|
||||
d_fake, d_real = self.forward(fake, real)
|
||||
|
||||
loss_g = 0
|
||||
for x_fake in d_fake:
|
||||
loss_g += torch.mean((1 - x_fake[-1]) ** 2)
|
||||
|
||||
loss_feature = 0
|
||||
|
||||
for i in range(len(d_fake)):
|
||||
for j in range(len(d_fake[i]) - 1):
|
||||
loss_feature += F.l1_loss(d_fake[i][j], d_real[i][j].detach())
|
||||
return loss_g, loss_feature
|
||||
+262
@@ -0,0 +1,262 @@
|
||||
from typing import Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
from torch.nn.utils import weight_norm
|
||||
|
||||
from .layers import WNConv1d
|
||||
#from .layers import WNConv1d
|
||||
|
||||
class VectorQuantize(nn.Module):
|
||||
"""
|
||||
Implementation of VQ similar to Karpathy's repo:
|
||||
https://github.com/karpathy/deep-vector-quantization
|
||||
Additionally uses following tricks from Improved VQGAN
|
||||
(https://arxiv.org/pdf/2110.04627.pdf):
|
||||
1. Factorized codes: Perform nearest neighbor lookup in low-dimensional space
|
||||
for improved codebook usage
|
||||
2. l2-normalized codes: Converts euclidean distance to cosine similarity which
|
||||
improves training stability
|
||||
"""
|
||||
|
||||
def __init__(self, input_dim: int, codebook_size: int, codebook_dim: int):
|
||||
super().__init__()
|
||||
self.codebook_size = codebook_size
|
||||
self.codebook_dim = codebook_dim
|
||||
|
||||
self.in_proj = WNConv1d(input_dim, codebook_dim, kernel_size=1)
|
||||
self.out_proj = WNConv1d(codebook_dim, input_dim, kernel_size=1)
|
||||
self.codebook = nn.Embedding(codebook_size, codebook_dim)
|
||||
|
||||
def forward(self, z):
|
||||
"""Quantized the input tensor using a fixed codebook and returns
|
||||
the corresponding codebook vectors
|
||||
|
||||
Parameters
|
||||
----------
|
||||
z : Tensor[B x D x T]
|
||||
|
||||
Returns
|
||||
-------
|
||||
Tensor[B x D x T]
|
||||
Quantized continuous representation of input
|
||||
Tensor[1]
|
||||
Commitment loss to train encoder to predict vectors closer to codebook
|
||||
entries
|
||||
Tensor[1]
|
||||
Codebook loss to update the codebook
|
||||
Tensor[B x T]
|
||||
Codebook indices (quantized discrete representation of input)
|
||||
Tensor[B x D x T]
|
||||
Projected latents (continuous representation of input before quantization)
|
||||
"""
|
||||
|
||||
# Factorized codes (ViT-VQGAN) Project input into low-dimensional space
|
||||
z_e = self.in_proj(z) # z_e : (B x D x T)
|
||||
z_q, indices = self.decode_latents(z_e)
|
||||
|
||||
commitment_loss = F.mse_loss(z_e, z_q.detach(), reduction="none").mean([1, 2])
|
||||
codebook_loss = F.mse_loss(z_q, z_e.detach(), reduction="none").mean([1, 2])
|
||||
|
||||
z_q = (
|
||||
z_e + (z_q - z_e).detach()
|
||||
) # noop in forward pass, straight-through gradient estimator in backward pass
|
||||
|
||||
z_q = self.out_proj(z_q)
|
||||
|
||||
return z_q, commitment_loss, codebook_loss, indices, z_e
|
||||
|
||||
def embed_code(self, embed_id):
|
||||
return F.embedding(embed_id, self.codebook.weight)
|
||||
|
||||
def decode_code(self, embed_id):
|
||||
return self.embed_code(embed_id).transpose(1, 2)
|
||||
|
||||
def decode_latents(self, latents):
|
||||
encodings = rearrange(latents, "b d t -> (b t) d")
|
||||
codebook = self.codebook.weight # codebook: (N x D)
|
||||
|
||||
# L2 normalize encodings and codebook (ViT-VQGAN)
|
||||
encodings = F.normalize(encodings)
|
||||
codebook = F.normalize(codebook)
|
||||
|
||||
# Compute euclidean distance with codebook
|
||||
dist = (
|
||||
encodings.pow(2).sum(1, keepdim=True)
|
||||
- 2 * encodings @ codebook.t()
|
||||
+ codebook.pow(2).sum(1, keepdim=True).t()
|
||||
)
|
||||
indices = rearrange((-dist).max(1)[1], "(b t) -> b t", b=latents.size(0))
|
||||
z_q = self.decode_code(indices)
|
||||
return z_q, indices
|
||||
|
||||
|
||||
class ResidualVectorQuantize(nn.Module):
|
||||
"""
|
||||
Introduced in SoundStream: An end2end neural audio codec
|
||||
https://arxiv.org/abs/2107.03312
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
input_dim: int = 512,
|
||||
n_codebooks: int = 9,
|
||||
codebook_size: int = 1024,
|
||||
codebook_dim: Union[int, list] = 8,
|
||||
quantizer_dropout: float = 0.0,
|
||||
):
|
||||
super().__init__()
|
||||
if isinstance(codebook_dim, int):
|
||||
codebook_dim = [codebook_dim for _ in range(n_codebooks)]
|
||||
|
||||
self.n_codebooks = n_codebooks
|
||||
self.codebook_dim = codebook_dim
|
||||
self.codebook_size = codebook_size
|
||||
|
||||
self.quantizers = nn.ModuleList(
|
||||
[
|
||||
VectorQuantize(input_dim, codebook_size, codebook_dim[i])
|
||||
for i in range(n_codebooks)
|
||||
]
|
||||
)
|
||||
self.quantizer_dropout = quantizer_dropout
|
||||
|
||||
def forward(self, z, n_quantizers: int = None):
|
||||
"""Quantized the input tensor using a fixed set of `n` codebooks and returns
|
||||
the corresponding codebook vectors
|
||||
Parameters
|
||||
----------
|
||||
z : Tensor[B x D x T]
|
||||
n_quantizers : int, optional
|
||||
No. of quantizers to use
|
||||
(n_quantizers < self.n_codebooks ex: for quantizer dropout)
|
||||
Note: if `self.quantizer_dropout` is True, this argument is ignored
|
||||
when in training mode, and a random number of quantizers is used.
|
||||
Returns
|
||||
-------
|
||||
dict
|
||||
A dictionary with the following keys:
|
||||
|
||||
"z" : Tensor[B x D x T]
|
||||
Quantized continuous representation of input
|
||||
"codes" : Tensor[B x N x T]
|
||||
Codebook indices for each codebook
|
||||
(quantized discrete representation of input)
|
||||
"latents" : Tensor[B x N*D x T]
|
||||
Projected latents (continuous representation of input before quantization)
|
||||
"vq/commitment_loss" : Tensor[1]
|
||||
Commitment loss to train encoder to predict vectors closer to codebook
|
||||
entries
|
||||
"vq/codebook_loss" : Tensor[1]
|
||||
Codebook loss to update the codebook
|
||||
"""
|
||||
z_q = 0
|
||||
residual = z
|
||||
commitment_loss = 0
|
||||
codebook_loss = 0
|
||||
|
||||
codebook_indices = []
|
||||
latents = []
|
||||
|
||||
if n_quantizers is None:
|
||||
n_quantizers = self.n_codebooks
|
||||
if self.training:
|
||||
n_quantizers = torch.ones((z.shape[0],)) * self.n_codebooks + 1
|
||||
dropout = torch.randint(1, self.n_codebooks + 1, (z.shape[0],))
|
||||
n_dropout = int(z.shape[0] * self.quantizer_dropout)
|
||||
n_quantizers[:n_dropout] = dropout[:n_dropout]
|
||||
n_quantizers = n_quantizers.to(z.device)
|
||||
|
||||
for i, quantizer in enumerate(self.quantizers):
|
||||
if self.training is False and i >= n_quantizers:
|
||||
break
|
||||
|
||||
z_q_i, commitment_loss_i, codebook_loss_i, indices_i, z_e_i = quantizer(
|
||||
residual
|
||||
)
|
||||
|
||||
# Create mask to apply quantizer dropout
|
||||
mask = (
|
||||
torch.full((z.shape[0],), fill_value=i, device=z.device) < n_quantizers
|
||||
)
|
||||
z_q = z_q + z_q_i * mask[:, None, None]
|
||||
residual = residual - z_q_i
|
||||
|
||||
# Sum losses
|
||||
commitment_loss += (commitment_loss_i * mask).mean()
|
||||
codebook_loss += (codebook_loss_i * mask).mean()
|
||||
|
||||
codebook_indices.append(indices_i)
|
||||
latents.append(z_e_i)
|
||||
|
||||
codes = torch.stack(codebook_indices, dim=1)
|
||||
latents = torch.cat(latents, dim=1)
|
||||
|
||||
return z_q, codes, latents, commitment_loss, codebook_loss
|
||||
|
||||
def from_codes(self, codes: torch.Tensor):
|
||||
"""Given the quantized codes, reconstruct the continuous representation
|
||||
Parameters
|
||||
----------
|
||||
codes : Tensor[B x N x T]
|
||||
Quantized discrete representation of input
|
||||
Returns
|
||||
-------
|
||||
Tensor[B x D x T]
|
||||
Quantized continuous representation of input
|
||||
"""
|
||||
z_q = 0.0
|
||||
z_p = []
|
||||
n_codebooks = codes.shape[1]
|
||||
for i in range(n_codebooks):
|
||||
z_p_i = self.quantizers[i].decode_code(codes[:, i, :])
|
||||
z_p.append(z_p_i)
|
||||
|
||||
z_q_i = self.quantizers[i].out_proj(z_p_i)
|
||||
z_q = z_q + z_q_i
|
||||
return z_q, torch.cat(z_p, dim=1), codes
|
||||
|
||||
def from_latents(self, latents: torch.Tensor):
|
||||
"""Given the unquantized latents, reconstruct the
|
||||
continuous representation after quantization.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
latents : Tensor[B x N x T]
|
||||
Continuous representation of input after projection
|
||||
|
||||
Returns
|
||||
-------
|
||||
Tensor[B x D x T]
|
||||
Quantized representation of full-projected space
|
||||
Tensor[B x D x T]
|
||||
Quantized representation of latent space
|
||||
"""
|
||||
z_q = 0
|
||||
z_p = []
|
||||
codes = []
|
||||
dims = np.cumsum([0] + [q.codebook_dim for q in self.quantizers])
|
||||
|
||||
n_codebooks = np.where(dims <= latents.shape[1])[0].max(axis=0, keepdims=True)[
|
||||
0
|
||||
]
|
||||
for i in range(n_codebooks):
|
||||
j, k = dims[i], dims[i + 1]
|
||||
z_p_i, codes_i = self.quantizers[i].decode_latents(latents[:, j:k, :])
|
||||
z_p.append(z_p_i)
|
||||
codes.append(codes_i)
|
||||
|
||||
z_q_i = self.quantizers[i].out_proj(z_p_i)
|
||||
z_q = z_q + z_q_i
|
||||
|
||||
return z_q, torch.cat(z_p, dim=1), torch.stack(codes, dim=1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
rvq = ResidualVectorQuantize(quantizer_dropout=True)
|
||||
x = torch.randn(16, 512, 80)
|
||||
y = rvq(x)
|
||||
print(y["latents"].shape)
|
||||
+127
@@ -0,0 +1,127 @@
|
||||
from pathlib import Path
|
||||
|
||||
import argbind
|
||||
from audiotools import ml
|
||||
|
||||
#from . import dac
|
||||
#import dac
|
||||
|
||||
from ..model import dac
|
||||
|
||||
#DAC = dac.model.DAC
|
||||
DAC = dac.DAC
|
||||
Accelerator = ml.Accelerator
|
||||
|
||||
__MODEL_LATEST_TAGS__ = {
|
||||
("44khz", "8kbps"): "0.0.1",
|
||||
("24khz", "8kbps"): "0.0.4",
|
||||
("16khz", "8kbps"): "0.0.5",
|
||||
("44khz", "16kbps"): "1.0.0",
|
||||
}
|
||||
|
||||
__MODEL_URLS__ = {
|
||||
(
|
||||
"44khz",
|
||||
"0.0.1",
|
||||
"8kbps",
|
||||
): "https://github.com/descriptinc/descript-audio-codec/releases/download/0.0.1/weights.pth",
|
||||
(
|
||||
"24khz",
|
||||
"0.0.4",
|
||||
"8kbps",
|
||||
): "https://github.com/descriptinc/descript-audio-codec/releases/download/0.0.4/weights_24khz.pth",
|
||||
(
|
||||
"16khz",
|
||||
"0.0.5",
|
||||
"8kbps",
|
||||
): "https://github.com/descriptinc/descript-audio-codec/releases/download/0.0.5/weights_16khz.pth",
|
||||
(
|
||||
"44khz",
|
||||
"1.0.0",
|
||||
"16kbps",
|
||||
): "https://github.com/descriptinc/descript-audio-codec/releases/download/1.0.0/weights_44khz_16kbps.pth",
|
||||
}
|
||||
|
||||
|
||||
@argbind.bind(group="download", positional=True, without_prefix=True)
|
||||
def download(
|
||||
model_type: str = "44khz", model_bitrate: str = "8kbps", tag: str = "latest"
|
||||
):
|
||||
"""
|
||||
Function that downloads the weights file from URL if a local cache is not found.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
model_type : str
|
||||
The type of model to download. Must be one of "44khz", "24khz", or "16khz". Defaults to "44khz".
|
||||
model_bitrate: str
|
||||
Bitrate of the model. Must be one of "8kbps", or "16kbps". Defaults to "8kbps".
|
||||
Only 44khz model supports 16kbps.
|
||||
tag : str
|
||||
The tag of the model to download. Defaults to "latest".
|
||||
|
||||
Returns
|
||||
-------
|
||||
Path
|
||||
Directory path required to load model via audiotools.
|
||||
"""
|
||||
model_type = model_type.lower()
|
||||
tag = tag.lower()
|
||||
|
||||
assert model_type in [
|
||||
"44khz",
|
||||
"24khz",
|
||||
"16khz",
|
||||
], "model_type must be one of '44khz', '24khz', or '16khz'"
|
||||
|
||||
assert model_bitrate in [
|
||||
"8kbps",
|
||||
"16kbps",
|
||||
], "model_bitrate must be one of '8kbps', or '16kbps'"
|
||||
|
||||
if tag == "latest":
|
||||
tag = __MODEL_LATEST_TAGS__[(model_type, model_bitrate)]
|
||||
|
||||
download_link = __MODEL_URLS__.get((model_type, tag, model_bitrate), None)
|
||||
|
||||
if download_link is None:
|
||||
raise ValueError(
|
||||
f"Could not find model with tag {tag} and model type {model_type}"
|
||||
)
|
||||
|
||||
local_path = (
|
||||
Path.home()
|
||||
/ ".cache"
|
||||
/ "descript"
|
||||
/ "dac"
|
||||
/ f"weights_{model_type}_{model_bitrate}_{tag}.pth"
|
||||
)
|
||||
if not local_path.exists():
|
||||
local_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Download the model
|
||||
import requests
|
||||
|
||||
response = requests.get(download_link)
|
||||
|
||||
if response.status_code != 200:
|
||||
raise ValueError(
|
||||
f"Could not download model. Received response code {response.status_code}"
|
||||
)
|
||||
local_path.write_bytes(response.content)
|
||||
|
||||
return local_path
|
||||
|
||||
|
||||
def load_model(
|
||||
model_type: str = "44khz",
|
||||
model_bitrate: str = "8kbps",
|
||||
tag: str = "latest",
|
||||
load_path: str = None,
|
||||
):
|
||||
if not load_path:
|
||||
load_path = download(
|
||||
model_type=model_type, model_bitrate=model_bitrate, tag=tag
|
||||
)
|
||||
generator = DAC.load(load_path)
|
||||
return generator
|
||||
+95
@@ -0,0 +1,95 @@
|
||||
import warnings
|
||||
from pathlib import Path
|
||||
|
||||
import argbind
|
||||
import numpy as np
|
||||
import torch
|
||||
from audiotools import AudioSignal
|
||||
from tqdm import tqdm
|
||||
|
||||
from dac import DACFile
|
||||
from dac.utils import load_model
|
||||
|
||||
warnings.filterwarnings("ignore", category=UserWarning)
|
||||
|
||||
|
||||
@argbind.bind(group="decode", positional=True, without_prefix=True)
|
||||
@torch.inference_mode()
|
||||
@torch.no_grad()
|
||||
def decode(
|
||||
input: str,
|
||||
output: str = "",
|
||||
weights_path: str = "",
|
||||
model_tag: str = "latest",
|
||||
model_bitrate: str = "8kbps",
|
||||
device: str = "cuda",
|
||||
model_type: str = "44khz",
|
||||
verbose: bool = False,
|
||||
):
|
||||
"""Decode audio from codes.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
input : str
|
||||
Path to input directory or file
|
||||
output : str, optional
|
||||
Path to output directory, by default "".
|
||||
If `input` is a directory, the directory sub-tree relative to `input` is re-created in `output`.
|
||||
weights_path : str, optional
|
||||
Path to weights file, by default "". If not specified, the weights file will be downloaded from the internet using the
|
||||
model_tag and model_type.
|
||||
model_tag : str, optional
|
||||
Tag of the model to use, by default "latest". Ignored if `weights_path` is specified.
|
||||
model_bitrate: str
|
||||
Bitrate of the model. Must be one of "8kbps", or "16kbps". Defaults to "8kbps".
|
||||
device : str, optional
|
||||
Device to use, by default "cuda". If "cpu", the model will be loaded on the CPU.
|
||||
model_type : str, optional
|
||||
The type of model to use. Must be one of "44khz", "24khz", or "16khz". Defaults to "44khz". Ignored if `weights_path` is specified.
|
||||
"""
|
||||
generator = load_model(
|
||||
model_type=model_type,
|
||||
model_bitrate=model_bitrate,
|
||||
tag=model_tag,
|
||||
load_path=weights_path,
|
||||
)
|
||||
generator.to(device)
|
||||
generator.eval()
|
||||
|
||||
# Find all .dac files in input directory
|
||||
_input = Path(input)
|
||||
input_files = list(_input.glob("**/*.dac"))
|
||||
|
||||
# If input is a .dac file, add it to the list
|
||||
if _input.suffix == ".dac":
|
||||
input_files.append(_input)
|
||||
|
||||
# Create output directory
|
||||
output = Path(output)
|
||||
output.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
for i in tqdm(range(len(input_files)), desc=f"Decoding files"):
|
||||
# Load file
|
||||
artifact = DACFile.load(input_files[i])
|
||||
|
||||
# Reconstruct audio from codes
|
||||
recons = generator.decompress(artifact, verbose=verbose)
|
||||
|
||||
# Compute output path
|
||||
relative_path = input_files[i].relative_to(input)
|
||||
output_dir = output / relative_path.parent
|
||||
if not relative_path.name:
|
||||
output_dir = output
|
||||
relative_path = input_files[i]
|
||||
output_name = relative_path.with_suffix(".wav").name
|
||||
output_path = output_dir / output_name
|
||||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Write to file
|
||||
recons.write(output_path)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
args = argbind.parse_args()
|
||||
with argbind.scope(args):
|
||||
decode()
|
||||
+94
@@ -0,0 +1,94 @@
|
||||
import math
|
||||
import warnings
|
||||
from pathlib import Path
|
||||
|
||||
import argbind
|
||||
import numpy as np
|
||||
import torch
|
||||
from audiotools import AudioSignal
|
||||
from audiotools.core import util
|
||||
from tqdm import tqdm
|
||||
|
||||
from dac.utils import load_model
|
||||
|
||||
warnings.filterwarnings("ignore", category=UserWarning)
|
||||
|
||||
|
||||
@argbind.bind(group="encode", positional=True, without_prefix=True)
|
||||
@torch.inference_mode()
|
||||
@torch.no_grad()
|
||||
def encode(
|
||||
input: str,
|
||||
output: str = "",
|
||||
weights_path: str = "",
|
||||
model_tag: str = "latest",
|
||||
model_bitrate: str = "8kbps",
|
||||
n_quantizers: int = None,
|
||||
device: str = "cuda",
|
||||
model_type: str = "44khz",
|
||||
win_duration: float = 5.0,
|
||||
verbose: bool = False,
|
||||
):
|
||||
"""Encode audio files in input path to .dac format.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
input : str
|
||||
Path to input audio file or directory
|
||||
output : str, optional
|
||||
Path to output directory, by default "". If `input` is a directory, the directory sub-tree relative to `input` is re-created in `output`.
|
||||
weights_path : str, optional
|
||||
Path to weights file, by default "". If not specified, the weights file will be downloaded from the internet using the
|
||||
model_tag and model_type.
|
||||
model_tag : str, optional
|
||||
Tag of the model to use, by default "latest". Ignored if `weights_path` is specified.
|
||||
model_bitrate: str
|
||||
Bitrate of the model. Must be one of "8kbps", or "16kbps". Defaults to "8kbps".
|
||||
n_quantizers : int, optional
|
||||
Number of quantizers to use, by default None. If not specified, all the quantizers will be used and the model will compress at maximum bitrate.
|
||||
device : str, optional
|
||||
Device to use, by default "cuda"
|
||||
model_type : str, optional
|
||||
The type of model to use. Must be one of "44khz", "24khz", or "16khz". Defaults to "44khz". Ignored if `weights_path` is specified.
|
||||
"""
|
||||
generator = load_model(
|
||||
model_type=model_type,
|
||||
model_bitrate=model_bitrate,
|
||||
tag=model_tag,
|
||||
load_path=weights_path,
|
||||
)
|
||||
generator.to(device)
|
||||
generator.eval()
|
||||
kwargs = {"n_quantizers": n_quantizers}
|
||||
|
||||
# Find all audio files in input path
|
||||
input = Path(input)
|
||||
audio_files = util.find_audio(input)
|
||||
|
||||
output = Path(output)
|
||||
output.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
for i in tqdm(range(len(audio_files)), desc="Encoding files"):
|
||||
# Load file
|
||||
signal = AudioSignal(audio_files[i])
|
||||
|
||||
# Encode audio to .dac format
|
||||
artifact = generator.compress(signal, win_duration, verbose=verbose, **kwargs)
|
||||
|
||||
# Compute output path
|
||||
relative_path = audio_files[i].relative_to(input)
|
||||
output_dir = output / relative_path.parent
|
||||
if not relative_path.name:
|
||||
output_dir = output
|
||||
relative_path = audio_files[i]
|
||||
output_name = relative_path.with_suffix(".dac").name
|
||||
output_path = output_dir / output_name
|
||||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
artifact.save(output_path)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
args = argbind.parse_args()
|
||||
with argbind.scope(args):
|
||||
encode()
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
+82
-43
@@ -6,14 +6,14 @@ import numpy as np
|
||||
import io
|
||||
import torchaudio
|
||||
from .node_utils import gc_clear
|
||||
from .generate import auto_prompt_type,pre_data,infer_stage2,inference_lowram_final
|
||||
from .generate import auto_prompt_type,infer_stage2,inference_lowram_final,build_model,Separator,song_infer_lowram
|
||||
import time
|
||||
import folder_paths
|
||||
|
||||
MAX_SEED = np.iinfo(np.int32).max
|
||||
current_node_path = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
|
||||
from .SongGeneration.codeclm.models import builders
|
||||
device = torch.device(
|
||||
"cuda:0") if torch.cuda.is_available() else torch.device(
|
||||
"mps") if torch.backends.mps.is_available() else torch.device(
|
||||
@@ -27,7 +27,7 @@ folder_paths.add_model_folder_path("SongGeneration", SongGeneration_Weigths_Path
|
||||
|
||||
|
||||
|
||||
class SongGeneration_Stage1:
|
||||
class SongGeneration_Loader:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@@ -35,24 +35,57 @@ class SongGeneration_Stage1:
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"demucs_pt": (["none"] + [i for i in folder_paths.get_filename_list("SongGeneration") if i.endswith(".pth")],),
|
||||
"auto_prompt_audio_type": (auto_prompt_type,),
|
||||
|
||||
"infer_model": (["none"] +[i for i in folder_paths.get_filename_list("SongGeneration") if i.endswith(".pt") ],),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("SongGeneration_Audiolm","SongGeneration_Cfg")
|
||||
RETURN_NAMES = ("model","cfg")
|
||||
FUNCTION = "main"
|
||||
CATEGORY = "SongGeneration"
|
||||
|
||||
def main(self, infer_model,):
|
||||
infer_model_path=folder_paths.get_full_path("SongGeneration", infer_model) if infer_model != "none" else None
|
||||
assert infer_model_path is not None ,"模型不能为空.need infer model"
|
||||
model,cfg=build_model(os.path.join(SongGeneration_Weigths_Path, "ckpt"),infer_model_path)
|
||||
return (model,cfg)
|
||||
|
||||
|
||||
|
||||
class SongGeneration_Stage1:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"vae": (folder_paths.get_filename_list("vae"),),
|
||||
"seperate_model": (["none"] + [i for i in folder_paths.get_filename_list("SongGeneration") if i.endswith(".safetensors") and not "fix" in i.lower()],),
|
||||
"prompt_pt": (["none"] + [i for i in folder_paths.get_filename_list("SongGeneration") if "prompt" in i.lower()],),
|
||||
"auto_prompt_audio_type": (auto_prompt_type,),
|
||||
"model_1rvq": (["none"] + [i for i in folder_paths.get_filename_list("SongGeneration") if i.endswith(".safetensors")],),
|
||||
"demucs_pt": (["none"] + [i for i in folder_paths.get_filename_list("SongGeneration") if i.endswith(".pth")],),
|
||||
},
|
||||
"optional": {
|
||||
"audio": ("AUDIO",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("SongGeneration_MODEL",)
|
||||
RETURN_NAMES = ("model",)
|
||||
FUNCTION = "loader_main"
|
||||
RETURN_TYPES = ("SongGeneration_Cond",)
|
||||
RETURN_NAMES = ("cond",)
|
||||
FUNCTION = "main"
|
||||
CATEGORY = "SongGeneration"
|
||||
|
||||
def loader_main(self, demucs_pt,auto_prompt_audio_type,**kwargs):
|
||||
def main(self,vae,seperate_model, auto_prompt_audio_type,prompt_pt,model_1rvq,demucs_pt,**kwargs):
|
||||
|
||||
audio=kwargs.get("audio", None)
|
||||
model_sep_path=folder_paths.get_full_path("SongGeneration", seperate_model) if seperate_model != "none" else None
|
||||
vae_model=folder_paths.get_full_path("vae", vae)
|
||||
prompt_pt_path=folder_paths.get_full_path("SongGeneration", prompt_pt) if prompt_pt != "none" else None
|
||||
|
||||
if audio is not None:
|
||||
|
||||
prompt_audio_path = os.path.join(folder_paths.get_input_directory(), f"audio_{time.strftime('%m%d%H%S')}_temp.wav")
|
||||
waveform=audio["waveform"].squeeze(0)
|
||||
buff = io.BytesIO()
|
||||
@@ -60,25 +93,25 @@ class SongGeneration_Stage1:
|
||||
with open(prompt_audio_path, 'wb') as f:
|
||||
f.write(buff.getbuffer())
|
||||
use_descriptions=False #不建议同时提供参考音频和描述文本
|
||||
|
||||
dm_model_path=folder_paths.get_full_path("SongGeneration", demucs_pt) if demucs_pt != "none" else None
|
||||
assert dm_model_path is not None ,"使用参考音频时,需要选择htdemucs模型, if use audio need htdemucs model"
|
||||
separator = Separator(dm_model_path, os.path.join(current_node_path, "SongGeneration/third_party/demucs/ckpt/htdemucs.yaml"))
|
||||
|
||||
model_1rvq_path=folder_paths.get_full_path("SongGeneration", model_1rvq) if model_1rvq != "none" else None
|
||||
assert model_1rvq_path is not None ,"使用参考音频时,需要选择model_模型, if use audio need model_odel"
|
||||
audio_tokenizer = builders.get_audio_tokenizer_model(f"Flow1dVAE1rvq_{model_1rvq_path}",os.path.join(current_node_path, f'SongGeneration/conf/stable_audio_1920_vae.json'),vae_model,'inference')
|
||||
|
||||
seperate_tokenizer = builders.get_audio_tokenizer_model(f"Flow1dVAESeparate_{model_sep_path}",os.path.join(current_node_path, f'SongGeneration/conf/stable_audio_1920_vae.json'),vae_model,'inference')
|
||||
|
||||
else:
|
||||
prompt_audio_path=None
|
||||
use_descriptions=True
|
||||
prompt_audio_path,use_descriptions,audio_tokenizer,separator,seperate_tokenizer=None,True,None,None,None
|
||||
|
||||
|
||||
if demucs_pt == "none":
|
||||
raise ValueError("No demucs_pt selected")
|
||||
|
||||
dm_model_path=folder_paths.get_full_path("SongGeneration", demucs_pt)
|
||||
dm_config_path=os.path.join(current_node_path, "SongGeneration/third_party/demucs/ckpt/htdemucs.yaml")
|
||||
Weigths_Path=os.path.join(SongGeneration_Weigths_Path, "ckpt")
|
||||
|
||||
item,max_duration,cfg=pre_data(Weigths_Path,dm_model_path,dm_config_path,folder_paths.get_output_directory(),prompt_audio_path,auto_prompt_audio_type)
|
||||
|
||||
original_item=song_infer_lowram(seperate_tokenizer,separator,audio_tokenizer,prompt_pt_path, folder_paths.get_output_directory(),prompt_audio_path,auto_prompt_audio_type,)
|
||||
|
||||
gc_clear()
|
||||
|
||||
return ({"item": item, "max_duration": max_duration,"use_descriptions": use_descriptions,"cfg":cfg,"Weigths_Path":Weigths_Path},)
|
||||
|
||||
print("Stage1 is done.")
|
||||
return ({"item": original_item, "use_descriptions": use_descriptions,"model_sep_path":model_sep_path,"vae_model":vae_model},)
|
||||
|
||||
class SongGeneration_Stage2:
|
||||
def __init__(self):
|
||||
@@ -88,7 +121,9 @@ class SongGeneration_Stage2:
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("SongGeneration_MODEL",),
|
||||
"model": ("SongGeneration_Audiolm",),
|
||||
"cfg": ("SongGeneration_Cfg",),
|
||||
"cond": ("SongGeneration_Cond",),
|
||||
"lyric": ("STRING", {"multiline": True, "default": "[intro-short] ;\n [verse]\n 雪花舞动在无尽的天际.情缘如同雪花般轻轻逝去.希望与真挚.永不磨灭.你的忧虑.随风而逝 ;\n [chorus]\n 我怀抱着守护这片梦境.在这世界中寻找爱与虚幻.苦辣酸甜.我们一起品尝.在雪的光芒中.紧紧相拥 ;\n [inst-short] ;\n [verse]\n雪花再次在风中飘扬.情愿如同雪花般消失无踪.希望与真挚.永不消失.在痛苦与喧嚣中.你找到解脱 ;\n [chorus]\n 我环绕着守护这片梦境.在这世界中感受爱与虚假.苦辣酸甜.我们一起分享.在白银的光芒中.我们同在 ;\n [outro-short]"}),
|
||||
"description": ("STRING", {"multiline": False, "default": "female, dark, pop, sad, piano and drums, the bpm is 125"}), #OPTIONAL
|
||||
"cfg_coef": ("FLOAT", {"default": 1.5, "min": 0.1, "max": 3.0, "step": 0.1}),
|
||||
@@ -96,24 +131,22 @@ class SongGeneration_Stage2:
|
||||
"top_k": ("INT", {"default": 50, "min": 1, "max": 100, "step": 1}),
|
||||
"top_p": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
"record_tokens": ("BOOLEAN", {"default": True}),
|
||||
"record_window": ("INT", {"default": 50, "min": 1, "max": 1000, "step": 1}),
|
||||
|
||||
"record_window": ("INT", {"default": 50, "min": 1, "max": 1000, "step": 1}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("SongGeneration_DICT",)
|
||||
RETURN_NAMES = ("model",)
|
||||
FUNCTION = "loader_main"
|
||||
RETURN_TYPES = ("SongGeneration_Cond",)
|
||||
RETURN_NAMES = ("cond",)
|
||||
FUNCTION = "main"
|
||||
CATEGORY = "SongGeneration"
|
||||
|
||||
def loader_main(self, model,lyric,description,cfg_coef,temp,top_k,top_p,record_tokens,record_window):
|
||||
def main(self, model,cfg,cond,lyric,description,cfg_coef,temp,top_k,top_p,record_tokens,record_window):
|
||||
|
||||
descriptions=description if model.get("use_descriptions") else None
|
||||
descriptions=description if cond.get("use_descriptions",False) else None
|
||||
|
||||
|
||||
items=infer_stage2(model.get("item"),model.get("cfg"),model.get("Weigths_Path"),model.get("max_duration"),lyric,descriptions,cfg_coef, temp,top_k,top_p,record_tokens ,record_window )
|
||||
items=infer_stage2(cond.get("item"),model,cfg.max_dur,lyric,descriptions,cfg_coef, temp,top_k,top_p,record_tokens ,record_window )
|
||||
gc_clear()
|
||||
return ({"items":items,"cfg":model.get("cfg"),"max_duration":model.get("max_duration"),},)
|
||||
return ({"items":items,"cfg":cfg,"model_sep_path":cond["model_sep_path"],"vae_model":cond["vae_model"]},)
|
||||
|
||||
|
||||
|
||||
@@ -125,9 +158,9 @@ class SongGeneration_Sampler:
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("SongGeneration_DICT",),
|
||||
"cond": ("SongGeneration_Cond",),
|
||||
"gen_type": (["mixed","bgm","vocal",],),
|
||||
"save_separate": ("BOOLEAN", {"default": True}),
|
||||
"save_separate": ("BOOLEAN", {"default": False}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -136,24 +169,30 @@ class SongGeneration_Sampler:
|
||||
FUNCTION = "sampler_main"
|
||||
CATEGORY = "SongGeneration"
|
||||
|
||||
def sampler_main(self, model,gen_type,save_separate):
|
||||
cfg=model.get("cfg")
|
||||
def sampler_main(self,cond,gen_type,save_separate):
|
||||
cfg=cond.get("cfg")
|
||||
cfg.gen_type=gen_type
|
||||
print("start inference final")
|
||||
audio=inference_lowram_final(cfg,model.get("max_duration"),model.get("items"),folder_paths.get_output_directory(),save_separate)
|
||||
|
||||
model_sep_path=cond["model_sep_path"]
|
||||
vae_model=cond["vae_model"]
|
||||
print("start inference final,loading model")
|
||||
seperate_tokenizer = builders.get_audio_tokenizer_model(f"Flow1dVAESeparate_{model_sep_path}",os.path.join(current_node_path, f'SongGeneration/conf/stable_audio_1920_vae.json'),vae_model,'inference')
|
||||
seperate_tokenizer = seperate_tokenizer.eval().cuda()
|
||||
audio=inference_lowram_final(cfg,seperate_tokenizer,cfg.max_dur,cond.get("items"),folder_paths.get_output_directory(),save_separate)
|
||||
del seperate_tokenizer
|
||||
gc_clear()
|
||||
return (audio,)
|
||||
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"SongGeneration_Loader":SongGeneration_Loader,
|
||||
"SongGeneration_Stage1": SongGeneration_Stage1,
|
||||
"SongGeneration_Stage2": SongGeneration_Stage2,
|
||||
"SongGeneration_Sampler": SongGeneration_Sampler,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"SongGeneration_Loader": "SongGeneration_Loader",
|
||||
"SongGeneration_Stage1": "SongGeneration_Stage1",
|
||||
"SongGeneration_Stage2": "SongGeneration_Stage2",
|
||||
"SongGeneration_Sampler": "SongGeneration_Sampler",
|
||||
|
||||
@@ -0,0 +1,598 @@
|
||||
{
|
||||
"id": "f7a5247b-ed1b-466d-91df-de458a19772c",
|
||||
"revision": 0,
|
||||
"last_node_id": 23,
|
||||
"last_link_id": 27,
|
||||
"nodes": [
|
||||
{
|
||||
"id": 9,
|
||||
"type": "SongGeneration_Stage2",
|
||||
"pos": [
|
||||
1413.7088623046875,
|
||||
273.0107727050781
|
||||
],
|
||||
"size": [
|
||||
441.0806579589844,
|
||||
600.54931640625
|
||||
],
|
||||
"flags": {},
|
||||
"order": 5,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "model",
|
||||
"type": "SongGeneration_Audiolm",
|
||||
"link": 11
|
||||
},
|
||||
{
|
||||
"name": "cfg",
|
||||
"type": "SongGeneration_Cfg",
|
||||
"link": 12
|
||||
},
|
||||
{
|
||||
"name": "cond",
|
||||
"type": "SongGeneration_Cond",
|
||||
"link": 18
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "cond",
|
||||
"type": "SongGeneration_Cond",
|
||||
"links": [
|
||||
19
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "SongGeneration_Stage2"
|
||||
},
|
||||
"widgets_values": [
|
||||
"[intro-medium]\n\n[verse]\n夜晚的街灯闪烁\n我漫步在熟悉的角落\n回忆像潮水般涌来\n你的笑容如此清晰\n在心头无法抹去\n那些曾经的甜蜜\n如今只剩我独自回忆\n\n[chorus]\n回忆的温度还在\n你却已不在\n我的心被爱填满\n却又被思念刺痛\n音乐的节奏奏响\n我的心却在流浪\n没有你的日子\n我该如何继续向前\n\n[inst-medium]\n\n[verse]\n手机屏幕亮起\n是你发来的消息\n简单的几个字\n却让我泪流满面\n曾经的拥抱温暖\n如今却变得遥远\n我多想回到从前\n重新拥有你的陪伴\n\n[chorus]\n回忆的温度还在\n你却已不在\n我的心被爱填满\n却又被思念刺痛\n音乐的节奏奏响\n我的心却在流浪\n没有你的日子\n我该如何继续向前\n\n[outro-medium]",
|
||||
"female, dark, pop, sad, piano and drums, the bpm is 125",
|
||||
1.5,
|
||||
0.9,
|
||||
50,
|
||||
0,
|
||||
true,
|
||||
50
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 14,
|
||||
"type": "SongGeneration_Sampler",
|
||||
"pos": [
|
||||
1981.21533203125,
|
||||
320.8858337402344
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
82
|
||||
],
|
||||
"flags": {},
|
||||
"order": 7,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "cond",
|
||||
"type": "SongGeneration_Cond",
|
||||
"link": 19
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "audio",
|
||||
"type": "AUDIO",
|
||||
"links": [
|
||||
21
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "SongGeneration_Sampler"
|
||||
},
|
||||
"widgets_values": [
|
||||
"mixed",
|
||||
false
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 16,
|
||||
"type": "PreviewAudio",
|
||||
"pos": [
|
||||
2089.89208984375,
|
||||
484.1910095214844
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
88
|
||||
],
|
||||
"flags": {},
|
||||
"order": 9,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "audio",
|
||||
"type": "AUDIO",
|
||||
"link": 21
|
||||
}
|
||||
],
|
||||
"outputs": [],
|
||||
"properties": {
|
||||
"Node name for S&R": "PreviewAudio"
|
||||
},
|
||||
"widgets_values": []
|
||||
},
|
||||
{
|
||||
"id": 13,
|
||||
"type": "SongGeneration_Stage1",
|
||||
"pos": [
|
||||
1035.6988525390625,
|
||||
456.20953369140625
|
||||
],
|
||||
"size": [
|
||||
296.767578125,
|
||||
178
|
||||
],
|
||||
"flags": {},
|
||||
"order": 0,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "audio",
|
||||
"shape": 7,
|
||||
"type": "AUDIO",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "cond",
|
||||
"type": "SongGeneration_Cond",
|
||||
"links": [
|
||||
18
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "SongGeneration_Stage1"
|
||||
},
|
||||
"widgets_values": [
|
||||
"autoencoder_music_1320k.ckpt",
|
||||
"model_2.safetensors",
|
||||
"new_prompt.pt",
|
||||
"Pop",
|
||||
"none",
|
||||
"none"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 19,
|
||||
"type": "SongGeneration_Sampler",
|
||||
"pos": [
|
||||
3454.367919921875,
|
||||
276.9729309082031
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
82
|
||||
],
|
||||
"flags": {},
|
||||
"order": 10,
|
||||
"mode": 2,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "cond",
|
||||
"type": "SongGeneration_Cond",
|
||||
"link": 25
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "audio",
|
||||
"type": "AUDIO",
|
||||
"links": [
|
||||
26
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "SongGeneration_Sampler"
|
||||
},
|
||||
"widgets_values": [
|
||||
"mixed",
|
||||
false
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 20,
|
||||
"type": "PreviewAudio",
|
||||
"pos": [
|
||||
3493.792724609375,
|
||||
490.2040710449219
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
88
|
||||
],
|
||||
"flags": {},
|
||||
"order": 11,
|
||||
"mode": 2,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "audio",
|
||||
"type": "AUDIO",
|
||||
"link": 26
|
||||
}
|
||||
],
|
||||
"outputs": [],
|
||||
"properties": {
|
||||
"Node name for S&R": "PreviewAudio"
|
||||
},
|
||||
"widgets_values": []
|
||||
},
|
||||
{
|
||||
"id": 17,
|
||||
"type": "SongGeneration_Stage2",
|
||||
"pos": [
|
||||
2920.681884765625,
|
||||
274.19219970703125
|
||||
],
|
||||
"size": [
|
||||
441.0806579589844,
|
||||
600.54931640625
|
||||
],
|
||||
"flags": {},
|
||||
"order": 8,
|
||||
"mode": 2,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "model",
|
||||
"type": "SongGeneration_Audiolm",
|
||||
"link": 22
|
||||
},
|
||||
{
|
||||
"name": "cfg",
|
||||
"type": "SongGeneration_Cfg",
|
||||
"link": 23
|
||||
},
|
||||
{
|
||||
"name": "cond",
|
||||
"type": "SongGeneration_Cond",
|
||||
"link": 24
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "cond",
|
||||
"type": "SongGeneration_Cond",
|
||||
"links": [
|
||||
25
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "SongGeneration_Stage2"
|
||||
},
|
||||
"widgets_values": [
|
||||
"[intro-medium]\n\n[verse]\n夜晚的街灯闪烁\n我漫步在熟悉的角落\n回忆像潮水般涌来\n你的笑容如此清晰\n在心头无法抹去\n那些曾经的甜蜜\n如今只剩我独自回忆\n\n[chorus]\n回忆的温度还在\n你却已不在\n我的心被爱填满\n却又被思念刺痛\n音乐的节奏奏响\n我的心却在流浪\n没有你的日子\n我该如何继续向前\n\n[inst-medium]\n\n[verse]\n手机屏幕亮起\n是你发来的消息\n简单的几个字\n却让我泪流满面\n曾经的拥抱温暖\n如今却变得遥远\n我多想回到从前\n重新拥有你的陪伴\n\n[chorus]\n回忆的温度还在\n你却已不在\n我的心被爱填满\n却又被思念刺痛\n音乐的节奏奏响\n我的心却在流浪\n没有你的日子\n我该如何继续向前\n\n[outro-medium]",
|
||||
"female, dark, pop, sad, piano and drums, the bpm is 125",
|
||||
1.5,
|
||||
0.9,
|
||||
50,
|
||||
0,
|
||||
true,
|
||||
50
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 18,
|
||||
"type": "SongGeneration_Loader",
|
||||
"pos": [
|
||||
2569.0380859375,
|
||||
243.3841094970703
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
78
|
||||
],
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 2,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "model",
|
||||
"type": "SongGeneration_Audiolm",
|
||||
"links": [
|
||||
22
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "cfg",
|
||||
"type": "SongGeneration_Cfg",
|
||||
"links": [
|
||||
23
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "SongGeneration_Loader"
|
||||
},
|
||||
"widgets_values": [
|
||||
"new_model.pt"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 1,
|
||||
"type": "SongGeneration_Loader",
|
||||
"pos": [
|
||||
1037.9068603515625,
|
||||
305.0125427246094
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
78
|
||||
],
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "model",
|
||||
"type": "SongGeneration_Audiolm",
|
||||
"links": [
|
||||
11
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "cfg",
|
||||
"type": "SongGeneration_Cfg",
|
||||
"links": [
|
||||
12
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "SongGeneration_Loader"
|
||||
},
|
||||
"widgets_values": [
|
||||
"new_model.pt"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 21,
|
||||
"type": "SongGeneration_Stage1",
|
||||
"pos": [
|
||||
2571.6611328125,
|
||||
412.2967834472656
|
||||
],
|
||||
"size": [
|
||||
296.767578125,
|
||||
178
|
||||
],
|
||||
"flags": {},
|
||||
"order": 6,
|
||||
"mode": 2,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "audio",
|
||||
"shape": 7,
|
||||
"type": "AUDIO",
|
||||
"link": 27
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "cond",
|
||||
"type": "SongGeneration_Cond",
|
||||
"links": [
|
||||
24
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "SongGeneration_Stage1"
|
||||
},
|
||||
"widgets_values": [
|
||||
"autoencoder_music_1320k.ckpt",
|
||||
"model_2.safetensors",
|
||||
"new_prompt.pt",
|
||||
"Pop",
|
||||
"model_2_fixed.safetensors",
|
||||
"htdemucs.pth"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 22,
|
||||
"type": "LoadAudio",
|
||||
"pos": [
|
||||
2541.430419921875,
|
||||
673.8790893554688
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
136
|
||||
],
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 2,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "AUDIO",
|
||||
"type": "AUDIO",
|
||||
"links": [
|
||||
27
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadAudio"
|
||||
},
|
||||
"widgets_values": [
|
||||
"AnimateDiff_00003.mp4",
|
||||
null,
|
||||
""
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 23,
|
||||
"type": "Note",
|
||||
"pos": [
|
||||
2131.3876953125,
|
||||
-117.82566833496094
|
||||
],
|
||||
"size": [
|
||||
697.4249877929688,
|
||||
185.1282958984375
|
||||
],
|
||||
"flags": {},
|
||||
"order": 4,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [],
|
||||
"properties": {},
|
||||
"widgets_values": [
|
||||
"rename model.pt if use large model.pt to large_model.pt 如果使用large模型,修改模型名字加入large,不然识别不到\nrename model.pt if use new base model.pt to new_model.pt 如果使用new模型,修改模型名字加入new,不然识别不到\nrename model.pt if use full model.pt to large_model.pt 如果使用full模型,修改模型名字加入full,不然识别不到\nnot rename model.pt if use origin base model.pt \n\n"
|
||||
],
|
||||
"color": "#432",
|
||||
"bgcolor": "#653"
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
11,
|
||||
1,
|
||||
0,
|
||||
9,
|
||||
0,
|
||||
"SongGeneration_Audiolm"
|
||||
],
|
||||
[
|
||||
12,
|
||||
1,
|
||||
1,
|
||||
9,
|
||||
1,
|
||||
"SongGeneration_Cfg"
|
||||
],
|
||||
[
|
||||
18,
|
||||
13,
|
||||
0,
|
||||
9,
|
||||
2,
|
||||
"SongGeneration_Cond"
|
||||
],
|
||||
[
|
||||
19,
|
||||
9,
|
||||
0,
|
||||
14,
|
||||
0,
|
||||
"SongGeneration_Cond"
|
||||
],
|
||||
[
|
||||
21,
|
||||
14,
|
||||
0,
|
||||
16,
|
||||
0,
|
||||
"AUDIO"
|
||||
],
|
||||
[
|
||||
22,
|
||||
18,
|
||||
0,
|
||||
17,
|
||||
0,
|
||||
"SongGeneration_Audiolm"
|
||||
],
|
||||
[
|
||||
23,
|
||||
18,
|
||||
1,
|
||||
17,
|
||||
1,
|
||||
"SongGeneration_Cfg"
|
||||
],
|
||||
[
|
||||
24,
|
||||
21,
|
||||
0,
|
||||
17,
|
||||
2,
|
||||
"SongGeneration_Cond"
|
||||
],
|
||||
[
|
||||
25,
|
||||
17,
|
||||
0,
|
||||
19,
|
||||
0,
|
||||
"SongGeneration_Cond"
|
||||
],
|
||||
[
|
||||
26,
|
||||
19,
|
||||
0,
|
||||
20,
|
||||
0,
|
||||
"AUDIO"
|
||||
],
|
||||
[
|
||||
27,
|
||||
22,
|
||||
0,
|
||||
21,
|
||||
0,
|
||||
"AUDIO"
|
||||
]
|
||||
],
|
||||
"groups": [
|
||||
{
|
||||
"id": 1,
|
||||
"title": "infer prompt pt",
|
||||
"bounding": [
|
||||
1010.257080078125,
|
||||
135.5015411376953,
|
||||
1365.45166015625,
|
||||
826.6909790039062
|
||||
],
|
||||
"color": "#3f789e",
|
||||
"font_size": 24,
|
||||
"flags": {}
|
||||
},
|
||||
{
|
||||
"id": 2,
|
||||
"title": "infer audio",
|
||||
"bounding": [
|
||||
2410.936767578125,
|
||||
139.90402221679688,
|
||||
1365.45166015625,
|
||||
826.6909790039062
|
||||
],
|
||||
"color": "#3f789e",
|
||||
"font_size": 24,
|
||||
"flags": {}
|
||||
}
|
||||
],
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 0.6209213230591556,
|
||||
"offset": [
|
||||
-904.784220650397,
|
||||
207.29318756335755
|
||||
]
|
||||
},
|
||||
"frontendVersion": "1.26.13",
|
||||
"VHS_latentpreview": false,
|
||||
"VHS_latentpreviewrate": 0,
|
||||
"VHS_MetadataImage": true,
|
||||
"VHS_KeepIntermediate": true
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 174 KiB |
@@ -1,195 +0,0 @@
|
||||
{
|
||||
"id": "30851f96-07c3-4422-bf5a-9cac438f486d",
|
||||
"revision": 0,
|
||||
"last_node_id": 5,
|
||||
"last_link_id": 4,
|
||||
"nodes": [
|
||||
{
|
||||
"id": 2,
|
||||
"type": "SongGeneration_Stage2",
|
||||
"pos": [
|
||||
21059.890625,
|
||||
-1126.8658447265625
|
||||
],
|
||||
"size": [
|
||||
377,
|
||||
411
|
||||
],
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "model",
|
||||
"type": "SongGeneration_MODEL",
|
||||
"link": 1
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "model",
|
||||
"type": "SongGeneration_DICT",
|
||||
"links": [
|
||||
2
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "SongGeneration_Stage2"
|
||||
},
|
||||
"widgets_values": [
|
||||
"[intro-short] ;\n [verse]\n 雪花舞动在无尽的天际.情缘如同雪花般轻轻逝去.希望与真挚.永不磨灭.你的忧虑.随风而逝 ;\n [chorus]\n 我怀抱着守护这片梦境.在这世界中寻找爱与虚幻.苦辣酸甜.我们一起品尝.在雪的光芒中.紧紧相拥 ;\n [inst-short] ;\n [verse]\n雪花再次在风中飘扬.情愿如同雪花般消失无踪.希望与真挚.永不消失.在痛苦与喧嚣中.你找到解脱 ;\n [chorus]\n 我环绕着守护这片梦境.在这世界中感受爱与虚假.苦辣酸甜.我们一起分享.在白银的光芒中.我们同在 ;\n [outro-short]",
|
||||
"female, dark, pop, sad, piano and drums, the bpm is 125",
|
||||
1.5,
|
||||
0.9,
|
||||
50,
|
||||
0,
|
||||
true,
|
||||
50
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 1,
|
||||
"type": "SongGeneration_Stage1",
|
||||
"pos": [
|
||||
20679.369140625,
|
||||
-1114.414306640625
|
||||
],
|
||||
"size": [
|
||||
296.767578125,
|
||||
82
|
||||
],
|
||||
"flags": {},
|
||||
"order": 0,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "audio",
|
||||
"shape": 7,
|
||||
"type": "AUDIO",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "model",
|
||||
"type": "SongGeneration_MODEL",
|
||||
"links": [
|
||||
1
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "SongGeneration_Stage1"
|
||||
},
|
||||
"widgets_values": [
|
||||
"htdemucs.pth",
|
||||
"Pop"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 3,
|
||||
"type": "SongGeneration_Sampler",
|
||||
"pos": [
|
||||
21540.087890625,
|
||||
-1061.4998779296875
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
82
|
||||
],
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "model",
|
||||
"type": "SongGeneration_DICT",
|
||||
"link": 2
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "audio",
|
||||
"type": "AUDIO",
|
||||
"links": [
|
||||
4
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "SongGeneration_Sampler"
|
||||
},
|
||||
"widgets_values": [
|
||||
"vocal",
|
||||
true
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 5,
|
||||
"type": "SaveAudio",
|
||||
"pos": [
|
||||
21841.2890625,
|
||||
-964.7002563476562
|
||||
],
|
||||
"size": [
|
||||
386.6000061035156,
|
||||
140.60000610351562
|
||||
],
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "audio",
|
||||
"type": "AUDIO",
|
||||
"link": 4
|
||||
}
|
||||
],
|
||||
"outputs": [],
|
||||
"properties": {},
|
||||
"widgets_values": [
|
||||
"audio/ComfyUI"
|
||||
]
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
1,
|
||||
1,
|
||||
0,
|
||||
2,
|
||||
0,
|
||||
"SongGeneration_MODEL"
|
||||
],
|
||||
[
|
||||
2,
|
||||
2,
|
||||
0,
|
||||
3,
|
||||
0,
|
||||
"SongGeneration_DICT"
|
||||
],
|
||||
[
|
||||
4,
|
||||
3,
|
||||
0,
|
||||
5,
|
||||
0,
|
||||
"AUDIO"
|
||||
]
|
||||
],
|
||||
"groups": [],
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 1,
|
||||
"offset": [
|
||||
-20532.38997582383,
|
||||
1343.5001986783957
|
||||
]
|
||||
},
|
||||
"frontendVersion": "1.23.4"
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
Binary file not shown.
|
Before Width: | Height: | Size: 96 KiB |
+96
-61
@@ -12,6 +12,7 @@ from .SongGeneration.codeclm.models import builders
|
||||
from .SongGeneration.codeclm.trainer.codec_song_pl import CodecLM_PL
|
||||
from .SongGeneration.codeclm.models import CodecLM
|
||||
from .SongGeneration.third_party.demucs.models.pretrained import get_model_from_yaml
|
||||
current_node_path = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
auto_prompt_type = ['Pop', 'R&B', 'Dance', 'Jazz', 'Folk', 'Rock', 'Chinese Style', 'Chinese Tradition', 'Metal', 'Reggae', 'Chinese Opera', 'Auto']
|
||||
|
||||
@@ -60,9 +61,10 @@ class Separator():
|
||||
return full_audio, vocal_audio, bgm_audio
|
||||
|
||||
|
||||
def pre_data(Weigths_Path,dm_model_path,dm_config_path,save_dir,prompt_audio_path,auto_prompt_audio_type):
|
||||
|
||||
def build_model(Weigths_Path,infer_model_path):
|
||||
torch.backends.cudnn.enabled = False
|
||||
curent_dir = os.path.join(folder_paths.base_path,"custom_nodes/ComfyUI_SongGeneration/SongGeneration")
|
||||
curent_dir = os.path.join(current_node_path,"SongGeneration")
|
||||
RESOLVERS = {
|
||||
"eval": lambda x: eval(x),
|
||||
"concat": lambda *x: [xxx for xx in x for xxx in xx],
|
||||
@@ -73,42 +75,78 @@ def pre_data(Weigths_Path,dm_model_path,dm_config_path,save_dir,prompt_audio_pat
|
||||
for name, func in RESOLVERS.items():
|
||||
if not OmegaConf.has_resolver(name):
|
||||
OmegaConf.register_new_resolver(name, func)
|
||||
np.random.seed(int(time.time()))
|
||||
|
||||
cfg_path = os.path.join(Weigths_Path, 'songgeneration_base/config.yaml')
|
||||
np.random.seed(int(time.time()))
|
||||
infer_model_type="new" if "new" in infer_model_path.lower() else "large" if "large" in infer_model_path.lower() else "full" if "full" in infer_model_path.lower() else "base"
|
||||
|
||||
cfg_path = os.path.join(current_node_path, f'SongGeneration/conf/{infer_model_type}_config.yaml')
|
||||
|
||||
cfg = OmegaConf.load(cfg_path)
|
||||
cfg.mode = 'inference'
|
||||
|
||||
|
||||
|
||||
cfg.vae_config=f"{Weigths_Path}/vae/stable_audio_1920_vae.json"
|
||||
cfg.vae_model=f"{Weigths_Path}/vae/autoencoder_music_1320k.ckpt"
|
||||
|
||||
cfg.audio_tokenizer_checkpoint=f"Flow1dVAE1rvq_{Weigths_Path}/model_1rvq/model_2_fixed.safetensors"
|
||||
cfg.audio_tokenizer_checkpoint_sep=f"Flow1dVAESeparate_{Weigths_Path}/model_septoken/model_2.safetensors"
|
||||
cfg.conditioners.type_info.QwTextTokenizer.token_path=os.path.join(folder_paths.base_path,"custom_nodes/ComfyUI_SongGeneration/SongGeneration/third_party/Qwen2-7B")
|
||||
max_duration = cfg.max_dur
|
||||
cfg.conditioners.type_info.QwTextTokenizer.token_path=os.path.join(current_node_path,"SongGeneration/third_party/Qwen2-7B")
|
||||
|
||||
auto_prompt = torch.load(os.path.join(Weigths_Path,'prompt.pt'),weights_only=False)
|
||||
merge_prompt = [x for sublist in auto_prompt.values() for x in sublist]
|
||||
|
||||
if prompt_audio_path is not None:
|
||||
separator = Separator(dm_model_path, dm_config_path)
|
||||
audio_tokenizer = builders.get_audio_tokenizer_model(cfg.audio_tokenizer_checkpoint, cfg)
|
||||
audio_tokenizer = audio_tokenizer.eval().cuda()
|
||||
|
||||
else:
|
||||
audio_tokenizer = None
|
||||
separator = None
|
||||
audiolm = builders.get_lm_model(cfg)
|
||||
checkpoint = torch.load(infer_model_path, map_location='cpu')
|
||||
audiolm_state_dict = {k.replace('audiolm.', ''): v for k, v in checkpoint.items() if k.startswith('audiolm')}
|
||||
audiolm.load_state_dict(audiolm_state_dict, strict=False)
|
||||
audiolm = audiolm.eval()
|
||||
#audiolm = audiolm.cuda().to(torch.float16)
|
||||
del audiolm_state_dict,checkpoint
|
||||
return audiolm,cfg
|
||||
|
||||
original_item=song_infer_lowram(cfg,separator,audio_tokenizer,merge_prompt,auto_prompt, save_dir,prompt_audio_path,auto_prompt_audio_type)
|
||||
print("step1 is done.")
|
||||
return copy.deepcopy(original_item),max_duration,cfg
|
||||
|
||||
|
||||
def infer_stage2(item,cfg,Weigths_Path,max_duration,lyric,descriptions,cfg_coef = 1.5, temp = 0.9,top_k = 50,top_p = 0.0,record_tokens = True,record_window = 50):
|
||||
ckpt_path = os.path.join(Weigths_Path, 'songgeneration_base/model.pt')
|
||||
# def pre_data(Weigths_Path,dm_model_path,dm_config_path,save_dir,prompt_audio_path,auto_prompt_audio_type,infer_model_type,prompt_pt_path):
|
||||
# torch.backends.cudnn.enabled = False
|
||||
# curent_dir = os.path.join(current_node_path,"SongGeneration")
|
||||
# RESOLVERS = {
|
||||
# "eval": lambda x: eval(x),
|
||||
# "concat": lambda *x: [xxx for xx in x for xxx in xx],
|
||||
# "get_fname": lambda: os.path.splitext(os.path.basename(sys.argv[1]))[0],
|
||||
# "load_yaml": lambda x: list(OmegaConf.load(os.path.join(curent_dir, x)))
|
||||
# }
|
||||
|
||||
# for name, func in RESOLVERS.items():
|
||||
# if not OmegaConf.has_resolver(name):
|
||||
# OmegaConf.register_new_resolver(name, func)
|
||||
# np.random.seed(int(time.time()))
|
||||
|
||||
# cfg_path = os.path.join(current_node_path, f'SongGeneration/conf/{infer_model_type}_config.yaml')
|
||||
|
||||
# cfg = OmegaConf.load(cfg_path)
|
||||
# cfg.mode = 'inference'
|
||||
|
||||
# cfg.vae_config=f"{Weigths_Path}/vae/stable_audio_1920_vae.json"
|
||||
# cfg.vae_model=f"{Weigths_Path}/vae/autoencoder_music_1320k.ckpt"
|
||||
|
||||
# cfg.audio_tokenizer_checkpoint=f"Flow1dVAE1rvq_{Weigths_Path}/model_1rvq/model_2_fixed.safetensors"
|
||||
# cfg.audio_tokenizer_checkpoint_sep=f"Flow1dVAESeparate_{Weigths_Path}/model_septoken/model_2.safetensors"
|
||||
# cfg.conditioners.type_info.QwTextTokenizer.token_path=os.path.join(current_node_path,"SongGeneration/third_party/Qwen2-7B")
|
||||
# max_duration = cfg.max_dur
|
||||
# vae_model=f"{Weigths_Path}/vae/autoencoder_music_1320k.ckpt"
|
||||
# vae_config=os.path.join(current_node_path, f'SongGeneration/conf/stable_audio_1920_vae.json')
|
||||
# auto_prompt = torch.load(prompt_pt_path,weights_only=False)
|
||||
# merge_prompt = [x for sublist in auto_prompt.values() for x in sublist]
|
||||
|
||||
# if prompt_audio_path is not None:
|
||||
# separator = Separator(dm_model_path, dm_config_path)
|
||||
# audio_tokenizer = builders.get_audio_tokenizer_model(f"Flow1dVAE1rvq_{Weigths_Path}/model_1rvq/model_2_fixed.safetensors", vae_config,vae_model)
|
||||
# audio_tokenizer = audio_tokenizer.eval().cuda()
|
||||
|
||||
# else:
|
||||
# audio_tokenizer = None
|
||||
# separator = None
|
||||
|
||||
# original_item=song_infer_lowram(cfg,separator,audio_tokenizer,merge_prompt,auto_prompt, save_dir,prompt_audio_path,auto_prompt_audio_type,)
|
||||
# print("step1 is done.")
|
||||
# return copy.deepcopy(original_item),max_duration,cfg
|
||||
|
||||
|
||||
def infer_stage2(item,audiolm,max_duration,lyric,descriptions,cfg_coef = 1.5, temp = 0.9,top_k = 50,top_p = 0.0,record_tokens = True,record_window = 50):
|
||||
#ckpt_path = os.path.join(Weigths_Path, 'songgeneration_base/model.pt')
|
||||
|
||||
item_copy = {
|
||||
'pmt_wav': item['pmt_wav'], # 这些是引用,但安全因为后续设为None不影响原始
|
||||
@@ -130,13 +168,11 @@ def infer_stage2(item,cfg,Weigths_Path,max_duration,lyric,descriptions,cfg_coef
|
||||
# seperate_tokenizer = None,
|
||||
# )
|
||||
# del model_light
|
||||
audiolm = builders.get_lm_model(cfg)
|
||||
checkpoint = torch.load(ckpt_path, map_location='cpu')
|
||||
audiolm_state_dict = {k.replace('audiolm.', ''): v for k, v in checkpoint.items() if k.startswith('audiolm')}
|
||||
audiolm.load_state_dict(audiolm_state_dict, strict=False)
|
||||
audiolm = audiolm.eval()
|
||||
audiolm = audiolm.cuda().to(torch.float16)
|
||||
del audiolm_state_dict,checkpoint
|
||||
# audiolm = builders.get_lm_model(cfg)
|
||||
# checkpoint = torch.load(ckpt_path, map_location='cpu')
|
||||
# audiolm_state_dict = {k.replace('audiolm.', ''): v for k, v in checkpoint.items() if k.startswith('audiolm')}
|
||||
# audiolm.load_state_dict(audiolm_state_dict, strict=False)
|
||||
audiolm=audiolm.cuda().to(torch.float16)
|
||||
torch.cuda.empty_cache()
|
||||
model = CodecLM(name = "tmp",
|
||||
lm = audiolm,
|
||||
@@ -147,6 +183,7 @@ def infer_stage2(item,cfg,Weigths_Path,max_duration,lyric,descriptions,cfg_coef
|
||||
|
||||
model.set_generation_params(duration=max_duration, extend_stride=5, temperature=temp,
|
||||
top_k=top_k, top_p=top_p,cfg_coef=cfg_coef, record_tokens=record_tokens, record_window=record_window)
|
||||
|
||||
print("model loaded,start inference step2")
|
||||
items=inference_lowram_step2(model,lyric,descriptions,item_copy,)
|
||||
audiolm = audiolm.cpu()
|
||||
@@ -154,25 +191,19 @@ def infer_stage2(item,cfg,Weigths_Path,max_duration,lyric,descriptions,cfg_coef
|
||||
model=None
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return items
|
||||
|
||||
|
||||
|
||||
def inference_lowram_step2(model,lyric,descriptions,item,):
|
||||
#print(item)
|
||||
pmt_wav = item['pmt_wav']
|
||||
vocal_wav = item['vocal_wav']
|
||||
bgm_wav = item['bgm_wav']
|
||||
melody_is_wav = item['melody_is_wav']
|
||||
|
||||
|
||||
generate_inp = {
|
||||
'lyrics': [lyric.replace(" ", " ")],
|
||||
'descriptions': [descriptions],
|
||||
'melody_wavs': pmt_wav,
|
||||
'vocal_wavs': vocal_wav,
|
||||
'bgm_wavs': bgm_wav,
|
||||
'melody_is_wav': melody_is_wav,
|
||||
'melody_wavs': item['pmt_wav'],
|
||||
'vocal_wavs': item['vocal_wav'],
|
||||
'bgm_wavs': item['bgm_wav'],
|
||||
'melody_is_wav': item['melody_is_wav'],
|
||||
}
|
||||
with torch.autocast(device_type="cuda", dtype=torch.float16):
|
||||
tokens = model.generate(**generate_inp, return_tokens=True)
|
||||
@@ -182,17 +213,17 @@ def inference_lowram_step2(model,lyric,descriptions,item,):
|
||||
|
||||
|
||||
|
||||
def inference_lowram_final(cfg,max_duration,item,save_dir,save_separate):
|
||||
def inference_lowram_final(cfg,seperate_tokenizer,max_duration,item,save_dir,save_separate):
|
||||
target_wav_name = f"{save_dir}/song_audios{time.strftime('%m%d%H%S')}.flac"
|
||||
seperate_tokenizer = builders.get_audio_tokenizer_model(cfg.audio_tokenizer_checkpoint_sep, cfg)
|
||||
seperate_tokenizer = seperate_tokenizer.eval().cuda()
|
||||
|
||||
model = CodecLM(name = "tmp",
|
||||
lm = None,
|
||||
audiotokenizer = None,
|
||||
max_duration = max_duration,
|
||||
seperate_tokenizer = seperate_tokenizer,
|
||||
)
|
||||
|
||||
print("model loaded,start inference final...")
|
||||
|
||||
with torch.no_grad():
|
||||
if item["melody_is_wav"]:
|
||||
if save_separate :
|
||||
@@ -222,30 +253,28 @@ def inference_lowram_final(cfg,max_duration,item,save_dir,save_separate):
|
||||
# item['vocal_wav']=None
|
||||
# item['bgm_wav']=None
|
||||
# item['melody_is_wav']=None
|
||||
|
||||
return {"waveform": wav_seperate[0].cpu().float().unsqueeze(0), "sample_rate": cfg.sample_rate}
|
||||
|
||||
|
||||
|
||||
def song_infer_lowram(cfg,separator,audio_tokenizer,merge_prompt,auto_prompt, save_dir,prompt_audio_path,auto_prompt_audio_type): #item dict
|
||||
def song_infer_lowram(seperate_tokenizer,separator,audio_tokenizer,prompt_pt_path, save_dir,prompt_audio_path,auto_prompt_audio_type): #item dict
|
||||
item = {}
|
||||
target_wav_name = f"{save_dir}/song_audios{time.strftime('%m%d%H%S')}.flac"
|
||||
melody_is_wav = False
|
||||
if prompt_audio_path:
|
||||
|
||||
if prompt_audio_path is not None:
|
||||
|
||||
pmt_wav, vocal_wav, bgm_wav = separator.run(prompt_audio_path)
|
||||
pmt_wav = pmt_wav.cuda()
|
||||
vocal_wav = vocal_wav.cuda()
|
||||
bgm_wav = bgm_wav.cuda()
|
||||
|
||||
audio_tokenizer = audio_tokenizer.eval().cuda()
|
||||
with torch.no_grad():
|
||||
pmt_wav, _ = audio_tokenizer.encode(pmt_wav)
|
||||
audio_tokenizer=None
|
||||
separator=None
|
||||
gc.collect()
|
||||
if "audio_tokenizer_checkpoint_sep" in cfg.keys():
|
||||
seperate_tokenizer = builders.get_audio_tokenizer_model(cfg.audio_tokenizer_checkpoint_sep, cfg)
|
||||
else:
|
||||
raise ValueError("No audio tokenizer checkpoint found")
|
||||
|
||||
seperate_tokenizer = seperate_tokenizer.eval().cuda()
|
||||
with torch.no_grad():
|
||||
vocal_wav, bgm_wav = seperate_tokenizer.encode(vocal_wav, bgm_wav)
|
||||
@@ -253,11 +282,16 @@ def song_infer_lowram(cfg,separator,audio_tokenizer,merge_prompt,auto_prompt, sa
|
||||
gc.collect()
|
||||
|
||||
elif auto_prompt_audio_type:
|
||||
assert prompt_pt_path is not None ,"prompt模型不能为空,need prmmpt model"
|
||||
auto_prompt = torch.load(prompt_pt_path,weights_only=False)
|
||||
#assert item["auto_prompt_audio_type"] in auto_prompt_type, f"auto_prompt_audio_type {item['auto_prompt_audio_type']} not found"
|
||||
if auto_prompt_audio_type == 'Auto':
|
||||
prompt_token = merge_prompt[np.random.randint(0, len(merge_prompt))]
|
||||
else:
|
||||
prompt_token = auto_prompt[auto_prompt_audio_type][np.random.randint(0, len(auto_prompt[auto_prompt_audio_type]))]
|
||||
# if auto_prompt_audio_type == 'Auto':
|
||||
# prompt_token = merge_prompt[np.random.randint(0, len(merge_prompt))]
|
||||
# else:
|
||||
# prompt_token = auto_prompt[auto_prompt_audio_type][np.random.randint(0, len(auto_prompt[auto_prompt_audio_type]))]
|
||||
prompt_token = auto_prompt[auto_prompt_audio_type][np.random.randint(0, len(auto_prompt[auto_prompt_audio_type]))]
|
||||
if torch.cuda.is_available():
|
||||
prompt_token = prompt_token.cuda()
|
||||
pmt_wav = prompt_token[:,[0],:]
|
||||
vocal_wav = prompt_token[:,[1],:]
|
||||
bgm_wav = prompt_token[:,[2],:]
|
||||
@@ -266,6 +300,7 @@ def song_infer_lowram(cfg,separator,audio_tokenizer,merge_prompt,auto_prompt, sa
|
||||
vocal_wav = None
|
||||
bgm_wav = None
|
||||
melody_is_wav = True
|
||||
|
||||
item['pmt_wav'] = pmt_wav
|
||||
item['vocal_wav'] = vocal_wav
|
||||
item['bgm_wav'] = bgm_wav
|
||||
|
||||
+1
-1
File diff suppressed because one or more lines are too long
Reference in New Issue
Block a user