diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/minimax_h3_qwen3vl.py b/python/sglang/multimodal_gen/runtime/models/encoders/minimax_h3_qwen3vl.py index edb51f98d..c3ff01c4c 100644 --- a/python/sglang/multimodal_gen/runtime/models/encoders/minimax_h3_qwen3vl.py +++ b/python/sglang/multimodal_gen/runtime/models/encoders/minimax_h3_qwen3vl.py @@ -182,10 +182,12 @@ class MiniMaxH3Qwen3VLEncoder(TextEncoder): for name, loaded_weight in weights: if not self.should_materialize_checkpoint_weight(name): continue - param = params.get(name) + param_name = name.replace(".attn.qkv.", ".attn.qkv_proj.") + param = params.get(param_name) if param is None: raise KeyError( - f"Unexpected MiniMax H3 Qwen3-VL checkpoint weight: {name}" + "Unexpected MiniMax H3 Qwen3-VL checkpoint weight: " + f"{name} (mapped to {param_name})" ) weight_loader = getattr(param, "weight_loader", default_weight_loader) try: @@ -196,7 +198,7 @@ class MiniMaxH3Qwen3VLEncoder(TextEncoder): f"{name!r}: checkpoint={tuple(loaded_weight.shape)}, " f"parameter={tuple(param.shape)}" ) from exc - loaded.add(name) + loaded.add(param_name) return loaded 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 eb029d3e5..d6bf3dbc4 100644 --- a/python/sglang/multimodal_gen/runtime/models/encoders/qwen2_5vl.py +++ b/python/sglang/multimodal_gen/runtime/models/encoders/qwen2_5vl.py @@ -5,7 +5,6 @@ from transformers import ( DynamicCache, PretrainedConfig, Qwen2_5_VLTextConfig, - Qwen2RMSNorm, ) from transformers.masking_utils import ( create_causal_mask, @@ -17,23 +16,30 @@ from transformers.utils import TransformersKwargs, is_torchdynamo_compiling from sglang.multimodal_gen.configs.models.encoders.qwen_image import Qwen2_5VLConfig from sglang.multimodal_gen.runtime.distributed import ( + get_tp_rank, get_tp_world_size, model_parallel_is_initialized, ) from sglang.multimodal_gen.runtime.layers.attention import LocalAttention -from sglang.multimodal_gen.runtime.layers.linear import ( - ColumnParallelLinear, - MergedColumnParallelLinear, - RowParallelLinear, -) 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.models.encoders.qwen_vl_rope import ( + apply_qwen_vl_text_rope, + build_qwen_vl_text_rope, +) from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum -from sglang.multimodal_gen.runtime.utils.common import add_prefix +from sglang.srt.layers.layernorm import RMSNorm +from sglang.srt.layers.linear import ( + ColumnParallelLinear, + ReplicatedLinear, + RowParallelLinear, +) +from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding +from sglang.srt.models.qwen2_5_vl import Qwen2_5_VLMLP # coding=utf-8 # Adapted from @@ -70,12 +76,9 @@ except ImportError: 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_VLCausalLMOutputWithPast, Qwen2_5_VLModelOutputWithPast, - Qwen2_5_VLRotaryEmbedding, - apply_multimodal_rotary_pos_emb, ) logger = logging.getLogger(__name__) @@ -134,6 +137,12 @@ def _tp_world_size() -> int: return get_tp_world_size() +def _tp_rank() -> int: + if not model_parallel_is_initialized(): + return 0 + return get_tp_rank() + + def _linear_output(linear: nn.Module, x: torch.Tensor) -> torch.Tensor: output = linear(x) return output[0] if isinstance(output, tuple) else output @@ -152,8 +161,10 @@ def _make_column_linear( out_features, bias=bias, gather_output=False, + tp_size=_tp_world_size(), + tp_rank=_tp_rank(), ) - return nn.Linear(in_features, out_features, bias=bias) + return ReplicatedLinear(in_features, out_features, bias=bias) def _make_row_linear( @@ -168,8 +179,10 @@ def _make_row_linear( in_features, out_features, bias=bias, + tp_size=_tp_world_size(), + tp_rank=_tp_rank(), ) - return nn.Linear(in_features, out_features, bias=bias) + return ReplicatedLinear(in_features, out_features, bias=bias) class Qwen2_5_VLAttention(nn.Module): @@ -183,10 +196,9 @@ class Qwen2_5_VLAttention(nn.Module): self.config = config self.layer_idx = layer_idx if layer_idx is None: - logger.warn( - f"Instantiating {self.__class__.__name__} without passing `layer_idx` is not recommended and will " - "to errors during the forward call, if caching is used. Please make sure to provide a `layer_idx` " - "when creating this class." + logger.warning( + "Instantiating %s without layer_idx disables correct cache updates", + self.__class__.__name__, ) self.hidden_size = config.hidden_size @@ -221,7 +233,6 @@ class Qwen2_5_VLAttention(nn.Module): self.num_key_value_groups = self.num_heads // self.num_key_value_heads self.is_causal = True self.attention_dropout = config.attention_dropout - self.rope_scaling = config.rope_scaling self.scaling = self.head_dim**-0.5 self.q_proj = _make_column_linear( @@ -254,7 +265,7 @@ class Qwen2_5_VLAttention(nn.Module): else None ) - self.rotary_emb = Qwen2_5_VLRotaryEmbedding(config=config) + self.rotary_emb = build_qwen_vl_text_rope(config) self.attn = LocalAttention( num_heads=self.num_heads, head_size=self.head_dim, @@ -276,9 +287,6 @@ class Qwen2_5_VLAttention(nn.Module): output_attentions: bool = False, use_cache: bool = False, cache_position: Optional[torch.LongTensor] = None, - position_embeddings: Optional[ - tuple[torch.Tensor, torch.Tensor] - ] = None, # necessary, but kept here for BC **kwargs: Unpack[FlashAttentionKwargs], ) -> tuple[torch.Tensor, Optional[torch.Tensor], Optional[tuple[torch.Tensor]]]: bsz, q_len, _ = hidden_states.size() @@ -291,17 +299,15 @@ class Qwen2_5_VLAttention(nn.Module): key_states = key_states.view(bsz, q_len, -1, self.head_dim).transpose(1, 2) value_states = value_states.view(bsz, q_len, -1, self.head_dim).transpose(1, 2) - cos, sin = position_embeddings - query_states, key_states = apply_multimodal_rotary_pos_emb( - query_states, key_states, cos, sin, self.rope_scaling["mrope_section"] + query_states, key_states = apply_qwen_vl_text_rope( + self.rotary_emb, + position_ids, + query_states, + key_states, ) if past_key_values is not None: - cache_kwargs = { - "sin": sin, - "cos": cos, - "cache_position": cache_position, - } # Specific to RoPE models + cache_kwargs = {"cache_position": cache_position} key_states, value_states = past_key_values.update( key_states, value_states, self.layer_idx, cache_kwargs ) @@ -324,38 +330,6 @@ class Qwen2_5_VLAttention(nn.Module): return attn_output -class Qwen2_5_VLTextMLP(nn.Module): - def __init__(self, config: Qwen2_5_VLTextConfig): - super().__init__() - tp_size = _tp_world_size() - use_tensor_parallel = tp_size > 1 and config.intermediate_size % tp_size == 0 - self.gate_proj = _make_column_linear( - config.hidden_size, - config.intermediate_size, - bias=False, - use_tensor_parallel=use_tensor_parallel, - ) - self.up_proj = _make_column_linear( - config.hidden_size, - config.intermediate_size, - bias=False, - use_tensor_parallel=use_tensor_parallel, - ) - self.down_proj = _make_row_linear( - config.intermediate_size, - config.hidden_size, - bias=False, - use_tensor_parallel=use_tensor_parallel, - ) - self.act_fn = ACT2FN[config.hidden_act] - - def forward(self, x: torch.Tensor) -> torch.Tensor: - x = self.act_fn(_linear_output(self.gate_proj, x)) * _linear_output( - self.up_proj, x - ) - return _linear_output(self.down_proj, x) - - class Qwen2_5_VLDecoderLayer(nn.Module): def __init__(self, config: Qwen2_5_VLTextConfig, layer_idx: int): super().__init__() @@ -371,11 +345,26 @@ class Qwen2_5_VLDecoderLayer(nn.Module): ) self.self_attn = Qwen2_5_VLAttention(config, layer_idx) - self.mlp = Qwen2_5_VLTextMLP(config) - self.input_layernorm = Qwen2RMSNorm(config.hidden_size, eps=config.rms_norm_eps) - self.post_attention_layernorm = Qwen2RMSNorm( - config.hidden_size, eps=config.rms_norm_eps + mlp_tp_size = _tp_world_size() + if config.intermediate_size % mlp_tp_size != 0: + mlp_tp_size = 1 + self.mlp = Qwen2_5_VLMLP( + config.hidden_size, + config.intermediate_size, + bias=False, + hidden_act=config.hidden_act, + prefix=f"model.language_model.layers.{layer_idx}.mlp", + fuse_gate_up=False, + tp_size=mlp_tp_size, + tp_rank=_tp_rank() if mlp_tp_size > 1 else 0, ) + norm_kwargs = dict( + eps=config.rms_norm_eps, + cast_x_before_out_mul=True, + force_native=True, + ) + self.input_layernorm = RMSNorm(config.hidden_size, **norm_kwargs) + self.post_attention_layernorm = RMSNorm(config.hidden_size, **norm_kwargs) self.attention_type = config.layer_types[layer_idx] def forward( @@ -387,9 +376,6 @@ class Qwen2_5_VLDecoderLayer(nn.Module): output_attentions: Optional[bool] = False, use_cache: Optional[bool] = False, cache_position: Optional[torch.LongTensor] = None, - position_embeddings: Optional[ - tuple[torch.Tensor, torch.Tensor] - ] = None, # necessary, but kept here for BC **kwargs: Unpack[FlashAttentionKwargs], ) -> tuple[ torch.FloatTensor, Optional[tuple[torch.FloatTensor, torch.FloatTensor]] @@ -408,9 +394,6 @@ class Qwen2_5_VLDecoderLayer(nn.Module): past_key_values (`Tuple(torch.FloatTensor)`, *optional*): cached past key and value projection states cache_position (`torch.LongTensor` of shape `(sequence_length)`, *optional*): Indices depicting the position of the input sequence tokens in the sequence. - position_embeddings (`tuple[torch.FloatTensor, torch.FloatTensor]`, *optional*): - Tuple containing the cosine and sine positional embeddings of shape `(batch_size, seq_len, head_dim)`, - with `head_dim` being the embedding dimension of each attention head. kwargs (`dict`, *optional*): Arbitrary kwargs to be ignored, used for FSDP and other methods that injects code into the model @@ -429,7 +412,6 @@ class Qwen2_5_VLDecoderLayer(nn.Module): output_attentions=output_attentions, use_cache=use_cache, cache_position=cache_position, - position_embeddings=position_embeddings, **kwargs, ) hidden_states = residual + hidden_states @@ -443,41 +425,6 @@ class Qwen2_5_VLDecoderLayer(nn.Module): return hidden_states -class Qwen2_5_VLMLP(nn.Module): - def __init__( - self, - in_features: int, - hidden_features: int = None, - bias: bool = True, - hidden_act="silu", - quant_config: Optional[QuantizationConfig] = None, - prefix: str = "", - ): - super().__init__() - self.gate_up_proj = MergedColumnParallelLinear( - input_size=in_features, - output_sizes=[hidden_features] * 2, # [gate_proj, up_proj] - bias=bias, - quant_config=quant_config, - prefix=add_prefix("gate_up_proj", prefix), - ) - self.down_proj = RowParallelLinear( - hidden_features, - in_features, - bias=bias, - quant_config=quant_config, - prefix=add_prefix("down_proj", prefix), - ) - self.act = ACT2FN[hidden_act] - - def forward(self, x: torch.Tensor) -> torch.Tensor: - gate_up, _ = self.gate_up_proj(x) - gate, up = gate_up.chunk(2, dim=-1) - x = self.act(gate) * up - x_down, _ = self.down_proj(x) - return x_down - - class Qwen2_5_VLTextModel(nn.Module): def __init__(self, config: PretrainedConfig): super().__init__() @@ -485,8 +432,11 @@ class Qwen2_5_VLTextModel(nn.Module): self.padding_idx = config.pad_token_id self.vocab_size = config.vocab_size - self.embed_tokens = nn.Embedding( - config.vocab_size, config.hidden_size, self.padding_idx + self.embed_tokens = VocabParallelEmbedding( + config.vocab_size, + config.hidden_size, + org_num_embeddings=config.vocab_size, + prefix="model.language_model.embed_tokens", ) self.layers = nn.ModuleList( [ @@ -495,8 +445,12 @@ class Qwen2_5_VLTextModel(nn.Module): ] ) self._attn_implementation = config._attn_implementation - self.norm = Qwen2RMSNorm(config.hidden_size, eps=config.rms_norm_eps) - self.rotary_emb = Qwen2_5_VLRotaryEmbedding(config=config) + self.norm = RMSNorm( + config.hidden_size, + eps=config.rms_norm_eps, + cast_x_before_out_mul=True, + force_native=True, + ) self.has_sliding_layers = "sliding_attention" in self.config.layer_types self.gradient_checkpointing = False @@ -600,9 +554,6 @@ class Qwen2_5_VLTextModel(nn.Module): hidden_states = inputs_embeds - # create position embeddings to be shared across the decoder layers - position_embeddings = self.rotary_emb(hidden_states, position_ids) - # decoder layers all_hidden_states = () if output_hidden_states else None all_self_attns = () if output_attentions else None @@ -614,12 +565,11 @@ class Qwen2_5_VLTextModel(nn.Module): hidden_states = decoder_layer( hidden_states, attention_mask=causal_mask_mapping[decoder_layer.attention_type], - position_ids=text_position_ids, + position_ids=position_ids, past_key_values=past_key_values, output_attentions=output_attentions, use_cache=use_cache, cache_position=cache_position, - position_embeddings=position_embeddings, **kwargs, ) @@ -1426,6 +1376,26 @@ class Qwen2_5_VLForConditionalGeneration(TextEncoder): if not self.enable_image_understanding: continue name = name.replace("visual.", "model.visual.") + name = name.replace(".attn.qkv.", ".attn.qkv_proj.") + + loaded_stacked_param = False + for weight_name, shard_id in ( + (".gate_proj.", 0), + (".up_proj.", 1), + ): + if weight_name not in name: + continue + fused_name = name.replace(weight_name, ".gate_up_proj.") + if fused_name not in params_dict: + continue + param = params_dict[fused_name] + loaded_weight = loaded_weight.to(param.dtype) + param.weight_loader(param, loaded_weight, shard_id) + loaded_params.add(fused_name) + loaded_stacked_param = True + break + if loaded_stacked_param: + continue try: # Skip loading extra bias for GPTQ models. if name.endswith(".bias") and name not in params_dict: 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 index acd3c4647..7d5978208 100644 --- a/python/sglang/multimodal_gen/runtime/models/encoders/qwen2_5vl_vision.py +++ b/python/sglang/multimodal_gen/runtime/models/encoders/qwen2_5vl_vision.py @@ -3,285 +3,75 @@ 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) - ) +from sglang.multimodal_gen.runtime.models.encoders.qwen_vl_vision import ( + PackedSequenceMetadata, + QwenVLVisionAttention, +) +from sglang.srt.layers.layernorm import RMSNorm +from sglang.srt.models.qwen2_5_vl import ( + Qwen2_5_VisionPatchEmbed as Qwen2_5VLVisionPatchEmbed, +) +from sglang.srt.models.qwen2_5_vl import ( + Qwen2_5_VisionPatchMerger as Qwen2_5VLVisionPatchMerger, +) +from sglang.srt.models.qwen2_5_vl import ( + Qwen2_5_VisionRotaryEmbedding as Qwen2_5VLVisionRotaryEmbedding, +) +from sglang.srt.models.qwen2_5_vl import ( + Qwen2_5_VLMLP, +) 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.norm1 = RMSNorm( + config.hidden_size, + eps=1e-6, + cast_x_before_out_mul=True, + force_native=True, + ) + self.norm2 = RMSNorm( + config.hidden_size, + eps=1e-6, + cast_x_before_out_mul=True, + force_native=True, + ) + self.attn = QwenVLVisionAttention( + config, + prefix=f"visual.blocks.{layer_idx}.attn", + model_name="Qwen2.5-VL", + ) + self.mlp = Qwen2_5_VLMLP( + config.hidden_size, + config.intermediate_size, + bias=True, + hidden_act=config.hidden_act, + prefix=f"visual.blocks.{layer_idx}.mlp", + fuse_gate_up=False, ) - self.mlp = Qwen2_5VLVisionMLP(config) def forward( self, hidden_states: torch.Tensor, *, - cu_seqlens: torch.Tensor, - cu_seqlens_host: tuple[int, ...], - max_seqlen: int, + metadata: PackedSequenceMetadata, 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, + metadata=metadata, 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: @@ -383,6 +173,7 @@ class Qwen2_5VLVisionTransformer(nn.Module): temporal_patch_size=config.temporal_patch_size, in_channels=config.in_channels, embed_dim=config.hidden_size, + disable_linear=True, ) head_dim = config.hidden_size // config.num_heads self.rotary_pos_emb = Qwen2_5VLVisionRotaryEmbedding(head_dim // 2) @@ -390,9 +181,13 @@ class Qwen2_5VLVisionTransformer(nn.Module): Qwen2_5VLVisionBlock(config, layer_idx) for layer_idx in range(config.depth) ) self.merger = Qwen2_5VLVisionPatchMerger( - output_dim=config.out_hidden_size, + dim=config.out_hidden_size, context_dim=config.hidden_size, + padded_context_dim=config.hidden_size, spatial_merge_size=config.spatial_merge_size, + prefix="visual.merger", + cast_x_before_out_mul=True, + force_native_norm=True, ) @property @@ -440,8 +235,8 @@ class Qwen2_5VLVisionTransformer(nn.Module): ).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) + 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 = ( @@ -451,9 +246,7 @@ class Qwen2_5VLVisionTransformer(nn.Module): ) hidden_states = block( hidden_states, - cu_seqlens=metadata.cu_seqlens, - cu_seqlens_host=metadata.cu_seqlens_host, - max_seqlen=metadata.max_seqlen, + metadata=metadata, position_embeddings=position_embeddings, ) diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/qwen3.py b/python/sglang/multimodal_gen/runtime/models/encoders/qwen3.py index 8c90cc11d..6e7edb8d1 100644 --- a/python/sglang/multimodal_gen/runtime/models/encoders/qwen3.py +++ b/python/sglang/multimodal_gen/runtime/models/encoders/qwen3.py @@ -7,9 +7,8 @@ from torch import nn from sglang.multimodal_gen.configs.models.encoders import BaseEncoderOutput from sglang.multimodal_gen.configs.models.encoders.qwen3 import Qwen3TextConfig from sglang.multimodal_gen.runtime.distributed import get_tp_world_size -from sglang.multimodal_gen.runtime.layers.activation import SiluAndMul from sglang.multimodal_gen.runtime.layers.attention import LocalAttention -from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm +from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm as MMGenRMSNorm from sglang.multimodal_gen.runtime.layers.linear import ( MergedColumnParallelLinear, QKVParallelLinear, @@ -25,6 +24,8 @@ from sglang.multimodal_gen.runtime.loader.weight_utils import ( maybe_remap_kv_scale_name, ) from sglang.multimodal_gen.runtime.models.encoders.base import TextEncoder +from sglang.srt.layers.activation import SiluAndMul +from sglang.srt.layers.layernorm import RMSNorm class Qwen3MLP(nn.Module): @@ -131,8 +132,9 @@ class Qwen3Attention(nn.Module): # QK-Norm: Key difference from LLaMA rms_norm_eps = getattr(config, "rms_norm_eps", 1e-6) - self.q_norm = RMSNorm(self.head_dim, eps=rms_norm_eps) - self.k_norm = RMSNorm(self.head_dim, eps=rms_norm_eps) + # Keep the small-hidden one-pass kernel used by diffusion QK norm. + self.q_norm = MMGenRMSNorm(self.head_dim, eps=rms_norm_eps) + self.k_norm = MMGenRMSNorm(self.head_dim, eps=rms_norm_eps) # Rotary embeddings self.rotary_emb = get_rope( diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/qwen3vl.py b/python/sglang/multimodal_gen/runtime/models/encoders/qwen3vl.py index 720ad375c..2a7f5e565 100644 --- a/python/sglang/multimodal_gen/runtime/models/encoders/qwen3vl.py +++ b/python/sglang/multimodal_gen/runtime/models/encoders/qwen3vl.py @@ -32,7 +32,12 @@ from sglang.multimodal_gen.runtime.models.encoders.base import TextEncoder from sglang.multimodal_gen.runtime.models.encoders.qwen3vl_vision import ( Qwen3VLVisionTransformer, ) +from sglang.multimodal_gen.runtime.models.encoders.qwen_vl_rope import ( + apply_qwen_vl_text_rope, + build_qwen_vl_text_rope, +) from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum +from sglang.srt.layers.layernorm import RMSNorm """Inference-only Qwen3-VL model compatible with HuggingFace weights.""" import logging @@ -57,12 +62,18 @@ from transformers.models.qwen3_vl.configuration_qwen3_vl import ( from transformers.models.qwen3_vl.modeling_qwen3_vl import ( Qwen3VLCausalLMOutputWithPast, Qwen3VLModelOutputWithPast, - Qwen3VLTextRMSNorm, - Qwen3VLTextRotaryEmbedding, - apply_rotary_pos_emb, ) +def _make_text_rms_norm(hidden_size: int, eps: float) -> RMSNorm: + return RMSNorm( + hidden_size, + eps=eps, + cast_x_before_out_mul=True, + force_native=True, + ) + + class Qwen3VLQuantizedLinear(ReplicatedLinear): def forward(self, x: torch.Tensor) -> torch.Tensor: return super().forward(x)[0] @@ -270,12 +281,9 @@ class Qwen3VLTextAttention(nn.Module): use_tensor_parallel=use_tensor_parallel, prefix=f"{prefix}.o_proj", ) - self.q_norm = Qwen3VLTextRMSNorm( - self.head_dim, eps=config.rms_norm_eps - ) # unlike olmo, only on the head dim! - self.k_norm = Qwen3VLTextRMSNorm( - self.head_dim, eps=config.rms_norm_eps - ) # thus post q_norm does not need reshape + self.q_norm = _make_text_rms_norm(self.head_dim, config.rms_norm_eps) + self.k_norm = _make_text_rms_norm(self.head_dim, config.rms_norm_eps) + self.rotary_emb = build_qwen_vl_text_rope(config, mrope_interleaved=True) self.attn = LocalAttention( num_heads=self.num_heads, @@ -292,7 +300,7 @@ class Qwen3VLTextAttention(nn.Module): def forward( self, hidden_states: torch.Tensor, - position_embeddings: tuple[torch.Tensor, torch.Tensor], + position_ids: torch.LongTensor, attention_mask: Optional[torch.Tensor], past_key_values: Optional[Cache] = None, cache_position: Optional[torch.LongTensor] = None, @@ -309,14 +317,15 @@ class Qwen3VLTextAttention(nn.Module): ).transpose(1, 2) value_states = self.v_proj(hidden_states).view(hidden_shape).transpose(1, 2) - cos, sin = position_embeddings - query_states, key_states = apply_rotary_pos_emb( - query_states, key_states, cos, sin + query_states, key_states = apply_qwen_vl_text_rope( + self.rotary_emb, + position_ids, + query_states, + key_states, ) if past_key_values is not None: - # sin and cos are specific to RoPE models; cache_position needed for the static cache - cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position} + cache_kwargs = {"cache_position": cache_position} key_states, value_states = past_key_values.update( key_states, value_states, self.layer_idx, cache_kwargs ) @@ -432,17 +441,16 @@ class Qwen3VLTextDecoderLayer(nn.Module): use_tensor_parallel=use_tensor_parallel, prefix=f"{prefix}.mlp", ) - self.input_layernorm = Qwen3VLTextRMSNorm( - config.hidden_size, eps=config.rms_norm_eps + self.input_layernorm = _make_text_rms_norm( + config.hidden_size, config.rms_norm_eps ) - self.post_attention_layernorm = Qwen3VLTextRMSNorm( - config.hidden_size, eps=config.rms_norm_eps + self.post_attention_layernorm = _make_text_rms_norm( + config.hidden_size, config.rms_norm_eps ) def forward( self, hidden_states: torch.Tensor, - position_embeddings: tuple[torch.Tensor, torch.Tensor], attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.LongTensor] = None, past_key_values: Optional[Cache] = None, @@ -460,7 +468,6 @@ class Qwen3VLTextDecoderLayer(nn.Module): past_key_values=past_key_values, use_cache=use_cache, cache_position=cache_position, - position_embeddings=position_embeddings, **kwargs, ) hidden_states = residual + hidden_states @@ -505,8 +512,7 @@ class Qwen3VLTextModel(nn.Module): for layer_idx in range(config.num_hidden_layers) ] ) - self.norm = Qwen3VLTextRMSNorm(config.hidden_size, eps=config.rms_norm_eps) - self.rotary_emb = Qwen3VLTextRotaryEmbedding(config=config) + self.norm = _make_text_rms_norm(config.hidden_size, config.rms_norm_eps) self.gradient_checkpointing = False # Initialize weights and apply final processing @@ -582,15 +588,10 @@ class Qwen3VLTextModel(nn.Module): position_ids = position_ids[None, ...].expand(3, position_ids.shape[0], -1) if position_ids.ndim == 3 and position_ids.shape[0] == 4: - text_position_ids = position_ids[0] position_ids = position_ids[1:] - else: - text_position_ids = position_ids[0] hidden_states = inputs_embeds - # create position embeddings to be shared across the decoder layers - position_embeddings = self.rotary_emb(hidden_states, position_ids) all_hidden_states = () if output_hidden_states else None all_self_attns = () if output_attentions else None # decoder layers @@ -598,11 +599,10 @@ class Qwen3VLTextModel(nn.Module): hidden_states = decoder_layer( hidden_states, attention_mask=attention_mask, - position_ids=text_position_ids, + position_ids=position_ids, past_key_values=past_key_values, cache_position=cache_position, output_attentions=output_attentions, - position_embeddings=position_embeddings, **kwargs, ) # hidden_states = layer_outputs @@ -1269,6 +1269,8 @@ class Qwen3VLForConditionalGeneration(TextEncoder): for name, loaded_weight in weights: if "rotary_emb.inv_freq" in name: continue + if "visual." in name: + name = name.replace(".attn.qkv.", ".attn.qkv_proj.") try: param = params_dict[name] diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/qwen3vl_vision.py b/python/sglang/multimodal_gen/runtime/models/encoders/qwen3vl_vision.py index fa0109896..fc60cf501 100644 --- a/python/sglang/multimodal_gen/runtime/models/encoders/qwen3vl_vision.py +++ b/python/sglang/multimodal_gen/runtime/models/encoders/qwen3vl_vision.py @@ -10,10 +10,16 @@ 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__) +from sglang.multimodal_gen.runtime.models.encoders.qwen_vl_vision import ( + PackedSequenceMetadata, + QwenVLVisionAttention, +) +from sglang.srt.models.qwen3_vl import ( + Qwen3_VisionMLP, + Qwen3VLMoeVisionPatchMerger, + Qwen3VLVisionPatchEmbed, +) +from sglang.srt.runtime_context import get_parallel @dataclass(frozen=True) @@ -23,57 +29,6 @@ class Qwen3VLVisionOutput: deepstack_features: list[torch.Tensor] -@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 Qwen3VLVisionPatchEmbed(nn.Module): - def __init__(self, config: Any) -> None: - super().__init__() - self.patch_size = config.patch_size - self.temporal_patch_size = config.temporal_patch_size - self.in_channels = config.in_channels - self.embed_dim = config.hidden_size - kernel_size = ( - config.temporal_patch_size, - config.patch_size, - config.patch_size, - ) - self.proj = nn.Conv3d( - config.in_channels, - config.hidden_size, - kernel_size=kernel_size, - stride=kernel_size, - bias=True, - ) - - 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 Qwen3VLVisionRotaryEmbedding(nn.Module): def __init__(self, dim: int, theta: float = 10000.0) -> None: super().__init__() @@ -89,140 +44,32 @@ class Qwen3VLVisionRotaryEmbedding(nn.Module): return torch.outer(positions, self.inv_freq) -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 Qwen3VLVisionAttention(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.qkv = nn.Linear(config.hidden_size, config.hidden_size * 3, bias=True) - self.proj = nn.Linear(config.hidden_size, config.hidden_size) - backend = get_attn_backend(self.head_dim, torch.get_default_dtype()) - self._attention_impl = None - 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=prefix, - ) - else: - logger.warning_once( - "Qwen3-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, - metadata: _PackedSequenceMetadata, - ) -> torch.Tensor: - if self._attention_impl is not None: - return self._attention_impl.forward_varlen( - query, - key, - value, - cu_seqlens=metadata.cu_seqlens, - cu_seqlens_host=metadata.cu_seqlens_host, - max_seqlen=metadata.max_seqlen, - ) - - output = torch.empty_like(query) - for start, stop in zip( - metadata.cu_seqlens_host[:-1], metadata.cu_seqlens_host[1:] - ): - if start == stop: - continue - segment = F.scaled_dot_product_attention( - query[start:stop].transpose(0, 1).unsqueeze(0), - key[start:stop].transpose(0, 1).unsqueeze(0), - value[start:stop].transpose(0, 1).unsqueeze(0), - 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, - *, - metadata: _PackedSequenceMetadata, - position_embeddings: tuple[torch.Tensor, torch.Tensor], - ) -> torch.Tensor: - sequence_length = hidden_states.shape[0] - query, key, value = ( - self.qkv(hidden_states) - .reshape(sequence_length, 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, metadata) - return self.proj(output.reshape(sequence_length, -1).contiguous()) - - -class Qwen3VLVisionMLP(nn.Module): - def __init__(self, config: Any) -> None: - super().__init__() - if config.hidden_act != "gelu_pytorch_tanh": - raise ValueError( - f"Unsupported Qwen3-VL vision activation: {config.hidden_act}" - ) - self.linear_fc1 = nn.Linear( - config.hidden_size, config.intermediate_size, bias=True - ) - self.linear_fc2 = nn.Linear( - config.intermediate_size, config.hidden_size, bias=True - ) - self.act_fn = nn.GELU(approximate="tanh") - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - return self.linear_fc2(self.act_fn(self.linear_fc1(hidden_states))) - - class Qwen3VLVisionBlock(nn.Module): def __init__(self, config: Any, layer_idx: int) -> None: super().__init__() + parallel = get_parallel() self.norm1 = nn.LayerNorm(config.hidden_size, eps=1e-6) self.norm2 = nn.LayerNorm(config.hidden_size, eps=1e-6) - self.attn = Qwen3VLVisionAttention( - config, prefix=f"visual.blocks.{layer_idx}.attn" + self.attn = QwenVLVisionAttention( + config, + prefix=f"visual.blocks.{layer_idx}.attn", + model_name="Qwen3-VL", + ) + self.mlp = Qwen3_VisionMLP( + config.hidden_size, + config.intermediate_size, + bias=True, + hidden_act=config.hidden_act, + prefix=f"visual.blocks.{layer_idx}.mlp", + tp_rank=parallel.tp_rank, + tp_size=parallel.tp_size, ) - self.mlp = Qwen3VLVisionMLP(config) def forward( self, hidden_states: torch.Tensor, *, - metadata: _PackedSequenceMetadata, + metadata: PackedSequenceMetadata, position_embeddings: tuple[torch.Tensor, torch.Tensor], ) -> torch.Tensor: hidden_states = hidden_states + self.attn( @@ -233,24 +80,6 @@ class Qwen3VLVisionBlock(nn.Module): return hidden_states + self.mlp(self.norm2(hidden_states)) -class Qwen3VLVisionPatchMerger(nn.Module): - def __init__(self, config: Any, *, use_postshuffle_norm: bool) -> None: - super().__init__() - self.hidden_size = config.hidden_size * config.spatial_merge_size**2 - self.use_postshuffle_norm = use_postshuffle_norm - norm_size = self.hidden_size if use_postshuffle_norm else config.hidden_size - self.norm = nn.LayerNorm(norm_size, eps=1e-6) - self.linear_fc1 = nn.Linear(self.hidden_size, self.hidden_size) - self.act_fn = nn.GELU() - self.linear_fc2 = nn.Linear(self.hidden_size, config.out_hidden_size) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if self.use_postshuffle_norm: - hidden_states = hidden_states.view(-1, self.hidden_size) - hidden_states = self.norm(hidden_states).view(-1, self.hidden_size) - return self.linear_fc2(self.act_fn(self.linear_fc1(hidden_states))) - - def _vision_position_ids( grid_thw: torch.Tensor, spatial_merge_size: int ) -> torch.Tensor: @@ -349,11 +178,12 @@ def _vision_cu_seqlens(grid_thw: torch.Tensor) -> torch.Tensor: class Qwen3VLVisionTransformer(nn.Module): def __init__(self, config: Any) -> None: super().__init__() + parallel = get_parallel() self.config = config self.spatial_merge_size = config.spatial_merge_size self.spatial_merge_unit = config.spatial_merge_size**2 self.patch_size = config.patch_size - self.patch_embed = Qwen3VLVisionPatchEmbed(config) + self.patch_embed = Qwen3VLVisionPatchEmbed(config, disable_linear=True) self.pos_embed = nn.Embedding( config.num_position_embeddings, config.hidden_size ) @@ -363,11 +193,29 @@ class Qwen3VLVisionTransformer(nn.Module): self.blocks = nn.ModuleList( Qwen3VLVisionBlock(config, layer_idx) for layer_idx in range(config.depth) ) - self.merger = Qwen3VLVisionPatchMerger(config, use_postshuffle_norm=False) + self.merger = Qwen3VLMoeVisionPatchMerger( + dim=config.out_hidden_size, + context_dim=config.hidden_size, + padded_context_dim=config.hidden_size, + spatial_merge_size=config.spatial_merge_size, + use_postshuffle_norm=False, + prefix="visual.merger", + tp_rank=parallel.tp_rank, + tp_size=parallel.tp_size, + ) self.deepstack_visual_indexes = tuple(config.deepstack_visual_indexes) self.deepstack_merger_list = nn.ModuleList( - Qwen3VLVisionPatchMerger(config, use_postshuffle_norm=True) - for _ in self.deepstack_visual_indexes + Qwen3VLMoeVisionPatchMerger( + dim=config.out_hidden_size, + context_dim=config.hidden_size, + padded_context_dim=config.hidden_size, + spatial_merge_size=config.spatial_merge_size, + use_postshuffle_norm=True, + prefix=f"visual.deepstack_merger_list.{merger_idx}", + tp_rank=parallel.tp_rank, + tp_size=parallel.tp_size, + ) + for merger_idx, _ in enumerate(self.deepstack_visual_indexes) ) self._deepstack_merger_by_layer = { layer_idx: merger_idx @@ -407,7 +255,7 @@ class Qwen3VLVisionTransformer(nn.Module): rotary = rotary.flatten(1) rotary = torch.cat((rotary, rotary), dim=-1) position_embeddings = (rotary.cos(), rotary.sin()) - metadata = _PackedSequenceMetadata.from_cu_seqlens(_vision_cu_seqlens(grid_thw)) + metadata = PackedSequenceMetadata.from_cu_seqlens(_vision_cu_seqlens(grid_thw)) deepstack_features = [] for layer_idx, block in enumerate(self.blocks): diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/qwen_vl_rope.py b/python/sglang/multimodal_gen/runtime/models/encoders/qwen_vl_rope.py new file mode 100644 index 000000000..85077a925 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/models/encoders/qwen_vl_rope.py @@ -0,0 +1,65 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Shared SRT rotary embedding adapter for Qwen-VL text encoders.""" + +from typing import Any + +import torch + +from sglang.srt.layers.rotary_embedding import get_rope +from sglang.srt.layers.rotary_embedding.base import RotaryEmbedding +from sglang.srt.utils.hf_transformers.common import get_rope_config + + +def build_qwen_vl_text_rope( + config: Any, *, mrope_interleaved: bool = False +) -> RotaryEmbedding: + head_dim = getattr(config, "head_dim", None) or ( + config.hidden_size // config.num_attention_heads + ) + rope_theta, rope_scaling = get_rope_config(config) + rope_scaling = dict(rope_scaling or {}) + rope_scaling["mrope_interleaved"] = mrope_interleaved + return get_rope( + head_size=head_dim, + rotary_dim=head_dim, + max_position=config.max_position_embeddings, + base=rope_theta, + is_neox_style=True, + rope_scaling=rope_scaling, + ) + + +def apply_qwen_vl_text_rope( + rotary_emb: RotaryEmbedding, + position_ids: torch.Tensor, + query: torch.Tensor, + key: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor]: + """Apply three-axis MRoPE to batched attention tensors.""" + if query.ndim != 4 or key.ndim != 4: + raise ValueError( + "Qwen-VL query and key must have shape " + "[batch, heads, sequence, head_dim]" + ) + if position_ids.ndim != 3 or position_ids.shape[0] != 3: + raise ValueError( + "Qwen-VL text position_ids must have shape [3, batch, sequence]" + ) + batch_size, num_query_heads, sequence_length, head_dim = query.shape + key_batch_size, num_key_value_heads, key_sequence_length, key_head_dim = key.shape + if (key_batch_size, key_sequence_length, key_head_dim) != ( + batch_size, + sequence_length, + head_dim, + ): + raise ValueError("Qwen-VL query and key shapes are incompatible") + if tuple(position_ids.shape[1:]) != (batch_size, sequence_length): + raise ValueError("Qwen-VL position_ids do not match the attention input") + + query = query.transpose(1, 2).reshape(-1, num_query_heads * head_dim) + key = key.transpose(1, 2).reshape(-1, num_key_value_heads * head_dim) + # Preserve HF's bf16 arithmetic order; fused MRoPE changes generated images. + query, key = rotary_emb.forward_native(position_ids.reshape(3, -1), query, key) + query = query.view(batch_size, sequence_length, num_query_heads, head_dim) + key = key.view(batch_size, sequence_length, num_key_value_heads, head_dim) + return query.transpose(1, 2), key.transpose(1, 2) diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/qwen_vl_vision.py b/python/sglang/multimodal_gen/runtime/models/encoders/qwen_vl_vision.py new file mode 100644 index 000000000..a8fcc624e --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/models/encoders/qwen_vl_vision.py @@ -0,0 +1,154 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Shared Qwen-VL vision attention for multimodal generation.""" + +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 +from sglang.srt.layers.linear import QKVParallelLinear, RowParallelLinear +from sglang.srt.runtime_context import get_parallel + +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:]) + ), + ) + + +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_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 QwenVLVisionAttention(nn.Module): + def __init__(self, config: Any, *, prefix: str, model_name: str) -> None: + super().__init__() + parallel = get_parallel() + self.num_heads = config.num_heads // parallel.tp_size + self.head_dim = config.hidden_size // config.num_heads + self.scaling = self.head_dim**-0.5 + self.qkv_proj = QKVParallelLinear( + hidden_size=config.hidden_size, + head_size=self.head_dim, + total_num_heads=config.num_heads, + bias=True, + prefix=f"{prefix}.qkv_proj", + tp_rank=parallel.tp_rank, + tp_size=parallel.tp_size, + ) + self.proj = RowParallelLinear( + input_size=config.hidden_size, + output_size=config.hidden_size, + bias=True, + prefix=f"{prefix}.proj", + tp_rank=parallel.tp_rank, + tp_size=parallel.tp_size, + ) + + backend = get_attn_backend(self.head_dim, torch.get_default_dtype()) + self._attention_impl = None + 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=prefix, + ) + else: + logger.warning_once( + f"{model_name} 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, + metadata: PackedSequenceMetadata, + ) -> torch.Tensor: + if self._attention_impl is not None: + return self._attention_impl.forward_varlen( + query, + key, + value, + cu_seqlens=metadata.cu_seqlens, + cu_seqlens_host=metadata.cu_seqlens_host, + max_seqlen=metadata.max_seqlen, + ) + + output = torch.empty_like(query) + for start, stop in zip( + metadata.cu_seqlens_host[:-1], metadata.cu_seqlens_host[1:] + ): + if start == stop: + continue + segment = F.scaled_dot_product_attention( + query[start:stop].transpose(0, 1).unsqueeze(0), + key[start:stop].transpose(0, 1).unsqueeze(0), + value[start:stop].transpose(0, 1).unsqueeze(0), + 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, + *, + metadata: PackedSequenceMetadata, + position_embeddings: tuple[torch.Tensor, torch.Tensor], + ) -> torch.Tensor: + sequence_length = hidden_states.shape[0] + qkv, _ = self.qkv_proj(hidden_states) + query, key, value = ( + qkv.reshape(sequence_length, 3, self.num_heads, self.head_dim) + .permute(1, 0, 2, 3) + .unbind(0) + ) + query, key = _apply_rotary_embedding(query, key, *position_embeddings) + output = self._packed_attention(query, key, value, metadata) + output, _ = self.proj(output.reshape(sequence_length, -1).contiguous()) + return output diff --git a/python/sglang/multimodal_gen/test/server/gpu_cases.py b/python/sglang/multimodal_gen/test/server/gpu_cases.py index 6719f0de8..59d2f0afd 100644 --- a/python/sglang/multimodal_gen/test/server/gpu_cases.py +++ b/python/sglang/multimodal_gen/test/server/gpu_cases.py @@ -1177,6 +1177,7 @@ STANDALONE_FILES = { "../single_test_file/test_disagg_server.py", "../single_test_file/test_ar_models.py", "../single_test_file/test_ipc_a2a_2_gpu.py", + "../single_test_file/test_encoder_fold_srt_linear_2_gpu.py", "../single_test_file/test_encoder_fold_srt_2_gpu.py", "../single_test_file/test_diffusion_bcg_tp2_zimage_turbo.py", "../single_test_file/test_dp_serving_2_gpu.py", @@ -1216,6 +1217,7 @@ STANDALONE_FILE_EST_TIMES = { "../single_test_file/test_ar_models.py": 600.0, # no model load; the cost is the one-time JIT build of the sync kernels "../single_test_file/test_ipc_a2a_2_gpu.py": 240.0, + "../single_test_file/test_encoder_fold_srt_linear_2_gpu.py": 120.0, "../single_test_file/test_encoder_fold_srt_2_gpu.py": 240.0, # ~60 s locally with a warm HF cache (load + one capture + 4 steps); # padded for cold-cache CI. diff --git a/python/sglang/multimodal_gen/test/single_test_file/component_accuracy/engine.py b/python/sglang/multimodal_gen/test/single_test_file/component_accuracy/engine.py index b500b21d1..72749d513 100644 --- a/python/sglang/multimodal_gen/test/single_test_file/component_accuracy/engine.py +++ b/python/sglang/multimodal_gen/test/single_test_file/component_accuracy/engine.py @@ -477,37 +477,40 @@ class AccuracyEngine: for name, tensor in target.named_parameters(): total += 1 src_tensor = None - for cand in generate_name_candidates(name, reverse_mapping): + candidates = generate_name_candidates(name, reverse_mapping) + for cand in candidates: if cand in lookup: src_tensor = lookup[cand] break if src_tensor is None: - for cand in generate_name_candidates(name, reverse_mapping): + for cand in candidates: src_tensor = fuse_qkv(lookup, cand) if src_tensor is not None: break if src_tensor is None: - for cand in generate_name_candidates(name, reverse_mapping): + for cand in candidates: src_tensor = fuse_gate_up_proj(lookup, cand) if src_tensor is not None: break - if src_tensor is None: - unmatched_details.append(f"{name}: no matching source tensor") - continue shard_context = shard_contexts.get(name) shard_world_size = ( shard_context.world_size if shard_context is not None else tp_world ) shard_rank = shard_context.rank if shard_context is not None else rank - # TP-sharded params must load via their own weight_loader; the - # generic narrow mis-slices fused QKV/gate_up projections. - needs_weight_loader = ( - shard_world_size > 1 or tensor.shape != src_tensor.shape + # Production loaders own fused projection sharding and alignment + # padding. Use them for TP parameters and whenever a direct copy + # cannot represent the source layout, including TP=1. + requires_weight_loader = ( + shard_world_size > 1 + or src_tensor is None + or src_tensor.shape != tensor.shape ) - if needs_weight_loader and load_param_with_weight_loader( + if requires_weight_loader and load_param_with_weight_loader( tensor, name, lookup, reverse_mapping ): matched += 1 + elif src_tensor is None: + unmatched_details.append(f"{name}: no matching source tensor") elif copy_tensor(tensor, src_tensor, shard_world_size, shard_rank): matched += 1 else: diff --git a/python/sglang/multimodal_gen/test/single_test_file/component_accuracy/utils.py b/python/sglang/multimodal_gen/test/single_test_file/component_accuracy/utils.py index 453e94b22..bb9df359c 100644 --- a/python/sglang/multimodal_gen/test/single_test_file/component_accuracy/utils.py +++ b/python/sglang/multimodal_gen/test/single_test_file/component_accuracy/utils.py @@ -915,10 +915,14 @@ def load_param_with_weight_loader(param, name, lookup, reverse_mapping) -> bool: loader(param, tensor.to(dtype=param.dtype), shard_id) return True for cand in candidates: - src = lookup.get(cand) - if src is not None: - loader(param, src.to(dtype=param.dtype)) - return True + source_names = [cand] + if "qkv_proj" in cand: + source_names.append(cand.replace("qkv_proj", "qkv")) + for source_name in source_names: + src = lookup.get(source_name) + if src is not None: + loader(param, src.to(dtype=param.dtype)) + return True except Exception: return False return False diff --git a/python/sglang/multimodal_gen/test/single_test_file/test_encoder_fold_srt_linear_2_gpu.py b/python/sglang/multimodal_gen/test/single_test_file/test_encoder_fold_srt_linear_2_gpu.py new file mode 100644 index 000000000..5d7a3bea8 --- /dev/null +++ b/python/sglang/multimodal_gen/test/single_test_file/test_encoder_fold_srt_linear_2_gpu.py @@ -0,0 +1,120 @@ +"""A folded encoder must run SRT collectives on its bound TP group.""" + +from __future__ import annotations + +import os +import subprocess +import sys +import unittest + +import torch +import torch.nn.functional as F +from torch import nn + +from sglang.multimodal_gen.runtime.platforms import current_platform +from sglang.test.test_utils import CustomTestCase + +_WORLD_SIZE = 2 + + +def _worker() -> int: + from sglang.multimodal_gen.runtime.distributed import ( + cleanup_dist_env_and_memory, + get_tp_group, + get_world_group, + init_distributed_environment, + initialize_model_parallel, + ) + from sglang.multimodal_gen.runtime.models.encoders.base import ( + EncoderTensorParallelMixin, + ) + from sglang.srt.distributed import parallel_state as srt_parallel_state + from sglang.srt.layers.linear import RowParallelLinear + + rank = int(os.environ["RANK"]) + world_size = int(os.environ["WORLD_SIZE"]) + device = torch.device(f"cuda:{rank}") + torch.cuda.set_device(device) + init_distributed_environment( + world_size=world_size, + rank=rank, + local_rank=rank, + ) + initialize_model_parallel( + tensor_parallel_degree=1, + sequence_parallel_degree=world_size, + ulysses_degree=world_size, + ring_degree=1, + ) + + class FoldedEncoder(EncoderTensorParallelMixin, nn.Module): + def __init__(self): + super().__init__() + self.bind_encoder_tp_group(get_world_group()) + self.proj = RowParallelLinear( + input_size=8, + output_size=6, + bias=False, + tp_rank=rank, + tp_size=world_size, + params_dtype=torch.float32, + ) + + def forward(self, inputs: torch.Tensor) -> torch.Tensor: + local_inputs = inputs.chunk(world_size, dim=-1)[rank].contiguous() + output, _ = self.proj(local_inputs) + return output + + full_weight = torch.arange(48, dtype=torch.float32, device=device).reshape(6, 8) + full_weight = (full_weight - 23.5) / 32 + inputs = torch.arange(24, dtype=torch.float32, device=device).reshape(3, 8) / 8 + model = FoldedEncoder().to(device).eval() + with torch.no_grad(): + model.proj.weight.copy_(full_weight[:, rank * 4 : (rank + 1) * 4].contiguous()) + expected = F.linear(inputs, full_weight) + actual = model(inputs) + + torch.testing.assert_close(actual, expected, rtol=1e-6, atol=1e-6) + assert get_tp_group().world_size == 1 + assert srt_parallel_state.get_tp_group().world_size == 1 + assert srt_parallel_state.get_attn_tp_group().world_size == 1 + + if rank == 0: + print("ENCODER_FOLD_SRT_LINEAR_PARITY PASS", flush=True) + torch.distributed.barrier() + cleanup_dist_env_and_memory() + return 0 + + +class TestEncoderFoldSrtLinearTwoGpu(CustomTestCase): + def test_folded_srt_linear_matches_unsharded_reference(self): + if not current_platform.is_cuda(): + self.skipTest("CUDA-only test") + if torch.cuda.device_count() < _WORLD_SIZE: + self.skipTest(f"needs {_WORLD_SIZE} GPUs") + + proc = subprocess.run( + [ + sys.executable, + "-m", + "torch.distributed.run", + f"--nproc-per-node={_WORLD_SIZE}", + "--master-port=29618", + __file__, + "--worker", + ], + capture_output=True, + text=True, + timeout=600, + ) + print(proc.stdout[-4000:]) + if proc.returncode != 0: + print(proc.stderr[-4000:], file=sys.stderr) + self.assertEqual(proc.returncode, 0, "folded SRT linear output diverged") + self.assertIn("ENCODER_FOLD_SRT_LINEAR_PARITY PASS", proc.stdout) + + +if __name__ == "__main__": + if "--worker" in sys.argv: + raise SystemExit(_worker()) + unittest.main() diff --git a/python/sglang/multimodal_gen/test/unit/test_component_accuracy_weight_transfer.py b/python/sglang/multimodal_gen/test/unit/test_component_accuracy_weight_transfer.py new file mode 100644 index 000000000..81ab4ae07 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_component_accuracy_weight_transfer.py @@ -0,0 +1,59 @@ +import torch +from torch import nn + +from sglang.multimodal_gen.test.single_test_file.component_accuracy.engine import ( + AccuracyEngine, +) + + +class _SourceProjectionSet(nn.Module): + def __init__(self) -> None: + super().__init__() + self.qkv = nn.Linear(2, 6, bias=False) + self.gate_proj = nn.Linear(2, 3, bias=False) + self.up_proj = nn.Linear(2, 3, bias=False) + self.down_proj = nn.Linear(3, 2, bias=False) + + +class _TargetProjectionSet(nn.Module): + def __init__(self) -> None: + super().__init__() + self.qkv_proj = nn.Linear(2, 6, bias=False) + self.gate_up_proj = nn.Linear(2, 8, bias=False) + self.down_proj = nn.Linear(4, 2, bias=False) + + self.qkv_proj.weight.weight_loader = self._load_qkv + self.gate_up_proj.weight.weight_loader = self._load_gate_up + self.down_proj.weight.weight_loader = self._load_down + + @staticmethod + def _load_qkv(param: nn.Parameter, source: torch.Tensor) -> None: + param.data.copy_(source) + + @staticmethod + def _load_gate_up(param: nn.Parameter, source: torch.Tensor, shard_id: int) -> None: + offset = shard_id * 4 + param.data[offset : offset + source.shape[0]].copy_(source) + + @staticmethod + def _load_down(param: nn.Parameter, source: torch.Tensor) -> None: + param.data[:, : source.shape[1]].copy_(source) + + +def test_transfer_weights_uses_loaders_for_fused_aliases_and_padding() -> None: + source = _SourceProjectionSet().to(dtype=torch.bfloat16) + target = _TargetProjectionSet() + with torch.no_grad(): + for index, parameter in enumerate(source.parameters(), start=1): + parameter.fill_(index) + for parameter in target.parameters(): + parameter.zero_() + + AccuracyEngine.transfer_weights(source, target, target_device=torch.device("cpu")) + + torch.testing.assert_close(target.qkv_proj.weight, source.qkv.weight) + torch.testing.assert_close(target.gate_up_proj.weight[:3], source.gate_proj.weight) + torch.testing.assert_close(target.gate_up_proj.weight[4:7], source.up_proj.weight) + assert torch.count_nonzero(target.gate_up_proj.weight[[3, 7]]) == 0 + torch.testing.assert_close(target.down_proj.weight[:, :3], source.down_proj.weight) + assert torch.count_nonzero(target.down_proj.weight[:, 3]) == 0 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 index 0a590cc81..b5d9d20ee 100644 --- a/python/sglang/multimodal_gen/test/unit/test_qwen2_5vl_generation.py +++ b/python/sglang/multimodal_gen/test/unit/test_qwen2_5vl_generation.py @@ -15,6 +15,8 @@ from sglang.multimodal_gen.runtime.models.encoders.qwen2_5vl import ( Qwen2_5_VLAttention, Qwen2_5_VLForConditionalGeneration, _apply_repetition_penalty, + _make_column_linear, + _make_row_linear, _select_next_token, ) from sglang.multimodal_gen.runtime.models.encoders.qwen2_5vl_vision import ( @@ -24,6 +26,17 @@ from sglang.multimodal_gen.runtime.models.encoders.qwen2_5vl_vision import ( _vision_window_index, ) from sglang.multimodal_gen.runtime.pipelines.longcat_image import LongCatImagePipeline +from sglang.srt.layers.linear import ( + ColumnParallelLinear, + ReplicatedLinear, + RowParallelLinear, +) +from sglang.srt.models.qwen2_5_vl import ( + Qwen2_5_VisionPatchEmbed, + Qwen2_5_VisionPatchMerger, + Qwen2_5_VLMLP, +) +from sglang.srt.runtime_context import get_parallel class _StubQwen2_5VL(Qwen2_5_VLForConditionalGeneration): @@ -66,6 +79,75 @@ class _AttentionRecorder(nn.Module): return query +def test_native_vision_reuses_srt_modules(): + config = SimpleNamespace( + hidden_size=16, + intermediate_size=24, + hidden_act="silu", + num_heads=2, + depth=0, + patch_size=2, + temporal_patch_size=1, + in_channels=3, + spatial_merge_size=2, + out_hidden_size=12, + fullatt_block_indexes=[], + window_size=8, + ) + with get_parallel().override(tp_size=1, tp_rank=0): + model = Qwen2_5VLVisionTransformer(config) + mlp = Qwen2_5_VLMLP( + 16, + 24, + fuse_gate_up=False, + ) + fused_mlp = Qwen2_5_VLMLP(16, 24) + + assert isinstance(model.patch_embed, Qwen2_5_VisionPatchEmbed) + assert isinstance(model.merger, Qwen2_5_VisionPatchMerger) + assert not mlp.fuse_gate_up + assert isinstance(mlp.gate_proj, ColumnParallelLinear) + assert isinstance(mlp.up_proj, ColumnParallelLinear) + assert mlp.gate_proj.tp_size == mlp.up_proj.tp_size == 1 + assert isinstance(mlp.down_proj, ReplicatedLinear) + assert isinstance(fused_mlp.down_proj, RowParallelLinear) + assert mlp.act is not None + assert isinstance( + _make_column_linear(16, 24, bias=False, use_tensor_parallel=False), + ReplicatedLinear, + ) + assert isinstance( + _make_row_linear(24, 16, bias=False, use_tensor_parallel=False), + ReplicatedLinear, + ) + + +def test_text_mlp_uses_single_rank_when_intermediate_size_is_not_tp_divisible( + monkeypatch, +): + monkeypatch.setattr(qwen2_5vl, "Qwen2_5_VLAttention", lambda *_args: nn.Identity()) + monkeypatch.setattr(qwen2_5vl, "_tp_world_size", lambda: 3) + monkeypatch.setattr(qwen2_5vl, "_tp_rank", lambda: 2) + config = SimpleNamespace( + hidden_size=16, + intermediate_size=25, + hidden_act="silu", + rms_norm_eps=1e-6, + use_sliding_window=False, + _attn_implementation="flash_attention_2", + layer_types=["full_attention"], + ) + + layer = qwen2_5vl.Qwen2_5_VLDecoderLayer(config, layer_idx=0) + + assert layer.mlp.tp_size == 1 + assert layer.mlp.tp_rank == 0 + assert isinstance(layer.mlp.gate_proj, ColumnParallelLinear) + assert isinstance(layer.mlp.up_proj, ColumnParallelLinear) + assert layer.mlp.gate_proj.tp_rank == layer.mlp.up_proj.tp_rank == 0 + assert isinstance(layer.mlp.down_proj, ReplicatedLinear) + + def test_explicit_attention_mask_is_limited_to_cached_generation(monkeypatch): attention = Qwen2_5_VLAttention.__new__(Qwen2_5_VLAttention) nn.Module.__init__(attention) @@ -76,12 +158,12 @@ def test_explicit_attention_mask_is_limited_to_cached_generation(monkeypatch): attention.num_heads = 1 attention.num_key_value_heads = 1 attention.head_dim = 4 - attention.rope_scaling = {"mrope_section": [1, 1, 0]} + attention.rotary_emb = object() attention.attn = _AttentionRecorder() monkeypatch.setattr( qwen2_5vl, - "apply_multimodal_rotary_pos_emb", - lambda query, key, *_args: (query, key), + "apply_qwen_vl_text_rope", + lambda _rotary_emb, _position_ids, query, key: (query, key), ) hidden_states = torch.randn(1, 2, 4) @@ -89,7 +171,7 @@ def test_explicit_attention_mask_is_limited_to_cached_generation(monkeypatch): kwargs = { "hidden_states": hidden_states, "attention_mask": explicit_mask, - "position_embeddings": (torch.empty(0), torch.empty(0)), + "position_ids": torch.zeros(3, 1, 2, dtype=torch.long), } attention(**kwargs, use_cache=False) diff --git a/python/sglang/multimodal_gen/test/unit/test_qwen3_encoder.py b/python/sglang/multimodal_gen/test/unit/test_qwen3_encoder.py index d7e159338..0cc34e9c2 100644 --- a/python/sglang/multimodal_gen/test/unit/test_qwen3_encoder.py +++ b/python/sglang/multimodal_gen/test/unit/test_qwen3_encoder.py @@ -2,7 +2,10 @@ from types import SimpleNamespace import torch +import sglang.multimodal_gen.runtime.models.encoders.qwen3 as qwen3 +import sglang.srt.layers.activation as srt_activation from sglang.multimodal_gen.runtime.models.encoders.qwen3 import Qwen3ForCausalLM +from sglang.srt.layers.activation import SiluAndMul class _CaptureLayer(torch.nn.Module): @@ -26,6 +29,53 @@ class _IdentityNorm(torch.nn.Module): return hidden_states, None +def test_mlp_reuses_srt_activation_without_server_context(monkeypatch): + def fail_get_exec(): + raise AssertionError("SiluAndMul must not read an unpublished context") + + monkeypatch.setattr(srt_activation, "publish_role", lambda: None) + monkeypatch.setattr(srt_activation, "get_exec", fail_get_exec) + + def make_linear(*_args, **_kwargs): + return torch.nn.Identity() + + monkeypatch.setattr(qwen3, "MergedColumnParallelLinear", make_linear) + monkeypatch.setattr(qwen3, "RowParallelLinear", make_linear) + + mlp = qwen3.Qwen3MLP(16, 24, "silu") + + assert isinstance(mlp.act_fn, SiluAndMul) + + +def test_attention_keeps_diffusion_one_pass_qk_norm(monkeypatch): + monkeypatch.setattr(qwen3, "get_tp_world_size", lambda: 1) + monkeypatch.setattr( + qwen3, "QKVParallelLinear", lambda **kwargs: torch.nn.Identity() + ) + monkeypatch.setattr( + qwen3, "RowParallelLinear", lambda **kwargs: torch.nn.Identity() + ) + monkeypatch.setattr(qwen3, "get_rope", lambda *args, **kwargs: torch.nn.Identity()) + monkeypatch.setattr( + qwen3, "LocalAttention", lambda *args, **kwargs: torch.nn.Identity() + ) + config = SimpleNamespace( + head_dim=128, + rms_norm_eps=1e-6, + _supported_attention_backends=(), + ) + + attention = qwen3.Qwen3Attention( + config, + hidden_size=256, + num_heads=2, + num_kv_heads=1, + ) + + assert isinstance(attention.q_norm, qwen3.MMGenRMSNorm) + assert isinstance(attention.k_norm, qwen3.MMGenRMSNorm) + + def test_default_position_ids_batch_shape(): model = Qwen3ForCausalLM.__new__(Qwen3ForCausalLM) torch.nn.Module.__init__(model) diff --git a/python/sglang/multimodal_gen/test/unit/test_qwen3vl_text.py b/python/sglang/multimodal_gen/test/unit/test_qwen3vl_text.py new file mode 100644 index 000000000..29bf13e1f --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_qwen3vl_text.py @@ -0,0 +1,78 @@ +from types import SimpleNamespace + +import torch +from torch import nn + +import sglang.multimodal_gen.runtime.models.encoders.qwen3vl as qwen3vl + + +class _IdentityAttention(nn.Module): + def forward(self, query, key, value): + return query + + +def test_qwen3vl_attention_uses_interleaved_mrope(monkeypatch): + captured_kwargs = {} + + def build_rope(_config, **kwargs): + captured_kwargs.update(kwargs) + return object() + + monkeypatch.setattr(qwen3vl, "build_qwen_vl_text_rope", build_rope) + monkeypatch.setattr( + qwen3vl, "_make_text_linear", lambda *args, **kwargs: nn.Identity() + ) + monkeypatch.setattr( + qwen3vl, "_make_text_row_linear", lambda *args, **kwargs: nn.Identity() + ) + monkeypatch.setattr( + qwen3vl, "_make_text_rms_norm", lambda *args, **kwargs: nn.Identity() + ) + monkeypatch.setattr(qwen3vl, "LocalAttention", lambda **kwargs: nn.Identity()) + config = SimpleNamespace( + head_dim=8, + hidden_size=8, + num_attention_heads=1, + num_key_value_heads=1, + attention_dropout=0.0, + attention_bias=False, + rms_norm_eps=1e-6, + ) + + qwen3vl.Qwen3VLTextAttention(config, layer_idx=0) + + assert captured_kwargs == {"mrope_interleaved": True} + + +def test_qwen3vl_attention_passes_three_axis_positions_to_srt_rope(monkeypatch): + attention = qwen3vl.Qwen3VLTextAttention.__new__(qwen3vl.Qwen3VLTextAttention) + 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.q_norm = nn.Identity() + attention.k_norm = nn.Identity() + attention.head_dim = 4 + attention.rotary_emb = object() + attention.attn = _IdentityAttention() + + captured_position_ids = None + + def apply_rope(_rotary_emb, position_ids, query, key): + nonlocal captured_position_ids + captured_position_ids = position_ids + return query, key + + monkeypatch.setattr(qwen3vl, "apply_qwen_vl_text_rope", apply_rope) + + hidden_states = torch.randn(1, 2, 4) + position_ids = torch.arange(6).view(3, 1, 2) + output = attention( + hidden_states, + position_ids=position_ids, + attention_mask=None, + ) + + assert captured_position_ids is position_ids + torch.testing.assert_close(output, hidden_states) diff --git a/python/sglang/multimodal_gen/test/unit/test_qwen3vl_vision.py b/python/sglang/multimodal_gen/test/unit/test_qwen3vl_vision.py index 9bdf75d37..cf0fe7e6f 100644 --- a/python/sglang/multimodal_gen/test/unit/test_qwen3vl_vision.py +++ b/python/sglang/multimodal_gen/test/unit/test_qwen3vl_vision.py @@ -9,6 +9,7 @@ from sglang.multimodal_gen.runtime.models.encoders.minimax_h3_qwen3vl import ( ) from sglang.multimodal_gen.runtime.models.encoders.qwen3vl import ( Qwen3VLForConditionalGeneration, + _make_text_rms_norm, ) from sglang.multimodal_gen.runtime.models.encoders.qwen3vl_vision import ( Qwen3VLVisionRotaryEmbedding, @@ -16,6 +17,12 @@ from sglang.multimodal_gen.runtime.models.encoders.qwen3vl_vision import ( _vision_cu_seqlens, _vision_position_ids, ) +from sglang.srt.layers.layernorm import RMSNorm +from sglang.srt.models.qwen3_vl import ( + Qwen3VLMoeVisionPatchMerger, + Qwen3VLVisionPatchEmbed, +) +from sglang.srt.runtime_context import get_parallel def test_native_vision_layout_matches_qwen3_merge_order(): @@ -38,6 +45,13 @@ def test_native_vision_layout_matches_qwen3_merge_order(): assert cu_seqlens.tolist() == [0, 24, 32, 40] +def test_qwen3vl_text_reuses_srt_rms_norm(): + norm = _make_text_rms_norm(16, 1e-6) + + assert isinstance(norm, RMSNorm) + assert norm.cast_x_before_out_mul + + def test_native_vision_keeps_checkpoint_parameter_names(): config = SimpleNamespace( hidden_size=16, @@ -53,7 +67,11 @@ def test_native_vision_keeps_checkpoint_parameter_names(): out_hidden_size=12, deepstack_visual_indexes=[], ) - model = Qwen3VLVisionTransformer(config) + with get_parallel().override(tp_size=1, tp_rank=0): + model = Qwen3VLVisionTransformer(config) + + assert isinstance(model.patch_embed, Qwen3VLVisionPatchEmbed) + assert isinstance(model.merger, Qwen3VLMoeVisionPatchMerger) assert set(model.state_dict()) == { "patch_embed.proj.weight", diff --git a/python/sglang/multimodal_gen/test/unit/test_qwen_vl_rope.py b/python/sglang/multimodal_gen/test/unit/test_qwen_vl_rope.py new file mode 100644 index 000000000..27695ef4e --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_qwen_vl_rope.py @@ -0,0 +1,150 @@ +from types import SimpleNamespace + +import pytest +import torch +from torch import nn + +import sglang.multimodal_gen.runtime.models.encoders.qwen_vl_rope as qwen_vl_rope +import sglang.srt.layers.rotary_embedding.base as rope_base +import sglang.srt.layers.rotary_embedding.factory as rope_factory +from sglang.multimodal_gen.runtime.models.encoders.qwen_vl_rope import ( + apply_qwen_vl_text_rope, + build_qwen_vl_text_rope, +) + + +class _RecordingRotaryEmbedding(nn.Module): + def __init__(self): + super().__init__() + self.positions = None + self.query_shape = None + self.key_shape = None + + def forward_native(self, positions, query, key): + self.positions = positions + self.query_shape = query.shape + self.key_shape = key.shape + return query + 1, key + 2 + + +def test_qwen_vl_rope_supports_transformers_v5_config(monkeypatch): + rope_parameters = { + "rope_type": "default", + "rope_theta": 1_000_000.0, + "mrope_section": [2, 1, 1], + } + config = SimpleNamespace( + head_dim=None, + hidden_size=32, + num_attention_heads=4, + max_position_embeddings=128, + rope_parameters=rope_parameters, + ) + captured_kwargs = {} + rotary_emb = object() + + def get_rope(**kwargs): + captured_kwargs.update(kwargs) + return rotary_emb + + monkeypatch.setattr(qwen_vl_rope, "get_rope", get_rope) + + assert build_qwen_vl_text_rope(config) is rotary_emb + assert captured_kwargs == { + "head_size": 8, + "rotary_dim": 8, + "max_position": 128, + "base": 1_000_000.0, + "is_neox_style": True, + "rope_scaling": {**rope_parameters, "mrope_interleaved": False}, + } + + +def test_qwen_vl_rope_enables_interleaved_layout_explicitly(monkeypatch): + config = SimpleNamespace( + head_dim=8, + max_position_embeddings=128, + rope_parameters={ + "rope_type": "default", + "rope_theta": 1_000_000.0, + "mrope_section": [2, 1, 1], + }, + ) + captured_kwargs = {} + + def get_rope(**kwargs): + captured_kwargs.update(kwargs) + return object() + + monkeypatch.setattr(qwen_vl_rope, "get_rope", get_rope) + + build_qwen_vl_text_rope(config, mrope_interleaved=True) + + assert captured_kwargs["rope_scaling"] == { + **config.rope_parameters, + "mrope_interleaved": True, + } + + +def test_qwen_vl_rope_does_not_require_srt_runtime_context(monkeypatch): + def fail_get_exec(): + raise AssertionError("Qwen-VL RoPE must not read an unpublished context") + + monkeypatch.setattr(rope_base, "get_exec", fail_get_exec) + monkeypatch.setattr(rope_base, "publish_role", lambda: None) + monkeypatch.setattr(rope_factory, "_ROPE_DICT", {}) + config = SimpleNamespace( + head_dim=None, + hidden_size=32, + num_attention_heads=4, + max_position_embeddings=37, + rope_parameters={ + "rope_type": "default", + "rope_theta": 123_457.0, + "mrope_section": [2, 1, 1], + }, + ) + + rotary_emb = build_qwen_vl_text_rope(config) + positions = torch.arange(9).view(3, 3) + query = torch.randn(3, 16) + key = torch.randn(3, 8) + + rotated_query, rotated_key = rotary_emb.forward_native(positions, query, key) + + assert rotated_query.shape == query.shape + assert rotated_key.shape == key.shape + + +def test_qwen_vl_rope_adapts_batched_gqa_layout(): + rotary_emb = _RecordingRotaryEmbedding() + query = torch.randn(2, 4, 5, 8) + key = torch.randn(2, 2, 5, 8) + position_ids = torch.arange(30).view(3, 2, 5) + + rotated_query, rotated_key = apply_qwen_vl_text_rope( + rotary_emb, position_ids, query, key + ) + + assert rotary_emb.query_shape == (10, 32) + assert rotary_emb.key_shape == (10, 16) + assert torch.equal(rotary_emb.positions, position_ids.reshape(3, -1)) + torch.testing.assert_close(rotated_query, query + 1) + torch.testing.assert_close(rotated_key, key + 2) + + +@pytest.mark.parametrize( + ("position_ids", "key"), + [ + (torch.zeros(2, 1, 3, dtype=torch.long), torch.zeros(1, 1, 3, 4)), + (torch.zeros(3, 1, 2, dtype=torch.long), torch.zeros(1, 1, 3, 4)), + ], +) +def test_qwen_vl_rope_rejects_incompatible_shapes(position_ids, key): + with pytest.raises(ValueError): + apply_qwen_vl_text_rope( + _RecordingRotaryEmbedding(), + position_ids, + torch.zeros(1, 1, 3, 4), + key, + ) diff --git a/python/sglang/multimodal_gen/test/unit/test_text_encoder_loader.py b/python/sglang/multimodal_gen/test/unit/test_text_encoder_loader.py index ac8871b78..72f0407d7 100644 --- a/python/sglang/multimodal_gen/test/unit/test_text_encoder_loader.py +++ b/python/sglang/multimodal_gen/test/unit/test_text_encoder_loader.py @@ -108,6 +108,26 @@ class TestMiniMaxH3CheckpointFilter(unittest.TestCase): expected, ) + def test_vision_qkv_checkpoint_name_maps_to_native_projection(self): + encoder = MiniMaxH3Qwen3VLEncoder.__new__(MiniMaxH3Qwen3VLEncoder) + torch.nn.Module.__init__(encoder) + encoder.model = torch.nn.Module() + encoder.model.visual = torch.nn.Module() + block = torch.nn.Module() + block.attn = torch.nn.Module() + block.attn.qkv_proj = torch.nn.Linear(2, 2) + encoder.model.visual.blocks = torch.nn.ModuleList([block]) + + loaded = encoder.load_weights( + [("model.visual.blocks.0.attn.qkv.bias", torch.tensor([1.0, 2.0]))] + ) + + self.assertEqual(loaded, {"model.visual.blocks.0.attn.qkv_proj.bias"}) + torch.testing.assert_close( + encoder.model.visual.blocks[0].attn.qkv_proj.bias, + torch.tensor([1.0, 2.0]), + ) + class TestTextEncoderQuantization(unittest.TestCase): def setUp(self): diff --git a/python/sglang/srt/layers/activation.py b/python/sglang/srt/layers/activation.py index fce55bbb0..8e5d6d048 100644 --- a/python/sglang/srt/layers/activation.py +++ b/python/sglang/srt/layers/activation.py @@ -33,7 +33,7 @@ from sglang.srt.model_executor.cuda_graph_config import ( Phase, check_cuda_graph_backend, ) -from sglang.srt.runtime_context import get_exec, get_parallel +from sglang.srt.runtime_context import get_exec, get_parallel, publish_role from sglang.srt.utils import ( cpu_has_amx_support, get_bool_env_var, @@ -130,7 +130,10 @@ logger = logging.getLogger(__name__) class SiluAndMul(BaseFusedOp): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) - if get_exec().deterministic.rl_on_policy_target is not None: + if ( + publish_role() is not None + and get_exec().deterministic.rl_on_policy_target is not None + ): self._forward_method = self.forward_native elif _use_aiter and envs.SGLANG_OPT_USE_AITER_SILU_MUL.get(): self._forward_method = self.forward_aiter diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py index c2a0ba5f7..6c4f92516 100644 --- a/python/sglang/srt/layers/layernorm.py +++ b/python/sglang/srt/layers/layernorm.py @@ -432,6 +432,7 @@ class RMSNorm(BaseFusedOp): weight_dtype: Optional = None, override_orig_dtype: Optional = None, x_pad_to_multiple: int = 0, + force_native: bool = False, ) -> None: super().__init__() self.has_weight = has_weight @@ -467,6 +468,8 @@ class RMSNorm(BaseFusedOp): except ImportError: self._fused_pad_kernel = None self._forward_method = self.forward_aiter + if force_native: + self._forward_method = self.forward_native def forward_cuda( self, @@ -481,11 +484,6 @@ class RMSNorm(BaseFusedOp): residual = residual + post_residual_addition return x, residual return x - # sgl_kernel rmsnorm requires 2D input; reshape higher-rank tensors - needs_reshape = x.dim() != 2 and residual is None - if needs_reshape: - original_shape = x.shape - x = x.contiguous().reshape(-1, original_shape[-1]) if self.variance_size_override is not None: return self.forward_native(x, residual, post_residual_addition) if is_batch_invariant_mode_enabled(): @@ -495,6 +493,10 @@ class RMSNorm(BaseFusedOp): or get_exec().deterministic.rl_on_policy_target == "fsdp" ): return self.forward_native(x, residual, post_residual_addition) + original_shape = x.shape + needs_reshape = x.dim() != 2 + if needs_reshape: + x = x.contiguous().reshape(-1, original_shape[-1]) out = rms_norm_batch_invariant( x, self.weight.data, @@ -517,6 +519,21 @@ class RMSNorm(BaseFusedOp): return self.forward_with_per_tensor_quant_fusion( x, scale, residual, post_residual_addition ) + + # CUDA RMSNorm kernels require 2D inputs. Flatten token dimensions for + # the kernel call and restore each returned tensor to its input shape. + original_shape = x.shape + residual_shape = residual.shape if residual is not None else original_shape + needs_reshape = x.dim() != 2 + if needs_reshape: + x = x.contiguous().reshape(-1, original_shape[-1]) + if residual is not None: + residual = residual.contiguous().reshape(-1, residual_shape[-1]) + if post_residual_addition is not None: + post_residual_addition = post_residual_addition.contiguous().reshape( + -1, post_residual_addition.shape[-1] + ) + if self.cast_x_before_out_mul and residual is None: # Use HF-semantics kernel (cast to dtype before weight multiply). if ( @@ -531,10 +548,8 @@ class RMSNorm(BaseFusedOp): else: # Fallback: pure-Python HF semantics (already implemented in forward_native). out = self.forward_native(x, None, None) - if needs_reshape: - out = out.reshape(original_shape) - return out - if residual is not None: + result = out + elif residual is not None: if self.cast_x_before_out_mul: if ( x.dtype in (torch.float16, torch.bfloat16) @@ -554,20 +569,28 @@ class RMSNorm(BaseFusedOp): self.variance_epsilon, cast_x_before_out_mul=self.cast_x_before_out_mul, ) - return x, residual - return self.forward_native(x, residual, post_residual_addition) - # TODO: Ideally we want to have (hidden_states+residual)+post_residual_addition. - # but right now we can only have hidden_states+(residual+post_residual_addition). - # (hidden_states+residual)+post_residual_addition != hidden_states+(residual+post_residual_addition), - # we probably need to add another parameter to fused_add_rmsnorm - if post_residual_addition is not None: - residual = residual + post_residual_addition - fused_add_rmsnorm(x, residual, self.weight.data, self.variance_epsilon) - return x, residual - out = rmsnorm(x, self.weight.data, self.variance_epsilon) + result = x, residual + else: + result = self.forward_native(x, residual, post_residual_addition) + else: + # TODO: Ideally we want to have (hidden_states+residual)+post_residual_addition. + # but right now we can only have hidden_states+(residual+post_residual_addition). + # (hidden_states+residual)+post_residual_addition != hidden_states+(residual+post_residual_addition), + # we probably need to add another parameter to fused_add_rmsnorm + if post_residual_addition is not None: + residual = residual + post_residual_addition + fused_add_rmsnorm(x, residual, self.weight.data, self.variance_epsilon) + result = x, residual + else: + result = rmsnorm(x, self.weight.data, self.variance_epsilon) + if needs_reshape: - out = out.reshape(original_shape) - return out + if residual is not None: + return result[0].reshape(original_shape), result[1].reshape( + residual_shape + ) + return result.reshape(original_shape) + return result def forward_npu( self, @@ -602,15 +625,6 @@ class RMSNorm(BaseFusedOp): # AITER's ROCm rmsnorm2d_fwd requires weight/activation dtypes to match; # FP32 weight + BF16 activation yields finite-but-corrupted output on gfx950. return self.forward_native(x, residual, post_residual_addition) - # Aiter's RMSNorm kernels expect 2D contiguous inputs. Keep the - # already-safe layout as a zero-copy path, and only normalize strided or - # higher-rank views such as Q/K slices from packed QKV projections. - needs_reshape = x.dim() != 2 and residual is None - if needs_reshape: - original_shape = x.shape - x = x.contiguous().reshape(-1, original_shape[-1]) - elif not x.is_contiguous(): - x = x.contiguous() if is_batch_invariant_mode_enabled(): if ( residual is not None @@ -619,6 +633,10 @@ class RMSNorm(BaseFusedOp): or (self._fused_pad_kernel is not None and self.x_pad_to_multiple > 0) ): return self.forward_native(x, residual, post_residual_addition) + original_shape = x.shape + needs_reshape = x.dim() != 2 + if needs_reshape: + x = x.contiguous().reshape(-1, original_shape[-1]) out = rms_norm_batch_invariant( x, self.weight.data, @@ -627,6 +645,25 @@ class RMSNorm(BaseFusedOp): if needs_reshape: out = out.reshape(original_shape) return out + + # AITER's RMSNorm kernels require 2D contiguous inputs. + original_shape = x.shape + residual_shape = residual.shape if residual is not None else original_shape + needs_reshape = x.dim() != 2 + if needs_reshape: + x = x.contiguous().reshape(-1, original_shape[-1]) + if residual is not None: + residual = residual.contiguous().reshape(-1, residual_shape[-1]) + if post_residual_addition is not None: + post_residual_addition = post_residual_addition.contiguous().reshape( + -1, post_residual_addition.shape[-1] + ) + else: + if not x.is_contiguous(): + x = x.contiguous() + if residual is not None and not residual.is_contiguous(): + residual = residual.contiguous() + # Fused (add +) rmsnorm + zero-pad path. Triggered when caller # constructed RMSNorm with x_pad_to_multiple > 0. Output last # dim is padded up; residual_out stays at original width. Used @@ -636,13 +673,20 @@ class RMSNorm(BaseFusedOp): if self._fused_pad_kernel is not None and self.x_pad_to_multiple > 0: if post_residual_addition is not None and residual is not None: residual = residual + post_residual_addition - return self._fused_pad_kernel( + result = self._fused_pad_kernel( x, self.weight.data, self.variance_epsilon, residual, self.x_pad_to_multiple, ) + if needs_reshape and residual is not None: + output, residual_out = result + output_shape = (*original_shape[:-1], output.shape[-1]) + return output.reshape(output_shape), residual_out.reshape( + residual_shape + ) + return result if residual is not None: residual_out = torch.empty_like(x) output = torch.empty_like(x) @@ -656,6 +700,10 @@ class RMSNorm(BaseFusedOp): self.weight.data, self.variance_epsilon, ) + if needs_reshape: + return output.reshape(original_shape), residual_out.reshape( + residual_shape + ) return output, residual_out output = rms_norm(x, self.weight.data, self.variance_epsilon) if needs_reshape: diff --git a/python/sglang/srt/layers/rotary_embedding/base.py b/python/sglang/srt/layers/rotary_embedding/base.py index 7f33f3159..c40c0186c 100644 --- a/python/sglang/srt/layers/rotary_embedding/base.py +++ b/python/sglang/srt/layers/rotary_embedding/base.py @@ -11,7 +11,7 @@ from sglang.kernels.fused_op import BaseFusedOp from sglang.srt.environ import envs from sglang.srt.layers.rotary_embedding.utils import apply_rotary_emb from sglang.srt.platforms import current_platform -from sglang.srt.runtime_context import get_exec +from sglang.srt.runtime_context import get_exec, publish_role from sglang.srt.utils import ( cpu_has_amx_support, get_bool_env_var, @@ -94,6 +94,10 @@ class RotaryEmbedding(BaseFusedOp): self.base = base self.is_neox_style = is_neox_style self.dtype = dtype + self._force_native = ( + publish_role() is not None + and get_exec().deterministic.rl_on_policy_target is not None + ) cache = self._compute_cos_sin_cache() # NOTE(ByronHsu): cache needs to be in FP32 for numerical stability. @@ -129,7 +133,7 @@ class RotaryEmbedding(BaseFusedOp): self._apply_rotary_emb_wrapped = apply_rotary_emb # XXX (MUSA): Implement sgl_kernel.rotary_embedding support for MUSA backend - if get_exec().deterministic.rl_on_policy_target is not None or _is_musa: + if self._force_native or _is_musa: self._forward_method = self.forward_native self._apply_rotary_emb_wrapped = torch.compile( dynamic=True, @@ -152,9 +156,7 @@ class RotaryEmbedding(BaseFusedOp): # use CPU to compute the cache and then move it to GPU. However, we # create the cache on GPU for faster initialization. This may cause # a slight numerical difference between the HF implementation and ours. - init_device = ( - "cpu" if get_exec().deterministic.rl_on_policy_target is not None else None - ) + init_device = "cpu" if self._force_native else None inv_freq = 1.0 / ( base ** ( @@ -164,7 +166,7 @@ class RotaryEmbedding(BaseFusedOp): / self.rotary_dim ) ) - if get_exec().deterministic.rl_on_policy_target is not None: + if self._force_native: inv_freq = inv_freq.cuda() return inv_freq diff --git a/python/sglang/srt/layers/rotary_embedding/mrope.py b/python/sglang/srt/layers/rotary_embedding/mrope.py index b8b57bf97..1948ff7b9 100644 --- a/python/sglang/srt/layers/rotary_embedding/mrope.py +++ b/python/sglang/srt/layers/rotary_embedding/mrope.py @@ -18,7 +18,7 @@ from sglang.srt.layers.rotary_embedding.yarn import ( yarn_get_mscale_simple, yarn_linear_ramp_mask, ) -from sglang.srt.runtime_context import attention_backends, get_exec +from sglang.srt.runtime_context import attention_backends from sglang.srt.utils import ( cpu_has_amx_support, is_cuda, @@ -131,7 +131,7 @@ class MRotaryEmbedding(RotaryEmbedding): self.register_buffer("axis_map", axis_map, persistent=False) else: self.axis_map = None - if get_exec().deterministic.rl_on_policy_target is not None: + if self._force_native: self._forward_method = self.forward_native def get_cos_sin_with_position(self, positions): diff --git a/python/sglang/srt/models/qwen2_5_vl.py b/python/sglang/srt/models/qwen2_5_vl.py index 34fca18ce..459489bd9 100644 --- a/python/sglang/srt/models/qwen2_5_vl.py +++ b/python/sglang/srt/models/qwen2_5_vl.py @@ -37,10 +37,6 @@ from transformers.models.qwen2_5_vl.configuration_qwen2_5_vl import ( Qwen2_5_VLConfig, Qwen2_5_VLVisionConfig, ) -from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import ( - Qwen2_5_VisionPatchEmbed, - Qwen2_5_VisionRotaryEmbedding, -) from sglang.srt.distributed.parallel_state import get_pp_group from sglang.srt.environ import envs @@ -50,10 +46,12 @@ from sglang.srt.layers.attention.vision import ( VisionAttentionMetadata, prepare_vision_attention_metadata, ) +from sglang.srt.layers.conv import Conv3dLayer from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( ColumnParallelLinear, MergedColumnParallelLinear, + ReplicatedLinear, RowParallelLinear, ) from sglang.srt.layers.logits_processor import LogitsProcessor @@ -85,6 +83,52 @@ _is_cpu = is_cpu() logger = logging.getLogger(__name__) +class Qwen2_5_VisionPatchEmbed(nn.Module): + def __init__( + self, + patch_size: int, + temporal_patch_size: int, + in_channels: int, + embed_dim: int, + disable_linear: bool = False, + ) -> 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 = Conv3dLayer( + in_channels, + embed_dim, + kernel_size=kernel_size, + stride=kernel_size, + bias=False, + disable_linear=disable_linear, + ) + + 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, + ) + hidden_states = self.proj(hidden_states.to(self.proj.weight.dtype)) + return hidden_states.view(-1, self.embed_dim) + + +class Qwen2_5_VisionRotaryEmbedding(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) + + class Qwen2_5_VLMLP(nn.Module): def __init__( self, @@ -95,31 +139,73 @@ class Qwen2_5_VLMLP(nn.Module): quant_config: Optional[QuantizationConfig] = None, prefix: str = "", use_data_parallel: bool = False, + fuse_gate_up: bool = True, + tp_size: Optional[int] = None, + tp_rank: Optional[int] = None, ): super().__init__() - self.tp_size = 1 if use_data_parallel else get_parallel().tp_size - self.tp_rank = 0 if use_data_parallel else get_parallel().tp_rank - self.gate_up_proj = MergedColumnParallelLinear( - input_size=in_features, - output_sizes=[hidden_features] * 2, # [gate_proj, up_proj] - bias=bias, - quant_config=quant_config, - prefix=add_prefix("gate_up_proj", prefix), - tp_size=self.tp_size, - tp_rank=self.tp_rank, - ) - self.down_proj = RowParallelLinear( - hidden_features, - in_features, - bias=bias, - quant_config=quant_config, - prefix=add_prefix("down_proj", prefix), - tp_size=self.tp_size, - tp_rank=self.tp_rank, - ) + if use_data_parallel: + if tp_size is not None or tp_rank is not None: + raise ValueError( + "Explicit MLP TP cannot be combined with data parallel" + ) + self.tp_size, self.tp_rank = 1, 0 + else: + if (tp_size is None) != (tp_rank is None): + raise ValueError("MLP tp_size and tp_rank must be set together") + self.tp_size = get_parallel().tp_size if tp_size is None else tp_size + self.tp_rank = get_parallel().tp_rank if tp_rank is None else tp_rank + self.fuse_gate_up = fuse_gate_up + if fuse_gate_up: + self.gate_up_proj = MergedColumnParallelLinear( + input_size=in_features, + output_sizes=[hidden_features] * 2, # [gate_proj, up_proj] + bias=bias, + quant_config=quant_config, + prefix=add_prefix("gate_up_proj", prefix), + tp_size=self.tp_size, + tp_rank=self.tp_rank, + ) + else: + projection_kwargs = dict( + input_size=in_features, + output_size=hidden_features, + bias=bias, + quant_config=quant_config, + tp_size=self.tp_size, + tp_rank=self.tp_rank, + ) + self.gate_proj = ColumnParallelLinear( + **projection_kwargs, + prefix=add_prefix("gate_proj", prefix), + ) + self.up_proj = ColumnParallelLinear( + **projection_kwargs, + prefix=add_prefix("up_proj", prefix), + ) + if not self.fuse_gate_up and self.tp_size == 1: + self.down_proj = ReplicatedLinear( + hidden_features, + in_features, + bias=bias, + quant_config=quant_config, + prefix=add_prefix("down_proj", prefix), + ) + else: + self.down_proj = RowParallelLinear( + hidden_features, + in_features, + bias=bias, + quant_config=quant_config, + prefix=add_prefix("down_proj", prefix), + tp_size=self.tp_size, + tp_rank=self.tp_rank, + ) self.hidden_act = hidden_act - if self.hidden_act == "silu": + if self.fuse_gate_up and self.hidden_act == "silu": self.act = SiluAndMul() + elif not self.fuse_gate_up: + self.act = ACT2FN[self.hidden_act] else: base_act = ACT2FN[self.hidden_act] @@ -130,8 +216,13 @@ class Qwen2_5_VLMLP(nn.Module): self.act = _act_fn def forward(self, x: torch.Tensor) -> torch.Tensor: - gate_up, _ = self.gate_up_proj(x) - x = self.act(gate_up) + if self.fuse_gate_up: + gate_up, _ = self.gate_up_proj(x) + x = self.act(gate_up) + else: + gate, _ = self.gate_proj(x) + up, _ = self.up_proj(x) + x = self.act(gate) * up x_down, _ = self.down_proj(x) return x_down @@ -225,11 +316,18 @@ class Qwen2_5_VisionPatchMerger(nn.Module): quant_config: Optional[QuantizationConfig] = None, prefix: str = "", use_data_parallel: bool = False, + cast_x_before_out_mul: bool = False, + force_native_norm: bool = False, ) -> None: super().__init__() self.hidden_size = context_dim * (spatial_merge_size**2) self.padded_context_dim = padded_context_dim * (spatial_merge_size**2) - self.ln_q = RMSNorm(context_dim, eps=1e-6) + self.ln_q = RMSNorm( + context_dim, + eps=1e-6, + cast_x_before_out_mul=cast_x_before_out_mul, + force_native=force_native_norm, + ) tp_size = 1 if use_data_parallel else get_parallel().tp_size tp_rank = 0 if use_data_parallel else get_parallel().tp_rank self.mlp = nn.ModuleList( @@ -257,10 +355,8 @@ class Qwen2_5_VisionPatchMerger(nn.Module): ) def forward(self, x: torch.Tensor) -> torch.Tensor: - # x expected shape: [S, B, context_dim] - S, B, D = x.shape - x2d = x.reshape(-1, D) - x2d = self.ln_q(x2d) # RMSNorm expects 2D + x2d = x.reshape(-1, x.shape[-1]) + x2d = self.ln_q(x2d) x2d = x2d.view(-1, self.hidden_size) # group into spatial_merge_unit mlp_fc1, mlp_act, mlp_fc2 = self.mlp x_parallel, _ = mlp_fc1(x2d) diff --git a/python/sglang/srt/models/qwen3_vl.py b/python/sglang/srt/models/qwen3_vl.py index afd7659f4..176bb40e4 100644 --- a/python/sglang/srt/models/qwen3_vl.py +++ b/python/sglang/srt/models/qwen3_vl.py @@ -102,6 +102,25 @@ _is_cpu = is_cpu() _VECTORIZED_VL_POS_EMBED_MIN_IMAGES = 6 +def _resolve_vision_tp( + *, + use_data_parallel: bool, + tp_size: Optional[int], + tp_rank: Optional[int], +) -> tuple[int, int]: + if use_data_parallel: + if tp_size is not None or tp_rank is not None: + raise ValueError("Explicit vision TP cannot be combined with data parallel") + return 1, 0 + if (tp_size is None) != (tp_rank is None): + raise ValueError("Vision tp_size and tp_rank must be set together") + if tp_size is None: + parallel = get_parallel() + return parallel.attn_tp_size, parallel.attn_tp_rank + assert tp_rank is not None + return tp_size, tp_rank + + class Qwen3_VisionMLP(nn.Module): def __init__( @@ -113,10 +132,15 @@ class Qwen3_VisionMLP(nn.Module): quant_config: Optional[QuantizationConfig] = None, prefix: str = "", use_data_parallel: bool = False, + tp_size: Optional[int] = None, + tp_rank: Optional[int] = None, ): super().__init__() - self.tp_size = 1 if use_data_parallel else get_parallel().attn_tp_size - self.tp_rank = 0 if use_data_parallel else get_parallel().attn_tp_rank + self.tp_size, self.tp_rank = _resolve_vision_tp( + use_data_parallel=use_data_parallel, + tp_size=tp_size, + tp_rank=tp_rank, + ) self.linear_fc1 = ColumnParallelLinear( in_features, hidden_features, @@ -145,7 +169,7 @@ class Qwen3_VisionMLP(nn.Module): class Qwen3VLVisionPatchEmbed(nn.Module): - def __init__(self, config) -> None: + def __init__(self, config, disable_linear: bool = False) -> None: super().__init__() self.patch_size = config.patch_size self.temporal_patch_size = config.temporal_patch_size @@ -159,6 +183,7 @@ class Qwen3VLVisionPatchEmbed(nn.Module): kernel_size=kernel_size, stride=kernel_size, bias=True, + disable_linear=disable_linear, ) def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: @@ -265,6 +290,8 @@ class Qwen3VLMoeVisionPatchMerger(nn.Module): quant_config: Optional[QuantizationConfig] = None, prefix: str = "", use_data_parallel: bool = False, + tp_size: Optional[int] = None, + tp_rank: Optional[int] = None, ) -> None: super().__init__() self.hidden_size = context_dim * (spatial_merge_size**2) @@ -277,8 +304,11 @@ class Qwen3VLMoeVisionPatchMerger(nn.Module): self.norm = norm_layer( self.hidden_size if use_postshuffle_norm else context_dim ) - self.tp_size = 1 if use_data_parallel else get_parallel().attn_tp_size - self.tp_rank = 0 if use_data_parallel else get_parallel().attn_tp_rank + self.tp_size, self.tp_rank = _resolve_vision_tp( + use_data_parallel=use_data_parallel, + tp_size=tp_size, + tp_rank=tp_rank, + ) self.linear_fc1 = ColumnParallelLinear( self.hidden_size, self.padded_context_dim, diff --git a/test/registered/layers/test_layernorm_fusion.py b/test/registered/layers/test_layernorm_fusion.py index 3cf4d581a..0a60e20d4 100644 --- a/test/registered/layers/test_layernorm_fusion.py +++ b/test/registered/layers/test_layernorm_fusion.py @@ -4,10 +4,43 @@ import unittest import torch from sglang.srt.layers.layernorm import RMSNorm -from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci from sglang.test.test_utils import CustomTestCase register_cuda_ci(est_time=15, stage="base-b", runner_config="1-gpu-large") +register_amd_ci(est_time=2, suite="stage-b-test-1-gpu-small-amd") + + +class TestRMSNormInputShape(CustomTestCase): + @classmethod + def setUpClass(cls): + if not torch.cuda.is_available(): + raise unittest.SkipTest("CUDA is not available") + + def test_higher_rank_residual(self): + torch.manual_seed(0) + shape = (2, 3, 512) + + cast_modes = (False,) if torch.version.hip is not None else (False, True) + for cast_x_before_out_mul in cast_modes: + with self.subTest(cast_x_before_out_mul=cast_x_before_out_mul): + layer = RMSNorm( + shape[-1], cast_x_before_out_mul=cast_x_before_out_mul + ).to(device="cuda", dtype=torch.bfloat16) + layer.weight.data.normal_(mean=1.0, std=0.1) + x = torch.randn(shape, device="cuda", dtype=torch.bfloat16) + residual = torch.randn_like(x) + + with torch.inference_mode(): + expected = layer.forward_native(x.clone(), residual.clone()) + actual = layer(x.clone(), residual.clone()) + + self.assertEqual(actual[0].shape, x.shape) + self.assertEqual(actual[1].shape, residual.shape) + torch.testing.assert_close( + actual[0], expected[0], atol=1e-2, rtol=1.5e-2 + ) + torch.testing.assert_close(actual[1], expected[1], atol=1e-2, rtol=1e-2) class TestRMSNormFp8QuantFusion(CustomTestCase):