diff --git a/acestep/cpu_offload.py b/acestep/cpu_offload.py new file mode 100644 index 0000000..efe6044 --- /dev/null +++ b/acestep/cpu_offload.py @@ -0,0 +1,41 @@ +import torch +import functools +from typing import Callable, TypeVar + + +class CpuOffloader: + def __init__(self, model, device="cpu"): + self.model = model + self.original_device = device + self.original_dtype = model.dtype + + def __enter__(self): + self.model.to(self.original_device, dtype=self.original_dtype) + return self.model + + def __exit__(self, *args): + self.model.to("cpu") + if torch.cuda.is_available(): + torch.cuda.empty_cache() + torch.cuda.synchronize() + + +T = TypeVar('T') + +def cpu_offload(model_attr: str): + def decorator(func: Callable[..., T]) -> Callable[..., T]: + @functools.wraps(func) + def wrapper(self, *args, **kwargs): + if not self.cpu_offload: + return func(self, *args, **kwargs) + + # Get the device from the class + device = self.device + # Get the model from the class attribute + model = getattr(self, model_attr) + + with CpuOffloader(model, device): + return func(self, *args, **kwargs) + + return wrapper + return decorator diff --git a/acestep/pipeline_ace_step.py b/acestep/pipeline_ace_step.py index 078e64b..62982fc 100644 --- a/acestep/pipeline_ace_step.py +++ b/acestep/pipeline_ace_step.py @@ -44,6 +44,7 @@ from acestep.apg_guidance import ( cfg_double_condition_forward, ) import torchaudio +from .cpu_offload import cpu_offload torch.backends.cudnn.benchmark = False @@ -96,6 +97,7 @@ class ACEStepPipeline: text_encoder_checkpoint_path=None, persistent_storage_path=None, torch_compile=False, + cpu_offload=False, **kwargs, ): if not checkpoint_dir: @@ -122,6 +124,7 @@ class ACEStepPipeline: self.device = device self.loaded = False self.torch_compile = torch_compile + self.cpu_offload = cpu_offload def load_checkpoint(self, checkpoint_dir=None): device = self.device @@ -274,12 +277,20 @@ class ACEStepPipeline: dcae_checkpoint_path=dcae_checkpoint_path, vocoder_checkpoint_path=vocoder_checkpoint_path, ) - self.music_dcae.to(device).eval().to(self.dtype) + # self.music_dcae.to(device).eval().to(self.dtype) + if self.cpu_offload: # might be redundant + self.music_dcae = self.music_dcae.to("cpu").eval().to(self.dtype) + else: + self.music_dcae = self.music_dcae.to(device).eval().to(self.dtype) self.ace_step_transformer = ACEStepTransformer2DModel.from_pretrained( ace_step_checkpoint_path, torch_dtype=self.dtype ) - self.ace_step_transformer.to(device).eval().to(self.dtype) + # self.ace_step_transformer.to(device).eval().to(self.dtype) + if self.cpu_offload: + self.ace_step_transformer = self.ace_step_transformer.to("cpu").eval().to(self.dtype) + else: + self.ace_step_transformer = self.ace_step_transformer.to(device).eval().to(self.dtype) lang_segment = LangSegment() @@ -389,7 +400,11 @@ class ACEStepPipeline: text_encoder_model = UMT5EncoderModel.from_pretrained( text_encoder_checkpoint_path, torch_dtype=self.dtype ).eval() - text_encoder_model = text_encoder_model.to(device).to(self.dtype) + # text_encoder_model = text_encoder_model.to(device).to(self.dtype) + if self.cpu_offload: + text_encoder_model = text_encoder_model.to("cpu").eval().to(self.dtype) + else: + text_encoder_model = text_encoder_model.to(device).eval().to(self.dtype) text_encoder_model.requires_grad_(False) self.text_encoder_model = text_encoder_model self.text_tokenizer = AutoTokenizer.from_pretrained( @@ -403,6 +418,7 @@ class ACEStepPipeline: self.ace_step_transformer = torch.compile(self.ace_step_transformer) self.text_encoder_model = torch.compile(self.text_encoder_model) + @cpu_offload("text_encoder_model") def get_text_embeddings(self, texts, device, text_max_length=256): inputs = self.text_tokenizer( texts, @@ -420,6 +436,7 @@ class ACEStepPipeline: attention_mask = inputs["attention_mask"] return last_hidden_states, attention_mask + @cpu_offload("text_encoder_model") def get_text_embeddings_null( self, texts, device, text_max_length=256, tau=0.01, l_min=8, l_max=10 ): @@ -531,6 +548,7 @@ class ACEStepPipeline: print("tokenize error", e, "for line", line, "major_language", lang) return lyric_token_idx + @cpu_offload("ace_step_transformer") def calc_v( self, zt_src, @@ -819,6 +837,7 @@ class ACEStepPipeline: target_latents = zt_edit if xt_tar is None else xt_tar return target_latents + @cpu_offload("ace_step_transformer") @torch.no_grad() def text2music_diffusion_process( self, @@ -1335,6 +1354,7 @@ class ACEStepPipeline: ) return target_latents + @cpu_offload("music_dcae") def latents2audio( self, latents, @@ -1382,6 +1402,7 @@ class ACEStepPipeline: ) return output_path_wav + @cpu_offload("music_dcae") def infer_latents(self, input_audio_path): if input_audio_path is None: return None