diff --git a/python/sglang/multimodal_gen/configs/models/encoders/base.py b/python/sglang/multimodal_gen/configs/models/encoders/base.py
index 0c4f86b36..514b1561a 100644
--- a/python/sglang/multimodal_gen/configs/models/encoders/base.py
+++ b/python/sglang/multimodal_gen/configs/models/encoders/base.py
@@ -13,6 +13,7 @@ from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
@dataclass
class EncoderArchConfig(ArchConfig):
+ _fsdp_shard_conditions: list = field(default_factory=lambda: [])
architectures: list[str] = field(default_factory=lambda: [])
_supported_attention_backends: set[AttentionBackendEnum] = field(
default_factory=lambda: {
diff --git a/python/sglang/multimodal_gen/configs/models/encoders/mistral.py b/python/sglang/multimodal_gen/configs/models/encoders/mistral.py
new file mode 100644
index 000000000..2ab6bbeac
--- /dev/null
+++ b/python/sglang/multimodal_gen/configs/models/encoders/mistral.py
@@ -0,0 +1,67 @@
+# SPDX-License-Identifier: Apache-2.0
+from dataclasses import dataclass, field
+
+from sglang.multimodal_gen.configs.models.encoders.base import (
+ TextEncoderArchConfig,
+ TextEncoderConfig,
+)
+
+
+def _is_transformer_layer(n: str, m) -> bool:
+ return "layers" in n and str.isdigit(n.split(".")[-1])
+
+
+def _is_embeddings(n: str, m) -> bool:
+ return n.endswith("embed_tokens")
+
+
+def _is_final_norm(n: str, m) -> bool:
+ return n.endswith("norm")
+
+
+@dataclass
+class Mistral3ArchConfig(TextEncoderArchConfig):
+ vocab_size: int = 32000
+ hidden_size: int = 4096
+ intermediate_size: int = 11008
+ num_hidden_layers: int = 32
+ num_attention_heads: int = 32
+ num_key_value_heads: int | None = None
+ hidden_act: str = "silu"
+ max_position_embeddings: int = 2048
+ initializer_range: float = 0.02
+ rms_norm_eps: float = 1e-6
+ use_cache: bool = True
+ pad_token_id: int = 0
+ bos_token_id: int = 1
+ eos_token_id: int = 2
+ pretraining_tp: int = 1
+ tie_word_embeddings: bool = False
+ rope_theta: float = 10000.0
+ rope_scaling: float | None = None
+ attention_bias: bool = False
+ attention_dropout: float = 0.0
+ mlp_bias: bool = False
+ head_dim: int | None = None
+ hidden_state_skip_layer: int = 2
+ text_len: int = 256
+ 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"),
+ (".gate_up_proj", ".gate_proj", 0), # type: ignore
+ (".gate_up_proj", ".up_proj", 1), # type: ignore
+ ]
+ )
+ _fsdp_shard_conditions: list = field(
+ default_factory=lambda: [_is_transformer_layer, _is_embeddings, _is_final_norm]
+ )
+
+
+@dataclass
+class Mistral3Config(TextEncoderConfig):
+ arch_config: TextEncoderArchConfig = field(default_factory=Mistral3ArchConfig)
+
+ prefix: str = "mistral3"
diff --git a/python/sglang/multimodal_gen/configs/models/vaes/flux.py b/python/sglang/multimodal_gen/configs/models/vaes/flux.py
index 0b56149d9..53bd44618 100644
--- a/python/sglang/multimodal_gen/configs/models/vaes/flux.py
+++ b/python/sglang/multimodal_gen/configs/models/vaes/flux.py
@@ -1,5 +1,3 @@
-# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
-
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
@@ -28,6 +26,11 @@ class FluxVAEArchConfig(VAEArchConfig):
clip_output: bool = True
+@dataclass
+class Flux2VAEArchConfig(FluxVAEArchConfig):
+ pass
+
+
@dataclass
class FluxVAEConfig(VAEConfig):
arch_config: FluxVAEArchConfig = field(default_factory=FluxVAEArchConfig)
@@ -48,3 +51,8 @@ class FluxVAEConfig(VAEConfig):
len(self.arch_config.block_out_channels) - 1
)
self.arch_config.spatial_compression_ratio = self.arch_config.vae_scale_factor
+
+
+@dataclass
+class Flux2VAEConfig(FluxVAEConfig):
+ arch_config: Flux2VAEArchConfig = field(default_factory=Flux2VAEArchConfig)
diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py
index 929201877..e7c2a8be3 100644
--- a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py
+++ b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py
@@ -149,6 +149,8 @@ class PipelineConfig:
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
default_factory=lambda: (preprocess_text,)
)
+
+ # get prompt_embeds from encoder output
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.tensor], ...] = (
field(default_factory=lambda: (postprocess_text,))
)
@@ -175,15 +177,19 @@ class PipelineConfig:
def slice_noise_pred(self, noise, latents):
return noise
- def adjust_size(self, width, height, image):
+ def maybe_resize_condition_image(self, width, height, image):
"""
image: input image
"""
- return width, height
+ return image, width, height
def adjust_num_frames(self, num_frames):
return num_frames
+ # tokenize the prompt
+ def tokenize_prompt(self, prompt: list[str], tokenizer, tok_kwargs) -> dict:
+ return tokenizer(prompt, **tok_kwargs)
+
# called in ImageEncodingStage, preprocess the image
def preprocess_image(self, image, image_processor: VaeImageProcessor):
return image
@@ -203,10 +209,32 @@ class PipelineConfig:
return shape
+ def get_decode_scale_and_shift(self, device, dtype, vae):
+ vae_arch_config = self.vae_config.arch_config
+ scaling_factor = getattr(vae_arch_config, "scaling_factor", None)
+ if scaling_factor is None:
+ scaling_factor = getattr(vae, "scaling_factor", None)
+
+ shift_factor = getattr(vae_arch_config, "shift_factor", None)
+ if shift_factor is None:
+ shift_factor = getattr(vae, "shift_factor", None)
+ return scaling_factor, shift_factor
+
# called after latents are prepared
def maybe_pack_latents(self, latents, batch_size, batch):
return latents
+ def maybe_prepare_latent_ids(self, latents):
+ return None
+
+ # called after vae encode
+ def post_process_vae_encode(self, image_latents, vae):
+ return image_latents
+
+ # called after scale_and_shift, before vae decoding
+ def preprocess_decoding(self, latents):
+ return latents
+
def gather_latents_for_sp(self, latents):
# For video latents [B, C, T_local, H, W], gather along time dim=2
latents = sequence_model_parallel_all_gather(latents, dim=2)
diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/flux.py b/python/sglang/multimodal_gen/configs/pipeline_configs/flux.py
index 60d194d9b..287b4b484 100644
--- a/python/sglang/multimodal_gen/configs/pipeline_configs/flux.py
+++ b/python/sglang/multimodal_gen/configs/pipeline_configs/flux.py
@@ -1,8 +1,8 @@
-# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
-
+import math
from dataclasses import dataclass, field
-from typing import Callable
+from typing import Callable, List, Optional
+import PIL.Image
import torch
from sglang.multimodal_gen.configs.models import DiTConfig, EncoderConfig, VAEConfig
@@ -11,8 +11,15 @@ from sglang.multimodal_gen.configs.models.encoders import (
BaseEncoderOutput,
CLIPTextConfig,
T5Config,
+ TextEncoderConfig,
)
-from sglang.multimodal_gen.configs.models.vaes.flux import FluxVAEConfig
+from sglang.multimodal_gen.configs.models.encoders.base import TextEncoderArchConfig
+from sglang.multimodal_gen.configs.models.encoders.mistral import (
+ Mistral3Config,
+ _is_embeddings,
+ _is_transformer_layer,
+)
+from sglang.multimodal_gen.configs.models.vaes.flux import Flux2VAEConfig, FluxVAEConfig
from sglang.multimodal_gen.configs.pipeline_configs.base import (
ImagePipelineConfig,
ModelTaskType,
@@ -122,7 +129,7 @@ class FluxPipelineConfig(ImagePipelineConfig):
return latent_image_ids
- def get_freqs_cis(self, prompt_embeds, width, height, device, rotary_emb):
+ def get_freqs_cis(self, prompt_embeds, width, height, device, rotary_emb, batch):
txt_ids = torch.zeros(prompt_embeds.shape[1], 3, device=device)
img_ids = self._prepare_latent_image_ids(
original_height=height,
@@ -156,7 +163,12 @@ class FluxPipelineConfig(ImagePipelineConfig):
def prepare_pos_cond_kwargs(self, batch, device, rotary_emb, dtype):
return {
"freqs_cis": self.get_freqs_cis(
- batch.prompt_embeds[1], batch.width, batch.height, device, rotary_emb
+ batch.prompt_embeds[1],
+ batch.width,
+ batch.height,
+ device,
+ rotary_emb,
+ batch,
),
"pooled_projections": (
batch.pooled_embeds[0] if batch.pooled_embeds else None
@@ -171,8 +183,362 @@ class FluxPipelineConfig(ImagePipelineConfig):
batch.height,
device,
rotary_emb,
+ batch,
),
"pooled_projections": (
batch.neg_pooled_embeds[0] if batch.neg_pooled_embeds else None
),
}
+
+
+def _prepare_latent_ids(
+ latents: torch.Tensor, # (B, C, H, W)
+):
+ r"""
+ Generates 4D position coordinates (T, H, W, L) for latent tensors.
+
+ Args:
+ latents (torch.Tensor):
+ Latent tensor of shape (B, C, H, W)
+
+ Returns:
+ torch.Tensor:
+ Position IDs tensor of shape (B, H*W, 4) All batches share the same coordinate structure: T=0,
+ H=[0..H-1], W=[0..W-1], L=0
+ """
+
+ batch_size, _, height, width = latents.shape
+
+ t = torch.arange(1) # [0] - time dimension
+ h = torch.arange(height)
+ w = torch.arange(width)
+ l = torch.arange(1) # [0] - layer dimension
+
+ # Create position IDs: (H*W, 4)
+ latent_ids = torch.cartesian_prod(t, h, w, l)
+
+ # Expand to batch: (B, H*W, 4)
+ latent_ids = latent_ids.unsqueeze(0).expand(batch_size, -1, -1)
+ return latent_ids
+
+
+def _unpack_latents_with_ids(
+ x: torch.Tensor, x_ids: torch.Tensor
+) -> list[torch.Tensor]:
+ """
+ using position ids to scatter tokens into place
+ """
+ x_list = []
+ x_ids = x_ids.to(device=x.device)
+ for data, pos in zip(x, x_ids):
+ _, ch = data.shape # noqa: F841
+ h_ids = pos[:, 1].to(torch.int64)
+ w_ids = pos[:, 2].to(torch.int64)
+
+ h = torch.max(h_ids) + 1
+ w = torch.max(w_ids) + 1
+
+ flat_ids = h_ids * w + w_ids
+
+ out = torch.zeros((h * w, ch), device=data.device, dtype=data.dtype)
+ out.scatter_(0, flat_ids.unsqueeze(1).expand(-1, ch), data)
+
+ # reshape from (H * W, C) to (H, W, C) and permute to (C, H, W)
+
+ out = out.view(h, w, ch).permute(2, 0, 1)
+ x_list.append(out)
+
+ return torch.stack(x_list, dim=0)
+
+
+def _patchify_latents(latents):
+ batch_size, num_channels_latents, height, width = latents.shape
+ latents = latents.view(
+ batch_size, num_channels_latents, height // 2, 2, width // 2, 2
+ )
+ latents = latents.permute(0, 1, 3, 5, 2, 4)
+ latents = latents.reshape(
+ batch_size, num_channels_latents * 4, height // 2, width // 2
+ )
+ return latents
+
+
+def _unpatchify_latents(latents):
+ batch_size, num_channels_latents, height, width = latents.shape
+ latents = latents.reshape(
+ batch_size, num_channels_latents // (2 * 2), 2, 2, height, width
+ )
+ latents = latents.permute(0, 1, 4, 2, 5, 3)
+ latents = latents.reshape(
+ batch_size, num_channels_latents // (2 * 2), height * 2, width * 2
+ )
+ return latents
+
+
+def _prepare_text_ids(
+ x: torch.Tensor, # (B, L, D) or (L, D)
+ t_coord: Optional[torch.Tensor] = None,
+):
+ B, L, _ = x.shape
+ out_ids = []
+
+ for i in range(B):
+ t = torch.arange(1) if t_coord is None else t_coord[i]
+ h = torch.arange(1)
+ w = torch.arange(1)
+ l = torch.arange(L)
+
+ coords = torch.cartesian_prod(t, h, w, l)
+ out_ids.append(coords)
+
+ return torch.stack(out_ids)
+
+
+def _prepare_image_ids(
+ image_latents: List[torch.Tensor], # [(1, C, H, W), (1, C, H, W), ...]
+ scale: int = 10,
+):
+ if not isinstance(image_latents, list):
+ raise ValueError(
+ f"Expected `image_latents` to be a list, got {type(image_latents)}."
+ )
+
+ # create time offset for each reference image
+ t_coords = [scale + scale * t for t in torch.arange(0, len(image_latents))]
+ t_coords = [t.view(-1) for t in t_coords]
+
+ image_latent_ids = []
+ for x, t in zip(image_latents, t_coords):
+ x = x.squeeze(0)
+ _, height, width = x.shape
+
+ x_ids = torch.cartesian_prod(
+ t, torch.arange(height), torch.arange(width), torch.arange(1)
+ )
+ image_latent_ids.append(x_ids)
+
+ image_latent_ids = torch.cat(image_latent_ids, dim=0)
+ image_latent_ids = image_latent_ids.unsqueeze(0)
+
+ return image_latent_ids
+
+
+def flux_2_postprocess_text(outputs: BaseEncoderOutput, _text_inputs) -> torch.Tensor:
+ hidden_states_layers: list[int] = [10, 20, 30]
+
+ out = torch.stack([outputs.hidden_states[k] for k in hidden_states_layers], dim=1)
+ batch_size, num_channels, seq_len, hidden_dim = out.shape
+ prompt_embeds = out.permute(0, 2, 1, 3).reshape(
+ batch_size, seq_len, num_channels * hidden_dim
+ )
+
+ 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, _is_embeddings]
+ )
+
+ 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):
+ print(f"{prompt=}")
+ 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
+ latents = latents.reshape(batch_size, num_channels, height * width).permute(0, 2, 1)
+
+ return latents
+
+
+@dataclass
+class Flux2PipelineConfig(FluxPipelineConfig):
+ embedded_cfg_scale: float = 4.0
+
+ task_type: ModelTaskType = ModelTaskType.I2I
+
+ text_encoder_configs: tuple[EncoderConfig, ...] = field(
+ default_factory=lambda: (Mistral3Config(),)
+ )
+ text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16",))
+
+ text_encoder_configs: tuple[EncoderConfig, ...] = field(
+ default_factory=lambda: (Flux2MistralTextConfig(),)
+ )
+ preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
+ default_factory=lambda: (flux_2_preprocess_text,),
+ )
+
+ postprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
+ default_factory=lambda: (flux_2_postprocess_text,)
+ )
+ vae_config: VAEConfig = field(default_factory=Flux2VAEConfig)
+
+ 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]
+ inputs = tokenizer.apply_chat_template(
+ prompts,
+ add_generation_prompt=False,
+ tokenize=True,
+ return_dict=True,
+ return_tensors="pt",
+ padding="max_length",
+ truncation=True,
+ # 2048 from official github repo, 512 from diffusers
+ max_length=512,
+ )
+
+ return inputs
+
+ def prepare_latent_shape(self, batch, batch_size, num_frames):
+ height = 2 * (
+ batch.height // (self.vae_config.arch_config.vae_scale_factor * 2)
+ )
+ width = 2 * (batch.width // (self.vae_config.arch_config.vae_scale_factor * 2))
+ num_channels_latents = self.dit_config.arch_config.in_channels
+ shape = (batch_size, num_channels_latents, height // 2, width // 2)
+ return shape
+
+ def get_pos_prompt_embeds(self, batch):
+ return batch.prompt_embeds[0]
+
+ def get_neg_prompt_embeds(self, batch):
+ return batch.negative_prompt_embeds[0]
+
+ def maybe_resize_condition_image(self, width, height, image):
+ target_area: int = 1024 * 1024
+
+ if width is not None and height is not None:
+ if width * height > target_area:
+ scale = math.sqrt(target_area / (width * height))
+ width = int(width * scale)
+ height = int(height * scale)
+ image = image.resize((width, height), PIL.Image.Resampling.LANCZOS)
+ width, height = image.size
+
+ return image, width, height
+
+ def get_freqs_cis(self, prompt_embeds, width, height, device, rotary_emb, batch):
+
+ txt_ids = _prepare_text_ids(prompt_embeds).to(device=device)
+
+ img_ids = batch.latent_ids
+ if batch.image_latent is not None:
+ image_latent_ids = batch.condition_image_latent_ids
+ img_ids = torch.cat([img_ids, image_latent_ids], dim=1).to(device=device)
+
+ if img_ids.ndim == 3:
+ img_ids = img_ids[0]
+ if txt_ids.ndim == 3:
+ txt_ids = txt_ids[0]
+
+ # NOTE(mick): prepare it here, to avoid unnecessary computations
+ img_cos, img_sin = rotary_emb.forward(img_ids)
+ img_cos = shard_rotary_emb_for_sp(img_cos)
+ img_sin = shard_rotary_emb_for_sp(img_sin)
+
+ txt_cos, txt_sin = rotary_emb.forward(txt_ids)
+
+ cos = torch.cat([txt_cos, img_cos], dim=0).to(device=device)
+ sin = torch.cat([txt_sin, img_sin], dim=0).to(device=device)
+ return cos, sin
+
+ def prepare_pos_cond_kwargs(self, batch, device, rotary_emb, dtype):
+ return {
+ "freqs_cis": self.get_freqs_cis(
+ batch.prompt_embeds[0],
+ batch.width,
+ batch.height,
+ device,
+ rotary_emb,
+ batch,
+ )
+ }
+
+ def prepare_neg_cond_kwargs(self, batch, device, rotary_emb, dtype):
+ return {}
+
+ def maybe_pack_latents(self, latents, batch_size, batch):
+ return flux2_pack_latents(latents)
+
+ def maybe_prepare_latent_ids(self, latents):
+ return _prepare_latent_ids(latents)
+
+ def post_process_vae_encode(self, image_latents, vae):
+ # patchify
+ image_latents = _patchify_latents(image_latents)
+ return image_latents
+
+ def preprocess_decoding(self, latents):
+ latents = _unpatchify_latents(latents)
+ return latents
+
+ def get_decode_scale_and_shift(self, device, dtype, vae):
+ vae_arch_config = self.vae_config.arch_config
+ latents_bn_mean = (
+ vae.bn.running_mean.view(1, -1, 1, 1).to(device=device).to(device, dtype)
+ )
+ latents_bn_std = torch.sqrt(
+ vae.bn.running_var.view(1, -1, 1, 1) + vae_arch_config.batch_norm_eps
+ ).to(device, dtype)
+ return 1 / latents_bn_std, latents_bn_mean
+
+ def post_denoising_loop(self, latents, batch):
+ latent_ids = batch.latent_ids
+ latents = _unpack_latents_with_ids(latents, latent_ids)
+
+ return latents
+
+ def slice_noise_pred(self, noise, latents):
+ # remove noise over input image
+ noise = noise[:, : latents.size(1) :]
+ return noise
diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/qwen_image.py b/python/sglang/multimodal_gen/configs/pipeline_configs/qwen_image.py
index 3441a0a17..fb39712a9 100644
--- a/python/sglang/multimodal_gen/configs/pipeline_configs/qwen_image.py
+++ b/python/sglang/multimodal_gen/configs/pipeline_configs/qwen_image.py
@@ -123,6 +123,18 @@ class QwenImagePipelineConfig(ImagePipelineConfig):
# pack latents
return _pack_latents(latents, batch_size, num_channels_latents, height, width)
+ def get_decode_scale_and_shift(self, device, dtype, vae):
+ vae_arch_config = self.vae_config.arch_config
+ scaling_factor = 1.0 / torch.tensor(
+ vae_arch_config.latents_std, device=device
+ ).view(1, vae_arch_config.z_dim, 1, 1, 1).to(device, dtype)
+ shift_factor = (
+ torch.tensor(vae_arch_config.latents_mean)
+ .view(1, vae_arch_config.z_dim, 1, 1, 1)
+ .to(device, dtype)
+ )
+ return scaling_factor, shift_factor
+
@staticmethod
def get_freqs_cis(img_shapes, txt_seq_lens, rotary_emb, device, dtype):
# img_shapes: for global entire image
@@ -203,7 +215,7 @@ class QwenImageEditPipelineConfig(QwenImagePipelineConfig):
assert batch_size == 1
height = batch.height
width = batch.width
- image = batch.pil_image
+ image = batch.condition_image
image_size = image[0].size if isinstance(image, list) else image.size
edit_width, edit_height, _ = calculate_dimensions(
1024 * 1024, image_size[0] / image_size[1]
@@ -268,7 +280,7 @@ class QwenImageEditPipelineConfig(QwenImagePipelineConfig):
image = image_processor.resize(image, calculated_height, calculated_width)
return image
- def adjust_size(self, width, height, image):
+ def maybe_resize_condition_image(self, width, height, image):
image_size = image[0].size if isinstance(image, list) else image.size
calculated_width, calculated_height, _ = calculate_dimensions(
1024 * 1024, image_size[0] / image_size[1]
diff --git a/python/sglang/multimodal_gen/configs/sample/base.py b/python/sglang/multimodal_gen/configs/sample/base.py
index 625defb76..09155b1dd 100644
--- a/python/sglang/multimodal_gen/configs/sample/base.py
+++ b/python/sglang/multimodal_gen/configs/sample/base.py
@@ -315,8 +315,10 @@ class SamplingParams:
sampling_params = SamplingParams.from_pretrained(model_path)
user_sampling_params = SamplingParams(*args, **kwargs)
+ # TODO: refactor
sampling_params._merge_with_user_params(user_sampling_params)
-
+ sampling_params.width_not_provided = user_sampling_params.width is None
+ sampling_params.height_not_provided = user_sampling_params.height is None
sampling_params.adjust(server_args)
return sampling_params
@@ -539,6 +541,8 @@ class SamplingParams:
if hasattr(self, field_name):
setattr(self, field_name, user_value)
+ self.height_not_provided = user_params.height_not_provided
+ self.width_not_provided = user_params.width_not_provided
self.__post_init__()
@property
diff --git a/python/sglang/multimodal_gen/configs/sample/flux.py b/python/sglang/multimodal_gen/configs/sample/flux.py
index 4c96467fb..f6b58fe85 100644
--- a/python/sglang/multimodal_gen/configs/sample/flux.py
+++ b/python/sglang/multimodal_gen/configs/sample/flux.py
@@ -16,3 +16,10 @@ class FluxSamplingParams(SamplingParams):
guidance_scale: float = 1.0
negative_prompt: str = None
num_inference_steps: int = 50
+
+ def __post_init__(self):
+ default_sample_size = 128
+ vae_scale_factor = 8
+ # FIXME
+ # self.height = default_sample_size * vae_scale_factor
+ # self.width = default_sample_size * vae_scale_factor
diff --git a/python/sglang/multimodal_gen/docs/support_matrix.md b/python/sglang/multimodal_gen/docs/support_matrix.md
index 99c5b2efa..d8eeb0374 100644
--- a/python/sglang/multimodal_gen/docs/support_matrix.md
+++ b/python/sglang/multimodal_gen/docs/support_matrix.md
@@ -10,34 +10,36 @@ The symbols used have the following meanings:
## Models x Optimization
-The `HuggingFace Model ID` can be passed directly to `from_pretrained()` methods, and sglang-diffusion will use the optimal
+The `HuggingFace Model ID` can be passed directly to `from_pretrained()` methods, and sglang-diffusion will use the
+optimal
default parameters when initializing and generating videos.
### Video Generation Models
-| Model Name | Hugging Face Model ID | Resolutions | TeaCache | Sliding Tile Attn | Sage Attn | Video Sparse Attention (VSA) |
-|:-----------------------------|:--------------------------------------------------|:---------------------------------------------|:--------:|:-----------------:|:---------:|:----------------------------:|
-| FastWan2.1 T2V 1.3B | `FastVideo/FastWan2.1-T2V-1.3B-Diffusers` | 480p | ⭕ | ⭕ | ⭕ | ✅ |
-| FastWan2.2 TI2V 5B Full Attn | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers` | 720p | ⭕ | ⭕ | ⭕ | ✅ |
-| Wan2.2 TI2V 5B | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | 720p | ⭕ | ⭕ | ✅ | ⭕ |
-| Wan2.2 T2V A14B | `Wan-AI/Wan2.2-T2V-A14B-Diffusers` | 480p
720p | ❌ | ❌ | ✅ | ⭕ |
-| Wan2.2 I2V A14B | `Wan-AI/Wan2.2-I2V-A14B-Diffusers` | 480p
720p | ❌ | ❌ | ✅ | ⭕ |
-| HunyuanVideo | `hunyuanvideo-community/HunyuanVideo` | 720×1280
544×960 | ❌ | ✅ | ✅ | ⭕ |
-| FastHunyuan | `FastVideo/FastHunyuan-diffusers` | 720×1280
544×960 | ❌ | ✅ | ✅ | ⭕ |
-| Wan2.1 T2V 1.3B | `Wan-AI/Wan2.1-T2V-1.3B-Diffusers` | 480p | ✅ | ✅ | ✅ | ⭕ |
-| Wan2.1 T2V 14B | `Wan-AI/Wan2.1-T2V-14B-Diffusers` | 480p, 720p | ✅ | ✅ | ✅ | ⭕ |
-| Wan2.1 I2V 480P | `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers` | 480p | ✅ | ✅ | ✅ | ⭕ |
-| Wan2.1 I2V 720P | `Wan-AI/Wan2.1-I2V-14B-720P-Diffusers` | 720p | ✅ | ✅ | ✅ | ⭕ |
+| Model Name | Hugging Face Model ID | Resolutions | TeaCache | Sliding Tile Attn | Sage Attn | Video Sparse Attention (VSA) |
+|:-----------------------------|:--------------------------------------------------|:--------------------|:--------:|:-----------------:|:---------:|:----------------------------:|
+| FastWan2.1 T2V 1.3B | `FastVideo/FastWan2.1-T2V-1.3B-Diffusers` | 480p | ⭕ | ⭕ | ⭕ | ✅ |
+| FastWan2.2 TI2V 5B Full Attn | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers` | 720p | ⭕ | ⭕ | ⭕ | ✅ |
+| Wan2.2 TI2V 5B | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | 720p | ⭕ | ⭕ | ✅ | ⭕ |
+| Wan2.2 T2V A14B | `Wan-AI/Wan2.2-T2V-A14B-Diffusers` | 480p
720p | ❌ | ❌ | ✅ | ⭕ |
+| Wan2.2 I2V A14B | `Wan-AI/Wan2.2-I2V-A14B-Diffusers` | 480p
720p | ❌ | ❌ | ✅ | ⭕ |
+| HunyuanVideo | `hunyuanvideo-community/HunyuanVideo` | 720×1280
544×960 | ❌ | ✅ | ✅ | ⭕ |
+| FastHunyuan | `FastVideo/FastHunyuan-diffusers` | 720×1280
544×960 | ❌ | ✅ | ✅ | ⭕ |
+| Wan2.1 T2V 1.3B | `Wan-AI/Wan2.1-T2V-1.3B-Diffusers` | 480p | ✅ | ✅ | ✅ | ⭕ |
+| Wan2.1 T2V 14B | `Wan-AI/Wan2.1-T2V-14B-Diffusers` | 480p, 720p | ✅ | ✅ | ✅ | ⭕ |
+| Wan2.1 I2V 480P | `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers` | 480p | ✅ | ✅ | ✅ | ⭕ |
+| Wan2.1 I2V 720P | `Wan-AI/Wan2.1-I2V-14B-720P-Diffusers` | 720p | ✅ | ✅ | ✅ | ⭕ |
**Note**: Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
### Image Generation Models
-| Model Name | HuggingFace Model ID | Resolutions | TeaCache | Sage Attn |
-|:----------------|:-------------------------------|:---------------|:--------:|:---------:|
-| FLUX.1-dev | `black-forest-labs/FLUX.1-dev` | Any resolution | ❌ | ❌ |
-| Qwen Image | `Qwen/Qwen-Image` | Any resolution | ❌ | ❌ |
-| Qwen Image Edit | `Qwen/Qwen-Image-Edit` | Any resolution | ❌ | ❌ |
+| Model Name | HuggingFace Model ID | Resolutions |
+|:----------------|:-------------------------------|:---------------|
+| FLUX.1-dev | `black-forest-labs/FLUX.1-dev` | Any resolution |
+| FLUX.2-dev | `black-forest-labs/FLUX.2-dev` | Any resolution |
+| Qwen Image | `Qwen/Qwen-Image` | Any resolution |
+| Qwen Image Edit | `Qwen/Qwen-Image-Edit` | Any resolution |
## Special requirements
diff --git a/python/sglang/multimodal_gen/registry.py b/python/sglang/multimodal_gen/registry.py
index 9600531ab..bf74b8c1e 100644
--- a/python/sglang/multimodal_gen/registry.py
+++ b/python/sglang/multimodal_gen/registry.py
@@ -26,6 +26,7 @@ from sglang.multimodal_gen.configs.pipeline_configs import (
WanT2V720PConfig,
)
from sglang.multimodal_gen.configs.pipeline_configs.base import PipelineConfig
+from sglang.multimodal_gen.configs.pipeline_configs.flux import Flux2PipelineConfig
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
QwenImageEditPipelineConfig,
QwenImagePipelineConfig,
@@ -394,6 +395,15 @@ def _register_configs():
],
model_detectors=[lambda id: "flux" in id.lower()],
)
+ register_configs(
+ model_name="flux-2",
+ sampling_param_cls=FluxSamplingParams,
+ pipeline_config_cls=Flux2PipelineConfig,
+ model_paths=[
+ "black-forest-labs/FLUX.2-dev",
+ ],
+ model_detectors=[lambda id: "flux.2" in id.lower()],
+ )
# Qwen-Image
register_configs(
diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loader.py
index 0975f8741..10c641383 100644
--- a/python/sglang/multimodal_gen/runtime/loader/component_loader.py
+++ b/python/sglang/multimodal_gen/runtime/loader/component_loader.py
@@ -337,7 +337,8 @@ class TextEncoderLoader(ComponentLoader):
cpu_offload=True,
reshard_after_forward=True,
mesh=mesh["offload"],
- fsdp_shard_conditions=model._fsdp_shard_conditions,
+ fsdp_shard_conditions=model_config.arch_config._fsdp_shard_conditions
+ or getattr(model, "_fsdp_shard_conditions", None),
pin_cpu_memory=server_args.pin_cpu_memory,
)
# We only enable strict check for non-quantized models
@@ -346,7 +347,7 @@ class TextEncoderLoader(ComponentLoader):
weights_not_loaded = weights_to_load - loaded_weights
if weights_not_loaded:
raise ValueError(
- "Following weights were not initialized from "
+ "Following model weights were not initialized from "
f"checkpoint: {weights_not_loaded}"
)
diff --git a/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py b/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py
new file mode 100644
index 000000000..adde9cd67
--- /dev/null
+++ b/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py
@@ -0,0 +1,860 @@
+# Copyright 2025 Black Forest Labs, The HuggingFace Team and The InstantX Team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# 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 typing import Any, Dict, Optional, Tuple
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from diffusers.models.attention import AttentionModuleMixin
+from diffusers.models.attention_dispatch import dispatch_attention_fn
+from diffusers.models.embeddings import TimestepEmbedding, Timesteps
+from diffusers.models.normalization import AdaLayerNormContinuous
+from diffusers.models.transformers.transformer_flux2 import Flux2PosEmbed
+
+from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig
+from sglang.multimodal_gen.runtime.layers.rotary_embedding import _apply_rotary_emb
+from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
+from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
+
+logger = init_logger(__name__) # pylint: disable=invalid-name
+
+
+def _get_projections(attn: "Flux2Attention", hidden_states, encoder_hidden_states=None):
+ query = attn.to_q(hidden_states)
+ key = attn.to_k(hidden_states)
+ value = attn.to_v(hidden_states)
+
+ encoder_query = encoder_key = encoder_value = None
+ if encoder_hidden_states is not None and attn.added_kv_proj_dim is not None:
+ encoder_query = attn.add_q_proj(encoder_hidden_states)
+ encoder_key = attn.add_k_proj(encoder_hidden_states)
+ encoder_value = attn.add_v_proj(encoder_hidden_states)
+
+ return query, key, value, encoder_query, encoder_key, encoder_value
+
+
+def _get_fused_projections(
+ attn: "Flux2Attention", hidden_states, encoder_hidden_states=None
+):
+ query, key, value = attn.to_qkv(hidden_states).chunk(3, dim=-1)
+
+ encoder_query = encoder_key = encoder_value = (None,)
+ if encoder_hidden_states is not None and hasattr(attn, "to_added_qkv"):
+ encoder_query, encoder_key, encoder_value = attn.to_added_qkv(
+ encoder_hidden_states
+ ).chunk(3, dim=-1)
+
+ return query, key, value, encoder_query, encoder_key, encoder_value
+
+
+def _get_qkv_projections(
+ attn: "Flux2Attention", hidden_states, encoder_hidden_states=None
+):
+ if attn.fused_projections:
+ return _get_fused_projections(attn, hidden_states, encoder_hidden_states)
+ return _get_projections(attn, hidden_states, encoder_hidden_states)
+
+
+class Flux2SwiGLU(nn.Module):
+ """
+ Flux 2 uses a SwiGLU-style activation in the transformer feedforward sub-blocks, but with the linear projection
+ layer fused into the first linear layer of the FF sub-block. Thus, this module has no trainable parameters.
+ """
+
+ def __init__(self):
+ super().__init__()
+ self.gate_fn = nn.SiLU()
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ x1, x2 = x.chunk(2, dim=-1)
+ x = self.gate_fn(x1) * x2
+ return x
+
+
+class Flux2FeedForward(nn.Module):
+ def __init__(
+ self,
+ dim: int,
+ dim_out: Optional[int] = None,
+ mult: float = 3.0,
+ inner_dim: Optional[int] = None,
+ bias: bool = False,
+ ):
+ super().__init__()
+ if inner_dim is None:
+ inner_dim = int(dim * mult)
+ dim_out = dim_out or dim
+
+ # Flux2SwiGLU will reduce the dimension by half
+ self.linear_in = nn.Linear(dim, inner_dim * 2, bias=bias)
+ self.act_fn = Flux2SwiGLU()
+ self.linear_out = nn.Linear(inner_dim, dim_out, bias=bias)
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ x = self.linear_in(x)
+ x = self.act_fn(x)
+ x = self.linear_out(x)
+ return x
+
+
+class Flux2AttnProcessor:
+ _attention_backend = None
+ _parallel_config = None
+
+ def __init__(self):
+ if not hasattr(F, "scaled_dot_product_attention"):
+ raise ImportError(
+ f"{self.__class__.__name__} requires PyTorch 2.0. Please upgrade your pytorch version."
+ )
+
+ def __call__(
+ self,
+ attn: "Flux2Attention",
+ hidden_states: torch.Tensor,
+ encoder_hidden_states: torch.Tensor = None,
+ attention_mask: Optional[torch.Tensor] = None,
+ freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
+ ) -> torch.Tensor:
+ query, key, value, encoder_query, encoder_key, encoder_value = (
+ _get_qkv_projections(attn, hidden_states, encoder_hidden_states)
+ )
+
+ query = query.unflatten(-1, (attn.heads, -1))
+ key = key.unflatten(-1, (attn.heads, -1))
+ value = value.unflatten(-1, (attn.heads, -1))
+
+ query = attn.norm_q(query)
+ key = attn.norm_k(key)
+
+ if attn.added_kv_proj_dim is not None:
+ encoder_query = encoder_query.unflatten(-1, (attn.heads, -1))
+ encoder_key = encoder_key.unflatten(-1, (attn.heads, -1))
+ encoder_value = encoder_value.unflatten(-1, (attn.heads, -1))
+
+ encoder_query = attn.norm_added_q(encoder_query)
+ encoder_key = attn.norm_added_k(encoder_key)
+
+ query = torch.cat([encoder_query, query], dim=1)
+ key = torch.cat([encoder_key, key], dim=1)
+ value = torch.cat([encoder_value, value], dim=1)
+
+ if freqs_cis is not None:
+ cos, sin = freqs_cis
+ query = _apply_rotary_emb(
+ query, cos, sin, is_neox_style=False, interleaved=True
+ )
+ key = _apply_rotary_emb(
+ key, cos, sin, is_neox_style=False, interleaved=True
+ )
+
+ hidden_states = dispatch_attention_fn(
+ query,
+ key,
+ value,
+ attn_mask=attention_mask,
+ backend=self._attention_backend,
+ parallel_config=self._parallel_config,
+ )
+ hidden_states = hidden_states.flatten(2, 3)
+ hidden_states = hidden_states.to(query.dtype)
+
+ if encoder_hidden_states is not None:
+ encoder_hidden_states, hidden_states = hidden_states.split_with_sizes(
+ [
+ encoder_hidden_states.shape[1],
+ hidden_states.shape[1] - encoder_hidden_states.shape[1],
+ ],
+ dim=1,
+ )
+ encoder_hidden_states = attn.to_add_out(encoder_hidden_states)
+
+ hidden_states = attn.to_out[0](hidden_states)
+ hidden_states = attn.to_out[1](hidden_states)
+
+ if encoder_hidden_states is not None:
+ return hidden_states, encoder_hidden_states
+ else:
+ return hidden_states
+
+
+class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
+ _default_processor_cls = Flux2AttnProcessor
+ _available_processors = [Flux2AttnProcessor]
+
+ def __init__(
+ self,
+ query_dim: int,
+ heads: int = 8,
+ dim_head: int = 64,
+ dropout: float = 0.0,
+ bias: bool = False,
+ added_kv_proj_dim: Optional[int] = None,
+ added_proj_bias: Optional[bool] = True,
+ out_bias: bool = True,
+ eps: float = 1e-5,
+ out_dim: int = None,
+ elementwise_affine: bool = True,
+ processor=None,
+ ):
+ super().__init__()
+
+ self.head_dim = dim_head
+ self.inner_dim = out_dim if out_dim is not None else dim_head * heads
+ self.query_dim = query_dim
+ self.out_dim = out_dim if out_dim is not None else query_dim
+ self.heads = out_dim // dim_head if out_dim is not None else heads
+
+ self.use_bias = bias
+ self.dropout = dropout
+
+ self.added_kv_proj_dim = added_kv_proj_dim
+ self.added_proj_bias = added_proj_bias
+
+ self.to_q = torch.nn.Linear(query_dim, self.inner_dim, bias=bias)
+ self.to_k = torch.nn.Linear(query_dim, self.inner_dim, bias=bias)
+ self.to_v = torch.nn.Linear(query_dim, self.inner_dim, bias=bias)
+
+ # QK Norm
+ self.norm_q = torch.nn.RMSNorm(
+ dim_head, eps=eps, elementwise_affine=elementwise_affine
+ )
+ self.norm_k = torch.nn.RMSNorm(
+ dim_head, eps=eps, elementwise_affine=elementwise_affine
+ )
+
+ self.to_out = torch.nn.ModuleList([])
+ self.to_out.append(torch.nn.Linear(self.inner_dim, self.out_dim, bias=out_bias))
+ self.to_out.append(torch.nn.Dropout(dropout))
+
+ if added_kv_proj_dim is not None:
+ self.norm_added_q = torch.nn.RMSNorm(dim_head, eps=eps)
+ self.norm_added_k = torch.nn.RMSNorm(dim_head, eps=eps)
+ self.add_q_proj = torch.nn.Linear(
+ added_kv_proj_dim, self.inner_dim, bias=added_proj_bias
+ )
+ self.add_k_proj = torch.nn.Linear(
+ added_kv_proj_dim, self.inner_dim, bias=added_proj_bias
+ )
+ self.add_v_proj = torch.nn.Linear(
+ added_kv_proj_dim, self.inner_dim, bias=added_proj_bias
+ )
+ self.to_add_out = torch.nn.Linear(self.inner_dim, query_dim, bias=out_bias)
+
+ if processor is None:
+ processor = self._default_processor_cls()
+ self.set_processor(processor)
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ encoder_hidden_states: Optional[torch.Tensor] = None,
+ attention_mask: Optional[torch.Tensor] = None,
+ freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
+ **kwargs,
+ ) -> torch.Tensor:
+ attn_parameters = set(
+ inspect.signature(self.processor.__call__).parameters.keys()
+ )
+ unused_kwargs = [k for k, _ in kwargs.items() if k not in attn_parameters]
+ if len(unused_kwargs) > 0:
+ logger.warning(
+ f"joint_attention_kwargs {unused_kwargs} are not expected by {self.processor.__class__.__name__} and will be ignored."
+ )
+ kwargs = {k: w for k, w in kwargs.items() if k in attn_parameters}
+ return self.processor(
+ self,
+ hidden_states,
+ encoder_hidden_states,
+ attention_mask,
+ freqs_cis,
+ **kwargs,
+ )
+
+
+class Flux2ParallelSelfAttnProcessor:
+ _attention_backend = None
+ _parallel_config = None
+
+ def __init__(self):
+ if not hasattr(F, "scaled_dot_product_attention"):
+ raise ImportError(
+ f"{self.__class__.__name__} requires PyTorch 2.0. Please upgrade your pytorch version."
+ )
+
+ def __call__(
+ self,
+ attn: "Flux2ParallelSelfAttention",
+ hidden_states: torch.Tensor,
+ attention_mask: Optional[torch.Tensor] = None,
+ freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
+ ) -> torch.Tensor:
+ # Parallel in (QKV + MLP in) projection
+ hidden_states = attn.to_qkv_mlp_proj(hidden_states)
+ qkv, mlp_hidden_states = torch.split(
+ hidden_states,
+ [3 * attn.inner_dim, attn.mlp_hidden_dim * attn.mlp_mult_factor],
+ dim=-1,
+ )
+
+ # Handle the attention logic
+ query, key, value = qkv.chunk(3, dim=-1)
+
+ query = query.unflatten(-1, (attn.heads, -1))
+ key = key.unflatten(-1, (attn.heads, -1))
+ value = value.unflatten(-1, (attn.heads, -1))
+
+ query = attn.norm_q(query)
+ key = attn.norm_k(key)
+
+ if freqs_cis is not None:
+ cos, sin = freqs_cis
+ query = _apply_rotary_emb(
+ query, cos, sin, is_neox_style=False, interleaved=True
+ )
+ key = _apply_rotary_emb(
+ key, cos, sin, is_neox_style=False, interleaved=True
+ )
+
+ hidden_states = dispatch_attention_fn(
+ query,
+ key,
+ value,
+ attn_mask=attention_mask,
+ backend=self._attention_backend,
+ parallel_config=self._parallel_config,
+ )
+ hidden_states = hidden_states.flatten(2, 3)
+ hidden_states = hidden_states.to(query.dtype)
+
+ # Handle the feedforward (FF) logic
+ mlp_hidden_states = attn.mlp_act_fn(mlp_hidden_states)
+
+ # Concatenate and parallel output projection
+ hidden_states = torch.cat([hidden_states, mlp_hidden_states], dim=-1)
+ hidden_states = attn.to_out(hidden_states)
+
+ return hidden_states
+
+
+class Flux2ParallelSelfAttention(torch.nn.Module, AttentionModuleMixin):
+ """
+ Flux 2 parallel self-attention for the Flux 2 single-stream transformer blocks.
+
+ This implements a parallel transformer block, where the attention QKV projections are fused to the feedforward (FF)
+ input projections, and the attention output projections are fused to the FF output projections. See the [ViT-22B
+ paper](https://arxiv.org/abs/2302.05442) for a visual depiction of this type of transformer block.
+ """
+
+ _default_processor_cls = Flux2ParallelSelfAttnProcessor
+ _available_processors = [Flux2ParallelSelfAttnProcessor]
+ # Does not support QKV fusion as the QKV projections are always fused
+ _supports_qkv_fusion = False
+
+ def __init__(
+ self,
+ query_dim: int,
+ heads: int = 8,
+ dim_head: int = 64,
+ dropout: float = 0.0,
+ bias: bool = False,
+ out_bias: bool = True,
+ eps: float = 1e-5,
+ out_dim: int = None,
+ elementwise_affine: bool = True,
+ mlp_ratio: float = 4.0,
+ mlp_mult_factor: int = 2,
+ processor=None,
+ ):
+ super().__init__()
+
+ self.head_dim = dim_head
+ self.inner_dim = out_dim if out_dim is not None else dim_head * heads
+ self.query_dim = query_dim
+ self.out_dim = out_dim if out_dim is not None else query_dim
+ self.heads = out_dim // dim_head if out_dim is not None else heads
+
+ self.use_bias = bias
+ self.dropout = dropout
+
+ self.mlp_ratio = mlp_ratio
+ self.mlp_hidden_dim = int(query_dim * self.mlp_ratio)
+ self.mlp_mult_factor = mlp_mult_factor
+
+ # Fused QKV projections + MLP input projection
+ self.to_qkv_mlp_proj = torch.nn.Linear(
+ self.query_dim,
+ self.inner_dim * 3 + self.mlp_hidden_dim * self.mlp_mult_factor,
+ bias=bias,
+ )
+ self.mlp_act_fn = Flux2SwiGLU()
+
+ # QK Norm
+ self.norm_q = torch.nn.RMSNorm(
+ dim_head, eps=eps, elementwise_affine=elementwise_affine
+ )
+ self.norm_k = torch.nn.RMSNorm(
+ dim_head, eps=eps, elementwise_affine=elementwise_affine
+ )
+
+ # Fused attention output projection + MLP output projection
+ self.to_out = torch.nn.Linear(
+ self.inner_dim + self.mlp_hidden_dim, self.out_dim, bias=out_bias
+ )
+
+ if processor is None:
+ processor = self._default_processor_cls()
+ self.set_processor(processor)
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ attention_mask: Optional[torch.Tensor] = None,
+ freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
+ **kwargs,
+ ) -> torch.Tensor:
+ attn_parameters = set(
+ inspect.signature(self.processor.__call__).parameters.keys()
+ )
+ unused_kwargs = [k for k, _ in kwargs.items() if k not in attn_parameters]
+ if len(unused_kwargs) > 0:
+ logger.warning(
+ f"joint_attention_kwargs {unused_kwargs} are not expected by {self.processor.__class__.__name__} and will be ignored."
+ )
+ kwargs = {k: w for k, w in kwargs.items() if k in attn_parameters}
+ return self.processor(self, hidden_states, attention_mask, freqs_cis, **kwargs)
+
+
+class Flux2SingleTransformerBlock(nn.Module):
+ def __init__(
+ self,
+ dim: int,
+ num_attention_heads: int,
+ attention_head_dim: int,
+ mlp_ratio: float = 3.0,
+ eps: float = 1e-6,
+ bias: bool = False,
+ ):
+ super().__init__()
+
+ self.norm = nn.LayerNorm(dim, elementwise_affine=False, eps=eps)
+
+ # Note that the MLP in/out linear layers are fused with the attention QKV/out projections, respectively; this
+ # is often called a "parallel" transformer block. See the [ViT-22B paper](https://arxiv.org/abs/2302.05442)
+ # for a visual depiction of this type of transformer block.
+ self.attn = Flux2ParallelSelfAttention(
+ query_dim=dim,
+ dim_head=attention_head_dim,
+ heads=num_attention_heads,
+ out_dim=dim,
+ bias=bias,
+ out_bias=bias,
+ eps=eps,
+ mlp_ratio=mlp_ratio,
+ mlp_mult_factor=2,
+ processor=Flux2ParallelSelfAttnProcessor(),
+ )
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ encoder_hidden_states: Optional[torch.Tensor],
+ temb_mod_params: Tuple[torch.Tensor, torch.Tensor, torch.Tensor],
+ freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
+ joint_attention_kwargs: Optional[Dict[str, Any]] = None,
+ split_hidden_states: bool = False,
+ text_seq_len: Optional[int] = None,
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
+ # If encoder_hidden_states is None, hidden_states is assumed to have encoder_hidden_states already
+ # concatenated
+ if encoder_hidden_states is not None:
+ text_seq_len = encoder_hidden_states.shape[1]
+ hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
+
+ mod_shift, mod_scale, mod_gate = temb_mod_params
+
+ norm_hidden_states = self.norm(hidden_states)
+ norm_hidden_states = (1 + mod_scale) * norm_hidden_states + mod_shift
+
+ joint_attention_kwargs = joint_attention_kwargs or {}
+ attn_output = self.attn(
+ hidden_states=norm_hidden_states,
+ freqs_cis=freqs_cis,
+ **joint_attention_kwargs,
+ )
+
+ hidden_states = hidden_states + mod_gate * attn_output
+ if hidden_states.dtype == torch.float16:
+ hidden_states = hidden_states.clip(-65504, 65504)
+
+ if split_hidden_states:
+ encoder_hidden_states, hidden_states = (
+ hidden_states[:, :text_seq_len],
+ hidden_states[:, text_seq_len:],
+ )
+ return encoder_hidden_states, hidden_states
+ else:
+ return hidden_states
+
+
+class Flux2TransformerBlock(nn.Module):
+ def __init__(
+ self,
+ dim: int,
+ num_attention_heads: int,
+ attention_head_dim: int,
+ mlp_ratio: float = 3.0,
+ eps: float = 1e-6,
+ bias: bool = False,
+ ):
+ super().__init__()
+ self.mlp_hidden_dim = int(dim * mlp_ratio)
+
+ self.norm1 = nn.LayerNorm(dim, elementwise_affine=False, eps=eps)
+ self.norm1_context = nn.LayerNorm(dim, elementwise_affine=False, eps=eps)
+
+ self.attn = Flux2Attention(
+ query_dim=dim,
+ added_kv_proj_dim=dim,
+ dim_head=attention_head_dim,
+ heads=num_attention_heads,
+ out_dim=dim,
+ bias=bias,
+ added_proj_bias=bias,
+ out_bias=bias,
+ eps=eps,
+ processor=Flux2AttnProcessor(),
+ )
+
+ self.norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=eps)
+ self.ff = Flux2FeedForward(dim=dim, dim_out=dim, mult=mlp_ratio, bias=bias)
+
+ self.norm2_context = nn.LayerNorm(dim, elementwise_affine=False, eps=eps)
+ self.ff_context = Flux2FeedForward(
+ dim=dim, dim_out=dim, mult=mlp_ratio, bias=bias
+ )
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ encoder_hidden_states: torch.Tensor,
+ temb_mod_params_img: Tuple[
+ Tuple[torch.Tensor, torch.Tensor, torch.Tensor], ...
+ ],
+ temb_mod_params_txt: Tuple[
+ Tuple[torch.Tensor, torch.Tensor, torch.Tensor], ...
+ ],
+ freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
+ joint_attention_kwargs: Optional[Dict[str, Any]] = None,
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
+ joint_attention_kwargs = joint_attention_kwargs or {}
+
+ # Modulation parameters shape: [1, 1, self.dim]
+ (shift_msa, scale_msa, gate_msa), (shift_mlp, scale_mlp, gate_mlp) = (
+ temb_mod_params_img
+ )
+ (c_shift_msa, c_scale_msa, c_gate_msa), (
+ c_shift_mlp,
+ c_scale_mlp,
+ c_gate_mlp,
+ ) = temb_mod_params_txt
+
+ # Img stream
+ norm_hidden_states = self.norm1(hidden_states)
+ norm_hidden_states = (1 + scale_msa) * norm_hidden_states + shift_msa
+
+ # Conditioning txt stream
+ norm_encoder_hidden_states = self.norm1_context(encoder_hidden_states)
+ norm_encoder_hidden_states = (
+ 1 + c_scale_msa
+ ) * norm_encoder_hidden_states + c_shift_msa
+
+ # Attention on concatenated img + txt stream
+ attention_outputs = self.attn(
+ hidden_states=norm_hidden_states,
+ encoder_hidden_states=norm_encoder_hidden_states,
+ freqs_cis=freqs_cis,
+ **joint_attention_kwargs,
+ )
+
+ attn_output, context_attn_output = attention_outputs
+
+ # Process attention outputs for the image stream (`hidden_states`).
+ attn_output = gate_msa * attn_output
+ hidden_states = hidden_states + attn_output
+
+ norm_hidden_states = self.norm2(hidden_states)
+ norm_hidden_states = norm_hidden_states * (1 + scale_mlp) + shift_mlp
+
+ ff_output = self.ff(norm_hidden_states)
+ hidden_states = hidden_states + gate_mlp * ff_output
+
+ # Process attention outputs for the text stream (`encoder_hidden_states`).
+ context_attn_output = c_gate_msa * context_attn_output
+ encoder_hidden_states = encoder_hidden_states + context_attn_output
+
+ norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states)
+ norm_encoder_hidden_states = (
+ norm_encoder_hidden_states * (1 + c_scale_mlp) + c_shift_mlp
+ )
+
+ context_ff_output = self.ff_context(norm_encoder_hidden_states)
+ encoder_hidden_states = encoder_hidden_states + c_gate_mlp * context_ff_output
+ if encoder_hidden_states.dtype == torch.float16:
+ encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504)
+
+ return encoder_hidden_states, hidden_states
+
+
+class Flux2TimestepGuidanceEmbeddings(nn.Module):
+ def __init__(
+ self, in_channels: int = 256, embedding_dim: int = 6144, bias: bool = False
+ ):
+ super().__init__()
+
+ self.time_proj = Timesteps(
+ num_channels=in_channels, flip_sin_to_cos=True, downscale_freq_shift=0
+ )
+ self.timestep_embedder = TimestepEmbedding(
+ in_channels=in_channels, time_embed_dim=embedding_dim, sample_proj_bias=bias
+ )
+
+ self.guidance_embedder = TimestepEmbedding(
+ in_channels=in_channels, time_embed_dim=embedding_dim, sample_proj_bias=bias
+ )
+
+ def forward(self, timestep: torch.Tensor, guidance: torch.Tensor) -> torch.Tensor:
+ timesteps_proj = self.time_proj(timestep)
+ timesteps_emb = self.timestep_embedder(
+ timesteps_proj.to(timestep.dtype)
+ ) # (N, D)
+
+ guidance_proj = self.time_proj(guidance)
+ guidance_emb = self.guidance_embedder(
+ guidance_proj.to(guidance.dtype)
+ ) # (N, D)
+
+ time_guidance_emb = timesteps_emb + guidance_emb
+
+ return time_guidance_emb
+
+
+class Flux2Modulation(nn.Module):
+ def __init__(self, dim: int, mod_param_sets: int = 2, bias: bool = False):
+ super().__init__()
+ self.mod_param_sets = mod_param_sets
+
+ self.linear = nn.Linear(dim, dim * 3 * self.mod_param_sets, bias=bias)
+ self.act_fn = nn.SiLU()
+
+ def forward(
+ self, temb: torch.Tensor
+ ) -> Tuple[Tuple[torch.Tensor, torch.Tensor, torch.Tensor], ...]:
+ mod = self.act_fn(temb)
+ mod = self.linear(mod)
+
+ if mod.ndim == 2:
+ mod = mod.unsqueeze(1)
+ mod_params = torch.chunk(mod, 3 * self.mod_param_sets, dim=-1)
+ # Return tuple of 3-tuples of modulation params shift/scale/gate
+ return tuple(
+ mod_params[3 * i : 3 * (i + 1)] for i in range(self.mod_param_sets)
+ )
+
+
+class Flux2Transformer2DModel(CachableDiT):
+ """
+ The Transformer model introduced in Flux 2.
+
+ Reference: https://blackforestlabs.ai/announcing-black-forest-labs/
+
+ """
+
+ def __init__(self, config: FluxConfig, hf_config: dict[str, Any]):
+ super().__init__(config=config, hf_config=hf_config)
+ patch_size: int = config.patch_size
+ in_channels: int = config.in_channels
+ out_channels: Optional[int] = config.out_channels
+ num_layers: int = config.num_layers
+ num_single_layers: int = config.num_single_layers
+ attention_head_dim: int = config.attention_head_dim
+ num_attention_heads: int = config.num_attention_heads
+ joint_attention_dim: int = config.joint_attention_dim
+ timestep_guidance_channels: int = config.timestep_guidance_channels
+ mlp_ratio: float = config.mlp_ratio
+ axes_dims_rope: Tuple[int, ...] = config.axes_dims_rope
+ rope_theta: int = config.rope_theta
+ eps: float = config.eps
+ self.out_channels = out_channels or in_channels
+ self.inner_dim = num_attention_heads * attention_head_dim
+
+ # 1. Sinusoidal positional embedding for RoPE on image and text tokens
+ # self.rotary_emb = Flux2PosEmbed(theta=rope_theta, axes_dim=axes_dims_rope)
+ self.rotary_emb = Flux2PosEmbed(theta=rope_theta, axes_dim=axes_dims_rope)
+
+ # 2. Combined timestep + guidance embedding
+ self.time_guidance_embed = Flux2TimestepGuidanceEmbeddings(
+ in_channels=timestep_guidance_channels,
+ embedding_dim=self.inner_dim,
+ bias=False,
+ )
+
+ # 3. Modulation (double stream and single stream blocks share modulation parameters, resp.)
+ # Two sets of shift/scale/gate modulation parameters for the double stream attn and FF sub-blocks
+ self.double_stream_modulation_img = Flux2Modulation(
+ self.inner_dim, mod_param_sets=2, bias=False
+ )
+ self.double_stream_modulation_txt = Flux2Modulation(
+ self.inner_dim, mod_param_sets=2, bias=False
+ )
+ # Only one set of modulation parameters as the attn and FF sub-blocks are run in parallel for single stream
+ self.single_stream_modulation = Flux2Modulation(
+ self.inner_dim, mod_param_sets=1, bias=False
+ )
+
+ # 4. Input projections
+ self.x_embedder = nn.Linear(in_channels, self.inner_dim, bias=False)
+ self.context_embedder = nn.Linear(
+ joint_attention_dim, self.inner_dim, bias=False
+ )
+
+ # 5. Double Stream Transformer Blocks
+ self.transformer_blocks = nn.ModuleList(
+ [
+ Flux2TransformerBlock(
+ dim=self.inner_dim,
+ num_attention_heads=num_attention_heads,
+ attention_head_dim=attention_head_dim,
+ mlp_ratio=mlp_ratio,
+ eps=eps,
+ bias=False,
+ )
+ for _ in range(num_layers)
+ ]
+ )
+
+ # 6. Single Stream Transformer Blocks
+ self.single_transformer_blocks = nn.ModuleList(
+ [
+ Flux2SingleTransformerBlock(
+ dim=self.inner_dim,
+ num_attention_heads=num_attention_heads,
+ attention_head_dim=attention_head_dim,
+ mlp_ratio=mlp_ratio,
+ eps=eps,
+ bias=False,
+ )
+ for _ in range(num_single_layers)
+ ]
+ )
+
+ # 7. Output layers
+ self.norm_out = AdaLayerNormContinuous(
+ self.inner_dim,
+ self.inner_dim,
+ elementwise_affine=False,
+ eps=eps,
+ bias=False,
+ )
+ self.proj_out = nn.Linear(
+ self.inner_dim, patch_size * patch_size * self.out_channels, bias=False
+ )
+
+ self.gradient_checkpointing = False
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ encoder_hidden_states: torch.Tensor = None,
+ timestep: torch.LongTensor = None,
+ guidance: torch.Tensor = None,
+ freqs_cis: torch.Tensor = None,
+ joint_attention_kwargs: Optional[Dict[str, Any]] = None,
+ ) -> torch.Tensor:
+ """
+ The [`FluxTransformer2DModel`] forward method.
+
+ Args:
+ hidden_states (`torch.Tensor` of shape `(batch_size, image_sequence_length, in_channels)`):
+ Input `hidden_states`.
+ encoder_hidden_states (`torch.Tensor` of shape `(batch_size, text_sequence_length, joint_attention_dim)`):
+ Conditional embeddings (embeddings computed from the input conditions such as prompts) to use.
+ timestep ( `torch.LongTensor`):
+ Used to indicate denoising step.
+ joint_attention_kwargs (`dict`, *optional*):
+ A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
+ `self.processor` in
+ [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
+
+ """
+ # 0. Handle input arguments
+ if joint_attention_kwargs is not None:
+ joint_attention_kwargs = joint_attention_kwargs.copy()
+ lora_scale = joint_attention_kwargs.pop("scale", 1.0)
+ else:
+ lora_scale = 1.0
+
+ num_txt_tokens = encoder_hidden_states.shape[1]
+
+ # 1. Calculate timestep embedding and modulation parameters
+ timestep = timestep.to(hidden_states.dtype)
+ guidance = guidance.to(hidden_states.dtype)
+
+ temb = self.time_guidance_embed(timestep, guidance)
+
+ double_stream_mod_img = self.double_stream_modulation_img(temb)
+ double_stream_mod_txt = self.double_stream_modulation_txt(temb)
+ single_stream_mod = self.single_stream_modulation(temb)[0]
+
+ # 2. Input projection for image (hidden_states) and conditioning text (encoder_hidden_states)
+ hidden_states = self.x_embedder(hidden_states)
+ encoder_hidden_states = self.context_embedder(encoder_hidden_states)
+
+ # 3. Calculate RoPE embeddings from image and text tokens
+ # NOTE: the below logic means that we can't support batched inference with images of different resolutions or
+ # text prompts of different lengths. Is this a use case we want to support?
+ # 4. Double Stream Transformer Blocks
+ for index_block, block in enumerate(self.transformer_blocks):
+ encoder_hidden_states, hidden_states = block(
+ hidden_states=hidden_states,
+ encoder_hidden_states=encoder_hidden_states,
+ temb_mod_params_img=double_stream_mod_img,
+ temb_mod_params_txt=double_stream_mod_txt,
+ freqs_cis=freqs_cis,
+ joint_attention_kwargs=joint_attention_kwargs,
+ )
+ # Concatenate text and image streams for single-block inference
+ hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
+
+ # 5. Single Stream Transformer Blocks
+ for index_block, block in enumerate(self.single_transformer_blocks):
+ hidden_states = block(
+ hidden_states=hidden_states,
+ encoder_hidden_states=None,
+ temb_mod_params=single_stream_mod,
+ freqs_cis=freqs_cis,
+ joint_attention_kwargs=joint_attention_kwargs,
+ )
+ # Remove text tokens from concatenated stream
+ hidden_states = hidden_states[:, num_txt_tokens:, ...]
+
+ # 6. Output layers
+ hidden_states = self.norm_out(hidden_states, temb)
+ output = self.proj_out(hidden_states)
+
+ return output
+
+
+EntryClass = Flux2Transformer2DModel
diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/base.py b/python/sglang/multimodal_gen/runtime/models/encoders/base.py
index a36c616cc..7b3eece96 100644
--- a/python/sglang/multimodal_gen/runtime/models/encoders/base.py
+++ b/python/sglang/multimodal_gen/runtime/models/encoders/base.py
@@ -25,7 +25,7 @@ class TextEncoder(nn.Module, ABC):
def __init__(self, config: TextEncoderConfig) -> None:
super().__init__()
self.config = config
- self._fsdp_shard_conditions = config._fsdp_shard_conditions
+ self._fsdp_shard_conditions = config.arch_config._fsdp_shard_conditions
self._stacked_params_mapping = config.arch_config.stacked_params_mapping
if not self.supported_attention_backends:
raise ValueError(
diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/mistral_3.py b/python/sglang/multimodal_gen/runtime/models/encoders/mistral_3.py
new file mode 100644
index 000000000..5554f005d
--- /dev/null
+++ b/python/sglang/multimodal_gen/runtime/models/encoders/mistral_3.py
@@ -0,0 +1,446 @@
+# coding=utf-8
+# Copyright 2025 HuggingFace Inc. team. All rights reserved.
+#
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# 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.
+from typing import Iterable, Optional, Union
+
+import torch
+from torch import nn
+from transformers import Cache, DynamicCache, LlavaConfig, Mistral3Config, MistralConfig
+from transformers.masking_utils import create_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,
+ MistralRMSNorm,
+ MistralRotaryEmbedding,
+ apply_rotary_pos_emb,
+)
+
+from sglang.multimodal_gen.runtime.loader.weight_utils import default_weight_loader
+from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
+
+logger = init_logger(__name__)
+
+
+def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
+ """
+ This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep).
+ The hidden states go from (batch, num_key_value_heads, seqlen, head_dim) to
+ (batch, num_attention_heads, seqlen, head_dim)
+ """
+ batch, num_key_value_heads, slen, head_dim = hidden_states.shape
+ if n_rep == 1:
+ return hidden_states
+ hidden_states = hidden_states[:, :, None, :, :].expand(
+ batch, num_key_value_heads, n_rep, slen, head_dim
+ )
+ return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)
+
+
+class MistralAttention(nn.Module):
+ """Multi-headed attention from 'Attention Is All You Need' paper"""
+
+ def __init__(self, config: MistralConfig, layer_idx: int):
+ 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
+ )
+ self.num_key_value_groups = (
+ config.num_attention_heads // config.num_key_value_heads
+ )
+ 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
+ )
+ self.k_proj = nn.Linear(
+ config.hidden_size, config.num_key_value_heads * self.head_dim, bias=False
+ )
+ self.v_proj = nn.Linear(
+ config.hidden_size, config.num_key_value_heads * self.head_dim, bias=False
+ )
+ self.o_proj = nn.Linear(
+ config.num_attention_heads * self.head_dim, config.hidden_size, bias=False
+ )
+ self.is_causal = True
+ self.num_heads = config.num_attention_heads
+ self.num_key_value_heads = config.num_key_value_heads
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ position_embeddings: tuple[torch.Tensor, torch.Tensor],
+ attention_mask: Optional[torch.Tensor],
+ past_key_values: Optional[Cache] = None,
+ cache_position: Optional[torch.LongTensor] = None,
+ **kwargs,
+ ) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
+ input_shape = hidden_states.shape[:-1]
+ hidden_shape = (*input_shape, -1, self.head_dim)
+
+ query_states = self.q_proj(hidden_states).view(hidden_shape).transpose(1, 2)
+ key_states = self.k_proj(hidden_states).view(hidden_shape).transpose(1, 2)
+ value_states = self.v_proj(hidden_states).view(hidden_shape).transpose(1, 2)
+
+ cos, sin = position_embeddings
+ query_states, key_states = apply_rotary_pos_emb(
+ query_states, key_states, cos, sin
+ )
+
+ if past_key_values is not None:
+ # sin and cos are specific to RoPE models; cache_position needed for the static cache
+ cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position}
+ key_states, value_states = past_key_values.update(
+ key_states, value_states, self.layer_idx, cache_kwargs
+ )
+
+ attention_interface = ALL_ATTENTION_FUNCTIONS["sdpa"]
+ attn_output, attn_weights = attention_interface(
+ self,
+ query_states,
+ key_states,
+ value_states,
+ attention_mask,
+ dropout=0.0,
+ scaling=self.scaling,
+ sliding_window=getattr(
+ self.config, "sliding_window", None
+ ), # main diff with Llama
+ **kwargs,
+ )
+
+ attn_output = attn_output.reshape(*input_shape, -1).contiguous()
+ attn_output = self.o_proj(attn_output)
+ return attn_output
+
+
+class MistralDecoderLayer(nn.Module):
+ def __init__(self, config: MistralConfig, layer_idx: int):
+ super().__init__()
+ self.hidden_size = config.hidden_size
+ self.self_attn = MistralAttention(config=config, layer_idx=layer_idx)
+ self.mlp = MistralMLP(config)
+ self.input_layernorm = MistralRMSNorm(
+ config.hidden_size, eps=config.rms_norm_eps
+ )
+ self.post_attention_layernorm = MistralRMSNorm(
+ config.hidden_size, eps=config.rms_norm_eps
+ )
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ attention_mask: Optional[torch.Tensor] = None,
+ position_ids: Optional[torch.LongTensor] = None,
+ past_key_values: Optional[Cache] = None,
+ use_cache: Optional[bool] = False,
+ cache_position: Optional[torch.LongTensor] = None,
+ position_embeddings: Optional[
+ tuple[torch.Tensor, torch.Tensor]
+ ] = None, # necessary, but kept here for BC
+ **kwargs,
+ ) -> torch.Tensor:
+ residual = hidden_states
+ hidden_states = self.input_layernorm(hidden_states)
+ # Self Attention
+ hidden_states = self.self_attn(
+ hidden_states=hidden_states,
+ attention_mask=attention_mask,
+ position_ids=position_ids,
+ past_key_values=past_key_values,
+ use_cache=use_cache,
+ cache_position=cache_position,
+ position_embeddings=position_embeddings,
+ **kwargs,
+ )
+ hidden_states = residual + hidden_states
+
+ # Fully Connected
+ residual = hidden_states
+ hidden_states = self.post_attention_layernorm(hidden_states)
+ hidden_states = self.mlp(hidden_states)
+ hidden_states = residual + hidden_states
+ return hidden_states
+
+
+class MistralModel(nn.Module):
+ def __init__(self, config: MistralConfig):
+ super().__init__()
+ self.config = config
+ self.padding_idx = config.pad_token_id
+ self.vocab_size = config.vocab_size
+
+ self.embed_tokens = nn.Embedding(
+ config.vocab_size, config.hidden_size, self.padding_idx
+ )
+ self.layers = nn.ModuleList(
+ [
+ MistralDecoderLayer(config, layer_idx)
+ for layer_idx in range(config.num_hidden_layers)
+ ]
+ )
+ 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"
+
+ def forward(
+ self,
+ input_ids: Optional[torch.LongTensor] = None,
+ attention_mask: Optional[torch.Tensor] = None,
+ position_ids: Optional[torch.LongTensor] = None,
+ past_key_values: Optional[Cache] = None,
+ inputs_embeds: Optional[torch.FloatTensor] = None,
+ use_cache: Optional[bool] = None,
+ cache_position: Optional[torch.LongTensor] = None,
+ output_hidden_states: Optional[bool] = None,
+ **kwargs,
+ ) -> BaseModelOutputWithPast:
+ if (input_ids is None) ^ (inputs_embeds is not None):
+ raise ValueError(
+ "You must specify exactly one of input_ids or inputs_embeds"
+ )
+
+ if inputs_embeds is None:
+ inputs_embeds = self.embed_tokens(input_ids)
+
+ if use_cache and past_key_values is None:
+ past_key_values = DynamicCache(config=self.config)
+
+ if cache_position is None:
+ past_seen_tokens = (
+ past_key_values.get_seq_length() if past_key_values is not None else 0
+ )
+ cache_position = torch.arange(
+ past_seen_tokens,
+ past_seen_tokens + inputs_embeds.shape[1],
+ device=inputs_embeds.device,
+ )
+
+ if position_ids is None:
+ position_ids = cache_position.unsqueeze(0)
+ mask_function = create_causal_mask
+ causal_mask = mask_function(
+ config=self.config,
+ input_embeds=inputs_embeds,
+ attention_mask=attention_mask,
+ cache_position=cache_position,
+ past_key_values=past_key_values,
+ position_ids=position_ids,
+ )
+
+ hidden_states = inputs_embeds
+ position_embeddings = self.rotary_emb(hidden_states, position_ids)
+
+ hidden_states_pool = []
+ for decoder_layer in self.layers[: self.config.num_hidden_layers]:
+ hidden_states = decoder_layer(
+ hidden_states,
+ attention_mask=causal_mask,
+ position_ids=position_ids,
+ past_key_values=past_key_values,
+ use_cache=use_cache,
+ cache_position=cache_position,
+ 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:
+ hidden_states_pool.append(hidden_states)
+
+ return BaseModelOutputWithPast(
+ hidden_states=hidden_states_pool,
+ last_hidden_state=hidden_states,
+ past_key_values=past_key_values if use_cache else None,
+ )
+
+
+class Mistral3Model(nn.Module):
+ _checkpoint_conversion_mapping = {"language_model.model": "language_model"}
+
+ def __init__(self, config: Mistral3Config):
+ super().__init__()
+ self.language_model = MistralModel(config.text_config)
+ self.config = config
+
+ def get_input_embeddings(self):
+ return self.language_model.embed_tokens
+
+ def set_decoder(self, decoder):
+ self.language_model = decoder
+
+ def get_decoder(self):
+ return self.language_model
+
+ def forward(
+ self,
+ input_ids: Optional[torch.LongTensor] = 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,
+ use_cache: Optional[bool] = None,
+ output_attentions: Optional[bool] = None,
+ output_hidoutput_hidden_statesden_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]:
+ output_attentions = False
+ output_hidden_states = True
+
+ if (input_ids is None) ^ (inputs_embeds is not None):
+ raise ValueError(
+ "You must specify exactly one of input_ids or inputs_embeds"
+ )
+
+ if inputs_embeds is None:
+ inputs_embeds = self.get_input_embeddings()(input_ids)
+
+ outputs: BaseModelOutputWithPast = self.language_model(
+ attention_mask=attention_mask,
+ position_ids=position_ids,
+ past_key_values=past_key_values,
+ inputs_embeds=inputs_embeds,
+ use_cache=use_cache,
+ output_attentions=output_attentions,
+ output_hidden_states=output_hidden_states,
+ return_dict=True,
+ cache_position=cache_position,
+ **kwargs,
+ )
+
+ return Mistral3ModelOutputWithPast(
+ last_hidden_state=outputs.last_hidden_state,
+ past_key_values=outputs.past_key_values,
+ hidden_states=outputs.hidden_states,
+ attentions=outputs.attentions,
+ )
+
+
+class Mistral3ForConditionalGeneration(nn.Module):
+ _checkpoint_conversion_mapping = {
+ "^language_model.model": "model.language_model",
+ "^multi_modal_projector": "model.multi_modal_projector",
+ "^language_model.lm_head": "lm_head",
+ }
+ _tied_weights_keys = ["lm_head.weight"]
+
+ def __init__(self, config: LlavaConfig):
+ super().__init__()
+ self.model = Mistral3Model(config.arch_config)
+
+ def get_input_embeddings(self):
+ return self.model.get_input_embeddings()
+
+ def set_decoder(self, decoder):
+ self.model.set_decoder(decoder)
+
+ def get_decoder(self):
+ return self.model.get_decoder()
+
+ # Make modules available through conditional class for BC
+ @property
+ def language_model(self):
+ return self.model.language_model
+
+ def forward(
+ self,
+ input_ids: Optional[torch.LongTensor] = None,
+ attention_mask: Optional[torch.Tensor] = None,
+ position_ids: Optional[torch.LongTensor] = None,
+ past_key_values: Optional[Cache] = None,
+ output_hidden_states: Optional[bool] = None,
+ inputs_embeds: Optional[torch.FloatTensor] = None,
+ labels: Optional[torch.LongTensor] = None,
+ use_cache: Optional[bool] = None,
+ return_dict: Optional[bool] = None,
+ cache_position: Optional[torch.LongTensor] = None,
+ logits_to_keep: Union[int, torch.Tensor] = 0,
+ image_sizes: Optional[torch.Tensor] = None,
+ **kwargs,
+ ) -> Union[tuple, Mistral3CausalLMOutputWithPast]:
+ r"""
+ labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
+ Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,
+ config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored
+ (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.
+
+ Example:
+
+ """
+ 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,
+ )
+
+ return Mistral3CausalLMOutputWithPast(
+ hidden_states=outputs.hidden_states,
+ )
+
+ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
+ # Define mapping for stacked parameters
+ params_dict = dict(self.named_parameters())
+ loaded_params: set[str] = set()
+ for name, loaded_weight in weights:
+ name_lower = name.lower()
+ if (
+ "vision" in name_lower
+ or "multi" in name_lower
+ or "lm_head" in name_lower
+ ):
+ continue
+ final_name = name.replace("language_model.model.", "model.language_model.")
+
+ if final_name in params_dict:
+ param = params_dict[final_name]
+ weight_loader = getattr(param, "weight_loader", default_weight_loader)
+ weight_loader(param, loaded_weight)
+ loaded_params.add(final_name)
+ else:
+ logger.warning(f"Param {name=} {final_name=} from weight is not loaded")
+
+ return loaded_params
+
+
+EntryClass = Mistral3ForConditionalGeneration
diff --git a/python/sglang/multimodal_gen/runtime/models/registry.py b/python/sglang/multimodal_gen/runtime/models/registry.py
index ea81be77b..2ee262571 100644
--- a/python/sglang/multimodal_gen/runtime/models/registry.py
+++ b/python/sglang/multimodal_gen/runtime/models/registry.py
@@ -237,10 +237,10 @@ def _try_inspect_model_cls(
@dataclass
class _ModelRegistry:
# Keyed by model_arch
- models: dict[str, _BaseRegisteredModel] = field(default_factory=dict)
+ registered_models: dict[str, _BaseRegisteredModel] = field(default_factory=dict)
def get_supported_archs(self) -> Set[str]:
- return self.models.keys()
+ return self.registered_models.keys()
def register_model(
self,
@@ -258,7 +258,7 @@ class _ModelRegistry:
when importing the model and thus the related error
:code:`RuntimeError: Cannot re-initialize CUDA in forked subprocess`.
"""
- if model_arch in self.models:
+ if model_arch in self.registered_models:
logger.warning(
"Model architecture %s is already registered, and will be "
"overwritten by the new model class %s.",
@@ -276,7 +276,7 @@ class _ModelRegistry:
else:
model = _RegisteredModel.from_model_cls(model_cls)
- self.models[model_arch] = model
+ self.registered_models[model_arch] = model
def _raise_for_unsupported(self, architectures: list[str]) -> NoReturn:
all_supported_archs = self.get_supported_archs()
@@ -293,16 +293,16 @@ class _ModelRegistry:
)
def _try_load_model_cls(self, model_arch: str) -> type[nn.Module] | None:
- if model_arch not in self.models:
+ if model_arch not in self.registered_models:
return None
- return _try_load_model_cls(model_arch, self.models[model_arch])
+ return _try_load_model_cls(model_arch, self.registered_models[model_arch])
def _try_inspect_model_cls(self, model_arch: str) -> _ModelInfo | None:
- if model_arch not in self.models:
+ if model_arch not in self.registered_models:
return None
- return _try_inspect_model_cls(model_arch, self.models[model_arch])
+ return _try_inspect_model_cls(model_arch, self.registered_models[model_arch])
def _normalize_archs(
self,
@@ -314,13 +314,12 @@ class _ModelRegistry:
logger.warning("No model architectures are specified")
normalized_arch = []
- for model in architectures:
- if model not in self.models:
+ for arch in architectures:
+ if arch not in self.registered_models:
raise Exception(
- f"Unsupported model architecture: {model}. Registered architectures: {architectures}"
+ f"Unsupported model architecture: {arch}. Registered architectures: {self.registered_models=}"
)
- model = "TransformersModel"
- normalized_arch.append(model)
+ normalized_arch.append(arch)
return normalized_arch
def inspect_model_cls(
diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_kl_flux2.py b/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_kl_flux2.py
new file mode 100644
index 000000000..9df72c09c
--- /dev/null
+++ b/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_kl_flux2.py
@@ -0,0 +1,524 @@
+import math
+from typing import Dict, Optional, Tuple, Union
+
+import torch
+import torch.nn as nn
+from diffusers.models.attention_processor import (
+ ADDED_KV_ATTENTION_PROCESSORS,
+ CROSS_ATTENTION_PROCESSORS,
+ AttentionProcessor,
+ AttnAddedKVProcessor,
+ AttnProcessor,
+)
+from diffusers.models.autoencoders.vae import (
+ Decoder,
+ DecoderOutput,
+ DiagonalGaussianDistribution,
+ Encoder,
+)
+from diffusers.models.modeling_outputs import AutoencoderKLOutput
+
+from sglang.multimodal_gen.configs.models.vaes.flux import Flux2VAEConfig
+from sglang.multimodal_gen.runtime.models.vaes.common import ParallelTiledVAE
+
+
+class AutoencoderKLFlux2(nn.Module, ParallelTiledVAE):
+ r"""
+ A VAE model with KL loss for encoding images into latents and decoding latent representations into images.
+
+ This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented
+ for all models (such as downloading or saving).
+
+ Parameters:
+ """
+
+ _supports_gradient_checkpointing = True
+ _no_split_modules = ["BasicTransformerBlock", "ResnetBlock2D"]
+
+ def __init__(
+ self,
+ config: Flux2VAEConfig,
+ ):
+ super().__init__()
+
+ self.config = config
+ arch_config = config.arch_config
+
+ in_channels: int = arch_config.in_channels
+ out_channels: int = arch_config.out_channels
+ down_block_types: Tuple[str, ...] = arch_config.down_block_types
+ up_block_types: Tuple[str, ...] = arch_config.up_block_types
+ block_out_channels: Tuple[int, ...] = arch_config.block_out_channels
+ layers_per_block: int = arch_config.layers_per_block
+ act_fn: str = arch_config.act_fn
+ latent_channels: int = arch_config.latent_channels
+ norm_num_groups: int = arch_config.norm_num_groups
+ sample_size: int = arch_config.sample_size
+ force_upcast: bool = arch_config.force_upcast
+ use_quant_conv: bool = arch_config.use_quant_conv
+ use_post_quant_conv: bool = arch_config.use_post_quant_conv
+ mid_block_add_attention: bool = arch_config.mid_block_add_attention
+ batch_norm_eps: float = arch_config.batch_norm_eps
+ batch_norm_momentum: float = arch_config.batch_norm_momentum
+ patch_size: Tuple[int, int] = arch_config.patch_size
+ # pass init params to Encoder
+ self.encoder = Encoder(
+ in_channels=in_channels,
+ out_channels=latent_channels,
+ down_block_types=down_block_types,
+ block_out_channels=block_out_channels,
+ layers_per_block=layers_per_block,
+ act_fn=act_fn,
+ norm_num_groups=norm_num_groups,
+ double_z=True,
+ mid_block_add_attention=mid_block_add_attention,
+ )
+
+ # pass init params to Decoder
+ self.decoder = Decoder(
+ in_channels=latent_channels,
+ out_channels=out_channels,
+ up_block_types=up_block_types,
+ block_out_channels=block_out_channels,
+ layers_per_block=layers_per_block,
+ norm_num_groups=norm_num_groups,
+ act_fn=act_fn,
+ mid_block_add_attention=mid_block_add_attention,
+ )
+
+ self.quant_conv = (
+ nn.Conv2d(2 * latent_channels, 2 * latent_channels, 1)
+ if use_quant_conv
+ else None
+ )
+ self.post_quant_conv = (
+ nn.Conv2d(latent_channels, latent_channels, 1)
+ if use_post_quant_conv
+ else None
+ )
+
+ self.bn = nn.BatchNorm2d(
+ math.prod(patch_size) * latent_channels,
+ eps=batch_norm_eps,
+ momentum=batch_norm_momentum,
+ affine=False,
+ track_running_stats=True,
+ )
+
+ self.use_slicing = False
+ self.use_tiling = False
+
+ # only relevant if vae tiling is enabled
+ self.tile_sample_min_size = self.config.sample_size
+ sample_size = (
+ self.config.sample_size[0]
+ if isinstance(self.config.sample_size, (list, tuple))
+ else self.config.sample_size
+ )
+ self.tile_latent_min_size = int(
+ sample_size / (2 ** (len(self.config.block_out_channels) - 1))
+ )
+ self.tile_overlap_factor = 0.25
+
+ @property
+ # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.attn_processors
+ def attn_processors(self) -> Dict[str, AttentionProcessor]:
+ r"""
+ Returns:
+ `dict` of attention processors: A dictionary containing all attention processors used in the model with
+ indexed by its weight name.
+ """
+ # set recursively
+ processors = {}
+
+ def fn_recursive_add_processors(
+ name: str,
+ module: torch.nn.Module,
+ processors: Dict[str, AttentionProcessor],
+ ):
+ if hasattr(module, "get_processor"):
+ processors[f"{name}.processor"] = module.get_processor()
+
+ for sub_name, child in module.named_children():
+ fn_recursive_add_processors(f"{name}.{sub_name}", child, processors)
+
+ return processors
+
+ for name, module in self.named_children():
+ fn_recursive_add_processors(name, module, processors)
+
+ return processors
+
+ # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_attn_processor
+ def set_attn_processor(
+ self, processor: Union[AttentionProcessor, Dict[str, AttentionProcessor]]
+ ):
+ r"""
+ Sets the attention processor to use to compute attention.
+
+ Parameters:
+ processor (`dict` of `AttentionProcessor` or only `AttentionProcessor`):
+ The instantiated processor class or a dictionary of processor classes that will be set as the processor
+ for **all** `Attention` layers.
+
+ If `processor` is a dict, the key needs to define the path to the corresponding cross attention
+ processor. This is strongly recommended when setting trainable attention processors.
+
+ """
+ count = len(self.attn_processors.keys())
+
+ if isinstance(processor, dict) and len(processor) != count:
+ raise ValueError(
+ f"A dict of processors was passed, but the number of processors {len(processor)} does not match the"
+ f" number of attention layers: {count}. Please make sure to pass {count} processor classes."
+ )
+
+ def fn_recursive_attn_processor(name: str, module: torch.nn.Module, processor):
+ if hasattr(module, "set_processor"):
+ if not isinstance(processor, dict):
+ module.set_processor(processor)
+ else:
+ module.set_processor(processor.pop(f"{name}.processor"))
+
+ for sub_name, child in module.named_children():
+ fn_recursive_attn_processor(f"{name}.{sub_name}", child, processor)
+
+ for name, module in self.named_children():
+ fn_recursive_attn_processor(name, module, processor)
+
+ # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_default_attn_processor
+ def set_default_attn_processor(self):
+ """
+ Disables custom attention processors and sets the default attention implementation.
+ """
+ if all(
+ proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS
+ for proc in self.attn_processors.values()
+ ):
+ processor = AttnAddedKVProcessor()
+ elif all(
+ proc.__class__ in CROSS_ATTENTION_PROCESSORS
+ for proc in self.attn_processors.values()
+ ):
+ processor = AttnProcessor()
+ else:
+ raise ValueError(
+ f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}"
+ )
+
+ self.set_attn_processor(processor)
+
+ def _encode(self, x: torch.Tensor) -> torch.Tensor:
+ batch_size, num_channels, height, width = x.shape
+
+ if self.use_tiling and (
+ width > self.tile_sample_min_size or height > self.tile_sample_min_size
+ ):
+ return self._tiled_encode(x)
+
+ enc = self.encoder(x)
+ if self.quant_conv is not None:
+ enc = self.quant_conv(enc)
+
+ return enc
+
+ def encode(
+ self, x: torch.Tensor, return_dict: bool = True
+ ) -> Union[DiagonalGaussianDistribution]:
+ """
+ Encode a batch of images into latents.
+
+ Args:
+ x (`torch.Tensor`): Input batch of images.
+ return_dict (`bool`, *optional*, defaults to `True`):
+ Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple.
+
+ Returns:
+ The latent representations of the encoded images. If `return_dict` is True, a
+ [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned.
+ """
+
+ if x.ndim == 5:
+ assert x.shape[2] == 1
+ x = x.squeeze(2)
+
+ if self.use_slicing and x.shape[0] > 1:
+ encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)]
+ h = torch.cat(encoded_slices)
+ else:
+ h = self._encode(x)
+
+ posterior = DiagonalGaussianDistribution(h)
+ return posterior
+
+ def _decode(
+ self, z: torch.Tensor, return_dict: bool = True
+ ) -> Union[DecoderOutput, torch.Tensor]:
+ if self.use_tiling and (
+ z.shape[-1] > self.tile_latent_min_size
+ or z.shape[-2] > self.tile_latent_min_size
+ ):
+ return self.tiled_decode(z, return_dict=return_dict)
+
+ if self.post_quant_conv is not None:
+ z = self.post_quant_conv(z)
+
+ dec = self.decoder(z)
+
+ if not return_dict:
+ return (dec,)
+
+ return DecoderOutput(sample=dec)
+
+ def decode(
+ self, z: torch.FloatTensor, return_dict: bool = True, generator=None
+ ) -> Union[DecoderOutput, torch.FloatTensor]:
+ """
+ Decode a batch of images.
+
+ Args:
+ z (`torch.Tensor`): Input batch of latent vectors.
+ return_dict (`bool`, *optional*, defaults to `True`):
+ Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple.
+
+ Returns:
+ [`~models.vae.DecoderOutput`] or `tuple`:
+ If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is
+ returned.
+
+ """
+ if self.use_slicing and z.shape[0] > 1:
+ decoded_slices = [self._decode(z_slice).sample for z_slice in z.split(1)]
+ decoded = torch.cat(decoded_slices)
+ else:
+ decoded = self._decode(z).sample
+
+ return decoded
+
+ def blend_v(
+ self, a: torch.Tensor, b: torch.Tensor, blend_extent: int
+ ) -> torch.Tensor:
+ blend_extent = min(a.shape[2], b.shape[2], blend_extent)
+ for y in range(blend_extent):
+ b[:, :, y, :] = a[:, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[
+ :, :, y, :
+ ] * (y / blend_extent)
+ return b
+
+ def blend_h(
+ self, a: torch.Tensor, b: torch.Tensor, blend_extent: int
+ ) -> torch.Tensor:
+ blend_extent = min(a.shape[3], b.shape[3], blend_extent)
+ for x in range(blend_extent):
+ b[:, :, :, x] = a[:, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[
+ :, :, :, x
+ ] * (x / blend_extent)
+ return b
+
+ def _tiled_encode(self, x: torch.Tensor) -> torch.Tensor:
+ r"""Encode a batch of images using a tiled encoder.
+
+ When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several
+ steps. This is useful to keep memory use constant regardless of image size. The end result of tiled encoding is
+ different from non-tiled encoding because each tile uses a different encoder. To avoid tiling artifacts, the
+ tiles overlap and are blended together to form a smooth output. You may still see tile-sized changes in the
+ output, but they should be much less noticeable.
+
+ Args:
+ x (`torch.Tensor`): Input batch of images.
+
+ Returns:
+ `torch.Tensor`:
+ The latent representation of the encoded videos.
+ """
+
+ overlap_size = int(self.tile_sample_min_size * (1 - self.tile_overlap_factor))
+ blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor)
+ row_limit = self.tile_latent_min_size - blend_extent
+
+ # Split the image into 512x512 tiles and encode them separately.
+ rows = []
+ for i in range(0, x.shape[2], overlap_size):
+ row = []
+ for j in range(0, x.shape[3], overlap_size):
+ tile = x[
+ :,
+ :,
+ i : i + self.tile_sample_min_size,
+ j : j + self.tile_sample_min_size,
+ ]
+ tile = self.encoder(tile)
+ if self.config.use_quant_conv:
+ tile = self.quant_conv(tile)
+ row.append(tile)
+ rows.append(row)
+ result_rows = []
+ for i, row in enumerate(rows):
+ result_row = []
+ for j, tile in enumerate(row):
+ # blend the above tile and the left tile
+ # to the current tile and add the current tile to the result row
+ if i > 0:
+ tile = self.blend_v(rows[i - 1][j], tile, blend_extent)
+ if j > 0:
+ tile = self.blend_h(row[j - 1], tile, blend_extent)
+ result_row.append(tile[:, :, :row_limit, :row_limit])
+ result_rows.append(torch.cat(result_row, dim=3))
+
+ enc = torch.cat(result_rows, dim=2)
+ return enc
+
+ def tiled_encode(
+ self, x: torch.Tensor, return_dict: bool = True
+ ) -> AutoencoderKLOutput:
+ r"""Encode a batch of images using a tiled encoder.
+
+ When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several
+ steps. This is useful to keep memory use constant regardless of image size. The end result of tiled encoding is
+ different from non-tiled encoding because each tile uses a different encoder. To avoid tiling artifacts, the
+ tiles overlap and are blended together to form a smooth output. You may still see tile-sized changes in the
+ output, but they should be much less noticeable.
+
+ Args:
+ x (`torch.Tensor`): Input batch of images.
+ return_dict (`bool`, *optional*, defaults to `True`):
+ Whether or not to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple.
+
+ Returns:
+ [`~models.autoencoder_kl.AutoencoderKLOutput`] or `tuple`:
+ If return_dict is True, a [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain
+ `tuple` is returned.
+ """
+ deprecation_message = (
+ "The tiled_encode implementation supporting the `return_dict` parameter is deprecated. In the future, the "
+ "implementation of this method will be replaced with that of `_tiled_encode` and you will no longer be able "
+ "to pass `return_dict`. You will also have to create a `DiagonalGaussianDistribution()` from the returned value."
+ )
+
+ overlap_size = int(self.tile_sample_min_size * (1 - self.tile_overlap_factor))
+ blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor)
+ row_limit = self.tile_latent_min_size - blend_extent
+
+ # Split the image into 512x512 tiles and encode them separately.
+ rows = []
+ for i in range(0, x.shape[2], overlap_size):
+ row = []
+ for j in range(0, x.shape[3], overlap_size):
+ tile = x[
+ :,
+ :,
+ i : i + self.tile_sample_min_size,
+ j : j + self.tile_sample_min_size,
+ ]
+ tile = self.encoder(tile)
+ if self.config.use_quant_conv:
+ tile = self.quant_conv(tile)
+ row.append(tile)
+ rows.append(row)
+ result_rows = []
+ for i, row in enumerate(rows):
+ result_row = []
+ for j, tile in enumerate(row):
+ # blend the above tile and the left tile
+ # to the current tile and add the current tile to the result row
+ if i > 0:
+ tile = self.blend_v(rows[i - 1][j], tile, blend_extent)
+ if j > 0:
+ tile = self.blend_h(row[j - 1], tile, blend_extent)
+ result_row.append(tile[:, :, :row_limit, :row_limit])
+ result_rows.append(torch.cat(result_row, dim=3))
+
+ moments = torch.cat(result_rows, dim=2)
+ posterior = DiagonalGaussianDistribution(moments)
+
+ if not return_dict:
+ return (posterior,)
+
+ return AutoencoderKLOutput(latent_dist=posterior)
+
+ def tiled_decode(
+ self, z: torch.Tensor, return_dict: bool = True
+ ) -> Union[DecoderOutput, torch.Tensor]:
+ r"""
+ Decode a batch of images using a tiled decoder.
+
+ Args:
+ z (`torch.Tensor`): Input batch of latent vectors.
+ return_dict (`bool`, *optional*, defaults to `True`):
+ Whether or not to return a [`~models.vae.DecoderOutput`] instead of a plain tuple.
+
+ Returns:
+ [`~models.vae.DecoderOutput`] or `tuple`:
+ If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is
+ returned.
+ """
+ overlap_size = int(self.tile_latent_min_size * (1 - self.tile_overlap_factor))
+ blend_extent = int(self.tile_sample_min_size * self.tile_overlap_factor)
+ row_limit = self.tile_sample_min_size - blend_extent
+
+ # Split z into overlapping 64x64 tiles and decode them separately.
+ # The tiles have an overlap to avoid seams between tiles.
+ rows = []
+ for i in range(0, z.shape[2], overlap_size):
+ row = []
+ for j in range(0, z.shape[3], overlap_size):
+ tile = z[
+ :,
+ :,
+ i : i + self.tile_latent_min_size,
+ j : j + self.tile_latent_min_size,
+ ]
+ if self.config.use_post_quant_conv:
+ tile = self.post_quant_conv(tile)
+ decoded = self.decoder(tile)
+ row.append(decoded)
+ rows.append(row)
+ result_rows = []
+ for i, row in enumerate(rows):
+ result_row = []
+ for j, tile in enumerate(row):
+ # blend the above tile and the left tile
+ # to the current tile and add the current tile to the result row
+ if i > 0:
+ tile = self.blend_v(rows[i - 1][j], tile, blend_extent)
+ if j > 0:
+ tile = self.blend_h(row[j - 1], tile, blend_extent)
+ result_row.append(tile[:, :, :row_limit, :row_limit])
+ result_rows.append(torch.cat(result_row, dim=3))
+
+ dec = torch.cat(result_rows, dim=2)
+ if not return_dict:
+ return (dec,)
+
+ return DecoderOutput(sample=dec)
+
+ def forward(
+ self,
+ sample: torch.Tensor,
+ sample_posterior: bool = False,
+ return_dict: bool = True,
+ generator: Optional[torch.Generator] = None,
+ ) -> Union[DecoderOutput, torch.Tensor]:
+ r"""
+ Args:
+ sample (`torch.Tensor`): Input sample.
+ sample_posterior (`bool`, *optional*, defaults to `False`):
+ Whether to sample from the posterior.
+ return_dict (`bool`, *optional*, defaults to `True`):
+ Whether or not to return a [`DecoderOutput`] instead of a plain tuple.
+ """
+ x = sample
+ posterior = self.encode(x).latent_dist
+ if sample_posterior:
+ z = posterior.sample(generator=generator)
+ else:
+ z = posterior.mode()
+ dec = self.decode(z).sample
+
+ if not return_dict:
+ return (dec,)
+
+ return DecoderOutput(sample=dec)
+
+
+EntryClass = AutoencoderKLFlux2
diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_kl_qwenimage.py b/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_kl_qwenimage.py
index 26d682f48..e860b313b 100644
--- a/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_kl_qwenimage.py
+++ b/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_kl_qwenimage.py
@@ -769,17 +769,6 @@ class AutoencoderKLQwenImage(nn.Module):
def __init__(
self,
config: QwenImageVAEConfig,
- # base_dim: int = 96,
- # z_dim: int = 16,
- # dim_mult: Tuple[int] = [1, 2, 4, 4],
- # num_res_blocks: int = 2,
- # attn_scales: List[float] = [],
- # temperal_downsample: List[bool] = [False, True, True],
- # dropout: float = 0.0,
- # latents_mean: List[float] = [-0.7571, -0.7089, -0.9113, 0.1075, -0.1745, 0.9653, -0.1517, 1.5508, 0.4134,
- # -0.0715, 0.5517, -0.3632, -0.1922, -0.9497, 0.2503, -0.2921],
- # latents_std: List[float] = [2.8184, 1.4541, 2.3275, 2.6558, 1.2196, 1.7708, 2.6052, 2.0743, 3.2687, 2.1526,
- # 2.8652, 1.5579, 1.6382, 1.1253, 2.8251, 1.9160],
) -> None:
# fmt: on
super().__init__()
@@ -961,8 +950,6 @@ class AutoencoderKLQwenImage(nn.Module):
h = self._encode(x)
posterior = DiagonalGaussianDistribution(h)
- if not return_dict:
- return (posterior,)
return posterior
def _decode(self, z: torch.Tensor, return_dict: bool = True):
diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/common.py b/python/sglang/multimodal_gen/runtime/models/vaes/common.py
index af1189d4f..58dad2507 100644
--- a/python/sglang/multimodal_gen/runtime/models/vaes/common.py
+++ b/python/sglang/multimodal_gen/runtime/models/vaes/common.py
@@ -10,6 +10,7 @@ from typing import Optional, cast
import numpy as np
import torch
import torch.distributed as dist
+from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution
from diffusers.utils.torch_utils import randn_tensor
from sglang.multimodal_gen.configs.models import VAEConfig
@@ -72,7 +73,7 @@ class ParallelTiledVAE(ABC):
def _decode(self, *args, **kwargs) -> torch.Tensor:
pass
- def encode(self, x: torch.Tensor) -> torch.Tensor:
+ def encode(self, x: torch.Tensor) -> DiagonalGaussianDistribution:
batch_size, num_channels, num_frames, height, width = x.shape
latent_num_frames = (num_frames - 1) // self.temporal_compression_ratio + 1
diff --git a/python/sglang/multimodal_gen/runtime/pipelines/flux_2.py b/python/sglang/multimodal_gen/runtime/pipelines/flux_2.py
new file mode 100644
index 000000000..92829fdba
--- /dev/null
+++ b/python/sglang/multimodal_gen/runtime/pipelines/flux_2.py
@@ -0,0 +1,118 @@
+# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
+# SPDX-License-Identifier: Apache-2.0
+
+from diffusers.image_processor import VaeImageProcessor
+
+from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline, Req
+from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
+ ComposedPipelineBase,
+)
+from sglang.multimodal_gen.runtime.pipelines_core.stages import (
+ ConditioningStage,
+ DecodingStage,
+ DenoisingStage,
+ ImageVAEEncodingStage,
+ InputValidationStage,
+ LatentPreparationStage,
+ TextEncodingStage,
+ TimestepPreparationStage,
+)
+from sglang.multimodal_gen.runtime.server_args import ServerArgs
+from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
+
+logger = init_logger(__name__)
+
+
+def compute_empirical_mu(batch: Req, server_args: ServerArgs):
+ num_steps = batch.num_inference_steps
+ image_seq_len = batch.raw_latent_shape[1]
+ a1, b1 = 8.73809524e-05, 1.89833333
+ a2, b2 = 0.00016927, 0.45666666
+
+ if image_seq_len > 4300:
+ mu = a2 * image_seq_len + b2
+ return "mu", float(mu)
+
+ m_200 = a2 * image_seq_len + b2
+ m_10 = a1 * image_seq_len + b1
+
+ a = (m_200 - m_10) / 190.0
+ b = m_200 - 200.0 * a
+ mu = a * num_steps + b
+
+ return "mu", float(mu)
+
+
+class Flux2Pipeline(LoRAPipeline, ComposedPipelineBase):
+ pipeline_name = "Flux2Pipeline"
+
+ _required_config_modules = [
+ "text_encoder",
+ "tokenizer",
+ "vae",
+ "transformer",
+ "scheduler",
+ ]
+
+ def create_pipeline_stages(self, server_args: ServerArgs):
+ """Set up pipeline stages with proper dependency injection."""
+
+ self.add_stage(
+ stage_name="input_validation_stage", stage=InputValidationStage()
+ )
+
+ self.add_stage(
+ stage_name="prompt_encoding_stage_primary",
+ stage=TextEncodingStage(
+ text_encoders=[
+ self.get_module("text_encoder"),
+ ],
+ tokenizers=[
+ self.get_module("tokenizer"),
+ ],
+ ),
+ )
+
+ self.add_stage(
+ stage_name="image_encoding_stage_primary",
+ stage=ImageVAEEncodingStage(
+ vae_image_processor=VaeImageProcessor(
+ vae_scale_factor=server_args.pipeline_config.vae_config.arch_config.vae_scale_factor
+ * 2
+ ),
+ vae=self.get_module("vae"),
+ ),
+ )
+
+ self.add_stage(stage_name="conditioning_stage", stage=ConditioningStage())
+
+ self.add_stage(
+ stage_name="latent_preparation_stage",
+ stage=LatentPreparationStage(
+ scheduler=self.get_module("scheduler"),
+ transformer=self.get_module("transformer"),
+ ),
+ )
+
+ self.add_stage(
+ stage_name="timestep_preparation_stage",
+ stage=TimestepPreparationStage(
+ scheduler=self.get_module("scheduler"),
+ prepare_extra_set_timesteps_kwargs=[compute_empirical_mu],
+ ),
+ )
+
+ self.add_stage(
+ stage_name="denoising_stage",
+ stage=DenoisingStage(
+ transformer=self.get_module("transformer"),
+ scheduler=self.get_module("scheduler"),
+ ),
+ )
+
+ self.add_stage(
+ stage_name="decoding_stage", stage=DecodingStage(vae=self.get_module("vae"))
+ )
+
+
+EntryClass = Flux2Pipeline
diff --git a/python/sglang/multimodal_gen/runtime/pipelines/qwen_image.py b/python/sglang/multimodal_gen/runtime/pipelines/qwen_image.py
index 458c00987..c87b72740 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines/qwen_image.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines/qwen_image.py
@@ -154,10 +154,6 @@ class QwenImageEditPipeline(LoRAPipeline, ComposedPipelineBase):
self.add_stage(
stage_name="image_encoding_stage_primary",
stage=ImageVAEEncodingStage(
- vae_image_processor=VaeImageProcessor(
- vae_scale_factor=server_args.pipeline_config.vae_config.arch_config.vae_scale_factor
- * 2
- ),
vae=self.get_module("vae"),
),
)
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py b/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py
index 8b9a7fe40..80d90c114 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py
@@ -55,7 +55,7 @@ class Req:
image_path: str | None = None
# Image encoder hidden states
image_embeds: list[torch.Tensor] = field(default_factory=list)
- pil_image: torch.Tensor | PIL.Image.Image | None = None
+ condition_image: torch.Tensor | PIL.Image.Image | None = None
pixel_values: torch.Tensor | PIL.Image.Image | None = None
preprocessed_image: torch.Tensor | None = None
@@ -93,9 +93,14 @@ class Req:
# Latent tensors
latents: torch.Tensor | None = None
+ # Flux-2
+ latent_ids: torch.Tensor | None = None
+
raw_latent_shape: torch.Tensor | None = None
noise_pred: torch.Tensor | None = None
- image_latent: torch.Tensor | None = None
+ # vae-encoded condition image
+ image_latent: torch.Tensor | list[torch.Tensor] | None = None
+ condition_image_latent_ids: torch.Tensor | list[torch.Tensor] | None = None
# Latent dimensions
height_latents: list[int] | int | None = None
@@ -212,8 +217,10 @@ class Req:
def adjust_size(self, server_args: ServerArgs):
if self.height is None or self.width is None:
- width, height = server_args.pipeline_config.adjust_size(
- self.width, self.height, self.pil_image
+ _image, width, height = (
+ server_args.pipeline_config.maybe_resize_condition_image(
+ self.width, self.height, self.condition_image
+ )
)
self.width = width
self.height = height
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py
index 5fa5bc16b..9a9e5fe0d 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py
@@ -9,11 +9,6 @@ import weakref
import torch
-from sglang.multimodal_gen.configs.models.vaes.base import VAEArchConfig
-from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
- QwenImageEditPipelineConfig,
- QwenImagePipelineConfig,
-)
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.loader.component_loader import VAELoader
from sglang.multimodal_gen.runtime.models.vaes.common import ParallelTiledVAE
@@ -64,36 +59,20 @@ class DecodingStage(PipelineStage):
# result.add_check("output", batch.output, [V.is_tensor, V.with_dims(5)])
return result
- def scale_and_shift(
- self, vae_arch_config: VAEArchConfig, latents: torch.Tensor, server_args
- ):
- # 1. scale
- is_qwen_image = isinstance(
- server_args.pipeline_config, QwenImagePipelineConfig
- ) or isinstance(server_args.pipeline_config, QwenImageEditPipelineConfig)
- if is_qwen_image:
- scaling_factor = 1.0 / torch.tensor(
- vae_arch_config.latents_std, device=latents.device
- ).view(1, vae_arch_config.z_dim, 1, 1, 1).to(latents.device, latents.dtype)
- else:
- scaling_factor = vae_arch_config.scaling_factor
+ def scale_and_shift(self, latents: torch.Tensor, server_args):
+ scaling_factor, shift_factor = (
+ server_args.pipeline_config.get_decode_scale_and_shift(
+ latents.device, latents.dtype, self.vae
+ )
+ )
+ # 1. scale
if isinstance(scaling_factor, torch.Tensor):
latents = latents / scaling_factor.to(latents.device, latents.dtype)
else:
latents = latents / scaling_factor
- # 2. shift
- if is_qwen_image:
- shift_factor = (
- torch.tensor(vae_arch_config.latents_mean)
- .view(1, vae_arch_config.z_dim, 1, 1, 1)
- .to(latents.device, latents.dtype)
- )
- else:
- shift_factor = getattr(vae_arch_config, "shift_factor", None)
-
- # Apply shifting if needed
+ # 2. apply shifting if needed
if shift_factor is not None:
if isinstance(shift_factor, torch.Tensor):
latents += shift_factor.to(latents.device, latents.dtype)
@@ -124,10 +103,10 @@ class DecodingStage(PipelineStage):
vae_autocast_enabled = (
vae_dtype != torch.float32
) and not server_args.disable_autocast
- vae_arch_config = server_args.pipeline_config.vae_config.arch_config
# scale and shift
- latents = self.scale_and_shift(vae_arch_config, latents, server_args)
+ latents = self.scale_and_shift(latents, server_args)
+ latents = server_args.pipeline_config.preprocess_decoding(latents)
# Decode latents
with torch.autocast(
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py
index e8d389870..6f498efd1 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py
@@ -275,14 +275,14 @@ class DenoisingStage(PipelineStage):
# FIXME: should probably move to latent preparation stage, to handle with offload
if (
server_args.pipeline_config.task_type == ModelTaskType.TI2V
- and batch.pil_image is not None
+ and batch.condition_image is not None
):
# Wan2.2 TI2V directly replaces the first frame of the latent with
# the image latent instead of appending along the channel dim
assert batch.image_latent is None, "TI2V task should not have image latents"
assert self.vae is not None, "VAE is not provided for TI2V task"
- self.vae = self.vae.to(batch.pil_image.device)
- z = self.vae.encode(batch.pil_image).mean.float()
+ self.vae = self.vae.to(batch.condition_image.device)
+ z = self.vae.encode(batch.condition_image).mean.float()
if self.vae.device != "cpu" and server_args.vae_cpu_offload:
self.vae = self.vae.to("cpu")
if hasattr(self.vae, "shift_factor") and self.vae.shift_factor is not None:
@@ -342,7 +342,7 @@ class DenoisingStage(PipelineStage):
# Shard z and reserved_frames_mask for TI2V if SP is enabled
if (
server_args.pipeline_config.task_type == ModelTaskType.TI2V
- and batch.pil_image is not None
+ and batch.condition_image is not None
and get_sp_world_size() > 1
):
sp_world_size = get_sp_world_size()
@@ -689,7 +689,7 @@ class DenoisingStage(PipelineStage):
# expand timestep
if (
server_args.pipeline_config.task_type == ModelTaskType.TI2V
- and batch.pil_image is not None
+ and batch.condition_image is not None
):
# Explicitly cast t_device to the target float type at the beginning.
# This ensures any precision-based rounding (e.g., float32(999.0) -> bfloat16(1000.0))
@@ -734,7 +734,7 @@ class DenoisingStage(PipelineStage):
"""
if (
server_args.pipeline_config.task_type == ModelTaskType.TI2V
- and batch.pil_image is not None
+ and batch.condition_image is not None
):
# Apply TI2V mask blending with SP-aware z and reserved_frames_mask.
# This ensures the first frame is always the condition image after each step.
@@ -1053,7 +1053,7 @@ class DenoisingStage(PipelineStage):
current_model: torch.nn.Module,
latent_model_input: torch.Tensor,
timestep,
- batch,
+ batch: Req,
timestep_index: int,
attn_metadata,
target_dtype,
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py
index 40112e68c..11bd59813 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py
@@ -9,7 +9,13 @@ This module contains implementations of image encoding stages for diffusion pipe
import PIL
import torch
+from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution
+from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType
+from sglang.multimodal_gen.configs.pipeline_configs.flux import (
+ Flux2PipelineConfig,
+ _prepare_image_ids,
+)
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
QwenImageEditPipelineConfig,
QwenImagePipelineConfig,
@@ -99,10 +105,12 @@ class ImageEncodingStage(PipelineStage):
The batch with encoded prompt embeddings.
"""
+ if batch.condition_image is None:
+ return batch
cuda_device = get_local_torch_device()
self.move_to_device(cuda_device)
- image = batch.pil_image
+ image = batch.condition_image
# preprocess via vae_image_processor
prompt_image = server_args.pipeline_config.preprocess_image(
@@ -177,9 +185,9 @@ class ImageEncodingStage(PipelineStage):
"""Verify image encoding stage inputs."""
result = VerificationResult()
if batch.debug:
- logger.debug(f"{batch.pil_image=}")
+ logger.debug(f"{batch.condition_image=}")
logger.debug(f"{batch.image_embeds=}")
- result.add_check("pil_image", batch.pil_image, V.not_none)
+ result.add_check("pil_image", batch.condition_image, V.not_none)
result.add_check("image_embeds", batch.image_embeds, V.is_list)
return result
@@ -217,10 +225,13 @@ class ImageVAEEncodingStage(PipelineStage):
Returns:
The batch with encoded outputs.
"""
- assert batch.pil_image is not None
+
+ if batch.condition_image is None:
+ return batch
+
if server_args.mode == ExecutionMode.INFERENCE:
- assert batch.pil_image is not None and isinstance(
- batch.pil_image, PIL.Image.Image
+ assert batch.condition_image is not None and isinstance(
+ batch.condition_image, PIL.Image.Image
)
assert batch.height is not None and isinstance(batch.height, int)
assert batch.width is not None and isinstance(batch.width, int)
@@ -229,8 +240,8 @@ class ImageVAEEncodingStage(PipelineStage):
width = batch.width
num_frames = batch.num_frames
elif server_args.mode == ExecutionMode.PREPROCESS:
- assert batch.pil_image is not None and isinstance(
- batch.pil_image, torch.Tensor
+ assert batch.condition_image is not None and isinstance(
+ batch.condition_image, torch.Tensor
)
assert batch.height is not None and isinstance(batch.height, list)
assert batch.width is not None and isinstance(batch.width, list)
@@ -241,10 +252,7 @@ class ImageVAEEncodingStage(PipelineStage):
self.vae = self.vae.to(get_local_torch_device())
- latent_height = height // self.vae.spatial_compression_ratio
- latent_width = width // self.vae.spatial_compression_ratio
-
- image = batch.pil_image
+ image = batch.condition_image
image = self.preprocess(
image,
vae_scale_factor=self.vae.spatial_compression_ratio,
@@ -255,19 +263,22 @@ class ImageVAEEncodingStage(PipelineStage):
# (B, C, H, W) -> (B, C, 1, H, W)
image = image.unsqueeze(2)
- video_condition = torch.cat(
- [
- image,
- image.new_zeros(
- image.shape[0],
- image.shape[1],
- num_frames - 1,
- image.shape[3],
- image.shape[4],
- ),
- ],
- dim=2,
- )
+ if num_frames == 1:
+ video_condition = image
+ else:
+ video_condition = torch.cat(
+ [
+ image,
+ image.new_zeros(
+ image.shape[0],
+ image.shape[1],
+ num_frames - 1,
+ image.shape[3],
+ image.shape[4],
+ ),
+ ],
+ dim=2,
+ )
video_condition = video_condition.to(
device=get_local_torch_device(), dtype=torch.float32
)
@@ -288,7 +299,9 @@ class ImageVAEEncodingStage(PipelineStage):
# self.vae.enable_parallel()
if not vae_autocast_enabled:
video_condition = video_condition.to(vae_dtype)
- encoder_output = self.vae.encode(video_condition)
+ encoder_output: DiagonalGaussianDistribution = self.vae.encode(
+ video_condition
+ )
if server_args.mode == ExecutionMode.PREPROCESS:
latent_condition = encoder_output.mean
@@ -296,29 +309,45 @@ class ImageVAEEncodingStage(PipelineStage):
generator = batch.generator
if generator is None:
raise ValueError("Generator must be provided")
- latent_condition = self.retrieve_latents(encoder_output, generator)
-
- # Apply shifting if needed
- if hasattr(self.vae, "shift_factor") and self.vae.shift_factor is not None:
- if isinstance(self.vae.shift_factor, torch.Tensor):
- latent_condition -= self.vae.shift_factor.to(
- latent_condition.device, latent_condition.dtype
- )
- else:
- latent_condition -= self.vae.shift_factor
-
- if isinstance(self.vae.scaling_factor, torch.Tensor):
- latent_condition = latent_condition * self.vae.scaling_factor.to(
- latent_condition.device, latent_condition.dtype
+ # TODO: verify
+ sample_mode = (
+ "argmax"
+ if server_args.pipeline_config.task_type == ModelTaskType.I2I
+ else "sample"
)
- else:
- latent_condition = latent_condition * self.vae.scaling_factor
+ latent_condition = self.retrieve_latents(
+ encoder_output, generator, sample_mode=sample_mode
+ )
+
+ latent_condition = self.server_args.pipeline_config.post_process_vae_encode(
+ latent_condition, self.vae
+ )
+
+ scaling_factor, shift_factor = (
+ self.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)
+
+ if isinstance(scaling_factor, torch.Tensor):
+ scaling_factor = scaling_factor.to(latent_condition.device)
+
+ latent_condition -= shift_factor
+ latent_condition = latent_condition * scaling_factor
+
+ batch_size = batch.batch_size
if server_args.mode == ExecutionMode.PREPROCESS:
batch.image_latent = latent_condition
else:
+ # TODO: abstract this
if isinstance(server_args.pipeline_config, QwenImageEditPipelineConfig):
- batch_size = batch.batch_size
if (
batch_size > latent_condition.shape[0]
and batch_size % latent_condition.shape[0] == 0
@@ -351,7 +380,31 @@ class ImageVAEEncodingStage(PipelineStage):
image_latent_height,
image_latent_width,
)
+ elif isinstance(server_args.pipeline_config, Flux2PipelineConfig):
+ # Pack each latent and concatenate
+ image_latents = [latent_condition]
+ # get image_latent_ids right after scale & shift
+ image_latent_ids = _prepare_image_ids(image_latents)
+ image_latent_ids = image_latent_ids.repeat(batch_size, 1, 1)
+ image_latent_ids = image_latent_ids.to(get_local_torch_device())
+ batch.condition_image_latent_ids = image_latent_ids
+
+ packed_latents = []
+ for latent in image_latents:
+ # latent: (1, 128, 32, 32)
+ packed = server_args.pipeline_config.maybe_pack_latents(
+ latent, None, None
+ ) # (1, 1024, 128)
+ packed = packed.squeeze(0) # (1024, 128) - remove batch dim
+ packed_latents.append(packed)
+
+ # Concatenate all reference tokens along sequence dimension
+ image_latents = torch.cat(packed_latents, dim=0) # (N*1024, 128)
+ image_latents = image_latents.unsqueeze(0) # (1, N*1024, 128)
+ image_latents = image_latents.repeat(batch_size, 1, 1)
else:
+ latent_height = height // self.vae.spatial_compression_ratio
+ latent_width = width // self.vae.spatial_compression_ratio
mask_lat_size = torch.ones(
1, 1, num_frames, latent_height, latent_width
)
@@ -388,7 +441,7 @@ class ImageVAEEncodingStage(PipelineStage):
def retrieve_latents(
self,
- encoder_output: torch.Tensor,
+ encoder_output: DiagonalGaussianDistribution,
generator: torch.Generator | None = None,
sample_mode: str = "sample",
):
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/input_validation.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/input_validation.py
index c92942e46..87381d384 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/input_validation.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/input_validation.py
@@ -11,9 +11,6 @@ from PIL import Image
from sglang.multimodal_gen.configs.pipeline_configs import WanI2V480PConfig
from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType
-from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
- QwenImageEditPipelineConfig,
-)
from sglang.multimodal_gen.runtime.models.vision_utils import load_image, load_video
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage
@@ -30,6 +27,7 @@ logger = init_logger(__name__)
# Alias for convenience
V = StageValidators
+
# TODO: since this might change sampling params after logging, should be do this beforehand?
@@ -117,23 +115,32 @@ class InputValidationStage(PipelineStage):
image = load_video(batch.image_path)[0]
else:
image = load_image(batch.image_path)
- batch.pil_image = image
+ batch.condition_image = image
+ condition_image_width, condition_image_height = image.width, image.height
+ else:
+ condition_image_width, condition_image_height = None, None
# NOTE: resizing needs to be bring in advance
- if isinstance(server_args.pipeline_config, QwenImageEditPipelineConfig):
- height = None if batch.height_not_provided else batch.height
- width = None if batch.width_not_provided else batch.width
- width, height = server_args.pipeline_config.adjust_size(
- height, width, batch.pil_image
- )
- batch.width = width
- batch.height = height
+ if server_args.pipeline_config.task_type == ModelTaskType.I2I:
+ if batch.condition_image is not None:
+ resized_image, resized_width, resized_height = (
+ server_args.pipeline_config.maybe_resize_condition_image(
+ condition_image_width,
+ condition_image_height,
+ batch.condition_image,
+ )
+ )
+ batch.condition_image = resized_image
+ batch.width = resized_width if batch.width_not_provided else batch.width
+ batch.height = (
+ resized_height if batch.height_not_provided else batch.height
+ )
elif (
server_args.pipeline_config.task_type == ModelTaskType.TI2V
- or server_args.pipeline_config.task_type == ModelTaskType.I2I
- ) and batch.pil_image is not None:
+ ) and batch.condition_image is not None:
+ # duplicate with vae_image_processor
# further processing for ti2v task
- img = batch.pil_image
+ img = batch.condition_image
ih, iw = img.height, img.width
patch_size = server_args.pipeline_config.dit_config.arch_config.patch_size
vae_stride = (
@@ -159,7 +166,7 @@ class InputValidationStage(PipelineStage):
batch.height = oh
batch.width = ow
# TODO: should we store in a new field: pixel values?
- batch.pil_image = img
+ batch.condition_image = img
if isinstance(server_args.pipeline_config, WanI2V480PConfig):
# TODO: could we merge with above?
@@ -173,7 +180,7 @@ class InputValidationStage(PipelineStage):
height = round(np.sqrt(max_area * aspect_ratio)) // mod_value * mod_value
width = round(np.sqrt(max_area / aspect_ratio)) // mod_value * mod_value
- batch.pil_image = batch.pil_image.resize((width, height))
+ batch.condition_image = batch.condition_image.resize((width, height))
batch.height = height
batch.width = width
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/latent_preparation.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/latent_preparation.py
index 4252da83b..acd901b3b 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/latent_preparation.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/latent_preparation.py
@@ -87,6 +87,12 @@ class LatentPreparationStage(PipelineStage):
latents = randn_tensor(
shape, generator=generator, device=device, dtype=dtype
)
+
+ latent_ids = server_args.pipeline_config.maybe_prepare_latent_ids(latents)
+
+ if latent_ids is not None:
+ batch.latent_ids = latent_ids.to(device=device)
+
latents = server_args.pipeline_config.maybe_pack_latents(
latents, batch_size, batch
)
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/text_encoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/text_encoding.py
index 47a5ef3c6..6423e45d2 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/text_encoding.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/text_encoding.py
@@ -11,6 +11,7 @@ import torch
from sglang.multimodal_gen.configs.models.encoders import BaseEncoderOutput
from sglang.multimodal_gen.configs.pipeline_configs import FluxPipelineConfig
+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.managers.forward_context import set_forward_context
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
@@ -239,9 +240,10 @@ class TextEncodingStage(PipelineStage):
else {}
)
- processed_texts: list[str] = []
+ processed_text_list: list[str] = []
for prompt_str in texts:
- processed_texts.append(preprocess_func(prompt_str))
+ preprocessed = preprocess_func(prompt_str)
+ processed_text_list.append(preprocessed)
# Prepare tokenizer args
tok_kwargs = self.prepare_tokenizer_kwargs(
@@ -249,10 +251,15 @@ class TextEncodingStage(PipelineStage):
**text_encoder_extra_arg,
)
- text_inputs = tokenizer(processed_texts, **tok_kwargs).to(target_device)
+ text_inputs: dict = server_args.pipeline_config.tokenize_prompt(
+ processed_text_list, tokenizer, tok_kwargs
+ ).to(target_device)
+
input_ids = text_inputs["input_ids"]
- is_flux = isinstance(server_args.pipeline_config, FluxPipelineConfig)
- is_flux_t5 = is_flux and i == 1
+ 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)
@@ -263,13 +270,14 @@ class TextEncodingStage(PipelineStage):
input_ids=input_ids,
attention_mask=attention_mask,
output_hidden_states=True,
+ use_cache=False,
)
prompt_embeds = postprocess_func(outputs, text_inputs)
if dtype is not None:
prompt_embeds = prompt_embeds.to(dtype=dtype)
embeds_list.append(prompt_embeds)
- if is_flux:
+ if is_flux_v1:
pooled_embeds_list.append(outputs.pooler_output)
if return_attention_mask:
attn_masks_list.append(attention_mask)
diff --git a/python/sglang/multimodal_gen/test/server/testcase_configs.py b/python/sglang/multimodal_gen/test/server/testcase_configs.py
index 018aab479..2bdc766ea 100644
--- a/python/sglang/multimodal_gen/test/server/testcase_configs.py
+++ b/python/sglang/multimodal_gen/test/server/testcase_configs.py
@@ -242,6 +242,19 @@ ONE_GPU_CASES_A: list[DiffusionTestCase] = [
output_size="1024x1024",
),
),
+ # DiffusionTestCase(
+ # "flux_2_image_t2i",
+ # DiffusionServerArgs(
+ # model_path="black-forest-labs/FLUX.2-dev",
+ # modality="image",
+ # warmup_text=1,
+ # warmup_edit=0,
+ # ),
+ # DiffusionSamplingParams(
+ # prompt="A futuristic cityscape at sunset with flying cars",
+ # output_size="1024x1024",
+ # ),
+ # ),
# === Text and Image to Image (TI2I) ===
# TODO: Timeout with Torch2.9. Add back when it can pass CI
# DiffusionTestCase(