[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
|
# handling DP attention
|
||||||
forward_batch: ForwardBatch = inputs["forward_batch"]
|
forward_batch: ForwardBatch = inputs["forward_batch"]
|
||||||
self._global_dp_buffer_len = forward_batch.global_dp_buffer_len
|
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._global_num_tokens = forward_batch.global_num_tokens_cpu
|
||||||
self._is_dp_max_padding = forward_batch.dp_padding_mode.is_max_len()
|
self._is_dp_max_padding = forward_batch.dp_padding_mode.is_max_len()
|
||||||
|
|
||||||
@@ -92,7 +92,9 @@ class _StageExecutor:
|
|||||||
|
|
||||||
stage = self._stages[self._index]
|
stage = self._stages[self._index]
|
||||||
|
|
||||||
if self._global_dp_buffer_len is not None:
|
# 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(
|
set_dp_buffer_len(
|
||||||
self._global_dp_buffer_len,
|
self._global_dp_buffer_len,
|
||||||
self._local_dp_buffer_len,
|
self._local_dp_buffer_len,
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ from sglang.srt.layers.communicator import (
|
|||||||
CommunicateSummableTensorPairFn,
|
CommunicateSummableTensorPairFn,
|
||||||
ScatterMode,
|
ScatterMode,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
||||||
from sglang.srt.layers.moe import (
|
from sglang.srt.layers.moe import (
|
||||||
get_deepep_mode,
|
get_deepep_mode,
|
||||||
get_moe_a2a_backend,
|
get_moe_a2a_backend,
|
||||||
@@ -630,6 +631,11 @@ class TboForwardBatchPreparer:
|
|||||||
), f"{key=} {old_value=} {num_tokens=} {batch=}"
|
), f"{key=} {old_value=} {num_tokens=} {batch=}"
|
||||||
output_dict[key] = old_value[start_token_index:end_token_index]
|
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 [
|
for key in [
|
||||||
"req_pool_indices",
|
"req_pool_indices",
|
||||||
"seq_lens",
|
"seq_lens",
|
||||||
@@ -840,6 +846,7 @@ def _model_forward_tbo(
|
|||||||
input_data_scatter_mode=input_data_scatter_mode,
|
input_data_scatter_mode=input_data_scatter_mode,
|
||||||
layer_input_scatter_mode=layer_input_scatter_mode,
|
layer_input_scatter_mode=layer_input_scatter_mode,
|
||||||
)
|
)
|
||||||
|
original_hidden_states_len = inputs["hidden_states"].shape[0]
|
||||||
del inputs
|
del inputs
|
||||||
|
|
||||||
context = (
|
context = (
|
||||||
@@ -857,7 +864,7 @@ def _model_forward_tbo(
|
|||||||
delta_stages=[0, operations_strategy.tbo_delta_stages],
|
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):
|
def _model_forward_non_tbo(inputs, operations_strategy: OperationsStrategy):
|
||||||
@@ -951,23 +958,49 @@ def _model_forward_filter_inputs(
|
|||||||
tbo_subbatch_index: int,
|
tbo_subbatch_index: int,
|
||||||
) -> Dict:
|
) -> Dict:
|
||||||
token_slice = slice(*output_forward_batch.tbo_parent_token_range)
|
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(
|
return dict(
|
||||||
hidden_states=hidden_states[token_slice],
|
hidden_states=_pad(hidden_states),
|
||||||
residual=None if residual is None else residual[token_slice],
|
residual=_pad(residual),
|
||||||
positions=positions[token_slice],
|
positions=_pad(positions),
|
||||||
forward_batch=output_forward_batch,
|
forward_batch=output_forward_batch,
|
||||||
tbo_subbatch_index=tbo_subbatch_index,
|
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):
|
def _handle_key(name):
|
||||||
value_a = output_a[name]
|
value_a = output_a[name]
|
||||||
value_b = output_b[name]
|
value_b = output_b[name]
|
||||||
assert (value_a is None) == (value_b is None)
|
assert (value_a is None) == (value_b is None)
|
||||||
if value_a is None:
|
if value_a is None:
|
||||||
return 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")
|
return _handle_key("hidden_states"), _handle_key("residual")
|
||||||
|
|
||||||
|
|||||||
@@ -217,14 +217,16 @@ class _LayerModeComputationContext:
|
|||||||
layer_id: int
|
layer_id: int
|
||||||
is_layer_sparse: bool
|
is_layer_sparse: bool
|
||||||
is_previous_layer_sparse: Optional[bool]
|
is_previous_layer_sparse: Optional[bool]
|
||||||
|
is_next_layer_sparse: Optional[bool]
|
||||||
|
|
||||||
def previous_layer(self):
|
def previous_layer(self):
|
||||||
assert self.is_previous_layer_sparse is not None
|
assert self.is_previous_layer_sparse is not None
|
||||||
return _LayerModeComputationContext(
|
return _LayerModeComputationContext(
|
||||||
|
num_layers=self.num_layers,
|
||||||
layer_id=self.layer_id - 1,
|
layer_id=self.layer_id - 1,
|
||||||
is_layer_sparse=self.is_previous_layer_sparse,
|
is_layer_sparse=self.is_previous_layer_sparse,
|
||||||
is_previous_layer_sparse=None,
|
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
|
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
|
@classmethod
|
||||||
def _compute_middle_residual_mode(cls, context: _LayerModeComputationContext):
|
def _compute_middle_residual_mode(cls, context: _LayerModeComputationContext):
|
||||||
mlp_mode = cls._compute_mlp_mode(context)
|
mlp_mode = cls._compute_mlp_mode(context)
|
||||||
@@ -288,6 +299,8 @@ class LayerScatterModes:
|
|||||||
if context.layer_id == context.num_layers - 1:
|
if context.layer_id == context.num_layers - 1:
|
||||||
return ScatterMode.model_input_output()
|
return ScatterMode.model_input_output()
|
||||||
if mlp_mode == ScatterMode.SCATTERED:
|
if mlp_mode == ScatterMode.SCATTERED:
|
||||||
|
if cls._should_gather_for_tbo(context):
|
||||||
|
return ScatterMode.TP_ATTN_FULL
|
||||||
return ScatterMode.SCATTERED
|
return ScatterMode.SCATTERED
|
||||||
if mlp_mode == ScatterMode.FULL:
|
if mlp_mode == ScatterMode.FULL:
|
||||||
return ScatterMode.TP_ATTN_FULL
|
return ScatterMode.TP_ATTN_FULL
|
||||||
|
|||||||
@@ -376,6 +376,7 @@ class ForwardBatch:
|
|||||||
# For two-batch overlap
|
# For two-batch overlap
|
||||||
tbo_split_seq_index: Optional[int] = None
|
tbo_split_seq_index: Optional[int] = None
|
||||||
tbo_parent_token_range: Optional[Tuple[int, int]] = None
|
tbo_parent_token_range: Optional[Tuple[int, int]] = None
|
||||||
|
tbo_padded_len: Optional[int] = None
|
||||||
tbo_children: Optional[List[ForwardBatch]] = None
|
tbo_children: Optional[List[ForwardBatch]] = None
|
||||||
|
|
||||||
# For matryoshka embeddings
|
# For matryoshka embeddings
|
||||||
@@ -852,6 +853,14 @@ class ForwardBatch:
|
|||||||
TboForwardBatchPreparer.prepare(
|
TboForwardBatchPreparer.prepare(
|
||||||
batch=self, is_draft_worker=model_runner.is_draft_worker
|
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):
|
def _pad_inputs_to_size(self, model_runner: ModelRunner, num_tokens, bs):
|
||||||
# padding
|
# padding
|
||||||
|
|||||||
@@ -582,12 +582,16 @@ class BailingMoEBlock(nn.Module):
|
|||||||
is_previous_layer_sparse = self._is_layer_sparse(
|
is_previous_layer_sparse = self._is_layer_sparse(
|
||||||
config, layer_id=layer_id - 1, is_nextn=False
|
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(
|
self.layer_scatter_modes = LayerScatterModes.init_new(
|
||||||
layer_id=layer_id,
|
layer_id=layer_id,
|
||||||
num_layers=config.num_hidden_layers,
|
num_layers=config.num_hidden_layers,
|
||||||
is_layer_sparse=self.is_layer_sparse,
|
is_layer_sparse=self.is_layer_sparse,
|
||||||
is_previous_layer_sparse=is_previous_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
|
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)
|
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_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(
|
self.layer_scatter_modes = LayerScatterModes.init_new(
|
||||||
layer_id=layer_id,
|
layer_id=layer_id,
|
||||||
num_layers=1 if is_nextn else config.num_hidden_layers,
|
num_layers=1 if is_nextn else config.num_hidden_layers,
|
||||||
is_layer_sparse=self.is_layer_sparse,
|
is_layer_sparse=self.is_layer_sparse,
|
||||||
is_previous_layer_sparse=is_previous_layer_sparse,
|
is_previous_layer_sparse=is_previous_layer_sparse,
|
||||||
|
is_next_layer_sparse=is_next_layer_sparse,
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.is_layer_sparse:
|
if self.is_layer_sparse:
|
||||||
|
|||||||
@@ -198,15 +198,17 @@ class FalconH1HybridAttentionDecoderLayer(nn.Module):
|
|||||||
prefix=f"{prefix}.mixer",
|
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
|
self.is_layer_sparse = False
|
||||||
is_previous_layer_sparse = False
|
is_previous_layer_sparse = False
|
||||||
|
is_next_layer_sparse = False
|
||||||
|
|
||||||
self.layer_scatter_modes = LayerScatterModes.init_new(
|
self.layer_scatter_modes = LayerScatterModes.init_new(
|
||||||
layer_id=layer_id,
|
layer_id=layer_id,
|
||||||
num_layers=config.num_hidden_layers,
|
num_layers=config.num_hidden_layers,
|
||||||
is_layer_sparse=self.is_layer_sparse,
|
is_layer_sparse=self.is_layer_sparse,
|
||||||
is_previous_layer_sparse=is_previous_layer_sparse,
|
is_previous_layer_sparse=is_previous_layer_sparse,
|
||||||
|
is_next_layer_sparse=is_next_layer_sparse,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.feed_forward = FalconH1MLP(
|
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)
|
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_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(
|
self.layer_scatter_modes = LayerScatterModes.init_new(
|
||||||
layer_id=layer_id,
|
layer_id=layer_id,
|
||||||
num_layers=1 if is_nextn else config.num_hidden_layers,
|
num_layers=1 if is_nextn else config.num_hidden_layers,
|
||||||
is_layer_sparse=self.is_layer_sparse,
|
is_layer_sparse=self.is_layer_sparse,
|
||||||
is_previous_layer_sparse=is_previous_layer_sparse,
|
is_previous_layer_sparse=is_previous_layer_sparse,
|
||||||
|
is_next_layer_sparse=is_next_layer_sparse,
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.is_layer_sparse:
|
if self.is_layer_sparse:
|
||||||
|
|||||||
@@ -395,12 +395,14 @@ class GptOssDecoderLayer(nn.Module):
|
|||||||
self.is_layer_sparse = True
|
self.is_layer_sparse = True
|
||||||
self.is_nextn = False
|
self.is_nextn = False
|
||||||
is_previous_layer_sparse = True
|
is_previous_layer_sparse = True
|
||||||
|
is_next_layer_sparse = True
|
||||||
|
|
||||||
self.layer_scatter_modes = LayerScatterModes.init_new(
|
self.layer_scatter_modes = LayerScatterModes.init_new(
|
||||||
layer_id=layer_id,
|
layer_id=layer_id,
|
||||||
num_layers=config.num_hidden_layers,
|
num_layers=config.num_hidden_layers,
|
||||||
is_layer_sparse=self.is_layer_sparse,
|
is_layer_sparse=self.is_layer_sparse,
|
||||||
is_previous_layer_sparse=is_previous_layer_sparse,
|
is_previous_layer_sparse=is_previous_layer_sparse,
|
||||||
|
is_next_layer_sparse=is_next_layer_sparse,
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.is_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)
|
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_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(
|
self.layer_scatter_modes = LayerScatterModes.init_new(
|
||||||
layer_id=layer_id,
|
layer_id=layer_id,
|
||||||
num_layers=config.num_hidden_layers,
|
num_layers=config.num_hidden_layers,
|
||||||
is_layer_sparse=self.is_layer_sparse,
|
is_layer_sparse=self.is_layer_sparse,
|
||||||
is_previous_layer_sparse=is_previous_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
|
self.is_last_layer = self.layer_id == config.num_hidden_layers - 1
|
||||||
|
|||||||
@@ -384,6 +384,7 @@ class Llama4DecoderLayer(nn.Module):
|
|||||||
self.config = config
|
self.config = config
|
||||||
is_moe_layer = self._is_moe_layer(layer_id)
|
is_moe_layer = self._is_moe_layer(layer_id)
|
||||||
is_previous_moe_layer = self._is_moe_layer(layer_id - 1)
|
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:
|
if is_moe_layer:
|
||||||
self.feed_forward = Llama4MoE(
|
self.feed_forward = Llama4MoE(
|
||||||
@@ -410,6 +411,7 @@ class Llama4DecoderLayer(nn.Module):
|
|||||||
num_layers=config.num_hidden_layers,
|
num_layers=config.num_hidden_layers,
|
||||||
is_layer_sparse=is_moe_layer,
|
is_layer_sparse=is_moe_layer,
|
||||||
is_previous_layer_sparse=is_previous_moe_layer,
|
is_previous_layer_sparse=is_previous_moe_layer,
|
||||||
|
is_next_layer_sparse=is_next_moe_layer,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.layer_communicator = LayerCommunicator(
|
self.layer_communicator = LayerCommunicator(
|
||||||
|
|||||||
@@ -380,6 +380,8 @@ class LongcatFlashDecoderLayer(nn.Module):
|
|||||||
num_layers=config.num_hidden_layers,
|
num_layers=config.num_hidden_layers,
|
||||||
is_layer_sparse=False,
|
is_layer_sparse=False,
|
||||||
is_previous_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)
|
for i in range(2)
|
||||||
]
|
]
|
||||||
@@ -398,6 +400,8 @@ class LongcatFlashDecoderLayer(nn.Module):
|
|||||||
num_layers=config.num_hidden_layers,
|
num_layers=config.num_hidden_layers,
|
||||||
is_layer_sparse=True,
|
is_layer_sparse=True,
|
||||||
is_previous_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(
|
self.moe_layer_communicator = LayerCommunicator(
|
||||||
layer_scatter_modes=self.moe_layer_scatter_modes,
|
layer_scatter_modes=self.moe_layer_scatter_modes,
|
||||||
|
|||||||
@@ -161,6 +161,7 @@ class LongcatFlashDenseDecoderLayer(nn.Module):
|
|||||||
num_layers=config.num_hidden_layers,
|
num_layers=config.num_hidden_layers,
|
||||||
is_layer_sparse=False,
|
is_layer_sparse=False,
|
||||||
is_previous_layer_sparse=False,
|
is_previous_layer_sparse=False,
|
||||||
|
is_next_layer_sparse=False,
|
||||||
)
|
)
|
||||||
self.layer_communicator = LayerCommunicator(
|
self.layer_communicator = LayerCommunicator(
|
||||||
layer_scatter_modes=self.layer_scatter_modes,
|
layer_scatter_modes=self.layer_scatter_modes,
|
||||||
|
|||||||
@@ -516,11 +516,13 @@ class MiniMaxM2DecoderLayer(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
is_previous_layer_sparse = True
|
is_previous_layer_sparse = True
|
||||||
|
is_next_layer_sparse = True
|
||||||
self.layer_scatter_modes = LayerScatterModes.init_new(
|
self.layer_scatter_modes = LayerScatterModes.init_new(
|
||||||
layer_id=layer_id,
|
layer_id=layer_id,
|
||||||
num_layers=config.num_hidden_layers,
|
num_layers=config.num_hidden_layers,
|
||||||
is_layer_sparse=self.is_layer_sparse,
|
is_layer_sparse=self.is_layer_sparse,
|
||||||
is_previous_layer_sparse=is_previous_layer_sparse,
|
is_previous_layer_sparse=is_previous_layer_sparse,
|
||||||
|
is_next_layer_sparse=is_next_layer_sparse,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.layer_communicator = LayerCommunicator(
|
self.layer_communicator = LayerCommunicator(
|
||||||
|
|||||||
@@ -461,12 +461,14 @@ class Qwen2MoeDecoderLayer(nn.Module):
|
|||||||
# Qwen2MoE all layers are sparse and have no nextn now
|
# Qwen2MoE all layers are sparse and have no nextn now
|
||||||
self.is_layer_sparse = True
|
self.is_layer_sparse = True
|
||||||
is_previous_layer_sparse = True
|
is_previous_layer_sparse = True
|
||||||
|
is_next_layer_sparse = True
|
||||||
|
|
||||||
self.layer_scatter_modes = LayerScatterModes.init_new(
|
self.layer_scatter_modes = LayerScatterModes.init_new(
|
||||||
layer_id=layer_id,
|
layer_id=layer_id,
|
||||||
num_layers=config.num_hidden_layers,
|
num_layers=config.num_hidden_layers,
|
||||||
is_layer_sparse=self.is_layer_sparse,
|
is_layer_sparse=self.is_layer_sparse,
|
||||||
is_previous_layer_sparse=is_previous_layer_sparse,
|
is_previous_layer_sparse=is_previous_layer_sparse,
|
||||||
|
is_next_layer_sparse=is_next_layer_sparse,
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.is_layer_sparse:
|
if self.is_layer_sparse:
|
||||||
|
|||||||
@@ -276,6 +276,7 @@ class Qwen3DecoderLayer(nn.Module):
|
|||||||
num_layers=config.num_hidden_layers,
|
num_layers=config.num_hidden_layers,
|
||||||
is_layer_sparse=False,
|
is_layer_sparse=False,
|
||||||
is_previous_layer_sparse=False,
|
is_previous_layer_sparse=False,
|
||||||
|
is_next_layer_sparse=False,
|
||||||
)
|
)
|
||||||
self.layer_communicator = LayerCommunicator(
|
self.layer_communicator = LayerCommunicator(
|
||||||
layer_scatter_modes=self.layer_scatter_modes,
|
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
|
# Qwen3MoE all layers are sparse and have no nextn now
|
||||||
self.is_layer_sparse = True
|
self.is_layer_sparse = True
|
||||||
is_previous_layer_sparse = True
|
is_previous_layer_sparse = True
|
||||||
|
is_next_layer_sparse = True
|
||||||
|
|
||||||
self.layer_scatter_modes = LayerScatterModes.init_new(
|
self.layer_scatter_modes = LayerScatterModes.init_new(
|
||||||
layer_id=layer_id,
|
layer_id=layer_id,
|
||||||
num_layers=config.num_hidden_layers,
|
num_layers=config.num_hidden_layers,
|
||||||
is_layer_sparse=self.is_layer_sparse,
|
is_layer_sparse=self.is_layer_sparse,
|
||||||
is_previous_layer_sparse=is_previous_layer_sparse,
|
is_previous_layer_sparse=is_previous_layer_sparse,
|
||||||
|
is_next_layer_sparse=is_next_layer_sparse,
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.is_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
|
# Qwen3Next all layers are sparse and have no nextn now
|
||||||
self.is_layer_sparse = True
|
self.is_layer_sparse = True
|
||||||
is_previous_layer_sparse = True
|
is_previous_layer_sparse = True
|
||||||
|
is_next_layer_sparse = True
|
||||||
self.layer_id = layer_id
|
self.layer_id = layer_id
|
||||||
|
|
||||||
self.layer_scatter_modes = LayerScatterModes.init_new(
|
self.layer_scatter_modes = LayerScatterModes.init_new(
|
||||||
@@ -499,6 +500,7 @@ class Qwen3HybridLinearDecoderLayer(nn.Module):
|
|||||||
num_layers=config.num_hidden_layers,
|
num_layers=config.num_hidden_layers,
|
||||||
is_layer_sparse=self.is_layer_sparse,
|
is_layer_sparse=self.is_layer_sparse,
|
||||||
is_previous_layer_sparse=is_previous_layer_sparse,
|
is_previous_layer_sparse=is_previous_layer_sparse,
|
||||||
|
is_next_layer_sparse=is_next_layer_sparse,
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.is_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
|
# Qwen3Next all layers are sparse and have no nextn now
|
||||||
self.is_layer_sparse = True
|
self.is_layer_sparse = True
|
||||||
is_previous_layer_sparse = True
|
is_previous_layer_sparse = True
|
||||||
|
is_next_layer_sparse = True
|
||||||
|
|
||||||
self.layer_scatter_modes = LayerScatterModes.init_new(
|
self.layer_scatter_modes = LayerScatterModes.init_new(
|
||||||
layer_id=layer_id,
|
layer_id=layer_id,
|
||||||
num_layers=config.num_hidden_layers,
|
num_layers=config.num_hidden_layers,
|
||||||
is_layer_sparse=self.is_layer_sparse,
|
is_layer_sparse=self.is_layer_sparse,
|
||||||
is_previous_layer_sparse=is_previous_layer_sparse,
|
is_previous_layer_sparse=is_previous_layer_sparse,
|
||||||
|
is_next_layer_sparse=is_next_layer_sparse,
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.is_layer_sparse:
|
if self.is_layer_sparse:
|
||||||
|
|||||||
@@ -337,12 +337,14 @@ class Step3TextDecoderLayer(nn.Module):
|
|||||||
self.is_previous_layer_sparse = (
|
self.is_previous_layer_sparse = (
|
||||||
True if layer_id - 1 in moe_layers_idx else False
|
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(
|
self.layer_scatter_modes = LayerScatterModes.init_new(
|
||||||
layer_id=layer_id,
|
layer_id=layer_id,
|
||||||
num_layers=config.num_hidden_layers,
|
num_layers=config.num_hidden_layers,
|
||||||
is_layer_sparse=self.is_layer_sparse,
|
is_layer_sparse=self.is_layer_sparse,
|
||||||
is_previous_layer_sparse=self.is_previous_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:
|
if not self.is_layer_sparse:
|
||||||
|
|||||||
@@ -251,6 +251,108 @@ class TestTBO(CustomTestCase):
|
|||||||
self.assertGreater(metrics["accuracy"], 0.60)
|
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")
|
@unittest.skip("covered in TestMTPWithTBO")
|
||||||
class TestMTP(CustomTestCase):
|
class TestMTP(CustomTestCase):
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -393,5 +495,81 @@ class TestMTPWithTBO(CustomTestCase):
|
|||||||
self.assertGreater(avg_spec_accept_length, 2.1)
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user