diff --git a/src/mflux/callbacks/instances/memory_saver.py b/src/mflux/callbacks/instances/memory_saver.py index 7a3512d..537245d 100644 --- a/src/mflux/callbacks/instances/memory_saver.py +++ b/src/mflux/callbacks/instances/memory_saver.py @@ -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"