[diffusion] fix: fix flux2 i2i accuracy (#22423)

This commit is contained in:
Mick
2026-04-10 16:16:51 +08:00
committed by GitHub
parent 6cf7f210bf
commit 7c6b5c095c
16 changed files with 402 additions and 116 deletions
@@ -352,6 +352,10 @@ class PipelineConfig:
def postprocess_vae_encode(self, image_latents, vae):
return image_latents
# called after postprocess_vae_encode, before generic scale/shift
def normalize_vae_encode(self, image_latents, vae):
return None
# called after scale_and_shift, before vae decoding
def preprocess_decoding(self, latents, server_args=None, vae=None):
return latents
@@ -364,6 +364,8 @@ class Flux2PipelineConfig(FluxPipelineConfig):
task_type: ModelTaskType = ModelTaskType.TI2I
vae_precision: str = "bf16"
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16",))
text_encoder_configs: tuple[EncoderConfig, ...] = field(
@@ -446,7 +448,22 @@ class Flux2PipelineConfig(FluxPipelineConfig):
def preprocess_condition_image(
self, image, target_width, target_height, vae_image_processor: VaeImageProcessor
):
img = image.resize((target_width, target_height), PIL.Image.Resampling.LANCZOS)
target_area = 1024 * 1024
img = image
if image.width * image.height > target_area:
resize_to_target_area = getattr(
vae_image_processor, "_resize_to_target_area", None
)
if callable(resize_to_target_area):
img = resize_to_target_area(image, target_area)
else:
scale = math.sqrt(target_area / (image.width * image.height))
resized_width = int(image.width * scale)
resized_height = int(image.height * scale)
img = image.resize(
(resized_width, resized_height), PIL.Image.Resampling.LANCZOS
)
image_width, image_height = img.size
vae_scale_factor = self.vae_config.arch_config.vae_scale_factor
multiple_of = vae_scale_factor * 2
@@ -531,6 +548,19 @@ class Flux2PipelineConfig(FluxPipelineConfig):
image_latents = _patchify_latents(image_latents)
return image_latents
def normalize_vae_encode(self, image_latents, vae):
if not self._check_vae_has_bn(vae):
return None
latents_bn_mean = vae.bn.running_mean.view(1, -1, 1, 1).to(
image_latents.device, image_latents.dtype
)
latents_bn_std = torch.sqrt(
vae.bn.running_var.view(1, -1, 1, 1)
+ self.vae_config.arch_config.batch_norm_eps
).to(image_latents.device, image_latents.dtype)
return (image_latents - latents_bn_mean) / latents_bn_std
def _check_vae_has_bn(self, vae):
"""Check if VAE has bn attribute (cached check to avoid repeated hasattr calls)."""
if not hasattr(self, "_vae_has_bn_cache"):
@@ -267,6 +267,9 @@ class SamplingParams:
diffusers_kwargs = getattr(self, "diffusers_kwargs", None)
if diffusers_kwargs:
extra["diffusers_kwargs"] = diffusers_kwargs
explicit_fields = getattr(self, "_explicit_fields", None)
if explicit_fields is not None:
extra["explicit_fields"] = sorted(explicit_fields)
return extra
def apply_request_extra(self, req: Any) -> None:
@@ -608,6 +611,7 @@ class SamplingParams:
sampling_params._merge_with_user_params(
user_sampling_params, explicit_fields=set(user_kwargs.keys())
)
sampling_params._explicit_fields = set(user_kwargs.keys())
sampling_params._adjust(server_args)
sampling_params._validate_with_pipeline_config(server_args.pipeline_config)
@@ -47,6 +47,13 @@ def add_multimodal_gen_generate_args(parser: argparse.ArgumentParser):
required=False,
help="Path to dump the performance metrics (JSON) for the run.",
)
parser.add_argument(
"--output-file-path",
type=str,
default=None,
required=False,
help="Convenience alias that sets both --output-path and --output-file-name.",
)
parser = ServerArgs.add_cli_args(parser)
parser = SamplingParams.add_cli_args(parser)
@@ -60,6 +67,18 @@ def add_multimodal_gen_generate_args(parser: argparse.ArgumentParser):
return parser
def _apply_output_file_path_override(
args: argparse.Namespace, sampling_params_kwargs: dict
):
output_file_path = args.output_file_path
if not output_file_path:
return
output_path = os.path.dirname(output_file_path) or "."
sampling_params_kwargs["output_path"] = output_path
sampling_params_kwargs["output_file_name"] = os.path.basename(output_file_path)
def maybe_dump_performance(
args: argparse.Namespace,
server_args,
@@ -129,6 +148,7 @@ def generate_cmd(args: argparse.Namespace, unknown_args: list[str] | None = None
)
sampling_params_kwargs.update(SamplingParams.get_cli_args(args))
_apply_output_file_path_override(args, sampling_params_kwargs)
sampling_params_kwargs["request_id"] = generate_request_id()
# Handle diffusers-specific kwargs passed via CLI
@@ -15,7 +15,6 @@ from torch import nn
from transformers import AutoImageProcessor, AutoProcessor, AutoTokenizer
from sglang.multimodal_gen.configs.models import ModelConfig
from sglang.multimodal_gen.configs.pipeline_configs.flux import Flux2PipelineConfig
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.loader.utils import (
_normalize_component_type,
@@ -51,6 +50,7 @@ class ComponentLoader(ABC):
def __init__(self, device=None) -> None:
self.device = device
self.component_architecture: str | None = None
def should_offload(
self, server_args: ServerArgs, model_config: ModelConfig | None = None
@@ -208,21 +208,9 @@ class ComponentLoader(ABC):
cls._loaders_registered = True
@classmethod
def for_component_type(
cls, component_name: str, transformers_or_diffusers: str
) -> "ComponentLoader":
"""
Factory method to create a component loader for a specific component type.
Args:
component_name: Type of component (e.g., "vae", "text_encoder", "transformer", "scheduler")
transformers_or_diffusers: Whether the component is from transformers or diffusers
"""
cls._ensure_loaders_registered()
# Map of component types to their loader classes and expected library
component_name = _normalize_component_type(component_name)
def resolve_transformers_or_diffusers(
self, transformers_or_diffusers: str, component_name: str
) -> str:
# NOTE(FlamingoPg): special for LTX-2 models
if component_name == "vocoder" or component_name == "connectors":
transformers_or_diffusers = "diffusers"
@@ -243,6 +231,31 @@ class ComponentLoader(ABC):
):
transformers_or_diffusers = "diffusers"
return transformers_or_diffusers
@classmethod
def for_component_type(
cls,
component_name: str,
transformers_or_diffusers: str,
component_architecture: str | None = None,
) -> "ComponentLoader":
"""
Factory method to create a component loader for a specific component type.
Args:
component_name: Type of component (e.g., "vae", "text_encoder", "transformer", "scheduler")
transformers_or_diffusers: Whether the component is from transformers or diffusers
"""
cls._ensure_loaders_registered()
# Map of component types to their loader classes and expected library
component_name = _normalize_component_type(component_name)
transformers_or_diffusers = cls.resolve_transformers_or_diffusers(
transformers_or_diffusers, component_name
)
if component_name in component_name_to_loader_cls:
loader_cls: Type[ComponentLoader] = component_name_to_loader_cls[
component_name
@@ -252,14 +265,16 @@ class ComponentLoader(ABC):
assert (
transformers_or_diffusers == expected_library
), f"{component_name} must be loaded from {expected_library}, got {transformers_or_diffusers}"
return loader_cls()
loader = loader_cls()
loader.component_architecture = component_architecture
return loader
# For unknown component types, use a generic loader
logger.warning(
"No specific loader found for component type: %s. Using generic loader.",
component_name,
)
return GenericComponentLoader(transformers_or_diffusers)
return GenericComponentLoader(transformers_or_diffusers, component_architecture)
class ImageProcessorLoader(ComponentLoader):
@@ -295,9 +310,14 @@ class TokenizerLoader(ComponentLoader):
def load_customized(
self, component_model_path: str, server_args: ServerArgs, component_name: str
) -> Any:
# Flux.2 aligns to the tokenizer defaults from the original baseline.
# TODO: abstract this
if isinstance(server_args.pipeline_config, Flux2PipelineConfig):
# Some pipelines keep the slot name `tokenizer` in model_index.json even
# when the declared class is a processor. e.g. FLUX.2:
# `tokenizer: ["transformers", "PixtralProcessor"]`.
# Honor the declared component class instead of guessing from the slot name.
if (
self.component_architecture is not None
and self.component_architecture.endswith("Processor")
):
return AutoProcessor.from_pretrained(component_model_path)
return AutoTokenizer.from_pretrained(
@@ -310,9 +330,12 @@ class TokenizerLoader(ComponentLoader):
class GenericComponentLoader(ComponentLoader):
"""Generic loader for components that don't have a specific loader."""
def __init__(self, library="transformers") -> None:
def __init__(
self, library="transformers", component_architecture: str | None = None
) -> None:
super().__init__()
self.library = library
self.component_architecture = component_architecture
class PipelineComponentLoader:
@@ -326,6 +349,7 @@ class PipelineComponentLoader:
component_model_path: str,
transformers_or_diffusers: str,
server_args: ServerArgs,
component_architecture: str | None = None,
):
"""
Load a pipeline component.
@@ -334,12 +358,12 @@ class PipelineComponentLoader:
component_name: Name of the component (e.g., "vae", "text_encoder", "transformer", "scheduler")
component_model_path: Path to the component model
transformers_or_diffusers: Whether the component is from transformers or diffusers
component_architecture: the class name of the module
"""
# Get the appropriate loader for this component type
loader = ComponentLoader.for_component_type(
component_name, transformers_or_diffusers
component_name, transformers_or_diffusers, component_architecture
)
try:
@@ -2,11 +2,10 @@ import dataclasses
import glob
import os
from collections.abc import Generator, Iterable
from typing import Generator, Iterable, cast
from typing import cast
import torch
import torch.distributed as dist
import torch.nn as nn
from torch import nn
from torch.distributed import init_device_mesh
from transformers import AutoModel
@@ -43,7 +43,10 @@ from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
apply_flashinfer_rope_qk_inplace,
)
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
from sglang.multimodal_gen.runtime.platforms import current_platform
from sglang.multimodal_gen.runtime.platforms import (
AttentionBackendEnum,
current_platform,
)
from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiTMixin
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
@@ -148,6 +151,7 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
eps: float = 1e-5,
out_dim: int = None,
elementwise_affine: bool = True,
supported_attention_backends: set[AttentionBackendEnum] | None = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
@@ -278,6 +282,7 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
dropout_rate=0,
softmax_scale=None,
causal=False,
supported_attention_backends=supported_attention_backends,
)
def forward(
@@ -400,6 +405,7 @@ class Flux2ParallelSelfAttention(torch.nn.Module, AttentionModuleMixin):
elementwise_affine: bool = True,
mlp_ratio: float = 4.0,
mlp_mult_factor: int = 2,
supported_attention_backends: set[AttentionBackendEnum] | None = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
@@ -459,6 +465,7 @@ class Flux2ParallelSelfAttention(torch.nn.Module, AttentionModuleMixin):
dropout_rate=0,
softmax_scale=None,
causal=False,
supported_attention_backends=supported_attention_backends,
)
def _patch_to_out_weight_loader(self) -> None:
@@ -545,6 +552,7 @@ class Flux2SingleTransformerBlock(nn.Module):
mlp_ratio: float = 3.0,
eps: float = 1e-6,
bias: bool = False,
supported_attention_backends: set[AttentionBackendEnum] | None = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
@@ -565,6 +573,7 @@ class Flux2SingleTransformerBlock(nn.Module):
eps=eps,
mlp_ratio=mlp_ratio,
mlp_mult_factor=2,
supported_attention_backends=supported_attention_backends,
quant_config=quant_config,
prefix=f"{prefix}.attn" if prefix else "attn",
)
@@ -621,6 +630,7 @@ class Flux2TransformerBlock(nn.Module):
mlp_ratio: float = 3.0,
eps: float = 1e-6,
bias: bool = False,
supported_attention_backends: set[AttentionBackendEnum] | None = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
@@ -640,6 +650,7 @@ class Flux2TransformerBlock(nn.Module):
added_proj_bias=bias,
out_bias=bias,
eps=eps,
supported_attention_backends=supported_attention_backends,
quant_config=quant_config,
prefix=f"{prefix}.attn" if prefix else "attn",
)
@@ -839,6 +850,12 @@ class Flux2Transformer2DModel(CachableDiT, OffloadableDiTMixin):
param_names_mapping = FluxConfig().arch_config.param_names_mapping
scale_shift_swap_params = ("norm_out.linear.weight", "norm_out.linear.bias")
# FLUX.2 stays closer to the official diffusers output with Torch SDPA.
# The generic FA path still produces a measurable image-level drift here.
_supported_attention_backends = {
AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.FA,
}
def post_load_weights(self) -> None:
if not isinstance(getattr(self, "quant_config", None), ModelOptFp4Config):
@@ -932,6 +949,7 @@ class Flux2Transformer2DModel(CachableDiT, OffloadableDiTMixin):
mlp_ratio=mlp_ratio,
eps=eps,
bias=False,
supported_attention_backends=self._supported_attention_backends,
quant_config=quant_config,
prefix=f"transformer_blocks.{i}",
)
@@ -949,6 +967,7 @@ class Flux2Transformer2DModel(CachableDiT, OffloadableDiTMixin):
mlp_ratio=mlp_ratio,
eps=eps,
bias=False,
supported_attention_backends=self._supported_attention_backends,
quant_config=quant_config,
prefix=f"single_transformer_blocks.{i}",
)
@@ -13,31 +13,42 @@
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import inspect
from contextlib import nullcontext
from typing import Iterable, Optional, Union
import torch
from torch import nn
from torch.nn.attention import SDPBackend, sdpa_kernel
from transformers import Cache, DynamicCache, LlavaConfig, Mistral3Config, MistralConfig
from transformers.integrations.sdpa_attention import sdpa_attention_forward
from transformers.masking_utils import create_causal_mask
from transformers.masking_utils import (
create_causal_mask,
create_sliding_window_causal_mask,
)
from transformers.modeling_outputs import BaseModelOutputWithPast
from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS
from transformers.models.mistral3.modeling_mistral3 import (
Mistral3CausalLMOutputWithPast,
Mistral3ModelOutputWithPast,
)
from transformers.models.mistral.modeling_mistral import (
MistralMLP,
MistralPreTrainedModel,
MistralRMSNorm,
MistralRotaryEmbedding,
apply_rotary_pos_emb,
eager_attention_forward,
)
from sglang.multimodal_gen.runtime.layers.attention import USPAttention
from sglang.multimodal_gen.runtime.loader.weight_utils import default_weight_loader
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
_CREATE_CAUSAL_MASK_ARG = (
"inputs_embeds"
if "inputs_embeds" in inspect.signature(create_causal_mask).parameters
else "input_embeds"
)
def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
@@ -62,10 +73,6 @@ class MistralAttention(nn.Module):
super().__init__()
self.config = config
self.layer_idx = layer_idx
self.num_key_value_groups = (
config.num_attention_heads // config.num_key_value_heads
)
self.head_dim = (
getattr(config, "head_dim", None)
or config.hidden_size // config.num_attention_heads
@@ -75,7 +82,6 @@ class MistralAttention(nn.Module):
)
self.scaling = self.head_dim**-0.5
self.attention_dropout = config.attention_dropout
self.is_causal = True
self.q_proj = nn.Linear(
config.hidden_size, config.num_attention_heads * self.head_dim, bias=False
)
@@ -91,17 +97,6 @@ class MistralAttention(nn.Module):
self.is_causal = True
self.num_heads = config.num_attention_heads
self.num_key_value_heads = config.num_key_value_heads
self.attn = USPAttention(
num_heads=self.num_heads,
head_size=self.head_dim,
dropout_rate=0,
softmax_scale=None,
causal=False,
supported_attention_backends={
AttentionBackendEnum.FA,
AttentionBackendEnum.TORCH_SDPA,
},
)
def forward(
self,
@@ -131,7 +126,15 @@ class MistralAttention(nn.Module):
key_states, value_states, self.layer_idx, cache_kwargs
)
attention_interface = sdpa_attention_forward
attn_implementation = getattr(self.config, "_attn_implementation", None)
attention_interface = eager_attention_forward
if attn_implementation and attn_implementation != "eager":
if hasattr(ALL_ATTENTION_FUNCTIONS, "get_interface"):
attention_interface = ALL_ATTENTION_FUNCTIONS.get_interface(
attn_implementation, eager_attention_forward
)
else:
attention_interface = ALL_ATTENTION_FUNCTIONS[attn_implementation]
attn_output, attn_weights = attention_interface(
self,
query_states,
@@ -148,7 +151,7 @@ class MistralAttention(nn.Module):
attn_output = attn_output.reshape(*input_shape, -1).contiguous()
attn_output = self.o_proj(attn_output)
return attn_output
return attn_output, attn_weights
class MistralDecoderLayer(nn.Module):
@@ -180,7 +183,7 @@ class MistralDecoderLayer(nn.Module):
residual = hidden_states
hidden_states = self.input_layernorm(hidden_states)
# Self Attention
hidden_states = self.self_attn(
hidden_states, _ = self.self_attn(
hidden_states=hidden_states,
attention_mask=attention_mask,
position_ids=position_ids,
@@ -200,10 +203,9 @@ class MistralDecoderLayer(nn.Module):
return hidden_states
class MistralModel(nn.Module):
class MistralModel(MistralPreTrainedModel):
def __init__(self, config: MistralConfig):
super().__init__()
self.config = config
super().__init__(config)
self.padding_idx = config.pad_token_id
self.vocab_size = config.vocab_size
@@ -219,7 +221,7 @@ class MistralModel(nn.Module):
self.norm = MistralRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.rotary_emb = MistralRotaryEmbedding(config=config)
self.gradient_checkpointing = False
self.config._attn_implementation = "sdpa"
self.post_init()
def forward(
self,
@@ -256,15 +258,20 @@ class MistralModel(nn.Module):
if position_ids is None:
position_ids = cache_position.unsqueeze(0)
mask_function = create_causal_mask
causal_mask = mask_function(
config=self.config,
inputs_embeds=inputs_embeds,
attention_mask=attention_mask,
cache_position=cache_position,
past_key_values=past_key_values,
position_ids=position_ids,
mask_function = (
create_causal_mask
if getattr(self.config, "sliding_window", None) is None
else create_sliding_window_causal_mask
)
mask_kwargs = {
"config": self.config,
_CREATE_CAUSAL_MASK_ARG: inputs_embeds,
"attention_mask": attention_mask,
"cache_position": cache_position,
"past_key_values": past_key_values,
"position_ids": position_ids,
}
causal_mask = mask_function(**mask_kwargs)
hidden_states = inputs_embeds
position_embeddings = self.rotary_emb(hidden_states, position_ids)
@@ -315,19 +322,21 @@ class Mistral3Model(nn.Module):
def forward(
self,
input_ids: Optional[torch.LongTensor] = None,
pixel_values: Optional[torch.FloatTensor] = None,
attention_mask: Optional[torch.Tensor] = None,
output_hidden_states: Optional[bool] = None,
position_ids: Optional[torch.LongTensor] = None,
past_key_values: Optional[Cache] = None,
inputs_embeds: Optional[torch.FloatTensor] = None,
vision_feature_layer: Optional[Union[int, list[int]]] = None,
use_cache: Optional[bool] = None,
output_attentions: Optional[bool] = None,
output_hidoutput_hidden_statesden_states: Optional[bool] = None,
output_hidden_states: Optional[bool] = None,
return_dict: Optional[bool] = None,
cache_position: Optional[torch.LongTensor] = None,
image_sizes: Optional[torch.Tensor] = None,
**kwargs,
) -> Union[tuple, Mistral3ModelOutputWithPast]:
del pixel_values, vision_feature_layer, return_dict
output_attentions = False
output_hidden_states = True
@@ -367,6 +376,7 @@ class Mistral3ForConditionalGeneration(nn.Module):
"^language_model.lm_head": "lm_head",
}
_tied_weights_keys = ["lm_head.weight"]
uses_sglang_forward_context = False
def __init__(self, config: LlavaConfig):
super().__init__()
@@ -413,19 +423,28 @@ class Mistral3ForConditionalGeneration(nn.Module):
"""
output_hidden_states = True
outputs = self.model(
input_ids=input_ids,
attention_mask=attention_mask,
position_ids=position_ids,
past_key_values=past_key_values,
inputs_embeds=inputs_embeds,
use_cache=use_cache,
output_hidden_states=output_hidden_states,
return_dict=True,
cache_position=cache_position,
image_sizes=image_sizes,
**kwargs,
execution_tensor = input_ids if input_ids is not None else inputs_embeds
sdpa_context = (
sdpa_kernel(SDPBackend.CUDNN_ATTENTION)
if execution_tensor is not None and execution_tensor.device.type == "cuda"
else nullcontext()
)
with sdpa_context:
# FLUX.2 uses the text-only Mistral3 path but still expects the
# same local SDPA kernel choice as the official HF implementation.
outputs = self.model(
input_ids=input_ids,
attention_mask=attention_mask,
position_ids=position_ids,
past_key_values=past_key_values,
inputs_embeds=inputs_embeds,
use_cache=use_cache,
output_hidden_states=output_hidden_states,
return_dict=True,
cache_position=cache_position,
image_sizes=image_sizes,
**kwargs,
)
return Mistral3CausalLMOutputWithPast(
hidden_states=outputs.hidden_states,
@@ -1,7 +1,7 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# SPDX-License-Identifier: Apache-2.0
from diffusers.image_processor import VaeImageProcessor
from diffusers.pipelines.flux2.image_processor import Flux2ImageProcessor
from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline, Req
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
@@ -45,7 +45,7 @@ class Flux2Pipeline(LoRAPipeline, ComposedPipelineBase):
]
def create_pipeline_stages(self, server_args: ServerArgs):
vae_image_processor = VaeImageProcessor(
vae_image_processor = Flux2ImageProcessor(
vae_scale_factor=server_args.pipeline_config.vae_config.arch_config.vae_scale_factor
* 2
)
@@ -304,6 +304,7 @@ class ComposedPipelineBase(ABC):
component_model_path=component_model_path,
transformers_or_diffusers=transformers_or_diffusers,
server_args=server_args,
component_architecture=architecture,
)
self.memory_usages[load_module_name] = memory_usage
@@ -333,24 +333,31 @@ class ImageVAEEncodingStage(PipelineStage):
latent_condition = server_args.pipeline_config.postprocess_vae_encode(
latent_condition, self.vae
)
scaling_factor, shift_factor = (
server_args.pipeline_config.get_decode_scale_and_shift(
device=latent_condition.device,
dtype=latent_condition.dtype,
vae=self.vae,
normalized_latent_condition = (
server_args.pipeline_config.normalize_vae_encode(
latent_condition, self.vae
)
)
if normalized_latent_condition is None:
scaling_factor, shift_factor = (
server_args.pipeline_config.get_decode_scale_and_shift(
device=latent_condition.device,
dtype=latent_condition.dtype,
vae=self.vae,
)
)
# apply shift & scale if needed
if isinstance(shift_factor, torch.Tensor):
shift_factor = shift_factor.to(latent_condition.device)
# apply shift & scale if needed
if isinstance(shift_factor, torch.Tensor):
shift_factor = shift_factor.to(latent_condition.device)
if isinstance(scaling_factor, torch.Tensor):
scaling_factor = scaling_factor.to(latent_condition.device)
if isinstance(scaling_factor, torch.Tensor):
scaling_factor = scaling_factor.to(latent_condition.device)
latent_condition -= shift_factor
latent_condition = latent_condition * scaling_factor
latent_condition -= shift_factor
latent_condition = latent_condition * scaling_factor
else:
latent_condition = normalized_latent_condition
if condition_latents is not None:
condition_latents.append(latent_condition)
@@ -134,8 +134,13 @@ class InputValidationStage(PipelineStage):
# adjust output image size
if calculated_size is not None:
calculated_width, calculated_height = calculated_size
width = batch.width or calculated_width
height = batch.height or calculated_height
explicit_fields = set(batch.extra.get("explicit_fields", []))
width_is_explicit = "width" in explicit_fields
height_is_explicit = "height" in explicit_fields
width = batch.width if width_is_explicit else calculated_width
height = batch.height if height_is_explicit else calculated_height
multiple_of = (
server_args.pipeline_config.vae_config.get_vae_scale_factor() * 2
)
@@ -134,6 +134,13 @@ class TextEncodingStage(PipelineStage):
return tok_kwargs
def _forward_text_encoder(self, text_encoder, encoder_forward_kwargs):
if not getattr(text_encoder, "uses_sglang_forward_context", True):
return text_encoder(**encoder_forward_kwargs)
with set_forward_context(current_timestep=0, attn_metadata=None):
return text_encoder(**encoder_forward_kwargs)
@torch.no_grad()
def encode_text(
self,
@@ -269,8 +276,9 @@ class TextEncodingStage(PipelineStage):
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):
outputs: BaseEncoderOutput = text_encoder(**encoder_forward_kwargs)
outputs: BaseEncoderOutput = self._forward_text_encoder(
text_encoder, encoder_forward_kwargs
)
postprocess_sig = inspect.signature(postprocess_func)
postprocess_kwargs = {}
@@ -279,14 +287,20 @@ class TextEncodingStage(PipelineStage):
postprocess_kwargs["pipeline_config"] = server_args.pipeline_config
prompt_embeds = postprocess_func(outputs, text_inputs, **postprocess_kwargs)
if dtype is not None:
prompt_embeds = prompt_embeds.to(dtype=dtype)
prompt_embeds = prompt_embeds.to(device=target_device, dtype=dtype)
else:
prompt_embeds = prompt_embeds.to(device=target_device)
embeds_list.append(prompt_embeds)
if is_flux_v1:
pooled_embeds_list.append(outputs.pooler_output)
if is_flux_v1 and outputs.pooler_output is not None:
# FLUX.1 only consumes the pooled CLIP projection. The T5
# encoder in the same pipeline has no pooler output.
pooled_embeds_list.append(
outputs.pooler_output.to(device=target_device)
)
if return_attention_mask:
mask_to_store = (
attention_mask
attention_mask.to(device=target_device)
if attention_mask is not None
else torch.ones(input_ids.shape[:2], device=target_device)
)
@@ -526,17 +526,15 @@ Repository: https://github.com/sglang-bot/sglang-ci-data (path: diffusion-ci/con
if not result.passed:
failed_frames = []
video_gt_info = ""
if is_video:
gt_remote_files = get_consistency_gt_remote_files(
case.id,
num_gpus,
is_video=True,
output_format=output_format,
)
video_gt_info = "\n".join(
f" - {filename}: {url}" for filename, url in gt_remote_files
)
gt_remote_files = get_consistency_gt_remote_files(
case.id,
num_gpus,
is_video=is_video,
output_format=output_format,
)
gt_remote_info = "\n".join(
f" - {filename}: {url}" for filename, url in gt_remote_files
)
for metric in result.frame_metrics:
failed_metrics = []
if not metric.clip_passed:
@@ -568,11 +566,7 @@ Repository: https://github.com/sglang-bot/sglang-ci-data (path: diffusion-ci/con
f"mean_abs_diff<={result.thresholds.mean_abs_diff_threshold}\n"
f" Failed frames:\n"
+ "\n".join(failed_frames)
+ (
f"\n Compared GT frame files and links:\n{video_gt_info}"
if video_gt_info
else ""
)
+ f"\n Compared GT files and links:\n{gt_remote_info}"
)
logger.info(
@@ -3,13 +3,19 @@
import unittest
from unittest.mock import MagicMock, patch
import numpy as np
import torch
from diffusers.pipelines.flux2.image_processor import Flux2ImageProcessor
from PIL import Image
from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType
from sglang.multimodal_gen.configs.pipeline_configs.flux import Flux2PipelineConfig
from sglang.multimodal_gen.configs.pipeline_configs.wan import (
WanI2V480PConfig,
WanI2V720PConfig,
)
from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
from sglang.multimodal_gen.runtime.pipelines.flux_2 import Flux2Pipeline
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines_core.stages.input_validation import (
InputValidationStage,
@@ -40,6 +46,28 @@ def _make_server_args(pipeline_config):
return sa
class _DummyTI2IConfig:
task_type = ModelTaskType.TI2I
def __init__(self):
self.vae_config = MagicMock()
self.vae_config.get_vae_scale_factor.return_value = 8
def preprocess_vae_image(self, batch, vae_image_processor):
return None
def calculate_condition_image_size(self, image, width, height):
return None
def preprocess_condition_image(
self, image, target_width, target_height, vae_image_processor
):
return image, (target_width, target_height)
def prepare_calculated_size(self, image):
return image.size
class TestCalculateDimensionsFromArea(unittest.TestCase):
"""Tests for InputValidationStage._calculate_dimensions_from_area."""
@@ -160,5 +188,88 @@ class TestPreprocessConditionImageResolution(unittest.TestCase):
self.assertEqual((batch.width, batch.height), (1280, 720))
class TestFlux2ConditionImagePreprocess(unittest.TestCase):
def test_matches_official_flux2_image_processor(self):
config = Flux2PipelineConfig()
config.vae_config.arch_config.vae_scale_factor = 8
processor = Flux2ImageProcessor(vae_scale_factor=16)
image = Image.fromarray(
np.arange(1792 * 1216 * 3, dtype=np.uint8).reshape(1216, 1792, 3),
mode="RGB",
)
size = config.calculate_condition_image_size(image, image.width, image.height)
self.assertEqual(size, (1232, 832))
processed, processed_size = config.preprocess_condition_image(
image, size[0], size[1], processor
)
official_image = processor._resize_to_target_area(image, 1024 * 1024)
expected_width = (official_image.width // 16) * 16
expected_height = (official_image.height // 16) * 16
expected = processor.preprocess(
official_image,
height=expected_height,
width=expected_width,
resize_mode="crop",
)
self.assertEqual(processed_size, (expected_width, expected_height))
self.assertTrue(torch.equal(processed, expected))
@patch.object(Flux2Pipeline, "add_standard_ti2i_stages")
def test_runtime_pipeline_uses_flux2_image_processor(self, mock_add_stages):
pipeline = object.__new__(Flux2Pipeline)
server_args = MagicMock()
server_args.pipeline_config.vae_config.arch_config.vae_scale_factor = 8
Flux2Pipeline.create_pipeline_stages(pipeline, server_args)
processor = mock_add_stages.call_args.kwargs["vae_image_processor"]
self.assertIsInstance(processor, Flux2ImageProcessor)
self.assertIs(
processor,
mock_add_stages.call_args.kwargs["image_vae_stage_kwargs"][
"vae_image_processor"
],
)
class TestFlux2TI2ISizeResolution(unittest.TestCase):
def setUp(self):
with patch(_GLOBAL_ARGS_PATCH, return_value=MagicMock()):
self.stage = InputValidationStage()
self.config = _DummyTI2IConfig()
def test_uses_condition_image_size_when_width_height_not_explicit(self):
image = Image.new("RGB", (1255, 833), color="red")
batch = _make_batch(image)
batch.extra = {}
self.stage.preprocess_condition_image(
batch,
_make_server_args(self.config),
image.width,
image.height,
)
self.assertEqual((batch.width, batch.height), (1248, 832))
def test_preserves_explicit_width_height_for_ti2i(self):
image = Image.new("RGB", (1255, 833), color="red")
batch = _make_batch(image, width=768, height=512)
batch.extra = {"explicit_fields": ["width", "height"]}
self.stage.preprocess_condition_image(
batch,
_make_server_args(self.config),
image.width,
image.height,
)
self.assertEqual((batch.width, batch.height), (768, 512))
if __name__ == "__main__":
unittest.main()
@@ -1,6 +1,7 @@
import argparse
import math
import unittest
from unittest.mock import MagicMock, patch
from sglang.multimodal_gen.configs.sample.diffusers_generic import (
DiffusersGenericSamplingParams,
@@ -196,6 +197,40 @@ class TestSamplingParamsCliArgs(unittest.TestCase):
self.assertEqual(target.negative_prompt, SamplingParams.negative_prompt)
def test_cli_path_tracks_explicit_width_height_fields(self):
server_args = MagicMock()
server_args.backend = "sglang"
server_args.model_id = None
server_args.pipeline_config = MagicMock()
with patch.object(
SamplingParams,
"from_pretrained",
side_effect=lambda *args, **kwargs: Flux2SamplingParams(),
):
implicit_size = SamplingParams.from_user_sampling_params_args(
"dummy-model",
server_args=server_args,
prompt="p",
image_path="/tmp/in.png",
)
explicit_size = SamplingParams.from_user_sampling_params_args(
"dummy-model",
server_args=server_args,
prompt="p",
image_path="/tmp/in.png",
width=768,
height=512,
)
implicit_fields = set(implicit_size.build_request_extra()["explicit_fields"])
explicit_fields = set(explicit_size.build_request_extra()["explicit_fields"])
self.assertNotIn("width", implicit_fields)
self.assertNotIn("height", implicit_fields)
self.assertIn("width", explicit_fields)
self.assertIn("height", explicit_fields)
if __name__ == "__main__":
unittest.main()