Merge pull request #23 from smthemex/Pr2

init
This commit is contained in:
smthemex
2025-10-23 14:23:35 +08:00
committed by GitHub
9 changed files with 109 additions and 84 deletions
+6 -7
View File
@@ -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
```
@@ -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
@@ -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']