diff --git a/python/sglang/srt/layers/attention/vision.py b/python/sglang/srt/layers/attention/vision.py index 19584fb1f..001934ed6 100644 --- a/python/sglang/srt/layers/attention/vision.py +++ b/python/sglang/srt/layers/attention/vision.py @@ -565,11 +565,25 @@ class VisionAttention(nn.Module): self.dummy_dim = (num_dummy_heads + num_heads) * self.head_size if self.qk_normalization: + norm_kwargs = ( + dict( + weight_dtype=torch.float32, + cast_x_before_out_mul=True, + ) + if get_global_server_args().rl_on_policy_target is not None + else {} + ) self.q_norm = RMSNorm( - self.dummy_dim, eps=layer_norm_eps, var_hidden_size=embed_dim + self.dummy_dim, + eps=layer_norm_eps, + var_hidden_size=embed_dim, + **norm_kwargs, ) self.k_norm = RMSNorm( - self.dummy_dim, eps=layer_norm_eps, var_hidden_size=embed_dim + self.dummy_dim, + eps=layer_norm_eps, + var_hidden_size=embed_dim, + **norm_kwargs, ) # Select attention backend via a unified method @@ -720,6 +734,15 @@ class VisionAttention(nn.Module): if x.dim() == 2: x = x.unsqueeze(0) assert x.dim() == 3, x.shape + if ( + get_global_server_args().rl_on_policy_target is not None + and position_embeddings is not None + ): + assert isinstance(position_embeddings, tuple), ( + "expected position_embeddings to be a tuple of two tensors,\n" + f"but got {type(position_embeddings)}, change if needed" + ) + position_embeddings = tuple(p.to(x.dtype) for p in position_embeddings) x_shape = x.shape bsz, s, _ = x_shape head = self.num_attention_heads_per_partition diff --git a/python/sglang/srt/layers/communicator.py b/python/sglang/srt/layers/communicator.py index a39dcef47..4ee1865e0 100644 --- a/python/sglang/srt/layers/communicator.py +++ b/python/sglang/srt/layers/communicator.py @@ -363,9 +363,10 @@ class LayerCommunicator: residual: torch.Tensor, forward_batch: ForwardBatch, captured_last_layer_outputs: Optional[List[torch.Tensor]] = None, + **kwargs, ): hidden_states, residual = self.prepare_attn( - hidden_states, residual, forward_batch + hidden_states, residual, forward_batch, **kwargs ) if captured_last_layer_outputs is not None: gathered_last_layer_output = self._communicate_simple_fn( @@ -385,6 +386,7 @@ class LayerCommunicator: residual: torch.Tensor, forward_batch: ForwardBatch, quant_format: str = "", + **kwargs, ): if get_attn_tp_context().input_scattered: hidden_states, residual = self._tp_reduce_scatter( @@ -434,7 +436,7 @@ class LayerCommunicator: ) else: - hidden_states = self.input_layernorm(hidden_states) + hidden_states = self.input_layernorm(hidden_states, **kwargs) else: if _use_aiter and _is_gfx95_supported and ("mxfp4" in quant_format): @@ -466,7 +468,9 @@ class LayerCommunicator: ) else: hidden_states, residual = self.input_layernorm( - hidden_states, residual + hidden_states, + residual, + **kwargs, ) hidden_states = self._communicate_simple_fn( diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py index 8e9d14910..b07164c53 100644 --- a/python/sglang/srt/layers/layernorm.py +++ b/python/sglang/srt/layers/layernorm.py @@ -104,21 +104,30 @@ class RMSNorm(MultiPlatformOp): self, x: torch.Tensor, residual: Optional[torch.Tensor] = None, + **kwargs, ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: if self.variance_size_override is not None: - return self.forward_native(x, residual) + return self.forward_native(x, residual, **kwargs) if is_batch_invariant_mode_enabled(): if ( residual is not None or get_global_server_args().rl_on_policy_target == "fsdp" ): - return self.forward_native(x, residual) + return self.forward_native(x, residual, **kwargs) return rms_norm_batch_invariant( x, self.weight.data, self.variance_epsilon, ) if residual is not None: + # 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 + post_residual_addition = kwargs.get("post_residual_addition") + residual = residual + ( + post_residual_addition if post_residual_addition is not None else 0.0 + ) fused_add_rmsnorm(x, residual, self.weight.data, self.variance_epsilon) return x, residual out = rmsnorm(x, self.weight.data, self.variance_epsilon) @@ -128,6 +137,7 @@ class RMSNorm(MultiPlatformOp): self, x: torch.Tensor, residual: Optional[torch.Tensor] = None, + **kwargs, ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: if residual is not None: out, _, residual_out = torch_npu.npu_add_rms_norm( @@ -140,6 +150,7 @@ class RMSNorm(MultiPlatformOp): self, x: torch.Tensor, residual: Optional[torch.Tensor] = None, + **kwargs, ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: if residual is not None: residual_out = torch.empty_like(x) @@ -159,6 +170,7 @@ class RMSNorm(MultiPlatformOp): self, x: torch.Tensor, residual: Optional[torch.Tensor] = None, + **kwargs, ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: if not x.is_contiguous(): # NOTE: Remove this if aiter kernel supports discontinuous input @@ -178,13 +190,23 @@ class RMSNorm(MultiPlatformOp): self, x: torch.Tensor, residual: Optional[torch.Tensor] = None, + **kwargs, ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: if not x.is_contiguous(): x = x.contiguous() orig_dtype = self.override_orig_dtype or x.dtype + post_residual_addition = kwargs.get("post_residual_addition") x = x.to(torch.float32) if residual is not None: - x = x + residual.to(torch.float32) + x = ( + x + + residual.to(torch.float32) + + ( + post_residual_addition.to(torch.float32) + if post_residual_addition is not None + else 0.0 + ) + ) if self.fp32_residual: residual = x.clone() else: @@ -225,6 +247,7 @@ class RMSNorm(MultiPlatformOp): self, x: torch.Tensor, residual: Optional[torch.Tensor] = None, + **kwargs, ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: if _is_cpu_amx_available: if residual is not None: @@ -236,15 +259,16 @@ class RMSNorm(MultiPlatformOp): x, self.weight.data, self.variance_epsilon ) else: - return self.forward_native(x, residual) + return self.forward_native(x, residual, **kwargs) def forward_xpu( self, x: torch.Tensor, residual: Optional[torch.Tensor] = None, + **kwargs, ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: if self.variance_size_override is not None: - return self.forward_native(x, residual) + return self.forward_native(x, residual, **kwargs) if residual is not None: fused_add_rmsnorm(x, residual, self.weight.data, self.variance_epsilon) return x, residual @@ -300,6 +324,7 @@ class LayerNorm(MultiPlatformOp): def forward_cuda( self, x: torch.Tensor, + **kwargs, ) -> torch.Tensor: if ( _flashinfer_layernorm_available @@ -308,11 +333,12 @@ class LayerNorm(MultiPlatformOp): ): return layernorm(x, self.weight, self.bias, self.variance_epsilon) else: - return self.forward_native(x) + return self.forward_native(x, **kwargs) def forward_native( self, x: torch.Tensor, + **kwargs, ) -> torch.Tensor: weight = self.weight if self.elementwise_affine else None bias = self.bias if self.use_bias else None @@ -329,25 +355,28 @@ class LayerNorm(MultiPlatformOp): def forward_hip( self, x: torch.Tensor, + **kwargs, ) -> torch.Tensor: - return self.forward_native(x) + return self.forward_native(x, **kwargs) def forward_npu( self, x: torch.Tensor, + **kwargs, ) -> torch.Tensor: - return self.forward_native(x) + return self.forward_native(x, **kwargs) def forward_cpu( self, x: torch.Tensor, + **kwargs, ) -> torch.Tensor: if _is_cpu_amx_available: return torch.ops.sgl_kernel.layernorm_cpu( x, self.weight.data, self.variance_epsilon ) else: - return self.forward_native(x) + return self.forward_native(x, **kwargs) class GemmaRMSNorm(MultiPlatformOp): @@ -368,6 +397,7 @@ class GemmaRMSNorm(MultiPlatformOp): self, x: torch.Tensor, residual: Optional[torch.Tensor] = None, + **kwargs, ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: if residual is not None: gemma_fused_add_rmsnorm( @@ -381,6 +411,7 @@ class GemmaRMSNorm(MultiPlatformOp): self, x: torch.Tensor, residual: Optional[torch.Tensor] = None, + **kwargs, ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: orig_dtype = x.dtype if residual is not None: @@ -398,13 +429,15 @@ class GemmaRMSNorm(MultiPlatformOp): self, x: torch.Tensor, residual: Optional[torch.Tensor] = None, + **kwargs, ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: - return self._forward_impl(x, residual) + return self._forward_impl(x, residual, **kwargs) def forward_cpu( self, x: torch.Tensor, residual: Optional[torch.Tensor] = None, + **kwargs, ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: if _is_cpu_amx_available: if residual is not None: @@ -415,12 +448,13 @@ class GemmaRMSNorm(MultiPlatformOp): return torch.ops.sgl_kernel.gemma_rmsnorm_cpu( x, self.weight.data, self.variance_epsilon ) - return self.forward_native(x, residual) + return self.forward_native(x, residual, **kwargs) def forward_npu( self, x: torch.Tensor, residual: Optional[torch.Tensor] = None, + **kwargs, ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: if residual is not None: x = x + residual @@ -433,8 +467,9 @@ class GemmaRMSNorm(MultiPlatformOp): self, x: torch.Tensor, residual: Optional[torch.Tensor] = None, + **kwargs, ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: - return self._forward_impl(x, residual) + return self._forward_impl(x, residual, **kwargs) class Gemma3RMSNorm(MultiPlatformOp): @@ -447,22 +482,22 @@ class Gemma3RMSNorm(MultiPlatformOp): def _norm(self, x): return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) - def forward_native(self, x): + def forward_native(self, x, **kwargs): output = self._norm(x.float()) # Llama does x.to(float16) * w whilst Gemma3 is (x * w).to(float16) # See https://github.com/huggingface/transformers/pull/29402 output = output * (1.0 + self.weight.float()) return output.type_as(x) - def forward_cpu(self, x): + def forward_cpu(self, x, **kwargs): if _is_cpu_amx_available and x.stride(-1) == 1: return torch.ops.sgl_kernel.gemma3_rmsnorm_cpu(x, self.weight, self.eps) - return self.forward_native(x) + return self.forward_native(x, **kwargs) - def forward_cuda(self, x): - return self.forward_native(x) + def forward_cuda(self, x, **kwargs): + return self.forward_native(x, **kwargs) - def forward_npu(self, x): + def forward_npu(self, x, **kwargs): output, _ = torch_npu.npu_gemma_rms_norm(x, self.weight, self.eps) return output diff --git a/python/sglang/srt/layers/rotary_embedding.py b/python/sglang/srt/layers/rotary_embedding.py index 7166ac8a5..56516b41b 100644 --- a/python/sglang/srt/layers/rotary_embedding.py +++ b/python/sglang/srt/layers/rotary_embedding.py @@ -16,7 +16,6 @@ from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( cpu_has_amx_support, get_bool_env_var, - get_compiler_backend, is_cpu, is_cuda, is_hip, @@ -1459,6 +1458,9 @@ class MRotaryEmbedding(RotaryEmbedding): f"Corrected mrope_section: {self.mrope_section} (sum={sum(self.mrope_section)})" ) + if get_global_server_args().rl_on_policy_target is not None: + self._forward_method = self.forward_native + def _match_cos_sin_cache_dtype(self, query: torch.Tensor) -> None: # __setattr__ in nn.Module (called by `self.cos_sin_cache = ...`) # is expensive, so avoid calling it if possible @@ -1468,8 +1470,7 @@ class MRotaryEmbedding(RotaryEmbedding): ): self.cos_sin_cache = self.cos_sin_cache.to(query.device, dtype=query.dtype) - @torch.compile(dynamic=True, backend=get_compiler_backend()) - def _forward_native( + def forward_native( self, positions: torch.Tensor, query: torch.Tensor, @@ -1526,7 +1527,7 @@ class MRotaryEmbedding(RotaryEmbedding): key = torch.cat((key_rot, key_pass), dim=-1).reshape(key_shape) return query, key - def forward( + def forward_cuda( self, positions: torch.Tensor, query: torch.Tensor, @@ -1543,14 +1544,12 @@ class MRotaryEmbedding(RotaryEmbedding): """ assert positions.ndim == 1 or positions.ndim == 2 - if positions.ndim == 2 and self.mrope_section and _is_cuda: - return self._forward_triton(positions, query, key) - elif _is_npu: - return self._forward_npu(positions, query, key) - else: - return self._forward_native(positions, query, key) + # Use Triton kernel for multimodal (2D positions) with mrope + if positions.ndim == 2 and self.mrope_section: + return self.forward_triton(positions, query, key) + return self.forward_native(positions, query, key, fused_set_kv_buffer_arg) - def _forward_triton( + def forward_triton( self, positions: torch.Tensor, query: torch.Tensor, @@ -1571,15 +1570,19 @@ class MRotaryEmbedding(RotaryEmbedding): ) return query, key - def _forward_npu( + def forward_npu( self, positions: torch.Tensor, query: torch.Tensor, key: torch.Tensor, + fused_set_kv_buffer_arg: Optional[FusedSetKVBufferArg] = None, ) -> Tuple[torch.Tensor, torch.Tensor]: # TODO: remove this when npu_mrope supports QNumHeads * QHeadSize > 4096 + assert ( + fused_set_kv_buffer_arg is None + ), "fused_set_kv_buffer_arg is not supported for npu implementation" if query.shape[1] > 4096: - return self._forward_native(positions, query, key) + return self.forward_native(positions, query, key, fused_set_kv_buffer_arg) rotary_mode = "half" if self.is_neox_style: rotary_mode = "half" diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index f43d690a3..b1b6b66ec 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -52,6 +52,7 @@ from sglang.srt.layers.dp_attention import ( set_dp_buffer_len, set_is_extend_in_batch, ) +from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import get_compiler_backend, is_npu, support_triton from sglang.srt.utils.common import ceil_align @@ -690,7 +691,10 @@ class ForwardBatch: mm_input = batch.multimodal_inputs[batch_idx] if self.forward_mode.is_decode(): # 3 * N - if mm_input is None: + if ( + mm_input is None + or get_global_server_args().rl_on_policy_target is not None + ): mrope_positions_list[batch_idx] = torch.full( (3, 1), self.seq_lens_cpu[batch_idx] - 1, @@ -707,7 +711,10 @@ class ForwardBatch: batch.extend_seq_lens[batch_idx], batch.extend_prefix_lens[batch_idx], ) - if mm_input is None: + if ( + mm_input is None + or get_global_server_args().rl_on_policy_target is not None + ): # text only mrope_positions = torch.tensor( [ diff --git a/python/sglang/srt/models/qwen2_moe.py b/python/sglang/srt/models/qwen2_moe.py index c978d2c11..3ad9f6736 100644 --- a/python/sglang/srt/models/qwen2_moe.py +++ b/python/sglang/srt/models/qwen2_moe.py @@ -506,6 +506,7 @@ class Qwen2MoeDecoderLayer(nn.Module): forward_batch: ForwardBatch, residual: Optional[torch.Tensor], captured_last_layer_outputs: Optional[List[torch.Tensor]] = None, + **kwargs, ) -> Tuple[torch.Tensor, torch.Tensor]: hidden_states, residual = ( @@ -514,6 +515,7 @@ class Qwen2MoeDecoderLayer(nn.Module): residual, forward_batch, captured_last_layer_outputs=captured_last_layer_outputs, + **kwargs, ) ) diff --git a/python/sglang/srt/models/qwen3.py b/python/sglang/srt/models/qwen3.py index f41235e21..9220831f6 100644 --- a/python/sglang/srt/models/qwen3.py +++ b/python/sglang/srt/models/qwen3.py @@ -276,10 +276,14 @@ class Qwen3DecoderLayer(nn.Module): hidden_states: torch.Tensor, forward_batch: ForwardBatch, residual: Optional[torch.Tensor], + **kwargs, ) -> Tuple[torch.Tensor, torch.Tensor]: # Self Attention hidden_states, residual = self.layer_communicator.prepare_attn( - hidden_states, residual, forward_batch + hidden_states, + residual, + forward_batch, + **kwargs, ) if hidden_states.shape[0] != 0: hidden_states = self.self_attn( diff --git a/python/sglang/srt/models/qwen3_moe.py b/python/sglang/srt/models/qwen3_moe.py index d88896bab..e11678a9e 100644 --- a/python/sglang/srt/models/qwen3_moe.py +++ b/python/sglang/srt/models/qwen3_moe.py @@ -756,6 +756,7 @@ class Qwen3MoeDecoderLayer(nn.Module): forward_batch: ForwardBatch, residual: Optional[torch.Tensor], captured_last_layer_outputs: Optional[List[torch.Tensor]] = None, + **kwargs, ) -> Tuple[torch.Tensor, torch.Tensor]: hidden_states, residual = ( @@ -764,6 +765,7 @@ class Qwen3MoeDecoderLayer(nn.Module): residual, forward_batch, captured_last_layer_outputs=captured_last_layer_outputs, + **kwargs, ) ) diff --git a/python/sglang/srt/models/qwen3_vl.py b/python/sglang/srt/models/qwen3_vl.py index bdf4c8312..891913078 100644 --- a/python/sglang/srt/models/qwen3_vl.py +++ b/python/sglang/srt/models/qwen3_vl.py @@ -32,13 +32,17 @@ from sglang.srt.distributed import ( from sglang.srt.distributed.parallel_state import get_pp_group from sglang.srt.environ import envs from sglang.srt.layers.attention.vision import VisionAttention +from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear from sglang.srt.layers.logits_processor import LogitsProcessor from sglang.srt.layers.pooler import Pooler, PoolingType from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.rotary_embedding import get_rope from sglang.srt.layers.utils import PPMissingLayer, get_layer_id -from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead +from sglang.srt.layers.vocab_parallel_embedding import ( + ParallelLMHead, + VocabParallelEmbedding, +) from sglang.srt.managers.mm_utils import ( MultiModalityDataPaddingPatternMultimodalTokens, general_mm_embed_routine, @@ -278,6 +282,7 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin): use_data_parallel: bool = False, ) -> None: super().__init__() + self.pp_group = get_pp_group() self.hidden_size = vision_config.hidden_size self.num_heads = vision_config.num_heads self.num_position_embeddings = vision_config.num_position_embeddings @@ -297,7 +302,17 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin): 1 + len(self.deepstack_visual_indexes) ) self.patch_embed = Qwen3VLVisionPatchEmbed(config=vision_config) - self.pos_embed = nn.Embedding(self.num_position_embeddings, self.hidden_size) + if self.pp_group.is_first_rank: + self.pos_embed = VocabParallelEmbedding( + self.num_position_embeddings, + self.hidden_size, + quant_config=quant_config, + enable_tp=not is_dp_attention_enabled(), + prefix=add_prefix("pos_embed", prefix), + ) + else: + self.pos_embed = PPMissingLayer() + norm_layer = partial(nn.LayerNorm, eps=norm_eps) head_dim = self.hidden_size // self.num_heads self.rotary_pos_emb = get_rope( @@ -549,6 +564,18 @@ class Qwen3LLMModel(Qwen3Model): len(config.vision_config.deepstack_visual_indexes) ) + def get_deepstack_embeds( + self, layer_idx: int, input_deepstack_embeds: Optional[torch.Tensor] + ) -> Optional[torch.Tensor]: + """Get deepstack embeddings for a given layer index, or None if not applicable.""" + if ( + input_deepstack_embeds is None + or layer_idx not in self.deepstack_embed_to_decoder_layer + ): + return None + sep = self.hidden_size * layer_idx + return input_deepstack_embeds[:, sep : sep + self.hidden_size] + def forward( self, input_ids: torch.Tensor, @@ -580,20 +607,26 @@ class Qwen3LLMModel(Qwen3Model): hidden_states + residual if residual is not None else hidden_states ) + # SGLang applies residual at the START of the next layer, not at the END like HuggingFace. + # See: https://github.com/huggingface/transformers/blob/v5.0.0rc0/src/transformers/models/qwen3_vl/modeling_qwen3_vl.py#L549 + # To match HF behavior, deepstack must be added AFTER residual: (hidden_states + residual) + deepstack + # The order matters because addition with different tensors is not associative in practice. + # Deepstack for prev_layer is applied at the start of current layer via post_residual_addition. + deepstack_embeds = self.get_deepstack_embeds( + layer_idx - 1, input_deepstack_embeds + ) hidden_states, residual = layer( positions, hidden_states, forward_batch, residual, + post_residual_addition=deepstack_embeds, ) - # process deepstack - if ( - input_deepstack_embeds is not None - and layer_idx in self.deepstack_embed_to_decoder_layer - ): - sep = self.hidden_size * layer_idx - hidden_states += input_deepstack_embeds[:, sep : sep + self.hidden_size] + # Handle deepstack for the last processed layer if it exists. + last_deepstack = self.get_deepstack_embeds( + self.end_layer - 1, input_deepstack_embeds + ) if not self.pp_group.is_last_rank: return PPProxyTensors( @@ -607,7 +640,9 @@ class Qwen3LLMModel(Qwen3Model): if residual is None: hidden_states = self.norm(hidden_states) else: - hidden_states, _ = self.norm(hidden_states, residual) + hidden_states, _ = self.norm( + hidden_states, residual, post_residual_addition=last_deepstack + ) if len(aux_hidden_states) == 0: return hidden_states diff --git a/python/sglang/srt/models/qwen3_vl_moe.py b/python/sglang/srt/models/qwen3_vl_moe.py index b13a221f8..ed64010bf 100644 --- a/python/sglang/srt/models/qwen3_vl_moe.py +++ b/python/sglang/srt/models/qwen3_vl_moe.py @@ -46,10 +46,26 @@ class Qwen3MoeLLMModel(Qwen3MoeModel): ): super().__init__(config=config, quant_config=quant_config, prefix=prefix) self.hidden_size = config.hidden_size + # Currently, we use 3 as len(config.vision_config.deepstack_visual_indexes) is not directly accessible here. + # This approach follows the original implementation. + # TODO: make config of type Qwen3VLMoeConfig, so that we can directly obtain deepstack_visual_indexes. + self.deepstack_embed_to_decoder_layer = range(3) def get_input_embeddings(self) -> nn.Embedding: return self.embed_tokens + def get_deepstack_embeds( + self, layer_idx: int, input_deepstack_embeds: Optional[torch.Tensor] + ) -> Optional[torch.Tensor]: + """Get deepstack embeddings for a given layer index, or None if not applicable.""" + if ( + input_deepstack_embeds is None + or layer_idx not in self.deepstack_embed_to_decoder_layer + ): + return None + sep = self.hidden_size * layer_idx + return input_deepstack_embeds[:, sep : sep + self.hidden_size] + def forward( self, input_ids: torch.Tensor, @@ -80,19 +96,26 @@ class Qwen3MoeLLMModel(Qwen3MoeModel): hidden_states + residual if residual is not None else hidden_states ) + # SGLang applies residual at the START of the next layer, not at the END like HuggingFace. + # See: https://github.com/huggingface/transformers/blob/v5.0.0rc0/src/transformers/models/qwen3_vl/modeling_qwen3_vl.py#L549 + # To match HF behavior, deepstack must be added AFTER residual: (hidden_states + residual) + deepstack + # The order matters because addition with different tensors is not associative in practice. + # Deepstack for prev_layer is applied at the start of current layer via post_residual_addition. + deepstack_embeds = self.get_deepstack_embeds( + layer_idx - 1, input_deepstack_embeds + ) hidden_states, residual = layer( positions, hidden_states, forward_batch, residual, + post_residual_addition=deepstack_embeds, ) - # process deepstack - if input_deepstack_embeds is not None and layer_idx < 3: - sep = self.hidden_size * layer_idx - hidden_states.add_( - input_deepstack_embeds[:, sep : sep + self.hidden_size] - ) + # Handle deepstack for the last processed layer if it exists. + last_deepstack = self.get_deepstack_embeds( + self.end_layer - 1, input_deepstack_embeds + ) if not self.pp_group.is_last_rank: return PPProxyTensors( @@ -106,7 +129,9 @@ class Qwen3MoeLLMModel(Qwen3MoeModel): if residual is None: hidden_states = self.norm(hidden_states) else: - hidden_states, _ = self.norm(hidden_states, residual) + hidden_states, _ = self.norm( + hidden_states, residual, post_residual_addition=last_deepstack + ) if len(aux_hidden_states) == 0: return hidden_states diff --git a/python/sglang/srt/multimodal/processors/base_processor.py b/python/sglang/srt/multimodal/processors/base_processor.py index ca0dc1f57..0fa304431 100644 --- a/python/sglang/srt/multimodal/processors/base_processor.py +++ b/python/sglang/srt/multimodal/processors/base_processor.py @@ -17,6 +17,7 @@ from sglang.srt.managers.schedule_batch import ( MultimodalDataItem, MultimodalInputFormat, ) +from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import envs, is_npu, load_audio, load_image, load_video, logger from sglang.srt.utils.cuda_ipc_transport_utils import ( MM_FEATURE_CACHE_SIZE, @@ -316,7 +317,9 @@ class BaseMultimodalProcessor(ABC): and isinstance(processor.image_processor, BaseImageProcessorFast) and not self.server_args.disable_fast_image_processor ): - if not _is_npu: + if get_global_server_args().rl_on_policy_target is not None: + kwargs["device"] = "cpu" + elif not _is_npu: kwargs["device"] = "cuda" elif processor.__class__.__name__ not in { "Qwen2_5_VLProcessor", diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index ad1bbe816..93609beb2 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -2323,6 +2323,9 @@ class ServerArgs: "Enable deterministic inference because of rl_on_policy_target." ) self.enable_deterministic_inference = True + + # For VLM + os.environ["SGLANG_VLM_CACHE_SIZE_MB"] = "0" # TODO remove this environment variable as a whole os.environ["SGLANG_ENABLE_DETERMINISTIC_INFERENCE"] = "1"