[bugfix] fix TBO crashes when attn_tp_size > 1 (#13730)
Co-authored-by: Cheng Wan <54331508+ch-wan@users.noreply.github.com>
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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")
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user