[diffusion] chore: reuse SRT CLIP encoder blocks (#35004)

This commit is contained in:
Mick
2026-08-17 19:51:51 +08:00
committed by GitHub
parent 6e8a4abb57
commit d97b796c16
21 changed files with 779 additions and 646 deletions
@@ -2,7 +2,7 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from typing import Any
from typing import Any, Literal
import torch
@@ -79,8 +79,8 @@ class EncoderConfig(ModelConfig):
# Parallel folding: during the encoding stage the whole DiT replica is idle,
# so TP-shard the encoder across those otherwise-unused GPUs instead of
# running it on a single rank. None = replicated, else the group to fold
# over ("sp"|"ulysses"|"ring"|"world"); resolved by finalize_encoder_folding.
parallel_folding_mode: str | None = None
# over ("sp"|"world"); resolved by finalize_encoder_folding.
parallel_folding_mode: Literal["sp", "world"] | None = None
@dataclass
@@ -595,42 +595,24 @@ def model_parallel_is_initialized() -> bool:
)
_TP_STATE_PATCHED = False
@contextmanager
def patch_tensor_parallel_group(tp_group: GroupCoordinator):
"""Patch the tp group temporarily until this function ends.
This method is for draft workers of speculative decoding to run draft model
with different tp degree from that of target model workers.
"""
global _TP_STATE_PATCHED
assert not _TP_STATE_PATCHED, "Should not call when it's already patched"
_TP_STATE_PATCHED = True
def use_tensor_parallel_group(tp_group: GroupCoordinator):
"""Use one TP group consistently across diffusion and reused SRT modules."""
old_tp_group = get_tp_group()
import sglang.srt.distributed.parallel_state as srt_parallel_state
patch_srt_tp = srt_parallel_state._TP is old_tp_group
patch_srt_attention_tp = srt_parallel_state._ATTN_TP is old_tp_group
old_srt_tp_group = srt_parallel_state._TP
old_srt_attention_tp_group = srt_parallel_state._ATTN_TP
global _TP
_TP = tp_group
if patch_srt_tp:
srt_parallel_state._TP = tp_group
if patch_srt_attention_tp:
srt_parallel_state._ATTN_TP = tp_group
srt_parallel_state._TP = tp_group
srt_parallel_state._ATTN_TP = tp_group
try:
yield
finally:
# restore the original state
_TP_STATE_PATCHED = False
_TP = old_tp_group
if patch_srt_tp and srt_parallel_state._TP is tp_group:
srt_parallel_state._TP = old_tp_group
if patch_srt_attention_tp and srt_parallel_state._ATTN_TP is tp_group:
srt_parallel_state._ATTN_TP = old_tp_group
srt_parallel_state._TP = old_srt_tp_group
srt_parallel_state._ATTN_TP = old_srt_attention_tp_group
def get_tp_world_size() -> int:
@@ -3,7 +3,6 @@ import glob
import os
import re
from collections.abc import Callable, Generator, Iterable
from contextlib import nullcontext
from typing import cast
import torch
@@ -16,11 +15,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
)
from sglang.multimodal_gen.runtime.distributed import (
get_local_torch_device,
get_tp_group,
)
from sglang.multimodal_gen.runtime.distributed.group_coordinator import GroupCoordinator
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
patch_tensor_parallel_group,
use_tensor_parallel_group,
)
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
ComponentLoader,
@@ -36,6 +33,7 @@ from sglang.multimodal_gen.runtime.loader.weight_utils import (
safetensors_weights_iterator,
)
from sglang.multimodal_gen.runtime.models.encoders.base import (
EncoderTensorParallelMixin,
TextEncoder,
finalize_encoder_folding,
get_folding_tp_group,
@@ -408,20 +406,10 @@ class TextEncoderLoader(ComponentLoader):
else:
model_device = local_torch_device
# Parallel folding: build + shard the encoder over the folding group (the
# idle DiT replica during the encoding stage) instead of the default TP
# group, so every encoder folds without threading the group through each layer.
fold_ctx = nullcontext()
if getattr(model_config, "parallel_folding_mode", None) is not None:
folding_group = get_folding_tp_group(model_config)
if (
isinstance(folding_group, GroupCoordinator)
and folding_group is not get_tp_group()
):
fold_ctx = patch_tensor_parallel_group(folding_group)
# patch tp group with folding group to achieve TP among folding group
with fold_ctx, set_default_torch_dtype(PRECISION_TO_TYPE[dtype]):
encoder_tp_group = get_folding_tp_group(model_config)
with use_tensor_parallel_group(encoder_tp_group), set_default_torch_dtype(
PRECISION_TO_TYPE[dtype]
):
with model_device, skip_init_modules():
architectures = getattr(model_config, "architectures", [])
model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
@@ -435,6 +423,13 @@ class TextEncoderLoader(ComponentLoader):
model_config.enable_image_understanding = enable_image_understanding
model = model_cls(model_config)
if not isinstance(model, EncoderTensorParallelMixin):
raise TypeError(
f"Native encoder {model_cls.__name__} must inherit "
"EncoderTensorParallelMixin"
)
model.bind_encoder_tp_group(encoder_tp_group)
weights_to_load = {name for name, _ in model.named_parameters()}
loaded_weights = model.load_weights(
self._get_all_weights(
@@ -18,6 +18,10 @@ from sglang.multimodal_gen.runtime.distributed import (
get_tp_group,
get_world_group,
)
from sglang.multimodal_gen.runtime.distributed.group_coordinator import GroupCoordinator
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
use_tensor_parallel_group,
)
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
LayerwiseOffloadableModuleMixin,
)
@@ -25,19 +29,16 @@ from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
def get_folding_tp_group(config: EncoderConfig):
"""group an encoder tensor-parallels over; the default TP group unless a
fold mode is set"""
"""Return the TP group selected for an encoder."""
mode = config.parallel_folding_mode
if mode == "sp":
return get_sp_group()
elif mode == "ulysses":
return get_sp_group().ulysses_group
elif mode == "ring":
return get_sp_group().ring_group
elif mode == "world":
if mode == "world":
# the whole single-replica DiT (all GPUs), regardless of tp/sp/cfg.
return get_world_group()
return get_tp_group()
if mode is None:
return get_tp_group()
raise ValueError(f"Unsupported encoder folding mode: {mode!r}")
# measured on 2/4xH100: folding wins only for wide encoders (T5-XXL 4096: -20%
@@ -148,7 +149,25 @@ def finalize_encoder_folding(
config.parallel_folding_mode = None
class TextEncoder(nn.Module, ABC, LayerwiseOffloadableModuleMixin):
class EncoderTensorParallelMixin:
"""Keep an encoder on the TP group that was used to build its shards."""
_encoder_tp_group: GroupCoordinator | None = None
def bind_encoder_tp_group(self, tp_group: GroupCoordinator) -> None:
self._encoder_tp_group = tp_group
def __call__(self, *args, **kwargs):
tp_group = self._encoder_tp_group
if tp_group is None:
return super().__call__(*args, **kwargs)
with use_tensor_parallel_group(tp_group):
return super().__call__(*args, **kwargs)
class TextEncoder(
EncoderTensorParallelMixin, nn.Module, ABC, LayerwiseOffloadableModuleMixin
):
# Opt in per encoder to data-parallel batched encoding: the gather rebuilds a
# BaseEncoderOutput, and subclasses are free to return their own output type
# instead (Qwen2_5_VLForConditionalGeneration returns
@@ -200,7 +219,9 @@ class TextEncoder(nn.Module, ABC, LayerwiseOffloadableModuleMixin):
return self._supported_attention_backends
class ImageEncoder(nn.Module, ABC, LayerwiseOffloadableModuleMixin):
class ImageEncoder(
EncoderTensorParallelMixin, nn.Module, ABC, LayerwiseOffloadableModuleMixin
):
layerwise_offload_dit_group_enabled = False
layer_names = [
"layers",
@@ -17,406 +17,22 @@ from sglang.multimodal_gen.configs.models.encoders import (
CLIPTextConfig,
CLIPVisionConfig,
)
from sglang.multimodal_gen.runtime.distributed import divide, get_tp_world_size
from sglang.multimodal_gen.runtime.layers.activation import get_act_fn
from sglang.multimodal_gen.runtime.layers.attention import LocalAttention
from sglang.multimodal_gen.runtime.layers.linear import (
ColumnParallelLinear,
QKVParallelLinear,
RowParallelLinear,
)
from sglang.multimodal_gen.runtime.layers.quantization import QuantizationConfig
# TODO: support quantization
# from vllm.model_executor.layers.quantization import QuantizationConfig
from sglang.multimodal_gen.runtime.loader.weight_utils import default_weight_loader
from sglang.multimodal_gen.runtime.models.encoders.base import ImageEncoder, TextEncoder
from sglang.multimodal_gen.runtime.models.encoders.vision import (
resolve_visual_encoder_outputs,
)
from sglang.multimodal_gen.runtime.platforms import (
AttentionBackendEnum,
current_platform,
from sglang.srt.models.clip import (
CLIPEncoder,
CLIPTextEmbeddings,
CLIPVisionEmbeddings,
prepare_clip_attention_mask,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
# Adapted from https://github.com/huggingface/transformers/blob/v4.39.0/src/transformers/models/clip/modeling_clip.py#L164 # noqa
class CLIPVisionEmbeddings(nn.Module):
def __init__(self, config: CLIPVisionConfig):
super().__init__()
self.config = config
self.embed_dim = config.hidden_size
self.image_size = config.image_size
self.patch_size = config.patch_size
assert self.image_size % self.patch_size == 0
self.class_embedding = nn.Parameter(torch.randn(self.embed_dim))
self.patch_embedding = nn.Conv2d(
in_channels=config.num_channels,
out_channels=self.embed_dim,
kernel_size=self.patch_size,
stride=self.patch_size,
bias=False,
)
self.num_patches = (self.image_size // self.patch_size) ** 2
self.num_positions = self.num_patches + 1
self.position_embedding = nn.Embedding(self.num_positions, self.embed_dim)
self.register_buffer(
"position_ids",
torch.arange(self.num_positions).expand((1, -1)),
persistent=False,
)
def forward(self, pixel_values: torch.Tensor) -> torch.Tensor:
batch_size = pixel_values.shape[0]
target_dtype = self.patch_embedding.weight.dtype
patch_embeds = self.patch_embedding(
pixel_values.to(dtype=target_dtype)
) # shape = [*, width, grid, grid]
patch_embeds = patch_embeds.flatten(2).transpose(1, 2)
class_embeds = self.class_embedding.expand(batch_size, 1, -1)
embeddings = torch.cat([class_embeds, patch_embeds], dim=1)
embeddings = embeddings + self.position_embedding(self.position_ids)
return embeddings
class CLIPTextEmbeddings(nn.Module):
def __init__(self, config: CLIPTextConfig):
super().__init__()
self.config = config
embed_dim = config.hidden_size
self.token_embedding = nn.Embedding(config.vocab_size, embed_dim)
self.position_embedding = nn.Embedding(
config.max_position_embeddings, embed_dim
)
# position_ids (1, len position emb) is contiguous in memory and exported when serialized
self.register_buffer(
"position_ids",
torch.arange(config.max_position_embeddings).expand((1, -1)),
persistent=False,
)
def forward(
self,
input_ids: torch.LongTensor | None = None,
position_ids: torch.LongTensor | None = None,
inputs_embeds: torch.FloatTensor | None = None,
) -> torch.Tensor:
if input_ids is not None:
seq_length = input_ids.shape[-1]
elif inputs_embeds is not None:
seq_length = inputs_embeds.shape[-2]
else:
raise ValueError("Either input_ids or inputs_embeds must be provided.")
max_position_embedding = self.position_embedding.weight.shape[0]
if seq_length > max_position_embedding:
raise ValueError(
f"Sequence length must be less than max_position_embeddings (got `sequence length`: "
f"{seq_length} and max_position_embeddings: {max_position_embedding}"
)
if position_ids is None:
position_ids = self.position_ids[:, :seq_length]
if inputs_embeds is None:
inputs_embeds = self.token_embedding(input_ids)
position_embeddings = self.position_embedding(position_ids)
embeddings = inputs_embeds + position_embeddings
return embeddings
class CLIPAttention(nn.Module):
"""Multi-headed attention from 'Attention Is All You Need' paper"""
def __init__(
self,
config: CLIPVisionConfig | CLIPTextConfig,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
):
super().__init__()
self.config = config
self.embed_dim = config.hidden_size
self.num_heads = config.num_attention_heads
self.head_dim = self.embed_dim // self.num_heads
if self.head_dim * self.num_heads != self.embed_dim:
raise ValueError(
"embed_dim must be divisible by num_heads "
f"(got `embed_dim`: {self.embed_dim} and `num_heads`:"
f" {self.num_heads})."
)
self.scale = self.head_dim**-0.5
self.dropout = config.attention_dropout
self.qkv_proj = QKVParallelLinear(
hidden_size=self.embed_dim,
head_size=self.head_dim,
total_num_heads=self.num_heads,
quant_config=quant_config,
prefix=f"{prefix}.qkv_proj",
)
self.out_proj = RowParallelLinear(
input_size=self.embed_dim,
output_size=self.embed_dim,
quant_config=quant_config,
prefix=f"{prefix}.out_proj",
)
self.tp_size = get_tp_world_size()
self.num_heads_per_partition = divide(self.num_heads, self.tp_size)
self.attn = LocalAttention(
self.num_heads_per_partition,
self.head_dim,
self.num_heads_per_partition,
softmax_scale=self.scale,
causal=True,
supported_attention_backends=config._supported_attention_backends,
)
def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):
return (
tensor.view(bsz, seq_len, self.num_heads, self.head_dim)
.transpose(1, 2)
.contiguous()
)
def forward(
self,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor | None = None,
):
"""Input shape: Batch x Time x Channel"""
qkv_states, _ = self.qkv_proj(hidden_states)
query_states, key_states, value_states = qkv_states.chunk(3, dim=-1)
# use flash_attn_func
query_states = query_states.reshape(
query_states.shape[0],
query_states.shape[1],
self.num_heads_per_partition,
self.head_dim,
)
key_states = key_states.reshape(
key_states.shape[0],
key_states.shape[1],
self.num_heads_per_partition,
self.head_dim,
)
value_states = value_states.reshape(
value_states.shape[0],
value_states.shape[1],
self.num_heads_per_partition,
self.head_dim,
)
if self.attn.backend == AttentionBackendEnum.TORCH_SDPA:
query_states = query_states.transpose(1, 2) # [B, H, S, D]
key_states = key_states.transpose(1, 2)
value_states = value_states.transpose(1, 2)
if (
current_platform.is_rocm()
or current_platform.is_musa()
or current_platform.is_xpu()
):
# ROCm: Using both is_causal=True and attn_mask causes NaN.
# Use is_causal=True alone (padding mask not needed for CLIP
# since pooler_output comes from EOS token before padding).
# XXX (MUSA): Torch SDPA on MUSA currently does not support
# using both `attn_mask` and `is_causal=True` simultaneously.
attn_output = torch.nn.functional.scaled_dot_product_attention(
query_states,
key_states,
value_states,
attn_mask=None,
is_causal=True,
scale=self.scale,
)
else:
if attention_mask is not None:
# SDPA requires [B, 1, 1, S] or [B, S, S] format mask
if attention_mask.dim() == 2:
attn_mask = attention_mask[:, None, None, :].to(
dtype=query_states.dtype
)
attn_mask = (1.0 - attn_mask) * torch.finfo(
query_states.dtype
).min
else:
attn_mask = attention_mask
else:
attn_mask = None
attn_output = torch.nn.functional.scaled_dot_product_attention(
query_states,
key_states,
value_states,
attn_mask=attn_mask,
is_causal=attention_mask is None,
scale=self.scale,
)
attn_output = attn_output.transpose(1, 2)
else:
# Use LocalAttention (doesn't support attention_mask, but maintains compatibility)
attn_output = self.attn(query_states, key_states, value_states)
attn_output = attn_output.reshape(
attn_output.shape[0],
attn_output.shape[1],
self.num_heads_per_partition * self.head_dim,
)
attn_output, _ = self.out_proj(attn_output)
return attn_output, None
class CLIPMLP(nn.Module):
def __init__(
self,
config: CLIPVisionConfig | CLIPTextConfig,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
) -> None:
super().__init__()
self.config = config
self.activation_fn = get_act_fn(config.hidden_act)
self.fc1 = ColumnParallelLinear(
config.hidden_size,
config.intermediate_size,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.fc1",
)
self.fc2 = RowParallelLinear(
config.intermediate_size,
config.hidden_size,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.fc2",
)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states, _ = self.fc1(hidden_states)
hidden_states = self.activation_fn(hidden_states)
hidden_states, _ = self.fc2(hidden_states)
return hidden_states
class CLIPEncoderLayer(nn.Module):
def __init__(
self,
config: CLIPTextConfig | CLIPVisionConfig,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
) -> None:
super().__init__()
self.self_attn = CLIPAttention(
config,
quant_config=quant_config,
prefix=f"{prefix}.self_attn",
)
self.layer_norm1 = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
self.mlp = CLIPMLP(config, quant_config=quant_config, prefix=f"{prefix}.mlp")
self.layer_norm2 = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
def forward(
self,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor | None = None,
) -> torch.Tensor:
residual = hidden_states
hidden_states = self.layer_norm1(hidden_states)
hidden_states, _ = self.self_attn(
hidden_states=hidden_states,
attention_mask=attention_mask,
)
hidden_states = residual + hidden_states
residual = hidden_states
hidden_states = self.layer_norm2(hidden_states)
hidden_states = self.mlp(hidden_states)
hidden_states = residual + hidden_states
return hidden_states
class CLIPEncoder(nn.Module):
"""
Transformer encoder consisting of `config.num_hidden_layers` self
attention layers. Each layer is a [`CLIPEncoderLayer`].
Args:
config: CLIPConfig
"""
def __init__(
self,
config: CLIPVisionConfig | CLIPTextConfig,
quant_config: QuantizationConfig | None = None,
num_hidden_layers_override: int | None = None,
prefix: str = "",
) -> None:
super().__init__()
self.config = config
if num_hidden_layers_override is None:
num_hidden_layers = config.num_hidden_layers
else:
num_hidden_layers = num_hidden_layers_override
self.layers = nn.ModuleList(
[
CLIPEncoderLayer(
config=config,
quant_config=quant_config,
prefix=f"{prefix}.layers.{layer_idx}",
)
for layer_idx in range(num_hidden_layers)
]
)
def forward(
self,
inputs_embeds: torch.Tensor,
return_all_hidden_states: bool,
attention_mask: torch.Tensor | None = None,
) -> torch.Tensor | list[torch.Tensor]:
hidden_states_pool = [inputs_embeds]
hidden_states = inputs_embeds
for idx, encoder_layer in enumerate(self.layers):
hidden_states = encoder_layer(
hidden_states,
attention_mask=attention_mask,
)
if return_all_hidden_states:
hidden_states_pool.append(hidden_states)
# If we have multiple feature sample layers, we return all hidden
# states in order and grab the ones we need by index.
if return_all_hidden_states:
return hidden_states_pool
return [hidden_states]
def _srt_clip_param_name(name: str) -> str:
return name.replace(".out_proj.", ".proj.")
class CLIPTextTransformer(nn.Module):
@@ -439,6 +55,7 @@ class CLIPTextTransformer(nn.Module):
quant_config=quant_config,
num_hidden_layers_override=num_hidden_layers_override,
prefix=prefix,
causal=True,
)
self.final_layer_norm = nn.LayerNorm(embed_dim, eps=config.layer_norm_eps)
@@ -468,17 +85,12 @@ class CLIPTextTransformer(nn.Module):
hidden_states = self.embeddings(input_ids=input_ids, position_ids=position_ids)
# CLIP's text model uses causal mask, prepare it here.
# https://github.com/openai/CLIP/blob/cfcffb90e69f37bf2ff1e988237a0fbe41f33c04/clip/model.py#L324
# causal_attention_mask = _create_4d_causal_attention_mask(
# input_shape, hidden_states.dtype, device=hidden_states.device
# )
# # expand attention_mask
# if attention_mask is not None and not self._use_flash_attention_2:
# raise NotImplementedError("attention_mask is not supported for CLIPTextTransformer")
# # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
# attention_mask = _prepare_4d_attention_mask(attention_mask, hidden_states.dtype)
attention_mask = prepare_clip_attention_mask(
input_shape,
hidden_states.dtype,
hidden_states.device,
attention_mask,
)
encoder_outputs = self.encoder(
inputs_embeds=hidden_states,
@@ -486,7 +98,12 @@ class CLIPTextTransformer(nn.Module):
attention_mask=attention_mask,
)
last_hidden_state = encoder_outputs[-1]
if output_hidden_states:
all_hidden_states = encoder_outputs
last_hidden_state = encoder_outputs[-1]
else:
last_hidden_state = encoder_outputs
all_hidden_states = [encoder_outputs]
last_hidden_state = self.final_layer_norm(last_hidden_state)
if self.eos_token_id == 2:
@@ -523,8 +140,7 @@ class CLIPTextTransformer(nn.Module):
return BaseEncoderOutput(
last_hidden_state=last_hidden_state,
pooler_output=pooled_output,
hidden_states=encoder_outputs,
# attentions=encoder_outputs.attentions,
hidden_states=all_hidden_states,
)
@@ -569,6 +185,7 @@ class CLIPTextModel(TextEncoder):
params_dict = dict(self.named_parameters())
loaded_params: set[str] = set()
for name, loaded_weight in weights:
name = _srt_clip_param_name(name)
# Handle q_proj, k_proj, v_proj -> qkv_proj mapping
for param_name, weight_name, shard_id in stacked_params_mapping:
if weight_name in name:
@@ -702,8 +319,6 @@ class CLIPVisionTransformer(nn.Module):
)
if not return_all_hidden_states:
encoder_outputs = encoder_outputs[0]
# Handle post-norm (if applicable) and stacks feature layers if needed
encoder_outputs = resolve_visual_encoder_outputs(
encoder_outputs,
@@ -763,6 +378,7 @@ class CLIPVisionModel(ImageEncoder):
for name, loaded_weight in weights:
if name.startswith("visual_projection"):
continue
name = _srt_clip_param_name(name)
# post_layernorm is not needed in CLIPVisionModel
if (
name.startswith("vision_model.post_layernorm")
@@ -38,6 +38,9 @@ from sglang.multimodal_gen.runtime.loader.weight_utils import default_weight_loa
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
LayerwiseOffloadableModuleMixin,
)
from sglang.multimodal_gen.runtime.models.encoders.base import (
EncoderTensorParallelMixin,
)
logger = logging.getLogger(__name__)
@@ -283,7 +286,9 @@ class Gemma2DecoderLayer(nn.Module):
return hidden_states
class Gemma2Model(nn.Module, LayerwiseOffloadableModuleMixin):
class Gemma2Model(
EncoderTensorParallelMixin, nn.Module, LayerwiseOffloadableModuleMixin
):
"""Gemma2 text encoder model for SANA pipeline."""
_fsdp_shard_conditions = []
@@ -4,7 +4,6 @@
# Adapted from sglang: python/sglang/srt/models/gemma3_causal.py
import logging
from contextlib import nullcontext
from typing import Any, Iterable, Optional, Set, Tuple
import torch
@@ -12,10 +11,7 @@ from torch import nn
from sglang.multimodal_gen.configs.models.encoders.base import BaseEncoderOutput
from sglang.multimodal_gen.configs.models.encoders.gemma_3 import Gemma3Config
from sglang.multimodal_gen.runtime.distributed import get_tp_group, get_tp_world_size
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
patch_tensor_parallel_group,
)
from sglang.multimodal_gen.runtime.distributed import get_tp_world_size
from sglang.multimodal_gen.runtime.layers.activation import GeluAndMul
from sglang.multimodal_gen.runtime.layers.linear import (
MergedColumnParallelLinear,
@@ -28,6 +24,9 @@ from sglang.multimodal_gen.runtime.loader.weight_utils import default_weight_loa
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
LayerwiseOffloadableModuleMixin,
)
from sglang.multimodal_gen.runtime.models.encoders.base import (
EncoderTensorParallelMixin,
)
from sglang.multimodal_gen.runtime.utils.common import add_prefix
from sglang.srt.models.siglip import SiglipVisionModel
@@ -678,7 +677,9 @@ class Gemma3TextModel(nn.Module):
return loaded_params
class Gemma3ForConditionalGeneration(nn.Module, LayerwiseOffloadableModuleMixin):
class Gemma3ForConditionalGeneration(
EncoderTensorParallelMixin, nn.Module, LayerwiseOffloadableModuleMixin
):
# transformers 5.6.0 flattened SiglipVisionModel, dropping the
# `vision_model` intermediate wrapper. Our reimpl keeps it, so remap
# HF source keys back into our nested namespace when transferring weights.
@@ -704,7 +705,6 @@ class Gemma3ForConditionalGeneration(nn.Module, LayerwiseOffloadableModuleMixin)
self.config = config
self.quant_config = quant_config
self.text_config = config.text_config
self._vision_tensor_parallel_group = get_tp_group()
# Vision Tower
self.vision_tower = SiglipVisionModel(
@@ -720,11 +720,6 @@ class Gemma3ForConditionalGeneration(nn.Module, LayerwiseOffloadableModuleMixin)
# Text Model
self.language_model = Gemma3TextModel(config)
def _vision_parallel_context(self):
if get_tp_group() is self._vision_tensor_parallel_group:
return nullcontext()
return patch_tensor_parallel_group(self._vision_tensor_parallel_group)
def get_placeholder_mask(
self,
input_ids: torch.LongTensor,
@@ -777,8 +772,7 @@ class Gemma3ForConditionalGeneration(nn.Module, LayerwiseOffloadableModuleMixin)
elif pixel_values.dim() != 4:
raise ValueError(f"Unexpected pixel_values shape: {pixel_values.shape}")
with self._vision_parallel_context():
vision_outputs = self.vision_tower(pixel_values)
vision_outputs = self.vision_tower(pixel_values)
image_features = self.multi_modal_projector(vision_outputs)
image_features = image_features.to(
device=inputs_embeds.device, dtype=inputs_embeds.dtype
@@ -110,7 +110,7 @@ class IdeogramQwen3VLTextEncoder(TextEncoder):
position_ids = pos_2d[None, ...].expand(4, 1, -1)
attention_mask = torch.ones_like(cur_token_ids)
with set_forward_context(current_timestep=0, attn_metadata=None):
outputs = self.forward(
outputs = self(
input_ids=cur_token_ids,
position_ids=position_ids,
attention_mask=attention_mask,
@@ -41,9 +41,6 @@ class MiniMaxH3Qwen3VLEncoder(TextEncoder):
eight otherwise-idle ranks during encoding.
"""
# encode_ids drives the forward pass; __call__ is never used, so FSDP2
# needs it registered or the root group (the vision tower) stays sharded.
_fsdp_forward_methods = ("encode_ids",)
layer_names = [*TextEncoder.layer_names, "model.visual.blocks"]
supports_dp_encode = True
@@ -150,10 +147,6 @@ class MiniMaxH3Qwen3VLEncoder(TextEncoder):
call_kwargs: dict[str, Any] = {
"input_ids": ids,
"attention_mask": torch.ones_like(ids),
"output_attentions": False,
"output_hidden_states": False,
"return_dict": True,
"use_cache": False,
}
if position_ids is not None:
call_kwargs["position_ids"] = position_ids.to(self.device)
@@ -166,7 +159,7 @@ class MiniMaxH3Qwen3VLEncoder(TextEncoder):
)
call_kwargs["video_grid_thw"] = host_video_grid_thw
hidden = self.model(**call_kwargs).last_hidden_state[0].to(torch.bfloat16)
hidden = self(**call_kwargs).last_hidden_state[0].to(torch.bfloat16)
expected_shape = [int(ids.shape[1]), self.hidden_dim]
if list(hidden.shape) != expected_shape:
raise ValueError(
@@ -65,6 +65,9 @@ from sglang.multimodal_gen.runtime.loader.weight_utils import default_weight_loa
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
LayerwiseOffloadableModuleMixin,
)
from sglang.multimodal_gen.runtime.models.encoders.base import (
EncoderTensorParallelMixin,
)
from sglang.multimodal_gen.runtime.platforms import (
AttentionBackendEnum,
current_platform,
@@ -635,7 +638,9 @@ class Mistral3Model(nn.Module):
)
class Mistral3ForConditionalGeneration(nn.Module, LayerwiseOffloadableModuleMixin):
class Mistral3ForConditionalGeneration(
EncoderTensorParallelMixin, nn.Module, LayerwiseOffloadableModuleMixin
):
_checkpoint_conversion_mapping = {
"^language_model.model": "model.language_model",
"^multi_modal_projector": "model.multi_modal_projector",
@@ -232,22 +232,22 @@ def _resolve_warmup_num_frames(
server_based_warmup: bool,
) -> int:
num_frames = getattr(sampling_defaults, "num_frames", 1)
if (
not server_based_warmup
or not _is_video_warmup_task(server_args)
or num_frames is None
):
# use default num frames
if not _is_video_warmup_task(server_args) or num_frames is None:
return num_frames
# Breakable CUDA graph replays only exact latent shapes: the warmup
# request must run the full serving frame count so its captured graphs
# match serving signatures (mirrors the uncapped-steps rule in
# _resolve_warmup_steps).
if getattr(server_args, "enable_breakable_cuda_graph", False) is True:
return num_frames
if (
not server_based_warmup
or getattr(server_args, "enable_breakable_cuda_graph", False) is True
):
warmup_num_frames = num_frames
else:
warmup_num_frames = min(num_frames, SERVER_WARMUP_MAX_VIDEO_FRAMES)
return min(num_frames, SERVER_WARMUP_MAX_VIDEO_FRAMES)
return server_args.pipeline_config.adjust_num_frames(warmup_num_frames)
def _effective_cfg_scale(sampling_defaults: SamplingParams) -> float | None:
@@ -1177,6 +1177,7 @@ STANDALONE_FILES = {
"../single_test_file/test_disagg_server.py",
"../single_test_file/test_ar_models.py",
"../single_test_file/test_ipc_a2a_2_gpu.py",
"../single_test_file/test_encoder_fold_srt_2_gpu.py",
"../single_test_file/test_diffusion_bcg_tp2_zimage_turbo.py",
"../single_test_file/test_dp_serving_2_gpu.py",
"../single_test_file/test_pynccl_a2a_capture_2_gpu.py",
@@ -1215,6 +1216,7 @@ STANDALONE_FILE_EST_TIMES = {
"../single_test_file/test_ar_models.py": 600.0,
# no model load; the cost is the one-time JIT build of the sync kernels
"../single_test_file/test_ipc_a2a_2_gpu.py": 240.0,
"../single_test_file/test_encoder_fold_srt_2_gpu.py": 240.0,
# ~60 s locally with a warm HF cache (load + one capture + 4 steps);
# padded for cold-cache CI.
"../single_test_file/test_diffusion_bcg_tp2_zimage_turbo.py": 180.0,
@@ -16,6 +16,7 @@ import time
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Callable, Sequence
from urllib.error import HTTPError, URLError
from urllib.request import urlopen
import pytest
@@ -404,6 +405,9 @@ class ServerManager:
]
if self.extra_args.strip():
command.extend(self.extra_args.strip().split())
access_log_exclude_flag = "--uvicorn-access-log-exclude-prefixes"
if not any(arg.startswith(access_log_exclude_flag) for arg in command):
command.extend(["--uvicorn-access-log-exclude-prefixes", "/health"])
env = os.environ.copy()
env["SGLANG_DIFFUSION_STAGE_LOGGING"] = "1"
@@ -471,9 +475,9 @@ class ServerManager:
)
def _wait_for_ready(self, process: subprocess.Popen, stdout_path: Path) -> None:
"""Wait for server to become ready."""
"""Wait until model warmup finishes and inference traffic is accepted."""
start = time.time()
ready_message = "Application startup complete."
health_url = f"http://127.0.0.1:{self.port}/health"
log_period = 30
prev_log_period_count = 0
@@ -484,14 +488,13 @@ class ServerManager:
f"Server exited early (code {process.returncode}).\n{tail}"
)
if stdout_path.exists():
try:
content = stdout_path.read_text(encoding="utf-8", errors="ignore")
if ready_message in content:
try:
with urlopen(health_url, timeout=1) as response:
if response.status == 200:
logger.info("[server-test] Server ready")
return
except Exception as e:
logger.debug("Could not read log yet: %s", e)
except (HTTPError, URLError, TimeoutError, OSError):
pass
elapsed = int(time.time() - start)
if (elapsed // log_period) > prev_log_period_count:
@@ -0,0 +1,275 @@
"""Two-rank encoder folding must preserve single-rank native output.
The focused CLIP check isolates the component loader and SRT tensor-parallel
layers. The tiny SD3 check covers the public server API and complete pipeline.
"""
from __future__ import annotations
import os
import subprocess
import sys
import unittest
from types import SimpleNamespace
import torch
from sglang.multimodal_gen.runtime.platforms import current_platform
from sglang.test.test_utils import CustomTestCase
_WORLD = 2
_TINY_SD3_MODEL = "yujiepan/stable-diffusion-3-tiny-random"
_TINY_SD3_REVISION = "abcdbb999b2d30c35d03efdce0be981e1efac0a4"
def _tiny_clip_config():
from sglang.multimodal_gen.configs.models.encoders.clip import (
CLIPTextArchConfig,
CLIPTextConfig,
)
return CLIPTextConfig(
arch_config=CLIPTextArchConfig(
architectures=["CLIPTextModel"],
vocab_size=32,
hidden_size=8,
intermediate_size=16,
projection_dim=8,
num_hidden_layers=1,
num_attention_heads=2,
max_position_embeddings=8,
pad_token_id=0,
bos_token_id=1,
eos_token_id=2,
text_len=8,
),
prefix="clip",
)
def _deterministic_state_dict(model: torch.nn.Module) -> dict[str, torch.Tensor]:
generator = torch.Generator(device="cpu").manual_seed(20260816)
state_dict = {}
for name, value in model.state_dict().items():
state_dict[name] = torch.randn(
value.shape,
dtype=value.dtype,
generator=generator,
).mul_(0.02)
return state_dict
def _clip_checkpoint_weights(
state_dict: dict[str, torch.Tensor],
) -> list[tuple[str, torch.Tensor]]:
weights = []
for name, value in state_dict.items():
if ".qkv_proj." not in name:
weights.append((name, value))
continue
for projection, shard in zip(("q", "k", "v"), value.chunk(3, dim=0)):
weights.append((name.replace("qkv_proj", f"{projection}_proj"), shard))
return weights
def _worker() -> int:
from sglang.multimodal_gen.runtime.distributed import (
cleanup_dist_env_and_memory,
get_tp_group,
get_world_group,
init_distributed_environment,
initialize_model_parallel,
)
from sglang.multimodal_gen.runtime.loader.component_loaders.text_encoder_loader import (
TextEncoderLoader,
)
from sglang.multimodal_gen.runtime.models.encoders.clip import CLIPTextModel
from sglang.srt.distributed import parallel_state as srt_parallel_state
rank = int(os.environ["RANK"])
world_size = int(os.environ["WORLD_SIZE"])
device = torch.device(f"cuda:{rank}")
torch.cuda.set_device(device)
init_distributed_environment(
world_size=world_size,
rank=rank,
local_rank=rank,
)
initialize_model_parallel(
tensor_parallel_degree=1,
sequence_parallel_degree=world_size,
ulysses_degree=world_size,
ring_degree=1,
)
config = _tiny_clip_config()
reference = CLIPTextModel(config).to(device).eval()
state_dict = _deterministic_state_dict(reference)
reference.load_state_dict(
{name: value.to(device) for name, value in state_dict.items()}
)
class InMemoryTextEncoderLoader(TextEncoderLoader):
def _get_all_weights(self, model, model_path, to_cpu):
del model, model_path, to_cpu
yield from _clip_checkpoint_weights(state_dict)
config.parallel_folding_mode = "world"
server_args = SimpleNamespace(
pipeline_config=SimpleNamespace(),
should_start_component_on_cpu=lambda component_name: False,
)
folded = InMemoryTextEncoderLoader().load_model(
"unused",
config,
server_args,
dtype="fp32",
component_starts_on_cpu=False,
)
folded.eval()
fold_group = get_world_group()
assert folded._encoder_tp_group is fold_group
assert folded.text_model.encoder.layers[0].mlp.fc2.tp_size == world_size
assert get_tp_group().world_size == 1
input_ids = torch.tensor([[1, 7, 11, 2]], device=device)
with torch.no_grad():
expected = reference(input_ids=input_ids).last_hidden_state
actual = folded(input_ids=input_ids).last_hidden_state
torch.testing.assert_close(actual, expected, rtol=2e-5, atol=2e-5)
assert get_tp_group().world_size == 1
assert srt_parallel_state.get_tp_group().world_size == 1
assert srt_parallel_state.get_attn_tp_group().world_size == 1
if rank == 0:
print("ENCODER_FOLD_SRT_PARITY PASS", flush=True)
torch.distributed.barrier()
cleanup_dist_env_and_memory()
return 0
def _generate_tiny_sd3(*, fold: bool):
from openai import OpenAI
from sglang.multimodal_gen.test.server.test_server_utils import (
ServerManager,
get_generate_fn,
)
from sglang.multimodal_gen.test.server.testcase_configs import (
DiffusionSamplingParams,
)
from sglang.multimodal_gen.test.test_utils import (
find_free_port,
image_bytes_to_numpy,
)
sampling_params = DiffusionSamplingParams(
prompt="a red cube",
output_size="64x64",
extras={"num_inference_steps": 2, "seed": 0, "guidance_scale": 1.0},
)
encoder_mode = "fold" if fold else "replicate"
parallel_args = f"--num-gpus 2 --ulysses-degree 2 --encoder-parallel {encoder_mode}"
extra_args = " ".join(
[
"--model-type diffusion",
"--backend sglang",
"--model-id stable-diffusion-3-medium",
f"--served-model-name {_TINY_SD3_MODEL}",
f"--revision {_TINY_SD3_REVISION}",
"--strict-ports",
parallel_args,
]
)
manager = ServerManager(
model=_TINY_SD3_MODEL,
port=find_free_port(),
wait_deadline=600,
extra_args=extra_args,
)
ctx = manager.start()
try:
client = OpenAI(
api_key="sglang-anything",
base_url=f"http://localhost:{ctx.port}/v1",
timeout=600,
max_retries=0,
)
model_ids = [model.id for model in client.models.list().data]
assert _TINY_SD3_MODEL in model_ids
generate = get_generate_fn(
model_path=_TINY_SD3_MODEL,
modality="image",
sampling_params=sampling_params,
)
_, content = generate("tiny_sd3_encoder_fold_e2e", client)
log = ctx.log_tail(lines=500)
assert "Using native sglang backend" in log
assert "[TextEncodingStage]" in log
return image_bytes_to_numpy(content)
finally:
ctx.cleanup()
class TestEncoderFoldSrtTwoGpu(CustomTestCase):
def test_folded_pipeline_matches_replicated_encoder(self):
if not current_platform.is_cuda():
self.skipTest("CUDA-only test")
if torch.cuda.device_count() < _WORLD:
self.skipTest(f"needs {_WORLD} GPUs")
from sglang.multimodal_gen.test.test_utils import (
compute_mean_abs_diff,
compute_psnr,
compute_ssim,
)
reference = _generate_tiny_sd3(fold=False)
folded = _generate_tiny_sd3(fold=True)
ssim = compute_ssim(folded, reference)
psnr = compute_psnr(folded, reference)
mean_abs_diff = compute_mean_abs_diff(folded, reference)
print(
"ENCODER_FOLD_E2E_PARITY "
f"ssim={ssim:.6f} psnr={psnr:.6f} mad={mean_abs_diff:.6f}",
flush=True,
)
# BF16 TP reductions may move the final uint8 output slightly. A wrong
# runtime group produces multi-level pixel drift, not this rounding noise.
self.assertGreaterEqual(ssim, 0.98)
self.assertLessEqual(mean_abs_diff, 2.0)
def test_folded_srt_clip_matches_single_rank(self):
if not current_platform.is_cuda():
self.skipTest("CUDA-only test")
if torch.cuda.device_count() < _WORLD:
self.skipTest(f"needs {_WORLD} GPUs")
proc = subprocess.run(
[
sys.executable,
"-m",
"torch.distributed.run",
f"--nproc-per-node={_WORLD}",
"--master-port=29617",
__file__,
"--worker",
],
capture_output=True,
text=True,
timeout=600,
)
print(proc.stdout[-4000:])
if proc.returncode != 0:
print(proc.stderr[-4000:], file=sys.stderr)
self.assertEqual(proc.returncode, 0, "folded SRT CLIP diverged")
self.assertIn("ENCODER_FOLD_SRT_PARITY PASS", proc.stdout)
if __name__ == "__main__":
if "--worker" in sys.argv:
raise SystemExit(_worker())
unittest.main()
@@ -39,7 +39,7 @@ logger = init_logger(__name__)
# NPU/ascend) is read from sgl-project/ci-data-diffusion, where the GT-gen workflows
# publish.
SGL_TEST_FILES_CI_DATA_REPO = "sgl-project/ci-data-diffusion"
SGL_TEST_FILES_CI_DATA_REVISION = "cc3f27fd2d1b4d8e1a7d5eec1247a215a502b9c1"
SGL_TEST_FILES_CI_DATA_REVISION = "8c3896984319c8d5628bf08df4b596baf2368ec7"
# The NPU pin is kept as a separate branch so ascend GT can be bumped independently
# when it's regenerated on its own cadence.
@@ -24,6 +24,10 @@ from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType
from sglang.multimodal_gen.configs.pipeline_configs.flux_finetuned import (
Flux2FinetunedPipelineConfig,
)
from sglang.multimodal_gen.configs.pipeline_configs.longlive2 import (
LongLive2T2VConfig,
)
from sglang.multimodal_gen.configs.sample.longlive2 import LongLive2SamplingParams
from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
from sglang.multimodal_gen.runtime.entrypoints.diffusion_generator import DiffGenerator
from sglang.multimodal_gen.runtime.entrypoints.utils import (
@@ -49,6 +53,7 @@ from sglang.multimodal_gen.runtime.server_warmup import (
from sglang.multimodal_gen.runtime.warmup_request_builder import (
DEFAULT_PLACEHOLDER_PROMPT,
SERVER_WARMUP_IMAGE_FALLBACK_RESOLUTION,
_resolve_warmup_num_frames,
build_warmup_reqs,
should_include_warmup_image,
supports_synthetic_warmup,
@@ -71,11 +76,7 @@ def _make_bare_scheduler(enable_cfg_parallel: bool) -> Scheduler:
server_args.enable_torch_compile = False
server_args.is_arg_explicitly_set.return_value = False
task_type = MagicMock()
task_type.requires_image_input.return_value = False
task_type.accepts_image_input.return_value = False
task_type.data_type.return_value = ModelTaskType.T2I.data_type()
server_args.pipeline_config.task_type = task_type
server_args.pipeline_config.task_type = ModelTaskType.T2I
scheduler.server_args = server_args
scheduler.req_based_warmup_scheduled = False
@@ -270,12 +271,8 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args.enable_cfg_parallel = False
server_args.enable_torch_compile = False
task_type = MagicMock()
task_type.requires_image_input.return_value = False
task_type.accepts_image_input.return_value = False
task_type.is_image_gen.return_value = False
task_type.data_type.return_value = ModelTaskType.T2V.data_type()
server_args.pipeline_config.task_type = task_type
server_args.pipeline_config.task_type = ModelTaskType.T2V
server_args.pipeline_config.adjust_num_frames.side_effect = lambda value: value
generator.server_args = server_args
sampling_defaults = SamplingParams(num_frames=81, num_inference_steps=50)
@@ -307,12 +304,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args.enable_cfg_parallel = False
server_args.enable_torch_compile = False
task_type = MagicMock()
task_type.requires_image_input.return_value = False
task_type.accepts_image_input.return_value = False
task_type.is_image_gen.return_value = True
task_type.data_type.return_value = ModelTaskType.T2I.data_type()
server_args.pipeline_config.task_type = task_type
server_args.pipeline_config.task_type = ModelTaskType.T2I
sampling_defaults = SamplingParams(
negative_prompt="model default negative",
@@ -348,12 +340,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args.enable_cfg_parallel = False
server_args.enable_torch_compile = False
task_type = MagicMock()
task_type.requires_image_input.return_value = False
task_type.accepts_image_input.return_value = False
task_type.is_image_gen.return_value = True
task_type.data_type.return_value = ModelTaskType.T2I.data_type()
server_args.pipeline_config.task_type = task_type
server_args.pipeline_config.task_type = ModelTaskType.T2I
sampling_defaults = SamplingParams(width=640, height=640)
with patch(
@@ -376,12 +363,8 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args.enable_cfg_parallel = False
server_args.enable_torch_compile = False
task_type = MagicMock()
task_type.requires_image_input.return_value = False
task_type.accepts_image_input.return_value = False
task_type.is_image_gen.return_value = False
task_type.data_type.return_value = ModelTaskType.T2V.data_type()
server_args.pipeline_config.task_type = task_type
server_args.pipeline_config.task_type = ModelTaskType.T2V
server_args.pipeline_config.adjust_num_frames.side_effect = lambda value: value
sampling_defaults = SamplingParams(
negative_prompt="model default negative",
@@ -418,12 +401,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args.enable_cfg_parallel = False
server_args.enable_torch_compile = False
task_type = MagicMock()
task_type.requires_image_input.return_value = False
task_type.accepts_image_input.return_value = False
task_type.is_image_gen.return_value = True
task_type.data_type.return_value = ModelTaskType.T2I.data_type()
server_args.pipeline_config.task_type = task_type
server_args.pipeline_config.task_type = ModelTaskType.T2I
sampling_defaults = SamplingParams(
width=1024,
@@ -452,12 +430,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args.enable_torch_compile = False
server_args.backend = "auto"
task_type = MagicMock()
task_type.requires_image_input.return_value = False
task_type.accepts_image_input.return_value = False
task_type.is_image_gen.return_value = True
task_type.data_type.return_value = ModelTaskType.T2I.data_type()
server_args.pipeline_config.task_type = task_type
server_args.pipeline_config.task_type = ModelTaskType.T2I
sampling_defaults = SamplingParams(width=1024, height=1024)
with patch(
@@ -481,12 +454,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args.enable_torch_compile = False
server_args.backend = "diffusers"
task_type = MagicMock()
task_type.requires_image_input.return_value = False
task_type.accepts_image_input.return_value = False
task_type.is_image_gen.return_value = True
task_type.data_type.return_value = ModelTaskType.T2I.data_type()
server_args.pipeline_config.task_type = task_type
server_args.pipeline_config.task_type = ModelTaskType.T2I
sampling_defaults = SamplingParams(width=1024, height=1024)
with patch(
@@ -507,12 +475,8 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args.enable_cfg_parallel = False
server_args.enable_torch_compile = False
task_type = MagicMock()
task_type.requires_image_input.return_value = False
task_type.accepts_image_input.return_value = False
task_type.is_image_gen.return_value = False
task_type.data_type.return_value = ModelTaskType.T2V.data_type()
server_args.pipeline_config.task_type = task_type
server_args.pipeline_config.task_type = ModelTaskType.T2V
server_args.pipeline_config.adjust_num_frames.side_effect = lambda value: value
sampling_defaults = SamplingParams(
width=832,
@@ -533,18 +497,35 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
self.assertEqual(reqs[0].num_inference_steps, 2)
self.assertEqual(reqs[0].num_frames, 17)
def test_video_warmup_preserves_model_frame_alignment(self):
pipeline_config = LongLive2T2VConfig()
server_args = SimpleNamespace(
pipeline_config=pipeline_config,
enable_breakable_cuda_graph=False,
)
num_frames = _resolve_warmup_num_frames(
server_args,
LongLive2SamplingParams(),
server_based_warmup=True,
)
temporal_scale = pipeline_config.vae_config.arch_config.scale_factor_temporal
latent_frames = (num_frames - 1) // temporal_scale + 1
self.assertEqual(num_frames, 29)
self.assertEqual(
latent_frames % pipeline_config.dit_config.arch_config.num_frames_per_block,
0,
)
def test_server_based_warmup_uses_video_supported_resolution_budget(self):
server_args = MagicMock()
server_args.warmup_steps = 1
server_args.enable_cfg_parallel = False
server_args.enable_torch_compile = False
task_type = MagicMock()
task_type.requires_image_input.return_value = False
task_type.accepts_image_input.return_value = False
task_type.is_image_gen.return_value = False
task_type.data_type.return_value = ModelTaskType.T2V.data_type()
server_args.pipeline_config.task_type = task_type
server_args.pipeline_config.task_type = ModelTaskType.T2V
server_args.pipeline_config.adjust_num_frames.side_effect = lambda value: value
sampling_defaults = SamplingParams(
width=1280,
@@ -580,13 +561,9 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args.enable_torch_compile = False
server_args.pipeline_class_name = "LTX2TwoStageHQPipeline"
task_type = MagicMock()
task_type.requires_image_input.return_value = False
task_type.accepts_image_input.return_value = False
task_type.is_image_gen.return_value = False
task_type.data_type.return_value = ModelTaskType.T2V.data_type()
server_args.pipeline_config.task_type = task_type
server_args.pipeline_config.task_type = ModelTaskType.T2V
server_args.pipeline_config.vae_scale_factor = 32
server_args.pipeline_config.adjust_num_frames.side_effect = lambda value: value
sampling_defaults = SamplingParams(
width=1920,
@@ -614,12 +591,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args.enable_cfg_parallel = False
server_args.enable_torch_compile = False
task_type = MagicMock()
task_type.requires_image_input.return_value = False
task_type.accepts_image_input.return_value = False
task_type.is_image_gen.return_value = True
task_type.data_type.return_value = ModelTaskType.T2I.data_type()
server_args.pipeline_config.task_type = task_type
server_args.pipeline_config.task_type = ModelTaskType.T2I
with patch(
"sglang.multimodal_gen.runtime.warmup_request_builder.get_model_sampling_defaults",
@@ -159,20 +159,50 @@ def test_srt_owned_groups_are_not_overwritten_or_cleared():
def test_srt_tp_groups_follow_encoder_folding_context():
original_tp_group = object()
original_diffusion_tp_group = object()
original_srt_tp_group = object()
original_srt_attention_tp_group = object()
folding_tp_group = object()
with (
patch.object(parallel_state, "_TP", original_tp_group),
patch.object(parallel_state, "_TP_STATE_PATCHED", False),
patch.object(srt_parallel_state, "_TP", original_tp_group),
patch.object(srt_parallel_state, "_ATTN_TP", original_tp_group),
patch.object(parallel_state, "_TP", original_diffusion_tp_group),
patch.object(srt_parallel_state, "_TP", original_srt_tp_group),
patch.object(
srt_parallel_state,
"_ATTN_TP",
original_srt_attention_tp_group,
),
):
with parallel_state.patch_tensor_parallel_group(folding_tp_group):
with parallel_state.use_tensor_parallel_group(folding_tp_group):
assert parallel_state._TP is folding_tp_group
assert srt_parallel_state._TP is folding_tp_group
assert srt_parallel_state._ATTN_TP is folding_tp_group
assert parallel_state._TP is original_diffusion_tp_group
assert srt_parallel_state._TP is original_srt_tp_group
assert srt_parallel_state._ATTN_TP is original_srt_attention_tp_group
def test_encoder_folding_context_is_nested_and_restores_each_group():
original_tp_group = object()
outer_tp_group = object()
inner_tp_group = object()
with (
patch.object(parallel_state, "_TP", original_tp_group),
patch.object(srt_parallel_state, "_TP", original_tp_group),
patch.object(srt_parallel_state, "_ATTN_TP", original_tp_group),
):
with parallel_state.use_tensor_parallel_group(outer_tp_group):
with parallel_state.use_tensor_parallel_group(inner_tp_group):
assert parallel_state._TP is inner_tp_group
assert srt_parallel_state._TP is inner_tp_group
assert srt_parallel_state._ATTN_TP is inner_tp_group
assert parallel_state._TP is outer_tp_group
assert srt_parallel_state._TP is outer_tp_group
assert srt_parallel_state._ATTN_TP is outer_tp_group
assert parallel_state._TP is original_tp_group
assert srt_parallel_state._TP is original_tp_group
assert srt_parallel_state._ATTN_TP is original_tp_group
@@ -5,9 +5,12 @@
"""
import asyncio
import tempfile
import unittest
from pathlib import Path
from types import SimpleNamespace
from unittest import mock
from urllib.error import URLError
from sglang.multimodal_gen.runtime.entrypoints import http_server
from sglang.multimodal_gen.runtime.entrypoints.http_server import (
@@ -15,6 +18,7 @@ from sglang.multimodal_gen.runtime.entrypoints.http_server import (
health_generate,
liveness,
)
from sglang.multimodal_gen.test.server.test_server_utils import ServerManager
def _make_request(warmup_done) -> SimpleNamespace:
@@ -90,5 +94,46 @@ class TestWaitUntilHttpLive(unittest.IsolatedAsyncioTestCase):
self.assertEqual(fake_client.urls, ["http://127.0.0.1:11000/liveness"] * 2)
class _ReadyResponse:
status = 200
def __enter__(self):
return self
def __exit__(self, *exc_info):
return False
class _RunningProcess:
returncode = None
def poll(self):
return None
class TestServerManagerReadiness(unittest.TestCase):
def test_waits_for_health_after_http_startup(self):
manager = ServerManager("test-model", port=11000, wait_deadline=1)
with tempfile.TemporaryDirectory() as temp_dir:
stdout_path = Path(temp_dir) / "server.log"
stdout_path.write_text("Application startup complete.\n", encoding="utf-8")
with (
mock.patch(
"sglang.multimodal_gen.test.server.test_server_utils.urlopen",
side_effect=[URLError("warming up"), _ReadyResponse()],
) as health_request,
mock.patch(
"sglang.multimodal_gen.test.server.test_server_utils.time.sleep"
),
):
manager._wait_for_ready(_RunningProcess(), stdout_path)
self.assertEqual(health_request.call_count, 2)
self.assertEqual(
[call.args[0] for call in health_request.call_args_list],
["http://127.0.0.1:11000/health"] * 2,
)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,116 @@
from types import SimpleNamespace
from unittest.mock import patch
import torch
from torch import nn
from sglang.multimodal_gen.runtime.models.encoders import clip as mmgen_clip
from sglang.srt.models import clip as srt_clip
def _clip_config():
return SimpleNamespace(
hidden_size=16,
intermediate_size=32,
num_attention_heads=2,
num_hidden_layers=1,
layer_norm_eps=1e-5,
hidden_act="quick_gelu",
vocab_size=32,
max_position_embeddings=8,
eos_token_id=2,
output_hidden_states=False,
attention_dropout=0.0,
)
class _FakeQKV(nn.Module):
def forward(self, hidden_states):
return torch.cat((hidden_states, hidden_states, hidden_states), dim=-1), None
class _FakeProjection(nn.Module):
def forward(self, hidden_states):
return hidden_states, None
def test_mmgen_clip_reuses_srt_components():
assert mmgen_clip.CLIPEncoder is srt_clip.CLIPEncoder
assert mmgen_clip.CLIPTextEmbeddings is srt_clip.CLIPTextEmbeddings
assert mmgen_clip.CLIPVisionEmbeddings is srt_clip.CLIPVisionEmbeddings
def test_clip_encoder_propagates_causal_semantics():
with (
patch.object(srt_clip, "CLIPAttention", return_value=nn.Identity()) as attn,
patch.object(srt_clip, "CLIPMLP", return_value=nn.Identity()),
):
srt_clip.CLIPEncoder(_clip_config(), causal=True)
assert attn.call_args.kwargs["causal"] is True
def test_mmgen_text_clip_requests_masked_srt_attention():
with patch.object(mmgen_clip, "CLIPEncoder", return_value=nn.Identity()) as encoder:
mmgen_clip.CLIPTextTransformer(_clip_config(), prefix="text_model.encoder")
assert encoder.call_args.kwargs["causal"] is True
def test_clip_attention_separates_text_and_vision_semantics():
parallel = SimpleNamespace(attn_tp_size=1, attn_tp_rank=0)
hidden_states = torch.randn(2, 3, 16)
padding_mask = torch.zeros(2, 1, 3, 3)
with (
patch.object(srt_clip, "get_parallel", return_value=parallel),
patch.object(srt_clip, "QKVParallelLinear", return_value=_FakeQKV()),
patch.object(srt_clip, "RowParallelLinear", return_value=_FakeProjection()),
patch.object(
srt_clip.F,
"scaled_dot_product_attention",
side_effect=lambda query, key, value, **kwargs: query,
) as sdpa,
):
text_attention = srt_clip.CLIPAttention(_clip_config(), causal=True)
vision_attention = srt_clip.CLIPAttention(_clip_config())
text_attention(hidden_states)
text_attention(hidden_states, attention_mask=padding_mask)
vision_attention(hidden_states)
assert sdpa.call_args_list[0].kwargs["is_causal"] is True
assert sdpa.call_args_list[1].kwargs["is_causal"] is False
assert sdpa.call_args_list[1].kwargs["attn_mask"] is padding_mask
assert sdpa.call_args_list[2].kwargs["is_causal"] is False
def test_prepare_clip_attention_mask_combines_causal_and_padding_masks():
mask = srt_clip.prepare_clip_attention_mask(
torch.Size((1, 3)),
torch.float32,
torch.device("cpu"),
torch.tensor([[1, 1, 0]]),
)
assert mask.shape == (1, 1, 3, 3)
assert mask[0, 0, 0, 0] == 0
assert mask[0, 0, 0, 1] < -1e20
assert torch.all(mask[..., 2] < -1e20)
def test_prepare_clip_attention_mask_keeps_unmasked_fast_path():
assert (
srt_clip.prepare_clip_attention_mask(
torch.Size((2, 3)), torch.float32, torch.device("cpu")
)
is None
)
def test_srt_clip_weight_name_mapping():
assert (
mmgen_clip._srt_clip_param_name(
"text_model.encoder.layers.0.self_attn.out_proj.weight"
)
== "text_model.encoder.layers.0.self_attn.proj.weight"
)
@@ -1,4 +1,3 @@
from contextlib import contextmanager
from types import SimpleNamespace
from unittest.mock import patch
@@ -6,6 +5,9 @@ from torch import nn
from sglang.multimodal_gen.runtime.loader.utils import get_param_names_mapping
from sglang.multimodal_gen.runtime.models.encoders import gemma_3
from sglang.multimodal_gen.runtime.models.encoders.base import (
EncoderTensorParallelMixin,
)
from sglang.srt.models import siglip
@@ -35,10 +37,8 @@ def test_siglip_encoder_propagates_attention_backend():
def test_gemma3_uses_srt_siglip_with_stable_backend():
config = SimpleNamespace(vision_config=object(), text_config=object())
folding_group = object()
with (
patch.object(gemma_3, "get_tp_group", return_value=folding_group),
patch.object(
gemma_3,
"SiglipVisionModel",
@@ -63,42 +63,8 @@ def test_gemma3_uses_srt_siglip_with_stable_backend():
quant_config=None,
prefix="vision_tower",
)
assert model._vision_tensor_parallel_group is folding_group
def test_gemma3_restores_vision_tensor_parallel_group():
model = gemma_3.Gemma3ForConditionalGeneration.__new__(
gemma_3.Gemma3ForConditionalGeneration
)
nn.Module.__init__(model)
folding_group = object()
active_group = object()
model._vision_tensor_parallel_group = folding_group
events = []
@contextmanager
def use_group(group):
events.append(("enter", group))
yield
events.append(("exit", group))
with (
patch.object(gemma_3, "get_tp_group", return_value=active_group),
patch.object(
gemma_3,
"patch_tensor_parallel_group",
side_effect=use_group,
) as patch_group,
):
with model._vision_parallel_context():
events.append(("forward", folding_group))
patch_group.assert_called_once_with(folding_group)
assert events == [
("enter", folding_group),
("forward", folding_group),
("exit", folding_group),
]
assert isinstance(model, EncoderTensorParallelMixin)
assert not hasattr(model, "_vision_tensor_parallel_group")
def test_gemma3_maps_hf_siglip_projection_name():
+136 -23
View File
@@ -6,21 +6,48 @@ from typing import Iterable, List, Optional, Tuple, Type, Union
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers import CLIPConfig, CLIPTextConfig, CLIPVisionConfig
from transformers.modeling_attn_mask_utils import _create_4d_causal_attention_mask
from sglang.srt.layers.activation import QuickGELU
from sglang.srt.layers.attention.vision import VisionAttention
from sglang.srt.layers.activation import QuickGELU, get_act_fn
from sglang.srt.layers.conv import Conv2dLayer
from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear
from sglang.srt.layers.linear import (
ColumnParallelLinear,
QKVParallelLinear,
RowParallelLinear,
)
from sglang.srt.layers.pooler import EmbeddingPoolerOutput, Pooler, PoolingType
from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.managers.schedule_batch import MultimodalInputs
from sglang.srt.model_executor.model_runner import ForwardBatch
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix, flatten_nested_list
def prepare_clip_attention_mask(
input_shape: torch.Size,
dtype: torch.dtype,
device: torch.device,
attention_mask: Optional[torch.Tensor] = None,
) -> Optional[torch.Tensor]:
if attention_mask is None:
return None
batch_size, sequence_length = input_shape
causal_mask = torch.full(
(sequence_length, sequence_length),
torch.finfo(dtype).min,
dtype=dtype,
device=device,
)
causal_mask = torch.triu(causal_mask, diagonal=1)
causal_mask = causal_mask[None, None].expand(batch_size, 1, -1, -1)
if attention_mask.dim() == 2:
attention_mask = attention_mask[:, None, None, :].to(dtype=dtype)
attention_mask = (1.0 - attention_mask) * torch.finfo(dtype).min
return causal_mask + attention_mask
class CLIPVisionEmbeddings(nn.Module):
def __init__(self, config: CLIPVisionConfig):
@@ -88,9 +115,18 @@ class CLIPTextEmbeddings(nn.Module):
position_ids: Optional[torch.LongTensor] = None,
inputs_embeds: Optional[torch.FloatTensor] = None,
) -> torch.Tensor:
seq_length = (
input_ids.shape[-1] if input_ids is not None else inputs_embeds.shape[-2]
)
if input_ids is not None:
seq_length = input_ids.shape[-1]
elif inputs_embeds is not None:
seq_length = inputs_embeds.shape[-2]
else:
raise ValueError("Either input_ids or inputs_embeds must be provided.")
max_positions = self.position_embedding.weight.shape[0]
if seq_length > max_positions:
raise ValueError(
f"Sequence length {seq_length} exceeds the maximum {max_positions}."
)
if position_ids is None:
position_ids = self.position_ids[:, :seq_length]
@@ -109,7 +145,7 @@ class CLIPMLP(nn.Module):
def __init__(
self,
config,
act_layer: Type[nn.Module] = QuickGELU,
act_layer: Optional[Type[nn.Module]] = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
@@ -120,7 +156,12 @@ class CLIPMLP(nn.Module):
quant_config=quant_config,
prefix=add_prefix("fc1", prefix),
)
self.act = act_layer()
if act_layer is not None:
self.act = act_layer()
elif config.hidden_act == "quick_gelu":
self.act = QuickGELU()
else:
self.act = get_act_fn(config.hidden_act)
self.fc2 = RowParallelLinear(
config.intermediate_size,
config.hidden_size,
@@ -135,29 +176,90 @@ class CLIPMLP(nn.Module):
return x
class CLIPAttention(nn.Module):
def __init__(
self,
config: Union[CLIPTextConfig, CLIPVisionConfig],
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
causal: bool = False,
) -> None:
super().__init__()
parallel = get_parallel()
self.num_heads = config.num_attention_heads // parallel.attn_tp_size
self.head_dim = config.hidden_size // config.num_attention_heads
self.causal = causal
self.dropout = config.attention_dropout
self.scale = self.head_dim**-0.5
self.qkv_proj = QKVParallelLinear(
hidden_size=config.hidden_size,
head_size=self.head_dim,
total_num_heads=config.num_attention_heads,
bias=True,
quant_config=quant_config,
prefix=add_prefix("qkv_proj", prefix),
tp_rank=parallel.attn_tp_rank,
tp_size=parallel.attn_tp_size,
)
self.proj = RowParallelLinear(
input_size=config.hidden_size,
output_size=config.hidden_size,
bias=True,
quant_config=quant_config,
prefix=add_prefix("proj", prefix),
tp_rank=parallel.attn_tp_rank,
tp_size=parallel.attn_tp_size,
)
def forward(
self,
hidden_states: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None,
) -> torch.Tensor:
batch_size, sequence_length, _ = hidden_states.shape
qkv, _ = self.qkv_proj(hidden_states)
query, key, value = qkv.chunk(3, dim=-1)
qkv_shape = (batch_size, sequence_length, self.num_heads, self.head_dim)
query = query.view(qkv_shape).transpose(1, 2)
key = key.view(qkv_shape).transpose(1, 2)
value = value.view(qkv_shape).transpose(1, 2)
output = F.scaled_dot_product_attention(
query,
key,
value,
attn_mask=attention_mask,
dropout_p=self.dropout if self.training else 0.0,
is_causal=self.causal and attention_mask is None,
scale=self.scale,
)
output = output.transpose(1, 2).reshape(
batch_size, sequence_length, self.num_heads * self.head_dim
)
output, _ = self.proj(output)
return output
class CLIPEncoderLayer(nn.Module):
def __init__(
self,
config: CLIPVisionConfig,
act_layer: Type[nn.Module] = QuickGELU,
act_layer: Optional[Type[nn.Module]] = None,
norm_layer: Type[nn.Module] = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
causal: bool = False,
) -> None:
super().__init__()
if norm_layer is None:
norm_layer = partial(nn.LayerNorm, eps=config.layer_norm_eps)
self.layer_norm1 = norm_layer(config.hidden_size)
self.layer_norm2 = norm_layer(config.hidden_size)
self.self_attn = VisionAttention(
embed_dim=config.hidden_size,
num_heads=config.num_attention_heads,
projection_size=config.hidden_size,
use_qkv_parallel=True,
flatten_batch=True,
self.self_attn = CLIPAttention(
config,
quant_config=quant_config,
prefix=add_prefix("self_attn", prefix),
causal=causal,
)
self.mlp = CLIPMLP(
config,
@@ -210,20 +312,29 @@ class CLIPEncoder(nn.Module):
config: CLIPVisionConfig,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
num_hidden_layers_override: Optional[int] = None,
act_layer: Optional[Type[nn.Module]] = None,
causal: bool = False,
) -> None:
super().__init__()
self.config = config
num_hidden_layers = config.num_hidden_layers
num_hidden_layers = (
config.num_hidden_layers
if num_hidden_layers_override is None
else num_hidden_layers_override
)
norm_layer = partial(nn.LayerNorm, eps=config.layer_norm_eps)
self.layers = nn.ModuleList(
[
CLIPEncoderLayer(
config=config,
act_layer=act_layer,
norm_layer=norm_layer,
quant_config=quant_config,
prefix=add_prefix(f"layers.{layer_idx}", prefix),
causal=causal,
)
for layer_idx in range(num_hidden_layers)
]
@@ -265,6 +376,7 @@ class CLIPTextTransformer(nn.Module):
config=config,
quant_config=quant_config,
prefix=add_prefix("encoder", prefix),
causal=True,
)
self.final_layer_norm = nn.LayerNorm(embed_dim, eps=config.layer_norm_eps)
@@ -281,12 +393,13 @@ class CLIPTextTransformer(nn.Module):
input_shape = input_ids.size()
input_ids = input_ids.view(-1, input_shape[-1])
hidden_states = self.embeddings(input_ids, position_ids)
causal_attention_mask = _create_4d_causal_attention_mask(
input_ids.shape, hidden_states.dtype, device=hidden_states.device
)
encoder_outputs = self.encoder(
hidden_states, attention_mask, causal_attention_mask
attention_mask = prepare_clip_attention_mask(
input_ids.shape,
hidden_states.dtype,
hidden_states.device,
attention_mask,
)
encoder_outputs = self.encoder(hidden_states, attention_mask=attention_mask)
last_hidden_state = self.final_layer_norm(encoder_outputs)
return last_hidden_state
@@ -311,7 +424,7 @@ class CLIPTextModel(nn.Module):
input_ids: torch.Tensor,
position_ids: torch.Tensor,
):
return self.text_model(input_ids, position_ids)
return self.text_model(input_ids, position_ids=position_ids)
class CLIPVisionTransformer(nn.Module):