diff --git a/docker/patch/latest/sglang.patch b/docker/patch/latest/sglang.patch index 89f99ab5a5..8d552fcb2c 100644 --- a/docker/patch/latest/sglang.patch +++ b/docker/patch/latest/sglang.patch @@ -219,7 +219,7 @@ index 1d8baf002..1672de78d 100644 if not hasattr(self, "polling_count"): diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py -index d0d4efd95..a5f06cd67 100644 +index d0d4efd95..e5e5f1b5f 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -260,6 +260,19 @@ class MooncakeKVManager(CommonKVManager): @@ -262,7 +262,7 @@ index d0d4efd95..a5f06cd67 100644 # Reuse _send_kvcache_generic interface to send extra pool data prefill_state_indices = np.array(prefill_state_indices, dtype=np.int32) dst_state_indices = np.array(req.dst_state_indices, dtype=np.int32) -@@ -880,13 +893,36 @@ class MooncakeKVManager(CommonKVManager): +@@ -880,13 +893,43 @@ class MooncakeKVManager(CommonKVManager): if kv_chunk.is_last: if kv_chunk.state_indices is not None: @@ -276,8 +276,15 @@ index d0d4efd95..a5f06cd67 100644 ) + if ret != 0: + with self.session_lock: -+ self.session_failures[req.mooncake_session_id] += 1 -+ if self.session_failures[req.mooncake_session_id] >= 1: ++ self.session_failures[ ++ req.mooncake_session_id ++ ] += 1 ++ if ( ++ self.session_failures[ ++ req.mooncake_session_id ++ ] ++ >= 1 ++ ): + self.failed_sessions.add( + req.mooncake_session_id + ) @@ -300,7 +307,7 @@ index d0d4efd95..a5f06cd67 100644 # Only the last chunk we need to send the aux data ret = self.send_aux( -@@ -895,6 +931,21 @@ class MooncakeKVManager(CommonKVManager): +@@ -895,6 +938,21 @@ class MooncakeKVManager(CommonKVManager): target_rank_registration_info.dst_aux_ptrs, ) polls.append(True if ret == 0 else False) @@ -322,25 +329,26 @@ index d0d4efd95..a5f06cd67 100644 dst_ranks_infos.append( (req.endpoint, req.dst_port, req.room) ) -@@ -977,15 +1028,18 @@ class MooncakeKVManager(CommonKVManager): +@@ -977,15 +1035,20 @@ class MooncakeKVManager(CommonKVManager): if status == KVPoll.Success: if bootstrap_room in self.request_status: - self.prefill_response_tracker[bootstrap_room].add(prefill_rank) -- expected_response_num = ( -- self.required_prefill_response_num_table[bootstrap_room] + # Guard against TOCTOU race: clear() may remove the entry + # between the request_status check and dict access here. -+ expected_response_num = self.required_prefill_response_num_table.get( -+ bootstrap_room - ) + expected_response_num = ( +- self.required_prefill_response_num_table[bootstrap_room] +- ) - arrived_response_num = len( - self.prefill_response_tracker[bootstrap_room] -- ) ++ self.required_prefill_response_num_table.get(bootstrap_room) + ) - if arrived_response_num == expected_response_num: - self.update_status(bootstrap_room, KVPoll.Success) + if expected_response_num is not None: -+ self.prefill_response_tracker[bootstrap_room].add(prefill_rank) ++ self.prefill_response_tracker[bootstrap_room].add( ++ prefill_rank ++ ) + arrived_response_num = len( + self.prefill_response_tracker[bootstrap_room] + ) @@ -350,7 +358,7 @@ index d0d4efd95..a5f06cd67 100644 self.record_failure( bootstrap_room, diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py -index fbc801635..4ea5638cd 100644 +index fbc801635..21fc1ce0d 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -20,6 +20,7 @@ Life cycle of a request in the prefill server @@ -470,15 +478,6 @@ index fbc801635..4ea5638cd 100644 elif poll == KVPoll.Success: # transfer done release_kv_cache(req, self.tree_cache) # unlock the tree req.finished_reason = FINISH_LENGTH(length=0) -@@ -743,7 +812,7 @@ class SchedulerDisaggregationPrefillMixin: - - page_indices = kv_to_page_indices(kv_indices, page_size) - if len(page_indices) == 0: -- logger.info( -+ logger.debug( - f"Skip sending kv chunk for request {req.rid=} {req.bootstrap_room=} because page_indices is empty" - ) - return diff --git a/python/sglang/srt/distributed/parallel_state.py b/python/sglang/srt/distributed/parallel_state.py index 8f1069c00..e47589295 100644 --- a/python/sglang/srt/distributed/parallel_state.py @@ -841,69 +840,6 @@ index 5bf5aa0c8..e52f39fd8 100644 @classmethod def set_metadata(cls, hidden_size: int, dtype: torch.dtype, device: torch.device): -diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py -index 39832c45a..c3d3c6b54 100644 ---- a/python/sglang/srt/layers/layernorm.py -+++ b/python/sglang/srt/layers/layernorm.py -@@ -93,20 +93,17 @@ class RMSNorm(MultiPlatformOp): - eps: float = 1e-6, - var_hidden_size: Optional[int] = None, - cast_x_before_out_mul: bool = False, -- fp32_residual: bool = False, -+ fp32_residual: bool = True, - has_weight: bool = True, -- weight_dtype: Optional = None, -- override_orig_dtype: Optional = None, - ) -> None: - super().__init__() - self.has_weight = has_weight - self.cast_x_before_out_mul = cast_x_before_out_mul - self.fp32_residual = fp32_residual -- self.override_orig_dtype = override_orig_dtype - if self.has_weight: -- self.weight = nn.Parameter(torch.ones(hidden_size, dtype=weight_dtype)) -+ self.weight = nn.Parameter(torch.ones(hidden_size)) - else: -- self.weight = torch.ones(hidden_size, dtype=weight_dtype) -+ self.weight = torch.ones(hidden_size) - self.variance_epsilon = eps - self.hidden_size = hidden_size - self.variance_size_override = ( -@@ -219,16 +216,19 @@ class RMSNorm(MultiPlatformOp): - ) -> 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 -+ orig_dtype = x.dtype -+ -+ if residual is not None and not self.fp32_residual: -+ x = x + residual -+ if post_residual_addition is not None: -+ x = x + post_residual_addition -+ residual = x.clone() - x = x.to(torch.float32) -- if residual is not None: -+ if residual is not None and self.fp32_residual: - x = x + residual.to(torch.float32) - if post_residual_addition is not None: - x = x + post_residual_addition.to(torch.float32) -- if self.fp32_residual: -- residual = x.clone() -- else: -- residual = x.to(orig_dtype) -+ residual = x.to(orig_dtype) - - hidden_size = x.shape[-1] - if hidden_size != self.hidden_size: -@@ -314,7 +314,7 @@ class RMSNorm(MultiPlatformOp): - - if get_tensor_model_parallel_world_size() > 1: - if post_residual_addition is not None: -- residual = residual + post_residual_addition -+ x = x + post_residual_addition - fused_result = flashinfer_allreduce_residual_rmsnorm( - input_tensor=x, - residual=residual, diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py index aff05bf42..130359232 100644 --- a/python/sglang/srt/layers/logits_processor.py @@ -1192,30 +1128,6 @@ index ebcc696ec..3b527021a 100644 def forward_npu( self, dispatch_output: Union[DeepEPNormalDispatchOutput, DeepEPLLDispatchOutput], -diff --git a/python/sglang/srt/layers/moe/fused_moe_triton/fused_moe.py b/python/sglang/srt/layers/moe/fused_moe_triton/fused_moe.py -index ebdbb42c6..714ffbe0e 100644 ---- a/python/sglang/srt/layers/moe/fused_moe_triton/fused_moe.py -+++ b/python/sglang/srt/layers/moe/fused_moe_triton/fused_moe.py -@@ -14,6 +14,7 @@ import torch.nn.functional as F - import triton.language as tl - - from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig -+from sglang.srt.server_args import get_global_server_args - from sglang.srt.utils import ( - cpu_has_amx_support, - get_bool_env_var, -@@ -617,7 +618,10 @@ def fused_experts_impl( - ).squeeze(dim=1) - else: - # According to micro benchmark results, torch.compile can get better performance for small token. -- if tokens_in_chunk <= 32: -+ if ( -+ not get_global_server_args().enable_deterministic_inference -+ and tokens_in_chunk <= 32 -+ ): - moe_sum_reduce_torch_compile( - intermediate_cache3.view(*intermediate_cache3.shape), - out_hidden_states[begin_chunk_idx:end_chunk_idx], diff --git a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py index de8a07ab3..5c9f4813a 100644 --- a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py @@ -1508,21 +1420,20 @@ index 6264f36d0..f0310e305 100644 is_k_full=self.is_k_full, routed_scaling_factor=self.moe_runner_config.routed_scaling_factor, diff --git a/python/sglang/srt/layers/rotary_embedding.py b/python/sglang/srt/layers/rotary_embedding.py -index ae0614635..32171c9c1 100644 +index ae0614635..3b6a8d254 100644 --- a/python/sglang/srt/layers/rotary_embedding.py +++ b/python/sglang/srt/layers/rotary_embedding.py -@@ -150,9 +150,7 @@ class RotaryEmbedding(MultiPlatformOp): - - if get_global_server_args().rl_on_policy_target is not None: - self._forward_method = self.forward_native -- self._apply_rotary_emb_wrapped = torch.compile(dynamic=True)( -- self._apply_rotary_emb_wrapped -- ) -+ - self.position_cos, self.position_sin = None, None - - def _compute_inv_freq(self, base: Union[int, float]) -> torch.Tensor: -@@ -1778,6 +1776,9 @@ class MRotaryEmbedding(RotaryEmbedding): +@@ -305,9 +305,6 @@ class RotaryEmbedding(MultiPlatformOp): + fused_set_kv_buffer_arg: Optional[FusedSetKVBufferArg] = None, + ) -> Tuple[torch.Tensor, torch.Tensor]: + """A PyTorch-npu implementation of forward().""" +- assert ( +- fused_set_kv_buffer_arg is None +- ), "fused_set_kv_buffer_arg is not supported for npu implementation" + if query.dtype == torch.bfloat16 and self.cos_sin_cache.dtype == torch.float: + return self.forward_native(positions, query, key, offsets) + if self.is_neox_style: +@@ -1778,6 +1775,9 @@ class MRotaryEmbedding(RotaryEmbedding): key: torch.Tensor, fused_set_kv_buffer_arg: Optional[FusedSetKVBufferArg] = None, ) -> Tuple[torch.Tensor, torch.Tensor]: @@ -1532,26 +1443,6 @@ index ae0614635..32171c9c1 100644 # TODO: remove this when npu_mrope supports QNumHeads * QHeadSize > 4096 assert ( fused_set_kv_buffer_arg is None -diff --git a/python/sglang/srt/layers/sampler.py b/python/sglang/srt/layers/sampler.py -index f78d83d79..f1d577453 100644 ---- a/python/sglang/srt/layers/sampler.py -+++ b/python/sglang/srt/layers/sampler.py -@@ -122,14 +122,9 @@ class Sampler(nn.Module): - # In RL on-policy mode, we use log_softmax to compute logprobs to match the trainer. - logprobs_via_logsoftmax_kernel = None - if self.rl_on_policy_target is not None: -- # TODO: use more inplace ops to save memory -- logits_div_temperature = ( -- logits.bfloat16().div(sampling_info.temperatures).bfloat16() -- ) - logprobs_via_logsoftmax_kernel = torch.log_softmax( -- logits_div_temperature, dim=-1 -+ logits / sampling_info.temperatures, dim=-1 - ) -- del logits_div_temperature - - if self.use_ascend_backend: - # Ascend backend: sample from logits directly. diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index ff1774567..42d27a82a 100644 --- a/python/sglang/srt/managers/io_struct.py @@ -1605,7 +1496,7 @@ index c07995798..dd8ca7167 100644 break diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py -index a9ff0ac94..ba0a75ee7 100644 +index a9ff0ac94..a50dd5122 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -114,6 +114,7 @@ from sglang.srt.managers.io_struct import ( @@ -1624,14 +1515,6 @@ index a9ff0ac94..ba0a75ee7 100644 (GetWeightsByNameReqInput, self.get_weights_by_name), (ReleaseMemoryOccupationReqInput, self.release_memory_occupation), (ResumeMemoryOccupationReqInput, self.resume_memory_occupation), -@@ -1304,7 +1306,6 @@ class Scheduler( - self.tp_cpu_group, - src=self.tp_group.ranks[0], - ) -- - # Process MM requests under EPD-disaggregation mode - if ( - self.pp_rank == 0 diff --git a/python/sglang/srt/managers/scheduler_metrics_mixin.py b/python/sglang/srt/managers/scheduler_metrics_mixin.py index 30b2732b9..68090b161 100644 --- a/python/sglang/srt/managers/scheduler_metrics_mixin.py @@ -1705,7 +1588,7 @@ index 482bc6ca6..857cfa6a3 100644 return self.send_to_detokenizer.send_output( diff --git a/python/sglang/srt/managers/scheduler_pp_mixin.py b/python/sglang/srt/managers/scheduler_pp_mixin.py -index 1a65a3c3d..245722295 100644 +index 1a65a3c3d..f76606469 100644 --- a/python/sglang/srt/managers/scheduler_pp_mixin.py +++ b/python/sglang/srt/managers/scheduler_pp_mixin.py @@ -20,6 +20,7 @@ from sglang.srt.layers.dp_attention import ( @@ -1764,9 +1647,9 @@ index 1a65a3c3d..245722295 100644 ) + # PP1 (last rank) stores its own batch outputs locally to avoid the + # PP1→PP0→PP1 round-trip that causes a deadlock in disagg prefill. -+ self.last_rank_local_result_queue: deque[Tuple[torch.cuda.Event, PPProxyTensors]] = ( -+ deque() -+ ) ++ self.last_rank_local_result_queue: deque[ ++ Tuple[torch.cuda.Event, PPProxyTensors] ++ ] = deque() self.send_req_work = [] self.send_proxy_work = [] @@ -1836,7 +1719,7 @@ index 1a65a3c3d..245722295 100644 if pp_outputs: with torch.profiler.record_function("send_res_dict_to_next_stage"): send_output_work = self._pp_send_dict_to_next_stage( -@@ -1034,20 +1079,36 @@ class SchedulerPPMixin: +@@ -1034,20 +1079,38 @@ class SchedulerPPMixin: ) if mbs[next_mb_id] is not None: @@ -1856,10 +1739,11 @@ index 1a65a3c3d..245722295 100644 - self.copy_stream.wait_stream(self.default_stream) - batch_result = self._pp_prep_batch_result( - mbs[next_mb_id], mb_metadata[next_mb_id], next_pp_outputs -- ) ++ q_event, next_pp_outputs = ( ++ self.last_rank_local_result_queue.popleft() + ) - d2h_event = torch.cuda.Event() - d2h_event.record(torch.cuda.current_stream()) -+ q_event, next_pp_outputs = self.last_rank_local_result_queue.popleft() + with self.copy_stream_ctx: + torch.cuda.current_stream().wait_event(q_event) + self.copy_stream.wait_stream(self.default_stream) @@ -1886,7 +1770,7 @@ index 1a65a3c3d..245722295 100644 return next_pp_outputs, batch_result, d2h_event, send_output_work -@@ -1085,9 +1146,12 @@ class SchedulerPPMixin: +@@ -1085,9 +1148,12 @@ class SchedulerPPMixin: """ Used by PP, get the required rids with the given poll statuses. """ @@ -1900,6 +1784,19 @@ index 1a65a3c3d..245722295 100644 ) rids: List = [] for poll_statuses in poll_statuses_group: +diff --git a/python/sglang/srt/managers/scheduler_profiler_mixin.py b/python/sglang/srt/managers/scheduler_profiler_mixin.py +index 7d08f12b3..afc045da2 100644 +--- a/python/sglang/srt/managers/scheduler_profiler_mixin.py ++++ b/python/sglang/srt/managers/scheduler_profiler_mixin.py +@@ -347,7 +347,7 @@ class SchedulerProfilerMixin: + if self.profiler_prefill_ct > self.profiler_target_prefill_ct: + if self.profile_in_progress: + self.stop_profile(stage=ForwardMode.EXTEND) +- elif batch.forward_mode.is_decode(): ++ elif batch.forward_mode.is_decode() or batch.forward_mode.is_prebuilt(): + if self.profiler_decode_ct == 0: + if self.profile_in_progress: + # force trace flush diff --git a/python/sglang/srt/managers/scheduler_update_weights_mixin.py b/python/sglang/srt/managers/scheduler_update_weights_mixin.py index 293a84350..244ea4eb1 100644 --- a/python/sglang/srt/managers/scheduler_update_weights_mixin.py @@ -2301,10 +2198,10 @@ index 1d917137c..669e5c518 100644 kv_size_bytes = super().get_kv_size_bytes() for index_k_cache in self.index_k_with_scale_buffer: diff --git a/python/sglang/srt/mem_cache/radix_cache.py b/python/sglang/srt/mem_cache/radix_cache.py -index 42b169728..fbefb0193 100644 +index 42b169728..8e799196a 100644 --- a/python/sglang/srt/mem_cache/radix_cache.py +++ b/python/sglang/srt/mem_cache/radix_cache.py -@@ -495,7 +495,13 @@ class RadixCache(BasePrefixCache): +@@ -495,7 +495,17 @@ class RadixCache(BasePrefixCache): if self.disable: return @@ -2315,11 +2212,15 @@ index 42b169728..fbefb0193 100644 + # req_to_token_pool initialization), leading to spurious tree nodes and memory + # leak when page-aligned token counts happen to cross a page boundary. + kv_committed_len = req.kv_committed_len -+ token_ids = req.fill_ids[:kv_committed_len] if kv_committed_len < len(req.fill_ids) else req.fill_ids ++ token_ids = ( ++ req.fill_ids[:kv_committed_len] ++ if kv_committed_len < len(req.fill_ids) ++ else req.fill_ids ++ ) kv_indices = self.req_to_token_pool.req_to_token[ req.req_pool_idx, : len(token_ids) ] -@@ -619,9 +625,8 @@ class RadixCache(BasePrefixCache): +@@ -619,9 +629,8 @@ class RadixCache(BasePrefixCache): node.lock_ref -= 1 self._update_leaf_status(node) if node.parent is None: @@ -2332,10 +2233,10 @@ index 42b169728..fbefb0193 100644 return delta diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py -index 275775a73..c67f3342e 100644 +index 275775a73..e4e2fdc39 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py -@@ -395,7 +395,10 @@ class ModelRunner(ModelRunnerKVCacheMixin): +@@ -395,7 +395,12 @@ class ModelRunner(ModelRunnerKVCacheMixin): self.forward_stream = torch.get_device_module(self.device).Stream() # CPU offload @@ -2343,11 +2244,13 @@ index 275775a73..c67f3342e 100644 + # For draft worker (e.g., MTP), do not set offloader to avoid overriding + # the main model's offloader. Draft worker uses NoopOffloader instead. + if not is_draft_worker: -+ set_offloader(create_offloader_from_server_args(server_args, dp_rank=dp_rank)) ++ set_offloader( ++ create_offloader_from_server_args(server_args, dp_rank=dp_rank) ++ ) self._weight_checker = WeightChecker(model_runner=self) -@@ -600,7 +603,8 @@ class ModelRunner(ModelRunnerKVCacheMixin): +@@ -600,7 +605,8 @@ class ModelRunner(ModelRunnerKVCacheMixin): ) # Init routed experts capturer @@ -2357,7 +2260,7 @@ index 275775a73..c67f3342e 100644 if self.device == "cuda" or self.device == "musa": self.init_cublas() -@@ -2429,11 +2433,19 @@ class ModelRunner(ModelRunnerKVCacheMixin): +@@ -2429,11 +2435,19 @@ class ModelRunner(ModelRunnerKVCacheMixin): output.expert_distribution_metrics = recorder_outputs.get("metrics") # Copy cached routing experts' buffers back to CPU cache @@ -2382,7 +2285,7 @@ index 275775a73..c67f3342e 100644 if self.eplb_manager is not None: self.eplb_manager.on_forward_pass_end() -@@ -2664,6 +2676,42 @@ class ModelRunner(ModelRunnerKVCacheMixin): +@@ -2664,6 +2678,42 @@ class ModelRunner(ModelRunnerKVCacheMixin): device=self.device, ) @@ -2558,203 +2461,6 @@ index 2cf813bce..1250c49e4 100644 def _canonicalize_weights(config, weights_in: Iterable[Tuple[str, torch.Tensor]]): weights_out_dict = dict(weights_in) -diff --git a/python/sglang/srt/models/qwen2.py b/python/sglang/srt/models/qwen2.py -index a3cfde4d6..18fb07b9d 100644 ---- a/python/sglang/srt/models/qwen2.py -+++ b/python/sglang/srt/models/qwen2.py -@@ -91,9 +91,6 @@ class Qwen2MLP(nn.Module): - self.act_fn = SiluAndMul() - - def forward(self, x): -- if get_global_server_args().rl_on_policy_target is not None: -- x = x.bfloat16() -- - gate_up, _ = self.gate_up_proj(x) - x = self.act_fn(gate_up) - x, _ = self.down_proj(x) -@@ -280,11 +277,6 @@ class Qwen2Model(nn.Module): - quant_config=quant_config, - use_attn_tp_group=is_dp_attention_enabled(), - prefix=add_prefix("embed_tokens", prefix), -- params_dtype=( -- torch.float32 -- if get_global_server_args().rl_on_policy_target is not None -- else None -- ), - ) - else: - self.embed_tokens = PPMissingLayer() -@@ -307,10 +299,8 @@ class Qwen2Model(nn.Module): - if self.pp_group.is_last_rank: - norm_kwargs = ( - dict( -- weight_dtype=torch.float32, - cast_x_before_out_mul=True, -- override_orig_dtype=torch.float32, -- fp32_residual=True, -+ fp32_residual=False, - ) - if get_global_server_args().rl_on_policy_target is not None - else {} -diff --git a/python/sglang/srt/models/qwen2_moe.py b/python/sglang/srt/models/qwen2_moe.py -index bbb883a2d..9bad2d1e0 100644 ---- a/python/sglang/srt/models/qwen2_moe.py -+++ b/python/sglang/srt/models/qwen2_moe.py -@@ -596,7 +596,17 @@ class Qwen2MoeModel(nn.Module): - prefix=add_prefix("layers", prefix), - ) - if self.pp_group.is_last_rank: -- self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) -+ norm_kwargs = ( -+ dict( -+ cast_x_before_out_mul=True, -+ fp32_residual=False, -+ ) -+ if get_global_server_args().rl_on_policy_target is not None -+ else {} -+ ) -+ self.norm = RMSNorm( -+ config.hidden_size, eps=config.rms_norm_eps, **norm_kwargs -+ ) - else: - self.norm = PPMissingLayer(return_tuple=True) - -diff --git a/python/sglang/srt/models/qwen3.py b/python/sglang/srt/models/qwen3.py -index b056317e4..351b18684 100644 ---- a/python/sglang/srt/models/qwen3.py -+++ b/python/sglang/srt/models/qwen3.py -@@ -90,8 +90,8 @@ class Qwen3Attention(nn.Module): - - norm_kwargs = ( - dict( -- weight_dtype=torch.float32, - cast_x_before_out_mul=True, -+ fp32_residual=False, - ) - if get_global_server_args().rl_on_policy_target is not None - else {} -@@ -242,10 +242,8 @@ class Qwen3DecoderLayer(nn.Module): - - norm_kwargs = ( - dict( -- weight_dtype=torch.float32, - cast_x_before_out_mul=True, -- override_orig_dtype=torch.float32, -- fp32_residual=True, -+ fp32_residual=False, - ) - if get_global_server_args().rl_on_policy_target is not None - else {} -diff --git a/python/sglang/srt/models/qwen3_moe.py b/python/sglang/srt/models/qwen3_moe.py -index 3fcf0cfa0..bfef7bcf8 100644 ---- a/python/sglang/srt/models/qwen3_moe.py -+++ b/python/sglang/srt/models/qwen3_moe.py -@@ -22,6 +22,7 @@ import math - from typing import Any, Dict, Iterable, List, Optional, Tuple, TypeVar - - import torch -+import torch.nn.functional as F - from torch import nn - from transformers import PretrainedConfig - -@@ -50,7 +51,7 @@ from sglang.srt.layers.moe import ( - ) - from sglang.srt.layers.moe.ep_moe.layer import get_moe_impl_class - from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE --from sglang.srt.layers.moe.topk import TopK -+from sglang.srt.layers.moe.topk import StandardTopKOutput, TopK - from sglang.srt.layers.moe.utils import ( - RoutingMethodType, - filter_moe_weight_param_global_expert, -@@ -233,6 +234,7 @@ class Qwen3MoeSparseMoeBlock(nn.Module): - use_grouped_topk=False, - layer_id=layer_id, - ) -+ self.top_k = config.num_experts_per_tok - - self.experts = get_moe_impl_class(quant_config)( - num_experts=config.num_experts -@@ -301,7 +303,22 @@ class Qwen3MoeSparseMoeBlock(nn.Module): - - # router_logits: (num_tokens, n_experts) - router_logits, _ = self.gate(hidden_states) -- topk_output = self.topk(hidden_states, router_logits) -+ -+ if get_global_server_args().rl_on_policy_target is not None: -+ routing_weights = F.softmax(router_logits, dim=1, dtype=torch.float) -+ routing_weights, selected_experts = torch.topk( -+ routing_weights, self.top_k, dim=-1 -+ ) -+ routing_weights /= routing_weights.sum(dim=-1, keepdim=True) -+ routing_weights = routing_weights.to(hidden_states.dtype) -+ topk_output = StandardTopKOutput( -+ topk_weights=routing_weights, -+ topk_ids=selected_experts, -+ router_logits=router_logits, -+ ) -+ else: -+ topk_output = self.topk(hidden_states, router_logits) -+ - final_hidden_states = self.experts(hidden_states, topk_output) - if ( - self.tp_size > 1 -@@ -482,13 +499,14 @@ class Qwen3MoeAttention(nn.Module): - ) - self.compatible_with_fused_kv_buffer = ( - False if isinstance(self.rotary_emb, MRotaryEmbedding) else True -- ) -+ ) and (get_global_server_args().rl_on_policy_target is None) - self.compatible_with_fused_qk_norm_rope = ( - not isinstance(self.rotary_emb, MRotaryEmbedding) - ) and self.head_dim in (64, 128, 256) - self.use_fused_qk_norm_rope = ( - get_global_server_args().enable_fused_qk_norm_rope - and self.compatible_with_fused_qk_norm_rope -+ and (get_global_server_args().rl_on_policy_target is None) - ) - self._used_fused_qk_norm_rope_last_call = False - -@@ -501,8 +519,16 @@ class Qwen3MoeAttention(nn.Module): - prefix=add_prefix("attn", prefix), - ) - -- self.q_norm = RMSNorm(self.head_dim, eps=rms_norm_eps) -- self.k_norm = RMSNorm(self.head_dim, eps=rms_norm_eps) -+ norm_kwargs = ( -+ dict( -+ cast_x_before_out_mul=True, -+ fp32_residual=False, -+ ) -+ if get_global_server_args().rl_on_policy_target is not None -+ else {} -+ ) -+ self.q_norm = RMSNorm(self.head_dim, eps=rms_norm_eps, **norm_kwargs) -+ self.k_norm = RMSNorm(self.head_dim, eps=rms_norm_eps, **norm_kwargs) - self.alt_stream = alt_stream - - def op_prepare(self, state): -@@ -743,9 +769,19 @@ class Qwen3MoeDecoderLayer(nn.Module): - quant_config=quant_config, - prefix=add_prefix("mlp", prefix), - ) -- self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) -+ norm_kwargs = ( -+ dict( -+ cast_x_before_out_mul=True, -+ fp32_residual=False, -+ ) -+ if get_global_server_args().rl_on_policy_target is not None -+ else {} -+ ) -+ self.input_layernorm = RMSNorm( -+ config.hidden_size, eps=config.rms_norm_eps, **norm_kwargs -+ ) - self.post_attention_layernorm = RMSNorm( -- config.hidden_size, eps=config.rms_norm_eps -+ config.hidden_size, eps=config.rms_norm_eps, **norm_kwargs - ) - - self.layer_communicator = LayerCommunicator( diff --git a/python/sglang/srt/models/qwen3_vl.py b/python/sglang/srt/models/qwen3_vl.py index d641826e3..3abc39ef3 100644 --- a/python/sglang/srt/models/qwen3_vl.py @@ -2912,9 +2618,18 @@ index 5fe45086c..b283d2e9b 100644 self.req_pool_indices[:raw_bs].copy_(forward_batch.req_pool_indices) diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py -index ac629c7ee..c039d2350 100644 +index ac629c7ee..904f54b4a 100644 --- a/python/sglang/srt/speculative/eagle_info.py +++ b/python/sglang/srt/speculative/eagle_info.py +@@ -337,7 +337,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): + sampling_info.top_ks, self.draft_token_num, dim=0 + ), + ) # (bs * draft_token_num, vocab_size) +- if not torch.all(sampling_info.top_ps == 1.0): ++ if sampling_info.need_top_p_sampling: + target_probs = top_p_renorm_prob( + target_probs, + torch.repeat_interleave( @@ -774,6 +774,10 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): self.topk_index = self.topk_index[: len(new_indices)] self.hidden_states = self.hidden_states[: len(new_indices)] diff --git a/docker/version.txt b/docker/version.txt index dee5f8538b..1d45a3aee8 100644 --- a/docker/version.txt +++ b/docker/version.txt @@ -1 +1 @@ -nightly-dev-20260303a +nightly-dev-20260303b