diff --git a/main.py b/main.py index 6d67901..aa11963 100644 --- a/main.py +++ b/main.py @@ -12,9 +12,10 @@ image = flux.generate_image( seed=3, prompt="Luxury food photograph of a birthday cake. In the middle it has three candles shaped like letters spelling the word 'MLX'. It has perfect lighting and a cozy background with big bokeh and shallow depth of field. The mood is a sunset balcony in tuscany. The photo is taken from the side of the cake. The scene is complemented by a warm, inviting light that highlights the textures and colors of the ingredients, giving it an appetizing and elegant look.", config=Config( - num_inference_steps=2, + num_inference_steps=20, width=256, height=256, + guidance=3.5, ) ) diff --git a/src/flux_1_schnell/config/config.py b/src/flux_1_schnell/config/config.py index a9ac241..df3950a 100644 --- a/src/flux_1_schnell/config/config.py +++ b/src/flux_1_schnell/config/config.py @@ -13,18 +13,26 @@ class Config: num_inference_steps: int = 4, width: int = 1024, height: int = 1024, + guidance: float = 4.0, ): if width %16 != 0 or height % 16 != 0: log.warning("Width and height should be multiples of 16. Rounding down.") self.width = 16 * (height // 16) self.height = 16 * (width // 16) - base_sigmas = Config.base_sigmas(num_inference_steps) + self.sigmas = Config.base_sigmas(num_inference_steps) self.num_inference_steps = num_inference_steps - self.time_steps = base_sigmas * self.num_train_steps - self.sigmas = mx.concatenate([base_sigmas, mx.zeros(1)]) + self.guidance = guidance @staticmethod def base_sigmas(num_inference_steps): - sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps) + sigmas = np.linspace(1.0, 0, num_inference_steps+1) sigmas = mx.array(sigmas).astype(mx.float32) return sigmas + + def shift_sigmas(self): + y1 = 0.5 + x1 = 256 + m = (1.15 - y1) / (4096 - x1) + b = y1 - m * x1 + mu = m + b + self.sigmas = mx.exp(mu) / (mx.exp(mu) + (1 / self.sigmas - 1)) diff --git a/src/flux_1_schnell/flux.py b/src/flux_1_schnell/flux.py index ef153b6..053c49b 100644 --- a/src/flux_1_schnell/flux.py +++ b/src/flux_1_schnell/flux.py @@ -19,8 +19,9 @@ from flux_1_schnell.weights.weight_handler import WeightHandler class Flux1: def __init__(self, repo_id: str): - is_dev = "FLUX.1-dev" in repo_id - max_t5_length = 512 if is_dev else 256 + self.is_dev = "FLUX.1-dev" in repo_id + # max_t5_length = 512 if self.is_dev else 256 + max_t5_length = 256 tokenizers = TokenizerHandler.load_from_disk_or_huggingface(repo_id, max_t5_length) self.t5_tokenizer = TokenizerT5(tokenizers.t5, max_length=max_t5_length) self.clip_tokenizer = TokenizerCLIP(tokenizers.clip) @@ -32,6 +33,8 @@ class Flux1: self.clip_text_encoder = CLIPEncoder(weights.clip_encoder) def generate_image(self, seed: int, prompt: str, config: Config = Config()) -> PIL.Image.Image: + if self.is_dev: + config.shift_sigmas() latents = LatentCreator.create(config.height, config.width, seed) t5_tokens = self.t5_tokenizer.tokenize(prompt) diff --git a/src/flux_1_schnell/models/transformer/time_text_embed.py b/src/flux_1_schnell/models/transformer/time_text_embed.py index abfc75d..596d271 100644 --- a/src/flux_1_schnell/models/transformer/time_text_embed.py +++ b/src/flux_1_schnell/models/transformer/time_text_embed.py @@ -15,15 +15,14 @@ class TimeTextEmbed(nn.Module): self.text_embedder = TextEmbedder() self.with_guidance_embed = with_guidance_embed if self.with_guidance_embed: - self.guidance = mx.broadcast_to(4.0, (1,)) self.guidance_embedder = GuidanceEmbedder() self.timestep_embedder = TimestepEmbedder() - def forward(self, time_step: mx.array, pooled_projection: mx.array) -> mx.array: + def forward(self, time_step: mx.array, pooled_projection: mx.array, guidance: mx.array) -> mx.array: time_steps_proj = self._time_proj(time_step) time_steps_emb = self.timestep_embedder.forward(time_steps_proj) if self.with_guidance_embed: - time_steps_emb += self.guidance_embedder.forward(self._time_proj(self.guidance)) + time_steps_emb += self.guidance_embedder.forward(self._time_proj(guidance)) pooled_projections = self.text_embedder.forward(pooled_projection) conditioning = time_steps_emb + pooled_projections return conditioning.astype(Config.precision) diff --git a/src/flux_1_schnell/models/transformer/transformer.py b/src/flux_1_schnell/models/transformer/transformer.py index 114c82c..0074e0d 100644 --- a/src/flux_1_schnell/models/transformer/transformer.py +++ b/src/flux_1_schnell/models/transformer/transformer.py @@ -34,10 +34,11 @@ class Transformer(nn.Module): hidden_states: mx.array, config: Config ) -> mx.array: - time_step = config.time_steps[t] + time_step = config.sigmas[t] * config.num_train_steps time_step = mx.broadcast_to(time_step, (1,)) hidden_states = self.x_embedder(hidden_states) - text_embeddings = self.time_text_embed.forward(time_step, pooled_prompt_embeds) + guidance = mx.broadcast_to(config.guidance, (1,)) + text_embeddings = self.time_text_embed.forward(time_step, pooled_prompt_embeds, guidance) encoder_hidden_states = self.context_embedder(prompt_embeds) txt_ids = Transformer._prepare_text_ids(seq_len = prompt_embeds.shape[1]) img_ids = Transformer._prepare_latent_image_ids(config.width, config.height)