diff --git a/README.md b/README.md index 8daa779..bf0eb26 100644 --- a/README.md +++ b/README.md @@ -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 ``` diff --git a/SongGeneration/codeclm/models/builders.py b/SongGeneration/codeclm/models/builders.py index 1944381..25075a3 100644 --- a/SongGeneration/codeclm/models/builders.py +++ b/SongGeneration/codeclm/models/builders.py @@ -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.""" diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/generate_1rvq.py b/SongGeneration/codeclm/tokenizer/Flow1dVAE/generate_1rvq.py index 4e1a9b3..f580811 100644 --- a/SongGeneration/codeclm/tokenizer/Flow1dVAE/generate_1rvq.py +++ b/SongGeneration/codeclm/tokenizer/Flow1dVAE/generate_1rvq.py @@ -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): diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/generate_septoken.py b/SongGeneration/codeclm/tokenizer/Flow1dVAE/generate_septoken.py index f47b670..585a354 100644 --- a/SongGeneration/codeclm/tokenizer/Flow1dVAE/generate_septoken.py +++ b/SongGeneration/codeclm/tokenizer/Flow1dVAE/generate_septoken.py @@ -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): diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/model_septoken.py b/SongGeneration/codeclm/tokenizer/Flow1dVAE/model_septoken.py index 61f64db..a08c063 100644 --- a/SongGeneration/codeclm/tokenizer/Flow1dVAE/model_septoken.py +++ b/SongGeneration/codeclm/tokenizer/Flow1dVAE/model_septoken.py @@ -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 \ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/models_gpt/models/gpt2_rope2_time_new_correct_mask_noncasual_reflow.py b/SongGeneration/codeclm/tokenizer/Flow1dVAE/models_gpt/models/gpt2_rope2_time_new_correct_mask_noncasual_reflow.py index a77a464..71ee6de 100644 --- a/SongGeneration/codeclm/tokenizer/Flow1dVAE/models_gpt/models/gpt2_rope2_time_new_correct_mask_noncasual_reflow.py +++ b/SongGeneration/codeclm/tokenizer/Flow1dVAE/models_gpt/models/gpt2_rope2_time_new_correct_mask_noncasual_reflow.py @@ -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, diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/data/__pycache__/__init__.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/data/__pycache__/__init__.cpython-311.pyc deleted file mode 100644 index e91fa38..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/data/__pycache__/__init__.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/data/__pycache__/ark_dataset.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/data/__pycache__/ark_dataset.cpython-311.pyc deleted file mode 100644 index 80f0f10..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/data/__pycache__/ark_dataset.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/data/__pycache__/mert_dataset.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/data/__pycache__/mert_dataset.cpython-311.pyc deleted file mode 100644 index f5b4b1f..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/data/__pycache__/mert_dataset.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/data/eat_data/__pycache__/__init__.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/data/eat_data/__pycache__/__init__.cpython-311.pyc deleted file mode 100644 index a9ea2bb..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/data/eat_data/__pycache__/__init__.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/data/eat_data/__pycache__/mae_image_dataset.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/data/eat_data/__pycache__/mae_image_dataset.cpython-311.pyc deleted file mode 100644 index 649d9f1..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/data/eat_data/__pycache__/mae_image_dataset.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/data/eat_data/__pycache__/raw_audio_dataset.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/data/eat_data/__pycache__/raw_audio_dataset.cpython-311.pyc deleted file mode 100644 index 3fa02fb..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/data/eat_data/__pycache__/raw_audio_dataset.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/data/utils/__pycache__/data_utils.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/data/utils/__pycache__/data_utils.cpython-311.pyc deleted file mode 100644 index a6eaddb..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/data/utils/__pycache__/data_utils.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/__pycache__/__init__.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/__pycache__/__init__.cpython-311.pyc deleted file mode 100644 index 59e7e13..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/__pycache__/__init__.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/eat/__pycache__/EAT_pretraining.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/eat/__pycache__/EAT_pretraining.cpython-311.pyc deleted file mode 100644 index 0604687..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/eat/__pycache__/EAT_pretraining.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/eat/__pycache__/__init__.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/eat/__pycache__/__init__.cpython-311.pyc deleted file mode 100644 index db3679e..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/eat/__pycache__/__init__.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/eat/__pycache__/base.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/eat/__pycache__/base.cpython-311.pyc deleted file mode 100644 index 56fca4d..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/eat/__pycache__/base.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/eat/__pycache__/images.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/eat/__pycache__/images.cpython-311.pyc deleted file mode 100644 index 8ddaeb1..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/eat/__pycache__/images.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/eat/__pycache__/mae.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/eat/__pycache__/mae.cpython-311.pyc deleted file mode 100644 index 119908b..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/eat/__pycache__/mae.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/eat/__pycache__/modules.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/eat/__pycache__/modules.cpython-311.pyc deleted file mode 100644 index 819059b..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/eat/__pycache__/modules.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/eat/__pycache__/moduless.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/eat/__pycache__/moduless.cpython-311.pyc deleted file mode 100644 index b848169..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/eat/__pycache__/moduless.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/mert/__pycache__/__init__.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/mert/__pycache__/__init__.cpython-311.pyc deleted file mode 100644 index 5cdd33c..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/mert/__pycache__/__init__.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/mert/__pycache__/mert_model.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/mert/__pycache__/mert_model.cpython-311.pyc deleted file mode 100644 index ecd8cad..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/mert/__pycache__/mert_model.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/mert/__pycache__/rvq.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/mert/__pycache__/rvq.cpython-311.pyc deleted file mode 100644 index 9ecb152..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/mert/__pycache__/rvq.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/__pycache__/__init__.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/__pycache__/__init__.cpython-311.pyc deleted file mode 100644 index ad3cd7a..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/__pycache__/__init__.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/__pycache__/musicfm_model.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/__pycache__/musicfm_model.cpython-311.pyc deleted file mode 100644 index 3a98e53..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/__pycache__/musicfm_model.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/model/__pycache__/__init__.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/model/__pycache__/__init__.cpython-311.pyc deleted file mode 100644 index 45c8343..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/model/__pycache__/__init__.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/model/__pycache__/musicfm_25hz.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/model/__pycache__/musicfm_25hz.cpython-311.pyc deleted file mode 100644 index 9dbfff6..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/model/__pycache__/musicfm_25hz.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/model/__pycache__/rvq.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/model/__pycache__/rvq.cpython-311.pyc deleted file mode 100644 index d0b3762..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/model/__pycache__/rvq.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/model/__pycache__/rvq_musicfm.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/model/__pycache__/rvq_musicfm.cpython-311.pyc deleted file mode 100644 index 26647b5..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/model/__pycache__/rvq_musicfm.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/modules/__pycache__/__init__.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/modules/__pycache__/__init__.cpython-311.pyc deleted file mode 100644 index e30e5cb..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/modules/__pycache__/__init__.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/modules/__pycache__/conv.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/modules/__pycache__/conv.cpython-311.pyc deleted file mode 100644 index ad6d5f2..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/modules/__pycache__/conv.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/modules/__pycache__/features.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/modules/__pycache__/features.cpython-311.pyc deleted file mode 100644 index 54a79b1..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/modules/__pycache__/features.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/modules/__pycache__/random_quantizer.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/modules/__pycache__/random_quantizer.cpython-311.pyc deleted file mode 100644 index 54846ac..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/models/musicfm/modules/__pycache__/random_quantizer.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/tasks/__pycache__/mert_pretraining.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/tasks/__pycache__/mert_pretraining.cpython-311.pyc deleted file mode 100644 index 66e58d2..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/tasks/__pycache__/mert_pretraining.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/tasks/__pycache__/pretraining_AS2M.cpython-311.pyc b/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/tasks/__pycache__/pretraining_AS2M.cpython-311.pyc deleted file mode 100644 index fc669c7..0000000 Binary files a/SongGeneration/codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq/tasks/__pycache__/pretraining_AS2M.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/codeclm/tokenizer/Flow1dVAE/tools/get_1dvae_large.py b/SongGeneration/codeclm/tokenizer/Flow1dVAE/tools/get_1dvae_large.py index 915dfc7..9230331 100644 --- a/SongGeneration/codeclm/tokenizer/Flow1dVAE/tools/get_1dvae_large.py +++ b/SongGeneration/codeclm/tokenizer/Flow1dVAE/tools/get_1dvae_large.py @@ -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 diff --git a/SongGeneration/codeclm/tokenizer/audio_tokenizer.py b/SongGeneration/codeclm/tokenizer/audio_tokenizer.py index 921670f..00d010c 100644 --- a/SongGeneration/codeclm/tokenizer/audio_tokenizer.py +++ b/SongGeneration/codeclm/tokenizer/audio_tokenizer.py @@ -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.") diff --git a/SongGeneration/conf/base_config.yaml b/SongGeneration/conf/base_config.yaml new file mode 100644 index 0000000..8246a1b --- /dev/null +++ b/SongGeneration/conf/base_config.yaml @@ -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 diff --git a/SongGeneration/conf/full_config.yaml b/SongGeneration/conf/full_config.yaml new file mode 100644 index 0000000..42d8b22 --- /dev/null +++ b/SongGeneration/conf/full_config.yaml @@ -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 diff --git a/SongGeneration/conf/large_config.yaml b/SongGeneration/conf/large_config.yaml new file mode 100644 index 0000000..3a951ea --- /dev/null +++ b/SongGeneration/conf/large_config.yaml @@ -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 diff --git a/SongGeneration/conf/new_config.yaml b/SongGeneration/conf/new_config.yaml new file mode 100644 index 0000000..8246a1b --- /dev/null +++ b/SongGeneration/conf/new_config.yaml @@ -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 diff --git a/SongGeneration/conf/stable_audio_1920_vae.json b/SongGeneration/conf/stable_audio_1920_vae.json new file mode 100644 index 0000000..0d74ab4 --- /dev/null +++ b/SongGeneration/conf/stable_audio_1920_vae.json @@ -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 + } + } +} diff --git a/SongGeneration/third_party/dac/__init__.py b/SongGeneration/third_party/dac/__init__.py new file mode 100644 index 0000000..51205ef --- /dev/null +++ b/SongGeneration/third_party/dac/__init__.py @@ -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 diff --git a/SongGeneration/third_party/dac/__main__.py b/SongGeneration/third_party/dac/__main__.py new file mode 100644 index 0000000..2fa8d15 --- /dev/null +++ b/SongGeneration/third_party/dac/__main__.py @@ -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) diff --git a/SongGeneration/third_party/dac/compare/__init__.py b/SongGeneration/third_party/dac/compare/__init__.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/SongGeneration/third_party/dac/compare/__init__.py @@ -0,0 +1 @@ + diff --git a/SongGeneration/third_party/dac/compare/encodec.py b/SongGeneration/third_party/dac/compare/encodec.py new file mode 100644 index 0000000..42877de --- /dev/null +++ b/SongGeneration/third_party/dac/compare/encodec.py @@ -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) diff --git a/SongGeneration/third_party/dac/model/__init__.py b/SongGeneration/third_party/dac/model/__init__.py new file mode 100644 index 0000000..02a75b7 --- /dev/null +++ b/SongGeneration/third_party/dac/model/__init__.py @@ -0,0 +1,4 @@ +from .base import CodecMixin +from .base import DACFile +from .dac import DAC +from .discriminator import Discriminator diff --git a/SongGeneration/third_party/dac/model/base.py b/SongGeneration/third_party/dac/model/base.py new file mode 100644 index 0000000..546b3cb --- /dev/null +++ b/SongGeneration/third_party/dac/model/base.py @@ -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 diff --git a/SongGeneration/third_party/dac/model/dac.py b/SongGeneration/third_party/dac/model/dac.py new file mode 100644 index 0000000..6d44a18 --- /dev/null +++ b/SongGeneration/third_party/dac/model/dac.py @@ -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) diff --git a/SongGeneration/third_party/dac/model/discriminator.py b/SongGeneration/third_party/dac/model/discriminator.py new file mode 100644 index 0000000..09c79d1 --- /dev/null +++ b/SongGeneration/third_party/dac/model/discriminator.py @@ -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() diff --git a/SongGeneration/third_party/dac/nn/__init__.py b/SongGeneration/third_party/dac/nn/__init__.py new file mode 100644 index 0000000..6718c8b --- /dev/null +++ b/SongGeneration/third_party/dac/nn/__init__.py @@ -0,0 +1,3 @@ +from . import layers +from . import loss +from . import quantize diff --git a/SongGeneration/third_party/dac/nn/layers.py b/SongGeneration/third_party/dac/nn/layers.py new file mode 100644 index 0000000..a4843de --- /dev/null +++ b/SongGeneration/third_party/dac/nn/layers.py @@ -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) + diff --git a/SongGeneration/third_party/dac/nn/loss.py b/SongGeneration/third_party/dac/nn/loss.py new file mode 100644 index 0000000..9bb3dd6 --- /dev/null +++ b/SongGeneration/third_party/dac/nn/loss.py @@ -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 diff --git a/SongGeneration/third_party/dac/nn/quantize.py b/SongGeneration/third_party/dac/nn/quantize.py new file mode 100644 index 0000000..a21b383 --- /dev/null +++ b/SongGeneration/third_party/dac/nn/quantize.py @@ -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) diff --git a/SongGeneration/third_party/dac/utils/__init__.py b/SongGeneration/third_party/dac/utils/__init__.py new file mode 100644 index 0000000..6a07a69 --- /dev/null +++ b/SongGeneration/third_party/dac/utils/__init__.py @@ -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 diff --git a/SongGeneration/third_party/dac/utils/decode.py b/SongGeneration/third_party/dac/utils/decode.py new file mode 100644 index 0000000..08d44e8 --- /dev/null +++ b/SongGeneration/third_party/dac/utils/decode.py @@ -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() diff --git a/SongGeneration/third_party/dac/utils/encode.py b/SongGeneration/third_party/dac/utils/encode.py new file mode 100644 index 0000000..aa3f6f4 --- /dev/null +++ b/SongGeneration/third_party/dac/utils/encode.py @@ -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() diff --git a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/data/__pycache__/__init__.cpython-311.pyc b/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/data/__pycache__/__init__.cpython-311.pyc deleted file mode 100644 index b47ad02..0000000 Binary files a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/data/__pycache__/__init__.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/data/__pycache__/utils.cpython-311.pyc b/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/data/__pycache__/utils.cpython-311.pyc deleted file mode 100644 index 921f49f..0000000 Binary files a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/data/__pycache__/utils.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/__init__.cpython-311.pyc b/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/__init__.cpython-311.pyc deleted file mode 100644 index 9364ae7..0000000 Binary files a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/__init__.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/adp.cpython-311.pyc b/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/adp.cpython-311.pyc deleted file mode 100644 index 8e5bfb6..0000000 Binary files a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/adp.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/autoencoders.cpython-311.pyc b/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/autoencoders.cpython-311.pyc deleted file mode 100644 index fa2d9ae..0000000 Binary files a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/autoencoders.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/blocks.cpython-311.pyc b/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/blocks.cpython-311.pyc deleted file mode 100644 index 8ad5515..0000000 Binary files a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/blocks.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/bottleneck.cpython-311.pyc b/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/bottleneck.cpython-311.pyc deleted file mode 100644 index 2de4b61..0000000 Binary files a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/bottleneck.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/conditioners.cpython-311.pyc b/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/conditioners.cpython-311.pyc deleted file mode 100644 index 5d22d63..0000000 Binary files a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/conditioners.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/diffusion.cpython-311.pyc b/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/diffusion.cpython-311.pyc deleted file mode 100644 index 5823e4b..0000000 Binary files a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/diffusion.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/dit.cpython-311.pyc b/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/dit.cpython-311.pyc deleted file mode 100644 index 7e7551e..0000000 Binary files a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/dit.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/factory.cpython-311.pyc b/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/factory.cpython-311.pyc deleted file mode 100644 index ae4567a..0000000 Binary files a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/factory.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/pretrained.cpython-311.pyc b/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/pretrained.cpython-311.pyc deleted file mode 100644 index 71724c1..0000000 Binary files a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/pretrained.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/pretransforms.cpython-311.pyc b/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/pretransforms.cpython-311.pyc deleted file mode 100644 index 54ce41d..0000000 Binary files a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/pretransforms.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/transformer.cpython-311.pyc b/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/transformer.cpython-311.pyc deleted file mode 100644 index ffbc775..0000000 Binary files a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/transformer.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/utils.cpython-311.pyc b/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/utils.cpython-311.pyc deleted file mode 100644 index 0383c08..0000000 Binary files a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/models/__pycache__/utils.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/training/__pycache__/__init__.cpython-311.pyc b/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/training/__pycache__/__init__.cpython-311.pyc deleted file mode 100644 index 8e64e99..0000000 Binary files a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/training/__pycache__/__init__.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/training/__pycache__/factory.cpython-311.pyc b/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/training/__pycache__/factory.cpython-311.pyc deleted file mode 100644 index 65c44be..0000000 Binary files a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/training/__pycache__/factory.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/training/__pycache__/utils.cpython-311.pyc b/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/training/__pycache__/utils.cpython-311.pyc deleted file mode 100644 index 61126d0..0000000 Binary files a/SongGeneration/third_party/stable_audio_tools/stable_audio_tools/training/__pycache__/utils.cpython-311.pyc and /dev/null differ diff --git a/SongGeneration_node.py b/SongGeneration_node.py index 8c62c86..c03d0c3 100644 --- a/SongGeneration_node.py +++ b/SongGeneration_node.py @@ -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", diff --git a/example_workflows/SongGeneration.json b/example_workflows/SongGeneration.json new file mode 100644 index 0000000..73fb78d --- /dev/null +++ b/example_workflows/SongGeneration.json @@ -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 +} \ No newline at end of file diff --git a/example_workflows/SongGeneration.png b/example_workflows/SongGeneration.png new file mode 100644 index 0000000..bf3efd6 Binary files /dev/null and b/example_workflows/SongGeneration.png differ diff --git a/example_workflows/Workflow_song.json b/example_workflows/Workflow_song.json deleted file mode 100644 index e738ac7..0000000 --- a/example_workflows/Workflow_song.json +++ /dev/null @@ -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 -} \ No newline at end of file diff --git a/example_workflows/example.png b/example_workflows/example.png deleted file mode 100644 index ec3f39d..0000000 Binary files a/example_workflows/example.png and /dev/null differ diff --git a/generate.py b/generate.py index 9cdce22..8cadec8 100644 --- a/generate.py +++ b/generate.py @@ -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 diff --git a/pyproject.toml b/pyproject.toml index 83ce36c..a94906c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "songgeneration" description = "SongGeneration:High-Quality Song Generation with Multi-Preference Alignment (SOTA),you can try VRAM>12G" -version = "1.0.0" +version = "1.0.1" license = {file = "LICENSE"} dependencies = ["# absl-py==2.0.0", "# accelerate==0.30.1", "accelerate", "# addict==2.4.0", "# aiofiles==23.2.1", "aiohttp", "# aiosignal==1.3.1", "alias-free-torch", "# aliyun-python-sdk-core==2.15.1", "# aliyun-python-sdk-kms==2.16.3", "# altair==5.3.0", "# annotated-types==0.6.0", "# antlr4-python3-runtime==4.8", "# anyio==4.3.0", "# argbind==0.3.9", "# asttokens==3.0.0", "# astunparse==1.6.3", "# async-timeout==4.0.3", "# attrs==23.1.0", "audiocraft", "#audioread==3.0.1", "av", "# backcall==0.2.0", "# beartype==0.18.5", "# bitarray==2.9.2", "# blis==0.7.11", "# boto3==1.29.6", "# botocore==1.32.6", "# braceexpand==0.1.7", "# cachetools==5.3.2", "# catalogue==2.0.10", "# certifi==2023.11.17", "# cffi==1.16.0", "# charset-normalizer==3.3.2", "# clean-fid==0.1.35", "# click==8.1.7", "# clip-anytorch==2.6.0", "# cloudpathlib==0.16.0", "cloudpickle", "# cn2an==0.5.22", "# colorama==0.4.6", "colorlog", "# confection==0.1.4", "# contourpy==1.1.1", "# crcmod==1.7", "# cryptography==43.0.0", "# cycler==0.12.1", "# cymem==2.0.8", "# Cython==3.0.10", "# dataclasses==0.6", "# datasets==2.18.0", "# dctorch==0.1.2", "# decorator==5.1.1", "# decord==0.6.0", "# deepspeed==0.14.0", "# demucs==4.0.1", "descript-audio-codec", "descript-audiotools", "diffusers", "# dill==0.3.8", "# Distance==0.1.3", "# docker-pycreds==0.4.0", "# docopt==0.6.2", "# docstring_parser==0.16", "dora_search", "einops", "einops-exts", "einx", "# ema-pytorch==0.5.1", "# encodec==0.1.1", "# exceptiongroup==1.2.0", "# executing==2.2.0", "# expecttest==0.1.6", "#fairseq==0.12.2 #pip install fairseq -f https://download.pytorch.org/whl/cu124/torch_stable.html", "fairseq", "# fastapi==0.110.3", "# ffmpy==0.3.2", "# filelock==3.13.1", "# fire==0.7.0", "flashy", "# flatten-dict==0.4.2", "# fonttools==4.49.0", "frozendict", "# frozenlist==1.4.1", "# fsspec==2023.10.0", "# ftfy==6.1.3", "# future==1.0.0", "# g2p-en", "# gitdb==4.0.11", "# GitPython==3.1.43", "# google-auth==2.23.4", "# google-auth-oauthlib==1.0.0", "# gradio==4.26.0", "# gradio_client==0.15.1", "# grpcio==1.59.3", "# h11==0.14.0", "# h5py==3.11.0", "# hf-xet==1.1.2", "# hjson==3.1.0", "# httpcore==1.0.5", "# httpx==0.27.0", "# huggingface-hub==0.25.2", "# hydra-colorlog==1.2.0", "# hydra-core==1.0.7", "# hypothesis==6.90.0", "# idna==3.4", "# imageio==2.35.1", "# importlib-metadata==6.8.0", "# importlib_resources==6.1.3", "# inflect==7.0.0", "# ipython==8.12.3", "# jedi==0.19.2", "# jieba-fast==0.53", "# Jinja2==3.1.2", "# jmespath==0.10.0", "# joblib==1.3.2", "# json5==0.9.25", "# jsonlines==4.0.0", "# jsonmerge==1.9.2", "# jsonschema==4.22.0", "# jsonschema-specifications==2023.12.1", "# julius==0.2.7", "# k-diffusion==0.1.1.post1", "kaldiio", "# kiwisolver==1.4.5", "# kornia==0.7.3", "# kornia_rs==0.1.9", "lameenc", "# langcodes==3.4.0", "# language_data==1.2.0", "# lazy_loader==0.3", "# librosa==0.9.1", "# lightning==2.2.1", "# lightning-utilities==0.10.1", "# lion-pytorch==0.2.2", "# llvmlite==0.41.1", "# local-attention==1.9.14", "# loguru==0.7.2", "# lxml==5.2.2", "# marisa-trie==1.1.1", "# Markdown==3.5.1", "# markdown-it-py==3.0.0", "# markdown2==2.5.1", "# MarkupSafe==2.1.3", "# matplotlib==3.7.5", "# matplotlib-inline==0.1.7", "# mdurl==0.1.2", "# modelscope==1.16.1", "# mpmath==1.3.0", "# msgpack==1.0.8", "# multidict==6.0.5", "# multiprocess==0.70.16", "# murmurhash==1.0.10", "# mypy-extensions==1.0.0", "# networkx==3.1", "# ninja==1.11.1.1", "# nltk==3.8.1", "nnAudio", "# num2words==0.5.13", "# numba==0.58.1", "# numpy==1.23.5", "# nvidia-cublas-cu11==11.11.3.6", "# nvidia-cuda-cupti-cu11==11.8.87", "# nvidia-cuda-nvrtc-cu11==11.8.89", "# nvidia-cuda-runtime-cu11==11.8.89", "# nvidia-cudnn-cu11==8.7.0.84", "# nvidia-cufft-cu11==10.9.0.58", "# nvidia-curand-cu11==10.3.0.86", "# nvidia-cusolver-cu11==11.4.1.48", "# nvidia-cusparse-cu11==11.7.5.86", "# nvidia-nccl-cu11==2.19.3", "# nvidia-nvtx-cu11==11.8.86", "# oauthlib==3.2.2", "# omegaconf==2.2.0 # fix", "omegaconf", "# opencv-contrib-python==4.8.1.78", "opencv-python", "openunmix", "# orjson==3.10.3", "# oss2==2.18.6", "# packaging==23.2", "# pandas==2.0.3", "# parso==0.8.4", "peft", "# pexpect==4.9.0", "# pickleshare==0.7.5", "# Pillow==10.1.0", "# pkgutil_resolve_name==1.3.10", "# platformdirs==4.2.0", "# pooch==1.8.1", "# portalocker==2.10.1", "# preshed==3.0.9", "# proces==0.1.7", "# prodict==0.8.18", "# progressbar==2.5", "# prompt_toolkit==3.0.51", "# protobuf==3.19.6", "# psutil==5.9.6", "# ptyprocess==0.7.0", "# pure_eval==0.2.3", "# py-cpuinfo==9.0.0", "# pyarrow==17.0.0", "# pyarrow-hotfix==0.6", "# pyasn1==0.5.1", "# pyasn1-modules==0.3.0", "# pybind11==2.11.1", "# pycparser==2.21", "# pycryptodome==3.20.0", "# pydantic==2.6.3", "# pydantic_core==2.16.3", "# pydub==0.25.1", "# Pygments==2.18.0", "# pyloudnorm==0.1.1", "# pynvml==11.5.0", "# pyparsing==3.1.2", "pypinyin", "# pyre-extensions==0.0.29", "# pyreaper==0.0.10", "# PySoundFile==0.9.0.post1", "# pystoi==0.4.1", "# python-dateutil==2.8.2", "# python-multipart==0.0.9", "# pytorch-lightning==2.2.1", "# pytz==2023.3.post1", "# PyWavelets==1.4.1", "# PyYAML==6.0.1", "# randomname==0.2.1", "# referencing==0.35.1", "# regex==2023.10.3", "# requests==2.32.3", "# requests-oauthlib==1.3.1", "# resampy==0.4.3", "# retrying==1.3.4", "# rich==13.7.1", "# rpds-py==0.18.1", "# rsa==4.9", "# ruamel.yaml==0.18.5", "# ruamel.yaml.clib==0.2.8", "# ruff==0.4.4", "# s3transfer==0.7.0", "# sacrebleu==2.4.2", "# safetensors==0.4.3", "# scikit-image==0.21.0", "# scikit-learn==1.3.2", "# scipy==1.10.1", "# semantic-version==2.10.0", "# sentencepiece==0.2.0", "# sentry-sdk==2.10.0", "# setproctitle==1.3.3", "# shellingham==1.5.4", "# six==1.16.0", "# smart-open==6.4.0", "# smmap==5.0.1", "# sniffio==1.3.1", "# sortedcontainers==2.4.0", "# SoundFile==0.10.3.post1", "# sox==1.4.1", "# soxr==0.3.7", "# spacy==3.7.4", "# spacy-legacy==3.0.12", "# spacy-loggers==1.0.5", "# srsly==2.4.8", "# stack-data==0.6.3", "# starlette==0.37.2", "submitit", "# sympy==1.12", "# tabulate==0.9.0", "# tensorboard==2.14.0", "# tensorboard-data-server==0.7.2", "# tensorboardX==2.6.2.2", "# termcolor==2.3.0", "# thinc==8.2.3", "# threadpoolctl==3.3.0", "# tifffile==2023.7.10", "# timm==0.9.11", "# tokenizers==0.15.2", "# tomlkit==0.12.0", "# toolz==0.12.1", "torch", "# torch-stoi==0.2.3", "torchaudio", "# torchdata==0.7.1", "# torchdiffeq==0.2.5", "# torchlibrosa==0.1.0", "# torchmetrics==1.3.1", "# torchsde==0.2.6", "# torchtext==0.17.0", "torchvision", "tqdm", "# traitlets==5.14.3", "# trampoline==0.1.2", "# transformers==4.37.2", "treetable", "triton", "# typeguard==2.13.0", "# typer==0.9.4", "# types-dataclasses==0.6.6", "# typing-inspect==0.9.0", "# typing_extensions==4.8.0", "# tzdata==2023.3", "# Unidecode==1.3.8", "# urllib3==1.26.18", "# uvicorn==0.29.0", "vector_quantize_pytorch", "# wandb==0.17.4", "# wasabi==1.1.2", "# wcwidth==0.2.12", "# weasel==0.3.4", "# webdataset==0.2.86", "# websockets==11.0.3", "# Werkzeug==3.0.1", "# wget==3.2", "# wordsegment==1.3.1", "# x-clip==0.14.4", "x-transformers", "xformers", "# yarl==1.9.4", "# zipp==3.17.0"]