cache models

This commit is contained in:
billwuhao
2025-05-08 12:58:23 +08:00
parent 1bf41c3f18
commit f472728f0b
3 changed files with 48 additions and 37 deletions
+24 -8
View File
@@ -266,6 +266,7 @@ class MultiLineLyrics:
def lyricsgen(self, multi_line_prompt: str):
return (multi_line_prompt.strip(),)
ap = None
class ACEStepGen:
@classmethod
@@ -276,7 +277,7 @@ class ACEStepGen:
"prompt": ("STRING", {"forceInput": True}),
"lyrics": ("STRING", {"forceInput": True}),
"parameters": ("STRING", {"forceInput": True}),
# "unload_model": ("BOOLEAN", {"default": False}),
"unload_model": ("BOOLEAN", {"default": True}),
},
}
@@ -288,16 +289,22 @@ class ACEStepGen:
def acestepgen(self, prompt: str, lyrics: str, parameters: str, unload_model=True):
parameters = ast.literal_eval(parameters)
ap = AP(model_path)
global ap
if ap is None:
ap = AP(model_path)
audio_output = ap(prompt=prompt, lyrics=lyrics, task="text2music", **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},)
class ACEStepRepainting:
def __init__(self):
ap = None
@classmethod
def INPUT_TYPES(cls):
@@ -311,7 +318,7 @@ 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": 4294967295, "step": 1}),
# "unload_model": ("BOOLEAN", {"default": False}),
"unload_model": ("BOOLEAN", {"default": True}),
},
}
@@ -331,8 +338,10 @@ class ACEStepRepainting:
parameters = ast.literal_eval(parameters)
parameters["audio_duration"] = audio_duration
global ap
if ap is None:
ap = AP(model_path)
ap = AP(model_path)
audio_output = ap(
prompt=prompt,
lyrics=lyrics,
@@ -349,6 +358,7 @@ class ACEStepRepainting:
ac.cleanup_file(src_audio_path)
if unload_model:
ap.cleanup()
ap = None
return ({"waveform": audio, "sample_rate": sr},)
@@ -368,7 +378,7 @@ 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": 4294967295, "step": 1}),
# "unload_model": ("BOOLEAN", {"default": False}),
"unload_model": ("BOOLEAN", {"default": True}),
},
}
@@ -385,8 +395,10 @@ 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(model_path)
ap = AP(model_path)
audio_output = ap(
prompt=prompt,
lyrics=lyrics,
@@ -404,6 +416,7 @@ class ACEStepEdit:
ac.cleanup_file(src_audio_path)
if unload_model:
ap.cleanup()
ap = None
return ({"waveform": audio, "sample_rate": sr},)
@@ -422,7 +435,7 @@ class ACEStepExtend:
"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": 4294967295, "step": 1}),
# "unload_model": ("BOOLEAN", {"default": False}),
"unload_model": ("BOOLEAN", {"default": True}),
},
}
@@ -442,8 +455,10 @@ class ACEStepExtend:
parameters = ast.literal_eval(parameters)
parameters["audio_duration"] = audio_duration
global ap
if ap is None:
ap = AP(model_path)
ap = AP(model_path)
audio_output = ap(
prompt=prompt,
lyrics=lyrics,
@@ -460,6 +475,7 @@ class ACEStepExtend:
ac.cleanup_file(src_audio_path)
if unload_model:
ap.cleanup()
ap = None
return ({"waveform": audio, "sample_rate": sr},)
+23 -28
View File
@@ -62,7 +62,6 @@ class ACEStepPipeline:
self.dtype = torch.float16
self.device = device
self.loaded = False
import platform
system = platform.system()
@@ -70,13 +69,19 @@ class ACEStepPipeline:
torch_compile = True
self.torch_compile = torch_compile
def load_checkpoint(self, checkpoint_dir=None):
start_time = time.time()
print("Checkpoint not loaded, loading checkpoint...")
self.load_checkpoint()
load_model_cost = time.time() - start_time
print(f"Model loaded in {load_model_cost:.2f} seconds.")
def load_checkpoint(self):
device = self.device
dcae_model_path = os.path.join(checkpoint_dir, "music_dcae_f8c8")
vocoder_model_path = os.path.join(checkpoint_dir, "music_vocoder")
ace_step_model_path = os.path.join(checkpoint_dir, "ace_step_transformer")
text_encoder_model_path = os.path.join(checkpoint_dir, "umt5-base")
dcae_model_path = os.path.join(self.checkpoint_dir, "music_dcae_f8c8")
vocoder_model_path = os.path.join(self.checkpoint_dir, "music_vocoder")
ace_step_model_path = os.path.join(self.checkpoint_dir, "ace_step_transformer")
text_encoder_model_path = os.path.join(self.checkpoint_dir, "umt5-base")
files_exist = (
os.path.exists(os.path.join(dcae_model_path, "config.json")) and
@@ -93,41 +98,41 @@ class ACEStepPipeline:
)
if not files_exist:
logger.info(f"Checkpoint directory {checkpoint_dir} is not complete, downloading from Hugging Face Hub")
logger.info(f"Checkpoint directory {self.checkpoint_dir} is not complete, downloading from Hugging Face Hub")
# download music dcae model
os.makedirs(dcae_model_path, exist_ok=True)
hf_hub_download(repo_id=REPO_ID, subfolder="music_dcae_f8c8",
filename="config.json", local_dir=checkpoint_dir, local_dir_use_symlinks=False)
filename="config.json", local_dir=self.checkpoint_dir, local_dir_use_symlinks=False)
hf_hub_download(repo_id=REPO_ID, subfolder="music_dcae_f8c8",
filename="diffusion_pytorch_model.safetensors", local_dir=checkpoint_dir, local_dir_use_symlinks=False)
filename="diffusion_pytorch_model.safetensors", local_dir=self.checkpoint_dir, local_dir_use_symlinks=False)
# download vocoder model
os.makedirs(vocoder_model_path, exist_ok=True)
hf_hub_download(repo_id=REPO_ID, subfolder="music_vocoder",
filename="config.json", local_dir=checkpoint_dir, local_dir_use_symlinks=False)
filename="config.json", local_dir=self.checkpoint_dir, local_dir_use_symlinks=False)
hf_hub_download(repo_id=REPO_ID, subfolder="music_vocoder",
filename="diffusion_pytorch_model.safetensors", local_dir=checkpoint_dir, local_dir_use_symlinks=False)
filename="diffusion_pytorch_model.safetensors", local_dir=self.checkpoint_dir, local_dir_use_symlinks=False)
# download ace_step transformer model
os.makedirs(ace_step_model_path, exist_ok=True)
hf_hub_download(repo_id=REPO_ID, subfolder="ace_step_transformer",
filename="config.json", local_dir=checkpoint_dir, local_dir_use_symlinks=False)
filename="config.json", local_dir=self.checkpoint_dir, local_dir_use_symlinks=False)
hf_hub_download(repo_id=REPO_ID, subfolder="ace_step_transformer",
filename="diffusion_pytorch_model.safetensors", local_dir=checkpoint_dir, local_dir_use_symlinks=False)
filename="diffusion_pytorch_model.safetensors", local_dir=self.checkpoint_dir, local_dir_use_symlinks=False)
# download text encoder model
os.makedirs(text_encoder_model_path, exist_ok=True)
hf_hub_download(repo_id=REPO_ID, subfolder="umt5-base",
filename="config.json", local_dir=checkpoint_dir, local_dir_use_symlinks=False)
filename="config.json", local_dir=self.checkpoint_dir, local_dir_use_symlinks=False)
hf_hub_download(repo_id=REPO_ID, subfolder="umt5-base",
filename="model.safetensors", local_dir=checkpoint_dir, local_dir_use_symlinks=False)
filename="model.safetensors", local_dir=self.checkpoint_dir, local_dir_use_symlinks=False)
hf_hub_download(repo_id=REPO_ID, subfolder="umt5-base",
filename="special_tokens_map.json", local_dir=checkpoint_dir, local_dir_use_symlinks=False)
filename="special_tokens_map.json", local_dir=self.checkpoint_dir, local_dir_use_symlinks=False)
hf_hub_download(repo_id=REPO_ID, subfolder="umt5-base",
filename="tokenizer_config.json", local_dir=checkpoint_dir, local_dir_use_symlinks=False)
filename="tokenizer_config.json", local_dir=self.checkpoint_dir, local_dir_use_symlinks=False)
hf_hub_download(repo_id=REPO_ID, subfolder="umt5-base",
filename="tokenizer.json", local_dir=checkpoint_dir, local_dir_use_symlinks=False)
filename="tokenizer.json", local_dir=self.checkpoint_dir, local_dir_use_symlinks=False)
logger.info("Models downloaded")
@@ -159,7 +164,6 @@ class ACEStepPipeline:
text_encoder_model.requires_grad_(False)
self.text_encoder_model = text_encoder_model
self.text_tokenizer = AutoTokenizer.from_pretrained(text_encoder_checkpoint_path)
self.loaded = True
# compile
if self.torch_compile:
@@ -177,7 +181,6 @@ class ACEStepPipeline:
self.text_tokenizer = None
gc.collect()
torch.cuda.empty_cache()
self.loaded = False
def get_text_embeddings(self, texts, device, text_max_length=256):
inputs = self.text_tokenizer(texts, return_tensors="pt", padding=True, truncation=True, max_length=text_max_length)
@@ -987,14 +990,6 @@ class ACEStepPipeline:
start_time = time.time()
if not self.loaded:
logger.warning("Checkpoint not loaded, loading checkpoint...")
self.load_checkpoint(self.checkpoint_dir)
load_model_cost = time.time() - start_time
logger.info(f"Model loaded in {load_model_cost:.2f} seconds.")
start_time = time.time()
random_generators, actual_seeds = self.set_seeds(batch_size, manual_seeds)
retake_random_generators, actual_retake_seeds = self.set_seeds(batch_size, retake_seeds)
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "ace-step"
description = "ACE-Step: A Step Towards Music Generation Foundation Model"
version = "1.0.1"
version = "1.0.2"
license = {file = "LICENSE"}
dependencies = ["diffusers", "librosa", "pypinyin", "pytorch_lightning", "soundfile", "hangul-romanize", "num2words", "spacy"]