From 627e162335627be46933847c498a71398df7881d Mon Sep 17 00:00:00 2001 From: Aditya Sharma <89210949+adityavaid@users.noreply.github.com> Date: Sat, 28 Mar 2026 14:58:02 +0530 Subject: [PATCH] [diffusion] fix: fix Flux2-Klein prompt tokenization length to 512 and add regression coverage (#21407) --- .../configs/pipeline_configs/flux.py | 15 ++++++++++++++- 1 file changed, 14 insertions(+), 1 deletion(-) diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/flux.py b/python/sglang/multimodal_gen/configs/pipeline_configs/flux.py index 9130dfc61..b24822154 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/flux.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/flux.py @@ -435,6 +435,17 @@ class Flux2PipelineConfig(FluxPipelineConfig): default_factory=lambda: (flux2_postprocess_text,) ) vae_config: VAEConfig = field(default_factory=Flux2VAEConfig) + text_encoder_extra_args: list[dict] = field( + default_factory=lambda: [ + dict( + max_length=512, + padding="max_length", + truncation=True, + return_overflowing_tokens=False, + return_length=False, + ) + ] + ) def tokenize_prompt(self, prompts: list[str], tokenizer, tok_kwargs) -> dict: # flatten to 1-d list @@ -679,7 +690,9 @@ class Flux2KleinPipelineConfig(Flux2PipelineConfig): texts = [_apply_chat_template(prompt) for prompt in prompts] tok_kwargs = dict(tok_kwargs or {}) - max_length = tok_kwargs.pop("max_length", 512) + tok_kwargs.pop("max_length", None) + # Flux2 Klein uses max_length 512. + max_length = 512 padding = tok_kwargs.pop("padding", "max_length") truncation = tok_kwargs.pop("truncation", True) return_tensors = tok_kwargs.pop("return_tensors", "pt")