This commit is contained in:
Michael Hedman
2025-05-18 20:22:46 +02:00
5 changed files with 35 additions and 23 deletions
+2 -18
View File
@@ -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:
+2 -2
View File
@@ -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
+18
View File
@@ -138,6 +138,21 @@ class ACEStepPipeline:
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
checkpoint_dir_models = None
@@ -1875,6 +1890,9 @@ class ACEStepPipeline:
format=format,
)
# Clean up memory after generation
self.cleanup_memory()
end_time = time.time()
latent2audio_time_cost = end_time - start_time
timecosts = {
+8 -1
View File
@@ -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"
]
},
)
+5 -2
View File
@@ -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,