Light refactoring: Move latent creation to separate class

This commit is contained in:
filipstrand 2024-10-20 11:31:26 +02:00
parent f08c667e1a
commit 2ed0628e58
5 changed files with 54 additions and 35 deletions

View File

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

View File

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

View File

View File

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

View File

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