Light refactoring: Move latent creation to separate class
This commit is contained in:
parent
f08c667e1a
commit
2ed0628e58
@ -13,6 +13,7 @@ from mflux.controlnet.controlnet_util import ControlnetUtil
|
|||||||
from mflux.controlnet.transformer_controlnet import TransformerControlnet
|
from mflux.controlnet.transformer_controlnet import TransformerControlnet
|
||||||
from mflux.controlnet.weight_handler_controlnet import WeightHandlerControlnet
|
from mflux.controlnet.weight_handler_controlnet import WeightHandlerControlnet
|
||||||
from mflux.error.exceptions import StopImageGenerationException
|
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.clip_encoder.clip_encoder import CLIPEncoder
|
||||||
from mflux.models.text_encoder.t5_encoder.t5_encoder import T5Encoder
|
from mflux.models.text_encoder.t5_encoder.t5_encoder import T5Encoder
|
||||||
from mflux.models.transformer.transformer import Transformer
|
from mflux.models.transformer.transformer import Transformer
|
||||||
@ -141,10 +142,7 @@ class Flux1Controlnet:
|
|||||||
controlnet_cond = ArrayUtil.pack_latents(controlnet_cond, config.height, config.width)
|
controlnet_cond = ArrayUtil.pack_latents(controlnet_cond, config.height, config.width)
|
||||||
|
|
||||||
# 1. Create the initial latents
|
# 1. Create the initial latents
|
||||||
latents = mx.random.normal(
|
latents = LatentCreator.create(seed=seed, height=config.height, width=config.width)
|
||||||
shape=[1, (config.height // 16) * (config.width // 16), 64],
|
|
||||||
key=mx.random.key(seed)
|
|
||||||
) # fmt: off
|
|
||||||
|
|
||||||
# 2. Embed the prompt
|
# 2. Embed the prompt
|
||||||
t5_tokens = self.t5_tokenizer.tokenize(prompt)
|
t5_tokens = self.t5_tokenizer.tokenize(prompt)
|
||||||
|
|||||||
@ -1,5 +1,4 @@
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Union
|
|
||||||
|
|
||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
from mlx import nn
|
from mlx import nn
|
||||||
@ -9,6 +8,7 @@ from mflux.config.config import Config
|
|||||||
from mflux.config.model_config import ModelConfig
|
from mflux.config.model_config import ModelConfig
|
||||||
from mflux.config.runtime_config import RuntimeConfig
|
from mflux.config.runtime_config import RuntimeConfig
|
||||||
from mflux.error.exceptions import StopImageGenerationException
|
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.clip_encoder.clip_encoder import CLIPEncoder
|
||||||
from mflux.models.text_encoder.t5_encoder.t5_encoder import T5Encoder
|
from mflux.models.text_encoder.t5_encoder.t5_encoder import T5Encoder
|
||||||
from mflux.models.transformer.transformer import Transformer
|
from mflux.models.transformer.transformer import Transformer
|
||||||
@ -75,27 +75,6 @@ class Flux1:
|
|||||||
if weights.quantization_level is not None:
|
if weights.quantization_level is not None:
|
||||||
self._set_model_weights(weights)
|
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(
|
def generate_image(
|
||||||
self,
|
self,
|
||||||
seed: int,
|
seed: int,
|
||||||
@ -105,13 +84,7 @@ class Flux1:
|
|||||||
) -> GeneratedImage:
|
) -> GeneratedImage:
|
||||||
# Create a new runtime config based on the model type and input parameters
|
# Create a new runtime config based on the model type and input parameters
|
||||||
config = RuntimeConfig(config, self.model_config)
|
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))
|
time_steps = tqdm(range(config.init_time_step, config.num_inference_steps))
|
||||||
|
|
||||||
stepwise_handler = StepwiseHandler(
|
stepwise_handler = StepwiseHandler(
|
||||||
flux=self,
|
flux=self,
|
||||||
config=config,
|
config=config,
|
||||||
@ -121,6 +94,9 @@ class Flux1:
|
|||||||
output_dir=stepwise_output_dir,
|
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
|
# 2. Embed the prompt
|
||||||
t5_tokens = self.t5_tokenizer.tokenize(prompt)
|
t5_tokens = self.t5_tokenizer.tokenize(prompt)
|
||||||
clip_tokens = self.clip_tokenizer.tokenize(prompt)
|
clip_tokens = self.clip_tokenizer.tokenize(prompt)
|
||||||
|
|||||||
0
src/mflux/latent_creator/__init__.py
Normal file
0
src/mflux/latent_creator/__init__.py
Normal file
45
src/mflux/latent_creator/latent_creator.py
Normal file
45
src/mflux/latent_creator/latent_creator.py
Normal 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
|
||||||
@ -96,14 +96,14 @@ class ImageUtil:
|
|||||||
return array
|
return array
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def load_image(path: str) -> PIL.Image.Image:
|
def load_image(path: str | pathlib.Path) -> PIL.Image.Image:
|
||||||
return PIL.Image.open(path)
|
return PIL.Image.open(path)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def scale_to_dimensions(
|
def scale_to_dimensions(
|
||||||
image: PIL.Image.Image,
|
image: PIL.Image.Image,
|
||||||
target_width: float,
|
target_width: int,
|
||||||
target_height: float,
|
target_height: int,
|
||||||
):
|
):
|
||||||
# for now, naively resize the user image to the output image size
|
# for now, naively resize the user image to the output image size
|
||||||
# todo: consider alt strategies: cropping / resize+padding
|
# todo: consider alt strategies: cropping / resize+padding
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user