diff --git a/ace_step_nodes.py b/ace_step_nodes.py index 8d1303e..17e2d47 100644 --- a/ace_step_nodes.py +++ b/ace_step_nodes.py @@ -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},) diff --git a/pipeline_ace_step.py b/pipeline_ace_step.py index d84f605..a062ed2 100644 --- a/pipeline_ace_step.py +++ b/pipeline_ace_step.py @@ -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) diff --git a/pyproject.toml b/pyproject.toml index 4b08fc1..ef4e21e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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"]