From 88c7ad05d660dc143c78e4293901d424013d5eba Mon Sep 17 00:00:00 2001 From: Copilot Date: Wed, 11 Mar 2026 12:37:03 +0000 Subject: [PATCH] [docker] store v0.5.9 patch --- docker/patch/v0.5.9/megatron.patch | 772 +++++ docker/patch/v0.5.9/sglang.patch | 3046 +++++++++++++++++ ...test_qwen2.5_0.5B_ppo_critic_only_short.py | 1 + 3 files changed, 3819 insertions(+) create mode 100644 docker/patch/v0.5.9/megatron.patch create mode 100644 docker/patch/v0.5.9/sglang.patch diff --git a/docker/patch/v0.5.9/megatron.patch b/docker/patch/v0.5.9/megatron.patch new file mode 100644 index 0000000000..6d2a233949 --- /dev/null +++ b/docker/patch/v0.5.9/megatron.patch @@ -0,0 +1,772 @@ +diff --git a/megatron/core/dist_checkpointing/strategies/common.py b/megatron/core/dist_checkpointing/strategies/common.py +index 41c21d93d..ef80f72d6 100644 +--- a/megatron/core/dist_checkpointing/strategies/common.py ++++ b/megatron/core/dist_checkpointing/strategies/common.py +@@ -86,7 +86,7 @@ class TorchCommonLoadStrategy(LoadCommonStrategy): + msc = MultiStorageClientFeature.import_package() + return msc.torch.load(load_path, map_location='cpu') + else: +- return torch.load(load_path, map_location='cpu') ++ return torch.load(load_path, map_location='cpu', weights_only=False) + except FileNotFoundError as e: + err_msg = f'Common file {load_path} does not exist' + if MultiStorageClientFeature.is_enabled(): +diff --git a/megatron/core/dist_checkpointing/strategies/torch.py b/megatron/core/dist_checkpointing/strategies/torch.py +index 5a1ea308d..aa701237f 100644 +--- a/megatron/core/dist_checkpointing/strategies/torch.py ++++ b/megatron/core/dist_checkpointing/strategies/torch.py +@@ -597,10 +597,12 @@ class MCoreLoadPlanner(DefaultLoadPlanner): + def _validate_global_shapes(self, metadata, sharded_tensors): + for sh_ten in sharded_tensors: + if sh_ten.key not in metadata.state_dict_metadata: +- raise KeyError( +- f"{sh_ten.key} from model not in state dict:" +- f" {sorted(metadata.state_dict_metadata.keys())}" +- ) ++ # raise KeyError( ++ # f"{sh_ten.key} from model not in state dict:" ++ # f" {sorted(metadata.state_dict_metadata.keys())}" ++ # ) ++ print(f"{sh_ten.key} from model not in state dict, will skip") ++ continue + loaded_shape = metadata.state_dict_metadata[sh_ten.key].size + expected_shape = self._expected_shape(sh_ten) + if loaded_shape != expected_shape: +@@ -630,7 +632,7 @@ class MCoreLoadPlanner(DefaultLoadPlanner): + tensor_metadata = self.metadata.state_dict_metadata + metadata_with_sizes = [ + (tensor_metadata[key], tensor_metadata[key].size, sharded_tensor) +- for key, sharded_tensor in self.allow_shape_mismatch_sharded_tensors.items() ++ for key, sharded_tensor in self.allow_shape_mismatch_sharded_tensors.items() if key in tensor_metadata + ] + try: + # Temporarily set sizes to expected shapes +@@ -959,6 +961,7 @@ class TorchDistLoadShardedStrategy(LoadShardedStrategy): + planner=MCoreLoadPlanner( + shapes_validation_sharded_tensors=flexible_shape_sharded_tensors, + allow_shape_mismatch_sharded_tensors=allow_shape_mismatch_sharded_tensors, ++ allow_partial_load=True, + ), + ) + +diff --git a/megatron/core/extensions/transformer_engine.py b/megatron/core/extensions/transformer_engine.py +index acb93ef78..d239db4ab 100644 +--- a/megatron/core/extensions/transformer_engine.py ++++ b/megatron/core/extensions/transformer_engine.py +@@ -408,6 +408,7 @@ class TELinear(te.pytorch.Linear): + ) + + for param in self.parameters(): ++ setattr(param, "parallel_mode", parallel_mode) + if is_expert: + # Reduce the gradient on the expert_data_parallel group for expert linear layers + setattr(param, "allreduce", not self.expert_parallel) +@@ -1161,6 +1162,61 @@ class TEDotProductAttention(te.pytorch.DotProductAttention): + + + if HAVE_TE and is_te_min_version("1.9.0.dev0"): ++ def ceil_div(x: int, y: int) -> int: ++ return (x + y - 1) // y ++ ++ class _FakeInt4QuantizationSTE(torch.autograd.Function): ++ @staticmethod ++ def forward(ctx, x, group_size): ++ m, n = x.shape ++ block_size_m, block_size_n = 1, group_size ++ ++ ++ m_padded = ceil_div(m, block_size_m) * block_size_m ++ n_padded = ceil_div(n, block_size_n) * block_size_n ++ ++ x_padded = torch.zeros( ++ (m_padded, n_padded), ++ dtype=x.dtype, device=x.device ++ ) ++ x_padded[:m, :n] = x ++ ++ x_view = x_padded.view( ++ m_padded // block_size_m, ++ block_size_m, ++ n_padded // block_size_n, ++ block_size_n ++ ) ++ ++ x_max = x_view.abs().float().amax(dim=(1, 3), keepdim=True) ++ q_max = 7 ++ x_scale = x_max / q_max ++ ++ x_scale = x_scale.clamp(min=1e-5) ++ ++ x_div = x_view / x_scale ++ x_round = torch.round(x_div) ++ ++ x_q_clamped = x_round.clamp(-q_max, q_max) ++ ++ x_dequant_view = x_q_clamped * x_scale ++ ++ x_dequant_full = x_dequant_view.view_as(x_padded) ++ x_out = x_dequant_full[:m, :n].contiguous().to(x.dtype) ++ ++ return x_out ++ ++ @staticmethod ++ def backward(ctx, grad_output): ++ return grad_output, None ++ ++ def fake_int4_quantization_ste(x, group_size): ++ x_out = _FakeInt4QuantizationSTE.apply(x, group_size) ++ ++ if hasattr(x, 'main_grad'): ++ x_out.main_grad = x.main_grad ++ ++ return x_out + + class TEGroupedLinear(te.pytorch.GroupedLinear): + """ +@@ -1351,6 +1407,7 @@ if HAVE_TE and is_te_min_version("1.9.0.dev0"): + _is_first_microbatch = ( + None if self.disable_parameter_transpose_cache else self.is_first_microbatch + ) ++ + out = super().forward(x, m_splits, is_first_microbatch=_is_first_microbatch) + self.is_first_microbatch = False + +@@ -1361,6 +1418,20 @@ if HAVE_TE and is_te_min_version("1.9.0.dev0"): + return out + return out, None + ++ def _get_weight_tensors(self): ++ """Get the weight tensors of the module.""" ++ weight_tensors = super()._get_weight_tensors() ++ ++ if os.getenv("OPEN_TRAINING_INT4_FAKE_QAT_FLAG", "0") == "1": ++ group_size = int(os.getenv("OPEN_TRAINING_INT4_GROUP_SIZE", "128")) ++ ++ weight_tensors = [ ++ fake_int4_quantization_ste(w, group_size) ++ for w in weight_tensors ++ ] ++ ++ return weight_tensors ++ + def _encode_extra_state(self, state): + # TE 2.0 changed the format of extra_state to be a byte tensor + if is_te_min_version("2.0.0"): +diff --git a/megatron/core/fusions/fused_mla_yarn_rope_apply.py b/megatron/core/fusions/fused_mla_yarn_rope_apply.py +index 1fd5dcfae..c9aeef1f0 100644 +--- a/megatron/core/fusions/fused_mla_yarn_rope_apply.py ++++ b/megatron/core/fusions/fused_mla_yarn_rope_apply.py +@@ -385,6 +385,7 @@ def rotary_fwd_kv_kernel( + SIN, + emb_dim: tl.constexpr, + k_dim: tl.constexpr, ++ k_dim_ceil: tl.constexpr, + v_dim: tl.constexpr, + head_num: tl.constexpr, + batch_size, +@@ -434,21 +435,27 @@ def rotary_fwd_kv_kernel( + cos_right = tl.load(COS + token_idx * emb_dim + emb_dim // 2 + tl.arange(0, emb_dim // 2)) + sin_right = tl.load(SIN + token_idx * emb_dim + emb_dim // 2 + tl.arange(0, emb_dim // 2)) + +- KV_ptr = KV + pid_m * stride_kv_seq + pid_head * BLOCK_H * stride_kv_nheads +- kv_off = tl.arange(0, BLOCK_H)[:, None] * stride_kv_nheads +- mask = kv_off < head_num * stride_kv_nheads +- k_in_off = kv_off + tl.arange(0, k_dim)[None, :] +- v_in_off = kv_off + k_dim + tl.arange(0, v_dim)[None, :] +- k = tl.load(KV_ptr + k_in_off, mask=mask) +- v = tl.load(KV_ptr + v_in_off, mask=mask) ++ KV_ptr = KV + pid_m * stride_kv_seq # + pid_head * BLOCK_H * stride_kv_nheads ++ ki_range = tl.arange(0, BLOCK_H)[:, None] + pid_head * BLOCK_H ++ kj_range = tl.arange(0, k_dim_ceil)[None, :] ++ mask_k = (ki_range < head_num) & (kj_range < k_dim) ++ mask_v = ki_range < head_num ++ k_off = ki_range * stride_kv_nheads + kj_range ++ if v_dim > 0: ++ v_off = ki_range * stride_kv_nheads + k_dim + tl.arange(0, v_dim)[None, :] ++ v = tl.load(KV_ptr + v_off, mask=mask_v) ++ else: ++ v = tl.zeros((BLOCK_H, 1), dtype=KV.dtype.element_ty) ++ k = tl.load(KV_ptr + k_off, mask=mask_k) + +- K_ptr = O_KEY + pid_m * stride_k_seq + pid_head * BLOCK_H * stride_k_nheads +- V_ptr = O_VALUE + pid_m * stride_v_seq + pid_head * BLOCK_H * stride_v_nheads ++ K_ptr = O_KEY + pid_m * stride_k_seq # + pid_head * BLOCK_H * stride_k_nheads ++ V_ptr = O_VALUE + pid_m * stride_v_seq # + pid_head * BLOCK_H * stride_v_nheads + +- k_out_off = tl.arange(0, BLOCK_H)[:, None] * stride_k_nheads + tl.arange(0, k_dim)[None, :] +- v_out_off = tl.arange(0, BLOCK_H)[:, None] * stride_v_nheads + tl.arange(0, v_dim)[None, :] +- tl.store(K_ptr + k_out_off, k, mask=mask) +- tl.store(V_ptr + v_out_off, v, mask=mask) ++ k_out_off = ki_range * stride_k_nheads + kj_range ++ tl.store(K_ptr + k_out_off, k, mask=mask_k) ++ if v_dim > 0: ++ v_out_off = ki_range * stride_v_nheads + tl.arange(0, v_dim)[None, :] ++ tl.store(V_ptr + v_out_off, v, mask=mask_v) + + EMB = K_POS_EMB + pid_m * stride_emb_seq + # x1 = t[..., 0::2], x2 = t[..., 1::2] +@@ -460,14 +467,16 @@ def rotary_fwd_kv_kernel( + x_left = x_left.expand_dims(0).broadcast_to(BLOCK_H, emb_dim // 2) + x_right = x_right.expand_dims(0).broadcast_to(BLOCK_H, emb_dim // 2) + ++ x_range = tl.arange(0, BLOCK_H)[:, None] + pid_head * BLOCK_H ++ mask_x = x_range < head_num + x_left_off = ( +- tl.arange(0, BLOCK_H)[:, None] * stride_k_nheads ++ x_range * stride_k_nheads + + k_dim + + tl.arange(0, emb_dim // 2)[None, :] + ) + x_right_off = x_left_off + emb_dim // 2 +- tl.store(K_ptr + x_left_off, x_left, mask=mask) +- tl.store(K_ptr + x_right_off, x_right, mask=mask) ++ tl.store(K_ptr + x_left_off, x_left, mask=mask_x) ++ tl.store(K_ptr + x_right_off, x_right, mask=mask_x) + + + @triton.autotune( +@@ -493,6 +502,7 @@ def rotary_bwd_kv_kernel( + SIN, + emb_dim: tl.constexpr, + k_dim: tl.constexpr, ++ k_dim_ceil: tl.constexpr, + v_dim: tl.constexpr, + head_num: tl.constexpr, + batch_size, +@@ -533,27 +543,32 @@ def rotary_bwd_kv_kernel( + else: + token_idx = _get_thd_token_idx(cu_seqlens_kv, pid_m, seq_num, cp_rank, cp_size) + +- dKV_ptr = dKV + pid_m * stride_dkv_seq + pid_head * BLOCK_H * stride_dkv_nheads +- dkv_off = tl.arange(0, BLOCK_H)[:, None] * stride_dkv_nheads +- mask = dkv_off < head_num * stride_dkv_nheads +- dk_out_off = dkv_off + tl.arange(0, k_dim)[None, :] +- dv_out_off = dkv_off + k_dim + tl.arange(0, v_dim)[None, :] +- +- dK_ptr = dK + pid_m * stride_dk_seq + pid_head * BLOCK_H * stride_dk_nheads +- dV_ptr = dV + pid_m * stride_dv_seq + pid_head * BLOCK_H * stride_dv_nheads +- dk_in_off = tl.arange(0, BLOCK_H)[:, None] * stride_dk_nheads + tl.arange(0, k_dim)[None, :] +- dv_in_off = tl.arange(0, BLOCK_H)[:, None] * stride_dv_nheads + tl.arange(0, v_dim)[None, :] +- dk = tl.load(dK_ptr + dk_in_off, mask=mask) +- dv = tl.load(dV_ptr + dv_in_off, mask=mask) +- tl.store(dKV_ptr + dk_out_off, dk, mask=mask) +- tl.store(dKV_ptr + dv_out_off, dv, mask=mask) ++ dKV_ptr = dKV + pid_m * stride_dkv_seq # + pid_head * BLOCK_H * stride_dkv_nheads ++ ki_range = tl.arange(0, BLOCK_H)[:, None] + pid_head * BLOCK_H ++ kj_range = tl.arange(0, k_dim_ceil)[None, :] ++ mask_k = (ki_range < head_num) & (kj_range < k_dim) ++ mask_v = ki_range < head_num ++ dk_out_off = ki_range * stride_dkv_nheads + kj_range ++ ++ dK_ptr = dK + pid_m * stride_dk_seq # + pid_head * BLOCK_H * stride_dk_nheads ++ dV_ptr = dV + pid_m * stride_dv_seq # + pid_head * BLOCK_H * stride_dv_nheads ++ dk_in_off = ki_range * stride_dk_nheads + kj_range ++ ++ dk = tl.load(dK_ptr + dk_in_off, mask=mask_k) ++ tl.store(dKV_ptr + dk_out_off, dk, mask=mask_k) ++ ++ if v_dim > 0: ++ dv_out_off = ki_range * stride_dkv_nheads + k_dim + tl.arange(0, v_dim)[None, :] ++ dv_in_off = ki_range * stride_dv_nheads + tl.arange(0, v_dim)[None, :] ++ dv = tl.load(dV_ptr + dv_in_off, mask=mask_v) ++ tl.store(dKV_ptr + dv_out_off, dv, mask=mask_v) + + if pid_head == 0: + x_left_accum = tl.zeros((BLOCK_H, emb_dim // 2), dtype=tl.float32) + x_right_accum = tl.zeros((BLOCK_H, emb_dim // 2), dtype=tl.float32) + for i in tl.static_range(triton.cdiv(head_num, BLOCK_H)): +- dK_ptr = dK + pid_m * stride_dk_seq + i * BLOCK_H * stride_dk_nheads +- x_off = tl.arange(0, BLOCK_H)[:, None] * stride_dk_nheads + k_dim ++ dK_ptr = dK + pid_m * stride_dk_seq # + i * BLOCK_H * stride_dk_nheads ++ x_off = tl.arange(0, BLOCK_H)[:, None] * stride_dk_nheads + k_dim + i * BLOCK_H * stride_dk_nheads + mask = x_off < head_num * stride_dk_nheads + x_left_off = x_off + tl.arange(0, emb_dim // 2)[None, :] + x_right_off = x_left_off + emb_dim // 2 +@@ -632,6 +647,7 @@ class ApplyMLARotaryEmbKV(torch.autograd.Function): + + o_key = kv.new_empty(total_seqlen, nheads, emb_dim + k_dim) + o_value = kv.new_empty(total_seqlen, nheads, v_dim) ++ k_dim_ceil = triton.next_power_of_2(k_dim) + + grid = lambda META: (total_seqlen, triton.cdiv(nheads, META["BLOCK_H"])) + rotary_fwd_kv_kernel[grid]( +@@ -643,6 +659,7 @@ class ApplyMLARotaryEmbKV(torch.autograd.Function): + sin, + emb_dim, + k_dim, ++ k_dim_ceil, + v_dim, + nheads, + batch_size, +@@ -700,6 +717,7 @@ class ApplyMLARotaryEmbKV(torch.autograd.Function): + + d_kv = dk.new_empty(total_seqlen, nheads, ctx.k_dim + ctx.v_dim) + d_emb = dk.new_empty(total_seqlen, 1, ctx.emb_dim) ++ k_dim_ceil = triton.next_power_of_2(ctx.k_dim) + + grid = lambda META: (total_seqlen, triton.cdiv(nheads, META["BLOCK_H"])) + rotary_bwd_kv_kernel[grid]( +@@ -711,6 +729,7 @@ class ApplyMLARotaryEmbKV(torch.autograd.Function): + sin, + ctx.emb_dim, + ctx.k_dim, ++ k_dim_ceil, + ctx.v_dim, + nheads, + batch_size, +diff --git a/megatron/core/models/common/language_module/language_module.py b/megatron/core/models/common/language_module/language_module.py +index 13d74aa52..060898a7a 100644 +--- a/megatron/core/models/common/language_module/language_module.py ++++ b/megatron/core/models/common/language_module/language_module.py +@@ -184,7 +184,15 @@ class LanguageModule(MegatronModule): + assert ( + column_parallel_linear is not None + ), "column_parallel_linear cannot be None when not using fused linear cross entropy." +- logits, _ = column_parallel_linear(hidden, **col_linear_kwargs) ++ # output ++ output_layer_params = {k: v.detach() for k, v in column_parallel_linear.named_parameters()} ++ output_layer_buffers = dict(column_parallel_linear.named_buffers()) ++ logits, _ = torch.func.functional_call( ++ column_parallel_linear, ++ {**output_layer_params, **output_layer_buffers}, ++ (hidden,), ++ col_linear_kwargs, ++ ) + + return self.compute_language_model_loss(labels, logits) + +diff --git a/megatron/core/models/gpt/gpt_layer_specs.py b/megatron/core/models/gpt/gpt_layer_specs.py +index e21127b87..712793853 100755 +--- a/megatron/core/models/gpt/gpt_layer_specs.py ++++ b/megatron/core/models/gpt/gpt_layer_specs.py +@@ -188,6 +188,8 @@ def get_gpt_layer_with_transformer_engine_spec( + use_kitchen: bool = False, + use_te_activation_func: bool = False, + fallback_to_eager_attn: bool = False, ++ post_self_attn_layernorm: bool = False, ++ post_mlp_layernorm: bool = False, + ) -> ModuleSpec: + """Use this spec to use lower-level Transformer Engine modules (required for fp8 training). + +@@ -260,6 +262,8 @@ def get_gpt_layer_with_transformer_engine_spec( + mlp=mlp, + sharded_state_dict_keys_map=sharded_state_dict_keys_map, + normalization=normalization, ++ post_self_attn_layernorm=post_self_attn_layernorm, ++ post_mlp_layernorm=post_mlp_layernorm, + ) + + +@@ -349,6 +353,8 @@ def get_transformer_layer_spec_for_backend( + mlp: ModuleSpec, + sharded_state_dict_keys_map: Optional[dict] = None, + normalization: Optional[str] = None, ++ post_self_attn_layernorm: bool = False, ++ post_mlp_layernorm: bool = False, + ) -> ModuleSpec: + """Helper function to get module spec for TransformerLayer""" + +@@ -371,9 +377,11 @@ def get_transformer_layer_spec_for_backend( + input_layernorm=input_layernorm, + self_attention=attention, + self_attn_bda=get_bias_dropout_add, ++ post_self_attn_layernorm=TENorm if post_self_attn_layernorm else IdentityOp, + pre_mlp_layernorm=pre_mlp_layernorm, + mlp=mlp, + mlp_bda=get_bias_dropout_add, ++ post_mlp_layernorm=TENorm if post_mlp_layernorm else IdentityOp, + sharded_state_dict_keys_map=sharded_state_dict_keys_map, + ), + ) +diff --git a/megatron/core/models/gpt/gpt_model.py b/megatron/core/models/gpt/gpt_model.py +index a1230568c..1fd52f65a 100644 +--- a/megatron/core/models/gpt/gpt_model.py ++++ b/megatron/core/models/gpt/gpt_model.py +@@ -446,6 +446,7 @@ class GPTModel(LanguageModule): + *, + inference_params: Optional[BaseInferenceContext] = None, + loss_mask: Optional[Tensor] = None, ++ mtp_kwargs: Optional[dict] = {}, + ) -> Tensor: + """Forward function of the GPT Model This function passes the input tensors + through the embedding layer, and then the decoder and finally into the post +@@ -508,6 +509,7 @@ class GPTModel(LanguageModule): + runtime_gather_output=runtime_gather_output, + extra_block_kwargs=extra_block_kwargs, + inference_context=inference_context, ++ mtp_kwargs=mtp_kwargs, + ) + + def _postprocess( +@@ -529,6 +531,7 @@ class GPTModel(LanguageModule): + runtime_gather_output=None, + extra_block_kwargs=None, + inference_context=None, ++ mtp_kwargs={}, + ): + """Postprocesses decoder hidden states to generate logits or compute loss. + +@@ -543,7 +546,8 @@ class GPTModel(LanguageModule): + output_weight = None + if self.share_embeddings_and_output_weights: + output_weight = self.shared_embedding_or_output_weight() +- if mtp_in_postprocess: ++ ++ if mtp_in_postprocess and mtp_kwargs.get('mtp_labels', None) is not None: + hidden_states = self.mtp( + input_ids=input_ids, + position_ids=position_ids, +@@ -563,13 +567,18 @@ class GPTModel(LanguageModule): + return hidden_states + + # Skip when mtp_num_layers is None or 0 +- if self.config.mtp_num_layers: +- mtp_labels = labels.clone() ++ if self.config.mtp_num_layers and mtp_kwargs.get('mtp_labels', None) is not None: ++ mtp_labels = mtp_kwargs['mtp_labels'].clone() ++ mtp_labels, _ = roll_tensor(mtp_labels, shifts=-1, dims=-1, cp_group=self.cp_group, packed_seq_params=packed_seq_params) ++ + hidden_states_list = torch.chunk(hidden_states, 1 + self.config.mtp_num_layers, dim=0) + hidden_states = hidden_states_list[0] + if loss_mask is None: + # if loss_mask is not provided, use all ones as loss_mask + loss_mask = torch.ones_like(mtp_labels) ++ else: ++ # Otherwise, roll the loss_mask to keep up with the mtp_labels ++ loss_mask, _ = roll_tensor(loss_mask, shifts=-1, dims=-1, cp_group=self.cp_group, packed_seq_params=packed_seq_params) + for mtp_layer_number in range(self.config.mtp_num_layers): + # Calc loss for the current Multi-Token Prediction (MTP) layers. + mtp_labels, _ = roll_tensor( +@@ -595,7 +604,7 @@ class GPTModel(LanguageModule): + sequence_parallel_enabled=self.output_layer.sequence_parallel, + column_parallel_linear=self.output_layer, + col_linear_kwargs={ +- 'weight': output_weight, ++ 'weight': output_weight.detach() if output_weight else None, + 'runtime_gather_output': runtime_gather_output, + }, + ) +diff --git a/megatron/core/optimizer/distrib_optimizer.py b/megatron/core/optimizer/distrib_optimizer.py +index 6e093f96f..eac21a3ea 100644 +--- a/megatron/core/optimizer/distrib_optimizer.py ++++ b/megatron/core/optimizer/distrib_optimizer.py +@@ -677,6 +677,8 @@ class DistributedOptimizer(MixedPrecisionOptimizer): + # TE FusedAdam will not accumulate step for empty param groups, so we need to + # align the step across param groups. + param_group["step"] = int(step) ++ if "step" in param_group and param_group["step"] is None: ++ del param_group["step"] + + # Grad scaler state. + if self.grad_scaler: +@@ -1646,6 +1648,8 @@ class DistributedOptimizer(MixedPrecisionOptimizer): + if key == 'padding': + tensors[key] = LocalNonpersistentObject(tensors[key]) + continue ++ if key == 'step': ++ continue + assert tensors[key].shape == (gbuf_local_end - gbuf_local_start,), ( + tensors[key].shape, + gbuf_local_start, +diff --git a/megatron/core/parallel_state.py b/megatron/core/parallel_state.py +index a273002b9..4f821cfd5 100644 +--- a/megatron/core/parallel_state.py ++++ b/megatron/core/parallel_state.py +@@ -11,6 +11,7 @@ from typing import Callable, List, Optional + + import numpy as np + import torch ++import torch.distributed as dist + + from .utils import GlobalMemoryBuffer, is_torch_min_version + +diff --git a/megatron/core/pipeline_parallel/p2p_communication.py b/megatron/core/pipeline_parallel/p2p_communication.py +index ac839c21f..f18309217 100644 +--- a/megatron/core/pipeline_parallel/p2p_communication.py ++++ b/megatron/core/pipeline_parallel/p2p_communication.py +@@ -26,22 +26,22 @@ def _batched_p2p_ops( + ops = [] + if tensor_send_prev is not None: + send_prev_op = torch.distributed.P2POp( +- torch.distributed.isend, tensor_send_prev, prev_pipeline_rank, group ++ torch.distributed.isend, tensor_send_prev, prev_pipeline_rank, + ) + ops.append(send_prev_op) + if tensor_recv_prev is not None: + recv_prev_op = torch.distributed.P2POp( +- torch.distributed.irecv, tensor_recv_prev, prev_pipeline_rank, group ++ torch.distributed.irecv, tensor_recv_prev, prev_pipeline_rank, + ) + ops.append(recv_prev_op) + if tensor_send_next is not None: + send_next_op = torch.distributed.P2POp( +- torch.distributed.isend, tensor_send_next, next_pipeline_rank, group ++ torch.distributed.isend, tensor_send_next, next_pipeline_rank, + ) + ops.append(send_next_op) + if tensor_recv_next is not None: + recv_next_op = torch.distributed.P2POp( +- torch.distributed.irecv, tensor_recv_next, next_pipeline_rank, group ++ torch.distributed.irecv, tensor_recv_next, next_pipeline_rank, + ) + ops.append(recv_next_op) + if len(ops) > 0: +diff --git a/megatron/core/transformer/moe/moe_utils.py b/megatron/core/transformer/moe/moe_utils.py +index 28cff06f5..58dc4bb70 100644 +--- a/megatron/core/transformer/moe/moe_utils.py ++++ b/megatron/core/transformer/moe/moe_utils.py +@@ -587,6 +587,9 @@ def topk_routing_with_score_function( + else: + return torch.topk(scores, k=topk, dim=1) + ++ from slime.utils.routing_replay import get_routing_replay_compute_topk ++ compute_topk = get_routing_replay_compute_topk(compute_topk) ++ + if score_function == "softmax": + if use_pre_softmax: + scores = torch.softmax(logits, dim=-1, dtype=torch.float32).type_as(logits) +diff --git a/megatron/core/transformer/moe/router.py b/megatron/core/transformer/moe/router.py +index 16fc9d9af..517944f25 100644 +--- a/megatron/core/transformer/moe/router.py ++++ b/megatron/core/transformer/moe/router.py +@@ -201,6 +201,9 @@ class TopKRouter(Router): + self.global_tokens_per_expert = None + self.ga_steps = None + ++ from slime.utils.routing_replay import register_routing_replay ++ register_routing_replay(self) ++ + def _maintain_float32_expert_bias(self): + """ + Maintain the expert bias in float32. +diff --git a/megatron/core/transformer/multi_token_prediction.py b/megatron/core/transformer/multi_token_prediction.py +index a8f4abfcd..f33f6f05e 100755 +--- a/megatron/core/transformer/multi_token_prediction.py ++++ b/megatron/core/transformer/multi_token_prediction.py +@@ -6,6 +6,7 @@ from typing import Callable, List, Optional, Union + + import torch + from torch import Tensor ++import warnings + + from megatron.core import InferenceParams, parallel_state, tensor_parallel + from megatron.core.dist_checkpointing.mapping import ShardedStateDict +@@ -714,17 +715,19 @@ class MultiTokenPredictionLayer(MegatronModule): + cp_group=self.cp_group, + packed_seq_params=packed_seq_params, + ) +- position_ids, _ = roll_tensor( +- position_ids, +- shifts=-1, +- dims=-1, +- cp_group=self.cp_group, +- packed_seq_params=packed_seq_params, +- ) ++ if position_ids is not None: ++ position_ids, _ = roll_tensor( ++ position_ids, ++ shifts=-1, ++ dims=-1, ++ cp_group=self.cp_group, ++ packed_seq_params=packed_seq_params, ++ ) + # embedding + decoder_input = embedding(input_ids=input_ids, position_ids=position_ids) ++ decoder_input = decoder_input.detach() + +- hidden_states = make_viewless_tensor(inp=hidden_states, requires_grad=True, keep_graph=True) ++ hidden_states = make_viewless_tensor(inp=hidden_states, requires_grad=True, keep_graph=False) + + return input_ids, position_ids, decoder_input, hidden_states + +@@ -826,6 +829,51 @@ class MultiTokenPredictionLayer(MegatronModule): + return hidden_states + + def _checkpointed_forward(self, forward_func, *args, **kwargs): ++ """Wrap `forward_func` with activation checkpointing while only passing tensors. ++ ++ Non-tensor arguments (e.g., configuration objects, None) are captured via closure so ++ that checkpoint implementations never receive them directly, avoiding save_for_backward ++ issues with non-tensor inputs. ++ """ ++ ++ # TODO(jiajun): Is there any better implementation here? ++ positional_specs = [] ++ kw_specs = [] ++ tensor_args: List[torch.Tensor] = [] ++ ++ for arg in args: ++ if torch.is_tensor(arg): ++ positional_specs.append(('tensor', len(tensor_args))) ++ tensor_args.append(arg) ++ else: ++ positional_specs.append(('const', arg)) ++ ++ for key, value in kwargs.items(): ++ if torch.is_tensor(value): ++ kw_specs.append((key, ('tensor', len(tensor_args)))) ++ tensor_args.append(value) ++ else: ++ kw_specs.append((key, ('const', value))) ++ ++ def run(*flat_tensor_args): ++ rebuilt_args = [] ++ for spec_type, payload in positional_specs: ++ if spec_type == 'tensor': ++ rebuilt_args.append(flat_tensor_args[payload]) ++ else: ++ rebuilt_args.append(payload) ++ ++ rebuilt_kwargs = {} ++ for key, (spec_type, payload) in kw_specs: ++ if spec_type == 'tensor': ++ rebuilt_kwargs[key] = flat_tensor_args[payload] ++ else: ++ rebuilt_kwargs[key] = payload ++ ++ return forward_func(*rebuilt_args, **rebuilt_kwargs) ++ ++ tensor_args_tuple = tuple(tensor_args) ++ + def checkpoint_handler(): + """Determines whether to use the `te_checkpoint` or `tensor_parallel.checkpoint`""" + if self.config.fp8: +@@ -836,12 +884,11 @@ class MultiTokenPredictionLayer(MegatronModule): + self.config.distribute_saved_activations, + tensor_parallel.random.get_cuda_rng_tracker, + parallel_state.get_tensor_model_parallel_group(), +- *args, +- **kwargs, ++ *tensor_args_tuple, + ) + else: + return tensor_parallel.checkpoint( +- forward_func, self.config.distribute_saved_activations, *args, *kwargs.values() ++ run, self.config.distribute_saved_activations, *tensor_args_tuple + ) + + if self.config.recompute_method == 'uniform': +diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py +index e2705bd9f..a0aa109b5 100644 +--- a/megatron/core/transformer/transformer_config.py ++++ b/megatron/core/transformer/transformer_config.py +@@ -210,6 +210,9 @@ class TransformerConfig(ModelParallelConfig): + attention_output_gate: bool = False + """Whether to apply output gate to the attention layers.""" + ++ post_self_attn_layernorm: bool = False ++ post_mlp_layernorm: bool = False ++ + test_mode: bool = False + """Whether to run real-time tests.""" + +diff --git a/megatron/core/transformer/transformer_layer.py b/megatron/core/transformer/transformer_layer.py +index 3ea405770..5a42001b9 100644 +--- a/megatron/core/transformer/transformer_layer.py ++++ b/megatron/core/transformer/transformer_layer.py +@@ -223,6 +223,7 @@ class TransformerLayerSubmodules: + input_layernorm: Union[ModuleSpec, type] = IdentityOp + self_attention: Union[ModuleSpec, type] = IdentityOp + self_attn_bda: Union[ModuleSpec, type] = IdentityFuncOp ++ post_self_attn_layernorm: Union[ModuleSpec, type] = IdentityOp + + pre_cross_attn_layernorm: Union[ModuleSpec, type] = IdentityOp + cross_attention: Union[ModuleSpec, type] = IdentityOp +@@ -231,6 +232,7 @@ class TransformerLayerSubmodules: + pre_mlp_layernorm: Union[ModuleSpec, type] = IdentityOp + mlp: Union[ModuleSpec, type] = IdentityOp + mlp_bda: Union[ModuleSpec, type] = IdentityFuncOp ++ post_mlp_layernorm: Union[ModuleSpec, type] = IdentityOp + + # Mapping for sharded tensor keys to be applied in `sharded_state_dict` method + sharded_state_dict_keys_map: Dict[str, str] = field(default_factory=dict) +@@ -310,6 +312,13 @@ class TransformerLayer(GraphableMegatronModule, BaseTransformerLayer): + # [Module 3: BiasDropoutFusion] + self.self_attn_bda = build_module(submodules.self_attn_bda) + ++ self.post_self_attn_layernorm = build_module( ++ submodules.post_self_attn_layernorm, ++ config=self.config, ++ hidden_size=self.config.hidden_size, ++ eps=self.config.layernorm_epsilon, ++ ) ++ + # [Module 4: Post SelfAttention] Optional Layernorm after self-attn + self.pre_cross_attn_layernorm = build_module( + submodules.pre_cross_attn_layernorm, +@@ -375,6 +384,13 @@ class TransformerLayer(GraphableMegatronModule, BaseTransformerLayer): + + self.is_moe_layer = isinstance(self.mlp, MoELayer) + ++ self.post_mlp_layernorm = build_module( ++ submodules.post_mlp_layernorm, ++ config=self.config, ++ hidden_size=self.config.hidden_size, ++ eps=self.config.layernorm_epsilon ++ ) ++ + self.recompute_input_layernorm = False + self.recompute_pre_mlp_layernorm = False + self.recompute_mlp = False +@@ -551,6 +567,10 @@ class TransformerLayer(GraphableMegatronModule, BaseTransformerLayer): + attention_output_with_bias[0] + ) + ++ attention_output, attention_output_bias = attention_output_with_bias ++ attention_output = self.post_self_attn_layernorm(attention_output) ++ attention_output_with_bias = (attention_output, attention_output_bias) ++ + # TODO: could we move `bias_dropout_add_exec_handler` itself + # inside the module provided in the `bias_dropout_add_spec` module? + nvtx_range_push(suffix="self_attn_bda") +@@ -677,6 +697,10 @@ class TransformerLayer(GraphableMegatronModule, BaseTransformerLayer): + else: + mlp_output_with_bias = self.mlp(pre_mlp_layernorm_output) + ++ mlp_output, mlp_output_bias = mlp_output_with_bias ++ mlp_output = self.post_mlp_layernorm(mlp_output) ++ mlp_output_with_bias = (mlp_output, mlp_output_bias) ++ + if self.recompute_pre_mlp_layernorm: + # discard the output of the pre-mlp layernorm and register the recompute + # as a gradient hook of mlp_output_with_bias[0] +diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py +index b267c8a81..83736acdc 100644 +--- a/megatron/training/arguments.py ++++ b/megatron/training/arguments.py +@@ -1398,6 +1398,9 @@ def core_transformer_config_from_args(args, config_class=None): + + kw_args['inference_sampling_seed'] = args.seed + ++ kw_args['post_self_attn_layernorm'] = args.post_self_attn_layernorm ++ kw_args['post_mlp_layernorm'] = args.post_mlp_layernorm ++ + # handle quantization config + # NOTE: Kitchen arguments are only added to the namespace when + # Kitchen library is available. +@@ -1764,6 +1767,12 @@ def _add_network_size_args(parser): + action='store_true', + help='If set, use original BERT residula connection ' + 'ordering.') ++ group.add_argument('--post-self-attn-layernorm', action='store_true', ++ help='If set, use post self attention layernorm.') ++ group.add_argument('--post-mlp-layernorm', action='store_true', ++ help='If set, use post MLP layernorm.') ++ group.add_argument('--use-gated-attention', action='store_true', ++ help='If set, use gated attention as in Qwen3Next') + group.add_argument('--openai-gelu', action='store_true', + help='Use OpenAIs GeLU implementation. This option' + 'should not be used unless for backward compatibility' +diff --git a/megatron/training/tokenizer/tokenizer.py b/megatron/training/tokenizer/tokenizer.py +index 13b7526ca..6c590f653 100644 +--- a/megatron/training/tokenizer/tokenizer.py ++++ b/megatron/training/tokenizer/tokenizer.py +@@ -136,7 +136,7 @@ class _HuggingFaceTokenizer(MegatronLegacyTokenizer): + # TODO(bnorick): download tokenizer once to lustre and use force offline to make sure all tasks read it from there + self._tokenizer = transformers.AutoTokenizer.from_pretrained( + pretrained_model_name_or_path=pretrained_model_name_or_path, +- trust_remote_code=trust_remote_code, ++ trust_remote_code=True, + **kwargs, + ) + self._vocab = self._tokenizer.get_vocab() diff --git a/docker/patch/v0.5.9/sglang.patch b/docker/patch/v0.5.9/sglang.patch new file mode 100644 index 0000000000..8cea544bd1 --- /dev/null +++ b/docker/patch/v0.5.9/sglang.patch @@ -0,0 +1,3046 @@ +diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py +index 6fbd1db82..f80ec11bb 100644 +--- a/python/sglang/srt/configs/model_config.py ++++ b/python/sglang/srt/configs/model_config.py +@@ -274,6 +274,7 @@ class ModelConfig: + + if is_draft_model and self.hf_config.architectures[0] in [ + "DeepseekV3ForCausalLM", ++ "DeepseekV32ForCausalLM", + "GlmMoeDsaForCausalLM", + ]: + self.hf_config.architectures[0] = "DeepseekV3ForCausalLMNextN" +diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py +index 67fe82ad6..2ef25c49b 100644 +--- a/python/sglang/srt/disaggregation/common/conn.py ++++ b/python/sglang/srt/disaggregation/common/conn.py +@@ -24,6 +24,7 @@ from sglang.srt.disaggregation.base.conn import ( + from sglang.srt.disaggregation.utils import DisaggregationMode + from sglang.srt.distributed import get_pp_group + from sglang.srt.layers.dp_attention import ( ++ get_attention_cp_size, + get_attention_dp_rank, + get_attention_dp_size, + get_attention_tp_rank, +@@ -116,10 +117,21 @@ class CommonKVManager(BaseKVManager): + + bootstrap_server_url = f"{host}:{self.bootstrap_port}" + url = f"http://{bootstrap_server_url}/route" ++ route_attn_tp_rank = self.attn_tp_rank ++ # In prefill CP mode, attention TP rank is flattened to 0, but requests are ++ # still routed by engine rank; register by engine rank to preserve all routes. ++ # Only apply this when actual CP is in use (cp_size > 1), not in pure DP ++ # attention mode (e.g. EP64) where each rank has its own dp_group already. ++ if ( ++ self.disaggregation_mode == DisaggregationMode.PREFILL ++ and self.attn_tp_size == 1 ++ and get_attention_cp_size() > 1 ++ ): ++ route_attn_tp_rank = self.kv_args.engine_rank + payload = { + "role": "Prefill", + "attn_tp_size": self.attn_tp_size, +- "attn_tp_rank": self.attn_tp_rank, ++ "attn_tp_rank": route_attn_tp_rank, + "attn_dp_size": self.attn_dp_size, + "attn_dp_rank": self.attn_dp_rank, + "pp_size": self.pp_size, +@@ -333,6 +345,10 @@ class CommonKVReceiver(BaseKVReceiver): + self.required_dst_info_num = ( + self.kv_mgr.attn_tp_size // self.prefill_attn_tp_size + ) ++ # With attention DP, one request is routed to one decode rank. ++ # Waiting for all TP shards to pre-allocate the same bootstrap room would stall forever. ++ if self.kv_mgr.attn_dp_size > 1: ++ self.required_dst_info_num = 1 + self.required_prefill_response_num = 1 * ( + self.prefill_pp_size // self.kv_mgr.pp_size + ) +@@ -357,6 +373,11 @@ class CommonKVReceiver(BaseKVReceiver): + # multiple connections in the connection pool and have to send dummy requests to other prefill ranks, + # or the KVPoll will never be set correctly + self.target_tp_rank = self.target_tp_ranks[0] ++ # For prefill CP mode (decode attention TP=1, prefill attention TP>1), ++ # route bootstrap to all prefill ranks as non-dummy so the serving rank ++ # always receives decode-side metadata. ++ if self.kv_mgr.attn_tp_size == 1 and self.prefill_attn_tp_size > 1: ++ self.target_tp_rank = None + self.required_dst_info_num = 1 + if self.kv_mgr.is_mla_backend: + self.required_prefill_response_num = ( +@@ -422,6 +443,7 @@ class CommonKVReceiver(BaseKVReceiver): + f"Could not fetch bootstrap info for engine rank: {self.kv_mgr.kv_args.engine_rank} and target_dp_group: {self.target_dp_group} and target_pp_rank {target_pp_rank}", + ) + self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed) ++ self.bootstrap_infos = None + return + + self.bootstrap_infos = bootstrap_infos +@@ -610,8 +632,12 @@ class CommonKVBootstrapServer(BaseKVBootstrapServer): + and int(target_dp_group) == -1 + and int(target_pp_rank) == -1 + ): ++ inferred_attn_tp_size = max( ++ (len(v) for v in self.prefill_port_table.values()), ++ default=self.attn_tp_size, ++ ) + prefill_parallel_info = { +- "prefill_attn_tp_size": self.attn_tp_size, ++ "prefill_attn_tp_size": inferred_attn_tp_size, + "prefill_dp_size": self.dp_size, + "prefill_pp_size": self.pp_size, + "prefill_page_size": self.page_size, +diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py +index 1d8baf002..1672de78d 100644 +--- a/python/sglang/srt/disaggregation/decode.py ++++ b/python/sglang/srt/disaggregation/decode.py +@@ -21,6 +21,7 @@ Life cycle of a request in the decode server + from __future__ import annotations + + import logging ++import os + import time + from collections import deque + from dataclasses import dataclass +@@ -336,6 +337,16 @@ class DecodePreallocQueue: + ) + return kv_manager + ++ def release_memory_occupation(self): ++ self.queue.clear() ++ self.retracted_queue.clear() ++ if hasattr(self.kv_manager, "deregister_buffer_to_engine"): ++ self.kv_manager.deregister_buffer_to_engine() ++ ++ def resume_memory_occupation(self): ++ if hasattr(self.kv_manager, "register_buffer_to_engine"): ++ self.kv_manager.register_buffer_to_engine() ++ + def add(self, req: Req, is_retracted: bool = False) -> None: + """Add a request to the pending queue.""" + if self._check_if_req_exceed_kv_capacity(req): +@@ -440,12 +451,37 @@ class DecodePreallocQueue: + [decode_req.kv_receiver for decode_req in self.queue], self.gloo_group + ) + ++ # Bootstrap timeout: if a request has been stuck in Bootstrapping for too long, treat it as failed. ++ bootstrap_timeout = float( ++ os.environ.get("SGLANG_DISAGGREGATION_TRANSFER_TIMEOUT", "600") ++ ) ++ now = time.perf_counter() ++ + for i, (decode_req, poll) in enumerate(zip(self.queue, polls)): + if rids_to_check is not None and decode_req.req.rid not in rids_to_check: + continue + + if poll == KVPoll.Bootstrapping: +- pass ++ # Check for bootstrap timeout ++ entry_time = getattr( ++ decode_req.req.time_stats, ++ "decode_prealloc_queue_entry_time", ++ None, ++ ) ++ if entry_time is not None and (now - entry_time) > bootstrap_timeout: ++ error_message = ( ++ f"Decode bootstrap timed out after {now - entry_time:.1f}s " ++ f"for request rank={self.tp_rank} " ++ f"{decode_req.req.rid=} {decode_req.req.bootstrap_room=}" ++ ) ++ logger.error(error_message) ++ prepare_abort( ++ decode_req.req, ++ error_message, ++ status_code=HTTPStatus.GATEWAY_TIMEOUT, ++ ) ++ if self.scheduler.enable_metrics: ++ self.scheduler.metrics_collector.increment_bootstrap_failed_reqs() + elif poll == KVPoll.WaitingForInput: + decode_req.waiting_for_input = True + elif poll == KVPoll.Failed: +@@ -830,6 +866,13 @@ class DecodeTransferQueue: + [decode_req.kv_receiver for decode_req in self.queue], self.gloo_group + ) + ++ # Transfer timeout: if a request has been in the transfer queue for too long ++ # (e.g., stuck in Bootstrapping/WaitingForInput/Transferring), treat it as failed. ++ transfer_timeout = float( ++ os.environ.get("SGLANG_DISAGGREGATION_TRANSFER_TIMEOUT", "600") ++ ) ++ now = time.perf_counter() ++ + transferred_reqs = [] + indices_to_remove = set() + for i, (decode_req, poll) in enumerate(zip(self.queue, polls)): +@@ -877,7 +920,31 @@ class DecodeTransferQueue: + KVPoll.WaitingForInput, + KVPoll.Transferring, + ]: +- pass ++ # Check for transfer timeout ++ entry_time = getattr( ++ decode_req.req.time_stats, ++ "decode_transfer_queue_entry_time", ++ None, ++ ) ++ if entry_time is not None and (now - entry_time) > transfer_timeout: ++ error_message = ( ++ f"Decode transfer timed out after {now - entry_time:.1f}s " ++ f"(state={poll}) for request rank={self.tp_rank} " ++ f"{decode_req.req.rid=} {decode_req.req.bootstrap_room=}" ++ ) ++ logger.error(error_message) ++ prepare_abort( ++ decode_req.req, ++ error_message, ++ status_code=HTTPStatus.GATEWAY_TIMEOUT, ++ ) ++ self.scheduler.stream_output( ++ [decode_req.req], decode_req.req.return_logprob ++ ) ++ release_kv_cache(decode_req.req, self.tree_cache, is_insert=False) ++ indices_to_remove.add(i) ++ if self.scheduler.enable_metrics: ++ self.scheduler.metrics_collector.increment_transfer_failed_reqs() + else: + raise ValueError(f"Unexpected poll case: {poll}") + +@@ -893,6 +960,14 @@ class DecodeTransferQueue: + + return transferred_reqs + ++ def release_memory_occupation(self): ++ """Clean up all in-flight transfers before releasing GPU memory.""" ++ self.queue.clear() ++ ++ def resume_memory_occupation(self): ++ """Resume after GPU memory re-allocation. Queue was already cleared on release.""" ++ pass ++ + + class SchedulerDisaggregationDecodeMixin: + +@@ -1072,7 +1147,15 @@ class SchedulerDisaggregationDecodeMixin: + resumed_reqs = self.disagg_decode_prealloc_queue.resume_retracted_reqs() + self.waiting_queue.extend(resumed_reqs) + if len(self.disagg_decode_prealloc_queue.retracted_queue) > 0: +- # if there are still retracted requests, we do not allocate new requests ++ # Still have retracted requests that couldn't resume (not enough memory). ++ # Don't accept new requests (pop_preallocated) — they would consume memory ++ # that retracted requests need. ++ # But DO drain completed transfers: their KV is already committed, and ++ # moving them to waiting_queue frees the reserved-decode-token budget ++ # in _allocatable_tokens(), which may unblock resume on the next iteration. ++ # Without this, completed transfers hold memory indefinitely → deadlock. ++ alloc_reqs = self.disagg_decode_transfer_queue.pop_transferred() ++ self.waiting_queue.extend(alloc_reqs) + return + + 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..fc4eef7b9 100644 +--- a/python/sglang/srt/disaggregation/mooncake/conn.py ++++ b/python/sglang/srt/disaggregation/mooncake/conn.py +@@ -260,6 +260,19 @@ class MooncakeKVManager(CommonKVManager): + self.kv_args.state_data_ptrs, self.kv_args.state_data_lens + ) + ++ def deregister_buffer_to_engine(self): ++ # Batch deregister KV data buffers ++ if self.kv_args.kv_data_ptrs: ++ self.engine.batch_deregister(self.kv_args.kv_data_ptrs) ++ ++ # Batch deregister auxiliary data buffers ++ if self.kv_args.aux_data_ptrs: ++ self.engine.batch_deregister(self.kv_args.aux_data_ptrs) ++ ++ # Batch deregister state/extra pool data buffers ++ if self.kv_args.state_data_ptrs: ++ self.engine.batch_deregister(self.kv_args.state_data_ptrs) ++ + def _transfer_data(self, mooncake_session_id, transfer_blocks): + if not transfer_blocks: + return 0 +@@ -643,13 +656,13 @@ class MooncakeKVManager(CommonKVManager): + raise RuntimeError( + f"PD Disaggregation does NOT support PD different TP sizes for non-MLA {state_type.upper()} hybrid models yet." + ) +- if len(prefill_state_indices) < len(req.dst_state_indices): +- logger.warning( +- f"len(prefill_state_indices) = {len(prefill_state_indices)}, len(dst_state_indices) = {len(req.dst_state_indices)}" ++ if len(prefill_state_indices) != len(req.dst_state_indices): ++ logger.error( ++ "PD extra-state index mismatch, reject transfer to avoid corrupted outputs: " ++ f"len(prefill_state_indices)={len(prefill_state_indices)}, " ++ f"len(dst_state_indices)={len(req.dst_state_indices)}" + ) +- prefill_state_indices = prefill_state_indices[ +- : len(req.dst_state_indices) +- ] ++ return -1 + # 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,43 @@ class MooncakeKVManager(CommonKVManager): + + if kv_chunk.is_last: + if kv_chunk.state_indices is not None: +- self.maybe_send_extra( ++ ret = self.maybe_send_extra( + req, + kv_chunk.state_indices, + target_rank_registration_info.dst_state_data_ptrs, + executor, + target_rank_registration_info, + ) ++ 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.failed_sessions.add( ++ req.mooncake_session_id ++ ) ++ logger.error( ++ f"Session {req.mooncake_session_id} failed." ++ ) ++ self.record_failure( ++ kv_chunk.room, ++ f"Failed to send extra state chunk of {kv_chunk.room} to {req.endpoint}:{req.dst_port}", ++ ) ++ self.update_status(kv_chunk.room, KVPoll.Failed) ++ self.sync_status_to_decode_endpoint( ++ req.endpoint, ++ req.dst_port, ++ req.room, ++ KVPoll.Failed, ++ local_rank, ++ ) ++ break + + # Only the last chunk we need to send the aux data + ret = self.send_aux( +@@ -895,6 +938,21 @@ class MooncakeKVManager(CommonKVManager): + target_rank_registration_info.dst_aux_ptrs, + ) + polls.append(True if ret == 0 else False) ++ if ret != 0: ++ # Mark session as failed to avoid hanging ++ # on subsequent batch_transfer_sync calls ++ with self.session_lock: ++ 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 ++ ) ++ logger.error( ++ f"Session {req.mooncake_session_id} failed (send_aux)." ++ ) + dst_ranks_infos.append( + (req.endpoint, req.dst_port, req.room) + ) +@@ -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) ++ # 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[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 ++ ) ++ arrived_response_num = len( ++ self.prefill_response_tracker[bootstrap_room] ++ ) ++ if arrived_response_num == expected_response_num: ++ self.update_status(bootstrap_room, KVPoll.Success) + elif status == KVPoll.Failed: + self.record_failure( + bootstrap_room, +@@ -1266,7 +1329,10 @@ class MooncakeKVReceiver(CommonKVReceiver): + super().__init__(mgr, bootstrap_addr, bootstrap_room, prefill_dp_rank) + + self.kv_mgr.addr_to_rooms_tracker[self.bootstrap_addr].add(self.bootstrap_room) +- self.kv_mgr.update_status(self.bootstrap_room, KVPoll.WaitingForInput) ++ # Only transition to WaitingForInput if bootstrap succeeded; ++ # if super().__init__() set status to Failed, do not override it. ++ if self.bootstrap_infos is not None: ++ self.kv_mgr.update_status(self.bootstrap_room, KVPoll.WaitingForInput) + + def _register_kv_args(self): + for bootstrap_info in self.bootstrap_infos: +diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py +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 + from __future__ import annotations + + import logging ++import os + import time + from collections import deque + from http import HTTPStatus +@@ -276,6 +277,12 @@ class PrefillBootstrapQueue: + [req.disagg_kv_sender for req in self.queue], self.gloo_group + ) + ++ # Bootstrap timeout: if a request has been stuck in Bootstrapping for too long, treat it as failed. ++ bootstrap_timeout = float( ++ os.environ.get("SGLANG_DISAGGREGATION_TRANSFER_TIMEOUT", "600") ++ ) ++ now = time.perf_counter() ++ + for i, (req, poll) in enumerate(zip(self.queue, polls)): + if rids_to_check is not None: + # if req not in reqs_info_to_check, skip +@@ -283,6 +290,27 @@ class PrefillBootstrapQueue: + continue + + if poll == KVPoll.Bootstrapping: ++ # Check for bootstrap timeout ++ entry_time = getattr( ++ req.time_stats, ++ "prefill_bootstrap_queue_entry_time", ++ None, ++ ) ++ if entry_time is not None and (now - entry_time) > bootstrap_timeout: ++ error_message = ( ++ f"Prefill bootstrap timed out after {now - entry_time:.1f}s " ++ f"for request rank={self.tp_rank} " ++ f"{req.rid=} {req.bootstrap_room=}" ++ ) ++ logger.error(error_message) ++ prepare_abort( ++ req, error_message, status_code=HTTPStatus.GATEWAY_TIMEOUT ++ ) ++ self.scheduler.stream_output([req], req.return_logprob) ++ indices_to_remove.add(i) ++ failed_reqs.append(req) ++ if self.scheduler.enable_metrics: ++ self.scheduler.metrics_collector.increment_bootstrap_failed_reqs() + continue + elif poll == KVPoll.Failed: + error_message = f"Prefill bootstrap failed for request rank={self.tp_rank} {req.rid=} {req.bootstrap_room=}" +@@ -335,6 +363,15 @@ class PrefillBootstrapQueue: + else: + return bootstrapped_reqs, failed_reqs + ++ def release_memory_occupation(self): ++ self.queue.clear() ++ if hasattr(self.kv_manager, "deregister_buffer_to_engine"): ++ self.kv_manager.deregister_buffer_to_engine() ++ ++ def resume_memory_occupation(self): ++ if hasattr(self.kv_manager, "register_buffer_to_engine"): ++ self.kv_manager.register_buffer_to_engine() ++ + + class SchedulerDisaggregationPrefillMixin: + """ +@@ -564,6 +601,13 @@ class SchedulerDisaggregationPrefillMixin: + self.attn_tp_cpu_group, + ) + ++ # Transfer timeout: if a request has been in the inflight queue for too long ++ # (e.g., stuck in WaitingForInput/Transferring), treat it as failed. ++ transfer_timeout = float( ++ os.environ.get("SGLANG_DISAGGREGATION_TRANSFER_TIMEOUT", "600") ++ ) ++ now = time.perf_counter() ++ + undone_reqs: List[Req] = [] + # Check .poll() for the reqs in disagg_prefill_inflight_queue. If Success, respond to the client and remove it from the queue + for req, poll in zip(self.disagg_prefill_inflight_queue, polls): +@@ -573,10 +617,35 @@ class SchedulerDisaggregationPrefillMixin: + undone_reqs.append(req) + continue + +- assert poll == KVPoll.Success or poll == KVPoll.Failed ++ if poll not in (KVPoll.Success, KVPoll.Failed): ++ undone_reqs.append(req) ++ continue + + if poll in [KVPoll.WaitingForInput, KVPoll.Transferring]: +- undone_reqs.append(req) ++ # Check for transfer timeout ++ entry_time = getattr( ++ req.time_stats, ++ "prefill_transfer_queue_entry_time", ++ None, ++ ) ++ if entry_time is not None and (now - entry_time) > transfer_timeout: ++ error_message = ( ++ f"Prefill transfer timed out after {now - entry_time:.1f}s " ++ f"(state={poll}) for request rank={self.tp_rank} " ++ f"{req.rid=} {req.bootstrap_room=}" ++ ) ++ logger.error(error_message) ++ release_kv_cache(req, self.tree_cache) # unlock the tree ++ prepare_abort( ++ req, error_message, status_code=HTTPStatus.GATEWAY_TIMEOUT ++ ) ++ if hasattr(req.disagg_kv_sender, "clear"): ++ req.disagg_kv_sender.clear() ++ done_reqs.append(req) ++ if self.enable_metrics: ++ self.metrics_collector.increment_transfer_failed_reqs() ++ else: ++ undone_reqs.append(req) + elif poll == KVPoll.Success: # transfer done + release_kv_cache(req, self.tree_cache) # unlock the tree + req.finished_reason = FINISH_LENGTH(length=0) +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 ++++ b/python/sglang/srt/distributed/parallel_state.py +@@ -1999,7 +1999,10 @@ def get_tensor_model_parallel_world_size(): + + def get_tensor_model_parallel_rank(): + """Return my rank for the tensor model parallel group.""" +- return get_tp_group().rank_in_group ++ try: ++ return get_tp_group().rank_in_group ++ except Exception: ++ return 0 + + + # ATTN_TP +diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py +index 0ed5a1b44..67e33c650 100644 +--- a/python/sglang/srt/entrypoints/engine.py ++++ b/python/sglang/srt/entrypoints/engine.py +@@ -52,6 +52,7 @@ from sglang.srt.managers.io_struct import ( + LoadLoRAAdapterReqInput, + MultimodalDataInputFormat, + OpenSessionReqInput, ++ PostProcessWeightsReqInput, + ReleaseMemoryOccupationReqInput, + ResumeMemoryOccupationReqInput, + RpcReqInput, +@@ -641,6 +642,24 @@ class Engine(EngineBase): + self.tokenizer_manager.update_weights_from_ipc(obj, None) + ) + ++ def post_process_weights( ++ self, ++ restore_weights_before_load: bool = False, ++ post_process_quantization: bool = False, ++ ): ++ """ ++ Optional post-processing for updated weights (e.g., Marlin conversion). ++ Should be called after weight update is finished. ++ """ ++ obj = PostProcessWeightsReqInput( ++ restore_weights_before_load=restore_weights_before_load, ++ post_process_quantization=post_process_quantization, ++ ) ++ ++ return self.loop.run_until_complete( ++ self.tokenizer_manager.post_process_weights(obj, None) ++ ) ++ + def get_weights_by_name(self, name: str, truncate_size: int = 100): + """Get weights by parameter name.""" + obj = GetWeightsByNameReqInput(name=name, truncate_size=truncate_size) +diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py +index 1d6816c01..402b42e05 100644 +--- a/python/sglang/srt/entrypoints/http_server.py ++++ b/python/sglang/srt/entrypoints/http_server.py +@@ -115,6 +115,7 @@ from sglang.srt.managers.io_struct import ( + OpenSessionReqInput, + ParseFunctionCallReq, + PauseGenerationReqInput, ++ PostProcessWeightsReqInput, + ProfileReqInput, + ReleaseMemoryOccupationReqInput, + ResumeMemoryOccupationReqInput, +@@ -574,10 +575,8 @@ async def model_info(): + @app.get("/weight_version") + async def weight_version(): + """Get the current weight version.""" +- raise HTTPException( +- status_code=404, +- detail="Endpoint '/get_weight_version' or '/weight_version' is deprecated. Please use '/model_info' instead.", +- ) ++ result = await model_info() ++ return {"weight_version": result.get("weight_version", None)} + + + @app.get("/get_server_info") +@@ -594,9 +593,19 @@ async def get_server_info(): + async def server_info(): + """Get the server information.""" + # Returns internal states per DP. +- internal_states: List[Dict[Any, Any]] = ( +- await _global_state.tokenizer_manager.get_internal_state() +- ) ++ # In large/disaggregated deployments this can occasionally block; keep endpoint responsive. ++ server_info_timeout = float(os.environ.get("SGLANG_SERVER_INFO_TIMEOUT", "2")) ++ try: ++ internal_states: List[Dict[Any, Any]] = await asyncio.wait_for( ++ _global_state.tokenizer_manager.get_internal_state(), ++ timeout=server_info_timeout, ++ ) ++ except asyncio.TimeoutError: ++ logger.warning( ++ "Timed out getting internal state for /server_info after %.1fs; returning empty internal_states", ++ server_info_timeout, ++ ) ++ internal_states = [] + + # This field is not serializable. + if hasattr(_global_state.tokenizer_manager.server_args, "model_config"): +@@ -1084,6 +1093,23 @@ async def update_weights_from_ipc(obj: UpdateWeightsFromIPCReqInput, request: Re + return ORJSONResponse(content, status_code=HTTPStatus.BAD_REQUEST) + + ++@app.post("/post_process_weights") ++@auth_level(AuthLevel.ADMIN_OPTIONAL) ++async def post_process_weights(req: PostProcessWeightsReqInput, request: Request): ++ """ ++ Optional post-processing for updated weights (e.g., Marlin conversion). ++ This should be called selectively after `update_weights_from_distributed/update_weights_from_tensor`. ++ """ ++ success, message = await _global_state.tokenizer_manager.post_process_weights( ++ req, request ++ ) ++ ++ content = {"success": success, "message": message} ++ return ORJSONResponse( ++ content, status_code=200 if success else HTTPStatus.BAD_REQUEST ++ ) ++ ++ + @app.post("/update_weight_version") + @auth_level(AuthLevel.ADMIN_OPTIONAL) + async def update_weight_version(obj: UpdateWeightVersionReqInput, request: Request): +diff --git a/python/sglang/srt/layers/attention/nsa/index_buf_accessor.py b/python/sglang/srt/layers/attention/nsa/index_buf_accessor.py +index 1cdf65b91..4783cd18f 100644 +--- a/python/sglang/srt/layers/attention/nsa/index_buf_accessor.py ++++ b/python/sglang/srt/layers/attention/nsa/index_buf_accessor.py +@@ -630,7 +630,6 @@ def _get_k_and_s_triton( + page_indices, + k_out, + s_out, +- seq_len, + page_size, + buf_numel_per_page, + index_head_dim, +@@ -647,7 +646,6 @@ def _get_k_and_s_triton_kernel( + page_indices_ptr, + k_out_ptr, + s_out_ptr, +- seq_len: tl.constexpr, + page_size: tl.constexpr, + buf_numel_per_page: tl.constexpr, + index_head_dim: tl.constexpr, +diff --git a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py +index ca54a931b..6c102a251 100644 +--- a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py ++++ b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py +@@ -1,6 +1,7 @@ + from __future__ import annotations + + import contextlib ++import os + from abc import ABC, abstractmethod + from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple + +@@ -207,7 +208,11 @@ class Indexer(MultiPlatformOp): + max_position=max_position_embeddings, + base=rope_theta, # type: ignore + rope_scaling=rope_scaling, +- is_neox_style=is_neox_style, ++ is_neox_style=( ++ os.environ.get("INDEXER_ROPE_NEOX_STYLE", "1") == "1" ++ if os.environ.get("INDEXER_ROPE_NEOX_STYLE", None) ++ else is_neox_style ++ ), + device=get_global_server_args().device, + ) + self.block_size = block_size +@@ -244,6 +249,11 @@ class Indexer(MultiPlatformOp): + x = x.to(self.weights_proj.weight.dtype) + weights, _ = self.weights_proj(x) + weights = weights.float() ++ if weights.shape[1] < q_scale.shape[1]: ++ assert q_scale.shape[1] % weights.shape[1] == 0 ++ weights = weights.repeat_interleave( ++ q_scale.shape[1] // weights.shape[1], dim=1 ++ ) + weights = weights * self.n_heads**-0.5 + weights = weights.unsqueeze(-1) * q_scale * self.softmax_scale + return weights +@@ -982,15 +992,26 @@ class Indexer(MultiPlatformOp): + query, key = self._get_q_k_bf16( + q_lora, x, positions, enable_dual_stream, forward_batch=forward_batch + ) ++ if query.shape[1] < 32: ++ assert 32 % query.shape[1] == 0 ++ query = query.repeat_interleave(32 // query.shape[1], dim=1) + q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt) + with torch.cuda.stream(self.alt_stream): + k_fp8, k_scale = act_quant(key, self.block_size, self.scale_fmt) + current_stream.wait_stream(self.alt_stream) ++ if weights.shape[1] < q_scale.shape[1]: ++ assert q_scale.shape[1] % weights.shape[1] == 0 ++ weights = weights.repeat_interleave( ++ q_scale.shape[1] // weights.shape[1], dim=1 ++ ) + weights = weights.unsqueeze(-1) * q_scale * self.softmax_scale + else: + query, key = self._get_q_k_bf16( + q_lora, x, positions, enable_dual_stream, forward_batch=forward_batch + ) ++ if query.shape[1] < 32: ++ assert 32 % query.shape[1] == 0 ++ query = query.repeat_interleave(32 // query.shape[1], dim=1) + + if enable_dual_stream: + current_stream = torch.cuda.current_stream() +diff --git a/python/sglang/srt/layers/attention/nsa/utils.py b/python/sglang/srt/layers/attention/nsa/utils.py +index 00ef96f9b..5adaec804 100644 +--- a/python/sglang/srt/layers/attention/nsa/utils.py ++++ b/python/sglang/srt/layers/attention/nsa/utils.py +@@ -54,7 +54,12 @@ def can_nsa_prefill_cp_round_robin_split(forward_batch: "ForwardBatch"): + return False + cp_size = get_attention_cp_size() + seq_len = sum(forward_batch.extend_seq_lens_cpu) +- return is_nsa_prefill_cp_round_robin_split() and seq_len > 0 and cp_size > 1 ++ return ( ++ is_nsa_prefill_cp_round_robin_split() ++ and seq_len >= cp_size ++ and seq_len % cp_size == 0 ++ and cp_size > 1 ++ ) + + + def nsa_cp_round_robin_split_data(input_: Union[torch.Tensor, List]): +@@ -91,20 +96,29 @@ def nsa_cp_round_robin_split_data(input_: Union[torch.Tensor, List]): + def cal_padded_tokens(forward_batch: "ForwardBatch"): + # Consistent with the padding calculation logic in ForwardBatch.prepare_mlp_sync_batch, + # calculate the actual token length after padding when attn_tp_size > 1 or in the MAX_LEN padding mode. +- global_num_tokens = forward_batch.global_num_tokens_cpu.copy() ++ if forward_batch.global_num_tokens_cpu is None: ++ # PD prefill CP+PP path can bypass MLP-sync metadata. Reconstruct a single-rank ++ # global token view from the local token count for NSA padding logic. ++ local_tokens = forward_batch.num_token_non_padded_cpu ++ if local_tokens is None: ++ local_tokens = len(forward_batch.input_ids) ++ global_num_tokens = [local_tokens * get_attention_cp_size()] ++ else: ++ global_num_tokens = forward_batch.global_num_tokens_cpu.copy() + sync_group_size = len(global_num_tokens) + attn_cp_size = get_attention_cp_size() + for i in range(sync_group_size): + global_num_tokens[i] = ceil_align(global_num_tokens[i], attn_cp_size) +- dp_padding_mode = DpPaddingMode.get_dp_padding_mode( +- forward_batch.is_extend_in_batch, global_num_tokens +- ) +- if dp_padding_mode.is_max_len(): +- tokens = max(global_num_tokens) +- elif len(global_num_tokens) > 1: +- tokens = global_num_tokens[get_attention_dp_rank()] +- else: ++ if len(global_num_tokens) == 1: + tokens = global_num_tokens[0] ++ else: ++ dp_padding_mode = DpPaddingMode.get_dp_padding_mode( ++ forward_batch.is_extend_in_batch, global_num_tokens ++ ) ++ if dp_padding_mode.is_max_len(): ++ tokens = max(global_num_tokens) ++ else: ++ tokens = global_num_tokens[get_attention_dp_rank()] + if can_nsa_prefill_cp_round_robin_split(forward_batch): + tokens = ceil_div(tokens, attn_cp_size) + return tokens +@@ -152,10 +166,20 @@ class NSAContextParallelMetadata: + + def can_cp_split(seq_len: int, cp_size: int, use_nsa: bool, forward_batch): + if is_nsa_prefill_cp_round_robin_split(): +- cur_cp_seq_len = seq_len // cp_size +- assert ( +- seq_len % cp_size == 0 +- ), f"seq_len {seq_len} is not divisible by cp_size {cp_size} when nsa_prefill_cp_mode is round-robin-split" ++ # Use actual extend sequence length instead of (possibly padded) input_ids ++ # length to stay consistent with can_nsa_prefill_cp_round_robin_split(), ++ # which also checks sum(extend_seq_lens_cpu). When prepare_mlp_sync_batch ++ # pads input_ids to ceil_align(n, attn_cp_size), len(input_ids) can become ++ # divisible by cp_size even though the real extend length is not, causing ++ # hidden_states to be CP-split while the attention metadata is not. ++ actual_seq_len = ( ++ sum(forward_batch.extend_seq_lens_cpu) ++ if forward_batch.extend_seq_lens_cpu is not None ++ else seq_len ++ ) ++ if actual_seq_len < cp_size or actual_seq_len % cp_size != 0: ++ return False ++ cur_cp_seq_len = actual_seq_len // cp_size + else: + # TODO current just support prefill batch=1 and len(input_ids) > self.cp_size * 2 + # Note: (self.cp_size * 2) To achieve load balancing for seq computation, +@@ -175,10 +199,6 @@ def can_cp_split(seq_len: int, cp_size: int, use_nsa: bool, forward_batch): + + def cp_split_and_rebuild_data(forward_batch, input_: torch.Tensor): + if is_nsa_prefill_cp_round_robin_split(): +- cp_size = get_attention_cp_size() +- assert ( +- input_.shape[0] % cp_size == 0 +- ), f"Expect input shape 0 can divided by cp size, but got input shape {input_.shape}, cp size {cp_size}" + return nsa_cp_round_robin_split_data(input_) + + input_list = list( +@@ -192,11 +212,6 @@ def cp_split_and_rebuild_data(forward_batch, input_: torch.Tensor): + + def cp_split_and_rebuild_position(forward_batch, positions: torch.Tensor): + if is_nsa_prefill_cp_round_robin_split(): +- cp_size = get_attention_cp_size() +- assert positions.shape[0] % cp_size == 0, ( +- f"Expect positions shape 0 can divided by cp size, but got positions shape {positions.shape}, " +- f"cp size {cp_size}" +- ) + return nsa_cp_round_robin_split_data(positions) + + position_id_list = list( +diff --git a/python/sglang/srt/layers/communicator_nsa_cp.py b/python/sglang/srt/layers/communicator_nsa_cp.py +index 296d14568..f4606a769 100644 +--- a/python/sglang/srt/layers/communicator_nsa_cp.py ++++ b/python/sglang/srt/layers/communicator_nsa_cp.py +@@ -34,7 +34,6 @@ from sglang.srt.layers.communicator import ( + from sglang.srt.layers.dp_attention import ( + attn_cp_all_gather_into_tensor, + attn_cp_reduce_scatter_tensor, +- get_local_dp_buffer, + ) + from sglang.srt.model_executor.forward_batch_info import ForwardBatch + +@@ -153,9 +152,23 @@ class NSACPCommunicateWithAllReduceAndLayerNormFn( + # for decode: attn tp full -> full + if nsa_use_prefill_cp(forward_batch): + assert context.attn_dp_size == 1 +- hidden_states, local_hidden_states = ( +- get_local_dp_buffer(), +- hidden_states, ++ local_hidden_states = hidden_states ++ total_tokens = ( ++ sum(forward_batch.extend_seq_lens_cpu) ++ if forward_batch.extend_seq_lens_cpu is not None ++ else local_hidden_states.shape[0] * context.attn_cp_size ++ ) ++ max_len = (total_tokens + context.attn_cp_size - 1) // context.attn_cp_size ++ if local_hidden_states.shape[0] < max_len: ++ pad = local_hidden_states.new_zeros( ++ ( ++ max_len - local_hidden_states.shape[0], ++ local_hidden_states.shape[1], ++ ) ++ ) ++ local_hidden_states = torch.cat([local_hidden_states, pad], dim=0) ++ hidden_states = local_hidden_states.new_empty( ++ (max_len * context.attn_cp_size, local_hidden_states.shape[1]) + ) + attn_cp_all_gather_into_tensor( + hidden_states, +diff --git a/python/sglang/srt/layers/dp_attention.py b/python/sglang/srt/layers/dp_attention.py +index 5bf5aa0c8..e52f39fd8 100644 +--- a/python/sglang/srt/layers/dp_attention.py ++++ b/python/sglang/srt/layers/dp_attention.py +@@ -90,11 +90,11 @@ class _DpGatheredBufferWrapper: + _hidden_size: int + _dtype: torch.dtype + _device: torch.device +- _global_dp_buffer_len: int +- _local_dp_buffer_len: int +- _dp_max_padding: bool +- _global_num_tokens: Optional[List[int]] +- _is_extend_in_batch: bool ++ _global_dp_buffer_len: int = 0 ++ _local_dp_buffer_len: int = 0 ++ _dp_max_padding: bool = False ++ _global_num_tokens: Optional[List[int]] = None ++ _is_extend_in_batch: bool = False + + @classmethod + def set_metadata(cls, hidden_size: int, dtype: torch.dtype, device: torch.device): +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 ++++ b/python/sglang/srt/layers/logits_processor.py +@@ -872,11 +872,6 @@ class LogitsProcessor(nn.Module): + None, # bias + True, # is_vnni + ) +- elif get_global_server_args().rl_on_policy_target is not None: +- # Due to tie-weight, we may not be able to change lm_head's weight dtype +- logits = torch.matmul( +- hidden_states.bfloat16(), lm_head.weight.T.bfloat16() +- ) + else: + logits = torch.matmul( + hidden_states.to(lm_head.weight.dtype), lm_head.weight.T +diff --git a/python/sglang/srt/layers/moe/ep_moe/deepep_bf16_kernels.py b/python/sglang/srt/layers/moe/ep_moe/deepep_bf16_kernels.py +new file mode 100644 +index 000000000..8d3d0f92e +--- /dev/null ++++ b/python/sglang/srt/layers/moe/ep_moe/deepep_bf16_kernels.py +@@ -0,0 +1,146 @@ ++"""Fused Triton kernels for DeepEP BF16 low-latency MoE decode. ++ ++Replaces the naive activation + masking pipeline (5+ CUDA kernels for silu+mul ++and arange+comparison+masked_fill+copy) with a single Triton elementwise kernel, ++while keeping cuBLAS batched GEMM for the matrix multiplies. ++ ++Pipeline: bmm → fused_act_mul_masked (in-place) → bmm(out=hidden) ++ (3 ops total: 2 cuBLAS + 1 Triton, vs original 7-8 separate CUDA kernels) ++""" ++ ++import torch ++import triton ++import triton.language as tl ++ ++ ++@triton.jit ++def _silu_mul_masked_kernel( ++ gate_up_ptr, ++ masked_m_ptr, ++ M, ++ N, ++ stride_ge, ++ stride_gm, ++ stride_gn, ++ BLOCK: tl.constexpr, ++): ++ """Fused SiLU(gate) * up with per-expert masking, written in-place. ++ ++ gate_up: [E, M, 2*N] — first N cols are gate, last N cols are up. ++ Writes SiLU(gate)*up to gate_up[:,:,:N] in-place. ++ Rows m >= masked_m[e] are zeroed. ++ """ ++ expert_id = tl.program_id(1) ++ pid = tl.program_id(0) ++ ++ expert_valid_m = tl.load(masked_m_ptr + expert_id) ++ ++ offs = pid * BLOCK + tl.arange(0, BLOCK) ++ total = M * N ++ mask = offs < total ++ ++ m = offs // N ++ n = offs % N ++ ++ gate_base = gate_up_ptr + expert_id * stride_ge ++ ++ gate_val = tl.load(gate_base + m * stride_gm + n * stride_gn, mask=mask, other=0.0) ++ up_val = tl.load( ++ gate_base + m * stride_gm + (n + N) * stride_gn, mask=mask, other=0.0 ++ ) ++ ++ gate_f32 = gate_val.to(tl.float32) ++ result = (gate_f32 * tl.sigmoid(gate_f32)) * up_val.to(tl.float32) ++ ++ # Zero invalid rows ++ valid = m < expert_valid_m ++ result = tl.where(valid, result, 0.0) ++ ++ tl.store( ++ gate_base + m * stride_gm + n * stride_gn, ++ result.to(gate_up_ptr.dtype.element_ty), ++ mask=mask, ++ ) ++ ++ ++@triton.jit ++def _gelu_mul_masked_kernel( ++ gate_up_ptr, ++ masked_m_ptr, ++ M, ++ N, ++ stride_ge, ++ stride_gm, ++ stride_gn, ++ BLOCK: tl.constexpr, ++): ++ """Fused GELU(gate) * up with per-expert masking, written in-place.""" ++ expert_id = tl.program_id(1) ++ pid = tl.program_id(0) ++ ++ expert_valid_m = tl.load(masked_m_ptr + expert_id) ++ ++ offs = pid * BLOCK + tl.arange(0, BLOCK) ++ total = M * N ++ mask = offs < total ++ ++ m = offs // N ++ n = offs % N ++ ++ gate_base = gate_up_ptr + expert_id * stride_ge ++ ++ gate_val = tl.load(gate_base + m * stride_gm + n * stride_gn, mask=mask, other=0.0) ++ up_val = tl.load( ++ gate_base + m * stride_gm + (n + N) * stride_gn, mask=mask, other=0.0 ++ ) ++ ++ g = gate_val.to(tl.float32) ++ kAlpha = 0.7978845608028654 ++ gate_act = 0.5 * g * (1.0 + tl.math.tanh(kAlpha * (g + 0.044715 * g * g * g))) ++ result = gate_act * up_val.to(tl.float32) ++ ++ valid = m < expert_valid_m ++ result = tl.where(valid, result, 0.0) ++ ++ tl.store( ++ gate_base + m * stride_gm + n * stride_gn, ++ result.to(gate_up_ptr.dtype.element_ty), ++ mask=mask, ++ ) ++ ++ ++def fused_act_mul_masked_inplace( ++ gate_up: torch.Tensor, ++ intermediate_size: int, ++ masked_m: torch.Tensor, ++ use_gelu: bool = False, ++) -> None: ++ """Fused activation + multiply + masking, written in-place to gate_up[:,:,:I]. ++ ++ After this call, gate_up[:, :, :intermediate_size] contains the masked ++ activated intermediate, suitable for the down projection GEMM. ++ ++ Args: ++ gate_up: [E, M, 2*I] output of bmm(tokens, w13.T), modified in-place ++ intermediate_size: I ++ masked_m: [E] per-expert valid token count ++ use_gelu: use GELU instead of SiLU ++ """ ++ E, M, _ = gate_up.shape ++ N = intermediate_size ++ ++ total = M * N ++ BLOCK = 1024 ++ grid = (triton.cdiv(total, BLOCK), E) ++ ++ kernel = _gelu_mul_masked_kernel if use_gelu else _silu_mul_masked_kernel ++ kernel[grid]( ++ gate_up, ++ masked_m, ++ M, ++ N, ++ gate_up.stride(0), ++ gate_up.stride(1), ++ gate_up.stride(2), ++ BLOCK=BLOCK, ++ ) +diff --git a/python/sglang/srt/layers/moe/ep_moe/layer.py b/python/sglang/srt/layers/moe/ep_moe/layer.py +index ebcc696ec..3b527021a 100644 +--- a/python/sglang/srt/layers/moe/ep_moe/layer.py ++++ b/python/sglang/srt/layers/moe/ep_moe/layer.py +@@ -132,11 +132,12 @@ class DeepEPMoE(FusedMoE): + and not _is_npu + and not ( + get_moe_runner_backend().is_flashinfer_cutedsl() ++ and self.quant_config is not None + and self.quant_config.get_name() == "modelopt_fp4" + ) ++ and (self.use_fp8_w8a8 or self.use_w4afp8) + ): +- # NPU supports low_latency deepep without deepgemm +- # FP4 quantization with flashinfer_cutedsl also supports low_latency deepep without deepgemm ++ # BF16 models don't need deep_gemm; they use per-expert torch.mm + assert ( + deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM + ), f"DeepEP {self.deepep_mode} mode requires deep_gemm" +@@ -154,6 +155,10 @@ class DeepEPMoE(FusedMoE): + # the last one is invalid rank_id + self.expert_mask[:-1] = 1 + ++ # Set bf16_weights flag on dispatcher so dispatch skips FP8 quantization ++ if not self.use_fp8_w8a8 and not self.use_w4afp8: ++ self.dispatcher.set_quant_config({"bf16_weights": True}) ++ + def forward( + self, + hidden_states: torch.Tensor, +@@ -228,6 +233,8 @@ class DeepEPMoE(FusedMoE): + elif DispatchOutputChecker.format_is_deepep_normal(dispatch_output): + if self.use_w4afp8: + output = self.forward_cutlass_w4afp8(dispatch_output) ++ elif not self.use_fp8_w8a8: ++ output = self.forward_bf16_normal(dispatch_output) + else: + assert False, "forward_deepgemm_contiguous is deprecated" + elif DispatchOutputChecker.format_is_deepep_ll(dispatch_output): +@@ -238,6 +245,8 @@ class DeepEPMoE(FusedMoE): + output = self.forward_flashinfer_cutedsl(dispatch_output) + elif self.use_w4afp8: + output = self.forward_cutlass_w4afp8_masked(dispatch_output) ++ elif not self.use_fp8_w8a8: ++ output = self.forward_bf16_ll(dispatch_output) + else: + assert False, "forward_deepgemm_masked is deprecated" + +@@ -341,6 +350,71 @@ class DeepEPMoE(FusedMoE): + dispatch_output=dispatch_output, + ) + ++ def forward_bf16_normal( ++ self, ++ dispatch_output: DeepEPNormalDispatchOutput, ++ ) -> torch.Tensor: ++ from sglang.srt.layers.moe.fused_moe_triton.fused_moe import fused_experts ++ ++ hidden_states = dispatch_output.hidden_states ++ topk_ids = dispatch_output.topk_ids ++ topk_weights = dispatch_output.topk_weights ++ ++ if hidden_states.shape[0] == 0: ++ return hidden_states ++ ++ # topk_ids uses local expert IDs (0..num_local_experts-1), -1 for remote. ++ # fused_experts handles -1 via moe_align_block_size filtering. ++ return fused_experts( ++ hidden_states=hidden_states, ++ w1=self.w13_weight, ++ w2=self.w2_weight, ++ topk_output=(topk_weights, topk_ids, None), ++ moe_runner_config=self.moe_runner_config, ++ ) ++ ++ def forward_bf16_ll( ++ self, ++ dispatch_output: DeepEPLLDispatchOutput, ++ ) -> torch.Tensor: ++ from sglang.srt.layers.moe.ep_moe.deepep_bf16_kernels import ( ++ fused_act_mul_masked_inplace, ++ ) ++ ++ hidden_states = dispatch_output.hidden_states ++ masked_m = dispatch_output.masked_m ++ expected_m = dispatch_output.expected_m ++ ++ _, max_tokens, _ = hidden_states.shape ++ if masked_m.numel() == 0 or max_tokens == 0: ++ return hidden_states ++ ++ expected_m = min(expected_m, max_tokens) ++ if expected_m <= 0: ++ return hidden_states ++ ++ tokens = hidden_states[:, :expected_m, :] ++ ++ # 1. Gate+Up GEMM (cuBLAS batched GEMM) ++ gate_up = torch.bmm(tokens, self.w13_weight.transpose(1, 2)) ++ ++ # 2. Fused SiLU(gate)*up + masking in-place (1 Triton kernel replaces 6 ops) ++ fused_act_mul_masked_inplace( ++ gate_up, ++ self.intermediate_size_per_partition, ++ masked_m, ++ use_gelu=(self.moe_runner_config.activation == "gelu"), ++ ) ++ ++ # 3. Down GEMM into hidden_states (cuBLAS, non-contiguous input is OK) ++ torch.bmm( ++ gate_up[:, :, : self.intermediate_size_per_partition], ++ self.w2_weight.transpose(1, 2), ++ out=hidden_states[:, :expected_m, :], ++ ) ++ ++ return hidden_states ++ + def forward_npu( + self, + dispatch_output: Union[DeepEPNormalDispatchOutput, DeepEPLLDispatchOutput], +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 ++++ b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py +@@ -697,6 +697,7 @@ class FusedMoE(torch.nn.Module): + "CompressedTensorsWNA16TritonMoE", + ] + ) ++ and "zero" not in weight_name + else loaded_weight + ) + +@@ -916,6 +917,7 @@ class FusedMoE(torch.nn.Module): + "CompressedTensorsWNA16TritonMoE", + ] + ) ++ and "zero" not in weight_name + else loaded_weight + ) + +diff --git a/python/sglang/srt/layers/moe/routed_experts_capturer.py b/python/sglang/srt/layers/moe/routed_experts_capturer.py +index 00bd68755..12d5577af 100644 +--- a/python/sglang/srt/layers/moe/routed_experts_capturer.py ++++ b/python/sglang/srt/layers/moe/routed_experts_capturer.py +@@ -8,10 +8,15 @@ import torch + + from sglang.srt.configs.model_config import ModelConfig + from sglang.srt.layers.dp_attention import ( ++ attn_tp_all_gather_into_tensor, + get_attention_dp_rank, ++ get_attention_tp_size, + get_dp_local_info, + is_dp_attention_enabled, + ) ++from sglang.srt.layers.moe import ( ++ get_moe_a2a_backend, ++) + from sglang.srt.mem_cache.memory_pool import ReqToTokenPool + from sglang.srt.model_executor.forward_batch_info import ForwardBatch + from sglang.srt.server_args import get_global_server_args +@@ -181,13 +186,26 @@ class _RoutedExpertsCapturerReal(RoutedExpertsCapturer): + device=device, + ) + ++ if get_moe_a2a_backend().is_deepep(): ++ attn_tp_size = get_attention_tp_size() if is_dp_attention_enabled() else 1 ++ self.gather_buffer = torch.empty( ++ ( ++ self.device_cache.buffer.shape[0] * attn_tp_size, ++ self.device_cache.buffer.shape[2], ++ ), ++ dtype=torch.int32, ++ device=device, ++ ) ++ + def _sync_fwd_experts_buffer_DtoH( + self, + forward_batch: ForwardBatch, + can_run_graph: bool, + cuda_graph_batch: int, + ): +- if is_dp_attention_enabled(): ++ # When DeepEP is enabled, capture() already does all_gather, so device_cache.buffer ++ # contains data from all DP ranks. We should not slice by DP rank in this case. ++ if is_dp_attention_enabled() and not get_moe_a2a_backend().is_deepep(): + local_start_pos, local_num_tokens = get_dp_local_info(forward_batch) + # handle with cuda graph padding + if can_run_graph: +@@ -206,6 +224,12 @@ class _RoutedExpertsCapturerReal(RoutedExpertsCapturer): + ].cpu() + + def capture(self, layer_id: int, topk_ids: torch.Tensor): ++ if get_moe_a2a_backend().is_deepep(): ++ local_topk_ids = topk_ids ++ topk_ids = self.gather_buffer[ ++ : local_topk_ids.size(0) * get_attention_tp_size() ++ ] ++ attn_tp_all_gather_into_tensor(topk_ids, local_topk_ids) + self.device_cache.capture_fwd_routed_experts(layer_id, topk_ids) + + def get_routed_experts( +diff --git a/python/sglang/srt/layers/moe/token_dispatcher/deepep.py b/python/sglang/srt/layers/moe/token_dispatcher/deepep.py +index 8539639d5..e7f5d1565 100644 +--- a/python/sglang/srt/layers/moe/token_dispatcher/deepep.py ++++ b/python/sglang/srt/layers/moe/token_dispatcher/deepep.py +@@ -388,6 +388,7 @@ class _DeepEPDispatcherImplNormal(_DeepEPDispatcherImplBase): + deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM + and not get_moe_runner_backend().is_cutlass() + and not envs.SGLANG_DEEPEP_BF16_DISPATCH.get() ++ and not self.quant_config.get("bf16_weights", False) + ): + # TODO hard code 128 block quant,use fp8 communication + hidden_states = sglang_per_token_group_quant_fp8( +@@ -466,7 +467,12 @@ class _DeepEPDispatcherImplNormal(_DeepEPDispatcherImplBase): + previous_event=previous_event, + async_finish=self.async_finish, + allocate_on_comm_stream=(previous_event is not None) and self.async_finish, +- expert_alignment=128 if deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM else 1, ++ expert_alignment=( ++ 128 ++ if deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM ++ and not self.quant_config.get("bf16_weights", False) ++ else 1 ++ ), + config=DeepEPConfig.get_instance().normal_dispatch_config, + ) + get_global_expert_distribution_recorder().on_deepep_dispatch_normal( +@@ -491,7 +497,12 @@ class _DeepEPDispatcherImplNormal(_DeepEPDispatcherImplBase): + topk_weights: torch.Tensor, + ): + +- if deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM or _use_aiter or _is_npu: ++ if ( ++ deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM ++ or _use_aiter ++ or _is_npu ++ or self.quant_config.get("bf16_weights", False) ++ ): + output = hidden_states + else: + raise NotImplementedError() # triton runner was supported but it's temporarily disabled +@@ -551,10 +562,12 @@ class _DeepEPDispatcherImplLowLatency(_DeepEPDispatcherImplBase): + buffer = self._get_buffer() + topk_weights, topk_ids = topk_output.topk_weights, topk_output.topk_ids + topk_ids = topk_ids.to(torch.int64) +- expected_m = ( +- hidden_states.shape[0] * buffer.group_size * topk_ids.shape[1] +- + self.num_experts +- ) // self.num_experts ++ # Use a correctness-preserving upper bound for per-expert token count. ++ # In the worst case, every rank routes all local tokens to the same expert. ++ expected_m = min( ++ hidden_states.shape[0] * buffer.group_size, ++ self.num_max_dispatch_tokens_per_rank * buffer.group_size, ++ ) + hidden_states, masked_m, event, hook = self._dispatch_core( + hidden_states, + topk_ids, +@@ -609,7 +622,9 @@ class _DeepEPDispatcherImplLowLatency(_DeepEPDispatcherImplBase): + input_global_scale = self.quant_config.get("input_global_scale", None) + if input_global_scale is not None: + use_nvfp4 = True +- elif not envs.SGLANG_DEEPEP_BF16_DISPATCH.get(): ++ elif not envs.SGLANG_DEEPEP_BF16_DISPATCH.get() and not self.quant_config.get( ++ "bf16_weights", False ++ ): + use_fp8 = True + + buffer = self._get_buffer() +diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py b/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py +index 4cbfed6f9..88b452744 100644 +--- a/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py ++++ b/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py +@@ -499,7 +499,7 @@ class CompressedTensorsConfig(QuantizationConfig): + ) + is_static = not weight_quant.dynamic + +- return is_channel_group and input_quant_none and is_symmetric and is_static ++ return is_channel_group and input_quant_none and is_static + + def _is_mxint4a16(self, weight_quant: BaseModel, input_quant: BaseModel) -> bool: + input_quant_none = input_quant is None +@@ -968,6 +968,9 @@ class CompressedTensorsFusedMoEMethod(FusedMoEMethodBase): + def process_weights_after_loading(self, layer: torch.nn.Module) -> None: + layer.scheme.process_weights_after_loading(layer) + ++ def restore_weights_before_loading(self, layer: torch.nn.Module) -> None: ++ layer.scheme.restore_weights_before_loading(layer) ++ + def create_weights( + self, + layer: torch.nn.Module, +diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py +index 6264f36d0..f0310e305 100644 +--- a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py ++++ b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py +@@ -17,7 +17,10 @@ from sglang.srt.layers.quantization.compressed_tensors.schemes import ( + CompressedTensorsMoEScheme, + ) + from sglang.srt.layers.quantization.gptq import gptq_marlin_moe_repack +-from sglang.srt.layers.quantization.marlin_utils import marlin_moe_permute_scales ++from sglang.srt.layers.quantization.marlin_utils import ( ++ marlin_moe_permute_scales, ++ moe_awq_to_marlin_zero_points, ++) + from sglang.srt.layers.quantization.utils import replace_parameter + from sglang.srt.utils import get_bool_env_var, is_cuda, is_hip, set_weight_attrs + +@@ -64,7 +67,7 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): + self.strategy = config.strategy + self.group_size = config.group_size + self.actorder = config.actorder +- assert config.symmetric, "Only symmetric quantization is supported for MoE" ++ self.sym = config.symmetric + + if not ( + self.quant_config.quant_format == CompressionFormat.pack_quantized.value +@@ -124,7 +127,7 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): + + # In the case where we have actorder/g_idx, + # we do not partition the w2 scales +- load_full_w2 = self.actorder and self.group_size != -1 ++ load_full_w2 = (self.actorder != "static") and self.group_size != -1 + + if load_full_w2: + w2_scales_size = intermediate_size_per_partition * layer.moe_tp_size +@@ -172,6 +175,32 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): + layer.register_parameter("w13_weight_shape", w13_weight_shape) + set_weight_attrs(w13_weight_shape, extra_weight_attrs) + ++ # add zero param ++ if not self.sym: ++ w13_qzeros = torch.nn.Parameter( ++ torch.empty( ++ num_experts, ++ num_groups_w13, ++ 2 * intermediate_size_per_partition // self.packed_factor, ++ dtype=torch.int32, ++ ), ++ requires_grad=False, ++ ) ++ layer.register_parameter("w13_weight_zero_point", w13_qzeros) ++ set_weight_attrs(w13_qzeros, extra_weight_attrs) ++ ++ w2_qzeros = torch.nn.Parameter( ++ torch.empty( ++ num_experts, ++ num_groups_w2, ++ hidden_size // self.packed_factor, ++ dtype=torch.int32, ++ ), ++ requires_grad=False, ++ ) ++ layer.register_parameter("w2_weight_zero_point", w2_qzeros) ++ set_weight_attrs(w2_qzeros, extra_weight_attrs) ++ + w13_g_idx = torch.nn.Parameter( + torch.empty( + num_experts, +@@ -225,11 +254,14 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): + + # Force record: these are the target GPTQ shapes for rollback. + layer._original_shapes["w13_weight_packed"] = tuple(w13_weight.shape) +- layer._original_shapes["w2_weight_packed"] = tuple(w2_weight.shape) ++ layer._original_shapes["w13_weight_scale"] = tuple(w13_scale.shape) ++ if not self.sym: ++ layer._original_shapes["w13_weight_zero_point"] = w13_qzeros.shape + +- # Also record the shapes of the scales. ++ layer._original_shapes["w2_weight_packed"] = tuple(w2_weight.shape) + layer._original_shapes["w2_weight_scale"] = tuple(w2_scale.shape) +- layer._original_shapes["w13_weight_scale"] = tuple(w13_scale.shape) ++ if not self.sym: ++ layer._original_shapes["w2_weight_zero_point"] = tuple(w2_qzeros.shape) + + def process_weights_after_loading(self, layer: torch.nn.Module) -> None: + +@@ -334,6 +366,24 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): + ) + replace_tensor("w2_weight_scale", marlin_w2_scales) + ++ # Repack zero ++ if not self.sym: ++ marlin_w13_zp = moe_awq_to_marlin_zero_points( ++ layer.w13_weight_zero_point, ++ size_k=layer.w13_weight_zero_point.shape[1], ++ size_n=layer.w13_weight_zero_point.shape[2] * self.packed_factor, ++ num_bits=self.num_bits, ++ ) ++ replace_tensor("w13_weight_zero_point", marlin_w13_zp) ++ ++ marlin_w2_zp = moe_awq_to_marlin_zero_points( ++ layer.w2_weight_zero_point, ++ size_k=layer.w2_weight_zero_point.shape[1], ++ size_n=layer.w2_weight_zero_point.shape[2] * self.packed_factor, ++ num_bits=self.num_bits, ++ ) ++ replace_tensor("w2_weight_zero_point", marlin_w2_zp) ++ + layer.is_marlin_converted = True + + def restore_weights_before_loading(self, layer: torch.nn.Module): +@@ -399,6 +449,8 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): + g_idx2=layer.w2_weight_g_idx, + sort_indices1=layer.w13_g_idx_sort_indices, + sort_indices2=layer.w2_g_idx_sort_indices, ++ w1_zeros=layer.w13_weight_zero_point if not self.sym else None, ++ w2_zeros=layer.w2_weight_zero_point if not self.sym else None, + num_bits=self.num_bits, + 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..3b6a8d254 100644 +--- a/python/sglang/srt/layers/rotary_embedding.py ++++ b/python/sglang/srt/layers/rotary_embedding.py +@@ -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]: ++ assert ( ++ fused_set_kv_buffer_arg is None ++ ), "fused_set_kv_buffer_arg is not supported for npu implementation" + # TODO: remove this when npu_mrope supports QNumHeads * QHeadSize > 4096 + assert ( + fused_set_kv_buffer_arg is None +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 ++++ b/python/sglang/srt/managers/io_struct.py +@@ -1403,6 +1403,20 @@ class UpdateWeightsFromIPCReqOutput(BaseReq): + message: str + + ++@dataclass ++class PostProcessWeightsReqInput(BaseReq): ++ # Whether to restore weights before loading new weights ++ restore_weights_before_load: bool = False ++ # Whether to enable quantization post-processing ++ post_process_quantization: bool = False ++ ++ ++@dataclass ++class PostProcessWeightsReqOutput(BaseReq): ++ success: bool ++ message: str ++ ++ + @dataclass + class InitWeightsSendGroupForRemoteInstanceReqOutput(BaseReq): + success: bool +@@ -1802,6 +1816,10 @@ class GetLoadReqOutput(BaseReq): + num_waiting_reqs: int + num_tokens: int + ts_tic: float ++ # Per-queue breakdown: list of {name, num_reqs, num_tokens, reqs: [{rid, seqlen, input_len, output_len}]} ++ queue_details: Optional[List[Dict[str, Any]]] = None ++ # Running batch info ++ running_details: Optional[Dict[str, Any]] = None + + + @dataclass +diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py +index c07995798..dd8ca7167 100644 +--- a/python/sglang/srt/managers/schedule_batch.py ++++ b/python/sglang/srt/managers/schedule_batch.py +@@ -1869,7 +1869,10 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): + while first_iter or ( + not self.check_decode_mem(selected_indices=sorted_indices) + ): +- if len(sorted_indices) == 1: ++ # We should allow all requests to be retracted in decode disaggregation mode ++ # because there call be prealloc prefill requests. ++ num_minimum_reqs = 0 if server_args.disaggregation_mode == "decode" else 1 ++ if len(sorted_indices) == num_minimum_reqs: + # Always keep at least one request + break + +diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py +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 ( + OpenSessionReqInput, + OpenSessionReqOutput, + PauseGenerationReqInput, ++ PostProcessWeightsReqInput, + ProfileReq, + ReleaseMemoryOccupationReqInput, + ResumeMemoryOccupationReqInput, +@@ -1063,6 +1064,7 @@ class Scheduler( + ), + (UpdateWeightsFromTensorReqInput, self.update_weights_from_tensor), + (UpdateWeightsFromIPCReqInput, self.update_weights_from_ipc), ++ (PostProcessWeightsReqInput, self.post_process_weights), + (GetWeightsByNameReqInput, self.get_weights_by_name), + (ReleaseMemoryOccupationReqInput, self.release_memory_occupation), + (ResumeMemoryOccupationReqInput, self.resume_memory_occupation), +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 ++++ b/python/sglang/srt/managers/scheduler_metrics_mixin.py +@@ -609,12 +609,54 @@ class SchedulerMetricsMixin: + num_tokens += sum(req.seqlen for queue in waiting_queues for req in queue) + num_waiting_reqs = sum(len(queue) for queue in waiting_queues) + ++ # Collect per-queue details ++ queue_names = ["waiting_queue"] ++ if self.disaggregation_mode == DisaggregationMode.PREFILL: ++ queue_names.append("bootstrap_queue") ++ elif self.disaggregation_mode == DisaggregationMode.DECODE: ++ queue_names.append("prealloc_queue") ++ queue_names.append("transfer_queue") ++ queue_names.append("retracted_queue") ++ ++ queue_details = [] ++ for name, queue in zip(queue_names, waiting_queues): ++ reqs_info = [] ++ for req in queue: ++ reqs_info.append( ++ { ++ "seqlen": req.seqlen, ++ } ++ ) ++ queue_details.append( ++ { ++ "name": name, ++ "num_reqs": len(queue), ++ "num_tokens": sum(r["seqlen"] for r in reqs_info), ++ "reqs": reqs_info, ++ } ++ ) ++ ++ # Collect running batch details ++ running_reqs_info = [] ++ for req in self.running_batch.reqs: ++ running_reqs_info.append( ++ { ++ "seqlen": req.seqlen, ++ } ++ ) ++ running_details = { ++ "num_reqs": len(self.running_batch.reqs), ++ "reqs": running_reqs_info, ++ } ++ + return GetLoadReqOutput( + dp_rank=self.dp_rank, + num_reqs=len(self.running_batch.reqs) + num_waiting_reqs, + num_waiting_reqs=num_waiting_reqs, + num_tokens=num_tokens, + ts_tic=time.perf_counter(), ++ queue_details=queue_details, ++ running_details=running_details, + ) + + def get_loads(self: Scheduler, req: GetLoadsReqInput = None) -> GetLoadsReqOutput: +diff --git a/python/sglang/srt/managers/scheduler_output_processor_mixin.py b/python/sglang/srt/managers/scheduler_output_processor_mixin.py +index 482bc6ca6..857cfa6a3 100644 +--- a/python/sglang/srt/managers/scheduler_output_processor_mixin.py ++++ b/python/sglang/srt/managers/scheduler_output_processor_mixin.py +@@ -1134,7 +1134,7 @@ class SchedulerOutputProcessorMixin: + req.log_time_stats() + + # Send to detokenizer +- if reqs or is_idle_batch: ++ if rids or is_idle_batch: + if self.model_config.is_multimodal_gen: + 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..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 ( + get_attention_dp_rank, + get_attention_dp_size, + is_dp_attention_enabled, ++ set_is_extend_in_batch, + ) + from sglang.srt.managers.schedule_batch import Req, ScheduleBatch + from sglang.srt.managers.utils import ( +@@ -224,7 +225,28 @@ class SchedulerPPMixin: + + self.process_prefill_chunk() + batch = self.get_new_batch_prefill() +- batch = self.maybe_prepare_mlp_sync_batch(batch) ++ need_mlp_sync = self.require_mlp_sync ++ skipped_mlp_sync = False ++ if ( ++ need_mlp_sync ++ and self.disaggregation_mode == DisaggregationMode.PREFILL ++ and self.server_args.enable_nsa_prefill_context_parallel ++ and self.pp_size > 1 ++ ): ++ # In PD prefill CP+PP, MLP sync all_gather can deadlock on idle micro-batches. ++ # Skip MLP sync here because decode-side MLP gather is not involved in this path. ++ need_mlp_sync = False ++ skipped_mlp_sync = True ++ batch = self.maybe_prepare_mlp_sync_batch( ++ batch, need_sync=need_mlp_sync ++ ) ++ if skipped_mlp_sync: ++ # MLP sync was skipped but set_is_extend_in_batch is still needed ++ # by the deepep dispatcher (called in model forward). ++ is_extend = ( ++ batch.forward_mode.is_extend() if batch is not None else False ++ ) ++ set_is_extend_in_batch(is_extend) + self.mbs[mb_id] = batch + self.running_mbs[mb_id] = self.running_batch + +@@ -288,6 +310,11 @@ class SchedulerPPMixin: + next_batch_result, + ) + self.last_mbs[next_mb_id] = self.mbs[next_mb_id] ++ if self.current_scheduler_metrics_enabled: ++ self.log_prefill_stats( ++ prefill_stats=self.mbs[next_mb_id].prefill_stats, ++ can_run_cuda_graph=next_batch_result.can_run_cuda_graph, ++ ) + + if tmbs[next_mb_id] is not None: + self.process_disagg_prefill_inflight_queue(next_release_rids) +@@ -524,6 +551,11 @@ class SchedulerPPMixin: + self.last_rank_comm_queue: deque[Tuple[torch.cuda.Event, PPProxyTensors]] = ( + deque() + ) ++ # 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.send_req_work = [] + self.send_proxy_work = [] +@@ -859,31 +891,39 @@ class SchedulerPPMixin: + + def _pp_send_pyobj_to_next_stage(self: Scheduler, data, async_send: bool = False): + p2p_work = [] +- if self.attn_tp_rank == 0: +- dp_offset = self.attn_dp_rank * self.attn_tp_size ++ if self.attn_tp_rank == 0 and self.attn_cp_rank == 0: ++ lane_offset = self.attn_dp_rank * self.attn_tp_size + p2p_work = point_to_point_pyobj( + data, +- self.pp_rank * self.tp_size + dp_offset, ++ self.pp_rank * self.tp_size + lane_offset, + self.world_group.cpu_group, +- self.pp_rank * self.tp_size + dp_offset, +- ((self.pp_rank + 1) % self.pp_size) * self.tp_size + dp_offset, ++ self.pp_rank * self.tp_size + lane_offset, ++ ((self.pp_rank + 1) % self.pp_size) * self.tp_size + lane_offset, + async_send=async_send, + ) + return p2p_work + + def _pp_recv_pyobj_from_prev_stage(self: Scheduler): +- if self.attn_tp_rank == 0: +- dp_offset = self.attn_dp_rank * self.attn_tp_size ++ if self.attn_tp_rank == 0 and self.attn_cp_rank == 0: ++ lane_offset = self.attn_dp_rank * self.attn_tp_size + data = point_to_point_pyobj( + [], +- self.pp_rank * self.tp_size + dp_offset, ++ self.pp_rank * self.tp_size + lane_offset, + self.world_group.cpu_group, +- ((self.pp_rank - 1) % self.pp_size) * self.tp_size + dp_offset, +- self.pp_rank * self.tp_size + dp_offset, ++ ((self.pp_rank - 1) % self.pp_size) * self.tp_size + lane_offset, ++ self.pp_rank * self.tp_size + lane_offset, + ) + else: + data = None + ++ if self.attn_cp_size > 1: ++ data = broadcast_pyobj( ++ data, ++ self.attn_cp_group.rank, ++ self.attn_cp_cpu_group, ++ src=self.attn_cp_group.ranks[0], ++ ) ++ + if self.attn_tp_size > 1: + data = broadcast_pyobj( + data, +@@ -1004,8 +1044,13 @@ class SchedulerPPMixin: + pp_outputs_to_send.tensors, + async_send=True, + ) +- # send the outputs from the last round to let the next stage worker run post processing +- if not self.pp_group.is_last_rank: ++ # Store locally so the last rank can process its own batch result ++ # without receiving from the second-to-last rank (avoids deadlock). ++ self.last_rank_local_result_queue.append((q_event, pp_outputs_to_send)) ++ elif self.pp_rank != self.pp_size - 2: ++ # Forward output through the chain: PP0→PP1→...→PP(last-2). ++ # The second-to-last rank does NOT forward to the last rank because ++ # the last rank uses last_rank_local_result_queue instead of receiving. + 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,38 @@ class SchedulerPPMixin: + ) + + if mbs[next_mb_id] is not None: +- with torch.profiler.record_function("recv_res_dict_from_prev_stage"): +- next_pp_outputs = None ++ if self.pp_group.is_last_rank: ++ # Last rank: use the locally-stored output instead of receiving ++ # from the second-to-last rank. Receiving would cause a deadlock ++ # because the chain PP_last→PP0→...→PP(last-2)→PP_last requires ++ # PP0 to have pp_outputs ready, which it doesn't on the first batch. + if not mbs[next_mb_id].forward_mode.is_prebuilt(): +- next_pp_outputs = PPProxyTensors( +- self._pp_recv_dict_from_prev_stage() +- ) +- if not mbs[next_mb_id].forward_mode.is_prebuilt(): +- with self.copy_stream_ctx: +- 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()) ++ with self.copy_stream_ctx: ++ torch.cuda.current_stream().wait_event(q_event) ++ 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 ++ ) ++ d2h_event = torch.cuda.Event() ++ d2h_event.record(torch.cuda.current_stream()) ++ else: ++ with torch.profiler.record_function("recv_res_dict_from_prev_stage"): ++ next_pp_outputs = None ++ if not mbs[next_mb_id].forward_mode.is_prebuilt(): ++ next_pp_outputs = PPProxyTensors( ++ self._pp_recv_dict_from_prev_stage() ++ ) ++ if not mbs[next_mb_id].forward_mode.is_prebuilt(): ++ with self.copy_stream_ctx: ++ 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 ++ ) ++ d2h_event = torch.cuda.Event() ++ d2h_event.record(torch.cuda.current_stream()) + + return next_pp_outputs, batch_result, d2h_event, send_output_work + +@@ -1085,9 +1148,12 @@ class SchedulerPPMixin: + """ + Used by PP, get the required rids with the given poll statuses. + """ ++ gloo_group = self.attn_tp_cpu_group ++ if self.attn_cp_size > 1: ++ gloo_group = self.tp_cpu_group + polls = poll_and_all_reduce( + [req.disagg_kv_sender if is_send else req.kv_receiver for req in req_queue], +- self.attn_tp_cpu_group, ++ gloo_group, + ) + 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 ++++ b/python/sglang/srt/managers/scheduler_update_weights_mixin.py +@@ -12,6 +12,7 @@ from sglang.srt.constants import ( + GPU_MEMORY_TYPE_KV_CACHE, + GPU_MEMORY_TYPE_WEIGHTS, + ) ++from sglang.srt.disaggregation.utils import DisaggregationMode + from sglang.srt.managers.io_struct import ( + CheckWeightsReqInput, + CheckWeightsReqOutput, +@@ -21,6 +22,8 @@ from sglang.srt.managers.io_struct import ( + GetWeightsByNameReqOutput, + InitWeightsUpdateGroupReqInput, + InitWeightsUpdateGroupReqOutput, ++ PostProcessWeightsReqInput, ++ PostProcessWeightsReqOutput, + ReleaseMemoryOccupationReqInput, + ReleaseMemoryOccupationReqOutput, + ResumeMemoryOccupationReqInput, +@@ -114,6 +117,11 @@ class SchedulerUpdateWeightsMixin: + torch.distributed.barrier(group=self.tp_cpu_group) + return UpdateWeightsFromIPCReqOutput(success, message) + ++ def post_process_weights(self, recv_req: PostProcessWeightsReqInput): ++ """Optional post-processing for updated weights (e.g., Marlin conversion).""" ++ success, message = self.tp_worker.post_process_weights(recv_req) ++ return PostProcessWeightsReqOutput(success, message) ++ + def get_weights_by_name(self: Scheduler, recv_req: GetWeightsByNameReqInput): + parameter = self.tp_worker.get_weights_by_name(recv_req) + return GetWeightsByNameReqOutput(parameter) +@@ -137,6 +145,15 @@ class SchedulerUpdateWeightsMixin: + self.memory_saver_adapter.pause(GPU_MEMORY_TYPE_KV_CACHE) + self.flush_cache() + ++ if self.disaggregation_mode == DisaggregationMode.DECODE: ++ if hasattr(self, "disagg_decode_transfer_queue"): ++ self.disagg_decode_transfer_queue.release_memory_occupation() ++ if hasattr(self, "disagg_decode_prealloc_queue"): ++ self.disagg_decode_prealloc_queue.release_memory_occupation() ++ elif self.disaggregation_mode == DisaggregationMode.PREFILL: ++ if hasattr(self, "disagg_prefill_bootstrap_queue"): ++ self.disagg_prefill_bootstrap_queue.release_memory_occupation() ++ + if GPU_MEMORY_TYPE_WEIGHTS in tags: + self.stashed_model_static_state = _export_static_state( + self.tp_worker.model_runner.model +@@ -177,6 +194,15 @@ class SchedulerUpdateWeightsMixin: + if GPU_MEMORY_TYPE_KV_CACHE in tags: + self.memory_saver_adapter.resume(GPU_MEMORY_TYPE_KV_CACHE) + ++ if self.disaggregation_mode == DisaggregationMode.DECODE: ++ if hasattr(self, "disagg_decode_transfer_queue"): ++ self.disagg_decode_transfer_queue.resume_memory_occupation() ++ if hasattr(self, "disagg_decode_prealloc_queue"): ++ self.disagg_decode_prealloc_queue.resume_memory_occupation() ++ elif self.disaggregation_mode == DisaggregationMode.PREFILL: ++ if hasattr(self, "disagg_prefill_bootstrap_queue"): ++ self.disagg_prefill_bootstrap_queue.resume_memory_occupation() ++ + return ResumeMemoryOccupationReqOutput() + + def check_weights(self: Scheduler, recv_req: CheckWeightsReqInput): +diff --git a/python/sglang/srt/managers/tokenizer_communicator_mixin.py b/python/sglang/srt/managers/tokenizer_communicator_mixin.py +index f2ffa9909..6e4d1d460 100644 +--- a/python/sglang/srt/managers/tokenizer_communicator_mixin.py ++++ b/python/sglang/srt/managers/tokenizer_communicator_mixin.py +@@ -59,6 +59,8 @@ from sglang.srt.managers.io_struct import ( + LoadLoRAAdapterReqOutput, + LoRAUpdateOutput, + OpenSessionReqInput, ++ PostProcessWeightsReqInput, ++ PostProcessWeightsReqOutput, + ProfileReq, + ProfileReqOutput, + ProfileReqType, +@@ -187,6 +189,9 @@ class TokenizerCommunicatorMixin: + self.update_weights_from_ipc_communicator = _Communicator( + self.send_to_scheduler, server_args.dp_size + ) ++ self.post_process_weights_communicator = _Communicator( ++ self.send_to_scheduler, server_args.dp_size ++ ) + self.get_weights_by_name_communicator = _Communicator( + self.send_to_scheduler, server_args.dp_size + ) +@@ -272,6 +277,10 @@ class TokenizerCommunicatorMixin: + UpdateWeightsFromIPCReqOutput, + self.update_weights_from_ipc_communicator.handle_recv, + ), ++ ( ++ PostProcessWeightsReqOutput, ++ self.post_process_weights_communicator.handle_recv, ++ ), + ( + GetWeightsByNameReqOutput, + self.get_weights_by_name_communicator.handle_recv, +@@ -522,6 +531,17 @@ class TokenizerCommunicatorMixin: + + return success, message + ++ async def post_process_weights( ++ self: TokenizerManager, ++ obj: PostProcessWeightsReqInput, ++ request: Optional[fastapi.Request] = None, ++ ) -> Tuple[bool, str]: ++ """Trigger post-processing hooks for weights after loading (e.g., Marlin conversion).""" ++ self.auto_create_handle_loop() ++ async with self.model_update_lock.writer_lock: ++ results = await self.post_process_weights_communicator(obj) ++ return _Communicator.merge_results(results) ++ + async def init_weights_send_group_for_remote_instance( + self, + obj: InitWeightsSendGroupForRemoteInstanceReqInput, +diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py +index 0914a5230..cce2d8a2b 100644 +--- a/python/sglang/srt/managers/tokenizer_manager.py ++++ b/python/sglang/srt/managers/tokenizer_manager.py +@@ -324,8 +324,12 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi + context, zmq.PULL, port_args.tokenizer_ipc_name, True + ) + if self.server_args.tokenizer_worker_num == 1: ++ self.send_to_scheduler_context = zmq.Context(1) + self.send_to_scheduler = get_zmq_socket( +- context, zmq.PUSH, port_args.scheduler_input_ipc_name, True ++ self.send_to_scheduler_context, ++ zmq.PUSH, ++ port_args.scheduler_input_ipc_name, ++ True, + ) + else: + from sglang.srt.managers.multi_tokenizer_mixin import SenderWrapper +@@ -1327,7 +1331,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi + async with self.is_pause_cond: + self.is_pause = True + if obj.mode != "abort": +- await self.send_to_scheduler.send_pyobj(obj) ++ self.send_to_scheduler.send_pyobj(obj) + else: + # we are using the model_update_lock to check if there is still on-going requests. + while True: +@@ -1341,7 +1345,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi + async def continue_generation(self, obj: ContinueGenerationReqInput): + async with self.is_pause_cond: + self.is_pause = False +- await self.send_to_scheduler.send_pyobj(obj) ++ self.send_to_scheduler.send_pyobj(obj) + self.is_pause_cond.notify_all() + + async def update_weights_from_disk( +diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py +index 86b009df4..16ebd52ae 100644 +--- a/python/sglang/srt/managers/tp_worker.py ++++ b/python/sglang/srt/managers/tp_worker.py +@@ -29,6 +29,7 @@ from sglang.srt.managers.io_struct import ( + InitWeightsUpdateGroupReqInput, + LoadLoRAAdapterFromTensorsReqInput, + LoadLoRAAdapterReqInput, ++ PostProcessWeightsReqInput, + SendWeightsToRemoteInstanceReqInput, + UnloadLoRAAdapterReqInput, + UpdateWeightFromDiskReqInput, +@@ -168,6 +169,11 @@ class BaseTpWorker(ABC): + success, message = self.model_runner.update_weights_from_ipc(recv_req) + return success, message + ++ def post_process_weights(self, recv_req: PostProcessWeightsReqInput): ++ """Perform optional post-processing on the updated model weights (e.g., Marlin conversion).""" ++ success, message = self.model_runner.post_process_weights(recv_req) ++ return success, message ++ + def get_weights_by_name(self, recv_req: GetWeightsByNameReqInput): + parameter = self.model_runner.get_weights_by_name( + recv_req.name, recv_req.truncate_size +diff --git a/python/sglang/srt/mem_cache/allocator.py b/python/sglang/srt/mem_cache/allocator.py +index fa08bb66a..fa539315c 100644 +--- a/python/sglang/srt/mem_cache/allocator.py ++++ b/python/sglang/srt/mem_cache/allocator.py +@@ -347,6 +347,84 @@ def alloc_decode_kernel( + tl.store(out_indices + pid, page * page_size) + + ++def alloc_extend_torch_fallback( ++ prefix_lens_cpu: torch.Tensor, ++ seq_lens_cpu: torch.Tensor, ++ last_loc: torch.Tensor, ++ free_pages: torch.Tensor, ++ out_indices: torch.Tensor, ++ page_size: int, ++ debug_mode: bool = False, ++): ++ extend_lens_cpu = (seq_lens_cpu - prefix_lens_cpu).to(torch.int64) ++ if extend_lens_cpu.numel() == 0: ++ return ++ ++ output_start_locs_cpu = torch.cumsum(extend_lens_cpu, dim=0) - extend_lens_cpu ++ num_pages_after = (seq_lens_cpu + page_size - 1) // page_size ++ num_pages_before = (prefix_lens_cpu + page_size - 1) // page_size ++ num_new_pages_cpu = num_pages_after - num_pages_before ++ page_start_locs_cpu = torch.cumsum(num_new_pages_cpu, dim=0) - num_new_pages_cpu ++ ++ total_new_pages = int(num_new_pages_cpu.sum().item()) ++ if total_new_pages > free_pages.numel(): ++ return ++ ++ if debug_mode: ++ assert int(extend_lens_cpu.sum().item()) == out_indices.numel() ++ ++ prefix_lens_list = prefix_lens_cpu.tolist() ++ seq_lens_list = seq_lens_cpu.tolist() ++ extend_lens_list = extend_lens_cpu.tolist() ++ out_start_list = output_start_locs_cpu.tolist() ++ page_start_list = page_start_locs_cpu.tolist() ++ num_new_pages_list = num_new_pages_cpu.tolist() ++ ++ device = out_indices.device ++ dtype = out_indices.dtype ++ offsets_page = torch.arange(page_size, device=device, dtype=dtype) ++ ++ for i, extend_len in enumerate(extend_lens_list): ++ if extend_len == 0: ++ continue ++ ++ pre_len = prefix_lens_list[i] ++ seq_len = seq_lens_list[i] ++ out_start = out_start_list[i] ++ page_start = page_start_list[i] ++ num_new_pages = num_new_pages_list[i] ++ ++ pre_mod = pre_len % page_size ++ part1 = min(extend_len, page_size - pre_mod) if pre_mod != 0 else 0 ++ if part1: ++ start_val = last_loc[i] + 1 ++ out_indices[out_start : out_start + part1] = start_val + torch.arange( ++ part1, device=device, dtype=dtype ++ ) ++ if part1 == extend_len: ++ continue ++ ++ ceil_pre_pages = (pre_len + page_size - 1) // page_size ++ full_pages_after = seq_len // page_size ++ num_full_pages = full_pages_after - ceil_pre_pages ++ if num_full_pages < 0: ++ num_full_pages = 0 ++ part2 = num_full_pages * page_size ++ if part2: ++ pages = free_pages[page_start : page_start + num_full_pages] ++ full_indices = (pages[:, None] * page_size + offsets_page).reshape(-1) ++ out_indices[out_start + part1 : out_start + part1 + part2] = full_indices ++ if part1 + part2 == extend_len: ++ continue ++ ++ part3 = extend_len - part1 - part2 ++ if part3: ++ last_page = free_pages[page_start + num_new_pages - 1] ++ out_indices[out_start + part1 + part2 : out_start + extend_len] = ( ++ last_page * page_size + torch.arange(part3, device=device, dtype=dtype) ++ ) ++ ++ + class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): + """ + An allocator managing the indices to kv cache data. +@@ -411,7 +489,7 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): + + self.seen_max_num_extend_tokens_next_power_of_2 = max( + self.seen_max_num_extend_tokens_next_power_of_2, +- min(tl.core.TRITON_MAX_TENSOR_NUMEL, next_power_of_2(extend_num_tokens)), ++ min(65536, next_power_of_2(extend_num_tokens)), + ) + + bs = len(prefix_lens) +@@ -424,7 +502,7 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): + (extend_num_tokens,), dtype=torch.int64, device=self.device + ) + +- if extend_num_tokens < tl.core.TRITON_MAX_TENSOR_NUMEL: ++ if extend_num_tokens < 65536: + alloc_extend_kernel[(bs,)]( + prefix_lens, + seq_lens, +diff --git a/python/sglang/srt/mem_cache/hiradix_cache.py b/python/sglang/srt/mem_cache/hiradix_cache.py +index d7cd472a9..81fae740f 100644 +--- a/python/sglang/srt/mem_cache/hiradix_cache.py ++++ b/python/sglang/srt/mem_cache/hiradix_cache.py +@@ -76,6 +76,7 @@ class HiRadixCache(RadixCache): + allocator_type=server_args.hicache_storage_backend, + ) + elif isinstance(self.kv_cache, NSATokenToKVPool): ++ # Check NSA before MLA since NSATokenToKVPool is a subclass of MLATokenToKVPool + self.token_to_kv_pool_host = NSATokenToKVPoolHost( + self.kv_cache, + server_args.hicache_ratio, +@@ -94,7 +95,7 @@ class HiRadixCache(RadixCache): + allocator_type=server_args.hicache_storage_backend, + ) + else: +- raise ValueError(f"HiRadixCache only supports MHA and MLA yet") ++ raise ValueError(f"HiRadixCache only supports MHA and MLA and NSA yet") + + self.tp_group = params.tp_cache_group + self.tp_world_size = torch.distributed.get_world_size(group=self.tp_group) +@@ -750,9 +751,8 @@ class HiRadixCache(RadixCache): + self._update_leaf_status(node) + self._update_host_leaf_status(node) + if node.parent is None: +- assert ( +- node is self.root_node +- ), f"This request holds the node from another tree" ++ # Node belongs to a stale (flushed) tree — stop traversal gracefully. ++ break + node = node.parent + return delta + +@@ -827,6 +827,7 @@ class HiRadixCache(RadixCache): + self._update_host_leaf_status(node) + # update leaf status for the parent because the node is evicted + self._update_leaf_status(node.parent) ++ self._update_host_leaf_status(node.parent) + return num_evicted + + def _evict_regular(self, node: TreeNode): +@@ -1330,6 +1331,7 @@ class HiRadixCache(RadixCache): + self._update_host_leaf_status(node) + # update parent status as a new leaf is added into device + self._update_leaf_status(node.parent) ++ self._update_host_leaf_status(node.parent) + else: + self._inc_hit_count(node, chunked) + total_prefix_length += prefix_len +@@ -1345,6 +1347,7 @@ class HiRadixCache(RadixCache): + self._update_host_leaf_status(new_node) + # update parent status as a new leaf is added into device + self._update_leaf_status(new_node.parent) ++ self._update_host_leaf_status(new_node.parent) + else: + self._inc_hit_count(new_node, chunked) + total_prefix_length += prefix_len +diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py +index 1d917137c..669e5c518 100644 +--- a/python/sglang/srt/mem_cache/memory_pool.py ++++ b/python/sglang/srt/mem_cache/memory_pool.py +@@ -1777,9 +1777,12 @@ class NSATokenToKVPool(MLATokenToKVPool): + else: + assert self.page_size == 64 + with ( +- torch.cuda.use_mem_pool(self.custom_mem_pool) +- if self.custom_mem_pool +- else nullcontext() ++ ( ++ torch.cuda.use_mem_pool(self.custom_mem_pool) ++ if self.custom_mem_pool ++ else nullcontext() ++ ), ++ self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE), + ): + self.index_k_with_scale_buffer = [ + torch.zeros( +@@ -1801,6 +1804,11 @@ class NSATokenToKVPool(MLATokenToKVPool): + ) + for _ in range(layer_num) + ] ++ self.index_k_with_scale_buffer_ptrs = torch.tensor( ++ [x.data_ptr() for x in self.index_k_with_scale_buffer], ++ dtype=torch.uint64, ++ device=self.device, ++ ) + self._finalize_allocation_log(size) + + def get_index_k_with_scale_buffer(self, layer_id: int) -> torch.Tensor: +@@ -1876,6 +1884,50 @@ class NSATokenToKVPool(MLATokenToKVPool): + ] + return data_ptrs, data_lens, item_lens + ++ def get_cpu_copy(self, indices): ++ # First, save the kv_buffer (inherited from MLATokenToKVPool) ++ kv_cache_cpu = super().get_cpu_copy(indices) ++ ++ # Additionally, save the index_k_with_scale_buffer (page-indexed) ++ page_indices = indices[:: self.page_size] // self.page_size ++ torch.cuda.synchronize() ++ index_k_cpu = [] ++ chunk_size = self.cpu_offloading_chunk_size ++ # Convert chunk_size from token-level to page-level ++ page_chunk_size = max(1, chunk_size // self.page_size) ++ for layer_id in range(self.layer_num): ++ index_k_cpu.append([]) ++ for i in range(0, len(page_indices), page_chunk_size): ++ chunk_page_indices = page_indices[i : i + page_chunk_size] ++ idx_cpu = self.index_k_with_scale_buffer[layer_id][ ++ chunk_page_indices ++ ].to("cpu", non_blocking=True) ++ index_k_cpu[-1].append(idx_cpu) ++ torch.cuda.synchronize() ++ ++ return {"kv": kv_cache_cpu, "index_k": index_k_cpu} ++ ++ def load_cpu_copy(self, kv_cache_cpu_dict, indices): ++ # Restore the kv_buffer (inherited from MLATokenToKVPool) ++ super().load_cpu_copy(kv_cache_cpu_dict["kv"], indices) ++ ++ # Restore the index_k_with_scale_buffer (page-indexed) ++ page_indices = indices[:: self.page_size] // self.page_size ++ index_k_cpu = kv_cache_cpu_dict["index_k"] ++ torch.cuda.synchronize() ++ chunk_size = self.cpu_offloading_chunk_size ++ page_chunk_size = max(1, chunk_size // self.page_size) ++ for layer_id in range(self.layer_num): ++ for i in range(0, len(page_indices), page_chunk_size): ++ chunk_page_indices = page_indices[i : i + page_chunk_size] ++ idx_cpu = index_k_cpu[layer_id][i // page_chunk_size] ++ assert idx_cpu.shape[0] == len(chunk_page_indices) ++ idx_chunk = idx_cpu.to( ++ self.index_k_with_scale_buffer[0].device, non_blocking=True ++ ) ++ self.index_k_with_scale_buffer[layer_id][chunk_page_indices] = idx_chunk ++ torch.cuda.synchronize() ++ + def get_kv_size_bytes(self): + 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..8e799196a 100644 +--- a/python/sglang/srt/mem_cache/radix_cache.py ++++ b/python/sglang/srt/mem_cache/radix_cache.py +@@ -495,7 +495,17 @@ class RadixCache(BasePrefixCache): + if self.disable: + return + +- token_ids = req.fill_ids ++ # Limit to kv_committed_len to avoid including tokens (e.g., the just-generated ++ # token in disagg prefill) that don't have computed KV yet. If fill_ids is longer ++ # than kv_committed_len, the extra tokens would produce stale values (0 from ++ # 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 ++ ) + kv_indices = self.req_to_token_pool.req_to_token[ + req.req_pool_idx, : len(token_ids) + ] +@@ -619,9 +629,8 @@ class RadixCache(BasePrefixCache): + node.lock_ref -= 1 + self._update_leaf_status(node) + if node.parent is None: +- assert ( +- node is self.root_node +- ), f"This request holds the node from another tree" ++ # Node belongs to a stale (flushed) tree — stop traversal gracefully. ++ break + node = node.parent + return delta + +diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py +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,12 @@ class ModelRunner(ModelRunnerKVCacheMixin): + self.forward_stream = torch.get_device_module(self.device).Stream() + + # CPU offload +- set_offloader(create_offloader_from_server_args(server_args, dp_rank=dp_rank)) ++ # 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) ++ ) + + self._weight_checker = WeightChecker(model_runner=self) + +@@ -600,7 +605,8 @@ class ModelRunner(ModelRunnerKVCacheMixin): + ) + + # Init routed experts capturer +- self.init_routed_experts_capturer() ++ if not self.is_draft_worker: ++ self.init_routed_experts_capturer() + + if self.device == "cuda" or self.device == "musa": + self.init_cublas() +@@ -2429,11 +2435,19 @@ class ModelRunner(ModelRunnerKVCacheMixin): + output.expert_distribution_metrics = recorder_outputs.get("metrics") + + # Copy cached routing experts' buffers back to CPU cache +- get_global_experts_capturer().on_forward_end( +- forward_batch=forward_batch, +- can_run_graph=output.can_run_graph, +- cuda_graph_batch=getattr(self.graph_runner, "bs", None), +- ) ++ if not self.is_draft_worker: ++ # In speculative decoding, num_tokens_per_bs > 1, so we need to pass ++ # the actual number of tokens per dp rank in cuda graph, not batch size. ++ cuda_graph_num_tokens = None ++ if getattr(self.graph_runner, "bs", None): ++ cuda_graph_num_tokens = ( ++ self.graph_runner.bs * self.graph_runner.num_tokens_per_bs ++ ) ++ get_global_experts_capturer().on_forward_end( ++ forward_batch=forward_batch, ++ can_run_graph=output.can_run_graph, ++ cuda_graph_batch=cuda_graph_num_tokens, ++ ) + + if self.eplb_manager is not None: + self.eplb_manager.on_forward_pass_end() +@@ -2664,6 +2678,42 @@ class ModelRunner(ModelRunnerKVCacheMixin): + device=self.device, + ) + ++ def post_process_weights(self, recv_req): ++ """ ++ Execute post-processing logic for model weights, such as Marlin quantization format conversion. ++ """ ++ from sglang.srt.model_loader.loader import device_loading_context ++ ++ target_device = torch.device("cuda", torch.cuda.current_device()) ++ ++ if recv_req.restore_weights_before_load: ++ for _, module in self.model.named_modules(): ++ quant_method = getattr(module, "quant_method", None) ++ ++ # Check if the module supports restoring weights ++ if quant_method is not None and hasattr( ++ quant_method, "restore_weights_before_loading" ++ ): ++ ++ with device_loading_context(module, target_device): ++ quant_method.restore_weights_before_loading(module) ++ ++ if recv_req.post_process_quantization: ++ # Iterate through all modules to apply specific post-loading processing ++ for _, module in self.model.named_modules(): ++ quant_method = getattr(module, "quant_method", None) ++ ++ # Check if the module supports quantization post-processing ++ if quant_method is not None and hasattr( ++ quant_method, "process_weights_after_loading" ++ ): ++ ++ # Apply the post-processing (e.g., repacking weights for Marlin kernel) ++ with device_loading_context(module, target_device): ++ quant_method.process_weights_after_loading(module) ++ ++ return True, "Success" ++ + + def _model_load_weights_direct(model, named_tensors: List[Tuple[str, torch.Tensor]]): + params_dict = dict(model.named_parameters()) +diff --git a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py +index cc673a9ca..06c430d2c 100644 +--- a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py ++++ b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py +@@ -1,4 +1,5 @@ + from sglang.srt.compilation.piecewise_context_manager import is_in_piecewise_cuda_graph ++from sglang.srt.layers.attention.hybrid_attn_backend import HybridAttnBackend + from sglang.srt.layers.attention.tbo_backend import TboAttnBackend + from sglang.srt.models.deepseek_common.attention_forward_methods.forward_methods import ( + AttnForwardMethod, +@@ -150,6 +151,8 @@ def handle_attention_nsa(attn, forward_batch): + backend = forward_batch.attn_backend + if isinstance(backend, TboAttnBackend): # if enable tbo, get primary backend + backend = backend.primary ++ if isinstance(backend, HybridAttnBackend): ++ backend = backend._select_backend(forward_batch.forward_mode) + if hasattr(backend, "use_mha") and backend.use_mha: + return AttnForwardMethod.MHA_ONE_SHOT + return AttnForwardMethod.MLA +diff --git a/python/sglang/srt/models/deepseek_nextn.py b/python/sglang/srt/models/deepseek_nextn.py +index cb13a7c67..d62111471 100644 +--- a/python/sglang/srt/models/deepseek_nextn.py ++++ b/python/sglang/srt/models/deepseek_nextn.py +@@ -29,6 +29,7 @@ from sglang.srt.layers.attention.nsa.utils import ( + can_cp_split, + cp_all_gather_rerange_output, + cp_split_and_rebuild_data, ++ cp_split_and_rebuild_position, + is_nsa_enable_prefill_cp, + nsa_use_prefill_cp, + prepare_input_dp_with_cp_dsa, +@@ -160,6 +161,7 @@ class DeepseekModelNextN(nn.Module): + + if nsa_use_prefill_cp(forward_batch, self.nsa_enable_prefill_cp): + hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states) ++ positions = cp_split_and_rebuild_position(forward_batch, positions) + residual = None + with get_global_expert_distribution_recorder().disable_this_region(): + hidden_states, residual = self.decoder( +diff --git a/python/sglang/srt/models/glm4v_moe.py b/python/sglang/srt/models/glm4v_moe.py +index 324de18b4..c99723f49 100644 +--- a/python/sglang/srt/models/glm4v_moe.py ++++ b/python/sglang/srt/models/glm4v_moe.py +@@ -52,11 +52,31 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration): + self.num_fused_shared_experts = 0 + self.determine_num_fused_shared_experts() + +- self.model = Glm4MoeModel( +- config, +- quant_config, +- prefix=add_prefix("language_model", prefix), +- ) ++ if not self.config.encoder_only: ++ self.model = Glm4MoeModel( ++ config, ++ quant_config, ++ prefix=add_prefix("language_model", prefix), ++ ) ++ ++ if self.pp_group.is_last_rank: ++ if self.pp_group.world_size == 1 and self.config.tie_word_embeddings: ++ self.lm_head = self.model.embed_tokens ++ else: ++ self.lm_head = ParallelLMHead( ++ config.vocab_size, ++ config.hidden_size, ++ quant_config=quant_config, ++ prefix=add_prefix("lm_head", prefix), ++ use_attn_tp_group=get_global_server_args().enable_dp_lm_head, ++ ) ++ else: ++ # ranks other than the last rank will have a placeholder layer ++ self.lm_head = PPMissingLayer() ++ else: ++ # encoder_only mode: no language model, so no lm_head needed ++ self.lm_head = None ++ + self.visual = Glm4vVisionModel( + config.vision_config, + quant_config=quant_config, +@@ -64,21 +84,6 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration): + use_data_parallel=self.use_data_parallel, + ) + +- if self.pp_group.is_last_rank: +- if self.pp_group.world_size == 1 and self.config.tie_word_embeddings: +- self.lm_head = self.model.embed_tokens +- else: +- self.lm_head = ParallelLMHead( +- config.vocab_size, +- config.hidden_size, +- quant_config=quant_config, +- prefix=add_prefix("lm_head", prefix), +- use_attn_tp_group=get_global_server_args().enable_dp_lm_head, +- ) +- else: +- # ranks other than the last rank will have a placeholder layer +- self.lm_head = PPMissingLayer() +- + self.logits_processor = LogitsProcessor(config) + self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True) + self.is_mrope_enabled = "mrope_section" in self.config.rope_scaling +@@ -219,6 +224,11 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration): + # Skip loading extra bias for GPTQ models. + if name.endswith(".bias") and name not in params_dict: + continue ++ # Skip loading visual/language model weights ++ if ( ++ self.config.encoder_only or self.config.language_only ++ ) and name not in params_dict: ++ continue + if name not in params_dict: + continue + +@@ -234,6 +244,8 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration): + param_name, weight_name, expert_id, shard_id = mapping + if weight_name not in name: + continue ++ if "visual" in name or self.config.encoder_only: ++ continue + + # Mark as expert weight regardless of whether we can process it + is_expert_weight = True +@@ -265,6 +277,11 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration): + # Skip loading extra bias for GPTQ models. + if name.endswith(".bias") and name not in params_dict: + continue ++ # Skip loading mm/language parameters ++ if ( ++ self.config.encoder_only or self.config.language_only ++ ) and name not in params_dict: ++ continue + if name not in params_dict: + continue + +diff --git a/python/sglang/srt/models/gpt_oss.py b/python/sglang/srt/models/gpt_oss.py +index 2cf813bce..1250c49e4 100644 +--- a/python/sglang/srt/models/gpt_oss.py ++++ b/python/sglang/srt/models/gpt_oss.py +@@ -17,6 +17,7 @@ + + import logging + import math ++import re + from collections.abc import Iterable + from functools import partial + from typing import Any, Dict, List, Optional, Tuple, Union +@@ -1065,6 +1066,12 @@ class GptOssForCausalLM(nn.Module): + weight_loader(param, loaded_weight, shard_id) + break + else: ++ # Try per-expert format: experts.{id}.{gate_proj|up_proj|down_proj}.{weight|bias} ++ per_expert_match = _PER_EXPERT_RE.match(name) ++ if per_expert_match: ++ _load_per_expert_param(per_expert_match, loaded_weight, params_dict) ++ continue ++ + for mapping in expert_params_mapping: + param_name, weight_name, shard_id = mapping + if weight_name not in name: +@@ -1143,6 +1150,88 @@ class GptOssForCausalLM(nn.Module): + return get_attention_sliding_window_size(self.config) + + ++# Regex for per-expert weight names: model.layers.X.mlp.experts.E.{proj}.{weight|bias} ++_PER_EXPERT_RE = re.compile( ++ r"(.+\.mlp\.experts\.)(\d+)\.(gate_proj|up_proj|down_proj)\.(weight|bias)" ++) ++ ++ ++def _load_per_expert_param(match, loaded_weight, params_dict): ++ """Load a per-expert weight/bias tensor into the fused FusedMoE parameter. ++ ++ Handles the mapping from per-expert names (e.g., experts.0.gate_proj.weight) ++ to fused parameters (e.g., experts.w13_weight). ++ """ ++ prefix, eid_str, proj, ptype = match.groups() ++ eid = int(eid_str) ++ ++ # Determine target fused parameter name ++ if proj in ("gate_proj", "up_proj"): ++ key = prefix + ("w13_weight" if ptype == "weight" else "w13_weight_bias") ++ else: # down_proj ++ key = prefix + ("w2_weight" if ptype == "weight" else "w2_weight_bias") ++ ++ if key not in params_dict: ++ return ++ ++ param = params_dict[key] ++ expert_slice = param.data[eid] # slice for this expert ++ ++ # Detect triton transposed layout from shape: ++ # w13: transposed=(E, hidden, 2*inter), non-transposed=(E, 2*inter, hidden) ++ # For w13, the larger dim is 2*intermediate; if it's dim -1, layout is transposed. ++ is_transposed = getattr(param, "is_transposed", False) ++ if not is_transposed and ptype == "weight" and "w13" in key: ++ # Infer from shape: transposed has shape[-1] > shape[-2] ++ is_transposed = param.data.shape[-1] > param.data.shape[-2] ++ ++ if ptype == "weight": ++ if proj in ("gate_proj", "up_proj"): ++ # w13_weight: gate in first half, up in second half ++ if is_transposed: ++ # Triton layout: (E, hidden, 2*intermediate) ++ half = expert_slice.shape[1] // 2 ++ dst = ( ++ expert_slice[:, :half] ++ if proj == "gate_proj" ++ else expert_slice[:, half:] ++ ) ++ else: ++ # Standard layout: (E, 2*intermediate, hidden) ++ half = expert_slice.shape[0] // 2 ++ dst = ( ++ expert_slice[:half] if proj == "gate_proj" else expert_slice[half:] ++ ) ++ # loaded_weight shape: (intermediate, hidden) ++ if is_transposed: ++ dst.copy_(loaded_weight.t()) ++ else: ++ dst.copy_(loaded_weight) ++ else: ++ # w2_weight: loaded_weight shape (hidden, intermediate) ++ # Detect transposition for w2 as well ++ w2_transposed = is_transposed ++ if not w2_transposed: ++ w2_transposed = param.data.shape[-1] > param.data.shape[-2] ++ if w2_transposed: ++ expert_slice.copy_(loaded_weight.t()) ++ else: ++ expert_slice.copy_(loaded_weight) ++ else: ++ # Bias handling ++ if proj in ("gate_proj", "up_proj"): ++ # w13_weight_bias: (E, 2*intermediate) ++ half = expert_slice.shape[0] // 2 ++ dst = expert_slice[:half] if proj == "gate_proj" else expert_slice[half:] ++ dst.copy_(loaded_weight) ++ else: ++ # w2_weight_bias: (E, hidden) - only rank 0 loads, others zero ++ if get_moe_tensor_parallel_rank() == 0: ++ expert_slice.copy_(loaded_weight) ++ else: ++ expert_slice.zero_() ++ ++ + def _canonicalize_weights(config, weights_in: Iterable[Tuple[str, torch.Tensor]]): + weights_out_dict = dict(weights_in) + +diff --git a/python/sglang/srt/models/qwen3_5.py b/python/sglang/srt/models/qwen3_5.py +index f01225487..1dad8bb8e 100644 +--- a/python/sglang/srt/models/qwen3_5.py ++++ b/python/sglang/srt/models/qwen3_5.py +@@ -372,6 +372,7 @@ class Qwen3_5LinearDecoderLayer(nn.Module): + input_layernorm=self.input_layernorm, + post_attention_layernorm=self.post_attention_layernorm, + allow_reduce_scatter=True, ++ is_last_layer=(layer_id == config.num_hidden_layers - 1), + ) + + def forward( +@@ -400,11 +401,24 @@ class Qwen3_5LinearDecoderLayer(nn.Module): + use_reduce_scatter = self.layer_communicator.should_use_reduce_scatter( + forward_batch + ) +- hidden_states = self.mlp(hidden_states, forward_batch, use_reduce_scatter) + +- hidden_states, residual = self.layer_communicator.postprocess_layer( +- hidden_states, residual, forward_batch ++ should_allreduce_fusion = ( ++ self.layer_communicator.should_fuse_mlp_allreduce_with_next_layer( ++ forward_batch ++ ) + ) ++ if isinstance(self.mlp, Qwen2MoeSparseMoeBlock): ++ hidden_states = self.mlp(hidden_states, forward_batch, use_reduce_scatter) ++ else: ++ hidden_states = self.mlp( ++ hidden_states, should_allreduce_fusion, use_reduce_scatter ++ ) ++ if should_allreduce_fusion: ++ hidden_states._sglang_needs_allreduce_fusion = True ++ else: ++ hidden_states, residual = self.layer_communicator.postprocess_layer( ++ hidden_states, residual, forward_batch ++ ) + + return hidden_states, residual + +@@ -549,6 +563,7 @@ class Qwen3_5AttentionDecoderLayer(nn.Module): + input_layernorm=self.input_layernorm, + post_attention_layernorm=self.post_attention_layernorm, + allow_reduce_scatter=True, ++ is_last_layer=(layer_id == config.num_hidden_layers - 1), + ) + + self.alt_stream = alt_stream +@@ -633,11 +648,24 @@ class Qwen3_5AttentionDecoderLayer(nn.Module): + use_reduce_scatter = self.layer_communicator.should_use_reduce_scatter( + forward_batch + ) +- hidden_states = self.mlp(hidden_states, forward_batch, use_reduce_scatter) + +- hidden_states, residual = self.layer_communicator.postprocess_layer( +- hidden_states, residual, forward_batch ++ should_allreduce_fusion = ( ++ self.layer_communicator.should_fuse_mlp_allreduce_with_next_layer( ++ forward_batch ++ ) + ) ++ if isinstance(self.mlp, Qwen2MoeSparseMoeBlock): ++ hidden_states = self.mlp(hidden_states, forward_batch, use_reduce_scatter) ++ else: ++ hidden_states = self.mlp( ++ hidden_states, should_allreduce_fusion, use_reduce_scatter ++ ) ++ if should_allreduce_fusion: ++ hidden_states._sglang_needs_allreduce_fusion = True ++ else: ++ hidden_states, residual = self.layer_communicator.postprocess_layer( ++ hidden_states, residual, forward_batch ++ ) + + return hidden_states, residual + +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 ++++ b/python/sglang/srt/models/qwen3_vl.py +@@ -711,14 +711,19 @@ class Qwen3LLMModel(Qwen3Model): + hidden_states + residual if residual is not None else hidden_states + ) + ++ deepstack_embeds = None ++ if input_deepstack_embeds is not None: ++ prev_layer_idx = layer_idx - 1 ++ if prev_layer_idx in self.deepstack_embed_to_decoder_layer: ++ sep = self.hidden_size * prev_layer_idx ++ deepstack_embeds = input_deepstack_embeds[ ++ :, sep : sep + self.hidden_size ++ ] ++ + # 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, +diff --git a/python/sglang/srt/multimodal/processors/glm4v.py b/python/sglang/srt/multimodal/processors/glm4v.py +index 33cce6fe2..0970c4550 100644 +--- a/python/sglang/srt/multimodal/processors/glm4v.py ++++ b/python/sglang/srt/multimodal/processors/glm4v.py +@@ -1,6 +1,9 @@ + from typing import List, Union + ++import torch ++ + from sglang.srt.layers.rotary_embedding import MRotaryEmbedding ++from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem + from sglang.srt.models.glm4v import Glm4vForConditionalGeneration + from sglang.srt.models.glm4v_moe import Glm4vMoeForConditionalGeneration + from sglang.srt.multimodal.processors.base_processor import ( +@@ -45,6 +48,8 @@ class Glm4vImageProcessor(SGLangBaseProcessor): + self.IMAGE_END_TOKEN_ID = hf_config.image_end_token_id + self.VIDEO_START_TOKEN_ID = hf_config.video_start_token_id + self.VIDEO_END_TOKEN_ID = hf_config.video_end_token_id ++ self.IM_START_TOKEN_ID = self.IMAGE_START_TOKEN_ID ++ self.IM_END_TOKEN_ID = self.IMAGE_END_TOKEN_ID + + # Vision config + self.IMAGE_FACTOR = 28 +@@ -59,6 +64,36 @@ class Glm4vImageProcessor(SGLangBaseProcessor): + video_token_id=self.IM_TOKEN_ID, + ).build(_processor) + ++ def get_mm_data(self, prompt, embeddings, img_grid_thw): ++ input_ids, offsets = self.build_input_ids(prompt, img_grid_thw) ++ mm_items = [ ++ MultimodalDataItem( ++ modality=Modality.IMAGE, ++ offsets=offsets, ++ precomputed_embeddings=embeddings, ++ ) ++ ] ++ ++ input_ids_tensor = torch.tensor(input_ids) ++ mrope_positions, mrope_position_delta = MRotaryEmbedding.get_rope_index_glm4v( ++ input_ids=input_ids_tensor.unsqueeze(0), ++ hf_config=self.hf_config, ++ image_grid_thw=img_grid_thw, ++ video_grid_thw=None, ++ attention_mask=None, ++ ) ++ mrope_positions = mrope_positions.squeeze(1) ++ ++ return { ++ "input_ids": input_ids, ++ "mm_items": mm_items, ++ "im_start_id": self.IM_START_TOKEN_ID, ++ "im_end_id": self.IM_END_TOKEN_ID, ++ "im_token_id": self.IM_TOKEN_ID, ++ "mrope_positions": mrope_positions, ++ "mrope_position_delta": mrope_position_delta, ++ } ++ + async def process_mm_data_async( + self, + image_data: List[Union[str, bytes]], +diff --git a/python/sglang/srt/multimodal/processors/qwen_vl.py b/python/sglang/srt/multimodal/processors/qwen_vl.py +index 4395654e4..f9b5ea4ab 100644 +--- a/python/sglang/srt/multimodal/processors/qwen_vl.py ++++ b/python/sglang/srt/multimodal/processors/qwen_vl.py +@@ -317,7 +317,7 @@ class QwenVLImageProcessor(SGLangBaseProcessor): + **kwargs, + ): + entry_time = time.perf_counter() +- base_output = self.load_mm_data( ++ base_output = self.legacy_load_mm_data( + prompt=input_text, + image_data=image_data, + video_data=request_obj.video_data, +diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py +index b080aeb16..957a613fa 100644 +--- a/python/sglang/srt/server_args.py ++++ b/python/sglang/srt/server_args.py +@@ -580,6 +580,7 @@ class ServerArgs: + cuda_graph_max_bs: Optional[int] = None + cuda_graph_bs: Optional[List[int]] = None + disable_cuda_graph: bool = False ++ disable_draft_cuda_graph: bool = False + disable_cuda_graph_padding: bool = False + enable_profile_cuda_graph: bool = False + enable_cudagraph_gc: bool = False +@@ -2089,7 +2090,16 @@ class ServerArgs: + assert ( + self.tp_size % (self.dp_size * self.attn_cp_size) == 0 + ), "tp_size must be divisible by dp_size * attn_cp_size" +- assert self.pp_size == 1, "PP is not supported with context parallelism" ++ if self.pp_size > 1: ++ assert ( ++ self.disaggregation_mode == "prefill" ++ and self.enable_nsa_prefill_context_parallel ++ and self.nsa_prefill_cp_mode == "round-robin-split" ++ ), ( ++ "PP with context parallelism is only supported for PD prefill " ++ "with --enable-nsa-prefill-context-parallel and " ++ "--nsa-prefill-cp-mode round-robin-split." ++ ) + + if self.moe_dp_size > 1: + # The tp_size is the world size, not the real tensor parallel size +@@ -4491,6 +4501,11 @@ class ServerArgs: + action="store_true", + help="Disable cuda graph.", + ) ++ parser.add_argument( ++ "--disable-draft-cuda-graph", ++ action="store_true", ++ help="Disable cuda graph for draft model in speculative decoding.", ++ ) + parser.add_argument( + "--disable-cuda-graph-padding", + action="store_true", +@@ -5636,6 +5651,54 @@ class PortArgs: + ) + + if not server_args.enable_dp_attention: ++ # In multi-node prefill PD with PP/CP, use TCP transport for tokenizer<->scheduler ++ # IPC occasionally stalls in this topology. ++ if ( ++ server_args.nnodes > 1 ++ and server_args.disaggregation_mode == "prefill" ++ and server_args.dist_init_addr is not None ++ ): ++ if server_args.dist_init_addr.startswith("["): # ipv6 address ++ port_num, host = configure_ipv6(server_args.dist_init_addr) ++ dist_init_addr = (host, str(port_num)) ++ else: ++ dist_init_addr = server_args.dist_init_addr.split(":") ++ ++ assert ( ++ len(dist_init_addr) == 2 ++ ), "please provide --dist-init-addr as host:port of head node" ++ ++ dist_init_host, dist_init_port = dist_init_addr ++ dist_init_port = int(dist_init_port) ++ port_base = dist_init_port + ZMQ_TCP_PORT_DELTA ++ detokenizer_port = port_base + 1 ++ rpc_port = port_base + 2 ++ metrics_port = port_base + 3 ++ scheduler_input_port = port_base + 4 ++ ++ try: ++ wait_port_available(dist_init_port, "dist_init_port") ++ wait_port_available(port_base, "port_base") ++ wait_port_available(detokenizer_port, "detokenizer_port") ++ wait_port_available(nccl_port, "nccl_port") ++ wait_port_available(rpc_port, "rpc_port") ++ wait_port_available(metrics_port, "metrics_port") ++ wait_port_available(scheduler_input_port, "scheduler_input_port") ++ except ValueError: ++ logger.exception( ++ f"Port is already in use. {dist_init_port=} {port_base=} {detokenizer_port=} {nccl_port=} {scheduler_input_port=}" ++ ) ++ raise ++ ++ return PortArgs( ++ tokenizer_ipc_name=f"tcp://{dist_init_host}:{port_base}", ++ scheduler_input_ipc_name=f"tcp://{dist_init_host}:{scheduler_input_port}", ++ detokenizer_ipc_name=f"tcp://{dist_init_host}:{detokenizer_port}", ++ nccl_port=nccl_port, ++ rpc_ipc_name=f"tcp://{dist_init_host}:{rpc_port}", ++ metrics_ipc_name=f"tcp://{dist_init_host}:{metrics_port}", ++ tokenizer_worker_ipc_name=tokenizer_worker_ipc_name, ++ ) + # Normal case, use IPC within a single node + return PortArgs( + tokenizer_ipc_name=f"ipc://{tempfile.NamedTemporaryFile(delete=False).name}", +diff --git a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py +index 5fe45086c..b283d2e9b 100644 +--- a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py ++++ b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py +@@ -341,7 +341,10 @@ class EAGLEDraftCudaGraphRunner: + self.seq_lens.fill_(self.seq_len_fill_value) + self.out_cache_loc.zero_() + self.positions.zero_() +- ++ self.topk_p.zero_() ++ self.topk_index.zero_() ++ self.hidden_states.zero_() ++ self.req_pool_indices.zero_() + num_tokens = bs * self.num_tokens_per_bs + + # Common inputs +@@ -350,8 +353,12 @@ class EAGLEDraftCudaGraphRunner: + forward_batch.out_cache_loc + ) + self.positions[:raw_num_token].copy_(forward_batch.positions) +- self.topk_p[:raw_bs].copy_(forward_batch.spec_info.topk_p) +- self.topk_index[:raw_bs].copy_(forward_batch.spec_info.topk_index) ++ self.topk_p[:raw_bs].copy_(forward_batch.spec_info.topk_p.clamp(0, 1)) ++ self.topk_index[:raw_bs].copy_( ++ forward_batch.spec_info.topk_index.clamp( ++ 0, self.model_runner.model_config.vocab_size - 1 ++ ) ++ ) + self.hidden_states[:raw_bs].copy_(forward_batch.spec_info.hidden_states) + 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..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)] + self.verified_id = self.verified_id[: len(new_indices)] ++ if self.accept_length is not None: ++ self.accept_length = self.accept_length[: len(new_indices)] ++ if self.accept_length_cpu is not None: ++ self.accept_length_cpu = self.accept_length_cpu[: len(new_indices)] + else: + # in some cases(e.g draft_extend), we have not filtered the batch by `unfinished_index` + self.topk_p = self.topk_p[new_indices] +@@ -805,6 +809,27 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): + self.verified_id = torch.cat([self.verified_id, spec_info.verified_id], axis=0) + self.topk_p = torch.cat([self.topk_p, spec_info.topk_p]) + self.topk_index = torch.cat([self.topk_index, spec_info.topk_index]) ++ if self.accept_length is not None and spec_info.accept_length is not None: ++ self.accept_length = torch.cat( ++ [self.accept_length, spec_info.accept_length] ++ ) ++ self.accept_length_cpu = self.accept_length.tolist() ++ elif self.accept_length is not None: ++ zeros = torch.zeros( ++ [spec_info.verified_id.shape[0]], ++ dtype=self.accept_length.dtype, ++ device=self.accept_length.device, ++ ) ++ self.accept_length = torch.cat([self.accept_length, zeros]) ++ self.accept_length_cpu = self.accept_length.tolist() ++ elif spec_info.accept_length is not None: ++ zeros = torch.zeros( ++ [self.verified_id.shape[0]], ++ dtype=self.accept_length.dtype, ++ device=self.accept_length.device, ++ ) ++ self.accept_length = torch.cat([zeros, spec_info.accept_length]) ++ self.accept_length_cpu = self.accept_length.tolist() + + + @dataclass +diff --git a/python/sglang/srt/speculative/eagle_worker.py b/python/sglang/srt/speculative/eagle_worker.py +index 32b3a520a..d7f940147 100644 +--- a/python/sglang/srt/speculative/eagle_worker.py ++++ b/python/sglang/srt/speculative/eagle_worker.py +@@ -234,7 +234,10 @@ class EAGLEWorker(TpModelWorker): + self.cuda_graph_runner = None + self.cuda_graph_runner_for_draft_extend = None + +- if self.server_args.disable_cuda_graph: ++ if ( ++ self.server_args.disable_cuda_graph ++ or self.server_args.disable_draft_cuda_graph ++ ): + return + + Device2DraftCudaGraphRunner = { +diff --git a/python/sglang/srt/utils/common.py b/python/sglang/srt/utils/common.py +index 4636128fa..a9b61df39 100644 +--- a/python/sglang/srt/utils/common.py ++++ b/python/sglang/srt/utils/common.py +@@ -2359,6 +2359,8 @@ class SafeUnpickler(pickle.Unpickler): + "sglang.srt.model_executor.model_runner.", + "sglang.srt.layers.", + "sglang.srt.utils.", ++ # --- slime --- ++ "slime.", + } + + DENY_CLASSES = { +diff --git a/python/sglang/srt/utils/weight_checker.py b/python/sglang/srt/utils/weight_checker.py +index 3be16446e..1b2371c83 100644 +--- a/python/sglang/srt/utils/weight_checker.py ++++ b/python/sglang/srt/utils/weight_checker.py +@@ -69,6 +69,9 @@ def _check_tensors( + actual_should_compare, + actual, + ) in zip(expect_tensors, actual_tensors, strict=True): ++ if ".cos_sin_cache" in expect_name: ++ # skip cos/sin cache which is deterministic from shape and dtype and may have different shapes due to different implementations. ++ continue + assert expect_name == actual_name, f"{expect_name=} {actual_name=}" + assert ( + expect_should_compare == actual_should_compare diff --git a/tests/test_qwen2.5_0.5B_ppo_critic_only_short.py b/tests/test_qwen2.5_0.5B_ppo_critic_only_short.py index a3c0cd15ce..533b2010bf 100644 --- a/tests/test_qwen2.5_0.5B_ppo_critic_only_short.py +++ b/tests/test_qwen2.5_0.5B_ppo_critic_only_short.py @@ -68,6 +68,7 @@ def execute(): sglang_args = ( "--rollout-num-gpus-per-engine 1 " + "--rollout-num-gpus 2 " f"--sglang-mem-fraction-static {0.6 if TIGHT_DEVICE_MEMORY else 0.7} " "--sglang-cuda-graph-max-bs 32 " "--sglang-enable-metrics "