[diffusion] model: support stable-diffusion-3-medium-diffusers (#19225)

Co-authored-by: zhaochenyang20 <zhaochen20@outlook.com>
Co-authored-by: Kangrui Du <kangruidu@gmail.com>
Co-authored-by: Xiaole Guo <gxlvera@gmail.com>
This commit is contained in:
Xiaole Guo
2026-04-13 16:07:06 +08:00
committed by GitHub
co-authored by zhaochenyang20 Kangrui Du Xiaole Guo
parent 9e6d1c066e
commit 4df60434d7
17 changed files with 800 additions and 29 deletions
@@ -5,6 +5,9 @@ from sglang.multimodal_gen.configs.models.dits.hunyuan3d import Hunyuan3DDiTConf
from sglang.multimodal_gen.configs.models.dits.hunyuanvideo import HunyuanVideoConfig from sglang.multimodal_gen.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
from sglang.multimodal_gen.configs.models.dits.mova_audio import MOVAAudioConfig from sglang.multimodal_gen.configs.models.dits.mova_audio import MOVAAudioConfig
from sglang.multimodal_gen.configs.models.dits.mova_video import MOVAVideoConfig from sglang.multimodal_gen.configs.models.dits.mova_video import MOVAVideoConfig
from sglang.multimodal_gen.configs.models.dits.stablediffusion3 import (
StableDiffusion3TransformerConfig,
)
from sglang.multimodal_gen.configs.models.dits.wanvideo import WanVideoConfig from sglang.multimodal_gen.configs.models.dits.wanvideo import WanVideoConfig
__all__ = [ __all__ = [
@@ -14,4 +17,5 @@ __all__ = [
"Hunyuan3DDiTConfig", "Hunyuan3DDiTConfig",
"MOVAAudioConfig", "MOVAAudioConfig",
"MOVAVideoConfig", "MOVAVideoConfig",
"StableDiffusion3TransformerConfig",
] ]
@@ -0,0 +1,37 @@
# SPDX-License-Identifier: Apache-2.0
"""StableDiffusion3 Transformer model configuration"""
from dataclasses import dataclass, field
from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig
@dataclass
class StableDiffusion3TransformerArchConfig(DiTArchConfig):
"""Architecture configuration for StableDiffusion3 Transformer, applicable to SD3-medium, SD3.5-medium, SD3.5-large."""
sample_size: int = 128
patch_size: int = 2
in_channels: int = 16
out_channels: int = 16
num_layers: int = 18
attention_head_dim: int = 64
num_attention_heads: int = 18
cross_attention_dim: int = 4096
joint_attention_dim: int = 4096
caption_projection_dim: int = 1152
pooled_projection_dim: int = 2048
pos_embed_max_size: int = 96
dual_attention_layers: tuple[int, ...] = ()
qk_norm: str | None = None
_class_name: str = "SD3Transformer2DModel"
@dataclass
class StableDiffusion3TransformerConfig(DiTConfig):
"""Configuration for StableDiffusion3 Transformer model."""
arch_config: StableDiffusion3TransformerArchConfig = field(
default_factory=StableDiffusion3TransformerArchConfig
)
@@ -3,11 +3,15 @@
from sglang.multimodal_gen.configs.models.vaes.dac import DacVAEConfig from sglang.multimodal_gen.configs.models.vaes.dac import DacVAEConfig
from sglang.multimodal_gen.configs.models.vaes.hunyuan3d import Hunyuan3DVAEConfig from sglang.multimodal_gen.configs.models.vaes.hunyuan3d import Hunyuan3DVAEConfig
from sglang.multimodal_gen.configs.models.vaes.hunyuanvae import HunyuanVAEConfig from sglang.multimodal_gen.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
from sglang.multimodal_gen.configs.models.vaes.stablediffusion3 import (
StableDiffusion3VAEConfig,
)
from sglang.multimodal_gen.configs.models.vaes.wanvae import WanVAEConfig from sglang.multimodal_gen.configs.models.vaes.wanvae import WanVAEConfig
__all__ = [ __all__ = [
"DacVAEConfig", "DacVAEConfig",
"HunyuanVAEConfig", "HunyuanVAEConfig",
"StableDiffusion3VAEConfig",
"WanVAEConfig", "WanVAEConfig",
"Hunyuan3DVAEConfig", "Hunyuan3DVAEConfig",
] ]
@@ -0,0 +1,72 @@
# SPDX-License-Identifier: Apache-2.0
"""StableDiffusion3 VAE configuration."""
from dataclasses import dataclass, field
from sglang.multimodal_gen.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class StableDiffusion3VAEArchConfig(VAEArchConfig):
"""Architecture configuration for StableDiffusion3 VAE."""
scaling_factor: float = 1.5305
shift_factor: float = 0.0609
spatial_compression_ratio: int = 8
temporal_compression_ratio: int = 1
in_channels: int = 3
out_channels: int = 3
latent_channels: int = 16
sample_size: int = 128
block_out_channels: tuple[int, ...] = (128, 256, 512, 512)
layers_per_block: int = 2
act_fn: str = "silu"
norm_num_groups: int = 32
down_block_types: tuple[str, ...] = (
"DownEncoderBlock2D",
"DownEncoderBlock2D",
"DownEncoderBlock2D",
"DownEncoderBlock2D",
)
up_block_types: tuple[str, ...] = (
"UpDecoderBlock2D",
"UpDecoderBlock2D",
"UpDecoderBlock2D",
"UpDecoderBlock2D",
)
attention_head_dim: int = 8
mid_block_add_attention: bool = True
use_quant_conv: bool = False
use_post_quant_conv: bool = False
@dataclass
class StableDiffusion3VAEConfig(VAEConfig):
"""Configuration for StableDiffusion3 VAE."""
arch_config: StableDiffusion3VAEArchConfig = field(
default_factory=StableDiffusion3VAEArchConfig
)
tile_sample_min_height: int = 512
tile_sample_min_width: int = 512
tile_sample_min_num_frames: int = 1
tile_sample_stride_height: int = 448
tile_sample_stride_width: int = 448
tile_sample_stride_num_frames: int = 1
use_tiling: bool = True
use_temporal_tiling: bool = False
use_parallel_tiling: bool = True
use_temporal_scaling_frames: bool = False
def __post_init__(self) -> None:
"""Post initialization for SD3 VAE specific setup."""
super().__post_init__()
self.update_model_arch({"_class_name": "AutoencoderKL"})
self.blend_num_frames = 0
@@ -30,6 +30,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.hunyuan3d import (
from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import LTX2PipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import LTX2PipelineConfig
from sglang.multimodal_gen.configs.pipeline_configs.mova import MOVAPipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.mova import MOVAPipelineConfig
from sglang.multimodal_gen.configs.pipeline_configs.sana import SanaPipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.sana import SanaPipelineConfig
from sglang.multimodal_gen.configs.pipeline_configs.stablediffusion3 import (
StableDiffusion3PipelineConfig,
)
from sglang.multimodal_gen.configs.pipeline_configs.wan import ( from sglang.multimodal_gen.configs.pipeline_configs.wan import (
SelfForcingWanT2V480PConfig, SelfForcingWanT2V480PConfig,
WanI2V480PConfig, WanI2V480PConfig,
@@ -55,6 +58,7 @@ __all__ = [
"SanaPipelineConfig", "SanaPipelineConfig",
"SlidingTileAttnConfig", "SlidingTileAttnConfig",
"MOVAPipelineConfig", "MOVAPipelineConfig",
"StableDiffusion3PipelineConfig",
"WanT2V480PConfig", "WanT2V480PConfig",
"WanI2V480PConfig", "WanI2V480PConfig",
"WanT2V720PConfig", "WanT2V720PConfig",
@@ -459,6 +459,33 @@ class PipelineConfig:
sharded_tensor = sharded_tensor[:, :, rank_in_sp_group, :, :, :] sharded_tensor = sharded_tensor[:, :, rank_in_sp_group, :, :, :]
return sharded_tensor, True return sharded_tensor, True
def get_text_encoder_attention_mask(
self, text_inputs: dict, encoder_index: int
) -> "torch.Tensor | None":
"""Return the attention mask for the given text encoder.
Override to suppress (return None) or modify the mask per model.
"""
return text_inputs.get("attention_mask")
def get_text_encoder_pooler_output(
self, outputs: "BaseEncoderOutput", encoder_index: int
) -> "torch.Tensor | None":
"""Return the pooler output for the given text encoder, or None to skip.
Override for models that need pooled embeddings (e.g. FLUX v1, SD3).
"""
return None
def select_vae_weight_files(
self,
safetensors_list: list[str],
component_model_path: str,
component_name: str,
vae_precision: str,
) -> list[str]:
return safetensors_list
def get_pos_prompt_embeds(self, batch): def get_pos_prompt_embeds(self, batch):
return batch.prompt_embeds return batch.prompt_embeds
@@ -81,6 +81,13 @@ class FluxPipelineConfig(ImagePipelineConfig):
] ]
) )
def get_text_encoder_attention_mask(self, text_inputs, encoder_index):
# Flux v1 does not use attention masks for text encoders.
return None
def get_text_encoder_pooler_output(self, outputs, encoder_index):
return outputs.pooler_output
def prepare_sigmas(self, sigmas, num_inference_steps): def prepare_sigmas(self, sigmas, num_inference_steps):
return self._prepare_sigmas(sigmas, num_inference_steps) return self._prepare_sigmas(sigmas, num_inference_steps)
@@ -391,6 +398,14 @@ class Flux2PipelineConfig(FluxPipelineConfig):
] ]
) )
def get_text_encoder_attention_mask(self, text_inputs, encoder_index):
# Flux2 uses standard attention masks (unlike Flux v1).
return text_inputs.get("attention_mask")
def get_text_encoder_pooler_output(self, outputs, encoder_index):
# Flux2 does not use pooler output.
return None
def tokenize_prompt(self, prompts: list[str], tokenizer, tok_kwargs) -> dict: def tokenize_prompt(self, prompts: list[str], tokenizer, tok_kwargs) -> dict:
messages = build_flux2_text_messages(prompts) messages = build_flux2_text_messages(prompts)
inputs = tokenizer.apply_chat_template( inputs = tokenizer.apply_chat_template(
@@ -0,0 +1,202 @@
# SPDX-License-Identifier: Apache-2.0
"""Stable Diffusion 3 pipeline configuration."""
import os
from dataclasses import dataclass, field
from typing import Callable
import torch
from sglang.multimodal_gen.configs.models import DiTConfig, EncoderConfig, VAEConfig
from sglang.multimodal_gen.configs.models.dits import StableDiffusion3TransformerConfig
from sglang.multimodal_gen.configs.models.encoders import BaseEncoderOutput
from sglang.multimodal_gen.configs.models.encoders.base import TextEncoderArchConfig
from sglang.multimodal_gen.configs.models.encoders.clip import (
CLIPTextArchConfig,
CLIPTextConfig,
)
from sglang.multimodal_gen.configs.models.encoders.t5 import (
T5ArchConfig,
T5Config,
)
from sglang.multimodal_gen.configs.models.vaes.stablediffusion3 import (
StableDiffusion3VAEConfig,
)
from sglang.multimodal_gen.configs.pipeline_configs.base import (
ModelTaskType,
SpatialImagePipelineConfig,
)
def sd3_clip_postprocess_text(outputs: BaseEncoderOutput, _text_inputs) -> torch.Tensor:
"""Extract pre-final hidden state for SD3 CLIP encoders."""
if outputs.hidden_states is None:
raise ValueError(
"SD3 CLIP postprocessing requires hidden_states from encoder output."
)
return outputs.hidden_states[-2]
def t5_postprocess_text(outputs: BaseEncoderOutput, _text_inputs) -> torch.Tensor:
return outputs.last_hidden_state
def select_sd3_vae_weight_files(
safetensors_list: list[str],
component_model_path: str,
component_name: str,
vae_precision: str,
) -> list[str]:
"""Select SD3 VAE checkpoint file candidates with minimal policy."""
if component_name not in ("vae", "video_vae"):
return safetensors_list
base_name = "diffusion_pytorch_model"
if vae_precision == "fp16":
fp16_path = os.path.join(component_model_path, f"{base_name}.fp16.safetensors")
if os.path.exists(fp16_path):
return [fp16_path]
full_path = os.path.join(component_model_path, f"{base_name}.safetensors")
if os.path.exists(full_path):
return [full_path]
return safetensors_list
@dataclass
class SD3CLIPTextArchConfig(CLIPTextArchConfig):
def __post_init__(self) -> None:
super().__post_init__()
self.tokenizer_kwargs.update(
{
"max_length": self.text_len,
"padding": "max_length",
}
)
@dataclass
class SD3CLIPTextConfig(CLIPTextConfig):
arch_config: TextEncoderArchConfig = field(default_factory=SD3CLIPTextArchConfig)
@dataclass
class SD3T5ArchConfig(T5ArchConfig):
def __post_init__(self) -> None:
super().__post_init__()
self.tokenizer_kwargs.update({"max_length": 256})
@dataclass
class SD3T5Config(T5Config):
arch_config: TextEncoderArchConfig = field(default_factory=SD3T5ArchConfig)
@dataclass
class StableDiffusion3PipelineConfig(SpatialImagePipelineConfig):
"""Configuration for SD3 image generation pipeline.
This config intentionally relies on SD3-specific encoder configs to provide
tokenizer kwargs, instead of stage-level tokenizer overrides.
"""
task_type: ModelTaskType = ModelTaskType.T2I
dit_config: DiTConfig = field(default_factory=StableDiffusion3TransformerConfig)
vae_config: VAEConfig = field(default_factory=StableDiffusion3VAEConfig)
text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (
SD3CLIPTextConfig(),
SD3CLIPTextConfig(),
SD3T5Config(),
)
)
text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("fp16", "fp16", "fp32")
)
preprocess_text_funcs: tuple[Callable[[str], str] | None, ...] = field(
default_factory=lambda: (
None,
None,
None,
)
)
postprocess_text_funcs: tuple[
Callable[[BaseEncoderOutput, dict], torch.Tensor], ...
] = field(
default_factory=lambda: (
sd3_clip_postprocess_text,
sd3_clip_postprocess_text,
t5_postprocess_text,
)
)
should_use_guidance: bool = False
guidance_scale: float = 7.0
def __post_init__(self) -> None:
configs = list(self.text_encoder_configs)
configs[0].update_model_arch({"_class_name": "CLIPTextModelWithProjection"})
configs[1].update_model_arch({"_class_name": "CLIPTextModelWithProjection"})
configs[2].update_model_arch({"_class_name": "T5EncoderModel"})
self.text_encoder_configs = tuple(configs)
def get_text_encoder_pooler_output(self, outputs, encoder_index):
# SD3 uses pooled embeddings only from the two CLIP encoders (indices 0 and 1).
if encoder_index <= 1:
return outputs.pooler_output
return None
def select_vae_weight_files(
self,
safetensors_list: list[str],
component_model_path: str,
component_name: str,
vae_precision: str,
) -> list[str]:
return select_sd3_vae_weight_files(
safetensors_list=safetensors_list,
component_model_path=component_model_path,
component_name=component_name,
vae_precision=vae_precision,
)
def tokenize_prompt(self, prompt: list[str], tokenizer, tok_kwargs) -> dict:
text_inputs = tokenizer(prompt, **tok_kwargs)
text_inputs["attention_mask"] = None
return text_inputs
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 prepare_pos_cond_kwargs(self, batch, device, rotary_emb, dtype):
return {
"pooled_projections": (
batch.pooled_embeds[0] if batch.pooled_embeds else None
)
}
def prepare_neg_cond_kwargs(self, batch, device, rotary_emb, dtype):
return {
"pooled_projections": (
batch.neg_pooled_embeds[0] if batch.neg_pooled_embeds else None
)
}
# SD3 image latents are spatial (B, C, H, W), not video-like (B, C, T, H, W).
def prepare_latent_shape(self, batch, batch_size, num_frames): # noqa: ARG002
spatial_ratio = self.vae_config.arch_config.spatial_compression_ratio
in_channels = self.dit_config.arch_config.in_channels
return (
batch_size,
in_channels,
batch.height // spatial_ratio,
batch.width // spatial_ratio,
)
@@ -0,0 +1,20 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# SPDX-License-Identifier: Apache-2.0
"""StableDiffusion3 sampling parameters configuration."""
from dataclasses import dataclass
from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
@dataclass
class StableDiffusion3SamplingParams(SamplingParams):
"""Sampling parameters for StableDiffusion3."""
# A single space ensures tokenizers produce valid (non-empty) input for CFG.
negative_prompt: str = " "
num_frames: int = 1
num_inference_steps: int = 50
guidance_scale: float = 7.0
guidance_rescale: float = 0.0
+26
View File
@@ -68,6 +68,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
QwenImagePipelineConfig, QwenImagePipelineConfig,
) )
from sglang.multimodal_gen.configs.pipeline_configs.sana import SanaPipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.sana import SanaPipelineConfig
from sglang.multimodal_gen.configs.pipeline_configs.stablediffusion3 import (
StableDiffusion3PipelineConfig,
)
from sglang.multimodal_gen.configs.pipeline_configs.wan import ( from sglang.multimodal_gen.configs.pipeline_configs.wan import (
FastWan2_1_T2V_480P_Config, FastWan2_1_T2V_480P_Config,
FastWan2_2_TI2V_5B_Config, FastWan2_2_TI2V_5B_Config,
@@ -109,6 +112,9 @@ from sglang.multimodal_gen.configs.sample.qwenimage import (
QwenImageSamplingParams, QwenImageSamplingParams,
) )
from sglang.multimodal_gen.configs.sample.sana import SanaSamplingParams from sglang.multimodal_gen.configs.sample.sana import SanaSamplingParams
from sglang.multimodal_gen.configs.sample.stablediffusion3 import (
StableDiffusion3SamplingParams,
)
from sglang.multimodal_gen.configs.sample.wan import ( from sglang.multimodal_gen.configs.sample.wan import (
FastWanT2V480PConfig, FastWanT2V480PConfig,
Turbo_Wan2_2_I2V_A14B_SamplingParam, Turbo_Wan2_2_I2V_A14B_SamplingParam,
@@ -849,6 +855,26 @@ def _register_configs():
hf_model_paths=["Qwen/Qwen-Image-Layered"], hf_model_paths=["Qwen/Qwen-Image-Layered"],
model_detectors=[lambda hf_id: "qwen-image-layered" in hf_id.lower()], model_detectors=[lambda hf_id: "qwen-image-layered" in hf_id.lower()],
) )
register_configs(
sampling_param_cls=StableDiffusion3SamplingParams,
pipeline_config_cls=StableDiffusion3PipelineConfig,
hf_model_paths=[
"stabilityai/stable-diffusion-3-medium",
"stabilityai/stable-diffusion-3-medium-diffusers",
"stabilityai/stable-diffusion-3.5-medium",
"stabilityai/stable-diffusion-3.5-medium-diffusers",
"stabilityai/stable-diffusion-3.5-large",
"stabilityai/stable-diffusion-3.5-large-diffusers",
],
model_detectors=[
lambda hf_id: "stable-diffusion-3-medium" in hf_id.lower()
or "stable-diffusion-3.5-medium" in hf_id.lower()
or "stable-diffusion-3.5-large" in hf_id.lower()
or "sd3-medium" in hf_id.lower()
or "sd3.5-medium" in hf_id.lower()
or "sd3.5-large" in hf_id.lower()
],
)
register_configs( register_configs(
sampling_param_cls=GlmImageSamplingParams, sampling_param_cls=GlmImageSamplingParams,
@@ -1,6 +1,7 @@
import dataclasses import dataclasses
import glob import glob
import os import os
import re
from collections.abc import Generator, Iterable from collections.abc import Generator, Iterable
from typing import cast from typing import cast
@@ -218,21 +219,21 @@ class TextEncoderLoader(ComponentLoader):
component_path=component_model_path component_path=component_model_path
) )
def is_not_first_encoder(module_name):
return "2" in module_name
# TODO(mick): had to throw an exception for different text-encoder arch # TODO(mick): had to throw an exception for different text-encoder arch
if not is_not_first_encoder(component_name): encoder_index = self._extract_encoder_index(component_name)
encoder_config = server_args.pipeline_config.text_encoder_configs[0] assert encoder_index < len(
encoder_config.update_model_arch(model_config) server_args.pipeline_config.text_encoder_configs
) and encoder_index < len(server_args.pipeline_config.text_encoder_precisions)
encoder_config = server_args.pipeline_config.text_encoder_configs[encoder_index]
encoder_config.update_model_arch(model_config)
if encoder_index == 0:
for key, value in diffusers_pretrained_config.__dict__.items(): for key, value in diffusers_pretrained_config.__dict__.items():
setattr(encoder_config.arch_config, key, value) setattr(encoder_config.arch_config, key, value)
encoder_dtype = server_args.pipeline_config.text_encoder_precisions[0] encoder_dtype = server_args.pipeline_config.text_encoder_precisions[
else: encoder_index
assert len(server_args.pipeline_config.text_encoder_configs) == 2 ]
encoder_config = server_args.pipeline_config.text_encoder_configs[1]
encoder_config.update_model_arch(model_config)
encoder_dtype = server_args.pipeline_config.text_encoder_precisions[1]
# TODO(will): add support for other dtypes # TODO(will): add support for other dtypes
return self.load_model( return self.load_model(
component_model_path, component_model_path,
@@ -242,6 +243,28 @@ class TextEncoderLoader(ComponentLoader):
cpu_offload_flag=cpu_offload_flag, cpu_offload_flag=cpu_offload_flag,
) )
@staticmethod
def _extract_encoder_index(component_name: str) -> int:
"""
Map text encoder component names to zero-based indices.
Examples:
- text_encoder -> 0
- text_encoder_2 -> 1
- text_encoder_3 -> 2
"""
match = re.search(r"_(\d+)$", component_name)
if match is None:
return 0
suffix_num = int(match.group(1))
if suffix_num <= 0:
raise ValueError(
f"Invalid text encoder component name '{component_name}': "
"numeric suffix must be >= 1."
)
return suffix_num - 1
def load_model( def load_model(
self, self,
model_path: str, model_path: str,
@@ -129,6 +129,13 @@ class VAELoader(ComponentLoader):
vae = vae_cls(vae_config).to(target_device) vae = vae_cls(vae_config).to(target_device)
safetensors_list = _list_safetensors_files(component_model_path) safetensors_list = _list_safetensors_files(component_model_path)
safetensors_list = server_args.pipeline_config.select_vae_weight_files(
safetensors_list=safetensors_list,
component_model_path=component_model_path,
component_name=component_name,
vae_precision=vae_precision,
)
assert ( assert (
len(safetensors_list) >= 1 len(safetensors_list) >= 1
), f"Found no safetensors files in {component_model_path}" ), f"Found no safetensors files in {component_model_path}"
@@ -181,9 +181,7 @@ class skip_init_modules:
def _normalize_component_type(module_type: str) -> str: def _normalize_component_type(module_type: str) -> str:
"""Normalize module types like 'text_encoder_2' -> 'text_encoder'.""" """Normalize module types like 'text_encoder_2' -> 'text_encoder'."""
if module_type.endswith("_2"): return re.sub(r"_\d+$", "", module_type)
return module_type[:-2]
return module_type
def _clean_hf_config_inplace(model_config: dict) -> None: def _clean_hf_config_inplace(model_config: dict) -> None:
@@ -0,0 +1,180 @@
# SPDX-License-Identifier: Apache-2.0
"""StableDiffusion3 Transformer model implementation.
NOTE: This initial implementation uses diffusers' JointTransformerBlock directly.
A native SGLang attention implementation is needed for FlashAttention, TP/SP,
quantization, and LoRA support.
"""
from typing import Any
import torch
import torch.nn as nn
from diffusers.models.attention import JointTransformerBlock
from diffusers.models.embeddings import CombinedTimestepTextProjEmbeddings, PatchEmbed
from diffusers.models.normalization import AdaLayerNormContinuous
from sglang.multimodal_gen.configs.models.dits.stablediffusion3 import (
StableDiffusion3TransformerConfig,
)
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__)
class SD3Transformer2DModel(CachableDiT):
_supports_gradient_checkpointing = True
_no_split_modules = ["JointTransformerBlock"]
_skip_layerwise_casting_patterns = ["pos_embed", "norm"]
def __init__(
self,
config: StableDiffusion3TransformerConfig,
hf_config: dict[str, Any] | None = None,
quant_config=None,
):
super().__init__(config=config, hf_config=hf_config)
self.config = config
arch_config = config.arch_config
sample_size = arch_config.sample_size
patch_size = arch_config.patch_size
in_channels = arch_config.in_channels
num_layers = arch_config.num_layers
attention_head_dim = arch_config.attention_head_dim
num_attention_heads = arch_config.num_attention_heads
joint_attention_dim = arch_config.joint_attention_dim
caption_projection_dim = arch_config.caption_projection_dim
pooled_projection_dim = arch_config.pooled_projection_dim
out_channels = arch_config.out_channels
pos_embed_max_size = arch_config.pos_embed_max_size
dual_attention_layers = arch_config.dual_attention_layers
qk_norm = arch_config.qk_norm
self.out_channels = out_channels if out_channels is not None else in_channels
self.inner_dim = num_attention_heads * attention_head_dim
self.patch_size = patch_size
self.pos_embed = PatchEmbed(
height=sample_size,
width=sample_size,
patch_size=patch_size,
in_channels=in_channels,
embed_dim=self.inner_dim,
pos_embed_max_size=pos_embed_max_size,
)
self.time_text_embed = CombinedTimestepTextProjEmbeddings(
embedding_dim=self.inner_dim, pooled_projection_dim=pooled_projection_dim
)
self.context_embedder = nn.Linear(joint_attention_dim, caption_projection_dim)
self.transformer_blocks = nn.ModuleList(
[
JointTransformerBlock(
dim=self.inner_dim,
num_attention_heads=num_attention_heads,
attention_head_dim=attention_head_dim,
context_pre_only=i == num_layers - 1,
qk_norm=qk_norm,
use_dual_attention=i in dual_attention_layers,
)
for i in range(num_layers)
]
)
self.norm_out = AdaLayerNormContinuous(
self.inner_dim, self.inner_dim, elementwise_affine=False, eps=1e-6
)
self.proj_out = nn.Linear(
self.inner_dim, patch_size * patch_size * self.out_channels, bias=True
)
self.gradient_checkpointing = False
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | None = None,
pooled_projections: torch.Tensor | None = None,
timestep: torch.LongTensor | None = None,
block_controlnet_hidden_states: list | None = None,
guidance: torch.Tensor | None = None,
joint_attention_kwargs: dict[str, Any] | None = None,
skip_layers: list[int] | None = None,
) -> torch.Tensor:
if encoder_hidden_states is None:
raise ValueError("encoder_hidden_states must be provided.")
if pooled_projections is None:
raise ValueError("pooled_projections must be provided.")
encoder_embeddings = encoder_hidden_states
height, width = hidden_states.shape[-2:]
hidden_states = self.pos_embed(hidden_states)
temb = self.time_text_embed(timestep, pooled_projections)
encoder_embeddings = self.context_embedder(encoder_embeddings)
skip_layer_set = set(skip_layers) if skip_layers else set()
if block_controlnet_hidden_states is not None:
interval_control = len(self.transformer_blocks) / len(
block_controlnet_hidden_states
)
else:
interval_control = 0
for index_block, block in enumerate(self.transformer_blocks):
if index_block not in skip_layer_set:
encoder_embeddings, hidden_states = block(
hidden_states=hidden_states,
encoder_hidden_states=encoder_embeddings,
temb=temb,
joint_attention_kwargs=joint_attention_kwargs,
)
# controlnet residual
if (
block_controlnet_hidden_states is not None
and block.context_pre_only is False
):
hidden_states = (
hidden_states
+ block_controlnet_hidden_states[
int(index_block / interval_control)
]
)
hidden_states = self.norm_out(hidden_states, temb)
hidden_states = self.proj_out(hidden_states)
# unpatchify
patch_size = self.patch_size
height = height // patch_size
width = width // patch_size
hidden_states = hidden_states.reshape(
shape=(
hidden_states.shape[0],
height,
width,
patch_size,
patch_size,
self.out_channels,
)
)
hidden_states = hidden_states.permute(0, 5, 1, 3, 2, 4)
output = hidden_states.reshape(
shape=(
hidden_states.shape[0],
self.out_channels,
height * patch_size,
width * patch_size,
)
)
return output
# Entry class for registry
EntryClass = SD3Transformer2DModel
@@ -594,6 +594,48 @@ class CLIPTextModel(TextEncoder):
return loaded_params return loaded_params
class CLIPTextModelWithProjection(CLIPTextModel):
"""
CLIP text encoder with projection head for pooled_output.
"""
def __init__(
self,
config: CLIPTextConfig,
) -> None:
super().__init__(config)
self.text_projection = nn.Linear(
config.hidden_size, config.projection_dim, bias=False
)
def forward(
self,
input_ids: torch.Tensor | None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
**kwargs,
) -> BaseEncoderOutput:
outputs: BaseEncoderOutput = self.text_model(
input_ids=input_ids,
attention_mask=attention_mask,
position_ids=position_ids,
output_hidden_states=output_hidden_states,
)
pooled_output = outputs.pooler_output
if pooled_output is not None:
pooled_output = self.text_projection(pooled_output)
return BaseEncoderOutput(
last_hidden_state=outputs.last_hidden_state,
pooler_output=pooled_output,
hidden_states=outputs.hidden_states,
attentions=outputs.attentions,
)
class CLIPVisionTransformer(nn.Module): class CLIPVisionTransformer(nn.Module):
def __init__( def __init__(
@@ -759,4 +801,4 @@ class BertModel(CLIPTextModel):
pass pass
EntryClass = [CLIPTextModel, CLIPVisionModel] EntryClass = [CLIPTextModel, CLIPTextModelWithProjection, CLIPVisionModel]
@@ -0,0 +1,111 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# SPDX-License-Identifier: Apache-2.0
"""StableDiffusion3 pipeline implementation."""
import torch
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
InputValidationStage,
PipelineStage,
TextEncodingStage,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
class SD3ConditioningStage(PipelineStage):
"""Merge CLIP-T, CLIP-G and T5 embeddings into unified prompt/pooled tensors."""
@torch.no_grad()
def forward(self, batch: Req, server_args: ServerArgs) -> Req:
batch.prompt_embeds, batch.pooled_embeds = self._merge(
batch.prompt_embeds, batch.pooled_embeds
)
if batch.do_classifier_free_guidance:
batch.negative_prompt_embeds, batch.neg_pooled_embeds = self._merge(
batch.negative_prompt_embeds, batch.neg_pooled_embeds
)
return batch
@staticmethod
def _merge(
embeds_list: list[torch.Tensor],
pooled_list: list[torch.Tensor],
) -> tuple[list[torch.Tensor], list[torch.Tensor]]:
"""Merge 3 encoder outputs into unified prompt/pooled tensors.
SD3-medium uses exactly 3 text encoders (CLIP-L, CLIP-G, T5).
Returns single-element lists to match the batch field format expected
by downstream stages (get_pos_prompt_embeds accesses index [0]).
"""
if len(embeds_list) != 3:
raise ValueError(
f"SD3 requires exactly 3 prompt embedding tensors, got {len(embeds_list)}."
)
if len(pooled_list) < 2:
raise ValueError(
f"SD3 requires at least 2 pooled embedding tensors, got {len(pooled_list)}."
)
clipt, clipg, t5 = embeds_list
clip_merged = torch.cat([clipt, clipg], dim=-1)
clip_merged = torch.nn.functional.pad(
clip_merged, (0, t5.shape[-1] - clip_merged.shape[-1])
)
merged_embeds = [torch.cat([clip_merged, t5], dim=-2)]
merged_pooled = [torch.cat([pooled_list[0], pooled_list[1]], dim=-1)]
return merged_embeds, merged_pooled
class StableDiffusion3Pipeline(ComposedPipelineBase):
"""StableDiffusion3 pipeline implementation."""
pipeline_name = "StableDiffusion3Pipeline"
_required_config_modules = [
"text_encoder",
"text_encoder_2",
"text_encoder_3",
"tokenizer",
"tokenizer_2",
"tokenizer_3",
"vae",
"transformer",
"scheduler",
]
def create_pipeline_stages(self, server_args: ServerArgs):
self.add_stage(InputValidationStage())
self.add_stage(
TextEncodingStage(
text_encoders=[
self.get_module("text_encoder"),
self.get_module("text_encoder_2"),
self.get_module("text_encoder_3"),
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer_2"),
self.get_module("tokenizer_3"),
],
),
"prompt_encoding_stage_primary",
)
self.add_stage(SD3ConditioningStage())
self.add_standard_timestep_preparation_stage()
self.add_standard_latent_preparation_stage()
self.add_standard_denoising_stage()
self.add_standard_decoding_stage()
EntryClass = StableDiffusion3Pipeline
@@ -12,8 +12,6 @@ import inspect
import torch import torch
from sglang.multimodal_gen.configs.models.encoders import BaseEncoderOutput 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.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
@@ -263,11 +261,11 @@ class TextEncodingStage(PipelineStage):
).to(target_device) ).to(target_device)
input_ids = text_inputs["input_ids"] input_ids = text_inputs["input_ids"]
is_flux_v1 = isinstance( attention_mask = (
server_args.pipeline_config, FluxPipelineConfig server_args.pipeline_config.get_text_encoder_attention_mask(
) and not isinstance(server_args.pipeline_config, Flux2PipelineConfig) text_inputs, i
)
attention_mask = None if is_flux_v1 else text_inputs["attention_mask"] )
encoder_forward_kwargs = { encoder_forward_kwargs = {
"input_ids": input_ids, "input_ids": input_ids,
"output_hidden_states": True, "output_hidden_states": True,
@@ -292,12 +290,13 @@ class TextEncodingStage(PipelineStage):
prompt_embeds = prompt_embeds.to(device=target_device) prompt_embeds = prompt_embeds.to(device=target_device)
embeds_list.append(prompt_embeds) embeds_list.append(prompt_embeds)
if is_flux_v1 and outputs.pooler_output is not None:
# FLUX.1 only consumes the pooled CLIP projection. The T5 pooled_output = server_args.pipeline_config.get_text_encoder_pooler_output(
# encoder in the same pipeline has no pooler output. outputs, i
pooled_embeds_list.append( )
outputs.pooler_output.to(device=target_device) if pooled_output is not None:
) pooled_embeds_list.append(pooled_output.to(device=target_device))
if return_attention_mask: if return_attention_mask:
mask_to_store = ( mask_to_store = (
attention_mask.to(device=target_device) attention_mask.to(device=target_device)