Qwen-Image-Layered-MRP-MLX/src/mflux/models/text_encoder/prompt_encoder.py

34 lines
1.1 KiB
Python

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