@@ -2,6 +2,7 @@
|
||||
[SongGeneration](https://github.com/tencent-ailab/SongGeneration):High-Quality Song Generation with Multi-Preference Alignment (SOTA),you can try VRAM>12G
|
||||
|
||||
# Update
|
||||
* 10/23 同步官方代码,删除fairseq库,已无安装难度;
|
||||
* 10/21同步官方代码,精简模型加载,删除hubert模型,优化lm模型加载顺序,避免转移到显存时峰值OOM;
|
||||
* 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);
|
||||
@@ -16,9 +17,8 @@ git clone https://github.com/smthemex/ComfyUI_SongGeneration.git
|
||||
```
|
||||
|
||||
# 2. Requirements
|
||||
* window平台最难装的就是[fairseq](https://github.com/facebookresearch/fairseq)库,python3.11的建议用轮子安装[liyaodev/fairseq](https://github.com/liyaodev/fairseq/releases/tag/v0.12.3.1);
|
||||
|
||||
* 如果缺失库,打开requirements_orgin.txt文件,看是少了哪个,手动安装;
|
||||
* The most difficult thing to install on the Windows platform is the Fairseq library. It is recommended to install it on wheels for version 3.11 [liyaodev/fairseq](https://github.com/liyaodev/fairseq/releases/tag/v0.12.3.1);
|
||||
* If the library is missing, open the ’requirements_orgin.txt‘ file and see which one is missing, then manually install it;
|
||||
|
||||
```
|
||||
@@ -32,17 +32,16 @@ pip install -r requirements.txt
|
||||
* 3.1.4 download htdemucs.pth [tencent/SongGeneration](https://huggingface.co/tencent/SongGeneration/tree/main/third_party/demucs/ckpt)
|
||||
* 文件结构如下,修改了加载流程,原来的结构也能用:
|
||||
```
|
||||
-- ComfyUI/models/SongGeneration/
|
||||
-- ComfyUI/models/SongGeneration/ # 24.4G all 整个文件夹的大小
|
||||
|-- 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 整个文件夹的大小
|
||||
|--new_model.pt # rename from model.pt #可选
|
||||
|--large_model.pt # rename from model.pt #可选
|
||||
|-- ckpt/
|
||||
|--encode-s12k.pt # 3.68G
|
||||
|--models--lengyue233--content-vec-best/
|
||||
-- ComfyUI/models/vae/
|
||||
|--autoencoder_music_1320k.ckpt
|
||||
```
|
||||
|
||||
BIN
Binary file not shown.
@@ -19,11 +19,11 @@ from .libs.rvq.descript_quantize3 import ResidualVectorQuantize
|
||||
import folder_paths
|
||||
from .models_gpt.models.gpt2_rope2_time_new_correct_mask_noncasual_reflow import GPT2Model
|
||||
from .models_gpt.models.gpt2_config import GPT2Config
|
||||
|
||||
from .our_MERT_BESTRQ.mert_fairseq.models.musicfm.musicfm_model import MusicFMModel, MusicFMConfig
|
||||
from torch.cuda.amp import autocast
|
||||
|
||||
|
||||
from .our_MERT_BESTRQ.test import load_model
|
||||
# from .our_MERT_BESTRQ.test import load_model
|
||||
|
||||
class HubertModelWithFinalProj(HubertModel):
|
||||
def __init__(self, config):
|
||||
@@ -272,6 +272,7 @@ class PromptCondAudioDiffusion(nn.Module):
|
||||
ssl_layer=None,
|
||||
uncondition=True,
|
||||
out_paint=False,
|
||||
ssl_path='ckpt/encode-s12k.pt'
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
@@ -294,30 +295,35 @@ class PromptCondAudioDiffusion(nn.Module):
|
||||
self.rsq48towav2vec = torchaudio.transforms.Resample(48000, 16000)
|
||||
# self.wav2vec = Wav2Vec2BertModel.from_pretrained("facebook/w2v-bert-2.0", trust_remote_code=True)
|
||||
# self.wav2vec_processor = AutoFeatureExtractor.from_pretrained("facebook/w2v-bert-2.0", trust_remote_code=True)
|
||||
self.bestrq = load_model(
|
||||
model_dir='codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq',
|
||||
checkpoint_dir='ckpt/encode-s12k.pt',
|
||||
)
|
||||
# self.bestrq = load_model(
|
||||
# model_dir='codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq',
|
||||
# checkpoint_dir='ckpt/encode-s12k.pt',
|
||||
# )
|
||||
ssl_path=os.path.join(folder_paths.models_dir,"SongGeneration/ckpt/encode-s12k.pt")
|
||||
self.bestrq = MusicFMModel(MusicFMConfig())
|
||||
bestrq_weights = torch.load(ssl_path, map_location='cpu',weights_only=False)["model"]
|
||||
self.bestrq.load_state_dict(bestrq_weights, strict=False)
|
||||
del bestrq_weights
|
||||
self.rsq48tobestrq = torchaudio.transforms.Resample(48000, 24000)
|
||||
self.rsq48tohubert = torchaudio.transforms.Resample(48000, 16000)
|
||||
for v in self.bestrq.parameters():v.requires_grad = False
|
||||
self.rvq_bestrq_emb = ResidualVectorQuantize(input_dim = 1024, n_codebooks = 1, codebook_size = 16_384, codebook_dim = 32, quantizer_dropout = 0.0, stale_tolerance=200)
|
||||
for v in self.rvq_bestrq_emb.parameters():v.requires_grad = False
|
||||
|
||||
|
||||
|
||||
# self.hubert = HubertModelWithFinalProj.from_pretrained(os.path.join(folder_paths.models_dir,"SongGeneration/ckpt/models--lengyue233--content-vec-best/snapshots/c0b9ba13db21beaa4053faae94c102ebe326fd68"))
|
||||
# for v in self.hubert.parameters():v.requires_grad = False
|
||||
self.zero_cond_embedding1 = nn.Parameter(torch.randn(32*32,))
|
||||
# self.xvecmodel = XVECModel()
|
||||
config = GPT2Config(n_positions=1000,n_layer=39,n_head=30,n_embd=1200)
|
||||
unet = GPT2Model(config)
|
||||
mlp = nn.Sequential(
|
||||
nn.Linear(1200, 1024),
|
||||
nn.SiLU(),
|
||||
nn.Linear(1024, 1024),
|
||||
nn.SiLU(),
|
||||
nn.Linear(1024, 768)
|
||||
)
|
||||
# config = GPT2Config(n_positions=1000,n_layer=39,n_head=30,n_embd=1200)
|
||||
# unet = GPT2Model(config)
|
||||
# mlp = nn.Sequential(
|
||||
# nn.Linear(1200, 1024),
|
||||
# nn.SiLU(),
|
||||
# nn.Linear(1024, 1024),
|
||||
# nn.SiLU(),
|
||||
# nn.Linear(1024, 768)
|
||||
# )
|
||||
self.set_from = "random"
|
||||
#self.cfm_wrapper = BASECFM(unet, mlp,self.ssl_layer)
|
||||
self.mask_emb = torch.nn.Embedding(3, 48)
|
||||
@@ -540,7 +546,7 @@ class PromptCondAudioDiffusion(nn.Module):
|
||||
input_audio_0 = self.preprocess_audio(input_audio_0)
|
||||
input_audio_1 = self.preprocess_audio(input_audio_1)
|
||||
|
||||
self.bestrq.eval()
|
||||
#self.bestrq.eval()
|
||||
|
||||
# bestrq_middle,bestrq_last = self.extract_bestrq_embeds(input_audios)
|
||||
# bestrq_middle = bestrq_middle.detach()
|
||||
@@ -577,7 +583,7 @@ class PromptCondAudioDiffusion(nn.Module):
|
||||
input_audio_0 = self.preprocess_audio(input_audio_0)
|
||||
input_audio_1 = self.preprocess_audio(input_audio_1)
|
||||
|
||||
self.bestrq.eval()
|
||||
#self.bestrq.eval()
|
||||
|
||||
# bestrq_middle,bestrq_last = self.extract_bestrq_embeds(input_audios)
|
||||
# bestrq_middle = bestrq_middle.detach()
|
||||
|
||||
@@ -21,9 +21,9 @@ from .libs.rvq.descript_quantize3 import ResidualVectorQuantize
|
||||
|
||||
from .models_gpt.models.gpt2_rope2_time_new_correct_mask_noncasual_reflow import GPT2Model
|
||||
from .models_gpt.models.gpt2_config import GPT2Config
|
||||
|
||||
from .our_MERT_BESTRQ.mert_fairseq.models.musicfm.musicfm_model import MusicFMModel, MusicFMConfig
|
||||
from torch.cuda.amp import autocast
|
||||
from .our_MERT_BESTRQ.test import load_model
|
||||
# from .our_MERT_BESTRQ.test import load_model
|
||||
|
||||
class HubertModelWithFinalProj(HubertModel):
|
||||
def __init__(self, config):
|
||||
@@ -147,41 +147,55 @@ class BASECFM(torch.nn.Module, ABC):
|
||||
mu (torch.Tensor): output of encoder
|
||||
shape: (batch_size, n_channels, mel_timesteps, n_feats)
|
||||
"""
|
||||
t, _, dt = t_span[0], t_span[-1], t_span[1] - t_span[0]
|
||||
#t, _, dt = t_span[0], t_span[-1], t_span[1] - t_span[0]
|
||||
dt = t_span[1:] - t_span[:-1]
|
||||
t = t_span[:-1]
|
||||
B = x.shape[0]
|
||||
|
||||
if guidance_scale > 1.0:
|
||||
def double(z):
|
||||
return torch.cat([z, z], 0) if z is not None else None
|
||||
attention_mask = double(attention_mask)
|
||||
|
||||
x_next = x.clone()
|
||||
noise = x.clone()
|
||||
|
||||
# I am storing this because I can later plot it by putting a debugger here and saving it to a file
|
||||
# Or in future might add like a return_all_steps flag
|
||||
sol = []
|
||||
for i in tqdm(range(len(dt))):
|
||||
ti = t[i]
|
||||
|
||||
for step in tqdm(range(1, len(t_span))):
|
||||
x[:,0:incontext_length,:] = (1 - (1 - self.sigma_min) * t) * noise[:,0:incontext_length,:] + t * incontext_x[:,0:incontext_length,:]
|
||||
if(guidance_scale > 1.0):
|
||||
x_next[:, :incontext_length] = (
|
||||
(1 - (1 - self.sigma_min) * ti) * noise[:, :incontext_length] +
|
||||
ti * incontext_x[:, :incontext_length]
|
||||
)
|
||||
|
||||
model_input = torch.cat([ \
|
||||
torch.cat([latent_mask_input, latent_mask_input], 0), \
|
||||
torch.cat([incontext_x, incontext_x], 0), \
|
||||
torch.cat([torch.zeros_like(mu), mu], 0), \
|
||||
torch.cat([x, x], 0), \
|
||||
], 2)
|
||||
timestep=t.unsqueeze(-1).repeat(2)
|
||||
|
||||
dphi_dt = self.estimator(inputs_embeds=model_input, attention_mask=attention_mask,time_step=timestep).last_hidden_state
|
||||
dphi_dt_uncond, dhpi_dt_cond = dphi_dt.chunk(2,0)
|
||||
dphi_dt = dphi_dt_uncond + guidance_scale * (dhpi_dt_cond - dphi_dt_uncond)
|
||||
else:
|
||||
model_input = torch.cat([latent_mask_input, incontext_x, mu, x], 2)
|
||||
timestep=t.unsqueeze(-1)
|
||||
dphi_dt = self.estimator(inputs_embeds=model_input, attention_mask=attention_mask,time_step=timestep).last_hidden_state
|
||||
if guidance_scale > 1.0:
|
||||
model_input = torch.cat([
|
||||
double(latent_mask_input),
|
||||
double(incontext_x),
|
||||
torch.cat([torch.zeros_like(mu), mu], 0),
|
||||
double(x_next),
|
||||
], dim=2)
|
||||
timestep = ti.expand(2 * B)
|
||||
|
||||
dphi_dt = dphi_dt[: ,:, -x.shape[2]:]
|
||||
x = x + dt * dphi_dt
|
||||
t = t + dt
|
||||
sol.append(x)
|
||||
if step < len(t_span) - 1:
|
||||
dt = t_span[step + 1] - t
|
||||
else:
|
||||
model_input = torch.cat([
|
||||
latent_mask_input, incontext_x, mu, x_next
|
||||
], dim=2)
|
||||
timestep = ti.expand(B)
|
||||
|
||||
v = self.estimator(inputs_embeds=model_input,
|
||||
attention_mask=attention_mask,
|
||||
time_step=timestep).last_hidden_state
|
||||
v = v[..., -x.shape[2]:]
|
||||
|
||||
if guidance_scale > 1.0:
|
||||
v_uncond, v_cond = v.chunk(2, 0)
|
||||
v = v_uncond + guidance_scale * (v_cond - v_uncond)
|
||||
|
||||
x_next = x_next + dt[i] * v
|
||||
|
||||
return x_next
|
||||
|
||||
return sol[-1]
|
||||
|
||||
def projection_loss(self,hidden_proj, bestrq_emb):
|
||||
bsz = hidden_proj.shape[0]
|
||||
@@ -243,6 +257,7 @@ class PromptCondAudioDiffusion(nn.Module):
|
||||
snr_gamma=None,
|
||||
uncondition=True,
|
||||
out_paint=False,
|
||||
ssl_path='ckpt/encode-s12k.pt'
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
@@ -263,10 +278,15 @@ class PromptCondAudioDiffusion(nn.Module):
|
||||
self.rsq48towav2vec = torchaudio.transforms.Resample(48000, 16000)
|
||||
# self.wav2vec = Wav2Vec2BertModel.from_pretrained("facebook/w2v-bert-2.0", trust_remote_code=True)
|
||||
# self.wav2vec_processor = AutoFeatureExtractor.from_pretrained("facebook/w2v-bert-2.0", trust_remote_code=True)
|
||||
self.bestrq = load_model(
|
||||
model_dir='codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq',
|
||||
checkpoint_dir='ckpt/encode-s12k.pt',
|
||||
)
|
||||
# self.bestrq = load_model(
|
||||
# model_dir='codeclm/tokenizer/Flow1dVAE/our_MERT_BESTRQ/mert_fairseq',
|
||||
# checkpoint_dir='ckpt/encode-s12k.pt',
|
||||
# )
|
||||
ssl_path=os.path.join(folder_paths.models_dir,"SongGeneration/ckpt/encode-s12k.pt")
|
||||
self.bestrq = MusicFMModel(MusicFMConfig())
|
||||
bestrq_weights = torch.load(ssl_path, map_location='cpu',weights_only=False)
|
||||
self.bestrq.load_state_dict(bestrq_weights, strict=False)
|
||||
del bestrq_weights
|
||||
self.rsq48tobestrq = torchaudio.transforms.Resample(48000, 24000)
|
||||
self.rsq48tohubert = torchaudio.transforms.Resample(48000, 16000)
|
||||
for v in self.bestrq.parameters():v.requires_grad = False
|
||||
|
||||
+24
-24
@@ -4,14 +4,14 @@ from .model.musicfm_25hz import MusicFM25Hz
|
||||
# import sys, os
|
||||
# sys.path.append(os.path.dirname(os.path.abspath(__file__)))
|
||||
# from model.musicfm_25hz import MusicFM25Hz
|
||||
try:
|
||||
from fairseq.fairseq.dataclass import FairseqDataclass
|
||||
from fairseq.fairseq.models import BaseFairseqModel, register_model
|
||||
from fairseq.fairseq.tasks.fairseq_task import FairseqTask
|
||||
except:
|
||||
from fairseq.dataclass import FairseqDataclass
|
||||
from fairseq.models import BaseFairseqModel, register_model
|
||||
from fairseq.tasks.fairseq_task import FairseqTask
|
||||
# try:
|
||||
# from fairseq.fairseq.dataclass import FairseqDataclass
|
||||
# from fairseq.fairseq.models import BaseFairseqModel, register_model
|
||||
# from fairseq.fairseq.tasks.fairseq_task import FairseqTask
|
||||
# except:
|
||||
# from fairseq.dataclass import FairseqDataclass
|
||||
# from fairseq.models import BaseFairseqModel, register_model
|
||||
# from fairseq.tasks.fairseq_task import FairseqTask
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import List, Tuple, Optional
|
||||
@@ -22,7 +22,7 @@ from logging import getLogger
|
||||
logger = getLogger(__name__)
|
||||
|
||||
@dataclass
|
||||
class MusicFMConfig(FairseqDataclass):
|
||||
class MusicFMConfig:
|
||||
label_rate:int = field(default=25)
|
||||
num_codebooks:int = field(default=1)
|
||||
codebook_dim:int = field(default=16)
|
||||
@@ -45,9 +45,9 @@ class MusicFMConfig(FairseqDataclass):
|
||||
|
||||
SAMPLE_RATE = 24_000
|
||||
|
||||
@register_model("musicfm", dataclass=MusicFMConfig)
|
||||
class MusicFMModel(BaseFairseqModel):
|
||||
def __init__(self, cfg: MusicFMConfig, task_cfg: FairseqTask):
|
||||
# @register_model("musicfm", dataclass=MusicFMConfig)
|
||||
class MusicFMModel(torch.nn.Module):
|
||||
def __init__(self, cfg: MusicFMConfig):
|
||||
super().__init__()
|
||||
self.cfg = cfg
|
||||
self.model = MusicFM25Hz(
|
||||
@@ -92,18 +92,18 @@ class MusicFMModel(BaseFairseqModel):
|
||||
result["hidden_emb"] = hidden_emb
|
||||
return result
|
||||
|
||||
@classmethod
|
||||
def build_model(cls, cfg: MusicFMConfig, task: FairseqTask):
|
||||
"""Build a new model instance."""
|
||||
# @classmethod
|
||||
# def build_model(cls, cfg: MusicFMConfig, task: FairseqTask):
|
||||
# """Build a new model instance."""
|
||||
|
||||
model = MusicFMModel(cfg, task.cfg)
|
||||
import numpy as np
|
||||
s = 0
|
||||
for param in model.parameters():
|
||||
s += np.product(param.size())
|
||||
print('# of parameters: '+str(s/1024.0/1024.0))
|
||||
return model
|
||||
# model = MusicFMModel(cfg, task.cfg)
|
||||
# import numpy as np
|
||||
# s = 0
|
||||
# for param in model.parameters():
|
||||
# s += np.product(param.size())
|
||||
# print('# of parameters: '+str(s/1024.0/1024.0))
|
||||
# return model
|
||||
|
||||
def get_losses(self, result, batch):
|
||||
return result['losses']
|
||||
# def get_losses(self, result, batch):
|
||||
# return result['losses']
|
||||
|
||||
Binary file not shown.
BIN
Binary file not shown.
Binary file not shown.
BIN
Binary file not shown.
Reference in New Issue
Block a user