From 4e2a10b1b9eedd61d8b9dfa06c56e6ac68f49b8e Mon Sep 17 00:00:00 2001 From: filipstrand Date: Tue, 29 Oct 2024 19:19:59 +0100 Subject: [PATCH] Rewrite img2img latent creation to emphasize the linear interp between latent and noise --- src/mflux/latent_creator/latent_creator.py | 22 ++++++++++++++-------- 1 file changed, 14 insertions(+), 8 deletions(-) diff --git a/src/mflux/latent_creator/latent_creator.py b/src/mflux/latent_creator/latent_creator.py index 90b4ee8..70332c9 100644 --- a/src/mflux/latent_creator/latent_creator.py +++ b/src/mflux/latent_creator/latent_creator.py @@ -12,7 +12,7 @@ class LatentCreator: seed: int, height: int, width: int, - ): + ) -> mx.array: return mx.random.normal( shape=[1, (height // 16) * (width // 16), 64], key=mx.random.key(seed) @@ -23,8 +23,8 @@ class LatentCreator: seed: int, runtime_conf: RuntimeConfig, vae: nn.Module, - ): - noise = LatentCreator.create( + ) -> mx.array: + pure_noise = LatentCreator.create( seed=seed, height=runtime_conf.height, width=runtime_conf.width, @@ -32,14 +32,20 @@ class LatentCreator: if runtime_conf.config.init_image_path is None: # Text2Image - return noise + return pure_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 + sigma = runtime_conf.sigmas[runtime_conf.init_time_step] + return LatentCreator.add_noise_by_interpolation( + clean=latents, + noise=pure_noise, + sigma=sigma + ) # fmt: off + + @staticmethod + def add_noise_by_interpolation(clean: mx.array, noise: mx.array, sigma: float) -> mx.array: + return (1 - sigma) * clean + sigma * noise