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():
|
||||
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:
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user