diff --git a/docs/docs/sglang-diffusion/compatibility_matrix.mdx b/docs/docs/sglang-diffusion/compatibility_matrix.mdx
index 34393e9d4..e13a09473 100644
--- a/docs/docs/sglang-diffusion/compatibility_matrix.mdx
+++ b/docs/docs/sglang-diffusion/compatibility_matrix.mdx
@@ -37,6 +37,10 @@ Rows are grouped when a family shares the same runtime path or optimization supp
Qwen-Image |
Qwen/Qwen-ImageQwen/Qwen-Image-2512Qwen/Qwen-Image-EditQwen/Qwen-Image-Edit-2509Qwen/Qwen-Image-Edit-2511Qwen/Qwen-Image-Layered
|
+
+ | LongCat-Image |
+ meituan-longcat/LongCat-Image
|
+
| SD3 / SD3.5 |
stabilityai/stable-diffusion-3-mediumstabilityai/stable-diffusion-3-medium-diffusersstabilityai/stable-diffusion-3.5-mediumstabilityai/stable-diffusion-3.5-medium-diffusersstabilityai/stable-diffusion-3.5-largestabilityai/stable-diffusion-3.5-large-diffusers
|
diff --git a/docs/docs/sglang-diffusion/models_with_ar.mdx b/docs/docs/sglang-diffusion/models_with_ar.mdx
index 2bff6aaf1..1d5c38167 100644
--- a/docs/docs/sglang-diffusion/models_with_ar.mdx
+++ b/docs/docs/sglang-diffusion/models_with_ar.mdx
@@ -1,11 +1,16 @@
---
-title: "Diffusion models with AR stage like GLM-Image"
-description: "Run diffusion pipelines that delegate an AR stage to a separate SGLang server, such as GLM-Image."
+title: "Diffusion models with autoregressive stages"
+description: "Run diffusion pipelines with in-process or separately deployed autoregressive encoders."
---
-## Quick Start
+SGLang Diffusion supports two AR execution paths. Qwen Image Layered and
+LongCat-Image use the native Qwen2.5-VL component in process. GLM-Image can use
+its bundled Transformers implementation or delegate AR inference to a separate
+SGLang server.
-Run model with transformers implementation for AR stage (default)
+## GLM-Image Quick Start
+
+Run GLM-Image with its bundled Transformers implementation (default):
```bash
# Terminal 1 : launch server
sglang serve --model-path zai-org/GLM-Image --port ${PORT}
@@ -20,7 +25,7 @@ curl http://${HOST}:${PORT}/v1/images/generations \
"size": "widthxheight"
}'
```
-Run model with SGLang srt implementation for AR stage (high performance)
+Run GLM-Image with a separate SGLang server for its AR stage:
```bash
# Terminal 1 : launch server with AR model
sglang serve --model-path /path/to/zai-org/GLM-Image/vision_language_encoder/ \
@@ -64,6 +69,16 @@ curl http://${HOST}:${PORT}/v1/images/generations \
T2I, I2I, V2I |
T2I |
+
+ | Qwen Image Layered |
+ Not used |
+ I2I, in-process native Qwen2.5-VL |
+
+
+ | LongCat-Image |
+ Not used |
+ T2I, in-process native Qwen2.5-VL |
+
diff --git a/python/sglang/multimodal_gen/configs/models/encoders/base.py b/python/sglang/multimodal_gen/configs/models/encoders/base.py
index 76180435f..8a659f5a0 100644
--- a/python/sglang/multimodal_gen/configs/models/encoders/base.py
+++ b/python/sglang/multimodal_gen/configs/models/encoders/base.py
@@ -86,6 +86,7 @@ class EncoderConfig(ModelConfig):
@dataclass
class TextEncoderConfig(EncoderConfig):
arch_config: ArchConfig = field(default_factory=TextEncoderArchConfig)
+ generation_config: dict[str, Any] = field(default_factory=dict)
@dataclass
diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/longcat_image.py b/python/sglang/multimodal_gen/configs/pipeline_configs/longcat_image.py
index f9cad23ab..8f620e9c0 100644
--- a/python/sglang/multimodal_gen/configs/pipeline_configs/longcat_image.py
+++ b/python/sglang/multimodal_gen/configs/pipeline_configs/longcat_image.py
@@ -8,6 +8,7 @@ from sglang.multimodal_gen.configs.models import DiTConfig, EncoderConfig, VAECo
from sglang.multimodal_gen.configs.models.dits.longcat_image import (
LongCatImageDitConfig,
)
+from sglang.multimodal_gen.configs.models.encoders.qwen_image import Qwen2_5VLConfig
from sglang.multimodal_gen.configs.models.vaes.longcat_image import (
LongCatImageVAEConfig,
)
@@ -16,6 +17,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.base import (
ModelTaskType,
TextConditioningOutput,
)
+from sglang.multimodal_gen.configs.pipeline_configs.model_deployment_config import (
+ ModelDeploymentConfig,
+)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
@@ -219,18 +223,6 @@ def longcat_postprocess_text(outputs, text_inputs, pipeline_config):
)
-@dataclass
-class LongCatImageEncoderConfig(EncoderConfig):
- """Encoder config for the in-stage-loaded HF Qwen2.5-VL text encoder.
-
- The encoder weights are loaded by `LongCatPromptRewriteStage` (not via
- `TextEncoderLoader`), so this config only supplies the fields the standard
- `TextEncodingStage` reads — primarily `tokenizer_kwargs`.
- """
-
- tokenizer_kwargs: dict = field(default_factory=lambda: {})
-
-
@dataclass
class LongCatImagePipelineConfig(ImagePipelineConfig):
"""Configuration for the LongCat-Image T2I pipeline."""
@@ -246,16 +238,20 @@ class LongCatImagePipelineConfig(ImagePipelineConfig):
dit_config: DiTConfig = field(default_factory=LongCatImageDitConfig)
vae_config: VAEConfig = field(default_factory=LongCatImageVAEConfig)
- # The Qwen2.5-VL text encoder (~7B) is loaded in bf16; the encoder is loaded
- # in-stage by LongCatPromptRewriteStage, not via TextEncoderLoader.
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16",))
text_encoder_configs: tuple[EncoderConfig, ...] = field(
- default_factory=lambda: (LongCatImageEncoderConfig(),)
+ default_factory=lambda: (Qwen2_5VLConfig(),)
)
postprocess_text_funcs: tuple[Callable, ...] = field(
default_factory=lambda: (longcat_postprocess_text,)
)
+ def get_model_deployment_config(self) -> ModelDeploymentConfig:
+ return ModelDeploymentConfig(
+ keep_resident_min_available_gb=70,
+ keep_resident_components=("text_encoder", "vae"),
+ )
+
# --- LatentPreparationStage hooks ---
def prepare_latent_shape(self, batch, batch_size, num_frames):
diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/qwen_image.py b/python/sglang/multimodal_gen/configs/pipeline_configs/qwen_image.py
index 610267b7c..91169c686 100644
--- a/python/sglang/multimodal_gen/configs/pipeline_configs/qwen_image.py
+++ b/python/sglang/multimodal_gen/configs/pipeline_configs/qwen_image.py
@@ -19,6 +19,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.base import (
pad_text_embeddings_with_mask,
shard_rotary_emb_for_sp,
)
+from sglang.multimodal_gen.configs.pipeline_configs.model_deployment_config import (
+ ModelDeploymentConfig,
+)
from sglang.multimodal_gen.configs.post_training.pipeline_configs import (
QwenImageRolloutPipelineMixin,
)
@@ -753,6 +756,12 @@ class QwenImageLayeredPipelineConfig(QwenImageEditPipelineConfig):
resolution: int = 640
vae_precision: str = "bf16"
+ def get_model_deployment_config(self) -> ModelDeploymentConfig:
+ return ModelDeploymentConfig(
+ keep_resident_min_available_gb=70,
+ keep_resident_components=("text_encoder", "vae"),
+ )
+
def postprocess_cfg_noise(
self,
batch,
diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py
index 5da710d42..d033cf811 100644
--- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py
+++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py
@@ -52,6 +52,7 @@ from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
get_config,
get_diffusers_component_config,
+ load_dict,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.runtime.utils.precision import precision_to_dtype
@@ -344,6 +345,9 @@ class TextEncoderLoader(ComponentLoader):
encoder_config = server_args.pipeline_config.text_encoder_configs[encoder_index]
encoder_config.update_model_arch(model_config)
+ encoder_config.generation_config = load_dict(
+ os.path.join(component_model_path, "generation_config.json")
+ )
if encoder_index == 0:
for key, value in diffusers_pretrained_config.__dict__.items():
diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/qwen2_5vl.py b/python/sglang/multimodal_gen/runtime/models/encoders/qwen2_5vl.py
index b7cd9e5f6..eb029d3e5 100644
--- a/python/sglang/multimodal_gen/runtime/models/encoders/qwen2_5vl.py
+++ b/python/sglang/multimodal_gen/runtime/models/encoders/qwen2_5vl.py
@@ -29,6 +29,9 @@ from sglang.multimodal_gen.runtime.layers.linear import (
from sglang.multimodal_gen.runtime.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 TextEncoder
+from sglang.multimodal_gen.runtime.models.encoders.qwen2_5vl_vision import (
+ Qwen2_5VLVisionTransformer,
+)
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
from sglang.multimodal_gen.runtime.utils.common import add_prefix
@@ -69,8 +72,6 @@ import torch
import torch.nn as nn
from transformers.activations import ACT2FN
from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import (
- Qwen2_5_VisionRotaryEmbedding,
- Qwen2_5_VisionTransformerPretrainedModel,
Qwen2_5_VLCausalLMOutputWithPast,
Qwen2_5_VLModelOutputWithPast,
Qwen2_5_VLRotaryEmbedding,
@@ -80,6 +81,53 @@ from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import (
logger = logging.getLogger(__name__)
+def _apply_repetition_penalty(
+ logits: torch.Tensor,
+ input_ids: torch.LongTensor,
+ penalty: float,
+) -> torch.Tensor:
+ if penalty == 1.0:
+ return logits
+ selected_logits = torch.gather(logits, 1, input_ids)
+ selected_logits = torch.where(
+ selected_logits < 0,
+ selected_logits * penalty,
+ selected_logits / penalty,
+ )
+ return logits.scatter(1, input_ids, selected_logits)
+
+
+def _select_next_token(
+ logits: torch.Tensor,
+ input_ids: torch.LongTensor,
+ *,
+ do_sample: bool,
+ temperature: float,
+ top_k: int,
+ top_p: float,
+ repetition_penalty: float,
+) -> torch.LongTensor:
+ scores = _apply_repetition_penalty(logits.float(), input_ids, repetition_penalty)
+ if not do_sample or top_k == 1:
+ return torch.argmax(scores, dim=-1)
+
+ scores = scores / temperature
+ if 0 < top_k < scores.shape[-1]:
+ top_k_threshold = torch.topk(scores, top_k, dim=-1).values[..., -1, None]
+ scores = scores.masked_fill(scores < top_k_threshold, -torch.inf)
+ if top_p < 1.0:
+ sorted_scores, sorted_indices = torch.sort(scores, descending=True)
+ cumulative_probs = torch.softmax(sorted_scores, dim=-1).cumsum(dim=-1)
+ sorted_remove = cumulative_probs > top_p
+ sorted_remove[..., 1:] = sorted_remove[..., :-1].clone()
+ sorted_remove[..., 0] = False
+ remove = torch.zeros_like(sorted_remove).scatter(
+ 1, sorted_indices, sorted_remove
+ )
+ scores = scores.masked_fill(remove, -torch.inf)
+ return torch.multinomial(torch.softmax(scores, dim=-1), num_samples=1).squeeze(1)
+
+
def _tp_world_size() -> int:
if not model_parallel_is_initialized():
return 1
@@ -261,7 +309,15 @@ class Qwen2_5_VLAttention(nn.Module):
query_states = query_states.transpose(1, 2)
key_states = key_states.transpose(1, 2)
value_states = value_states.transpose(1, 2)
- attn_output = self.attn(query_states, key_states, value_states)
+ # Diffusion text encoding is cache-free and historically uses the native
+ # causal kernel; its trailing padding is removed during postprocessing.
+ # Cached generation still needs the explicit mask for padded batches.
+ attn_output = self.attn(
+ query_states,
+ key_states,
+ value_states,
+ attn_mask=attention_mask if use_cache else None,
+ )
attn_output = attn_output.reshape(bsz, q_len, -1).contiguous()
attn_output = _linear_output(self.o_proj, attn_output)
@@ -597,28 +653,14 @@ class Qwen2_5_VLModel(nn.Module):
_checkpoint_conversion_mapping = {"^model": "language_model"}
# Reference: fix gemma3 grad acc #37208
accepts_loss_kwargs = False
- _no_split_modules = ["Qwen2_5_VLDecoderLayer", "Qwen2_5_VLVisionBlock"]
+ _no_split_modules = ["Qwen2_5_VLDecoderLayer", "Qwen2_5VLVisionBlock"]
def __init__(self, config, enable_image_understanding: bool = False):
super().__init__()
self.language_model = Qwen2_5_VLTextModel(config.text_config)
if enable_image_understanding:
- self.visual = Qwen2_5_VisionTransformerPretrainedModel._from_config(
- config.vision_config
- )
- self.visual.to(torch.get_default_dtype())
- # keeps the vision rotary frequencies in fp32 even when weights are bf16 (as HF does)
- head_dim = (
- config.vision_config.hidden_size // config.vision_config.num_heads
- )
- rotary_dim = head_dim // 2
- inv_freq = Qwen2_5_VisionRotaryEmbedding(rotary_dim).inv_freq
- self.visual.rotary_pos_emb.register_buffer(
- "inv_freq",
- inv_freq,
- persistent=False,
- )
+ self.visual = Qwen2_5VLVisionTransformer(config.vision_config)
self.rope_deltas = None # cache rope_deltas here
self.config = config
# Initialize weights and apply final processing
@@ -902,11 +944,6 @@ class Qwen2_5_VLModel(nn.Module):
"""
pixel_values = pixel_values.type(self.visual.dtype)
image_embeds = self.visual(pixel_values, grid_thw=image_grid_thw)
- if not isinstance(image_embeds, torch.Tensor):
- # In transformers v5, the visual encoder returns BaseModelOutputWithPooling.
- # pooler_output contains the spatially merged embeddings (what we need),
- # while last_hidden_state contains the raw unmerged output.
- image_embeds = image_embeds.pooler_output
split_sizes = (
image_grid_thw.prod(-1) // self.visual.spatial_merge_size**2
).tolist()
@@ -1106,6 +1143,9 @@ class Qwen2_5_VLModel(nn.Module):
class Qwen2_5_VLForConditionalGeneration(TextEncoder):
+ layer_names = [*TextEncoder.layer_names, "model.visual.blocks"]
+ _fsdp_forward_methods = ("generate",)
+
# BitandBytes specific attributes
default_bitsandbytes_target_modules = [
".gate_up_proj.",
@@ -1132,6 +1172,7 @@ class Qwen2_5_VLForConditionalGeneration(TextEncoder):
) -> None:
super().__init__(config)
enable_image_understanding = config.enable_image_understanding
+ generation_config = config.generation_config
config = config.arch_config
self.model = Qwen2_5_VLModel(
config, enable_image_understanding=enable_image_understanding
@@ -1141,6 +1182,7 @@ class Qwen2_5_VLForConditionalGeneration(TextEncoder):
)
self.enable_image_understanding = enable_image_understanding
+ self.generation_config = generation_config
self.config = config
@@ -1225,6 +1267,152 @@ class Qwen2_5_VLForConditionalGeneration(TextEncoder):
rope_deltas=outputs.rope_deltas,
)
+ @torch.no_grad()
+ def generate(
+ self,
+ input_ids: torch.LongTensor,
+ attention_mask: Optional[torch.Tensor] = None,
+ *,
+ max_new_tokens: int,
+ pixel_values: Optional[torch.Tensor] = None,
+ pixel_values_videos: Optional[torch.FloatTensor] = None,
+ image_grid_thw: Optional[torch.LongTensor] = None,
+ video_grid_thw: Optional[torch.LongTensor] = None,
+ second_per_grid_ts: Optional[torch.Tensor] = None,
+ mm_token_type_ids: Optional[torch.IntTensor] = None,
+ do_sample: Optional[bool] = None,
+ temperature: Optional[float] = None,
+ top_k: Optional[int] = None,
+ top_p: Optional[float] = None,
+ repetition_penalty: Optional[float] = None,
+ eos_token_id: Optional[Union[int, list[int]]] = None,
+ pad_token_id: Optional[int] = None,
+ ) -> torch.LongTensor:
+ """Generate tokens with Qwen2.5-VL's native decoder and KV cache."""
+ # Transformers 5 processors emit this field. The Qwen2.5-VL checkpoint
+ # derives modality positions from image/video placeholder token IDs.
+ del mm_token_type_ids
+
+ if input_ids.ndim != 2:
+ raise ValueError("input_ids must have shape [batch_size, sequence_length]")
+ if max_new_tokens < 0:
+ raise ValueError("max_new_tokens must be non-negative")
+ if max_new_tokens == 0:
+ return input_ids
+
+ generation_config = self.generation_config
+ do_sample = (
+ generation_config.get("do_sample", False)
+ if do_sample is None
+ else do_sample
+ )
+ temperature = (
+ generation_config.get("temperature", 1.0)
+ if temperature is None
+ else temperature
+ )
+ top_k = generation_config.get("top_k", 0) if top_k is None else top_k
+ top_p = generation_config.get("top_p", 1.0) if top_p is None else top_p
+ repetition_penalty = (
+ generation_config.get("repetition_penalty", 1.0)
+ if repetition_penalty is None
+ else repetition_penalty
+ )
+ eos_token_id = (
+ generation_config.get("eos_token_id", self.config.eos_token_id)
+ if eos_token_id is None
+ else eos_token_id
+ )
+ pad_token_id = (
+ generation_config.get("pad_token_id", self.config.pad_token_id)
+ if pad_token_id is None
+ else pad_token_id
+ )
+
+ if do_sample and temperature <= 0:
+ raise ValueError("temperature must be positive when sampling")
+ if top_k < 0:
+ raise ValueError("top_k must be non-negative")
+ if not 0 < top_p <= 1:
+ raise ValueError("top_p must be in (0, 1]")
+ if repetition_penalty <= 0:
+ raise ValueError("repetition_penalty must be positive")
+
+ eos_token_ids = (
+ []
+ if eos_token_id is None
+ else [eos_token_id] if isinstance(eos_token_id, int) else list(eos_token_id)
+ )
+ if pad_token_id is None:
+ raise ValueError("pad_token_id must be set for generation")
+ eos_tokens = torch.tensor(
+ eos_token_ids, dtype=input_ids.dtype, device=input_ids.device
+ )
+ if attention_mask is None:
+ attention_mask = torch.ones_like(input_ids)
+
+ generated_ids = input_ids
+ unfinished = torch.ones(
+ input_ids.shape[0], dtype=torch.bool, device=input_ids.device
+ )
+ past_key_values = None
+ model_input_ids = input_ids
+ cache_position = torch.arange(input_ids.shape[1], device=input_ids.device)
+ self.model.rope_deltas = None
+
+ for _ in range(max_new_tokens):
+ outputs = self(
+ input_ids=model_input_ids,
+ attention_mask=attention_mask,
+ past_key_values=past_key_values,
+ use_cache=True,
+ cache_position=cache_position,
+ pixel_values=pixel_values,
+ pixel_values_videos=pixel_values_videos,
+ image_grid_thw=image_grid_thw,
+ video_grid_thw=video_grid_thw,
+ second_per_grid_ts=second_per_grid_ts,
+ logits_to_keep=1,
+ )
+ next_tokens = _select_next_token(
+ outputs.logits[:, -1, :],
+ generated_ids,
+ do_sample=do_sample,
+ temperature=temperature,
+ top_k=top_k,
+ top_p=top_p,
+ repetition_penalty=repetition_penalty,
+ )
+ next_tokens = torch.where(
+ unfinished,
+ next_tokens,
+ torch.full_like(next_tokens, pad_token_id),
+ )
+ generated_ids = torch.cat([generated_ids, next_tokens[:, None]], dim=-1)
+
+ if eos_tokens.numel() > 0:
+ reached_eos = (next_tokens[:, None] == eos_tokens[None, :]).any(dim=-1)
+ unfinished = unfinished & ~reached_eos
+ if not unfinished.any():
+ break
+
+ past_key_values = outputs.past_key_values
+ model_input_ids = next_tokens[:, None]
+ attention_mask = torch.cat(
+ [attention_mask, attention_mask.new_ones((input_ids.shape[0], 1))],
+ dim=-1,
+ )
+ cache_position = torch.tensor(
+ [generated_ids.shape[1] - 1], device=input_ids.device
+ )
+ pixel_values = None
+ pixel_values_videos = None
+ image_grid_thw = None
+ video_grid_thw = None
+ second_per_grid_ts = None
+
+ return generated_ids
+
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
loaded_params: set[str] = set()
diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/qwen2_5vl_vision.py b/python/sglang/multimodal_gen/runtime/models/encoders/qwen2_5vl_vision.py
new file mode 100644
index 000000000..acd3c4647
--- /dev/null
+++ b/python/sglang/multimodal_gen/runtime/models/encoders/qwen2_5vl_vision.py
@@ -0,0 +1,461 @@
+# SPDX-License-Identifier: Apache-2.0
+"""Native Qwen2.5-VL vision encoder."""
+
+from __future__ import annotations
+
+from dataclasses import dataclass
+from typing import Any
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+
+from sglang.multimodal_gen.runtime.layers.attention.selector import get_attn_backend
+from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
+
+logger = init_logger(__name__)
+
+
+@dataclass(frozen=True)
+class _PackedSequenceMetadata:
+ cu_seqlens: torch.Tensor
+ cu_seqlens_host: tuple[int, ...]
+ max_seqlen: int
+
+ @classmethod
+ def from_cu_seqlens(cls, cu_seqlens: torch.Tensor) -> _PackedSequenceMetadata:
+ bounds = tuple(int(value) for value in cu_seqlens.tolist())
+ return cls(
+ cu_seqlens=cu_seqlens,
+ cu_seqlens_host=bounds,
+ max_seqlen=max(
+ stop - start for start, stop in zip(bounds[:-1], bounds[1:])
+ ),
+ )
+
+
+class Qwen2_5VLVisionRMSNorm(nn.Module):
+ def __init__(self, hidden_size: int, eps: float = 1e-6) -> None:
+ super().__init__()
+ self.weight = nn.Parameter(torch.ones(hidden_size))
+ self.variance_epsilon = eps
+
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
+ input_dtype = hidden_states.dtype
+ hidden_states = hidden_states.float()
+ variance = hidden_states.square().mean(dim=-1, keepdim=True)
+ hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
+ return self.weight * hidden_states.to(input_dtype)
+
+
+class Qwen2_5VLVisionPatchEmbed(nn.Module):
+ def __init__(
+ self,
+ patch_size: int,
+ temporal_patch_size: int,
+ in_channels: int,
+ embed_dim: int,
+ ) -> None:
+ super().__init__()
+ self.patch_size = patch_size
+ self.temporal_patch_size = temporal_patch_size
+ self.in_channels = in_channels
+ self.embed_dim = embed_dim
+ kernel_size = (temporal_patch_size, patch_size, patch_size)
+ self.proj = nn.Conv3d(
+ in_channels,
+ embed_dim,
+ kernel_size=kernel_size,
+ stride=kernel_size,
+ bias=False,
+ )
+
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
+ hidden_states = hidden_states.view(
+ -1,
+ self.in_channels,
+ self.temporal_patch_size,
+ self.patch_size,
+ self.patch_size,
+ )
+ return self.proj(hidden_states.to(self.proj.weight.dtype)).view(
+ -1, self.embed_dim
+ )
+
+
+class Qwen2_5VLVisionRotaryEmbedding(nn.Module):
+ def __init__(self, dim: int, theta: float = 10000.0) -> None:
+ super().__init__()
+ inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim))
+ self.register_buffer("inv_freq", inv_freq, persistent=False)
+
+ def forward(self, position_ids: torch.Tensor) -> torch.Tensor:
+ return (position_ids.unsqueeze(-1) * self.inv_freq).flatten(1)
+
+
+def _rotate_half(hidden_states: torch.Tensor) -> torch.Tensor:
+ first, second = hidden_states.chunk(2, dim=-1)
+ return torch.cat((-second, first), dim=-1)
+
+
+def _apply_vision_rotary_embedding(
+ query: torch.Tensor,
+ key: torch.Tensor,
+ cos: torch.Tensor,
+ sin: torch.Tensor,
+) -> tuple[torch.Tensor, torch.Tensor]:
+ query_dtype = query.dtype
+ key_dtype = key.dtype
+ query = query.float()
+ key = key.float()
+ cos = cos.unsqueeze(-2).float()
+ sin = sin.unsqueeze(-2).float()
+ query = query * cos + _rotate_half(query) * sin
+ key = key * cos + _rotate_half(key) * sin
+ return query.to(query_dtype), key.to(key_dtype)
+
+
+class Qwen2_5VLVisionAttention(nn.Module):
+ def __init__(self, config: Any, prefix: str) -> None:
+ super().__init__()
+ self.num_heads = config.num_heads
+ self.head_dim = config.hidden_size // config.num_heads
+ self.scaling = self.head_dim**-0.5
+ self.prefix = prefix
+ self.qkv = nn.Linear(config.hidden_size, config.hidden_size * 3, bias=True)
+ self.proj = nn.Linear(config.hidden_size, config.hidden_size)
+ self._attention_impl = None
+ self._initialize_attention(torch.get_default_dtype())
+
+ def _initialize_attention(self, dtype: torch.dtype) -> None:
+ backend = get_attn_backend(self.head_dim, dtype)
+ if backend.supports_packed_varlen():
+ self._attention_impl = backend.get_impl_cls()(
+ num_heads=self.num_heads,
+ head_size=self.head_dim,
+ num_kv_heads=self.num_heads,
+ softmax_scale=self.scaling,
+ causal=False,
+ prefix=self.prefix,
+ )
+ else:
+ logger.warning_once(
+ "Qwen2.5-VL vision attention uses torch SDPA because "
+ f"{backend.get_enum().name.lower()} does not support packed sequences"
+ )
+
+ def _packed_attention(
+ self,
+ query: torch.Tensor,
+ key: torch.Tensor,
+ value: torch.Tensor,
+ cu_seqlens: torch.Tensor,
+ cu_seqlens_host: tuple[int, ...],
+ max_seqlen: int,
+ ) -> torch.Tensor:
+ if self._attention_impl is not None:
+ return self._attention_impl.forward_varlen(
+ query,
+ key,
+ value,
+ cu_seqlens=cu_seqlens,
+ cu_seqlens_host=cu_seqlens_host,
+ max_seqlen=max_seqlen,
+ )
+
+ output = torch.empty_like(query)
+ for start, stop in zip(cu_seqlens_host[:-1], cu_seqlens_host[1:]):
+ if start == stop:
+ continue
+ query_segment = query[start:stop].transpose(0, 1).unsqueeze(0)
+ key_segment = key[start:stop].transpose(0, 1).unsqueeze(0)
+ value_segment = value[start:stop].transpose(0, 1).unsqueeze(0)
+ segment = F.scaled_dot_product_attention(
+ query_segment,
+ key_segment,
+ value_segment,
+ dropout_p=0.0,
+ is_causal=False,
+ scale=self.scaling,
+ )
+ output[start:stop] = segment.squeeze(0).transpose(0, 1)
+ return output
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ *,
+ cu_seqlens: torch.Tensor,
+ cu_seqlens_host: tuple[int, ...],
+ max_seqlen: int,
+ position_embeddings: tuple[torch.Tensor, torch.Tensor],
+ ) -> torch.Tensor:
+ seq_len = hidden_states.shape[0]
+ query, key, value = (
+ self.qkv(hidden_states)
+ .reshape(seq_len, 3, self.num_heads, self.head_dim)
+ .permute(1, 0, 2, 3)
+ .unbind(0)
+ )
+ query, key = _apply_vision_rotary_embedding(query, key, *position_embeddings)
+ output = self._packed_attention(
+ query,
+ key,
+ value,
+ cu_seqlens,
+ cu_seqlens_host,
+ max_seqlen,
+ )
+ return self.proj(output.reshape(seq_len, -1).contiguous())
+
+
+class Qwen2_5VLVisionMLP(nn.Module):
+ def __init__(self, config: Any) -> None:
+ super().__init__()
+ if config.hidden_act != "silu":
+ raise ValueError(
+ f"Unsupported Qwen2.5-VL vision activation: {config.hidden_act}"
+ )
+ self.gate_proj = nn.Linear(
+ config.hidden_size, config.intermediate_size, bias=True
+ )
+ self.up_proj = nn.Linear(
+ config.hidden_size, config.intermediate_size, bias=True
+ )
+ self.down_proj = nn.Linear(
+ config.intermediate_size, config.hidden_size, bias=True
+ )
+
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
+ return self.down_proj(
+ F.silu(self.gate_proj(hidden_states)) * self.up_proj(hidden_states)
+ )
+
+
+class Qwen2_5VLVisionBlock(nn.Module):
+ def __init__(self, config: Any, layer_idx: int) -> None:
+ super().__init__()
+ self.norm1 = Qwen2_5VLVisionRMSNorm(config.hidden_size)
+ self.norm2 = Qwen2_5VLVisionRMSNorm(config.hidden_size)
+ self.attn = Qwen2_5VLVisionAttention(
+ config, prefix=f"visual.blocks.{layer_idx}.attn"
+ )
+ self.mlp = Qwen2_5VLVisionMLP(config)
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ *,
+ cu_seqlens: torch.Tensor,
+ cu_seqlens_host: tuple[int, ...],
+ max_seqlen: int,
+ position_embeddings: tuple[torch.Tensor, torch.Tensor],
+ ) -> torch.Tensor:
+ hidden_states = hidden_states + self.attn(
+ self.norm1(hidden_states),
+ cu_seqlens=cu_seqlens,
+ cu_seqlens_host=cu_seqlens_host,
+ max_seqlen=max_seqlen,
+ position_embeddings=position_embeddings,
+ )
+ return hidden_states + self.mlp(self.norm2(hidden_states))
+
+
+class Qwen2_5VLVisionPatchMerger(nn.Module):
+ def __init__(
+ self,
+ output_dim: int,
+ context_dim: int,
+ spatial_merge_size: int,
+ ) -> None:
+ super().__init__()
+ self.hidden_size = context_dim * spatial_merge_size**2
+ self.ln_q = Qwen2_5VLVisionRMSNorm(context_dim)
+ self.mlp = nn.Sequential(
+ nn.Linear(self.hidden_size, self.hidden_size),
+ nn.GELU(),
+ nn.Linear(self.hidden_size, output_dim),
+ )
+
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
+ hidden_states = self.ln_q(hidden_states).view(-1, self.hidden_size)
+ return self.mlp(hidden_states)
+
+
+def _vision_position_ids(
+ grid_thw: torch.Tensor, spatial_merge_size: int
+) -> torch.Tensor:
+ position_ids = []
+ for t, h, w in grid_thw.tolist():
+ h_positions = torch.arange(h, device=grid_thw.device)[:, None].expand(h, w)
+ h_positions = (
+ h_positions.reshape(
+ h // spatial_merge_size,
+ spatial_merge_size,
+ w // spatial_merge_size,
+ spatial_merge_size,
+ )
+ .transpose(1, 2)
+ .flatten()
+ )
+ w_positions = torch.arange(w, device=grid_thw.device)[None, :].expand(h, w)
+ w_positions = (
+ w_positions.reshape(
+ h // spatial_merge_size,
+ spatial_merge_size,
+ w // spatial_merge_size,
+ spatial_merge_size,
+ )
+ .transpose(1, 2)
+ .flatten()
+ )
+ positions = torch.stack((h_positions, w_positions), dim=-1)
+ position_ids.append(positions.repeat(t, 1))
+ return torch.cat(position_ids, dim=0)
+
+
+def _vision_window_index(
+ grid_thw: torch.Tensor,
+ *,
+ spatial_merge_size: int,
+ window_size: int,
+ patch_size: int,
+) -> tuple[torch.Tensor, torch.Tensor]:
+ window_indices = []
+ cumulative_window_lengths = [0]
+ window_index_offset = 0
+ merger_window_size = window_size // spatial_merge_size // patch_size
+ spatial_merge_unit = spatial_merge_size**2
+
+ for grid_t, grid_h, grid_w in grid_thw.tolist():
+ merged_height = grid_h // spatial_merge_size
+ merged_width = grid_w // spatial_merge_size
+ index = torch.arange(
+ grid_t * merged_height * merged_width, device=grid_thw.device
+ ).reshape(grid_t, merged_height, merged_width)
+ pad_height = merger_window_size - merged_height % merger_window_size
+ pad_width = merger_window_size - merged_width % merger_window_size
+ num_windows_height = (merged_height + pad_height) // merger_window_size
+ num_windows_width = (merged_width + pad_width) // merger_window_size
+ index = F.pad(index, (0, pad_width, 0, pad_height), value=-100)
+ index = index.reshape(
+ grid_t,
+ num_windows_height,
+ merger_window_size,
+ num_windows_width,
+ merger_window_size,
+ )
+ index = index.permute(0, 1, 3, 2, 4).reshape(
+ grid_t,
+ num_windows_height * num_windows_width,
+ merger_window_size,
+ merger_window_size,
+ )
+ sequence_lengths = (index != -100).sum(dim=(2, 3)).reshape(-1)
+ index = index.flatten()
+ window_indices.append(index[index != -100] + window_index_offset)
+ cumulative = (
+ sequence_lengths.cumsum(0) * spatial_merge_unit
+ + cumulative_window_lengths[-1]
+ )
+ cumulative_window_lengths.extend(cumulative.tolist())
+ window_index_offset += grid_t * merged_height * merged_width
+
+ window_index = torch.cat(window_indices)
+ cu_window_seqlens = torch.tensor(
+ cumulative_window_lengths, device=grid_thw.device, dtype=torch.int32
+ )
+ return window_index, torch.unique_consecutive(cu_window_seqlens)
+
+
+class Qwen2_5VLVisionTransformer(nn.Module):
+ def __init__(self, config: Any) -> None:
+ super().__init__()
+ self.spatial_merge_size = config.spatial_merge_size
+ self.spatial_merge_unit = config.spatial_merge_size**2
+ self.patch_size = config.patch_size
+ self.window_size = config.window_size
+ self.full_attention_layers = frozenset(
+ int(layer_idx) for layer_idx in config.fullatt_block_indexes
+ )
+ self.patch_embed = Qwen2_5VLVisionPatchEmbed(
+ patch_size=config.patch_size,
+ temporal_patch_size=config.temporal_patch_size,
+ in_channels=config.in_channels,
+ embed_dim=config.hidden_size,
+ )
+ head_dim = config.hidden_size // config.num_heads
+ self.rotary_pos_emb = Qwen2_5VLVisionRotaryEmbedding(head_dim // 2)
+ self.blocks = nn.ModuleList(
+ Qwen2_5VLVisionBlock(config, layer_idx) for layer_idx in range(config.depth)
+ )
+ self.merger = Qwen2_5VLVisionPatchMerger(
+ output_dim=config.out_hidden_size,
+ context_dim=config.hidden_size,
+ spatial_merge_size=config.spatial_merge_size,
+ )
+
+ @property
+ def dtype(self) -> torch.dtype:
+ return self.patch_embed.proj.weight.dtype
+
+ @property
+ def device(self) -> torch.device:
+ return self.patch_embed.proj.weight.device
+
+ def forward(
+ self, hidden_states: torch.Tensor, grid_thw: torch.Tensor
+ ) -> torch.Tensor:
+ hidden_states = hidden_states.to(device=self.device, dtype=self.dtype)
+ grid_thw = grid_thw.to(self.device)
+ hidden_states = self.patch_embed(hidden_states)
+
+ position_ids = _vision_position_ids(grid_thw, self.spatial_merge_size)
+ window_index, cu_window_seqlens = _vision_window_index(
+ grid_thw,
+ spatial_merge_size=self.spatial_merge_size,
+ window_size=self.window_size,
+ patch_size=self.patch_size,
+ )
+
+ seq_len = hidden_states.shape[0]
+ hidden_states = hidden_states.reshape(
+ seq_len // self.spatial_merge_unit, self.spatial_merge_unit, -1
+ )[window_index]
+ hidden_states = hidden_states.reshape(seq_len, -1)
+
+ rotary = self.rotary_pos_emb(position_ids)
+ rotary = rotary.reshape(
+ seq_len // self.spatial_merge_unit, self.spatial_merge_unit, -1
+ )[window_index]
+ rotary = rotary.reshape(seq_len, -1)
+ rotary = torch.cat((rotary, rotary), dim=-1)
+ position_embeddings = (
+ rotary.cos(),
+ rotary.sin(),
+ )
+
+ cu_seqlens = torch.repeat_interleave(
+ grid_thw[:, 1] * grid_thw[:, 2], grid_thw[:, 0]
+ ).cumsum(dim=0, dtype=torch.int32)
+ cu_seqlens = F.pad(cu_seqlens, (1, 0), value=0)
+
+ full_metadata = _PackedSequenceMetadata.from_cu_seqlens(cu_seqlens)
+ window_metadata = _PackedSequenceMetadata.from_cu_seqlens(cu_window_seqlens)
+
+ for layer_idx, block in enumerate(self.blocks):
+ metadata = (
+ full_metadata
+ if layer_idx in self.full_attention_layers
+ else window_metadata
+ )
+ hidden_states = block(
+ hidden_states,
+ cu_seqlens=metadata.cu_seqlens,
+ cu_seqlens_host=metadata.cu_seqlens_host,
+ max_seqlen=metadata.max_seqlen,
+ position_embeddings=position_embeddings,
+ )
+
+ merged = self.merger(hidden_states)
+ return merged[torch.argsort(window_index)]
diff --git a/python/sglang/multimodal_gen/runtime/pipelines/longcat_image.py b/python/sglang/multimodal_gen/runtime/pipelines/longcat_image.py
index dd0a31067..87222a1e7 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines/longcat_image.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines/longcat_image.py
@@ -27,11 +27,8 @@ class LongCatImagePipeline(LoRAPipeline, ComposedPipelineBase):
pipeline_name = "LongCatImagePipeline"
- # The Qwen2.5-VL text encoder is loaded in-stage by LongCatPromptRewriteStage
- # (not via TextEncoderLoader), so "text_encoder" is intentionally absent;
- # the stage registers the loaded module via add_module("text_encoder", ...)
- # so the standard TextEncodingStage can fetch the same instance.
_required_config_modules = [
+ "text_encoder",
"tokenizer",
"text_processor",
"vae",
@@ -41,23 +38,17 @@ class LongCatImagePipeline(LoRAPipeline, ComposedPipelineBase):
def create_pipeline_stages(self, server_args: ServerArgs):
# 1. Prompt rewriting (optional) + request-level setup (generator, cfg renorm).
- # Loads the HF Qwen2.5-VL encoder and shares it with TextEncodingStage.
rewrite_stage = LongCatPromptRewriteStage(
- tokenizer=self.get_module("tokenizer"),
+ text_encoder=self.get_module("text_encoder"),
text_processor=self.get_module("text_processor"),
- model_path=self.model_path,
text_encoder_dtype=PRECISION_TO_TYPE[
server_args.pipeline_config.text_encoder_precisions[0]
],
)
self.add_stage(rewrite_stage)
- self.add_module("text_encoder", rewrite_stage.text_encoder)
# 2. Text encoding via the standard stage (tokenize_prompt +
- # postprocess_text_funcs hooks on the pipeline config). Shares the
- # encoder instance registered above; both stages declare a
- # "text_encoder" ComponentUse so the residency manager keeps it
- # resident across rewrite->encode and offloads after the last use.
+ # postprocess_text_funcs hooks on the pipeline config).
self.add_standard_text_encoding_stage()
# 3. Latent preparation (batch-size-aware via pipeline config hooks)
diff --git a/python/sglang/multimodal_gen/runtime/pipelines/qwen_image.py b/python/sglang/multimodal_gen/runtime/pipelines/qwen_image.py
index e92ebe540..9c45ea37d 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines/qwen_image.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines/qwen_image.py
@@ -95,6 +95,7 @@ class QwenImageLayeredPipeline(QwenImageEditPipeline):
pipeline_name = "QwenImageLayeredPipeline"
_required_config_modules = [
+ "text_encoder",
"vae",
"tokenizer",
"processor",
@@ -106,12 +107,11 @@ class QwenImageLayeredPipeline(QwenImageEditPipeline):
def create_before_denoising_stage():
return QwenImageLayeredBeforeDenoisingStage(
vae=self.get_module("vae"),
- text_encoder=None,
+ text_encoder=self.get_module("text_encoder"),
tokenizer=self.get_module("tokenizer"),
processor=self.get_module("processor"),
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
- model_path=self.model_path,
vae_dtype=PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision],
text_encoder_dtype=PRECISION_TO_TYPE[
server_args.pipeline_config.text_encoder_precisions[0]
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py
index 4ac1fe22b..380a26cae 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py
@@ -322,6 +322,7 @@ class ImageEncodingStage(PipelineStage):
pixel_values=image_inputs.pixel_values,
image_grid_thw=image_inputs.image_grid_thw,
output_hidden_states=True,
+ use_cache=False,
)
if batch.do_classifier_free_guidance:
neg_outputs = self.text_encoder(
@@ -330,6 +331,7 @@ class ImageEncodingStage(PipelineStage):
pixel_values=neg_image_inputs.pixel_values,
image_grid_thw=neg_image_inputs.image_grid_thw,
output_hidden_states=True,
+ use_cache=False,
)
prompt_embeds, prompt_embeds_mask, prompt_seq_lens = (
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/longcat_image.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/longcat_image.py
index 5ea41183b..41de73d46 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/longcat_image.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/longcat_image.py
@@ -1,7 +1,7 @@
"""Prompt-rewriting stage for LongCat-Image (T2I).
-`LongCatPromptRewriteStage` optionally rewrites the prompt via the HuggingFace
-`Qwen2_5_VLForConditionalGeneration.generate()` and sets the CPU generator for
+`LongCatPromptRewriteStage` optionally rewrites the prompt via the native
+Qwen2.5-VL encoder and sets the CPU generator for
seed reproducibility. Text encoding, latent preparation, RoPE and denoising are
all handled by the standard stages + `LongCatImagePipelineConfig` hooks, so this
is the only model-specific stage.
@@ -181,15 +181,13 @@ def _get_prompt_language(prompt):
class LongCatPromptRewriteStage(PipelineStage):
"""Optional prompt rewriting + request-level setup for LongCat-Image.
- Loads the Qwen2.5-VL text encoder (HuggingFace) in-stage and, when
- `enable_prompt_rewrite` is set, rewrites the prompt via `.generate()`
+ When `enable_prompt_rewrite` is set, rewrites the prompt via `.generate()`
(using the checkpoint's generation_config.json sampling params). Always
sets the CPU generator for seed reproducibility. CFG-renorm params are
read directly from sampling_params in `postprocess_cfg_noise`, not set here.
- The same encoder instance is shared with the standard `TextEncodingStage`
- (the pipeline registers it via `add_module("text_encoder", ...)`), so
- rewrite and encode run on one set of weights. Both stages declare a
+ The same encoder instance is shared with the standard `TextEncodingStage`,
+ so rewrite and encode run on one set of weights. Both stages declare a
`text_encoder` ComponentUse; the residency manager keeps the encoder
resident across the adjacent rewrite->encode uses and offloads it only
after the last use (when `--text-encoder-cpu-offload` is enabled).
@@ -197,35 +195,19 @@ class LongCatPromptRewriteStage(PipelineStage):
def __init__(
self,
- tokenizer,
+ text_encoder,
text_processor,
- model_path: str,
text_encoder_dtype: torch.dtype,
):
super().__init__()
self.text_encoder_dtype = text_encoder_dtype
- from transformers import Qwen2_5_VLForConditionalGeneration
-
- cpu_offload = self.server_args.text_encoder_cpu_offload
- init_device = torch.device("cpu") if cpu_offload else get_local_torch_device()
- self.text_encoder = (
- Qwen2_5_VLForConditionalGeneration.from_pretrained(
- model_path, subfolder="text_encoder"
- )
- .to(init_device)
- .to(dtype=self.text_encoder_dtype)
- )
- self.tokenizer = tokenizer
+ self.text_encoder = text_encoder
self.text_processor = text_processor
def component_uses(
self, server_args: ServerArgs, stage_name: str | None = None
) -> list[ComponentUse]:
stage_name = self._component_stage_name(stage_name)
- # "text_encoder" matches is_text_encoder_component_name, so
- # text_encoder_cpu_offload routes through VanillaD2HStrategy (the encoder
- # is a plain HF nn.Module, not FSDP-sharded). memory_intensive triggers
- # torch.cuda.empty_cache() after the encoder leaves CUDA.
return [
ComponentUse(
stage_name,
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/qwen_image_layered.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/qwen_image_layered.py
index c76e678cb..19f47a9ac 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/qwen_image_layered.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/qwen_image_layered.py
@@ -164,7 +164,6 @@ class QwenImageLayeredBeforeDenoisingStage(PipelineStage):
processor,
transformer,
scheduler,
- model_path,
vae_dtype: torch.dtype,
text_encoder_dtype: torch.dtype,
) -> None:
@@ -172,15 +171,7 @@ class QwenImageLayeredBeforeDenoisingStage(PipelineStage):
self.vae = vae.to(dtype=vae_dtype)
self.vae_dtype = vae_dtype
self.text_encoder_dtype = text_encoder_dtype
- if text_encoder is None:
- from transformers import Qwen2_5_VLForConditionalGeneration
-
- text_encoder = Qwen2_5_VLForConditionalGeneration.from_pretrained(
- model_path, subfolder="text_encoder"
- )
- self.text_encoder = text_encoder.to(
- device=get_local_torch_device(), dtype=self.text_encoder_dtype
- )
+ self.text_encoder = text_encoder
self.tokenizer = tokenizer
self.processor = processor
self.transformer = transformer
@@ -285,11 +276,12 @@ the image\n<|vision_start|><|image_pad|><|vision_end|><|im_end|>\n<|im_start|>as
padding=True,
return_tensors="pt",
).to(device)
- encoder_hidden_states = self.text_encoder(
- input_ids=txt_tokens.input_ids,
- attention_mask=txt_tokens.attention_mask,
- output_hidden_states=True,
- )
+ with set_forward_context(current_timestep=0, attn_metadata=None):
+ encoder_hidden_states = self.text_encoder(
+ input_ids=txt_tokens.input_ids,
+ attention_mask=txt_tokens.attention_mask,
+ output_hidden_states=True,
+ )
hidden_states = encoder_hidden_states.hidden_states[-1]
split_hidden_states = self._extract_masked_hidden(
hidden_states, txt_tokens.attention_mask
diff --git a/python/sglang/multimodal_gen/test/unit/test_disagg_roles.py b/python/sglang/multimodal_gen/test/unit/test_disagg_roles.py
index 2c0b6fc5c..ac0d9ec23 100644
--- a/python/sglang/multimodal_gen/test/unit/test_disagg_roles.py
+++ b/python/sglang/multimodal_gen/test/unit/test_disagg_roles.py
@@ -345,9 +345,6 @@ class TestPipelineSpecificExtraModules(unittest.TestCase):
extra_allowed_modules=extras,
)
self.assertEqual(extras, {"vae", "transformer"})
- self.assertNotIn(
- "text_encoder", QwenImageLayeredPipeline._required_config_modules
- )
self.assertEqual(
set(filtered),
{
@@ -356,6 +353,7 @@ class TestPipelineSpecificExtraModules(unittest.TestCase):
"processor",
"transformer",
"scheduler",
+ "text_encoder",
},
)
@@ -464,21 +462,16 @@ class TestQwenImageLayeredDtype(_GlobalStageArgsMixin, unittest.TestCase):
def to(self, *args, **kwargs):
return self
- with patch(
- "sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.qwen_image_layered.get_local_torch_device",
- return_value=torch.device("cpu"),
- ):
- stage = QwenImageLayeredBeforeDenoisingStage(
- vae=_DummyVAE(),
- text_encoder=torch.nn.Linear(1, 1),
- tokenizer=object(),
- processor=object(),
- transformer=object(),
- scheduler=object(),
- model_path="/unused",
- vae_dtype=torch.float32,
- text_encoder_dtype=torch.float16,
- )
+ stage = QwenImageLayeredBeforeDenoisingStage(
+ vae=_DummyVAE(),
+ text_encoder=torch.nn.Linear(1, 1),
+ tokenizer=object(),
+ processor=object(),
+ transformer=object(),
+ scheduler=object(),
+ vae_dtype=torch.float32,
+ text_encoder_dtype=torch.float16,
+ )
uses = stage.component_uses(SimpleNamespace(), "qwen_layered")
self.assertEqual(
diff --git a/python/sglang/multimodal_gen/test/unit/test_qwen2_5vl_generation.py b/python/sglang/multimodal_gen/test/unit/test_qwen2_5vl_generation.py
new file mode 100644
index 000000000..0a590cc81
--- /dev/null
+++ b/python/sglang/multimodal_gen/test/unit/test_qwen2_5vl_generation.py
@@ -0,0 +1,226 @@
+from types import SimpleNamespace
+
+import torch
+from torch import nn
+
+import sglang.multimodal_gen.runtime.models.encoders.qwen2_5vl as qwen2_5vl
+from sglang.multimodal_gen.configs.models.encoders.qwen_image import Qwen2_5VLConfig
+from sglang.multimodal_gen.configs.pipeline_configs.longcat_image import (
+ LongCatImagePipelineConfig,
+)
+from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
+ QwenImageLayeredPipelineConfig,
+)
+from sglang.multimodal_gen.runtime.models.encoders.qwen2_5vl import (
+ Qwen2_5_VLAttention,
+ Qwen2_5_VLForConditionalGeneration,
+ _apply_repetition_penalty,
+ _select_next_token,
+)
+from sglang.multimodal_gen.runtime.models.encoders.qwen2_5vl_vision import (
+ Qwen2_5VLVisionRotaryEmbedding,
+ Qwen2_5VLVisionTransformer,
+ _vision_position_ids,
+ _vision_window_index,
+)
+from sglang.multimodal_gen.runtime.pipelines.longcat_image import LongCatImagePipeline
+
+
+class _StubQwen2_5VL(Qwen2_5_VLForConditionalGeneration):
+ def __init__(self, next_tokens: list[list[int]], eos_token_id=5):
+ nn.Module.__init__(self)
+ self.model = nn.Module()
+ self.model.rope_deltas = torch.tensor([99])
+ self.config = SimpleNamespace(eos_token_id=eos_token_id, pad_token_id=0)
+ self.generation_config = {
+ "do_sample": True,
+ "temperature": 0.1,
+ "top_k": 1,
+ "top_p": 0.001,
+ "repetition_penalty": 1.05,
+ "eos_token_id": [eos_token_id, 6],
+ "pad_token_id": 0,
+ }
+ self.next_tokens = next_tokens
+ self.calls = []
+
+ def forward(self, input_ids, **kwargs):
+ call_index = len(self.calls)
+ self.calls.append((input_ids.clone(), kwargs))
+ logits = torch.full((input_ids.shape[0], 1, 8), -100.0)
+ for batch_index, token_id in enumerate(self.next_tokens[call_index]):
+ logits[batch_index, 0, token_id] = 100.0
+ return SimpleNamespace(
+ logits=logits,
+ past_key_values=f"cache-{call_index}",
+ )
+
+
+class _AttentionRecorder(nn.Module):
+ def __init__(self):
+ super().__init__()
+ self.masks = []
+
+ def forward(self, query, key, value, attn_mask=None):
+ self.masks.append(attn_mask)
+ return query
+
+
+def test_explicit_attention_mask_is_limited_to_cached_generation(monkeypatch):
+ attention = Qwen2_5_VLAttention.__new__(Qwen2_5_VLAttention)
+ nn.Module.__init__(attention)
+ attention.q_proj = nn.Identity()
+ attention.k_proj = nn.Identity()
+ attention.v_proj = nn.Identity()
+ attention.o_proj = nn.Identity()
+ attention.num_heads = 1
+ attention.num_key_value_heads = 1
+ attention.head_dim = 4
+ attention.rope_scaling = {"mrope_section": [1, 1, 0]}
+ attention.attn = _AttentionRecorder()
+ monkeypatch.setattr(
+ qwen2_5vl,
+ "apply_multimodal_rotary_pos_emb",
+ lambda query, key, *_args: (query, key),
+ )
+
+ hidden_states = torch.randn(1, 2, 4)
+ explicit_mask = torch.zeros(1, 1, 2, 2)
+ kwargs = {
+ "hidden_states": hidden_states,
+ "attention_mask": explicit_mask,
+ "position_embeddings": (torch.empty(0), torch.empty(0)),
+ }
+
+ attention(**kwargs, use_cache=False)
+ attention(**kwargs, use_cache=True)
+
+ assert attention.attn.masks[0] is None
+ assert attention.attn.masks[1] is explicit_mask
+
+
+def test_repetition_penalty_matches_sign_dependent_scaling():
+ logits = torch.tensor([[2.0, -3.0, 4.0]])
+ penalized = _apply_repetition_penalty(logits, torch.tensor([[0, 1]]), penalty=2.0)
+ torch.testing.assert_close(penalized, torch.tensor([[1.0, -6.0, 4.0]]))
+
+
+def test_top_k_one_sampling_is_deterministic():
+ token = _select_next_token(
+ torch.tensor([[1.0, 3.0, 2.0]]),
+ torch.tensor([[0]]),
+ do_sample=True,
+ temperature=0.1,
+ top_k=1,
+ top_p=0.001,
+ repetition_penalty=1.0,
+ )
+ assert token.tolist() == [1]
+
+
+def test_generate_reuses_cache_and_only_prefills_vision_once():
+ model = _StubQwen2_5VL([[4], [5]])
+ pixel_values = torch.ones(1, 3)
+ generated = model.generate(
+ torch.tensor([[1, 2]]),
+ attention_mask=torch.ones(1, 2, dtype=torch.long),
+ pixel_values=pixel_values,
+ image_grid_thw=torch.tensor([[1, 1, 1]]),
+ mm_token_type_ids=torch.tensor([[0, 1]]),
+ max_new_tokens=3,
+ )
+
+ assert generated.tolist() == [[1, 2, 4, 5]]
+ assert len(model.calls) == 2
+ assert model.calls[0][1]["pixel_values"] is pixel_values
+ assert model.calls[1][1]["pixel_values"] is None
+ assert model.calls[0][1]["past_key_values"] is None
+ assert model.calls[1][1]["past_key_values"] == "cache-0"
+ assert model.calls[0][1]["cache_position"].tolist() == [0, 1]
+ assert model.calls[1][1]["cache_position"].tolist() == [2]
+ assert model.model.rope_deltas is None
+
+
+def test_generate_pads_finished_rows_until_the_batch_stops():
+ model = _StubQwen2_5VL([[5, 4], [7, 6]])
+ generated = model.generate(
+ torch.tensor([[1], [2]]),
+ max_new_tokens=3,
+ )
+
+ assert generated.tolist() == [[1, 5, 0], [2, 4, 6]]
+
+
+def test_native_vision_indices_preserve_merged_token_groups():
+ grid_thw = torch.tensor([[1, 4, 4], [2, 2, 4]])
+ position_ids = _vision_position_ids(grid_thw, spatial_merge_size=2)
+ window_index, cu_window_seqlens = _vision_window_index(
+ grid_thw,
+ spatial_merge_size=2,
+ window_size=8,
+ patch_size=2,
+ )
+
+ assert position_ids.shape == (32, 2)
+ assert sorted(window_index.tolist()) == list(range(8))
+ assert cu_window_seqlens[0].item() == 0
+ assert cu_window_seqlens[-1].item() == 32
+
+
+def test_native_vision_keeps_rotary_trigonometry_in_fp32():
+ class PatchEmbed(nn.Module):
+ def __init__(self):
+ super().__init__()
+ self.proj = nn.Linear(1, 8, bias=False, dtype=torch.bfloat16)
+
+ def forward(self, hidden_states):
+ return self.proj(hidden_states)
+
+ class BlockRecorder(nn.Module):
+ def __init__(self):
+ super().__init__()
+ self.position_embedding_dtypes = None
+
+ def forward(self, hidden_states, *, position_embeddings, **_kwargs):
+ self.position_embedding_dtypes = tuple(
+ embedding.dtype for embedding in position_embeddings
+ )
+ return hidden_states
+
+ class Merger(nn.Module):
+ def forward(self, hidden_states):
+ return hidden_states.reshape(-1, 4, hidden_states.shape[-1])[:, 0]
+
+ model = Qwen2_5VLVisionTransformer.__new__(Qwen2_5VLVisionTransformer)
+ nn.Module.__init__(model)
+ model.spatial_merge_size = 2
+ model.spatial_merge_unit = 4
+ model.patch_size = 2
+ model.window_size = 8
+ model.full_attention_layers = frozenset({0})
+ model.patch_embed = PatchEmbed()
+ model.rotary_pos_emb = Qwen2_5VLVisionRotaryEmbedding(2)
+ block = BlockRecorder()
+ model.blocks = nn.ModuleList([block])
+ model.merger = Merger()
+ output = model(
+ torch.zeros(16, 1, dtype=torch.bfloat16),
+ grid_thw=torch.tensor([[1, 4, 4]]),
+ )
+
+ assert output.dtype == torch.bfloat16
+ assert block.position_embedding_dtypes == (torch.float32, torch.float32)
+
+
+def test_qwen_generation_pipelines_load_the_native_component():
+ longcat_config = LongCatImagePipelineConfig()
+
+ assert isinstance(longcat_config.text_encoder_configs[0], Qwen2_5VLConfig)
+ assert "text_encoder" in LongCatImagePipeline._required_config_modules
+ assert Qwen2_5_VLForConditionalGeneration._fsdp_forward_methods == ("generate",)
+ assert "model.visual.blocks" in Qwen2_5_VLForConditionalGeneration.layer_names
+
+ for pipeline_config in (longcat_config, QwenImageLayeredPipelineConfig()):
+ deployment = pipeline_config.get_model_deployment_config()
+ assert deployment.keep_resident_min_available_gb == 70
+ assert deployment.keep_resident_components == ("text_encoder", "vae")
diff --git a/python/sglang/multimodal_gen/test/unit/test_server_args.py b/python/sglang/multimodal_gen/test/unit/test_server_args.py
index ac67027e1..ab9119ad1 100644
--- a/python/sglang/multimodal_gen/test/unit/test_server_args.py
+++ b/python/sglang/multimodal_gen/test/unit/test_server_args.py
@@ -20,6 +20,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.hunyuan import FastHunyuanCo
from sglang.multimodal_gen.configs.pipeline_configs.lingbot_world import (
LingBotWorldCausalDMDConfig,
)
+from sglang.multimodal_gen.configs.pipeline_configs.longcat_image import (
+ LongCatImagePipelineConfig,
+)
from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import (
LTX2PipelineConfig,
LTX23PipelineConfig,
@@ -32,6 +35,7 @@ from sglang.multimodal_gen.configs.pipeline_configs.model_deployment_config impo
)
from sglang.multimodal_gen.configs.pipeline_configs.mova import MOVAPipelineConfig
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
+ QwenImageLayeredPipelineConfig,
QwenImagePipelineConfig,
)
from sglang.multimodal_gen.configs.pipeline_configs.sana_wm import (
@@ -1146,6 +1150,32 @@ class TestOffloadDefaults(unittest.TestCase):
self.assertEqual(qwen_deployment.keep_resident_components, ("vae",))
self.assertIsNone(qwen_deployment.keep_resident_min_available_gb)
+ def test_qwen_ar_generation_residency_scales_with_available_memory(self):
+ pipeline_configs = (
+ QwenImageLayeredPipelineConfig(),
+ LongCatImagePipelineConfig(),
+ )
+
+ for pipeline_config in pipeline_configs:
+ high_memory_args = self._from_dict_with_pipeline_config(
+ pipeline_config,
+ memory_gb=80,
+ kwargs={"performance_mode": "auto"},
+ )
+ self.assertNotIn(
+ "text_encoder", high_memory_args.layerwise_offload_components or []
+ )
+ self.assertFalse(high_memory_args.text_encoder_cpu_offload)
+
+ constrained_args = self._from_dict_with_pipeline_config(
+ pipeline_config,
+ memory_gb=60,
+ kwargs={"performance_mode": "auto"},
+ )
+ self.assertIn(
+ "text_encoder", constrained_args.layerwise_offload_components or []
+ )
+
def test_auto_multi_gpu_sana_wm_prefers_fsdp_and_cfg_parallel(self):
args = self._from_dict_with_pipeline_config(
SanaWMPipelineConfig(),