From 7b1eba235d849f3121b6203582a46959e88f4a88 Mon Sep 17 00:00:00 2001 From: Fabio Date: Sat, 17 Aug 2024 22:33:11 +0200 Subject: [PATCH 01/20] WIP include flux1 dev --- .gitignore | 1 + main.py | 4 ++-- src/flux_1_schnell/flux.py | 13 ++++--------- src/flux_1_schnell/scheduler/__init__.py | 0 src/flux_1_schnell/scheduler/scheduler.py | 20 -------------------- 5 files changed, 7 insertions(+), 31 deletions(-) delete mode 100644 src/flux_1_schnell/scheduler/__init__.py delete mode 100644 src/flux_1_schnell/scheduler/scheduler.py diff --git a/.gitignore b/.gitignore index fd92c93..ee9ae73 100644 --- a/.gitignore +++ b/.gitignore @@ -10,3 +10,4 @@ .venv *.png *.jpg +*.pyc diff --git a/main.py b/main.py index 60273f3..c5a0469 100644 --- a/main.py +++ b/main.py @@ -4,9 +4,9 @@ import sys sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), 'src'))) from flux_1_schnell.config.config import Config -from flux_1_schnell.flux import Flux1Schnell +from flux_1_schnell.flux import Flux1 -flux = Flux1Schnell("black-forest-labs/FLUX.1-schnell") +flux = Flux1("black-forest-labs/FLUX.1-schnell") image = flux.generate_image( seed=3, diff --git a/src/flux_1_schnell/flux.py b/src/flux_1_schnell/flux.py index d2d3299..1fd564f 100644 --- a/src/flux_1_schnell/flux.py +++ b/src/flux_1_schnell/flux.py @@ -10,14 +10,13 @@ from flux_1_schnell.models.text_encoder.t5_encoder.t5_encoder import T5Encoder from flux_1_schnell.models.transformer.transformer import Transformer from flux_1_schnell.models.vae.vae import VAE from flux_1_schnell.post_processing.image_util import ImageUtil -from flux_1_schnell.scheduler.scheduler import FlowMatchEulerDiscreteNoiseScheduler from flux_1_schnell.tokenizer.clip_tokenizer import TokenizerCLIP from flux_1_schnell.tokenizer.t5_tokenizer import TokenizerT5 from flux_1_schnell.tokenizer.tokenizer_handler import TokenizerHandler from flux_1_schnell.weights.weight_handler import WeightHandler -class Flux1Schnell: +class Flux1: def __init__(self, repo_id: str): tokenizers = TokenizerHandler.load_from_disk_or_huggingface(repo_id) @@ -47,16 +46,12 @@ class Flux1Schnell: config=config ) - latents = FlowMatchEulerDiscreteNoiseScheduler.denoise( - t=t, - noise=noise, - latent=latents, - config=config - ) + dt = config.sigmas[t + 1] - config.sigmas[t] + latents += noise * dt mx.eval(latents) - latents = Flux1Schnell._unpack_latents(latents, config.width, config.height) + latents = Flux1._unpack_latents(latents, config.width, config.height) decoded = self.vae.decode(latents) return ImageUtil.to_image(decoded) diff --git a/src/flux_1_schnell/scheduler/__init__.py b/src/flux_1_schnell/scheduler/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/src/flux_1_schnell/scheduler/scheduler.py b/src/flux_1_schnell/scheduler/scheduler.py deleted file mode 100644 index 3fb3e7b..0000000 --- a/src/flux_1_schnell/scheduler/scheduler.py +++ /dev/null @@ -1,20 +0,0 @@ -import mlx.core as mx - -from flux_1_schnell.config.config import Config - - -class FlowMatchEulerDiscreteNoiseScheduler: - - @staticmethod - def denoise( - t: int, - noise: mx.array, - latent: mx.array, - config: Config, - ) -> mx.array: - sigma = config.sigmas[t] - denoised = latent - noise * sigma - derivative = (latent - denoised) / sigma - dt = config.sigmas[t + 1] - sigma - prev_sample = latent + derivative * dt - return prev_sample From 14c5b2bd0e354a22a3dd8930e519f655e763004a Mon Sep 17 00:00:00 2001 From: Fabio Date: Sun, 18 Aug 2024 01:01:01 +0200 Subject: [PATCH 02/20] fix architecture for dev --- main.py | 2 +- src/flux_1_schnell/flux.py | 6 ++++-- .../text_encoder/t5_encoder/t5_self_attention.py | 9 +++++---- .../models/transformer/guidance_embedder.py | 16 ++++++++++++++++ .../models/transformer/time_text_embed.py | 11 +++++++++-- .../models/transformer/transformer.py | 9 +++++---- src/flux_1_schnell/tokenizer/t5_tokenizer.py | 6 +++--- .../tokenizer/tokenizer_handler.py | 8 ++++---- 8 files changed, 47 insertions(+), 20 deletions(-) create mode 100644 src/flux_1_schnell/models/transformer/guidance_embedder.py diff --git a/main.py b/main.py index c5a0469..6d67901 100644 --- a/main.py +++ b/main.py @@ -6,7 +6,7 @@ sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), 'src'))) from flux_1_schnell.config.config import Config from flux_1_schnell.flux import Flux1 -flux = Flux1("black-forest-labs/FLUX.1-schnell") +flux = Flux1("black-forest-labs/FLUX.1-dev") image = flux.generate_image( seed=3, diff --git a/src/flux_1_schnell/flux.py b/src/flux_1_schnell/flux.py index 1fd564f..ef153b6 100644 --- a/src/flux_1_schnell/flux.py +++ b/src/flux_1_schnell/flux.py @@ -19,8 +19,10 @@ from flux_1_schnell.weights.weight_handler import WeightHandler class Flux1: def __init__(self, repo_id: str): - tokenizers = TokenizerHandler.load_from_disk_or_huggingface(repo_id) - self.t5_tokenizer = TokenizerT5(tokenizers.t5) + is_dev = "FLUX.1-dev" in repo_id + max_t5_length = 512 if is_dev else 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) weights = WeightHandler.load_from_disk_or_huggingface(repo_id) diff --git a/src/flux_1_schnell/models/text_encoder/t5_encoder/t5_self_attention.py b/src/flux_1_schnell/models/text_encoder/t5_encoder/t5_self_attention.py index dc2794d..a628d85 100644 --- a/src/flux_1_schnell/models/text_encoder/t5_encoder/t5_self_attention.py +++ b/src/flux_1_schnell/models/text_encoder/t5_encoder/t5_self_attention.py @@ -19,7 +19,8 @@ class T5SelfAttention(nn.Module): key_states = T5SelfAttention.shape(self.k(hidden_states)) value_states = T5SelfAttention.shape(self.v(hidden_states)) scores = mx.matmul(query_states, mx.transpose(key_states, (0, 1, 3, 2))) - position_bias = self._compute_bias() + seq_length = hidden_states.shape[1] + position_bias = self._compute_bias(seq_length=seq_length) scores += position_bias attn_weights = nn.softmax(scores, axis=-1) attn_output = T5SelfAttention.un_shape(mx.matmul(attn_weights, value_states)) @@ -34,9 +35,9 @@ class T5SelfAttention(nn.Module): def un_shape(states): return mx.reshape(mx.transpose(states, (0, 2, 1, 3)), (1, -1, 4096)) - def _compute_bias(self): - context_position = mx.arange(start=0, stop=256, step=1)[:, None] - memory_position = mx.arange(start=0, stop=256, step=1)[None, :] + def _compute_bias(self, seq_length): + context_position = mx.arange(start=0, stop=seq_length, step=1)[:, None] + memory_position = mx.arange(start=0, stop=seq_length, step=1)[None, :] relative_position = memory_position - context_position relative_position_bucket = T5SelfAttention._relative_position_bucket(relative_position) values = self.relative_attention_bias(relative_position_bucket) diff --git a/src/flux_1_schnell/models/transformer/guidance_embedder.py b/src/flux_1_schnell/models/transformer/guidance_embedder.py new file mode 100644 index 0000000..5ed915d --- /dev/null +++ b/src/flux_1_schnell/models/transformer/guidance_embedder.py @@ -0,0 +1,16 @@ +from mlx import nn +import mlx.core as mx + + +class GuidanceEmbedder(nn.Module): + + def __init__(self): + super().__init__() + self.linear_1 = nn.Linear(256, 3072) + self.linear_2 = nn.Linear(3072, 3072) + + def forward(self, sample: mx.array) -> mx.array: + sample = self.linear_1(sample) + sample = nn.silu(sample) + sample = self.linear_2(sample) + return sample \ No newline at end of file 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 a825740..abfc75d 100644 --- a/src/flux_1_schnell/models/transformer/time_text_embed.py +++ b/src/flux_1_schnell/models/transformer/time_text_embed.py @@ -5,18 +5,25 @@ import mlx.core as mx from flux_1_schnell.config.config import Config from flux_1_schnell.models.transformer.text_embedder import TextEmbedder from flux_1_schnell.models.transformer.timestep_embedder import TimestepEmbedder +from flux_1_schnell.models.transformer.guidance_embedder import GuidanceEmbedder class TimeTextEmbed(nn.Module): - def __init__(self): + def __init__(self, with_guidance_embed: bool = False): super().__init__() 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: - time_steps_proj = TimeTextEmbed._time_proj(time_step) + 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)) 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 4bb779a..114c82c 100644 --- a/src/flux_1_schnell/models/transformer/transformer.py +++ b/src/flux_1_schnell/models/transformer/transformer.py @@ -15,7 +15,8 @@ class Transformer(nn.Module): super().__init__() self.pos_embed = EmbedND() self.x_embedder = nn.Linear(64, 3072) - self.time_text_embed = TimeTextEmbed() + with_guidance_embed = "guidance_embedder" in weights["time_text_embed"].keys() + self.time_text_embed = TimeTextEmbed(with_guidance_embed = with_guidance_embed) self.context_embedder = nn.Linear(4096, 3072) self.transformer_blocks = [JointTransformerBlock(i) for i in range(19)] self.single_transformer_blocks = [SingleTransformerBlock(i) for i in range(38)] @@ -38,7 +39,7 @@ class Transformer(nn.Module): hidden_states = self.x_embedder(hidden_states) text_embeddings = self.time_text_embed.forward(time_step, pooled_prompt_embeds) encoder_hidden_states = self.context_embedder(prompt_embeds) - txt_ids = Transformer._prepare_text_ids() + txt_ids = Transformer._prepare_text_ids(seq_len = prompt_embeds.shape[1]) img_ids = Transformer._prepare_latent_image_ids(config.width, config.height) ids = mx.concatenate((txt_ids, img_ids), axis=1) image_rotary_emb = self.pos_embed.forward(ids) @@ -78,5 +79,5 @@ class Transformer(nn.Module): return latent_image_ids @staticmethod - def _prepare_text_ids() -> mx.array: - return mx.zeros((1, 256, 3)) + def _prepare_text_ids(seq_len) -> mx.array: + return mx.zeros((1, seq_len, 3)) diff --git a/src/flux_1_schnell/tokenizer/t5_tokenizer.py b/src/flux_1_schnell/tokenizer/t5_tokenizer.py index e227505..2ce659f 100644 --- a/src/flux_1_schnell/tokenizer/t5_tokenizer.py +++ b/src/flux_1_schnell/tokenizer/t5_tokenizer.py @@ -3,16 +3,16 @@ from transformers import T5Tokenizer class TokenizerT5: - MAX_TOKEN_LENGTH = 256 - def __init__(self, tokenizer: T5Tokenizer): + def __init__(self, tokenizer: T5Tokenizer, max_length: int = 256): self.tokenizer = tokenizer + self.max_length = max_length def tokenize(self, prompt: str) -> mx.array: return self.tokenizer( [prompt], padding="max_length", - max_length=TokenizerT5.MAX_TOKEN_LENGTH, + max_length=self.max_length, truncation=True, return_length=False, return_overflowing_tokens=False, diff --git a/src/flux_1_schnell/tokenizer/tokenizer_handler.py b/src/flux_1_schnell/tokenizer/tokenizer_handler.py index 8c754bb..0269e53 100644 --- a/src/flux_1_schnell/tokenizer/tokenizer_handler.py +++ b/src/flux_1_schnell/tokenizer/tokenizer_handler.py @@ -9,7 +9,7 @@ from flux_1_schnell.tokenizer.t5_tokenizer import TokenizerT5 class TokenizerHandler: - def __init__(self, repo_id: str): + def __init__(self, repo_id: str, max_t5_length: int = 256): root_path = TokenizerHandler._download_or_get_cached_tokenizers(repo_id) self.clip = transformers.CLIPTokenizer.from_pretrained( @@ -20,12 +20,12 @@ class TokenizerHandler: self.t5 = transformers.T5Tokenizer.from_pretrained( pretrained_model_name_or_path=root_path / "tokenizer_2", local_files_only=True, - max_length=TokenizerT5.MAX_TOKEN_LENGTH + max_length=max_t5_length ) @staticmethod - def load_from_disk_or_huggingface(repo_id: str) -> "TokenizerHandler": - return TokenizerHandler(repo_id) + def load_from_disk_or_huggingface(repo_id: str, max_t5_length: int = 256) -> "TokenizerHandler": + return TokenizerHandler(repo_id, max_t5_length) @staticmethod def _download_or_get_cached_tokenizers(repo_id: str) -> Path: From 40b40b554116f47b5ed476581b8029a83c5f6243 Mon Sep 17 00:00:00 2001 From: Fabio Date: Sun, 18 Aug 2024 13:30:19 +0200 Subject: [PATCH 03/20] add shift to schedule --- main.py | 3 ++- src/flux_1_schnell/config/config.py | 16 ++++++++++++---- src/flux_1_schnell/flux.py | 7 +++++-- .../models/transformer/time_text_embed.py | 5 ++--- .../models/transformer/transformer.py | 5 +++-- 5 files changed, 24 insertions(+), 12 deletions(-) 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) From 6977e9206147c0e95c3b3b9dc60ae55f8e3030ce Mon Sep 17 00:00:00 2001 From: Fabio Date: Sun, 18 Aug 2024 17:15:58 +0200 Subject: [PATCH 04/20] fix sigmas shift --- src/flux_1_schnell/config/config.py | 2 +- src/flux_1_schnell/flux.py | 3 +-- 2 files changed, 2 insertions(+), 3 deletions(-) diff --git a/src/flux_1_schnell/config/config.py b/src/flux_1_schnell/config/config.py index 5c63d07..ca5a9f6 100644 --- a/src/flux_1_schnell/config/config.py +++ b/src/flux_1_schnell/config/config.py @@ -34,5 +34,5 @@ class Config: x1 = 256 m = (1.15 - y1) / (4096 - x1) b = y1 - m * x1 - mu = m + b + mu = m * self.width * self.height / 256 + 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 bbed325..7a02bdc 100644 --- a/src/flux_1_schnell/flux.py +++ b/src/flux_1_schnell/flux.py @@ -20,8 +20,7 @@ class Flux1: def __init__(self, repo_id: str): self.is_dev = "FLUX.1-dev" in repo_id - # max_t5_length = 512 if self.is_dev else 256 - max_t5_length = 256 + max_t5_length = 512 if self.is_dev else 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) From acfc448f0c51744661d329ed6822f7d9c1be5856 Mon Sep 17 00:00:00 2001 From: Fabio Date: Sun, 18 Aug 2024 17:56:31 +0200 Subject: [PATCH 05/20] fix guidance scale --- src/flux_1_schnell/models/transformer/transformer.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/flux_1_schnell/models/transformer/transformer.py b/src/flux_1_schnell/models/transformer/transformer.py index 5c1b2b7..fa4e356 100644 --- a/src/flux_1_schnell/models/transformer/transformer.py +++ b/src/flux_1_schnell/models/transformer/transformer.py @@ -37,7 +37,7 @@ class Transformer(nn.Module): time_step = config.sigmas[t] * config.num_train_steps time_step = mx.broadcast_to(time_step, (1,)) hidden_states = self.x_embedder(hidden_states) - guidance = mx.broadcast_to(config.guidance, (1,)) + guidance = mx.broadcast_to(config.guidance * config.num_train_steps, (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]) From 9a7854218d9016158ea65831c5aacb23d996a64e Mon Sep 17 00:00:00 2001 From: Fabio Date: Sun, 18 Aug 2024 19:17:14 +0200 Subject: [PATCH 06/20] changes to gitignore --- .gitignore | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/.gitignore b/.gitignore index ee9ae73..e7b1ada 100644 --- a/.gitignore +++ b/.gitignore @@ -11,3 +11,7 @@ *.png *.jpg *.pyc + + +*.pt +trial.py \ No newline at end of file From d14f49f758b99e371f45f3351e6c807487ef6cba Mon Sep 17 00:00:00 2001 From: Fabio Date: Sun, 18 Aug 2024 21:23:52 +0200 Subject: [PATCH 07/20] revert changes to config --- src/flux_1_schnell/config/config.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/src/flux_1_schnell/config/config.py b/src/flux_1_schnell/config/config.py index ca5a9f6..0335128 100644 --- a/src/flux_1_schnell/config/config.py +++ b/src/flux_1_schnell/config/config.py @@ -19,13 +19,14 @@ class Config: log.warning("Width and height should be multiples of 16. Rounding down.") self.width = 16 * (height // 16) self.height = 16 * (width // 16) - self.sigmas = Config.base_sigmas(num_inference_steps) + base_sigmas = Config.base_sigmas(num_inference_steps) self.num_inference_steps = num_inference_steps self.guidance = guidance + self.sigmas = mx.concatenate([base_sigmas, mx.zeros(1)]) @staticmethod def base_sigmas(num_inference_steps): - sigmas = np.linspace(1.0, 0, num_inference_steps+1) + sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps) sigmas = mx.array(sigmas).astype(mx.float32) return sigmas From 68cbaab81645a5fc772be61308e5e77472e6ced1 Mon Sep 17 00:00:00 2001 From: Fabio Date: Sun, 18 Aug 2024 23:31:54 +0200 Subject: [PATCH 08/20] fix type and add main_dev --- .gitignore | 4 ---- main.py | 6 +++--- src/flux_1_schnell/config/config.py | 3 ++- src/flux_1_schnell/flux.py | 7 +++---- src/flux_1_schnell/models/transformer/transformer.py | 4 ++-- 5 files changed, 10 insertions(+), 14 deletions(-) diff --git a/.gitignore b/.gitignore index e7b1ada..ee9ae73 100644 --- a/.gitignore +++ b/.gitignore @@ -11,7 +11,3 @@ *.png *.jpg *.pyc - - -*.pt -trial.py \ No newline at end of file diff --git a/main.py b/main.py index aa11963..6d24b9b 100644 --- a/main.py +++ b/main.py @@ -6,15 +6,15 @@ sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), 'src'))) from flux_1_schnell.config.config import Config from flux_1_schnell.flux import Flux1 -flux = Flux1("black-forest-labs/FLUX.1-dev") +flux = Flux1("black-forest-labs/FLUX.1-dev", max_sequence_length=512) 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=20, - width=256, - height=256, + height=768, + width=1360, guidance=3.5, ) ) diff --git a/src/flux_1_schnell/config/config.py b/src/flux_1_schnell/config/config.py index 0335128..a88532e 100644 --- a/src/flux_1_schnell/config/config.py +++ b/src/flux_1_schnell/config/config.py @@ -5,7 +5,7 @@ import logging log = logging.getLogger(__name__) class Config: - precision: mx.Dtype = mx.float16 + precision: mx.Dtype = mx.bfloat16 num_train_steps = 1000 def __init__( @@ -37,3 +37,4 @@ class Config: b = y1 - m * x1 mu = m * self.width * self.height / 256 + b self.sigmas = mx.exp(mu) / (mx.exp(mu) + (1 / self.sigmas - 1)) + self.sigmas[-1] = 0 diff --git a/src/flux_1_schnell/flux.py b/src/flux_1_schnell/flux.py index 7a02bdc..39cbbea 100644 --- a/src/flux_1_schnell/flux.py +++ b/src/flux_1_schnell/flux.py @@ -18,11 +18,10 @@ from flux_1_schnell.weights.weight_handler import WeightHandler class Flux1: - def __init__(self, repo_id: str): + def __init__(self, repo_id: str, max_sequence_length: int = 512): self.is_dev = "FLUX.1-dev" in repo_id - max_t5_length = 512 if self.is_dev else 256 - tokenizers = TokenizerHandler.load_from_disk_or_huggingface(repo_id, max_t5_length) - self.t5_tokenizer = TokenizerT5(tokenizers.t5, max_length=max_t5_length) + tokenizers = TokenizerHandler.load_from_disk_or_huggingface(repo_id, max_sequence_length) + self.t5_tokenizer = TokenizerT5(tokenizers.t5, max_length=max_sequence_length) self.clip_tokenizer = TokenizerCLIP(tokenizers.clip) weights = WeightHandler.load_from_disk_or_huggingface(repo_id) diff --git a/src/flux_1_schnell/models/transformer/transformer.py b/src/flux_1_schnell/models/transformer/transformer.py index fa4e356..f5ce147 100644 --- a/src/flux_1_schnell/models/transformer/transformer.py +++ b/src/flux_1_schnell/models/transformer/transformer.py @@ -35,9 +35,9 @@ class Transformer(nn.Module): config: Config ) -> mx.array: time_step = config.sigmas[t] * config.num_train_steps - time_step = mx.broadcast_to(time_step, (1,)) + time_step = mx.broadcast_to(time_step, (1,)).astype(config.precision) hidden_states = self.x_embedder(hidden_states) - guidance = mx.broadcast_to(config.guidance * config.num_train_steps, (1,)) + guidance = mx.broadcast_to(config.guidance * config.num_train_steps, (1,)).astype(config.precision) 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]) From c2b600dd9e9298c7e7910ded44cda7ed95a778bd Mon Sep 17 00:00:00 2001 From: Fabio Date: Sun, 18 Aug 2024 23:31:59 +0200 Subject: [PATCH 09/20] fix type and add main_dev --- main.py | 10 +++++----- main_dev.py | 22 ++++++++++++++++++++++ 2 files changed, 27 insertions(+), 5 deletions(-) create mode 100644 main_dev.py diff --git a/main.py b/main.py index 6d24b9b..00567f2 100644 --- a/main.py +++ b/main.py @@ -4,19 +4,19 @@ import sys sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), 'src'))) from flux_1_schnell.config.config import Config -from flux_1_schnell.flux import Flux1 +from flux_1_schnell.flux import Flux1Schnell +from flux_1_schnell.post_processing.image_util import ImageUtil -flux = Flux1("black-forest-labs/FLUX.1-dev", max_sequence_length=512) +flux = Flux1Schnell("black-forest-labs/FLUX.1-schnell") 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=20, + num_inference_steps=2, height=768, width=1360, - guidance=3.5, ) ) -image.save("image.png") +ImageUtil.save_image(image, "image.png") \ No newline at end of file diff --git a/main_dev.py b/main_dev.py new file mode 100644 index 0000000..6d24b9b --- /dev/null +++ b/main_dev.py @@ -0,0 +1,22 @@ +import os +import sys + +sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), 'src'))) + +from flux_1_schnell.config.config import Config +from flux_1_schnell.flux import Flux1 + +flux = Flux1("black-forest-labs/FLUX.1-dev", max_sequence_length=512) + +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=20, + height=768, + width=1360, + guidance=3.5, + ) +) + +image.save("image.png") From 45309d0129b5f58e43543eced38c789bfd5dcb2f Mon Sep 17 00:00:00 2001 From: Fabio Date: Sun, 18 Aug 2024 23:41:49 +0200 Subject: [PATCH 10/20] fix main --- README.md | 2 +- main.py | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/README.md b/README.md index 06bebbe..f4b057e 100644 --- a/README.md +++ b/README.md @@ -20,7 +20,7 @@ like [Numpy](https://numpy.org) and [Pillow](https://pypi.org/project/pillow/) f ### Models - [x] FLUX.1-Scnhell -- [ ] FLUX.1-Dev +- [x] FLUX.1-Dev ### Installation 1. Clone the repo: diff --git a/main.py b/main.py index 5cf39bb..8b26c5b 100644 --- a/main.py +++ b/main.py @@ -4,10 +4,10 @@ import sys sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), 'src'))) from flux_1_schnell.config.config import Config -from flux_1_schnell.flux import Flux1Schnell +from flux_1_schnell.flux import Flux1 from flux_1_schnell.post_processing.image_util import ImageUtil -flux = Flux1Schnell("black-forest-labs/FLUX.1-schnell") +flux = Flux1("black-forest-labs/FLUX.1-schnell", max_sequence_length=256) image = flux.generate_image( seed=3, From 40b77e81bc548bc1115c521e5172a4b958b085c8 Mon Sep 17 00:00:00 2001 From: Fabio Date: Mon, 19 Aug 2024 00:05:37 +0200 Subject: [PATCH 11/20] remove main dev, add hints to main --- main.py | 3 +++ main_dev.py | 22 ---------------------- 2 files changed, 3 insertions(+), 22 deletions(-) delete mode 100644 main_dev.py diff --git a/main.py b/main.py index 8b26c5b..51b12ea 100644 --- a/main.py +++ b/main.py @@ -8,6 +8,8 @@ from flux_1_schnell.flux import Flux1 from flux_1_schnell.post_processing.image_util import ImageUtil flux = Flux1("black-forest-labs/FLUX.1-schnell", max_sequence_length=256) +# flux = Flux1("black-forest-labs/FLUX.1-dev", max_sequence_length=512) + image = flux.generate_image( seed=3, @@ -16,6 +18,7 @@ image = flux.generate_image( num_inference_steps=2, height=768, width=1360, + guidance=3.5, ) ) diff --git a/main_dev.py b/main_dev.py deleted file mode 100644 index 6d24b9b..0000000 --- a/main_dev.py +++ /dev/null @@ -1,22 +0,0 @@ -import os -import sys - -sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), 'src'))) - -from flux_1_schnell.config.config import Config -from flux_1_schnell.flux import Flux1 - -flux = Flux1("black-forest-labs/FLUX.1-dev", max_sequence_length=512) - -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=20, - height=768, - width=1360, - guidance=3.5, - ) -) - -image.save("image.png") From e60cd79f0ca024e0485d600b8cf9e8c3e2635cf7 Mon Sep 17 00:00:00 2001 From: Fabio Date: Mon, 19 Aug 2024 20:41:11 +0200 Subject: [PATCH 12/20] make config static and freeze after initialisation --- src/flux_1_schnell/config/config.py | 40 +++++++++++-------- src/flux_1_schnell/flux.py | 10 +++-- .../models/transformer/transformer.py | 5 ++- 3 files changed, 33 insertions(+), 22 deletions(-) diff --git a/src/flux_1_schnell/config/config.py b/src/flux_1_schnell/config/config.py index 3524b9c..13b63d8 100644 --- a/src/flux_1_schnell/config/config.py +++ b/src/flux_1_schnell/config/config.py @@ -1,3 +1,4 @@ +from dataclasses import dataclass import mlx.core as mx import numpy as np import logging @@ -5,37 +6,44 @@ import logging log = logging.getLogger(__name__) +def get_sigmas(num_inference_steps): + sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps) + sigmas = mx.array(sigmas).astype(mx.float32) + return mx.concatenate([sigmas, mx.zeros(1)]) + +def shift_sigmas(sigmas, width, height): + y1 = 0.5 + x1 = 256 + m = (1.15 - y1) / (4096 - x1) + b = y1 - m * x1 + mu = m * width * height / 256 + b + shifted_sigmas = mx.exp(mu) / (mx.exp(mu) + (1 / sigmas - 1)) + shifted_sigmas[-1] = 0 + return shifted_sigmas + + +@dataclass class Config: precision: mx.Dtype = mx.bfloat16 - num_train_steps = 1000 def __init__( self, + num_train_steps: int = 1000, num_inference_steps: int = 4, width: int = 1024, height: int = 1024, guidance: float = 4.0, ): + self.num_train_steps = num_train_steps 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.num_inference_steps = num_inference_steps self.guidance = guidance - self.sigmas = mx.concatenate([base_sigmas, mx.zeros(1)]) - @staticmethod - def base_sigmas(num_inference_steps): - sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps) - sigmas = mx.array(sigmas).astype(mx.float32) - return sigmas + def __post_init__(self, **data): + super().__init__(**data) + self.__config__.frozen = True + - def shift_sigmas(self): - y1 = 0.5 - x1 = 256 - m = (1.15 - y1) / (4096 - x1) - b = y1 - m * x1 - mu = m * self.width * self.height / 256 + b - self.sigmas = mx.exp(mu) / (mx.exp(mu) + (1 / self.sigmas - 1)) - self.sigmas[-1] = 0 diff --git a/src/flux_1_schnell/flux.py b/src/flux_1_schnell/flux.py index f5eb5dd..a034c68 100644 --- a/src/flux_1_schnell/flux.py +++ b/src/flux_1_schnell/flux.py @@ -3,7 +3,7 @@ import mlx.core as mx from PIL import Image from tqdm import tqdm -from flux_1_schnell.config.config import Config +from flux_1_schnell.config.config import Config, get_sigmas, shift_sigmas from flux_1_schnell.latent_creator.latent_creator import LatentCreator from flux_1_schnell.models.text_encoder.clip_encoder.clip_encoder import CLIPEncoder from flux_1_schnell.models.text_encoder.t5_encoder.t5_encoder import T5Encoder @@ -31,8 +31,9 @@ class Flux1: self.clip_text_encoder = CLIPEncoder(weights.clip_encoder) def generate_image(self, seed: int, prompt: str, config: Config = Config()) -> PIL.Image.Image: + sigmas = get_sigmas(config.num_inference_steps) if self.is_dev: - config.shift_sigmas() + sigmas = shift_sigmas(sigmas) latents = LatentCreator.create(config.height, config.width, seed) t5_tokens = self.t5_tokenizer.tokenize(prompt) @@ -46,10 +47,11 @@ class Flux1: prompt_embeds=prompt_embeds, pooled_prompt_embeds=pooled_prompt_embeds, hidden_states=latents, - config=config + config=config, + sigmas=sigmas ) - dt = config.sigmas[t + 1] - config.sigmas[t] + dt = sigmas[t + 1] - sigmas[t] latents += noise * dt mx.eval(latents) diff --git a/src/flux_1_schnell/models/transformer/transformer.py b/src/flux_1_schnell/models/transformer/transformer.py index 6638501..11daa28 100644 --- a/src/flux_1_schnell/models/transformer/transformer.py +++ b/src/flux_1_schnell/models/transformer/transformer.py @@ -32,9 +32,10 @@ class Transformer(nn.Module): prompt_embeds: mx.array, pooled_prompt_embeds: mx.array, hidden_states: mx.array, - config: Config + config: Config, + sigmas: mx.array, ) -> mx.array: - time_step = config.sigmas[t] * config.num_train_steps + time_step = sigmas[t] * config.num_train_steps time_step = mx.broadcast_to(time_step, (1,)).astype(config.precision) hidden_states = self.x_embedder(hidden_states) guidance = mx.broadcast_to(config.guidance * config.num_train_steps, (1,)).astype(config.precision) From 3710f101e5791104f60240afe15240cb95f55d77 Mon Sep 17 00:00:00 2001 From: Fabio Date: Mon, 19 Aug 2024 21:09:41 +0200 Subject: [PATCH 13/20] small fix and update readme --- README.md | 1 - src/flux_1_schnell/flux.py | 2 +- 2 files changed, 1 insertion(+), 2 deletions(-) diff --git a/README.md b/README.md index f4b057e..e4510cd 100644 --- a/README.md +++ b/README.md @@ -133,6 +133,5 @@ Luxury food photograph of an italian Linguine pasta alle vongole dish with lots ### TODO -- FLUX Dev implementation - LoRA adapters - Command line args \ No newline at end of file diff --git a/src/flux_1_schnell/flux.py b/src/flux_1_schnell/flux.py index a034c68..e4da8b2 100644 --- a/src/flux_1_schnell/flux.py +++ b/src/flux_1_schnell/flux.py @@ -33,7 +33,7 @@ class Flux1: def generate_image(self, seed: int, prompt: str, config: Config = Config()) -> PIL.Image.Image: sigmas = get_sigmas(config.num_inference_steps) if self.is_dev: - sigmas = shift_sigmas(sigmas) + sigmas = shift_sigmas(sigmas, config.width, config.height) latents = LatentCreator.create(config.height, config.width, seed) t5_tokens = self.t5_tokenizer.tokenize(prompt) From d1b4557d2ef22e327dbe2c4d677b3da5a9f2bd23 Mon Sep 17 00:00:00 2001 From: Fabio Date: Mon, 19 Aug 2024 21:23:32 +0200 Subject: [PATCH 14/20] small fix to main --- main.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main.py b/main.py index 552b9f2..3320690 100644 --- a/main.py +++ b/main.py @@ -28,7 +28,7 @@ def main(): flux = Flux1("black-forest-labs/FLUX.1-schnell", max_sequence_length=args.max_sequence_length) image = flux.generate_image( - seed=args.seed, + seed=seed, prompt=args.prompt, config=Config( num_inference_steps=args.steps, From c0cc8059e976de9f8b1de20df07e67698a5e3ed5 Mon Sep 17 00:00:00 2001 From: Fabio Date: Mon, 19 Aug 2024 21:41:19 +0200 Subject: [PATCH 15/20] fix cli --- main.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/main.py b/main.py index 3320690..2e859ea 100644 --- a/main.py +++ b/main.py @@ -15,7 +15,7 @@ def main(): parser.add_argument('--output', type=str, default="image.png", help='The filename for the output image. Default is "image.png".') parser.add_argument('--model', type=str, default="black-forest-labs/FLUX.1-schnell", help='The model to use. Default is "black-forest-labs/FLUX.1-schnell".') parser.add_argument('--max_sequence_length', type=int, default=256, help='Max Sequence Length (Default is 256)') - parser.add_argument('--seed', type=int, default=0, help='Entropy Seed (Default is time-based random-seed)') + parser.add_argument('--seed', type=int, default=None, help='Entropy Seed (Default is time-based random-seed)') parser.add_argument('--height', type=int, default=1024, help='Image height (Default is 1024)') parser.add_argument('--width', type=int, default=1024, help='Image width (Default is 1024)') parser.add_argument('--steps', type=int, default=4, help='Inference Steps') @@ -23,9 +23,9 @@ def main(): args = parser.parse_args() - seed = args.seed or int(time.time()) + seed = int(time.time()) if args.seed is None else args.seed - flux = Flux1("black-forest-labs/FLUX.1-schnell", max_sequence_length=args.max_sequence_length) + flux = Flux1(args.model, max_sequence_length=args.max_sequence_length) image = flux.generate_image( seed=seed, From 4b42e50257e86158e1c0ef0ff551eb6c904ccd7f Mon Sep 17 00:00:00 2001 From: filipstrand Date: Tue, 20 Aug 2024 09:32:05 +0200 Subject: [PATCH 16/20] Add separate model config and other small updates --- README.md | 6 +- main.py | 6 +- src/flux_1_schnell/config/config.py | 61 +++++++++++-------- src/flux_1_schnell/config/model_config.py | 29 +++++++++ src/flux_1_schnell/flux.py | 37 ++++++++--- .../models/transformer/guidance_embedder.py | 2 +- .../models/transformer/transformer.py | 9 ++- 7 files changed, 104 insertions(+), 46 deletions(-) create mode 100644 src/flux_1_schnell/config/model_config.py diff --git a/README.md b/README.md index d32d41f..bed2eb2 100644 --- a/README.md +++ b/README.md @@ -49,7 +49,7 @@ python main.py --prompt "Luxury food photograph" --steps 2 --seed 2 - **`--output`** (optional, `str`, default: `"image.png"`): Output image filename. -- **`--model`** (optional, `str`, default: `"black-forest-labs/FLUX.1-schnell"`): Model to use for generation. +- **`--model`** (optional, `str`, default: `"schnell"`): Model to use for generation. - **`--seed`** (optional, `int`, default: `0`): Seed for random number generation. Default is time-based. @@ -67,10 +67,10 @@ import sys sys.path.append("/path/to/mflux/src") from flux_1_schnell.config.config import Config -from flux_1_schnell.flux import Flux1Schnell +from flux_1_schnell.flux import Flux1 from flux_1_schnell.post_processing.image_util import ImageUtil -flux = Flux1Schnell("black-forest-labs/FLUX.1-schnell") +flux = Flux1.from_repo("black-forest-labs/FLUX.1-schnell") image = flux.generate_image( seed=3, diff --git a/main.py b/main.py index 2e859ea..ff903b0 100644 --- a/main.py +++ b/main.py @@ -9,12 +9,12 @@ from flux_1_schnell.config.config import Config from flux_1_schnell.flux import Flux1 from flux_1_schnell.post_processing.image_util import ImageUtil + def main(): parser = argparse.ArgumentParser(description='Generate an image based on a prompt.') parser.add_argument('--prompt', type=str, required=True, help='The textual description of the image to generate.') parser.add_argument('--output', type=str, default="image.png", help='The filename for the output image. Default is "image.png".') - parser.add_argument('--model', type=str, default="black-forest-labs/FLUX.1-schnell", help='The model to use. Default is "black-forest-labs/FLUX.1-schnell".') - parser.add_argument('--max_sequence_length', type=int, default=256, help='Max Sequence Length (Default is 256)') + parser.add_argument('--model', type=str, default="schnell", help='The model to use. Default is "schnell".') parser.add_argument('--seed', type=int, default=None, help='Entropy Seed (Default is time-based random-seed)') parser.add_argument('--height', type=int, default=1024, help='Image height (Default is 1024)') parser.add_argument('--width', type=int, default=1024, help='Image width (Default is 1024)') @@ -25,7 +25,7 @@ def main(): seed = int(time.time()) if args.seed is None else args.seed - flux = Flux1(args.model, max_sequence_length=args.max_sequence_length) + flux = Flux1.from_alias(args.model) image = flux.generate_image( seed=seed, diff --git a/src/flux_1_schnell/config/config.py b/src/flux_1_schnell/config/config.py index 13b63d8..0ea5122 100644 --- a/src/flux_1_schnell/config/config.py +++ b/src/flux_1_schnell/config/config.py @@ -1,28 +1,13 @@ -from dataclasses import dataclass +import logging + import mlx.core as mx import numpy as np -import logging + +from flux_1_schnell.config.model_config import ModelConfig log = logging.getLogger(__name__) -def get_sigmas(num_inference_steps): - sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps) - sigmas = mx.array(sigmas).astype(mx.float32) - return mx.concatenate([sigmas, mx.zeros(1)]) - -def shift_sigmas(sigmas, width, height): - y1 = 0.5 - x1 = 256 - m = (1.15 - y1) / (4096 - x1) - b = y1 - m * x1 - mu = m * width * height / 256 + b - shifted_sigmas = mx.exp(mu) / (mx.exp(mu) + (1 / sigmas - 1)) - shifted_sigmas[-1] = 0 - return shifted_sigmas - - -@dataclass class Config: precision: mx.Dtype = mx.bfloat16 @@ -33,6 +18,7 @@ class Config: width: int = 1024, height: int = 1024, guidance: float = 4.0, + sigmas: mx.array | None = None, ): self.num_train_steps = num_train_steps if width % 16 != 0 or height % 16 != 0: @@ -41,9 +27,36 @@ class Config: self.height = 16 * (width // 16) self.num_inference_steps = num_inference_steps self.guidance = guidance + self.sigmas = sigmas - def __post_init__(self, **data): - super().__init__(**data) - self.__config__.frozen = True - - + def copy_with_sigmas(self, model: ModelConfig) -> "Config": + sigmas = Config._get_sigmas(self.num_inference_steps) + if model == ModelConfig.FLUX1_DEV: + sigmas = Config._shift_sigmas(sigmas, self.width, self.height) + + return Config( + num_train_steps=self.num_train_steps, + num_inference_steps=self.num_inference_steps, + width=self.width, + height=self.height, + guidance=self.guidance, + sigmas=sigmas, + ) + + @staticmethod + def _get_sigmas(num_inference_steps): + sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps) + sigmas = mx.array(sigmas).astype(mx.float32) + return mx.concatenate([sigmas, mx.zeros(1)]) + + @staticmethod + def _shift_sigmas(sigmas: mx.array, width: int, height: int): + y1 = 0.5 + x1 = 256 + m = (1.15 - y1) / (4096 - x1) + b = y1 - m * x1 + mu = m * width * height / 256 + b + mu = mx.array(mu) + shifted_sigmas = mx.exp(mu) / (mx.exp(mu) + (1 / sigmas - 1)) + shifted_sigmas[-1] = 0 + return shifted_sigmas diff --git a/src/flux_1_schnell/config/model_config.py b/src/flux_1_schnell/config/model_config.py new file mode 100644 index 0000000..0f70952 --- /dev/null +++ b/src/flux_1_schnell/config/model_config.py @@ -0,0 +1,29 @@ +from enum import Enum + + +class ModelConfig(Enum): + FLUX1_DEV = ("black-forest-labs/FLUX.1-dev", "dev", 512) + FLUX1_SCHNELL = ("black-forest-labs/FLUX.1-schnell", "schnell", 256) + + def __init__(self, model_name: str, alias: str, max_sequence_length: int): + self.alias = alias + self.model_name = model_name + self.max_sequence_length = max_sequence_length + + @staticmethod + def from_repo(model_name: str) -> "ModelConfig": + try: + for model in ModelConfig: + if model.model_name == model_name: + return model + except KeyError: + raise ValueError(f"'{model_name}' is not a valid model") + + @staticmethod + def from_alias(alias: str) -> "ModelConfig": + try: + for model in ModelConfig: + if model.alias == alias: + return model + except KeyError: + raise ValueError(f"'{alias}' is not a valid model") diff --git a/src/flux_1_schnell/flux.py b/src/flux_1_schnell/flux.py index e4da8b2..cbea48f 100644 --- a/src/flux_1_schnell/flux.py +++ b/src/flux_1_schnell/flux.py @@ -3,8 +3,9 @@ import mlx.core as mx from PIL import Image from tqdm import tqdm -from flux_1_schnell.config.config import Config, get_sigmas, shift_sigmas +from flux_1_schnell.config.config import Config from flux_1_schnell.latent_creator.latent_creator import LatentCreator +from flux_1_schnell.config.model_config import ModelConfig from flux_1_schnell.models.text_encoder.clip_encoder.clip_encoder import CLIPEncoder from flux_1_schnell.models.text_encoder.t5_encoder.t5_encoder import T5Encoder from flux_1_schnell.models.transformer.transformer import Transformer @@ -18,44 +19,60 @@ from flux_1_schnell.weights.weight_handler import WeightHandler class Flux1: - def __init__(self, repo_id: str, max_sequence_length: int = 512): - self.is_dev = "FLUX.1-dev" in repo_id - tokenizers = TokenizerHandler.load_from_disk_or_huggingface(repo_id, max_sequence_length) - self.t5_tokenizer = TokenizerT5(tokenizers.t5, max_length=max_sequence_length) + def __init__(self, repo_id: str): + self.model_config = ModelConfig.from_repo(repo_id) + + # Initialize the tokenizers + tokenizers = TokenizerHandler.load_from_disk_or_huggingface(repo_id, self.model_config.max_sequence_length) + self.t5_tokenizer = TokenizerT5(tokenizers.t5, max_length=self.model_config.max_sequence_length) self.clip_tokenizer = TokenizerCLIP(tokenizers.clip) + # Initialize the models weights = WeightHandler.load_from_disk_or_huggingface(repo_id) self.vae = VAE(weights.vae) self.transformer = Transformer(weights.transformer) self.t5_text_encoder = T5Encoder(weights.t5_encoder) self.clip_text_encoder = CLIPEncoder(weights.clip_encoder) + @staticmethod + def from_repo(repo_id: str) -> "Flux1": + return Flux1(repo_id) + + @staticmethod + def from_alias(alias: str) -> "Flux1": + return Flux1(ModelConfig.from_alias(alias).model_name) + def generate_image(self, seed: int, prompt: str, config: Config = Config()) -> PIL.Image.Image: - sigmas = get_sigmas(config.num_inference_steps) - if self.is_dev: - sigmas = shift_sigmas(sigmas, config.width, config.height) + # Create a new config with sigmas based on what model we are running + config = config.copy_with_sigmas(self.model_config) + + # Create the latents latents = LatentCreator.create(config.height, config.width, seed) + # Embedd 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 tqdm(range(config.num_inference_steps)): + # Predict the noise noise = self.transformer.predict( t=t, prompt_embeds=prompt_embeds, pooled_prompt_embeds=pooled_prompt_embeds, hidden_states=latents, config=config, - sigmas=sigmas ) - dt = sigmas[t + 1] - sigmas[t] + # Take one denoise step + dt = config.sigmas[t + 1] - config.sigmas[t] latents += noise * dt + # To enable progress tracking mx.eval(latents) + # Decode the latent array latents = Flux1._unpack_latents(latents, config.height, config.width) decoded = self.vae.decode(latents) return ImageUtil.to_image(decoded) diff --git a/src/flux_1_schnell/models/transformer/guidance_embedder.py b/src/flux_1_schnell/models/transformer/guidance_embedder.py index 5ed915d..ca21128 100644 --- a/src/flux_1_schnell/models/transformer/guidance_embedder.py +++ b/src/flux_1_schnell/models/transformer/guidance_embedder.py @@ -13,4 +13,4 @@ class GuidanceEmbedder(nn.Module): sample = self.linear_1(sample) sample = nn.silu(sample) sample = self.linear_2(sample) - return sample \ No newline at end of file + return sample diff --git a/src/flux_1_schnell/models/transformer/transformer.py b/src/flux_1_schnell/models/transformer/transformer.py index 11daa28..22d1f28 100644 --- a/src/flux_1_schnell/models/transformer/transformer.py +++ b/src/flux_1_schnell/models/transformer/transformer.py @@ -16,7 +16,7 @@ class Transformer(nn.Module): self.pos_embed = EmbedND() self.x_embedder = nn.Linear(64, 3072) with_guidance_embed = "guidance_embedder" in weights["time_text_embed"].keys() - self.time_text_embed = TimeTextEmbed(with_guidance_embed = with_guidance_embed) + self.time_text_embed = TimeTextEmbed(with_guidance_embed=with_guidance_embed) self.context_embedder = nn.Linear(4096, 3072) self.transformer_blocks = [JointTransformerBlock(i) for i in range(19)] self.single_transformer_blocks = [SingleTransformerBlock(i) for i in range(38)] @@ -33,15 +33,14 @@ class Transformer(nn.Module): pooled_prompt_embeds: mx.array, hidden_states: mx.array, config: Config, - sigmas: mx.array, ) -> mx.array: - time_step = sigmas[t] * config.num_train_steps + time_step = config.sigmas[t] * config.num_train_steps time_step = mx.broadcast_to(time_step, (1,)).astype(config.precision) hidden_states = self.x_embedder(hidden_states) guidance = mx.broadcast_to(config.guidance * config.num_train_steps, (1,)).astype(config.precision) 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]) + txt_ids = Transformer._prepare_text_ids(seq_len=prompt_embeds.shape[1]) img_ids = Transformer._prepare_latent_image_ids(config.height, config.width) ids = mx.concatenate((txt_ids, img_ids), axis=1) image_rotary_emb = self.pos_embed.forward(ids) @@ -81,5 +80,5 @@ class Transformer(nn.Module): return latent_image_ids @staticmethod - def _prepare_text_ids(seq_len) -> mx.array: + def _prepare_text_ids(seq_len: mx.array) -> mx.array: return mx.zeros((1, seq_len, 3)) From 7795d47cd3a0bfe1a5f8934c85af6527100bd5d1 Mon Sep 17 00:00:00 2001 From: filipstrand Date: Tue, 20 Aug 2024 10:57:14 +0200 Subject: [PATCH 17/20] Add a separate runtime config --- src/flux_1_schnell/config/config.py | 39 ------------- src/flux_1_schnell/config/runtime_config.py | 58 +++++++++++++++++++ src/flux_1_schnell/flux.py | 5 +- .../models/transformer/transformer.py | 3 +- 4 files changed, 63 insertions(+), 42 deletions(-) create mode 100644 src/flux_1_schnell/config/runtime_config.py diff --git a/src/flux_1_schnell/config/config.py b/src/flux_1_schnell/config/config.py index 0ea5122..ab8a2d5 100644 --- a/src/flux_1_schnell/config/config.py +++ b/src/flux_1_schnell/config/config.py @@ -1,9 +1,6 @@ import logging import mlx.core as mx -import numpy as np - -from flux_1_schnell.config.model_config import ModelConfig log = logging.getLogger(__name__) @@ -13,50 +10,14 @@ class Config: def __init__( self, - num_train_steps: int = 1000, num_inference_steps: int = 4, width: int = 1024, height: int = 1024, guidance: float = 4.0, - sigmas: mx.array | None = None, ): - self.num_train_steps = num_train_steps 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) self.num_inference_steps = num_inference_steps self.guidance = guidance - self.sigmas = sigmas - - def copy_with_sigmas(self, model: ModelConfig) -> "Config": - sigmas = Config._get_sigmas(self.num_inference_steps) - if model == ModelConfig.FLUX1_DEV: - sigmas = Config._shift_sigmas(sigmas, self.width, self.height) - - return Config( - num_train_steps=self.num_train_steps, - num_inference_steps=self.num_inference_steps, - width=self.width, - height=self.height, - guidance=self.guidance, - sigmas=sigmas, - ) - - @staticmethod - def _get_sigmas(num_inference_steps): - sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps) - sigmas = mx.array(sigmas).astype(mx.float32) - return mx.concatenate([sigmas, mx.zeros(1)]) - - @staticmethod - def _shift_sigmas(sigmas: mx.array, width: int, height: int): - y1 = 0.5 - x1 = 256 - m = (1.15 - y1) / (4096 - x1) - b = y1 - m * x1 - mu = m * width * height / 256 + b - mu = mx.array(mu) - shifted_sigmas = mx.exp(mu) / (mx.exp(mu) + (1 / sigmas - 1)) - shifted_sigmas[-1] = 0 - return shifted_sigmas diff --git a/src/flux_1_schnell/config/runtime_config.py b/src/flux_1_schnell/config/runtime_config.py new file mode 100644 index 0000000..3807ec1 --- /dev/null +++ b/src/flux_1_schnell/config/runtime_config.py @@ -0,0 +1,58 @@ +import mlx.core as mx +import numpy as np + +from flux_1_schnell.config.config import Config +from flux_1_schnell.config.model_config import ModelConfig + + +class RuntimeConfig: + + def __init__(self, config: Config, model_config: ModelConfig): + self.config = config + self.num_train_steps = 1000 + self.sigmas = self._create_sigmas(config, model_config) + + @property + def height(self): + return self.config.height + + @property + def width(self): + return self.config.height + + @property + def guidance(self): + return self.config.guidance + + @property + def num_inference_steps(self): + return self.config.num_inference_steps + + @property + def precision(self): + return self.config.precision + + @staticmethod + def _create_sigmas(config, model): + sigmas = RuntimeConfig._create_sigmas_values(config.num_inference_steps) + if model == ModelConfig.FLUX1_DEV: + sigmas = RuntimeConfig._shift_sigmas(sigmas, config.width, config.height) + return sigmas + + @staticmethod + def _create_sigmas_values(num_inference_steps: int) -> mx.array: + sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps) + sigmas = mx.array(sigmas).astype(mx.float32) + return mx.concatenate([sigmas, mx.zeros(1)]) + + @staticmethod + def _shift_sigmas(sigmas: mx.array, width: int, height: int) -> mx.array: + y1 = 0.5 + x1 = 256 + m = (1.15 - y1) / (4096 - x1) + b = y1 - m * x1 + mu = m * width * height / 256 + b + mu = mx.array(mu) + shifted_sigmas = mx.exp(mu) / (mx.exp(mu) + (1 / sigmas - 1)) + shifted_sigmas[-1] = 0 + return shifted_sigmas diff --git a/src/flux_1_schnell/flux.py b/src/flux_1_schnell/flux.py index cbea48f..4b38c2c 100644 --- a/src/flux_1_schnell/flux.py +++ b/src/flux_1_schnell/flux.py @@ -4,6 +4,7 @@ from PIL import Image from tqdm import tqdm from flux_1_schnell.config.config import Config +from flux_1_schnell.config.runtime_config import RuntimeConfig from flux_1_schnell.latent_creator.latent_creator import LatentCreator from flux_1_schnell.config.model_config import ModelConfig from flux_1_schnell.models.text_encoder.clip_encoder.clip_encoder import CLIPEncoder @@ -43,8 +44,8 @@ class Flux1: return Flux1(ModelConfig.from_alias(alias).model_name) def generate_image(self, seed: int, prompt: str, config: Config = Config()) -> PIL.Image.Image: - # Create a new config with sigmas based on what model we are running - config = config.copy_with_sigmas(self.model_config) + # Create a new runtime config based on the model type and input parameters + config = RuntimeConfig(config, self.model_config) # Create the latents latents = LatentCreator.create(config.height, config.width, seed) diff --git a/src/flux_1_schnell/models/transformer/transformer.py b/src/flux_1_schnell/models/transformer/transformer.py index 22d1f28..15b6119 100644 --- a/src/flux_1_schnell/models/transformer/transformer.py +++ b/src/flux_1_schnell/models/transformer/transformer.py @@ -2,6 +2,7 @@ import mlx.core as mx from mlx import nn from flux_1_schnell.config.config import Config +from flux_1_schnell.config.runtime_config import RuntimeConfig from flux_1_schnell.models.transformer.ada_layer_norm_continous import AdaLayerNormContinuous from flux_1_schnell.models.transformer.embed_nd import EmbedND from flux_1_schnell.models.transformer.joint_transformer_block import JointTransformerBlock @@ -32,7 +33,7 @@ class Transformer(nn.Module): prompt_embeds: mx.array, pooled_prompt_embeds: mx.array, hidden_states: mx.array, - config: Config, + config: RuntimeConfig, ) -> mx.array: time_step = config.sigmas[t] * config.num_train_steps time_step = mx.broadcast_to(time_step, (1,)).astype(config.precision) From 4d613c7390530ee30500cc610500722594331a14 Mon Sep 17 00:00:00 2001 From: filipstrand Date: Tue, 20 Aug 2024 11:45:36 +0200 Subject: [PATCH 18/20] Move `num_train_steps` to ModelConfig --- src/flux_1_schnell/config/model_config.py | 13 ++++++++++--- src/flux_1_schnell/config/runtime_config.py | 6 +++++- 2 files changed, 15 insertions(+), 4 deletions(-) diff --git a/src/flux_1_schnell/config/model_config.py b/src/flux_1_schnell/config/model_config.py index 0f70952..d236cc2 100644 --- a/src/flux_1_schnell/config/model_config.py +++ b/src/flux_1_schnell/config/model_config.py @@ -2,12 +2,19 @@ from enum import Enum class ModelConfig(Enum): - FLUX1_DEV = ("black-forest-labs/FLUX.1-dev", "dev", 512) - FLUX1_SCHNELL = ("black-forest-labs/FLUX.1-schnell", "schnell", 256) + FLUX1_DEV = ("black-forest-labs/FLUX.1-dev", "dev", 1000, 512) + FLUX1_SCHNELL = ("black-forest-labs/FLUX.1-schnell", "schnell", 1000, 256) - def __init__(self, model_name: str, alias: str, max_sequence_length: int): + def __init__( + self, + model_name: str, + alias: str, + num_train_steps: int, + max_sequence_length: int, + ): self.alias = alias self.model_name = model_name + self.num_train_steps = num_train_steps self.max_sequence_length = max_sequence_length @staticmethod diff --git a/src/flux_1_schnell/config/runtime_config.py b/src/flux_1_schnell/config/runtime_config.py index 3807ec1..e4b3937 100644 --- a/src/flux_1_schnell/config/runtime_config.py +++ b/src/flux_1_schnell/config/runtime_config.py @@ -9,7 +9,7 @@ class RuntimeConfig: def __init__(self, config: Config, model_config: ModelConfig): self.config = config - self.num_train_steps = 1000 + self.model_config = model_config self.sigmas = self._create_sigmas(config, model_config) @property @@ -32,6 +32,10 @@ class RuntimeConfig: def precision(self): return self.config.precision + @property + def num_train_steps(self): + return self.model_config.num_train_steps + @staticmethod def _create_sigmas(config, model): sigmas = RuntimeConfig._create_sigmas_values(config.num_inference_steps) From f9f032a220dff98d16248df9a03e19f7f32c01f1 Mon Sep 17 00:00:00 2001 From: filipstrand Date: Tue, 20 Aug 2024 18:17:33 +0200 Subject: [PATCH 19/20] Update flag description --- main.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main.py b/main.py index bedaa55..0372c54 100644 --- a/main.py +++ b/main.py @@ -14,7 +14,7 @@ def main(): parser = argparse.ArgumentParser(description='Generate an image based on a prompt.') parser.add_argument('--prompt', type=str, required=True, help='The textual description of the image to generate.') parser.add_argument('--output', type=str, default="image.png", help='The filename for the output image. Default is "image.png".') - parser.add_argument('--model', type=str, default="schnell", help='The model to use. Default is "schnell".') + parser.add_argument('--model', type=str, default="schnell", help='The model to use ("schnell" or "dev"). Default is "schnell".') parser.add_argument('--seed', type=int, default=None, help='Entropy Seed (Default is time-based random-seed)') parser.add_argument('--height', type=int, default=1024, help='Image height (Default is 1024)') parser.add_argument('--width', type=int, default=1024, help='Image width (Default is 1024)') From 01c8a372e6e422cd7b8e3b23c410a7983c357185 Mon Sep 17 00:00:00 2001 From: filipstrand Date: Tue, 20 Aug 2024 21:50:37 +0200 Subject: [PATCH 20/20] Update readme with dev model description --- README.md | 14 ++++++++++---- 1 file changed, 10 insertions(+), 4 deletions(-) diff --git a/README.md b/README.md index bed2eb2..6f84615 100644 --- a/README.md +++ b/README.md @@ -37,19 +37,27 @@ like [Numpy](https://numpy.org) and [Pillow](https://pypi.org/project/pillow/) f ``` ### Generating an image -Run the provided [main.py](main.py) by specifying a prompt and some optional arguments like so: +Run the provided [main.py](main.py) by specifying a prompt and some optional arguments like so using the default `Schnell` model: ``` python main.py --prompt "Luxury food photograph" --steps 2 --seed 2 ``` +or use the slower, but more powerful `Dev` model and run it with more time steps: + +``` +python main.py --model dev --prompt "Luxury food photograph" --steps 25 --seed 2 +``` + +⚠️ *If the specific model is not already downloaded on your machine, it will start the download process and fetch the model weights (~34GB in size for the Schnell or Dev model respectively).* ⚠️ + #### Full list of Command-Line Arguments - **`--prompt`** (required, `str`): Text description of the image to generate. - **`--output`** (optional, `str`, default: `"image.png"`): Output image filename. -- **`--model`** (optional, `str`, default: `"schnell"`): Model to use for generation. +- **`--model`** (optional, `str`, default: `"schnell"`): Model to use for generation (`"schnell"` or `"dev"`). - **`--seed`** (optional, `int`, default: `0`): Seed for random number generation. Default is time-based. @@ -85,8 +93,6 @@ image = flux.generate_image( ImageUtil.save_image(image, "image.png") ``` -If the model is not already downloaded on your machine, it will start the download process and fetch the model weights (~34GB in size for the Schnell model). - ### Image generation speed (updated)