[diffusion] fix: fix accuracy for flux series (#22059)
Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
@@ -10,6 +10,11 @@ from sglang.multimodal_gen.configs.models.encoders.clip import (
|
||||
CLIPTextConfig,
|
||||
CLIPVisionConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.models.encoders.flux_2 import (
|
||||
FLUX_2_SYSTEM_MESSAGE,
|
||||
Flux2MistralTextConfig,
|
||||
build_flux2_text_messages,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.models.encoders.gemma2 import Gemma2Config
|
||||
from sglang.multimodal_gen.configs.models.encoders.gemma_3 import Gemma3Config
|
||||
from sglang.multimodal_gen.configs.models.encoders.llama import LlamaConfig
|
||||
@@ -23,6 +28,9 @@ __all__ = [
|
||||
"BaseEncoderOutput",
|
||||
"CLIPTextConfig",
|
||||
"CLIPVisionConfig",
|
||||
"FLUX_2_SYSTEM_MESSAGE",
|
||||
"Flux2MistralTextConfig",
|
||||
"build_flux2_text_messages",
|
||||
"LlamaConfig",
|
||||
"Qwen3TextConfig",
|
||||
"T5Config",
|
||||
|
||||
@@ -0,0 +1,63 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""FLUX.2 Mistral text encoder configuration and prompt formatting."""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from sglang.multimodal_gen.configs.models.encoders.base import (
|
||||
TextEncoderArchConfig,
|
||||
TextEncoderConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.models.encoders.qwen_image import (
|
||||
_is_transformer_layer,
|
||||
)
|
||||
|
||||
FLUX_2_SYSTEM_MESSAGE = (
|
||||
"You are an AI that reasons about image descriptions. You give structured responses focusing on object relationships, object\n"
|
||||
"attribution and actions without speculation."
|
||||
)
|
||||
|
||||
|
||||
def build_flux2_text_messages(prompts: list[str]) -> list[list[dict]]:
|
||||
cleaned_prompts = [prompt.replace("[IMG]", "") for prompt in prompts]
|
||||
return [
|
||||
[
|
||||
{
|
||||
"role": "system",
|
||||
"content": [{"type": "text", "text": FLUX_2_SYSTEM_MESSAGE}],
|
||||
},
|
||||
{"role": "user", "content": [{"type": "text", "text": prompt}]},
|
||||
]
|
||||
for prompt in cleaned_prompts
|
||||
]
|
||||
|
||||
|
||||
@dataclass
|
||||
class Flux2MistralTextArchConfig(TextEncoderArchConfig):
|
||||
stacked_params_mapping: list[tuple[str, str, str]] = field(
|
||||
default_factory=lambda: [
|
||||
("qkv_proj", "q_proj", "q"),
|
||||
("qkv_proj", "k_proj", "k"),
|
||||
("qkv_proj", "v_proj", "v"),
|
||||
]
|
||||
)
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda: [_is_transformer_layer]
|
||||
)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.tokenizer_kwargs = {
|
||||
"padding": "max_length",
|
||||
"truncation": True,
|
||||
"max_length": 512,
|
||||
"add_special_tokens": True,
|
||||
"return_attention_mask": True,
|
||||
"return_tensors": "pt",
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class Flux2MistralTextConfig(TextEncoderConfig):
|
||||
arch_config: TextEncoderArchConfig = field(
|
||||
default_factory=Flux2MistralTextArchConfig
|
||||
)
|
||||
prefix: str = "flux_2_mistral"
|
||||
@@ -11,14 +11,11 @@ from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig
|
||||
from sglang.multimodal_gen.configs.models.encoders import (
|
||||
BaseEncoderOutput,
|
||||
CLIPTextConfig,
|
||||
Flux2MistralTextConfig,
|
||||
T5Config,
|
||||
TextEncoderConfig,
|
||||
build_flux2_text_messages,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.models.encoders.base import TextEncoderArchConfig
|
||||
from sglang.multimodal_gen.configs.models.encoders.qwen3 import Qwen3TextConfig
|
||||
from sglang.multimodal_gen.configs.models.encoders.qwen_image import (
|
||||
_is_transformer_layer,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.models.vaes.flux import Flux2VAEConfig, FluxVAEConfig
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.base import (
|
||||
ImagePipelineConfig,
|
||||
@@ -353,61 +350,6 @@ def flux2_klein_postprocess_text(
|
||||
return prompt_embeds
|
||||
|
||||
|
||||
@dataclass
|
||||
class Flux2MistralTextArchConfig(TextEncoderArchConfig):
|
||||
stacked_params_mapping: list[tuple[str, str, str]] = field(
|
||||
default_factory=lambda: [
|
||||
# (param_name, shard_name, shard_id)
|
||||
("qkv_proj", "q_proj", "q"),
|
||||
("qkv_proj", "k_proj", "k"),
|
||||
("qkv_proj", "v_proj", "v"),
|
||||
]
|
||||
)
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda: [_is_transformer_layer]
|
||||
)
|
||||
|
||||
def __post_init__(self):
|
||||
self.tokenizer_kwargs = {
|
||||
"padding": "max_length",
|
||||
"truncation": True,
|
||||
"max_length": 512,
|
||||
"add_special_tokens": True,
|
||||
"return_attention_mask": True,
|
||||
"return_tensors": "pt",
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class Flux2MistralTextConfig(TextEncoderConfig):
|
||||
arch_config: TextEncoderArchConfig = field(
|
||||
default_factory=Flux2MistralTextArchConfig
|
||||
)
|
||||
|
||||
|
||||
def format_text_input(prompts: List[str], system_message: str = None):
|
||||
# Remove [IMG] tokens from prompts to avoid Pixtral validation issues
|
||||
# when truncation is enabled. The processor counts [IMG] tokens and fails
|
||||
# if the count changes after truncation.
|
||||
cleaned_txt = [prompt.replace("[IMG]", "") for prompt in prompts]
|
||||
|
||||
return [
|
||||
[
|
||||
{
|
||||
"role": "system",
|
||||
"content": [{"type": "text", "text": system_message}],
|
||||
},
|
||||
{"role": "user", "content": [{"type": "text", "text": prompt}]},
|
||||
]
|
||||
for prompt in cleaned_txt
|
||||
]
|
||||
|
||||
|
||||
def flux_2_preprocess_text(prompt: str):
|
||||
system_message = "You are an AI that reasons about image descriptions. You give structured responses focusing on object relationships, object attribution and actions without speculation."
|
||||
return format_text_input([prompt], system_message=system_message)
|
||||
|
||||
|
||||
# Copied from diffusers.pipelines.qwenimage.pipeline_qwenimage.QwenImagePipeline._pack_latents
|
||||
def flux2_pack_latents(latents):
|
||||
batch_size, num_channels, height, width = latents.shape
|
||||
@@ -428,7 +370,7 @@ class Flux2PipelineConfig(FluxPipelineConfig):
|
||||
default_factory=lambda: (Flux2MistralTextConfig(),)
|
||||
)
|
||||
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
|
||||
default_factory=lambda: (flux_2_preprocess_text,),
|
||||
default_factory=lambda: (None,),
|
||||
)
|
||||
|
||||
postprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
|
||||
@@ -448,10 +390,9 @@ class Flux2PipelineConfig(FluxPipelineConfig):
|
||||
)
|
||||
|
||||
def tokenize_prompt(self, prompts: list[str], tokenizer, tok_kwargs) -> dict:
|
||||
# flatten to 1-d list
|
||||
prompts = [p for prompt in prompts for p in prompt]
|
||||
messages = build_flux2_text_messages(prompts)
|
||||
inputs = tokenizer.apply_chat_template(
|
||||
prompts,
|
||||
messages,
|
||||
add_generation_prompt=False,
|
||||
tokenize=True,
|
||||
return_dict=True,
|
||||
|
||||
@@ -14,12 +14,18 @@ class FluxSamplingParams(SamplingParams):
|
||||
|
||||
num_frames: int = 1
|
||||
# Denoising stage
|
||||
guidance_scale: float = 1.0
|
||||
guidance_scale: float = 3.5
|
||||
negative_prompt: str = None
|
||||
num_inference_steps: int = 50
|
||||
|
||||
|
||||
@dataclass
|
||||
class Flux2KleinSamplingParams(FluxSamplingParams):
|
||||
class Flux2SamplingParams(FluxSamplingParams):
|
||||
guidance_scale: float = 4.0
|
||||
|
||||
|
||||
@dataclass
|
||||
class Flux2KleinSamplingParams(Flux2SamplingParams):
|
||||
# Klein is step-distilled, so default to 4 steps
|
||||
guidance_scale: float = 1.0
|
||||
num_inference_steps: int = 4
|
||||
|
||||
@@ -76,6 +76,7 @@ from sglang.multimodal_gen.configs.pipeline_configs.wan import (
|
||||
)
|
||||
from sglang.multimodal_gen.configs.sample.flux import (
|
||||
Flux2KleinSamplingParams,
|
||||
Flux2SamplingParams,
|
||||
FluxSamplingParams,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.sample.glmimage import GlmImageSamplingParams
|
||||
@@ -747,7 +748,7 @@ def _register_configs():
|
||||
],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=FluxSamplingParams,
|
||||
sampling_param_cls=Flux2SamplingParams,
|
||||
pipeline_config_cls=Flux2PipelineConfig,
|
||||
hf_model_paths=[
|
||||
"black-forest-labs/FLUX.2-dev",
|
||||
|
||||
@@ -298,7 +298,7 @@ class TokenizerLoader(ComponentLoader):
|
||||
# Flux.2 aligns to the tokenizer defaults from the original baseline.
|
||||
# TODO: abstract this
|
||||
if isinstance(server_args.pipeline_config, Flux2PipelineConfig):
|
||||
return AutoTokenizer.from_pretrained(component_model_path)
|
||||
return AutoProcessor.from_pretrained(component_model_path)
|
||||
|
||||
return AutoTokenizer.from_pretrained(
|
||||
component_model_path,
|
||||
|
||||
@@ -1010,7 +1010,7 @@ class Flux2Transformer2DModel(CachableDiT, OffloadableDiTMixin):
|
||||
# 1. Calculate timestep embedding and modulation parameters
|
||||
timestep = timestep.to(hidden_states.dtype)
|
||||
if guidance is not None:
|
||||
guidance = guidance.to(hidden_states.dtype)
|
||||
guidance = guidance.to(hidden_states.dtype) * 1000
|
||||
|
||||
temb = self.time_guidance_embed(timestep, guidance)
|
||||
|
||||
|
||||
@@ -269,8 +269,10 @@ class MistralModel(nn.Module):
|
||||
hidden_states = inputs_embeds
|
||||
position_embeddings = self.rotary_emb(hidden_states, position_ids)
|
||||
|
||||
hidden_states_pool = []
|
||||
hidden_states_pool = [] if output_hidden_states else None
|
||||
for decoder_layer in self.layers[: self.config.num_hidden_layers]:
|
||||
if output_hidden_states:
|
||||
hidden_states_pool.append(hidden_states)
|
||||
hidden_states = decoder_layer(
|
||||
hidden_states,
|
||||
attention_mask=causal_mask,
|
||||
@@ -281,8 +283,6 @@ class MistralModel(nn.Module):
|
||||
position_embeddings=position_embeddings,
|
||||
**kwargs,
|
||||
)
|
||||
if output_hidden_states:
|
||||
hidden_states_pool.append(hidden_states)
|
||||
|
||||
hidden_states = self.norm(hidden_states)
|
||||
if output_hidden_states:
|
||||
|
||||
@@ -21,6 +21,10 @@ from tqdm.auto import tqdm
|
||||
|
||||
from sglang.multimodal_gen import envs
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType, STA_Mode
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.flux import (
|
||||
Flux2PipelineConfig,
|
||||
FluxPipelineConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.wan import (
|
||||
Wan2_2_TI2V_5B_Config,
|
||||
)
|
||||
@@ -353,14 +357,15 @@ class DenoisingStage(PipelineStage):
|
||||
@lru_cache(maxsize=8)
|
||||
def _build_guidance(self, batch_size, target_dtype, device, guidance_val):
|
||||
"""Builds a guidance tensor. This method is cached."""
|
||||
return (
|
||||
torch.full(
|
||||
(batch_size,),
|
||||
guidance_val,
|
||||
dtype=target_dtype,
|
||||
device=device,
|
||||
)
|
||||
* 1000.0
|
||||
if isinstance(
|
||||
self.server_args.pipeline_config, FluxPipelineConfig
|
||||
) and not isinstance(self.server_args.pipeline_config, Flux2PipelineConfig):
|
||||
guidance_val = guidance_val * 1000.0
|
||||
return torch.full(
|
||||
(batch_size,),
|
||||
guidance_val,
|
||||
dtype=target_dtype,
|
||||
device=device,
|
||||
)
|
||||
|
||||
def get_or_build_guidance(self, bsz: int, dtype, device):
|
||||
|
||||
@@ -259,17 +259,14 @@ class TextEncodingStage(PipelineStage):
|
||||
is_flux_v1 = isinstance(
|
||||
server_args.pipeline_config, FluxPipelineConfig
|
||||
) and not isinstance(server_args.pipeline_config, Flux2PipelineConfig)
|
||||
is_flux_t5 = is_flux_v1 and i == 1
|
||||
|
||||
if is_flux_t5:
|
||||
attention_mask = torch.ones(input_ids.shape[:2], device=target_device)
|
||||
else:
|
||||
attention_mask = text_inputs["attention_mask"]
|
||||
attention_mask = None if is_flux_v1 else text_inputs["attention_mask"]
|
||||
encoder_forward_kwargs = {
|
||||
"input_ids": input_ids,
|
||||
"attention_mask": attention_mask,
|
||||
"output_hidden_states": True,
|
||||
}
|
||||
if attention_mask is not None:
|
||||
encoder_forward_kwargs["attention_mask"] = attention_mask
|
||||
if "use_cache" in inspect.signature(text_encoder.forward).parameters:
|
||||
encoder_forward_kwargs["use_cache"] = False
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
@@ -288,7 +285,12 @@ class TextEncodingStage(PipelineStage):
|
||||
if is_flux_v1:
|
||||
pooled_embeds_list.append(outputs.pooler_output)
|
||||
if return_attention_mask:
|
||||
attn_masks_list.append(attention_mask)
|
||||
mask_to_store = (
|
||||
attention_mask
|
||||
if attention_mask is not None
|
||||
else torch.ones(input_ids.shape[:2], device=target_device)
|
||||
)
|
||||
attn_masks_list.append(mask_to_store)
|
||||
|
||||
# Shape results according to return_type
|
||||
if return_type == "list":
|
||||
|
||||
@@ -5,7 +5,11 @@ import unittest
|
||||
from sglang.multimodal_gen.configs.sample.diffusers_generic import (
|
||||
DiffusersGenericSamplingParams,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.sample.flux import FluxSamplingParams
|
||||
from sglang.multimodal_gen.configs.sample.flux import (
|
||||
Flux2KleinSamplingParams,
|
||||
Flux2SamplingParams,
|
||||
FluxSamplingParams,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.sample.qwenimage import QwenImageSamplingParams
|
||||
from sglang.multimodal_gen.configs.sample.sampling_params import (
|
||||
SamplingParams,
|
||||
@@ -74,6 +78,11 @@ class TestSamplingParamsSubclass(unittest.TestCase):
|
||||
self.assertEqual(params.height, 640)
|
||||
self.assertEqual(params.width, 768)
|
||||
|
||||
def test_flux_guidance_defaults_match_model_defaults(self):
|
||||
self.assertEqual(FluxSamplingParams().guidance_scale, 3.5)
|
||||
self.assertEqual(Flux2SamplingParams().guidance_scale, 4.0)
|
||||
self.assertEqual(Flux2KleinSamplingParams().guidance_scale, 1.0)
|
||||
|
||||
def test_diffusers_generic_calls_base_post_init(self):
|
||||
with self.assertRaises(AssertionError):
|
||||
DiffusersGenericSamplingParams(num_frames=0)
|
||||
|
||||
Reference in New Issue
Block a user