Merge branch 'main' of https://github.com/ace-step/ACE-Step.git
This commit is contained in:
@@ -360,10 +360,6 @@ class ACEStepTransformer2DModel(
|
|||||||
for module in self.children():
|
for module in self.children():
|
||||||
fn_recursive_feed_forward(module, chunk_size, dim)
|
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(
|
def forward_lyric_encoder(
|
||||||
self,
|
self,
|
||||||
lyric_token_idx: Optional[torch.LongTensor] = None,
|
lyric_token_idx: Optional[torch.LongTensor] = None,
|
||||||
@@ -456,20 +452,8 @@ class ACEStepTransformer2DModel(
|
|||||||
|
|
||||||
if self.training and self.gradient_checkpointing:
|
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(
|
hidden_states = torch.utils.checkpoint.checkpoint(
|
||||||
create_custom_forward(block),
|
block,
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
attention_mask=attention_mask,
|
attention_mask=attention_mask,
|
||||||
encoder_hidden_states=encoder_hidden_states,
|
encoder_hidden_states=encoder_hidden_states,
|
||||||
@@ -477,7 +461,7 @@ class ACEStepTransformer2DModel(
|
|||||||
rotary_freqs_cis=rotary_freqs_cis,
|
rotary_freqs_cis=rotary_freqs_cis,
|
||||||
rotary_freqs_cis_cross=encoder_rotary_freqs_cis,
|
rotary_freqs_cis_cross=encoder_rotary_freqs_cis,
|
||||||
temb=temb,
|
temb=temb,
|
||||||
**ckpt_kwargs,
|
use_reentrant=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -1030,8 +1030,8 @@ class ConformerEncoder(torch.nn.Module):
|
|||||||
mask_pad: torch.Tensor,
|
mask_pad: torch.Tensor,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
for layer in self.encoders:
|
for layer in self.encoders:
|
||||||
xs, chunk_masks, _, _ = ckpt.checkpoint(
|
xs, chunk_masks, _, _ = torch.utils.checkpoint.checkpoint(
|
||||||
layer.__call__, xs, chunk_masks, pos_emb, mask_pad
|
layer.__call__, xs, chunk_masks, pos_emb, mask_pad, use_reentrant=False
|
||||||
)
|
)
|
||||||
return xs
|
return xs
|
||||||
|
|
||||||
|
|||||||
@@ -137,6 +137,21 @@ class ACEStepPipeline:
|
|||||||
self.cpu_offload = cpu_offload
|
self.cpu_offload = cpu_offload
|
||||||
self.quantized = quantized
|
self.quantized = quantized
|
||||||
self.overlapped_decode = overlapped_decode
|
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):
|
def load_checkpoint(self, checkpoint_dir=None, export_quantized_weights=False):
|
||||||
device = self.device
|
device = self.device
|
||||||
@@ -1874,6 +1889,9 @@ class ACEStepPipeline:
|
|||||||
save_path=save_path,
|
save_path=save_path,
|
||||||
format=format,
|
format=format,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Clean up memory after generation
|
||||||
|
self.cleanup_memory()
|
||||||
|
|
||||||
end_time = time.time()
|
end_time = time.time()
|
||||||
latent2audio_time_cost = end_time - start_time
|
latent2audio_time_cost = end_time - start_time
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ setup(
|
|||||||
description="ACE Step: A Step Towards Music Generation Foundation Model",
|
description="ACE Step: A Step Towards Music Generation Foundation Model",
|
||||||
long_description=open("README.md", encoding="utf-8").read(),
|
long_description=open("README.md", encoding="utf-8").read(),
|
||||||
long_description_content_type="text/markdown",
|
long_description_content_type="text/markdown",
|
||||||
version="0.1.2",
|
version="0.2.0",
|
||||||
packages=find_namespace_packages(),
|
packages=find_namespace_packages(),
|
||||||
install_requires=open("requirements.txt", encoding="utf-8").read().splitlines(),
|
install_requires=open("requirements.txt", encoding="utf-8").read().splitlines(),
|
||||||
author="ACE Studio, StepFun AI",
|
author="ACE Studio, StepFun AI",
|
||||||
@@ -24,4 +24,11 @@ setup(
|
|||||||
package_data={
|
package_data={
|
||||||
"acestep.models.lyrics_utils": ["vocab.json"], # Specify the relative path to vocab.json
|
"acestep.models.lyrics_utils": ["vocab.json"], # Specify the relative path to vocab.json
|
||||||
},
|
},
|
||||||
|
extras_require={
|
||||||
|
"train": [
|
||||||
|
"peft",
|
||||||
|
"tensorboard",
|
||||||
|
"tensorboardX"
|
||||||
|
]
|
||||||
|
},
|
||||||
)
|
)
|
||||||
|
|||||||
+5
-2
@@ -49,7 +49,7 @@ class Pipeline(LightningModule):
|
|||||||
ssl_coeff: float = 1.0,
|
ssl_coeff: float = 1.0,
|
||||||
checkpoint_dir=None,
|
checkpoint_dir=None,
|
||||||
max_steps: int = 200000,
|
max_steps: int = 200000,
|
||||||
warmup_steps: int = 4000,
|
warmup_steps: int = 10,
|
||||||
dataset_path: str = "./data/your_dataset_path",
|
dataset_path: str = "./data/your_dataset_path",
|
||||||
lora_config_path: str = None,
|
lora_config_path: str = None,
|
||||||
adapter_name: str = "lora_adapter",
|
adapter_name: str = "lora_adapter",
|
||||||
@@ -68,6 +68,7 @@ class Pipeline(LightningModule):
|
|||||||
acestep_pipeline.load_checkpoint(acestep_pipeline.checkpoint_dir)
|
acestep_pipeline.load_checkpoint(acestep_pipeline.checkpoint_dir)
|
||||||
|
|
||||||
transformers = acestep_pipeline.ace_step_transformer.float().cpu()
|
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"
|
assert lora_config_path is not None, "Please provide a LoRA config path"
|
||||||
if lora_config_path is not None:
|
if lora_config_path is not None:
|
||||||
@@ -76,6 +77,7 @@ class Pipeline(LightningModule):
|
|||||||
except ImportError:
|
except ImportError:
|
||||||
raise ImportError("Please install peft library to use LoRA training")
|
raise ImportError("Please install peft library to use LoRA training")
|
||||||
with open(lora_config_path, encoding="utf-8") as f:
|
with open(lora_config_path, encoding="utf-8") as f:
|
||||||
|
import json
|
||||||
lora_config = json.load(f)
|
lora_config = json.load(f)
|
||||||
lora_config = LoraConfig(**lora_config)
|
lora_config = LoraConfig(**lora_config)
|
||||||
transformers.add_adapter(adapter_config=lora_config, adapter_name=adapter_name)
|
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(
|
lr_scheduler = torch.optim.lr_scheduler.LambdaLR(
|
||||||
optimizer, lr_lambda, last_epoch=-1
|
optimizer, lr_lambda, last_epoch=-1
|
||||||
)
|
)
|
||||||
return [optimizer], lr_scheduler
|
return [optimizer], [{"scheduler": lr_scheduler, "interval": "step"}]
|
||||||
|
|
||||||
def train_dataloader(self):
|
def train_dataloader(self):
|
||||||
self.train_dataset = Text2MusicDataset(
|
self.train_dataset = Text2MusicDataset(
|
||||||
@@ -825,6 +827,7 @@ def main(args):
|
|||||||
dataset_path=args.dataset_path,
|
dataset_path=args.dataset_path,
|
||||||
checkpoint_dir=args.checkpoint_dir,
|
checkpoint_dir=args.checkpoint_dir,
|
||||||
adapter_name=args.exp_name,
|
adapter_name=args.exp_name,
|
||||||
|
lora_config_path=args.lora_config_path
|
||||||
)
|
)
|
||||||
checkpoint_callback = ModelCheckpoint(
|
checkpoint_callback = ModelCheckpoint(
|
||||||
monitor=None,
|
monitor=None,
|
||||||
|
|||||||
Reference in New Issue
Block a user