From 2ed0628e581d39991cd2362860f25507e3ee31c1 Mon Sep 17 00:00:00 2001 From: filipstrand Date: Sun, 20 Oct 2024 11:31:26 +0200 Subject: [PATCH] Light refactoring: Move latent creation to separate class --- src/mflux/controlnet/flux_controlnet.py | 6 +-- src/mflux/flux/flux.py | 32 ++------------- src/mflux/latent_creator/__init__.py | 0 src/mflux/latent_creator/latent_creator.py | 45 ++++++++++++++++++++++ src/mflux/post_processing/image_util.py | 6 +-- 5 files changed, 54 insertions(+), 35 deletions(-) create mode 100644 src/mflux/latent_creator/__init__.py create mode 100644 src/mflux/latent_creator/latent_creator.py diff --git a/src/mflux/controlnet/flux_controlnet.py b/src/mflux/controlnet/flux_controlnet.py index 5f4687e..d0ab954 100644 --- a/src/mflux/controlnet/flux_controlnet.py +++ b/src/mflux/controlnet/flux_controlnet.py @@ -13,6 +13,7 @@ from mflux.controlnet.controlnet_util import ControlnetUtil from mflux.controlnet.transformer_controlnet import TransformerControlnet from mflux.controlnet.weight_handler_controlnet import WeightHandlerControlnet from mflux.error.exceptions import StopImageGenerationException +from mflux.latent_creator.latent_creator import LatentCreator from mflux.models.text_encoder.clip_encoder.clip_encoder import CLIPEncoder from mflux.models.text_encoder.t5_encoder.t5_encoder import T5Encoder from mflux.models.transformer.transformer import Transformer @@ -141,10 +142,7 @@ class Flux1Controlnet: controlnet_cond = ArrayUtil.pack_latents(controlnet_cond, config.height, config.width) # 1. Create the initial latents - latents = mx.random.normal( - shape=[1, (config.height // 16) * (config.width // 16), 64], - key=mx.random.key(seed) - ) # fmt: off + latents = LatentCreator.create(seed=seed, height=config.height, width=config.width) # 2. Embed the prompt t5_tokens = self.t5_tokenizer.tokenize(prompt) diff --git a/src/mflux/flux/flux.py b/src/mflux/flux/flux.py index fd59a7f..b714d91 100644 --- a/src/mflux/flux/flux.py +++ b/src/mflux/flux/flux.py @@ -1,5 +1,4 @@ from pathlib import Path -from typing import Union import mlx.core as mx from mlx import nn @@ -9,6 +8,7 @@ from mflux.config.config import Config from mflux.config.model_config import ModelConfig from mflux.config.runtime_config import RuntimeConfig from mflux.error.exceptions import StopImageGenerationException +from mflux.latent_creator.latent_creator import LatentCreator from mflux.models.text_encoder.clip_encoder.clip_encoder import CLIPEncoder from mflux.models.text_encoder.t5_encoder.t5_encoder import T5Encoder from mflux.models.transformer.transformer import Transformer @@ -75,27 +75,6 @@ class Flux1: if weights.quantization_level is not None: self._set_model_weights(weights) - def create_initial_latents(self, seed: int, config: Config, init_image_path: Union[str | Path | None] = None): - noise = mx.random.normal( - shape=[1, (config.height // 16) * (config.width // 16), 64], - key=mx.random.key(seed) - ) # fmt: off - if init_image_path is not None: - user_image = ImageUtil.load_image(init_image_path).convert("RGB") - latents = ArrayUtil.pack_latents( - self.vae.encode( - ImageUtil.to_array(ImageUtil.scale_to_dimensions(user_image, config.width, config.height)) - ), - config.width, - config.height, - ) - sigmas_for_init_image_strength = config.sigmas[config.init_time_step] - latents_adjusted = latents * (1.0 - sigmas_for_init_image_strength) - noise_adjusted = noise * sigmas_for_init_image_strength - return latents_adjusted + noise_adjusted - else: - return noise - def generate_image( self, seed: int, @@ -105,13 +84,7 @@ class Flux1: ) -> GeneratedImage: # Create a new runtime config based on the model type and input parameters config = RuntimeConfig(config, self.model_config) - - # 1. Create the initial latents: text-to-image and image-to-image - # todo: the nested config.config and confusing Config relations - # should be refactored to bring more clarity - latents = self.create_initial_latents(seed, config, config.config.init_image_path) time_steps = tqdm(range(config.init_time_step, config.num_inference_steps)) - stepwise_handler = StepwiseHandler( flux=self, config=config, @@ -121,6 +94,9 @@ class Flux1: output_dir=stepwise_output_dir, ) + # 1. Create the initial latents + latents = LatentCreator.create_for_txt2img_or_img2img(seed, config, self.vae) + # 2. Embed the prompt t5_tokens = self.t5_tokenizer.tokenize(prompt) clip_tokens = self.clip_tokenizer.tokenize(prompt) diff --git a/src/mflux/latent_creator/__init__.py b/src/mflux/latent_creator/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/mflux/latent_creator/latent_creator.py b/src/mflux/latent_creator/latent_creator.py new file mode 100644 index 0000000..90b4ee8 --- /dev/null +++ b/src/mflux/latent_creator/latent_creator.py @@ -0,0 +1,45 @@ +import mlx.core as mx +from mlx import nn + +from mflux.config.runtime_config import RuntimeConfig +from mflux.post_processing.array_util import ArrayUtil +from mflux.post_processing.image_util import ImageUtil + + +class LatentCreator: + @staticmethod + def create( + seed: int, + height: int, + width: int, + ): + return mx.random.normal( + shape=[1, (height // 16) * (width // 16), 64], + key=mx.random.key(seed) + ) # fmt: off + + @staticmethod + def create_for_txt2img_or_img2img( + seed: int, + runtime_conf: RuntimeConfig, + vae: nn.Module, + ): + noise = LatentCreator.create( + seed=seed, + height=runtime_conf.height, + width=runtime_conf.width, + ) + + if runtime_conf.config.init_image_path is None: + # Text2Image + return noise + else: + # Image2Image + user_image = ImageUtil.load_image(runtime_conf.config.init_image_path).convert("RGB") + scaled_user_image = ImageUtil.scale_to_dimensions(user_image, runtime_conf.width, runtime_conf.height) + encoded = vae.encode(ImageUtil.to_array(scaled_user_image)) + latents = ArrayUtil.pack_latents(encoded, runtime_conf.width, runtime_conf.height) + sigmas_for_init_image_strength = runtime_conf.sigmas[runtime_conf.init_time_step] + latents_adjusted = latents * (1.0 - sigmas_for_init_image_strength) + noise_adjusted = noise * sigmas_for_init_image_strength + return latents_adjusted + noise_adjusted diff --git a/src/mflux/post_processing/image_util.py b/src/mflux/post_processing/image_util.py index 24cbadb..f8c6af4 100644 --- a/src/mflux/post_processing/image_util.py +++ b/src/mflux/post_processing/image_util.py @@ -96,14 +96,14 @@ class ImageUtil: return array @staticmethod - def load_image(path: str) -> PIL.Image.Image: + def load_image(path: str | pathlib.Path) -> PIL.Image.Image: return PIL.Image.open(path) @staticmethod def scale_to_dimensions( image: PIL.Image.Image, - target_width: float, - target_height: float, + target_width: int, + target_height: int, ): # for now, naively resize the user image to the output image size # todo: consider alt strategies: cropping / resize+padding