More low-ram improvements
This commit is contained in:
parent
eab6091ead
commit
d225d1bbb9
@ -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.
|
||||
|
||||
|
||||
@ -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"
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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
|
||||
|
||||
Loading…
Reference in New Issue
Block a user