From 14c5b2bd0e354a22a3dd8930e519f655e763004a Mon Sep 17 00:00:00 2001 From: Fabio Date: Sun, 18 Aug 2024 01:01:01 +0200 Subject: [PATCH] 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: