add memory saver to controlnet
This commit is contained in:
parent
50752f58f4
commit
f65286e1fb
@ -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()
|
||||
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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():
|
||||
|
||||
@ -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()
|
||||
|
||||
Loading…
Reference in New Issue
Block a user