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:
|
def _delete_transformer(self) -> None:
|
||||||
self.flux.transformer = None
|
self.flux.transformer = None
|
||||||
|
if hasattr(self.flux, 'transformer_controlnet'):
|
||||||
|
self.flux.transformer_controlnet = None
|
||||||
gc.collect()
|
gc.collect()
|
||||||
mx.metal.clear_cache()
|
mx.metal.clear_cache()
|
||||||
|
|
||||||
|
|||||||
@ -119,7 +119,7 @@ class Flux1Controlnet(nn.Module):
|
|||||||
dt = config.sigmas[t + 1] - config.sigmas[t]
|
dt = config.sigmas[t + 1] - config.sigmas[t]
|
||||||
latents += noise * dt
|
latents += noise * dt
|
||||||
|
|
||||||
# (Optional) Call subscribes at end of loop
|
# (Optional) Call subscribers at end of loop
|
||||||
Callbacks.in_loop(
|
Callbacks.in_loop(
|
||||||
t=t,
|
t=t,
|
||||||
seed=seed,
|
seed=seed,
|
||||||
@ -142,6 +142,14 @@ class Flux1Controlnet(nn.Module):
|
|||||||
time_steps=time_steps,
|
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
|
# 7. Decode the latent array and return the image
|
||||||
latents = ArrayUtil.unpack_latents(latents=latents, height=config.height, width=config.width)
|
latents = ArrayUtil.unpack_latents(latents=latents, height=config.height, width=config.width)
|
||||||
decoded = self.vae.decode(latents)
|
decoded = self.vae.decode(latents)
|
||||||
|
|||||||
@ -99,7 +99,7 @@ class Flux1(nn.Module):
|
|||||||
dt = config.sigmas[t + 1] - config.sigmas[t]
|
dt = config.sigmas[t + 1] - config.sigmas[t]
|
||||||
latents += noise * dt
|
latents += noise * dt
|
||||||
|
|
||||||
# (Optional) Call subscribes at end of loop
|
# (Optional) Call subscribers at end of loop
|
||||||
Callbacks.in_loop(
|
Callbacks.in_loop(
|
||||||
t=t,
|
t=t,
|
||||||
seed=seed,
|
seed=seed,
|
||||||
@ -122,7 +122,7 @@ class Flux1(nn.Module):
|
|||||||
time_steps=time_steps,
|
time_steps=time_steps,
|
||||||
)
|
)
|
||||||
|
|
||||||
# (Optional) Call subscribes at end of loop
|
# (Optional) Call subscribers at end of loop
|
||||||
Callbacks.after_loop(
|
Callbacks.after_loop(
|
||||||
seed=seed,
|
seed=seed,
|
||||||
prompt=prompt,
|
prompt=prompt,
|
||||||
|
|||||||
@ -1,8 +1,8 @@
|
|||||||
from mflux import Config, Flux1, ModelConfig, StopImageGenerationException
|
from mflux import Config, Flux1, ModelConfig, StopImageGenerationException
|
||||||
from mflux.callbacks.callback_registry import CallbackRegistry
|
from mflux.callbacks.callback_registry import CallbackRegistry
|
||||||
from mflux.callbacks.instances.stepwise_handler import StepwiseHandler
|
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.callbacks.instances.memory_saver import MemorySaver
|
||||||
|
from mflux.ui.cli.parsers import CommandLineParser
|
||||||
|
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
|
|||||||
@ -2,12 +2,14 @@ from mflux import Config, Flux1Controlnet, ModelConfig, StopImageGenerationExcep
|
|||||||
from mflux.callbacks.callback_registry import CallbackRegistry
|
from mflux.callbacks.callback_registry import CallbackRegistry
|
||||||
from mflux.callbacks.instances.canny_saver import CannyImageSaver
|
from mflux.callbacks.instances.canny_saver import CannyImageSaver
|
||||||
from mflux.callbacks.instances.stepwise_handler import StepwiseHandler
|
from mflux.callbacks.instances.stepwise_handler import StepwiseHandler
|
||||||
|
from mflux.callbacks.instances.memory_saver import MemorySaver
|
||||||
from mflux.ui.cli.parsers import CommandLineParser
|
from mflux.ui.cli.parsers import CommandLineParser
|
||||||
|
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
# 0. Parse command line arguments
|
# 0. Parse command line arguments
|
||||||
parser = CommandLineParser(description="Generate an image based on a prompt and a controlnet reference image.") # fmt: off
|
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_model_arguments(require_model_arg=True)
|
||||||
parser.add_lora_arguments()
|
parser.add_lora_arguments()
|
||||||
parser.add_image_generator_arguments(supports_metadata_config=False)
|
parser.add_image_generator_arguments(supports_metadata_config=False)
|
||||||
@ -33,6 +35,13 @@ def main():
|
|||||||
CallbackRegistry.register_in_loop(handler)
|
CallbackRegistry.register_in_loop(handler)
|
||||||
CallbackRegistry.register_interrupt(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:
|
try:
|
||||||
for seed in args.seed:
|
for seed in args.seed:
|
||||||
# 3. Generate an image for each seed value
|
# 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)
|
image.save(path=args.output.format(seed=seed), export_json_metadata=args.metadata)
|
||||||
except StopImageGenerationException as stop_exc:
|
except StopImageGenerationException as stop_exc:
|
||||||
print(stop_exc)
|
print(stop_exc)
|
||||||
|
finally:
|
||||||
|
if memory_saver:
|
||||||
|
print(memory_saver.memory_stats())
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
main()
|
main()
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user