[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:
Yuhao Yao
2025-12-12 10:18:40 +08:00
committed by GitHub
parent 0aa3dec5c7
commit e9e7f15eb5
20 changed files with 285 additions and 16 deletions

View File

@@ -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:

View File

@@ -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")

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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:

View File

@@ -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(

View File

@@ -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:

View File

@@ -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:

View File

@@ -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

View File

@@ -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(

View File

@@ -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,

View File

@@ -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,

View File

@@ -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(

View File

@@ -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:

View File

@@ -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,

View File

@@ -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:

View File

@@ -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:

View File

@@ -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:

View File

@@ -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()