Add PromptCache: Small optimization for repeated generations with same prompt
This commit is contained in:
parent
2d481e77b2
commit
09d75e46b8
@ -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(
|
||||
|
||||
@ -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(
|
||||
|
||||
@ -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
|
||||
|
||||
33
src/mflux/models/text_encoder/prompt_encoder.py
Normal file
33
src/mflux/models/text_encoder/prompt_encoder.py
Normal 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
|
||||
Loading…
Reference in New Issue
Block a user