lora weight

This commit is contained in:
billwuhao
2025-05-26 23:30:40 +08:00
parent cfc4d54f75
commit bcdfd58934
2 changed files with 23 additions and 45 deletions
+22 -44
View File
@@ -8,6 +8,7 @@ import sys
import librosa
from transformers import UMT5EncoderModel, AutoTokenizer
from diffusers.utils.peft_utils import set_weights_and_activate_adapters
current_dir = os.path.dirname(os.path.abspath(__file__))
if current_dir not in sys.path:
@@ -205,10 +206,10 @@ class ACEModelLoader:
CATEGORY = "🎤MW/MW-ACE-Step"
def load(self, dcae_checkpoint, vocoder_checkpoint, ace_step_checkpoint, text_encoder_checkpoint, quantized=False, cpu_offload=False, torch_compile=False):
dcae_checkpoint = os.path.join(model_path, "music_dcae_f8c8")
vocoder_checkpoint = os.path.join(model_path, "music_vocoder")
ace_step_checkpoint = os.path.join(model_path, "ace_step_transformer")
text_encoder_checkpoint = os.path.join(model_path, "umt5-base")
dcae_checkpoint = os.path.join(model_path, dcae_checkpoint)
vocoder_checkpoint = os.path.join(model_path, vocoder_checkpoint)
ace_step_checkpoint = os.path.join(model_path, ace_step_checkpoint)
text_encoder_checkpoint = os.path.join(model_path, text_encoder_checkpoint)
for path in [dcae_checkpoint, vocoder_checkpoint, ace_step_checkpoint, text_encoder_checkpoint]:
if not os.path.exists(path):
@@ -316,6 +317,9 @@ class ACEModelLoader:
class ACELoRALoader:
def __init__(self):
self.lora_weight = None
self.lora_name = None
@classmethod
def INPUT_TYPES(cls):
loras_path = os.path.join(model_path, "loras")
@@ -324,6 +328,7 @@ class ACELoRALoader:
"required": {
"models": ("ACE_MODELS",),
"lora_name": (models, {"default": "ACE-Step-v1-chinese-rap-LoRA"}),
"lora_weight": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
},
}
@@ -332,17 +337,20 @@ class ACELoRALoader:
FUNCTION = "load"
CATEGORY = "🎤MW/MW-ACE-Step"
def load(self, models, lora_name):
def load(self, models, lora_name, lora_weight):
lora_path = os.path.join(model_path, "loras", lora_name)
if not all((lora_name, self.lora_weight)) or self.lora_name != lora_name or self.lora_weight != lora_weight:
models[1].unload_lora()
models[1].load_lora_adapter(
os.path.join(lora_path, "pytorch_lora_weights.safetensors"),
adapter_name="zh_rap_lora",
adapter_name="ace_step_lora",
with_alpha=True,
)
set_weights_and_activate_adapters(models[1], ["ace_step_lora"], [lora_weight])
return (models,)
ap = None
class ACEStepGen:
files = DataSampler().input_params_files
songs = {os.path.basename(file): file for file in files}
@@ -353,7 +361,6 @@ class ACEStepGen:
return {
"required": {
"models": ("ACE_MODELS",),
# "unload_model": ("BOOLEAN", {"default": True}),
},
"optional": {
"prompt": ("STRING", {"forceInput": True}),
@@ -382,7 +389,6 @@ class ACEStepGen:
ref_audio_strength=None,
overlapped_decode=False,
delicious_song="None",
# unload_model=True
):
if delicious_song != "None":
@@ -395,9 +401,7 @@ class ACEStepGen:
assert parameters and prompt and lyrics, "parameters, prompt and lyrics are required"
parameters = ast.literal_eval(parameters)
global ap
if ap is None:
ap = AP(*models, overlapped_decode=overlapped_decode)
ap = AP(*models, overlapped_decode=overlapped_decode)
audio2audio_enable = False
ref_audio_input = None
@@ -419,10 +423,6 @@ class ACEStepGen:
**parameters
)
audio, sr = audio_output[0][0].unsqueeze(0), audio_output[0][1]
# if unload_model:
# ap.cleanup()
# ap = None
return ({"waveform": audio, "sample_rate": sr}, prompt, lyrics)
@@ -443,7 +443,6 @@ class ACEStepRepainting:
"repaint_end": ("INT", {"default": 0, "min": 0, "max": 1000, "step": 1}),
"repaint_variance": ("FLOAT", {"default": 0.01, "min": 0.01, "max": 1.0, "step": 0.01}),
"seed": ("INT", {"default":0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF, "step": 1}),
# "unload_model": ("BOOLEAN", {"default": True}),
"overlapped_decode": ("BOOLEAN", {"default": False}),
},
}
@@ -463,7 +462,6 @@ class ACEStepRepainting:
repaint_end,
repaint_variance,
seed,
# unload_model=True,
negative_prompt: str="",
overlapped_decode=False
):
@@ -479,9 +477,8 @@ class ACEStepRepainting:
parameters = ast.literal_eval(parameters)
parameters["audio_duration"] = audio_duration
global ap
if ap is None:
ap = AP(*models, overlapped_decode=overlapped_decode)
ap = AP(*models, overlapped_decode=overlapped_decode)
audio_output = ap(
prompt=prompt,
@@ -496,10 +493,6 @@ class ACEStepRepainting:
**parameters)
audio, sr = audio_output[0][0].unsqueeze(0), audio_output[0][1]
# if unload_model:
# ap.cleanup()
# ap = None
return ({"waveform": audio, "sample_rate": sr},)
@@ -520,7 +513,6 @@ class ACEStepEdit:
"edit_n_min": ("FLOAT", {"default": 0.6, "min": 0.0, "max": 1.0, "step": 0.01}),
"edit_n_max": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"seed": ("INT", {"default":0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF, "step": 1}),
# "unload_model": ("BOOLEAN", {"default": True}),
"overlapped_decode": ("BOOLEAN", {"default": False}),
},
}
@@ -541,7 +533,6 @@ class ACEStepEdit:
edit_n_min,
edit_n_max,
seed,
# unload_model=True,
overlapped_decode=False
):
if seed!= 0:
@@ -553,9 +544,8 @@ class ACEStepEdit:
audio_duration = librosa.get_duration(filename=src_audio_path)
parameters = ast.literal_eval(parameters)
parameters["audio_duration"] = audio_duration
global ap
if ap is None:
ap = AP(*models, overlapped_decode=overlapped_decode)
ap = AP(*models, overlapped_decode=overlapped_decode)
audio_output = ap(
prompt=prompt,
@@ -570,10 +560,6 @@ class ACEStepEdit:
**parameters)
audio, sr = audio_output[0][0].unsqueeze(0), audio_output[0][1]
# if unload_model:
# ap.cleanup()
# ap = None
return ({"waveform": audio, "sample_rate": sr},)
@@ -592,9 +578,7 @@ class ACEStepExtend:
"parameters": ("STRING", {"forceInput": True}),
"left_extend_length": ("INT", {"default": 0, "min": 0, "max": 1000, "step": 1}),
"right_extend_length": ("INT", {"default": 0, "min": 0, "max": 1000, "step": 1}),
# "repaint_variance": ("FLOAT", {"default": 0.01, "min": 0.01, "max": 1.0, "step": 0.01}),
"seed": ("INT", {"default":0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF, "step": 1}),
# "unload_model": ("BOOLEAN", {"default": True}),
"overlapped_decode": ("BOOLEAN", {"default": False}),
},
}
@@ -613,7 +597,6 @@ class ACEStepExtend:
left_extend_length,
right_extend_length,
seed,
# unload_model=True,
negative_prompt: str="",
overlapped_decode=False
):
@@ -629,9 +612,8 @@ class ACEStepExtend:
parameters = ast.literal_eval(parameters)
parameters["audio_duration"] = audio_duration
global ap
if ap is None:
ap = AP(*models, overlapped_decode=overlapped_decode)
ap = AP(*models, overlapped_decode=overlapped_decode)
audio_output = ap(
prompt=prompt,
@@ -646,10 +628,6 @@ class ACEStepExtend:
**parameters)
audio, sr = audio_output[0][0].unsqueeze(0), audio_output[0][1]
# if unload_model:
# ap.cleanup()
# ap = None
return ({"waveform": audio, "sample_rate": sr},)
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "ace-step"
description = "ACE-Step: A Step Towards Music Generation Foundation Model"
version = "1.1.4"
version = "1.1.5"
license = {file = "LICENSE"}
dependencies = ["diffusers", "librosa", "pypinyin", "pytorch_lightning", "soundfile", "hangul-romanize", "num2words", "spacy"]