[diffusion] chore: reuse SRT CLIP encoder blocks (#35004)
This commit is contained in:
@@ -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:
|
||||
|
||||
+13
-18
@@ -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():
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user