[diffusion] chore: use native qwen2.5-vl generation (#34896)

This commit is contained in:
Mick
2026-08-16 09:57:39 +08:00
committed by GitHub
parent 4a6dc267e1
commit eb6b773149
16 changed files with 1010 additions and 116 deletions
@@ -37,6 +37,10 @@ Rows are grouped when a family shares the same runtime path or optimization supp
<td>Qwen-Image</td> <td>Qwen-Image</td>
<td><div className="sgd-id-list"><code>Qwen/Qwen-Image</code><code>Qwen/Qwen-Image-2512</code><code>Qwen/Qwen-Image-Edit</code><code>Qwen/Qwen-Image-Edit-2509</code><code>Qwen/Qwen-Image-Edit-2511</code><code>Qwen/Qwen-Image-Layered</code></div></td> <td><div className="sgd-id-list"><code>Qwen/Qwen-Image</code><code>Qwen/Qwen-Image-2512</code><code>Qwen/Qwen-Image-Edit</code><code>Qwen/Qwen-Image-Edit-2509</code><code>Qwen/Qwen-Image-Edit-2511</code><code>Qwen/Qwen-Image-Layered</code></div></td>
</tr> </tr>
<tr>
<td>LongCat-Image</td>
<td><div className="sgd-id-list"><code>meituan-longcat/LongCat-Image</code></div></td>
</tr>
<tr> <tr>
<td>SD3 / SD3.5</td> <td>SD3 / SD3.5</td>
<td><div className="sgd-id-list"><code>stabilityai/stable-diffusion-3-medium</code><code>stabilityai/stable-diffusion-3-medium-diffusers</code><code>stabilityai/stable-diffusion-3.5-medium</code><code>stabilityai/stable-diffusion-3.5-medium-diffusers</code><code>stabilityai/stable-diffusion-3.5-large</code><code>stabilityai/stable-diffusion-3.5-large-diffusers</code></div></td> <td><div className="sgd-id-list"><code>stabilityai/stable-diffusion-3-medium</code><code>stabilityai/stable-diffusion-3-medium-diffusers</code><code>stabilityai/stable-diffusion-3.5-medium</code><code>stabilityai/stable-diffusion-3.5-medium-diffusers</code><code>stabilityai/stable-diffusion-3.5-large</code><code>stabilityai/stable-diffusion-3.5-large-diffusers</code></div></td>
+20 -5
View File
@@ -1,11 +1,16 @@
--- ---
title: "Diffusion models with AR stage like GLM-Image" title: "Diffusion models with autoregressive stages"
description: "Run diffusion pipelines that delegate an AR stage to a separate SGLang server, such as GLM-Image." 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 ```bash
# Terminal 1 : launch server # Terminal 1 : launch server
sglang serve --model-path zai-org/GLM-Image --port ${PORT} sglang serve --model-path zai-org/GLM-Image --port ${PORT}
@@ -20,7 +25,7 @@ curl http://${HOST}:${PORT}/v1/images/generations \
"size": "widthxheight" "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 ```bash
# Terminal 1 : launch server with AR model # Terminal 1 : launch server with AR model
sglang serve --model-path /path/to/zai-org/GLM-Image/vision_language_encoder/ \ sglang serve --model-path /path/to/zai-org/GLM-Image/vision_language_encoder/ \
@@ -64,6 +69,16 @@ curl http://${HOST}:${PORT}/v1/images/generations \
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)", whiteSpace: "nowrap"}}>T2I, I2I, V2I</td> <td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)", whiteSpace: "nowrap"}}>T2I, I2I, V2I</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>T2I</td> <td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>T2I</td>
</tr> </tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>Qwen Image Layered</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Not used</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>I2I, in-process native Qwen2.5-VL</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>LongCat-Image</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Not used</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>T2I, in-process native Qwen2.5-VL</td>
</tr>
</tbody> </tbody>
</table> </table>
@@ -86,6 +86,7 @@ class EncoderConfig(ModelConfig):
@dataclass @dataclass
class TextEncoderConfig(EncoderConfig): class TextEncoderConfig(EncoderConfig):
arch_config: ArchConfig = field(default_factory=TextEncoderArchConfig) arch_config: ArchConfig = field(default_factory=TextEncoderArchConfig)
generation_config: dict[str, Any] = field(default_factory=dict)
@dataclass @dataclass
@@ -8,6 +8,7 @@ from sglang.multimodal_gen.configs.models import DiTConfig, EncoderConfig, VAECo
from sglang.multimodal_gen.configs.models.dits.longcat_image import ( from sglang.multimodal_gen.configs.models.dits.longcat_image import (
LongCatImageDitConfig, LongCatImageDitConfig,
) )
from sglang.multimodal_gen.configs.models.encoders.qwen_image import Qwen2_5VLConfig
from sglang.multimodal_gen.configs.models.vaes.longcat_image import ( from sglang.multimodal_gen.configs.models.vaes.longcat_image import (
LongCatImageVAEConfig, LongCatImageVAEConfig,
) )
@@ -16,6 +17,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.base import (
ModelTaskType, ModelTaskType,
TextConditioningOutput, TextConditioningOutput,
) )
from sglang.multimodal_gen.configs.pipeline_configs.model_deployment_config import (
ModelDeploymentConfig,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__) 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 @dataclass
class LongCatImagePipelineConfig(ImagePipelineConfig): class LongCatImagePipelineConfig(ImagePipelineConfig):
"""Configuration for the LongCat-Image T2I pipeline.""" """Configuration for the LongCat-Image T2I pipeline."""
@@ -246,16 +238,20 @@ class LongCatImagePipelineConfig(ImagePipelineConfig):
dit_config: DiTConfig = field(default_factory=LongCatImageDitConfig) dit_config: DiTConfig = field(default_factory=LongCatImageDitConfig)
vae_config: VAEConfig = field(default_factory=LongCatImageVAEConfig) 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_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16",))
text_encoder_configs: tuple[EncoderConfig, ...] = field( text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (LongCatImageEncoderConfig(),) default_factory=lambda: (Qwen2_5VLConfig(),)
) )
postprocess_text_funcs: tuple[Callable, ...] = field( postprocess_text_funcs: tuple[Callable, ...] = field(
default_factory=lambda: (longcat_postprocess_text,) 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 --- # --- LatentPreparationStage hooks ---
def prepare_latent_shape(self, batch, batch_size, num_frames): def prepare_latent_shape(self, batch, batch_size, num_frames):
@@ -19,6 +19,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.base import (
pad_text_embeddings_with_mask, pad_text_embeddings_with_mask,
shard_rotary_emb_for_sp, 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 ( from sglang.multimodal_gen.configs.post_training.pipeline_configs import (
QwenImageRolloutPipelineMixin, QwenImageRolloutPipelineMixin,
) )
@@ -753,6 +756,12 @@ class QwenImageLayeredPipelineConfig(QwenImageEditPipelineConfig):
resolution: int = 640 resolution: int = 640
vae_precision: str = "bf16" 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( def postprocess_cfg_noise(
self, self,
batch, batch,
@@ -52,6 +52,7 @@ from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import ( from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
get_config, get_config,
get_diffusers_component_config, get_diffusers_component_config,
load_dict,
) )
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.runtime.utils.precision import precision_to_dtype 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 = server_args.pipeline_config.text_encoder_configs[encoder_index]
encoder_config.update_model_arch(model_config) 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: if encoder_index == 0:
for key, value in diffusers_pretrained_config.__dict__.items(): for key, value in diffusers_pretrained_config.__dict__.items():
@@ -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.layers.quantization import QuantizationConfig
from sglang.multimodal_gen.runtime.loader.weight_utils import default_weight_loader 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.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.platforms import AttentionBackendEnum
from sglang.multimodal_gen.runtime.utils.common import add_prefix from sglang.multimodal_gen.runtime.utils.common import add_prefix
@@ -69,8 +72,6 @@ import torch
import torch.nn as nn import torch.nn as nn
from transformers.activations import ACT2FN from transformers.activations import ACT2FN
from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import ( from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import (
Qwen2_5_VisionRotaryEmbedding,
Qwen2_5_VisionTransformerPretrainedModel,
Qwen2_5_VLCausalLMOutputWithPast, Qwen2_5_VLCausalLMOutputWithPast,
Qwen2_5_VLModelOutputWithPast, Qwen2_5_VLModelOutputWithPast,
Qwen2_5_VLRotaryEmbedding, Qwen2_5_VLRotaryEmbedding,
@@ -80,6 +81,53 @@ from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import (
logger = logging.getLogger(__name__) 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: def _tp_world_size() -> int:
if not model_parallel_is_initialized(): if not model_parallel_is_initialized():
return 1 return 1
@@ -261,7 +309,15 @@ class Qwen2_5_VLAttention(nn.Module):
query_states = query_states.transpose(1, 2) query_states = query_states.transpose(1, 2)
key_states = key_states.transpose(1, 2) key_states = key_states.transpose(1, 2)
value_states = value_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 = attn_output.reshape(bsz, q_len, -1).contiguous()
attn_output = _linear_output(self.o_proj, attn_output) 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"} _checkpoint_conversion_mapping = {"^model": "language_model"}
# Reference: fix gemma3 grad acc #37208 # Reference: fix gemma3 grad acc #37208
accepts_loss_kwargs = False 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): def __init__(self, config, enable_image_understanding: bool = False):
super().__init__() super().__init__()
self.language_model = Qwen2_5_VLTextModel(config.text_config) self.language_model = Qwen2_5_VLTextModel(config.text_config)
if enable_image_understanding: if enable_image_understanding:
self.visual = Qwen2_5_VisionTransformerPretrainedModel._from_config( self.visual = Qwen2_5VLVisionTransformer(config.vision_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.rope_deltas = None # cache rope_deltas here self.rope_deltas = None # cache rope_deltas here
self.config = config self.config = config
# Initialize weights and apply final processing # Initialize weights and apply final processing
@@ -902,11 +944,6 @@ class Qwen2_5_VLModel(nn.Module):
""" """
pixel_values = pixel_values.type(self.visual.dtype) pixel_values = pixel_values.type(self.visual.dtype)
image_embeds = self.visual(pixel_values, grid_thw=image_grid_thw) 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 = ( split_sizes = (
image_grid_thw.prod(-1) // self.visual.spatial_merge_size**2 image_grid_thw.prod(-1) // self.visual.spatial_merge_size**2
).tolist() ).tolist()
@@ -1106,6 +1143,9 @@ class Qwen2_5_VLModel(nn.Module):
class Qwen2_5_VLForConditionalGeneration(TextEncoder): class Qwen2_5_VLForConditionalGeneration(TextEncoder):
layer_names = [*TextEncoder.layer_names, "model.visual.blocks"]
_fsdp_forward_methods = ("generate",)
# BitandBytes specific attributes # BitandBytes specific attributes
default_bitsandbytes_target_modules = [ default_bitsandbytes_target_modules = [
".gate_up_proj.", ".gate_up_proj.",
@@ -1132,6 +1172,7 @@ class Qwen2_5_VLForConditionalGeneration(TextEncoder):
) -> None: ) -> None:
super().__init__(config) super().__init__(config)
enable_image_understanding = config.enable_image_understanding enable_image_understanding = config.enable_image_understanding
generation_config = config.generation_config
config = config.arch_config config = config.arch_config
self.model = Qwen2_5_VLModel( self.model = Qwen2_5_VLModel(
config, enable_image_understanding=enable_image_understanding config, enable_image_understanding=enable_image_understanding
@@ -1141,6 +1182,7 @@ class Qwen2_5_VLForConditionalGeneration(TextEncoder):
) )
self.enable_image_understanding = enable_image_understanding self.enable_image_understanding = enable_image_understanding
self.generation_config = generation_config
self.config = config self.config = config
@@ -1225,6 +1267,152 @@ class Qwen2_5_VLForConditionalGeneration(TextEncoder):
rope_deltas=outputs.rope_deltas, 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]]): def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
loaded_params: set[str] = set() loaded_params: set[str] = set()
@@ -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)]
@@ -27,11 +27,8 @@ class LongCatImagePipeline(LoRAPipeline, ComposedPipelineBase):
pipeline_name = "LongCatImagePipeline" 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 = [ _required_config_modules = [
"text_encoder",
"tokenizer", "tokenizer",
"text_processor", "text_processor",
"vae", "vae",
@@ -41,23 +38,17 @@ class LongCatImagePipeline(LoRAPipeline, ComposedPipelineBase):
def create_pipeline_stages(self, server_args: ServerArgs): def create_pipeline_stages(self, server_args: ServerArgs):
# 1. Prompt rewriting (optional) + request-level setup (generator, cfg renorm). # 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( rewrite_stage = LongCatPromptRewriteStage(
tokenizer=self.get_module("tokenizer"), text_encoder=self.get_module("text_encoder"),
text_processor=self.get_module("text_processor"), text_processor=self.get_module("text_processor"),
model_path=self.model_path,
text_encoder_dtype=PRECISION_TO_TYPE[ text_encoder_dtype=PRECISION_TO_TYPE[
server_args.pipeline_config.text_encoder_precisions[0] server_args.pipeline_config.text_encoder_precisions[0]
], ],
) )
self.add_stage(rewrite_stage) self.add_stage(rewrite_stage)
self.add_module("text_encoder", rewrite_stage.text_encoder)
# 2. Text encoding via the standard stage (tokenize_prompt + # 2. Text encoding via the standard stage (tokenize_prompt +
# postprocess_text_funcs hooks on the pipeline config). Shares the # postprocess_text_funcs hooks on the pipeline config).
# 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.
self.add_standard_text_encoding_stage() self.add_standard_text_encoding_stage()
# 3. Latent preparation (batch-size-aware via pipeline config hooks) # 3. Latent preparation (batch-size-aware via pipeline config hooks)
@@ -95,6 +95,7 @@ class QwenImageLayeredPipeline(QwenImageEditPipeline):
pipeline_name = "QwenImageLayeredPipeline" pipeline_name = "QwenImageLayeredPipeline"
_required_config_modules = [ _required_config_modules = [
"text_encoder",
"vae", "vae",
"tokenizer", "tokenizer",
"processor", "processor",
@@ -106,12 +107,11 @@ class QwenImageLayeredPipeline(QwenImageEditPipeline):
def create_before_denoising_stage(): def create_before_denoising_stage():
return QwenImageLayeredBeforeDenoisingStage( return QwenImageLayeredBeforeDenoisingStage(
vae=self.get_module("vae"), vae=self.get_module("vae"),
text_encoder=None, text_encoder=self.get_module("text_encoder"),
tokenizer=self.get_module("tokenizer"), tokenizer=self.get_module("tokenizer"),
processor=self.get_module("processor"), processor=self.get_module("processor"),
transformer=self.get_module("transformer"), transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"), scheduler=self.get_module("scheduler"),
model_path=self.model_path,
vae_dtype=PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision], vae_dtype=PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision],
text_encoder_dtype=PRECISION_TO_TYPE[ text_encoder_dtype=PRECISION_TO_TYPE[
server_args.pipeline_config.text_encoder_precisions[0] server_args.pipeline_config.text_encoder_precisions[0]
@@ -322,6 +322,7 @@ class ImageEncodingStage(PipelineStage):
pixel_values=image_inputs.pixel_values, pixel_values=image_inputs.pixel_values,
image_grid_thw=image_inputs.image_grid_thw, image_grid_thw=image_inputs.image_grid_thw,
output_hidden_states=True, output_hidden_states=True,
use_cache=False,
) )
if batch.do_classifier_free_guidance: if batch.do_classifier_free_guidance:
neg_outputs = self.text_encoder( neg_outputs = self.text_encoder(
@@ -330,6 +331,7 @@ class ImageEncodingStage(PipelineStage):
pixel_values=neg_image_inputs.pixel_values, pixel_values=neg_image_inputs.pixel_values,
image_grid_thw=neg_image_inputs.image_grid_thw, image_grid_thw=neg_image_inputs.image_grid_thw,
output_hidden_states=True, output_hidden_states=True,
use_cache=False,
) )
prompt_embeds, prompt_embeds_mask, prompt_seq_lens = ( prompt_embeds, prompt_embeds_mask, prompt_seq_lens = (
@@ -1,7 +1,7 @@
"""Prompt-rewriting stage for LongCat-Image (T2I). """Prompt-rewriting stage for LongCat-Image (T2I).
`LongCatPromptRewriteStage` optionally rewrites the prompt via the HuggingFace `LongCatPromptRewriteStage` optionally rewrites the prompt via the native
`Qwen2_5_VLForConditionalGeneration.generate()` and sets the CPU generator for Qwen2.5-VL encoder and sets the CPU generator for
seed reproducibility. Text encoding, latent preparation, RoPE and denoising are seed reproducibility. Text encoding, latent preparation, RoPE and denoising are
all handled by the standard stages + `LongCatImagePipelineConfig` hooks, so this all handled by the standard stages + `LongCatImagePipelineConfig` hooks, so this
is the only model-specific stage. is the only model-specific stage.
@@ -181,15 +181,13 @@ def _get_prompt_language(prompt):
class LongCatPromptRewriteStage(PipelineStage): class LongCatPromptRewriteStage(PipelineStage):
"""Optional prompt rewriting + request-level setup for LongCat-Image. """Optional prompt rewriting + request-level setup for LongCat-Image.
Loads the Qwen2.5-VL text encoder (HuggingFace) in-stage and, when When `enable_prompt_rewrite` is set, rewrites the prompt via `.generate()`
`enable_prompt_rewrite` is set, rewrites the prompt via `.generate()`
(using the checkpoint's generation_config.json sampling params). Always (using the checkpoint's generation_config.json sampling params). Always
sets the CPU generator for seed reproducibility. CFG-renorm params are sets the CPU generator for seed reproducibility. CFG-renorm params are
read directly from sampling_params in `postprocess_cfg_noise`, not set here. read directly from sampling_params in `postprocess_cfg_noise`, not set here.
The same encoder instance is shared with the standard `TextEncodingStage` The same encoder instance is shared with the standard `TextEncodingStage`,
(the pipeline registers it via `add_module("text_encoder", ...)`), so so rewrite and encode run on one set of weights. Both stages declare a
rewrite and encode run on one set of weights. Both stages declare a
`text_encoder` ComponentUse; the residency manager keeps the encoder `text_encoder` ComponentUse; the residency manager keeps the encoder
resident across the adjacent rewrite->encode uses and offloads it only resident across the adjacent rewrite->encode uses and offloads it only
after the last use (when `--text-encoder-cpu-offload` is enabled). after the last use (when `--text-encoder-cpu-offload` is enabled).
@@ -197,35 +195,19 @@ class LongCatPromptRewriteStage(PipelineStage):
def __init__( def __init__(
self, self,
tokenizer, text_encoder,
text_processor, text_processor,
model_path: str,
text_encoder_dtype: torch.dtype, text_encoder_dtype: torch.dtype,
): ):
super().__init__() super().__init__()
self.text_encoder_dtype = text_encoder_dtype self.text_encoder_dtype = text_encoder_dtype
from transformers import Qwen2_5_VLForConditionalGeneration self.text_encoder = text_encoder
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_processor = text_processor self.text_processor = text_processor
def component_uses( def component_uses(
self, server_args: ServerArgs, stage_name: str | None = None self, server_args: ServerArgs, stage_name: str | None = None
) -> list[ComponentUse]: ) -> list[ComponentUse]:
stage_name = self._component_stage_name(stage_name) 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 [ return [
ComponentUse( ComponentUse(
stage_name, stage_name,
@@ -164,7 +164,6 @@ class QwenImageLayeredBeforeDenoisingStage(PipelineStage):
processor, processor,
transformer, transformer,
scheduler, scheduler,
model_path,
vae_dtype: torch.dtype, vae_dtype: torch.dtype,
text_encoder_dtype: torch.dtype, text_encoder_dtype: torch.dtype,
) -> None: ) -> None:
@@ -172,15 +171,7 @@ class QwenImageLayeredBeforeDenoisingStage(PipelineStage):
self.vae = vae.to(dtype=vae_dtype) self.vae = vae.to(dtype=vae_dtype)
self.vae_dtype = vae_dtype self.vae_dtype = vae_dtype
self.text_encoder_dtype = text_encoder_dtype self.text_encoder_dtype = text_encoder_dtype
if text_encoder is None: self.text_encoder = text_encoder
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.tokenizer = tokenizer self.tokenizer = tokenizer
self.processor = processor self.processor = processor
self.transformer = transformer self.transformer = transformer
@@ -285,11 +276,12 @@ the image\n<|vision_start|><|image_pad|><|vision_end|><|im_end|>\n<|im_start|>as
padding=True, padding=True,
return_tensors="pt", return_tensors="pt",
).to(device) ).to(device)
encoder_hidden_states = self.text_encoder( with set_forward_context(current_timestep=0, attn_metadata=None):
input_ids=txt_tokens.input_ids, encoder_hidden_states = self.text_encoder(
attention_mask=txt_tokens.attention_mask, input_ids=txt_tokens.input_ids,
output_hidden_states=True, attention_mask=txt_tokens.attention_mask,
) output_hidden_states=True,
)
hidden_states = encoder_hidden_states.hidden_states[-1] hidden_states = encoder_hidden_states.hidden_states[-1]
split_hidden_states = self._extract_masked_hidden( split_hidden_states = self._extract_masked_hidden(
hidden_states, txt_tokens.attention_mask hidden_states, txt_tokens.attention_mask
@@ -345,9 +345,6 @@ class TestPipelineSpecificExtraModules(unittest.TestCase):
extra_allowed_modules=extras, extra_allowed_modules=extras,
) )
self.assertEqual(extras, {"vae", "transformer"}) self.assertEqual(extras, {"vae", "transformer"})
self.assertNotIn(
"text_encoder", QwenImageLayeredPipeline._required_config_modules
)
self.assertEqual( self.assertEqual(
set(filtered), set(filtered),
{ {
@@ -356,6 +353,7 @@ class TestPipelineSpecificExtraModules(unittest.TestCase):
"processor", "processor",
"transformer", "transformer",
"scheduler", "scheduler",
"text_encoder",
}, },
) )
@@ -464,21 +462,16 @@ class TestQwenImageLayeredDtype(_GlobalStageArgsMixin, unittest.TestCase):
def to(self, *args, **kwargs): def to(self, *args, **kwargs):
return self return self
with patch( stage = QwenImageLayeredBeforeDenoisingStage(
"sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.qwen_image_layered.get_local_torch_device", vae=_DummyVAE(),
return_value=torch.device("cpu"), text_encoder=torch.nn.Linear(1, 1),
): tokenizer=object(),
stage = QwenImageLayeredBeforeDenoisingStage( processor=object(),
vae=_DummyVAE(), transformer=object(),
text_encoder=torch.nn.Linear(1, 1), scheduler=object(),
tokenizer=object(), vae_dtype=torch.float32,
processor=object(), text_encoder_dtype=torch.float16,
transformer=object(), )
scheduler=object(),
model_path="/unused",
vae_dtype=torch.float32,
text_encoder_dtype=torch.float16,
)
uses = stage.component_uses(SimpleNamespace(), "qwen_layered") uses = stage.component_uses(SimpleNamespace(), "qwen_layered")
self.assertEqual( self.assertEqual(
@@ -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")
@@ -20,6 +20,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.hunyuan import FastHunyuanCo
from sglang.multimodal_gen.configs.pipeline_configs.lingbot_world import ( from sglang.multimodal_gen.configs.pipeline_configs.lingbot_world import (
LingBotWorldCausalDMDConfig, LingBotWorldCausalDMDConfig,
) )
from sglang.multimodal_gen.configs.pipeline_configs.longcat_image import (
LongCatImagePipelineConfig,
)
from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import ( from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import (
LTX2PipelineConfig, LTX2PipelineConfig,
LTX23PipelineConfig, 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.mova import MOVAPipelineConfig
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import ( from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
QwenImageLayeredPipelineConfig,
QwenImagePipelineConfig, QwenImagePipelineConfig,
) )
from sglang.multimodal_gen.configs.pipeline_configs.sana_wm import ( 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.assertEqual(qwen_deployment.keep_resident_components, ("vae",))
self.assertIsNone(qwen_deployment.keep_resident_min_available_gb) 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): def test_auto_multi_gpu_sana_wm_prefers_fsdp_and_cfg_parallel(self):
args = self._from_dict_with_pipeline_config( args = self._from_dict_with_pipeline_config(
SanaWMPipelineConfig(), SanaWMPipelineConfig(),