Update calls in MemorySaver

This commit is contained in:
filipstrand 2025-05-07 22:00:13 +02:00
parent 393a60fb2b
commit 5f0f39f362

View File

@ -18,9 +18,9 @@ class MemorySaver(BeforeLoopCallback, InLoopCallback, AfterLoopCallback):
self.flux = flux
self.keep_transformer = keep_transformer
self.peak_memory: int = 0
mx.metal.set_cache_limit(cache_limit_bytes)
mx.metal.clear_cache()
mx.metal.reset_peak_memory()
mx.set_cache_limit(cache_limit_bytes)
mx.clear_cache()
mx.reset_peak_memory()
def call_before_loop(
self,
@ -31,7 +31,7 @@ class MemorySaver(BeforeLoopCallback, InLoopCallback, AfterLoopCallback):
canny_image: PIL.Image.Image | None = None,
depth_image: PIL.Image.Image | None = None,
) -> None:
self.peak_memory = mx.metal.get_peak_memory()
self.peak_memory = mx.get_peak_memory()
self._delete_encoders()
def call_in_loop(
@ -43,7 +43,7 @@ class MemorySaver(BeforeLoopCallback, InLoopCallback, AfterLoopCallback):
config: RuntimeConfig,
time_steps: tqdm,
) -> None:
self.peak_memory = mx.metal.get_peak_memory()
self.peak_memory = mx.get_peak_memory()
def call_after_loop(
self,
@ -52,7 +52,7 @@ class MemorySaver(BeforeLoopCallback, InLoopCallback, AfterLoopCallback):
latents: mx.array,
config: RuntimeConfig,
) -> None:
self.peak_memory = mx.metal.get_peak_memory()
self.peak_memory = mx.get_peak_memory()
if not self.keep_transformer:
self._delete_transformer()
@ -61,15 +61,15 @@ class MemorySaver(BeforeLoopCallback, InLoopCallback, AfterLoopCallback):
self.flux.clip_text_encoder = None
self.flux.t5_text_encoder = None
gc.collect()
mx.metal.clear_cache()
mx.clear_cache()
def _delete_transformer(self) -> None:
self.flux.transformer = None
if hasattr(self.flux, "transformer_controlnet"):
self.flux.transformer_controlnet = None
gc.collect()
mx.metal.clear_cache()
mx.clear_cache()
def memory_stats(self) -> str:
self.peak_memory = mx.metal.get_peak_memory()
self.peak_memory = mx.get_peak_memory()
return f"Peak MLX memory: {self.peak_memory / 10**9:.2f} GB"