Merge pull request #19 from smthemex/Pr1

init
This commit is contained in:
smthemex
2025-10-18 11:48:49 +08:00
committed by GitHub
83 changed files with 3490 additions and 344 deletions
+17 -16
View File
@@ -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
![](https://github.com/smthemex/ComfyUI_SongGeneration/blob/main/example_workflows/example.png)
![](https://github.com/smthemex/ComfyUI_SongGeneration/blob/main/example_workflows/SongGeneration.png)
# 5 Citation
```
+3 -4
View File
@@ -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 \
@@ -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,
@@ -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.")
+141
View File
@@ -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
+141
View File
@@ -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
+141
View File
@@ -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
+141
View File
@@ -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
View File
@@ -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
View File
@@ -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)
+1
View File
@@ -0,0 +1 @@
+54
View File
@@ -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)
+4
View File
@@ -0,0 +1,4 @@
from .base import CodecMixin
from .base import DACFile
from .dac import DAC
from .discriminator import Discriminator
+294
View File
@@ -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
View File
@@ -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)
+228
View File
@@ -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()
+3
View File
@@ -0,0 +1,3 @@
from . import layers
from . import loss
from . import quantize
+34
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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()
+82 -43
View File
@@ -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",
+598
View File
@@ -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

-195
View File
@@ -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
View File
@@ -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
View File
File diff suppressed because one or more lines are too long