Update calls in MemorySaver
This commit is contained in:
parent
393a60fb2b
commit
5f0f39f362
@ -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"
|
||||
|
||||
Loading…
Reference in New Issue
Block a user