diff --git a/README.md b/README.md index 613aad7..aff15fb 100644 --- a/README.md +++ b/README.md @@ -196,7 +196,7 @@ mflux-generate --model dev --prompt "Luxury food photograph" --steps 25 --seed 2 - **`--config-from-metadata`** or **`-C`** (optional, `str`): [EXPERIMENTAL] Path to a prior file saved via `--metadata`, or a compatible handcrafted config file adhering to the expected args schema. -- **`--low-ram`** (optional): Reduces GPU memory usage by constraining the MLX cache size and releasing text encoders and transformer components after use. This option is only compatible with single image generation. While it may slightly decrease performance, it helps prevent system memory swapping to disk, allowing generation on systems with limited RAM. +- **`--low-ram`** (optional): Reduces GPU memory usage by limiting the MLX cache size and releasing text encoders and transformer components after use (single image generation only). While this may slightly decrease performance, it helps prevent system memory swapping to disk, allowing image generation on systems with limited RAM. - **`--lora-name`** (optional, `str`, default: `None`): The name of the LoRA to download from Hugging Face. diff --git a/src/mflux/callbacks/instances/memory_saver.py b/src/mflux/callbacks/instances/memory_saver.py index bd928b9..932c369 100644 --- a/src/mflux/callbacks/instances/memory_saver.py +++ b/src/mflux/callbacks/instances/memory_saver.py @@ -14,8 +14,9 @@ class MemorySaver(BeforeLoopCallback, InLoopCallback, AfterLoopCallback): components at strategic points in the execution cycle. """ - def __init__(self, flux, cache_limit_bytes: int = 1000**3): + def __init__(self, flux, keep_transformer: bool = True, cache_limit_bytes: int = 1000**3): 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() @@ -51,9 +52,11 @@ class MemorySaver(BeforeLoopCallback, InLoopCallback, AfterLoopCallback): config: RuntimeConfig, ) -> None: self.peak_memory = mx.metal.get_peak_memory() - self._delete_transformer() + if not self.keep_transformer: + self._delete_transformer() def _delete_encoders(self) -> None: + # repeated image generation only works with the same prompt (cache) self.flux.clip_text_encoder = None self.flux.t5_text_encoder = None gc.collect() @@ -67,4 +70,5 @@ class MemorySaver(BeforeLoopCallback, InLoopCallback, AfterLoopCallback): mx.metal.clear_cache() def memory_stats(self) -> str: + self.peak_memory = mx.metal.get_peak_memory() return f"Peak MLX memory: {self.peak_memory / 10**9:.2f} GB" diff --git a/src/mflux/generate.py b/src/mflux/generate.py index 49508da..dbd7765 100644 --- a/src/mflux/generate.py +++ b/src/mflux/generate.py @@ -34,7 +34,7 @@ def main(): memory_saver = None if args.low_ram: - memory_saver = MemorySaver(flux) + memory_saver = MemorySaver(flux=flux, keep_transformer=len(args.seed) > 1) CallbackRegistry.register_before_loop(memory_saver) CallbackRegistry.register_in_loop(memory_saver) CallbackRegistry.register_after_loop(memory_saver) diff --git a/src/mflux/generate_controlnet.py b/src/mflux/generate_controlnet.py index 90efa40..41dfd79 100644 --- a/src/mflux/generate_controlnet.py +++ b/src/mflux/generate_controlnet.py @@ -37,7 +37,7 @@ def main(): memory_saver = None if args.low_ram: - memory_saver = MemorySaver(flux) + memory_saver = MemorySaver(flux=flux, keep_transformer=len(args.seed) > 1) CallbackRegistry.register_before_loop(memory_saver) CallbackRegistry.register_in_loop(memory_saver) CallbackRegistry.register_after_loop(memory_saver) diff --git a/src/mflux/generate_fill.py b/src/mflux/generate_fill.py index b4e86f2..fda51d9 100644 --- a/src/mflux/generate_fill.py +++ b/src/mflux/generate_fill.py @@ -39,7 +39,7 @@ def main(): memory_saver = None if args.low_ram: - memory_saver = MemorySaver(flux) + memory_saver = MemorySaver(flux, keep_transformer=len(args.seed) > 1) CallbackRegistry.register_before_loop(memory_saver) CallbackRegistry.register_in_loop(memory_saver) CallbackRegistry.register_after_loop(memory_saver) diff --git a/src/mflux/generate_in_context.py b/src/mflux/generate_in_context.py index 4a79894..02482e6 100644 --- a/src/mflux/generate_in_context.py +++ b/src/mflux/generate_in_context.py @@ -38,7 +38,7 @@ def main(): memory_saver = None if args.low_ram: - memory_saver = MemorySaver(flux) + memory_saver = MemorySaver(flux, keep_transformer=len(args.seed) > 1) CallbackRegistry.register_before_loop(memory_saver) CallbackRegistry.register_in_loop(memory_saver) CallbackRegistry.register_after_loop(memory_saver) diff --git a/src/mflux/ui/cli/parsers.py b/src/mflux/ui/cli/parsers.py index 39e2e18..f38dbae 100644 --- a/src/mflux/ui/cli/parsers.py +++ b/src/mflux/ui/cli/parsers.py @@ -220,7 +220,4 @@ class CommandLineParser(argparse.ArgumentParser): namespace.image_outpaint_padding = box_values.parse_box_value(namespace.image_outpaint_padding) print(f"{namespace.image_outpaint_padding=}") - if getattr(namespace, 'low_ram', False) and len(namespace.seed) > 1: - self.error("--low-ram cannot be used with multiple seeds") - return namespace