From e9e7f15eb50efd683febf8b3139de71cf38a54ef Mon Sep 17 00:00:00 2001 From: Yuhao Yao <37280700+yuhyao@users.noreply.github.com> Date: Fri, 12 Dec 2025 10:18:40 +0800 Subject: [PATCH] [bugfix] fix TBO crashes when attn_tp_size > 1 (#13730) Co-authored-by: Cheng Wan <54331508+ch-wan@users.noreply.github.com> --- python/sglang/srt/batch_overlap/operations.py | 18 +- .../srt/batch_overlap/two_batch_overlap.py | 45 ++++- python/sglang/srt/layers/communicator.py | 15 +- .../srt/model_executor/forward_batch_info.py | 9 + python/sglang/srt/models/bailing_moe.py | 4 + python/sglang/srt/models/deepseek_v2.py | 2 + python/sglang/srt/models/falcon_h1.py | 4 +- python/sglang/srt/models/glm4_moe.py | 2 + python/sglang/srt/models/gpt_oss.py | 2 + python/sglang/srt/models/llada2.py | 2 + python/sglang/srt/models/llama4.py | 2 + python/sglang/srt/models/longcat_flash.py | 4 + .../sglang/srt/models/longcat_flash_nextn.py | 1 + python/sglang/srt/models/minimax_m2.py | 2 + python/sglang/srt/models/qwen2_moe.py | 2 + python/sglang/srt/models/qwen3.py | 1 + python/sglang/srt/models/qwen3_moe.py | 2 + python/sglang/srt/models/qwen3_next.py | 4 + python/sglang/srt/models/step3_vl.py | 2 + test/srt/ep/test_deepep_small.py | 178 ++++++++++++++++++ 20 files changed, 285 insertions(+), 16 deletions(-) diff --git a/python/sglang/srt/batch_overlap/operations.py b/python/sglang/srt/batch_overlap/operations.py index 9d824587c..7701c347f 100644 --- a/python/sglang/srt/batch_overlap/operations.py +++ b/python/sglang/srt/batch_overlap/operations.py @@ -83,7 +83,7 @@ class _StageExecutor: # handling DP attention forward_batch: ForwardBatch = inputs["forward_batch"] self._global_dp_buffer_len = forward_batch.global_dp_buffer_len - self._local_dp_buffer_len = forward_batch.input_ids.shape[0] + self._local_dp_buffer_len = forward_batch.tbo_padded_len self._global_num_tokens = forward_batch.global_num_tokens_cpu self._is_dp_max_padding = forward_batch.dp_padding_mode.is_max_len() @@ -92,13 +92,15 @@ class _StageExecutor: stage = self._stages[self._index] - if self._global_dp_buffer_len is not None: - set_dp_buffer_len( - self._global_dp_buffer_len, - self._local_dp_buffer_len, - self._is_dp_max_padding, - self._global_num_tokens, - ) + # TODO: We currently always call set_dp_buffer_len here because sub-batches + # may have different padded lengths. It can likely be removed after TBO slice & + # pad logic is refactored. + set_dp_buffer_len( + self._global_dp_buffer_len, + self._local_dp_buffer_len, + self._is_dp_max_padding, + self._global_num_tokens, + ) with _annotate_region(debug_name=f"{self._debug_name}{self._index}"): for op in stage: diff --git a/python/sglang/srt/batch_overlap/two_batch_overlap.py b/python/sglang/srt/batch_overlap/two_batch_overlap.py index 4d5746d26..da266528e 100644 --- a/python/sglang/srt/batch_overlap/two_batch_overlap.py +++ b/python/sglang/srt/batch_overlap/two_batch_overlap.py @@ -20,6 +20,7 @@ from sglang.srt.layers.communicator import ( CommunicateSummableTensorPairFn, ScatterMode, ) +from sglang.srt.layers.dp_attention import get_attention_tp_size from sglang.srt.layers.moe import ( get_deepep_mode, get_moe_a2a_backend, @@ -630,6 +631,11 @@ class TboForwardBatchPreparer: ), f"{key=} {old_value=} {num_tokens=} {batch=}" output_dict[key] = old_value[start_token_index:end_token_index] + attention_tp_size = get_attention_tp_size() + output_dict["tbo_padded_len"] = ( + (end_token_index - start_token_index - 1) // attention_tp_size + 1 + ) * attention_tp_size + for key in [ "req_pool_indices", "seq_lens", @@ -840,6 +846,7 @@ def _model_forward_tbo( input_data_scatter_mode=input_data_scatter_mode, layer_input_scatter_mode=layer_input_scatter_mode, ) + original_hidden_states_len = inputs["hidden_states"].shape[0] del inputs context = ( @@ -857,7 +864,7 @@ def _model_forward_tbo( delta_stages=[0, operations_strategy.tbo_delta_stages], ) - return _model_forward_tbo_merge_outputs(*outputs_arr) + return _model_forward_tbo_merge_outputs(*outputs_arr, original_hidden_states_len) def _model_forward_non_tbo(inputs, operations_strategy: OperationsStrategy): @@ -951,23 +958,49 @@ def _model_forward_filter_inputs( tbo_subbatch_index: int, ) -> Dict: token_slice = slice(*output_forward_batch.tbo_parent_token_range) + hidden_states = hidden_states[token_slice] + residual = None if residual is None else residual[token_slice] + positions = positions[token_slice] + + assert output_forward_batch.tbo_padded_len is not None + padded_len = output_forward_batch.tbo_padded_len + + def _pad(x): + nonlocal padded_len + if x is None: + return None + if x.shape[0] == padded_len: + return x + res = torch.zeros((padded_len, *x.shape[1:]), dtype=x.dtype, device=x.device) + res[: x.shape[0]] = x + return res + return dict( - hidden_states=hidden_states[token_slice], - residual=None if residual is None else residual[token_slice], - positions=positions[token_slice], + hidden_states=_pad(hidden_states), + residual=_pad(residual), + positions=_pad(positions), forward_batch=output_forward_batch, tbo_subbatch_index=tbo_subbatch_index, ) -def _model_forward_tbo_merge_outputs(output_a, output_b): +def _model_forward_tbo_merge_outputs(output_a, output_b, original_len): def _handle_key(name): value_a = output_a[name] value_b = output_b[name] assert (value_a is None) == (value_b is None) if value_a is None: return None - return torch.concat([value_a, value_b], dim=0) + s0, t0 = output_a["forward_batch"].tbo_parent_token_range + s1, t1 = output_b["forward_batch"].tbo_parent_token_range + res = torch.zeros( + (original_len, *value_a.shape[1:]), + dtype=value_a.dtype, + device=value_a.device, + ) + res[slice(s0, t0)] = value_a[: t0 - s0] + res[slice(s1, t1)] = value_b[: t1 - s1] + return res return _handle_key("hidden_states"), _handle_key("residual") diff --git a/python/sglang/srt/layers/communicator.py b/python/sglang/srt/layers/communicator.py index 932f52aeb..33ce75364 100644 --- a/python/sglang/srt/layers/communicator.py +++ b/python/sglang/srt/layers/communicator.py @@ -217,14 +217,16 @@ class _LayerModeComputationContext: layer_id: int is_layer_sparse: bool is_previous_layer_sparse: Optional[bool] + is_next_layer_sparse: Optional[bool] def previous_layer(self): assert self.is_previous_layer_sparse is not None return _LayerModeComputationContext( + num_layers=self.num_layers, layer_id=self.layer_id - 1, is_layer_sparse=self.is_previous_layer_sparse, is_previous_layer_sparse=None, - num_layers=self.num_layers, + is_next_layer_sparse=self.is_layer_sparse, ) @@ -273,6 +275,15 @@ class LayerScatterModes: else ScatterMode.FULL ) + @classmethod + def _should_gather_for_tbo(cls, context: _LayerModeComputationContext): + return ( + not context.is_layer_sparse + and context.is_next_layer_sparse + and enable_moe_dense_fully_dp() + and get_global_server_args().enable_two_batch_overlap + ) + @classmethod def _compute_middle_residual_mode(cls, context: _LayerModeComputationContext): mlp_mode = cls._compute_mlp_mode(context) @@ -288,6 +299,8 @@ class LayerScatterModes: if context.layer_id == context.num_layers - 1: return ScatterMode.model_input_output() if mlp_mode == ScatterMode.SCATTERED: + if cls._should_gather_for_tbo(context): + return ScatterMode.TP_ATTN_FULL return ScatterMode.SCATTERED if mlp_mode == ScatterMode.FULL: return ScatterMode.TP_ATTN_FULL diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 0b3d74b78..931445cc5 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -376,6 +376,7 @@ class ForwardBatch: # For two-batch overlap tbo_split_seq_index: Optional[int] = None tbo_parent_token_range: Optional[Tuple[int, int]] = None + tbo_padded_len: Optional[int] = None tbo_children: Optional[List[ForwardBatch]] = None # For matryoshka embeddings @@ -852,6 +853,14 @@ class ForwardBatch: TboForwardBatchPreparer.prepare( batch=self, is_draft_worker=model_runner.is_draft_worker ) + # TODO: The following is added to make sure sub-batch input_ids are padded + # to the multiple of attn_tp_size. It can likely be removed after this + # function is refactored and merged into the Scheduler. + if self.tbo_children: + for child in self.tbo_children: + child._pad_inputs_to_size( + model_runner, child.tbo_padded_len, child.batch_size + ) def _pad_inputs_to_size(self, model_runner: ModelRunner, num_tokens, bs): # padding diff --git a/python/sglang/srt/models/bailing_moe.py b/python/sglang/srt/models/bailing_moe.py index 248839be6..02048850d 100644 --- a/python/sglang/srt/models/bailing_moe.py +++ b/python/sglang/srt/models/bailing_moe.py @@ -582,12 +582,16 @@ class BailingMoEBlock(nn.Module): is_previous_layer_sparse = self._is_layer_sparse( config, layer_id=layer_id - 1, is_nextn=False ) + is_next_layer_sparse = self._is_layer_sparse( + config, layer_id=layer_id + 1, is_nextn=False + ) self.layer_scatter_modes = LayerScatterModes.init_new( layer_id=layer_id, num_layers=config.num_hidden_layers, is_layer_sparse=self.is_layer_sparse, is_previous_layer_sparse=is_previous_layer_sparse, + is_next_layer_sparse=is_next_layer_sparse, ) self.is_last_layer = self.layer_id == config.num_hidden_layers - 1 diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 2b1c3c04d..e210a83af 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -2734,12 +2734,14 @@ class DeepseekV2DecoderLayer(nn.Module): self.is_layer_sparse = self._is_layer_sparse(layer_id, is_nextn=is_nextn) is_previous_layer_sparse = self._is_layer_sparse(layer_id - 1, is_nextn=False) + is_next_layer_sparse = self._is_layer_sparse(layer_id + 1, is_nextn=False) self.layer_scatter_modes = LayerScatterModes.init_new( layer_id=layer_id, num_layers=1 if is_nextn else config.num_hidden_layers, is_layer_sparse=self.is_layer_sparse, is_previous_layer_sparse=is_previous_layer_sparse, + is_next_layer_sparse=is_next_layer_sparse, ) if self.is_layer_sparse: diff --git a/python/sglang/srt/models/falcon_h1.py b/python/sglang/srt/models/falcon_h1.py index 0fab9e410..551307cb8 100644 --- a/python/sglang/srt/models/falcon_h1.py +++ b/python/sglang/srt/models/falcon_h1.py @@ -198,15 +198,17 @@ class FalconH1HybridAttentionDecoderLayer(nn.Module): prefix=f"{prefix}.mixer", ) - # FalconH1 all layers are sparse and have no nextn now + # FalconH1 all layers are dense and have no nextn now self.is_layer_sparse = False is_previous_layer_sparse = False + is_next_layer_sparse = False self.layer_scatter_modes = LayerScatterModes.init_new( layer_id=layer_id, num_layers=config.num_hidden_layers, is_layer_sparse=self.is_layer_sparse, is_previous_layer_sparse=is_previous_layer_sparse, + is_next_layer_sparse=is_next_layer_sparse, ) self.feed_forward = FalconH1MLP( diff --git a/python/sglang/srt/models/glm4_moe.py b/python/sglang/srt/models/glm4_moe.py index 280f5602c..6483e41c1 100644 --- a/python/sglang/srt/models/glm4_moe.py +++ b/python/sglang/srt/models/glm4_moe.py @@ -714,12 +714,14 @@ class Glm4MoeDecoderLayer(nn.Module): self.is_layer_sparse = self._is_layer_sparse(layer_id, is_nextn=is_nextn) is_previous_layer_sparse = self._is_layer_sparse(layer_id - 1, is_nextn=False) + is_next_layer_sparse = self._is_layer_sparse(layer_id + 1, is_nextn=False) self.layer_scatter_modes = LayerScatterModes.init_new( layer_id=layer_id, num_layers=1 if is_nextn else config.num_hidden_layers, is_layer_sparse=self.is_layer_sparse, is_previous_layer_sparse=is_previous_layer_sparse, + is_next_layer_sparse=is_next_layer_sparse, ) if self.is_layer_sparse: diff --git a/python/sglang/srt/models/gpt_oss.py b/python/sglang/srt/models/gpt_oss.py index 9474700c4..3e7283dd4 100644 --- a/python/sglang/srt/models/gpt_oss.py +++ b/python/sglang/srt/models/gpt_oss.py @@ -395,12 +395,14 @@ class GptOssDecoderLayer(nn.Module): self.is_layer_sparse = True self.is_nextn = False is_previous_layer_sparse = True + is_next_layer_sparse = True self.layer_scatter_modes = LayerScatterModes.init_new( layer_id=layer_id, num_layers=config.num_hidden_layers, is_layer_sparse=self.is_layer_sparse, is_previous_layer_sparse=is_previous_layer_sparse, + is_next_layer_sparse=is_next_layer_sparse, ) if self.is_layer_sparse: diff --git a/python/sglang/srt/models/llada2.py b/python/sglang/srt/models/llada2.py index b8bfc81ef..94d771224 100644 --- a/python/sglang/srt/models/llada2.py +++ b/python/sglang/srt/models/llada2.py @@ -579,12 +579,14 @@ class LLaDA2MoeBlock(nn.Module): self.is_layer_sparse = self._is_layer_sparse(config, layer_id=layer_id) is_previous_layer_sparse = self._is_layer_sparse(config, layer_id=layer_id - 1) + is_next_layer_sparse = self._is_layer_sparse(config, layer_id=layer_id + 1) self.layer_scatter_modes = LayerScatterModes.init_new( layer_id=layer_id, num_layers=config.num_hidden_layers, is_layer_sparse=self.is_layer_sparse, is_previous_layer_sparse=is_previous_layer_sparse, + is_next_layer_sparse=is_next_layer_sparse, ) self.is_last_layer = self.layer_id == config.num_hidden_layers - 1 diff --git a/python/sglang/srt/models/llama4.py b/python/sglang/srt/models/llama4.py index ca46534a7..4a4e309bf 100644 --- a/python/sglang/srt/models/llama4.py +++ b/python/sglang/srt/models/llama4.py @@ -384,6 +384,7 @@ class Llama4DecoderLayer(nn.Module): self.config = config is_moe_layer = self._is_moe_layer(layer_id) is_previous_moe_layer = self._is_moe_layer(layer_id - 1) + is_next_moe_layer = self._is_moe_layer(layer_id + 1) if is_moe_layer: self.feed_forward = Llama4MoE( @@ -410,6 +411,7 @@ class Llama4DecoderLayer(nn.Module): num_layers=config.num_hidden_layers, is_layer_sparse=is_moe_layer, is_previous_layer_sparse=is_previous_moe_layer, + is_next_layer_sparse=is_next_moe_layer, ) self.layer_communicator = LayerCommunicator( diff --git a/python/sglang/srt/models/longcat_flash.py b/python/sglang/srt/models/longcat_flash.py index 3530609ba..4c2a625f5 100644 --- a/python/sglang/srt/models/longcat_flash.py +++ b/python/sglang/srt/models/longcat_flash.py @@ -380,6 +380,8 @@ class LongcatFlashDecoderLayer(nn.Module): num_layers=config.num_hidden_layers, is_layer_sparse=False, is_previous_layer_sparse=False, + # TODO: Check if the following is correct. + is_next_layer_sparse=False, ) for i in range(2) ] @@ -398,6 +400,8 @@ class LongcatFlashDecoderLayer(nn.Module): num_layers=config.num_hidden_layers, is_layer_sparse=True, is_previous_layer_sparse=True, + # TODO: Check if the following is correct. + is_next_layer_sparse=True, ) self.moe_layer_communicator = LayerCommunicator( layer_scatter_modes=self.moe_layer_scatter_modes, diff --git a/python/sglang/srt/models/longcat_flash_nextn.py b/python/sglang/srt/models/longcat_flash_nextn.py index b51417c65..f43946600 100644 --- a/python/sglang/srt/models/longcat_flash_nextn.py +++ b/python/sglang/srt/models/longcat_flash_nextn.py @@ -161,6 +161,7 @@ class LongcatFlashDenseDecoderLayer(nn.Module): num_layers=config.num_hidden_layers, is_layer_sparse=False, is_previous_layer_sparse=False, + is_next_layer_sparse=False, ) self.layer_communicator = LayerCommunicator( layer_scatter_modes=self.layer_scatter_modes, diff --git a/python/sglang/srt/models/minimax_m2.py b/python/sglang/srt/models/minimax_m2.py index 4ab5adbba..ca9059f89 100644 --- a/python/sglang/srt/models/minimax_m2.py +++ b/python/sglang/srt/models/minimax_m2.py @@ -516,11 +516,13 @@ class MiniMaxM2DecoderLayer(nn.Module): ) is_previous_layer_sparse = True + is_next_layer_sparse = True self.layer_scatter_modes = LayerScatterModes.init_new( layer_id=layer_id, num_layers=config.num_hidden_layers, is_layer_sparse=self.is_layer_sparse, is_previous_layer_sparse=is_previous_layer_sparse, + is_next_layer_sparse=is_next_layer_sparse, ) self.layer_communicator = LayerCommunicator( diff --git a/python/sglang/srt/models/qwen2_moe.py b/python/sglang/srt/models/qwen2_moe.py index ea33e81ef..03dad1845 100644 --- a/python/sglang/srt/models/qwen2_moe.py +++ b/python/sglang/srt/models/qwen2_moe.py @@ -461,12 +461,14 @@ class Qwen2MoeDecoderLayer(nn.Module): # Qwen2MoE all layers are sparse and have no nextn now self.is_layer_sparse = True is_previous_layer_sparse = True + is_next_layer_sparse = True self.layer_scatter_modes = LayerScatterModes.init_new( layer_id=layer_id, num_layers=config.num_hidden_layers, is_layer_sparse=self.is_layer_sparse, is_previous_layer_sparse=is_previous_layer_sparse, + is_next_layer_sparse=is_next_layer_sparse, ) if self.is_layer_sparse: diff --git a/python/sglang/srt/models/qwen3.py b/python/sglang/srt/models/qwen3.py index 7fcfa3fe4..c527201c0 100644 --- a/python/sglang/srt/models/qwen3.py +++ b/python/sglang/srt/models/qwen3.py @@ -276,6 +276,7 @@ class Qwen3DecoderLayer(nn.Module): num_layers=config.num_hidden_layers, is_layer_sparse=False, is_previous_layer_sparse=False, + is_next_layer_sparse=False, ) self.layer_communicator = LayerCommunicator( layer_scatter_modes=self.layer_scatter_modes, diff --git a/python/sglang/srt/models/qwen3_moe.py b/python/sglang/srt/models/qwen3_moe.py index 9737ac719..9388b974a 100644 --- a/python/sglang/srt/models/qwen3_moe.py +++ b/python/sglang/srt/models/qwen3_moe.py @@ -728,12 +728,14 @@ class Qwen3MoeDecoderLayer(nn.Module): # Qwen3MoE all layers are sparse and have no nextn now self.is_layer_sparse = True is_previous_layer_sparse = True + is_next_layer_sparse = True self.layer_scatter_modes = LayerScatterModes.init_new( layer_id=layer_id, num_layers=config.num_hidden_layers, is_layer_sparse=self.is_layer_sparse, is_previous_layer_sparse=is_previous_layer_sparse, + is_next_layer_sparse=is_next_layer_sparse, ) if self.is_layer_sparse: diff --git a/python/sglang/srt/models/qwen3_next.py b/python/sglang/srt/models/qwen3_next.py index f1f965ccd..9cde50468 100644 --- a/python/sglang/srt/models/qwen3_next.py +++ b/python/sglang/srt/models/qwen3_next.py @@ -492,6 +492,7 @@ class Qwen3HybridLinearDecoderLayer(nn.Module): # Qwen3Next all layers are sparse and have no nextn now self.is_layer_sparse = True is_previous_layer_sparse = True + is_next_layer_sparse = True self.layer_id = layer_id self.layer_scatter_modes = LayerScatterModes.init_new( @@ -499,6 +500,7 @@ class Qwen3HybridLinearDecoderLayer(nn.Module): num_layers=config.num_hidden_layers, is_layer_sparse=self.is_layer_sparse, is_previous_layer_sparse=is_previous_layer_sparse, + is_next_layer_sparse=is_next_layer_sparse, ) if self.is_layer_sparse: @@ -647,12 +649,14 @@ class Qwen3HybridAttentionDecoderLayer(nn.Module): # Qwen3Next all layers are sparse and have no nextn now self.is_layer_sparse = True is_previous_layer_sparse = True + is_next_layer_sparse = True self.layer_scatter_modes = LayerScatterModes.init_new( layer_id=layer_id, num_layers=config.num_hidden_layers, is_layer_sparse=self.is_layer_sparse, is_previous_layer_sparse=is_previous_layer_sparse, + is_next_layer_sparse=is_next_layer_sparse, ) if self.is_layer_sparse: diff --git a/python/sglang/srt/models/step3_vl.py b/python/sglang/srt/models/step3_vl.py index 4474f62d5..8d1673a05 100644 --- a/python/sglang/srt/models/step3_vl.py +++ b/python/sglang/srt/models/step3_vl.py @@ -337,12 +337,14 @@ class Step3TextDecoderLayer(nn.Module): self.is_previous_layer_sparse = ( True if layer_id - 1 in moe_layers_idx else False ) + self.is_next_layer_sparse = True if layer_id + 1 in moe_layers_idx else False self.layer_scatter_modes = LayerScatterModes.init_new( layer_id=layer_id, num_layers=config.num_hidden_layers, is_layer_sparse=self.is_layer_sparse, is_previous_layer_sparse=self.is_previous_layer_sparse, + is_next_layer_sparse=self.is_next_layer_sparse, ) if not self.is_layer_sparse: diff --git a/test/srt/ep/test_deepep_small.py b/test/srt/ep/test_deepep_small.py index 2dd42b189..126adab6b 100644 --- a/test/srt/ep/test_deepep_small.py +++ b/test/srt/ep/test_deepep_small.py @@ -251,6 +251,108 @@ class TestTBO(CustomTestCase): self.assertGreater(metrics["accuracy"], 0.60) +class TestTBOWithTPAttn(CustomTestCase): + @classmethod + def setUpClass(cls): + cls.model = DEFAULT_MODEL_NAME_FOR_TEST_MLA + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=[ + "--trust-remote-code", + "--tp", + "4", + "--moe-a2a-backend", + "deepep", + "--enable-two-batch-overlap", + "--cuda-graph-max-bs", + "128", + "--max-running-requests", + "512", + "--mem-fraction-static", # temp fix as DeepEP buffer is too large. + "0.7", + ], + env={ + **os.environ, + "SGLANG_TBO_DEBUG": "1", + }, + ) + + @classmethod + def tearDownClass(cls): + kill_process_tree(cls.process.pid) + + def test_gsm8k(self): + args = SimpleNamespace( + num_shots=5, + data_path=None, + num_questions=200, + max_new_tokens=512, + parallel=128, + host="http://127.0.0.1", + port=int(self.base_url.split(":")[-1]), + ) + metrics = run_eval_few_shot_gsm8k(args) + print(metrics) + + self.assertGreater(metrics["accuracy"], 0.60) + + +# There exists bug when using MTP + TBO + attn_tp_size > 1, currently skip that case. +# @unittest.skip("covered in TestMTPWithTPAttnAndTBO") +class TestTBOWithTPAttnAndDenseDP(CustomTestCase): + @classmethod + def setUpClass(cls): + cls.model = DEFAULT_MODEL_NAME_FOR_TEST_MLA + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=[ + "--trust-remote-code", + "--tp", + "4", + "--moe-dense-tp-size", + "1", + "--moe-a2a-backend", + "deepep", + "--enable-two-batch-overlap", + "--cuda-graph-max-bs", + "128", + "--max-running-requests", + "512", + "--mem-fraction-static", # temp fix as DeepEP buffer is too large. + "0.7", + ], + env={ + **os.environ, + "SGLANG_TBO_DEBUG": "1", + }, + ) + + @classmethod + def tearDownClass(cls): + kill_process_tree(cls.process.pid) + + def test_gsm8k(self): + args = SimpleNamespace( + num_shots=5, + data_path=None, + num_questions=200, + max_new_tokens=512, + parallel=128, + host="http://127.0.0.1", + port=int(self.base_url.split(":")[-1]), + ) + metrics = run_eval_few_shot_gsm8k(args) + print(metrics) + + self.assertGreater(metrics["accuracy"], 0.60) + + @unittest.skip("covered in TestMTPWithTBO") class TestMTP(CustomTestCase): @classmethod @@ -393,5 +495,81 @@ class TestMTPWithTBO(CustomTestCase): self.assertGreater(avg_spec_accept_length, 2.1) +@unittest.skip("skipped due to bug when using MTP & TBO & attn_tp_size > 1") +class TestMTPWithTPAttnAndTBO(CustomTestCase): + @classmethod + def setUpClass(cls): + + cls.model = DEFAULT_MODEL_NAME_FOR_TEST_MLA + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=[ + "--tp-size", + "4", + "--moe-dense-tp-size", + "1", + "--enable-two-batch-overlap", + "--moe-a2a-backend", + "deepep", + "--trust-remote-code", + "--speculative-algorithm", + "EAGLE", + "--speculative-num-steps", + "2", + "--speculative-eagle-topk", + "3", + "--speculative-num-draft-tokens", + "3", + "--speculative-draft-model-path", + DEFAULT_MODEL_NAME_FOR_TEST_MLA_NEXTN, + "--chunked-prefill-size", + "256", + "--cuda-graph-max-bs", + "32", + "--max-running-requests", + "128", + "--mem-fraction-static", # temp fix as DeepEP buffer is too large. + "0.7", + ], + env={ + **os.environ, + "SGLANG_TBO_DEBUG": "1", + }, + ) + + @classmethod + def tearDownClass(cls): + kill_process_tree(cls.process.pid) + + def test_gsm8k(self): + args = SimpleNamespace( + num_shots=5, + data_path=None, + num_questions=200, + max_new_tokens=512, + parallel=128, + host="http://127.0.0.1", + port=int(self.base_url.split(":")[-1]), + ) + metrics = run_eval_few_shot_gsm8k(args) + print(metrics) + + self.assertGreater(metrics["accuracy"], 0.60) + + server_info = requests.get(self.base_url + "/get_server_info") + avg_spec_accept_length = server_info.json()["internal_states"][0][ + "avg_spec_accept_length" + ] + print( + f"###test_gsm8k (deepseek-v3 mtp + dp + tbo):\n" + f"accuracy={metrics['accuracy']=:.3f}\n" + f"{avg_spec_accept_length=:.3f}\n" + ) + self.assertGreater(avg_spec_accept_length, 2.1) + + if __name__ == "__main__": unittest.main()