diff --git a/acestep/models/ace_step_transformer.py b/acestep/models/ace_step_transformer.py index 391c985..3798042 100644 --- a/acestep/models/ace_step_transformer.py +++ b/acestep/models/ace_step_transformer.py @@ -360,10 +360,6 @@ class ACEStepTransformer2DModel( for module in self.children(): fn_recursive_feed_forward(module, chunk_size, dim) - def _set_gradient_checkpointing(self, module, value=False): - if hasattr(module, "gradient_checkpointing"): - module.gradient_checkpointing = value - def forward_lyric_encoder( self, lyric_token_idx: Optional[torch.LongTensor] = None, @@ -456,20 +452,8 @@ class ACEStepTransformer2DModel( if self.training and self.gradient_checkpointing: - def create_custom_forward(module, return_dict=None): - def custom_forward(*inputs): - if return_dict is not None: - return module(*inputs, return_dict=return_dict) - else: - return module(*inputs) - - return custom_forward - - ckpt_kwargs: Dict[str, Any] = ( - {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} - ) hidden_states = torch.utils.checkpoint.checkpoint( - create_custom_forward(block), + block, hidden_states=hidden_states, attention_mask=attention_mask, encoder_hidden_states=encoder_hidden_states, @@ -477,7 +461,7 @@ class ACEStepTransformer2DModel( rotary_freqs_cis=rotary_freqs_cis, rotary_freqs_cis_cross=encoder_rotary_freqs_cis, temb=temb, - **ckpt_kwargs, + use_reentrant=False, ) else: diff --git a/acestep/models/lyrics_utils/lyric_encoder.py b/acestep/models/lyrics_utils/lyric_encoder.py index bd849bd..6fe8f27 100644 --- a/acestep/models/lyrics_utils/lyric_encoder.py +++ b/acestep/models/lyrics_utils/lyric_encoder.py @@ -1030,8 +1030,8 @@ class ConformerEncoder(torch.nn.Module): mask_pad: torch.Tensor, ) -> torch.Tensor: for layer in self.encoders: - xs, chunk_masks, _, _ = ckpt.checkpoint( - layer.__call__, xs, chunk_masks, pos_emb, mask_pad + xs, chunk_masks, _, _ = torch.utils.checkpoint.checkpoint( + layer.__call__, xs, chunk_masks, pos_emb, mask_pad, use_reentrant=False ) return xs diff --git a/acestep/pipeline_ace_step.py b/acestep/pipeline_ace_step.py index 4072ae4..0f6c6ef 100644 --- a/acestep/pipeline_ace_step.py +++ b/acestep/pipeline_ace_step.py @@ -137,6 +137,21 @@ class ACEStepPipeline: self.cpu_offload = cpu_offload self.quantized = quantized self.overlapped_decode = overlapped_decode + + def cleanup_memory(self): + """Clean up GPU and CPU memory to prevent VRAM overflow during multiple generations.""" + # Clear CUDA cache + if torch.cuda.is_available(): + torch.cuda.empty_cache() + + # Log memory usage if in verbose mode + allocated = torch.cuda.memory_allocated() / (1024 ** 3) + reserved = torch.cuda.memory_reserved() / (1024 ** 3) + logger.info(f"GPU Memory: {allocated:.2f}GB allocated, {reserved:.2f}GB reserved") + + # Collect Python garbage + import gc + gc.collect() def load_checkpoint(self, checkpoint_dir=None, export_quantized_weights=False): device = self.device @@ -1874,6 +1889,9 @@ class ACEStepPipeline: save_path=save_path, format=format, ) + + # Clean up memory after generation + self.cleanup_memory() end_time = time.time() latent2audio_time_cost = end_time - start_time diff --git a/setup.py b/setup.py index 89ade60..b58461b 100644 --- a/setup.py +++ b/setup.py @@ -5,7 +5,7 @@ setup( description="ACE Step: A Step Towards Music Generation Foundation Model", long_description=open("README.md", encoding="utf-8").read(), long_description_content_type="text/markdown", - version="0.1.2", + version="0.2.0", packages=find_namespace_packages(), install_requires=open("requirements.txt", encoding="utf-8").read().splitlines(), author="ACE Studio, StepFun AI", @@ -24,4 +24,11 @@ setup( package_data={ "acestep.models.lyrics_utils": ["vocab.json"], # Specify the relative path to vocab.json }, + extras_require={ + "train": [ + "peft", + "tensorboard", + "tensorboardX" + ] + }, ) diff --git a/trainer.py b/trainer.py index b713d50..a3ed9cb 100644 --- a/trainer.py +++ b/trainer.py @@ -49,7 +49,7 @@ class Pipeline(LightningModule): ssl_coeff: float = 1.0, checkpoint_dir=None, max_steps: int = 200000, - warmup_steps: int = 4000, + warmup_steps: int = 10, dataset_path: str = "./data/your_dataset_path", lora_config_path: str = None, adapter_name: str = "lora_adapter", @@ -68,6 +68,7 @@ class Pipeline(LightningModule): acestep_pipeline.load_checkpoint(acestep_pipeline.checkpoint_dir) transformers = acestep_pipeline.ace_step_transformer.float().cpu() + transformers.enable_gradient_checkpointing() assert lora_config_path is not None, "Please provide a LoRA config path" if lora_config_path is not None: @@ -76,6 +77,7 @@ class Pipeline(LightningModule): except ImportError: raise ImportError("Please install peft library to use LoRA training") with open(lora_config_path, encoding="utf-8") as f: + import json lora_config = json.load(f) lora_config = LoraConfig(**lora_config) transformers.add_adapter(adapter_config=lora_config, adapter_name=adapter_name) @@ -438,7 +440,7 @@ class Pipeline(LightningModule): lr_scheduler = torch.optim.lr_scheduler.LambdaLR( optimizer, lr_lambda, last_epoch=-1 ) - return [optimizer], lr_scheduler + return [optimizer], [{"scheduler": lr_scheduler, "interval": "step"}] def train_dataloader(self): self.train_dataset = Text2MusicDataset( @@ -825,6 +827,7 @@ def main(args): dataset_path=args.dataset_path, checkpoint_dir=args.checkpoint_dir, adapter_name=args.exp_name, + lora_config_path=args.lora_config_path ) checkpoint_callback = ModelCheckpoint( monitor=None,