diff --git a/src/mflux/callbacks/instances/memory_saver.py b/src/mflux/callbacks/instances/memory_saver.py index e8c1a95..2adabd6 100644 --- a/src/mflux/callbacks/instances/memory_saver.py +++ b/src/mflux/callbacks/instances/memory_saver.py @@ -65,6 +65,8 @@ class MemorySaver(BeforeLoopCallback, InLoopCallback, AfterLoopCallback): 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() diff --git a/src/mflux/controlnet/flux_controlnet.py b/src/mflux/controlnet/flux_controlnet.py index 73479e3..27c4230 100644 --- a/src/mflux/controlnet/flux_controlnet.py +++ b/src/mflux/controlnet/flux_controlnet.py @@ -119,7 +119,7 @@ class Flux1Controlnet(nn.Module): dt = config.sigmas[t + 1] - config.sigmas[t] latents += noise * dt - # (Optional) Call subscribes at end of loop + # (Optional) Call subscribers at end of loop Callbacks.in_loop( t=t, seed=seed, @@ -142,6 +142,14 @@ class Flux1Controlnet(nn.Module): time_steps=time_steps, ) + # (Optional) Call subscribers at end of loop + Callbacks.after_loop( + seed=seed, + prompt=prompt, + latents=latents, + config=config, + ) # fmt: off + # 7. Decode the latent array and return the image latents = ArrayUtil.unpack_latents(latents=latents, height=config.height, width=config.width) decoded = self.vae.decode(latents) diff --git a/src/mflux/flux/flux.py b/src/mflux/flux/flux.py index b1426c0..78cfd59 100644 --- a/src/mflux/flux/flux.py +++ b/src/mflux/flux/flux.py @@ -99,7 +99,7 @@ class Flux1(nn.Module): dt = config.sigmas[t + 1] - config.sigmas[t] latents += noise * dt - # (Optional) Call subscribes at end of loop + # (Optional) Call subscribers at end of loop Callbacks.in_loop( t=t, seed=seed, @@ -122,7 +122,7 @@ class Flux1(nn.Module): time_steps=time_steps, ) - # (Optional) Call subscribes at end of loop + # (Optional) Call subscribers at end of loop Callbacks.after_loop( seed=seed, prompt=prompt, diff --git a/src/mflux/generate.py b/src/mflux/generate.py index c1c0589..ca2f209 100644 --- a/src/mflux/generate.py +++ b/src/mflux/generate.py @@ -1,8 +1,8 @@ from mflux import Config, Flux1, ModelConfig, StopImageGenerationException from mflux.callbacks.callback_registry import CallbackRegistry from mflux.callbacks.instances.stepwise_handler import StepwiseHandler -from mflux.ui.cli.parsers import CommandLineParser from mflux.callbacks.instances.memory_saver import MemorySaver +from mflux.ui.cli.parsers import CommandLineParser def main(): diff --git a/src/mflux/generate_controlnet.py b/src/mflux/generate_controlnet.py index bf92231..9df5e83 100644 --- a/src/mflux/generate_controlnet.py +++ b/src/mflux/generate_controlnet.py @@ -2,12 +2,14 @@ from mflux import Config, Flux1Controlnet, ModelConfig, StopImageGenerationExcep from mflux.callbacks.callback_registry import CallbackRegistry from mflux.callbacks.instances.canny_saver import CannyImageSaver from mflux.callbacks.instances.stepwise_handler import StepwiseHandler +from mflux.callbacks.instances.memory_saver import MemorySaver from mflux.ui.cli.parsers import CommandLineParser def main(): # 0. Parse command line arguments parser = CommandLineParser(description="Generate an image based on a prompt and a controlnet reference image.") # fmt: off + parser.add_general_arguments() parser.add_model_arguments(require_model_arg=True) parser.add_lora_arguments() parser.add_image_generator_arguments(supports_metadata_config=False) @@ -33,6 +35,13 @@ def main(): CallbackRegistry.register_in_loop(handler) CallbackRegistry.register_interrupt(handler) + memory_saver = None + if args.low_ram: + memory_saver = MemorySaver(flux) + CallbackRegistry.register_before_loop(memory_saver) + CallbackRegistry.register_in_loop(memory_saver) + CallbackRegistry.register_after_loop(memory_saver) + try: for seed in args.seed: # 3. Generate an image for each seed value @@ -53,7 +62,9 @@ def main(): image.save(path=args.output.format(seed=seed), export_json_metadata=args.metadata) except StopImageGenerationException as stop_exc: print(stop_exc) - + finally: + if memory_saver: + print(memory_saver.memory_stats()) if __name__ == "__main__": main()