Merge pull request #143 from ssakar/low_ram_improvements

More low-ram improvements
This commit is contained in:
Filip Strand 2025-03-22 10:59:24 +01:00 committed by GitHub
commit d3c2972fec
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
7 changed files with 11 additions and 10 deletions

View File

@ -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.

View File

@ -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"

View File

@ -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)

View File

@ -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)

View File

@ -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)

View File

@ -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)

View File

@ -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