cache models
This commit is contained in:
+24
-8
@@ -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
@@ -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
@@ -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"]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user