Add PromptCache: Small optimization for repeated generations with same prompt

This commit is contained in:
filipstrand 2025-02-09 20:38:08 +01:00
parent 2d481e77b2
commit 09d75e46b8
4 changed files with 53 additions and 9 deletions

View File

@ -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(

View File

@ -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(

View File

@ -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

View File

@ -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