diff --git a/src/mflux/config/config.py b/src/mflux/config/config.py index 19f8974..2285b3a 100644 --- a/src/mflux/config/config.py +++ b/src/mflux/config/config.py @@ -1,4 +1,6 @@ import logging +import time +from pathlib import Path import mlx.core as mx @@ -14,6 +16,9 @@ class Config: width: int = 1024, height: int = 1024, guidance: float = 4.0, + init_image: Path | None = None, + init_image_strength: float | None = None, + seed: float | None = None, ): if width % 16 != 0 or height % 16 != 0: log.warning("Width and height should be multiples of 16. Rounding down.") @@ -21,6 +26,9 @@ class Config: self.height = 16 * (height // 16) self.num_inference_steps = num_inference_steps self.guidance = guidance + self.init_image = init_image + self.init_image_strength = init_image_strength + self.seed = seed or int(time.time()) class ConfigControlnet(Config): diff --git a/src/mflux/config/runtime_config.py b/src/mflux/config/runtime_config.py index 6b770a1..ece21f2 100644 --- a/src/mflux/config/runtime_config.py +++ b/src/mflux/config/runtime_config.py @@ -1,15 +1,25 @@ +import logging +import typing as t + import mlx.core as mx import numpy as np from mflux.config.config import Config, ConfigControlnet from mflux.config.model_config import ModelConfig +from mflux.post_processing.array_util import ArrayUtil +from mflux.post_processing.image_util import ImageUtil + +logger = logging.getLogger(__name__) + +VAE: t.TypeAlias = "mflux.models.vae.vae.VAE" # noqa: F821 class RuntimeConfig: - def __init__(self, config: Config | ConfigControlnet, model_config: ModelConfig): + def __init__(self, config: Config | ConfigControlnet, model_config: ModelConfig, vae: VAE): self.config = config self.model_config = model_config self.sigmas = self._create_sigmas(config, model_config) + self.vae = vae @property def height(self) -> int: @@ -19,6 +29,10 @@ class RuntimeConfig: def width(self) -> int: return self.config.width + @property + def seed(self) -> int: + return self.config.seed + @property def guidance(self) -> float: return self.config.guidance @@ -35,6 +49,40 @@ class RuntimeConfig: def num_train_steps(self) -> int: return self.model_config.num_train_steps + @property + def init_time_step(self) -> int: + if self.config.init_image is None: + # text to image, always begin at time step 0 + return 0 + else: + # we skip to the time step as informed by the init_image_strength + # the higher the strength number, the more time steps we skip + strength = max(0.0, min(1.0, self.config.init_image_strength)) + # if the strength is too small to even influence the image + # help the user round up so the init_image has influence at step 1 + t = max(1, int(self.num_inference_steps * strength)) + return t + + @property + def init_latents(self) -> mx.array: + noise = mx.random.normal( + shape=[1, (self.height // 16) * (self.width // 16), 64], + key=mx.random.key(self.seed) + ) # fmt: off + if self.config.init_image is not None: + user_image = ImageUtil.load_image(self.config.init_image).convert("RGB") + latents = ArrayUtil.pack_latents( + self.vae.encode(ImageUtil.to_array(ImageUtil.scale_to_dimensions(user_image, self.width, self.height))), + self.width, + self.height, + ) + sigmas_for_init_image_strength = self.sigmas[self.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 + @property def controlnet_strength(self) -> float: if isinstance(self.config, ConfigControlnet): diff --git a/src/mflux/controlnet/controlnet_util.py b/src/mflux/controlnet/controlnet_util.py index 64cb27c..e0dffd7 100644 --- a/src/mflux/controlnet/controlnet_util.py +++ b/src/mflux/controlnet/controlnet_util.py @@ -17,13 +17,6 @@ class ControlnetUtil: image_to_canny = np.concatenate([image_to_canny, image_to_canny, image_to_canny], axis=2) return PIL.Image.fromarray(image_to_canny) - @staticmethod - def scale_image(height: int, width: int, img: PIL.Image) -> PIL.Image: - if height != img.height or width != img.width: - log.warning(f"Control image has different dimensions than the model. Resizing to {width}x{height}") - img = img.resize((width, height), PIL.Image.LANCZOS) - return img - @staticmethod def save_canny_image(control_image: PIL.Image, path: str): from mflux import ImageUtil diff --git a/src/mflux/controlnet/flux_controlnet.py b/src/mflux/controlnet/flux_controlnet.py index 5f4687e..604a93f 100644 --- a/src/mflux/controlnet/flux_controlnet.py +++ b/src/mflux/controlnet/flux_controlnet.py @@ -131,7 +131,7 @@ class Flux1Controlnet: # Embed the controlnet reference image control_image = ImageUtil.load_image(controlnet_image_path) - control_image = ControlnetUtil.scale_image(config.height, config.width, control_image) + control_image = ImageUtil.scale_to_dimensions(control_image, config.height, config.width) control_image = ControlnetUtil.preprocess_canny(control_image) if controlnet_save_canny: ControlnetUtil.save_canny_image(control_image, output) diff --git a/src/mflux/flux/flux.py b/src/mflux/flux/flux.py index be0b367..e2a8a52 100644 --- a/src/mflux/flux/flux.py +++ b/src/mflux/flux/flux.py @@ -82,8 +82,12 @@ class Flux1: stepwise_output_dir: Path = None, ) -> GeneratedImage: # Create a new runtime config based on the model type and input parameters - config = RuntimeConfig(config, self.model_config) - time_steps = tqdm(range(config.num_inference_steps)) + config = RuntimeConfig(config, self.model_config, self.vae) + + # 1. Create the initial latents: text-to-image and image-to-image + latents = config.init_latents + time_steps = tqdm(range(config.init_time_step, config.num_inference_steps)) + stepwise_handler = StepwiseHandler( flux=self, config=config, @@ -93,19 +97,13 @@ class Flux1: output_dir=stepwise_output_dir, ) - # 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 - # 2. Embed the prompt t5_tokens = self.t5_tokenizer.tokenize(prompt) clip_tokens = self.clip_tokenizer.tokenize(prompt) prompt_embeds = self.t5_text_encoder.forward(t5_tokens) pooled_prompt_embeds = self.clip_text_encoder.forward(clip_tokens) - for t in time_steps: + for gen_step, t in enumerate(time_steps, 1): try: # 3.t Predict the noise noise = self.transformer.predict( @@ -121,7 +119,7 @@ class Flux1: latents += noise * dt # Handle stepwise output if enabled - stepwise_handler.process_step(t, latents) + stepwise_handler.process_step(gen_step, latents) # Evaluate to enable progress tracking mx.eval(latents) diff --git a/src/mflux/generate.py b/src/mflux/generate.py index b89b17b..1a0457c 100644 --- a/src/mflux/generate.py +++ b/src/mflux/generate.py @@ -11,6 +11,7 @@ def main(): parser.add_model_arguments() parser.add_lora_arguments() parser.add_image_generator_arguments() + parser.add_image_to_image_arguments(required=False) parser.add_output_arguments() args = parser.parse_args() @@ -34,6 +35,9 @@ def main(): height=args.height, width=args.width, guidance=args.guidance, + init_image=args.init_image, + init_image_strength=args.init_image_strength, + seed=args.seed ), ) diff --git a/src/mflux/post_processing/image_util.py b/src/mflux/post_processing/image_util.py index 61a1d40..24cbadb 100644 --- a/src/mflux/post_processing/image_util.py +++ b/src/mflux/post_processing/image_util.py @@ -6,16 +6,15 @@ import typing as t import mlx.core as mx import numpy as np import piexif -import PIL import PIL.Image -from PIL import Image from mflux.config.config import ConfigControlnet -from mflux.config.runtime_config import RuntimeConfig from mflux.post_processing.generated_image import GeneratedImage log = logging.getLogger(__name__) +RuntimeConfig: t.TypeAlias = "mflux.config.runtime_config.RuntimeConfig" # noqa: F821 + class ImageUtil: @staticmethod @@ -50,11 +49,11 @@ class ImageUtil: ) @staticmethod - def to_composite_image(generated_images: t.List[GeneratedImage]) -> Image: + def to_composite_image(generated_images: t.List[GeneratedImage]) -> PIL.Image.Image: # stitch horizontally total_width = sum(gen_img.image.width for gen_img in generated_images) max_height = max(gen_img.image.height for gen_img in generated_images) - composite_img = Image.new("RGB", (total_width, max_height)) + composite_img = PIL.Image.new("RGB", (total_width, max_height)) current_x = 0 for index, gen_img in enumerate(generated_images): composite_img.paste(gen_img.image, (current_x, 0)) @@ -97,8 +96,21 @@ class ImageUtil: return array @staticmethod - def load_image(path: str) -> Image.Image: - return Image.open(path) + def load_image(path: str) -> PIL.Image.Image: + return PIL.Image.open(path) + + @staticmethod + def scale_to_dimensions( + image: PIL.Image.Image, + target_width: float, + target_height: float, + ): + # for now, naively resize the user image to the output image size + # todo: consider alt strategies: cropping / resize+padding + if (image.width, image.height) != (target_width, target_height): + return image.resize((target_width, target_height), PIL.Image.LANCZOS) + else: + return image @staticmethod def save_image( diff --git a/src/mflux/post_processing/stepwise_handler.py b/src/mflux/post_processing/stepwise_handler.py index d3e9cd6..ee604f6 100644 --- a/src/mflux/post_processing/stepwise_handler.py +++ b/src/mflux/post_processing/stepwise_handler.py @@ -1,6 +1,7 @@ from pathlib import Path import mlx.core as mx +import tqdm from mflux.config.runtime_config import RuntimeConfig from mflux.post_processing.array_util import ArrayUtil @@ -14,7 +15,7 @@ class StepwiseHandler: config: RuntimeConfig, seed: int, prompt: str, - time_steps, + time_steps: tqdm.std.tqdm, output_dir: Path | None = None, ): self.flux = flux @@ -28,7 +29,12 @@ class StepwiseHandler: if self.output_dir: self.output_dir.mkdir(parents=True, exist_ok=True) - def process_step(self, step: int, latents: mx.array): + def save_composite(self): + if self.step_wise_images: + composite_img = ImageUtil.to_composite_image(self.step_wise_images) + composite_img.save(self.output_dir / f"seed_{self.seed}_composite.png") + + def process_step(self, gen_step: int, latents: mx.array): if self.output_dir: unpack_latents = ArrayUtil.unpack_latents(latents, self.config.height, self.config.width) stepwise_decoded = self.flux.vae.decode(unpack_latents) @@ -45,11 +51,10 @@ class StepwiseHandler: self.step_wise_images.append(stepwise_img) stepwise_img.save( - path=self.output_dir / f"seed_{self.seed}_step{step + 1}of{len(self.time_steps)}.png", + path=self.output_dir / f"seed_{self.seed}_step{gen_step}of{len(self.time_steps)}.png", export_json_metadata=False, ) + self.save_composite() def handle_interruption(self): - if self.step_wise_images: - composite_img = ImageUtil.to_composite_image(self.step_wise_images) - composite_img.save(self.output_dir / f"seed_{self.seed}_composite.png") + self.save_composite() diff --git a/src/mflux/ui/cli/parsers.py b/src/mflux/ui/cli/parsers.py index 7b51b7b..0814c8e 100644 --- a/src/mflux/ui/cli/parsers.py +++ b/src/mflux/ui/cli/parsers.py @@ -27,6 +27,10 @@ class CommandLineParser(argparse.ArgumentParser): self.add_argument("--seed", type=int, default=None, help="Entropy Seed (Default is time-based random-seed)") self._add_image_generator_common_arguments() + def add_image_to_image_arguments(self, required=False) -> None: + self.add_argument("--init-image", type=Path, required=required, help="Local path to init image") + self.add_argument("--init-image-strength", type=float, required=False, default=ui_defaults.INIT_IMAGE_STRENGTH, help=f"Controls how strongly the init image influences the output image. A value of 0.0 means no influence. (Default is {ui_defaults.INIT_IMAGE_STRENGTH})") + def add_batch_image_generator_arguments(self) -> None: self.add_argument("--prompts-file", type=Path, required=True, help="Local path for a file that holds a batch of prompts.") self.add_argument("--global-seed", type=int, default=None, help="Entropy Seed (used for all prompts in the batch)") diff --git a/src/mflux/ui/defaults.py b/src/mflux/ui/defaults.py index fec8d7d..7ed4bd7 100644 --- a/src/mflux/ui/defaults.py +++ b/src/mflux/ui/defaults.py @@ -1,6 +1,7 @@ CONTROLNET_STRENGTH = 0.4 GUIDANCE_SCALE = 3.5 HEIGHT, WIDTH = 1024, 1024 +INIT_IMAGE_STRENGTH = 0.4 # for image-to-image init_image MODEL_CHOICES = ["dev", "schnell"] MODEL_INFERENCE_STEPS = { "dev": 14,