diff --git a/src/mflux/controlnet/flux_controlnet.py b/src/mflux/controlnet/flux_controlnet.py index 5d38a11..043515d 100644 --- a/src/mflux/controlnet/flux_controlnet.py +++ b/src/mflux/controlnet/flux_controlnet.py @@ -11,6 +11,7 @@ from mflux.controlnet.transformer_controlnet import TransformerControlnet from mflux.flux.flux_initializer import FluxInitializer from mflux.latent_creator.latent_creator import LatentCreator from mflux.models.text_encoder.clip_encoder.clip_encoder import CLIPEncoder +from mflux.models.text_encoder.prompt_encoder import PromptEncoder from mflux.models.text_encoder.t5_encoder.t5_encoder import T5Encoder from mflux.models.transformer.transformer import Transformer from mflux.models.vae.vae import VAE @@ -73,10 +74,14 @@ class Flux1Controlnet(nn.Module): ) # fmt: off # 3. Encode the prompt - t5_tokens = self.t5_tokenizer.tokenize(prompt) - clip_tokens = self.clip_tokenizer.tokenize(prompt) - prompt_embeds = self.t5_text_encoder(t5_tokens) - pooled_prompt_embeds = self.clip_text_encoder(clip_tokens) + prompt_embeds, pooled_prompt_embeds = PromptEncoder.encode_prompt( + prompt=prompt, + prompt_cache=self.prompt_cache, + t5_tokenizer=self.t5_tokenizer, + clip_tokenizer=self.clip_tokenizer, + t5_text_encoder=self.t5_text_encoder, + clip_text_encoder=self.clip_text_encoder, + ) # (Optional) Call subscribers for beginning of loop Callbacks.before_loop( diff --git a/src/mflux/flux/flux.py b/src/mflux/flux/flux.py index 178b68e..c5befb2 100644 --- a/src/mflux/flux/flux.py +++ b/src/mflux/flux/flux.py @@ -9,6 +9,7 @@ from mflux.config.runtime_config import RuntimeConfig from mflux.flux.flux_initializer import FluxInitializer from mflux.latent_creator.latent_creator import Img2Img, LatentCreator from mflux.models.text_encoder.clip_encoder.clip_encoder import CLIPEncoder +from mflux.models.text_encoder.prompt_encoder import PromptEncoder from mflux.models.text_encoder.t5_encoder.t5_encoder import T5Encoder from mflux.models.transformer.transformer import Transformer from mflux.models.vae.vae import VAE @@ -66,10 +67,14 @@ class Flux1(nn.Module): ) # 2. Encode the prompt - t5_tokens = self.t5_tokenizer.tokenize(prompt) - clip_tokens = self.clip_tokenizer.tokenize(prompt) - prompt_embeds = self.t5_text_encoder(t5_tokens) - pooled_prompt_embeds = self.clip_text_encoder(clip_tokens) + prompt_embeds, pooled_prompt_embeds = PromptEncoder.encode_prompt( + prompt=prompt, + prompt_cache=self.prompt_cache, + t5_tokenizer=self.t5_tokenizer, + clip_tokenizer=self.clip_tokenizer, + t5_text_encoder=self.t5_text_encoder, + clip_text_encoder=self.clip_text_encoder, + ) # (Optional) Call subscribers for beginning of loop Callbacks.before_loop( diff --git a/src/mflux/flux/flux_initializer.py b/src/mflux/flux/flux_initializer.py index 2e57ab1..75a7c6e 100644 --- a/src/mflux/flux/flux_initializer.py +++ b/src/mflux/flux/flux_initializer.py @@ -22,7 +22,8 @@ class FluxInitializer: lora_paths: list[str] | None, lora_scales: list[float] | None, ) -> None: - # 0. Set paths and config for later + # 0. Set paths, configs and prompt_cache for later + flux_model.prompt_cache = {} flux_model.lora_paths = lora_paths flux_model.lora_scales = lora_scales flux_model.model_config = model_config diff --git a/src/mflux/models/text_encoder/prompt_encoder.py b/src/mflux/models/text_encoder/prompt_encoder.py new file mode 100644 index 0000000..48fa189 --- /dev/null +++ b/src/mflux/models/text_encoder/prompt_encoder.py @@ -0,0 +1,33 @@ +import mlx.core as mx + +from mflux.models.text_encoder.clip_encoder.clip_encoder import CLIPEncoder +from mflux.models.text_encoder.t5_encoder.t5_encoder import T5Encoder +from mflux.tokenizer.clip_tokenizer import TokenizerCLIP +from mflux.tokenizer.t5_tokenizer import TokenizerT5 + + +class PromptEncoder: + @staticmethod + def encode_prompt( + prompt: str, + prompt_cache: dict[str, (mx.array, mx.array)], + t5_tokenizer: TokenizerT5, + clip_tokenizer: TokenizerCLIP, + t5_text_encoder: T5Encoder, + clip_text_encoder: CLIPEncoder, + ) -> (mx.array, mx.array): + # 1. Return prompt encodings if already cached + if prompt in prompt_cache: + return prompt_cache[prompt] + + # 1. Encode the prompt + t5_tokens = t5_tokenizer.tokenize(prompt) + clip_tokens = clip_tokenizer.tokenize(prompt) + prompt_embeds = t5_text_encoder(t5_tokens) + pooled_prompt_embeds = clip_text_encoder(clip_tokens) + + # 2. Cache the encoded prompt + prompt_cache[prompt] = (prompt_embeds, pooled_prompt_embeds) + + # 3. Return prompt encodings + return prompt_embeds, pooled_prompt_embeds