add memory saver to controlnet

This commit is contained in:
Serkan Sakar 2025-03-02 19:14:23 +01:00
parent 50752f58f4
commit f65286e1fb
5 changed files with 26 additions and 5 deletions

View File

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

View File

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

View File

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

View File

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

View File

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