Clean up imports and move files (#14317)
This commit is contained in:
211
python/sglang/srt/batch_overlap/operations.py
Normal file
211
python/sglang/srt/batch_overlap/operations.py
Normal file
@@ -0,0 +1,211 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any, Callable, Dict, Generator, List, Sequence, Union
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.dp_attention import set_dp_buffer_len
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
|
||||
_ENABLE_PROFILE = bool(int(os.environ.get("SGLANG_OPERATIONS_ENABLE_PROFILE", "0")))
|
||||
|
||||
if _ENABLE_PROFILE:
|
||||
import nvtx
|
||||
|
||||
|
||||
def execute_operations(inputs, operations):
|
||||
stages = _convert_operations_to_stages(operations)
|
||||
executor = _StageExecutor("primary", stages, inputs=inputs)
|
||||
for _ in range(executor.num_stages):
|
||||
executor.next()
|
||||
assert executor.done
|
||||
return executor.output
|
||||
|
||||
|
||||
def execute_overlapped_operations(
|
||||
inputs_arr: Sequence,
|
||||
operations_arr: Sequence,
|
||||
delta_stages: Sequence[int],
|
||||
) -> Sequence:
|
||||
# Make it explicit for clarity; if we need multi-batch overlap, this can be generalized
|
||||
inputs_a, inputs_b = inputs_arr
|
||||
operations_a, operations_b = operations_arr
|
||||
delta_stage_a, delta_stage_b = delta_stages
|
||||
assert delta_stage_a == 0
|
||||
delta_stage = delta_stage_b
|
||||
|
||||
stages_a = _convert_operations_to_stages(operations_a)
|
||||
stages_b = _convert_operations_to_stages(operations_b)
|
||||
executor_a = _StageExecutor("a", stages_a, inputs=inputs_a)
|
||||
executor_b = _StageExecutor("b", stages_b, inputs=inputs_b)
|
||||
|
||||
for _ in range(delta_stage):
|
||||
executor_a.next()
|
||||
|
||||
for _ in range(executor_a.num_stages - delta_stage):
|
||||
executor_a.next()
|
||||
executor_b.next()
|
||||
|
||||
for _ in range(delta_stage):
|
||||
executor_b.next()
|
||||
|
||||
assert executor_a.done and executor_b.done
|
||||
return [executor_a.output, executor_b.output]
|
||||
|
||||
|
||||
class YieldOperation:
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class ExecutionOperation:
|
||||
debug_name: str
|
||||
fn: Callable
|
||||
|
||||
|
||||
Operation = Union[YieldOperation, ExecutionOperation, Callable]
|
||||
Stage = List[ExecutionOperation]
|
||||
|
||||
|
||||
class _StageExecutor:
|
||||
def __init__(self, debug_name: str, stages: List[Stage], inputs: dict):
|
||||
self._debug_name = debug_name
|
||||
self._stages = stages
|
||||
self._index = 0
|
||||
self._stage_state = _StateDict()
|
||||
self._stage_output = inputs
|
||||
|
||||
# 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._global_num_tokens = forward_batch.global_num_tokens_cpu
|
||||
self._is_dp_max_padding = forward_batch.dp_padding_mode.is_max_len()
|
||||
|
||||
def next(self):
|
||||
assert not self.done
|
||||
|
||||
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,
|
||||
)
|
||||
|
||||
with _annotate_region(debug_name=f"{self._debug_name}{self._index}"):
|
||||
for op in stage:
|
||||
with _annotate_region(debug_name=op.debug_name):
|
||||
self._stage_output = op.fn(
|
||||
state=self._stage_state,
|
||||
**(
|
||||
self._stage_output if self._stage_output is not None else {}
|
||||
),
|
||||
)
|
||||
|
||||
self._index += 1
|
||||
|
||||
@property
|
||||
def output(self):
|
||||
assert self.done
|
||||
return self._stage_output
|
||||
|
||||
@property
|
||||
def done(self):
|
||||
return self._index >= self.num_stages
|
||||
|
||||
@property
|
||||
def num_stages(self):
|
||||
return len(self._stages)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _annotate_region(debug_name):
|
||||
if _ENABLE_PROFILE:
|
||||
with torch.autograd.profiler.record_function(debug_name):
|
||||
with nvtx.annotate(debug_name):
|
||||
yield
|
||||
else:
|
||||
yield
|
||||
|
||||
|
||||
class _StateDict:
|
||||
def __init__(self):
|
||||
self._data = {}
|
||||
|
||||
def __setattr__(self, key, value):
|
||||
if key == "_data":
|
||||
super().__setattr__(key, value)
|
||||
return
|
||||
assert (
|
||||
key not in self._data
|
||||
), f"`{key}` already exist, are you sure you want to override it?"
|
||||
self._data[key] = value
|
||||
|
||||
def __getattr__(self, item):
|
||||
return self._data[item]
|
||||
|
||||
def __delattr__(self, item):
|
||||
del self._data[item]
|
||||
|
||||
def pop(self, item):
|
||||
return self._data.pop(item)
|
||||
|
||||
def update(self, values: Dict[str, Any]):
|
||||
for k, v in values.items():
|
||||
setattr(self, k, v)
|
||||
|
||||
def get(self, item):
|
||||
return self._data.get(item)
|
||||
|
||||
def clear(self, expect_keys: Sequence[str]):
|
||||
if set(self._data.keys()) != set(expect_keys):
|
||||
raise Exception(
|
||||
f"Unexpected keys when clearning. This may indicate you do not release memory early enough but leave it to here. {list(self._data.keys())=} {expect_keys=}"
|
||||
)
|
||||
|
||||
self._data.clear()
|
||||
|
||||
|
||||
def _convert_operations_to_stages(operations: List[Operation]) -> List[Stage]:
|
||||
operations = _decorate_operations(operations)
|
||||
operation_chunks = list(
|
||||
_chunk_by_separator(operations, lambda op: isinstance(op, YieldOperation))
|
||||
)
|
||||
assert all(len(chunk) > 0 for chunk in operation_chunks)
|
||||
return operation_chunks
|
||||
|
||||
|
||||
def _chunk_by_separator(
|
||||
items: List[Any], is_separator: Callable[[Any], bool]
|
||||
) -> Generator[List[Any], None, None]:
|
||||
pending_items = []
|
||||
for item in items:
|
||||
if is_separator(item):
|
||||
yield pending_items
|
||||
pending_items = []
|
||||
else:
|
||||
pending_items.append(item)
|
||||
if len(pending_items) > 0:
|
||||
yield pending_items
|
||||
|
||||
|
||||
def _decorate_operations(operations: List[Operation], debug_name_prefix: str = ""):
|
||||
return [_decorate_operation(op, debug_name_prefix) for op in operations]
|
||||
|
||||
|
||||
def _decorate_operation(operation: Operation, debug_name_prefix: str):
|
||||
if isinstance(operation, YieldOperation):
|
||||
return operation
|
||||
return ExecutionOperation(
|
||||
debug_name=debug_name_prefix
|
||||
+ getattr(operation, "__name__", "unknown").replace("op_", ""),
|
||||
fn=operation,
|
||||
)
|
||||
211
python/sglang/srt/batch_overlap/operations_strategy.py
Normal file
211
python/sglang/srt/batch_overlap/operations_strategy.py
Normal file
@@ -0,0 +1,211 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.batch_overlap import operations
|
||||
from sglang.srt.batch_overlap.operations import Operation
|
||||
from sglang.srt.layers.moe.token_dispatcher import DeepEPConfig
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||
|
||||
|
||||
@dataclass
|
||||
class OperationsStrategy:
|
||||
operations: List[Operation]
|
||||
deep_gemm_num_sms: Optional[int] = None
|
||||
tbo_delta_stages: Optional[int] = None
|
||||
|
||||
@classmethod
|
||||
def concat(cls, items: List["OperationsStrategy"]) -> "OperationsStrategy":
|
||||
return OperationsStrategy(
|
||||
operations=[x for item in items for x in item.operations],
|
||||
deep_gemm_num_sms=_assert_all_same(
|
||||
[item.deep_gemm_num_sms for item in items]
|
||||
),
|
||||
tbo_delta_stages=_assert_all_same(
|
||||
[item.tbo_delta_stages for item in items]
|
||||
),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def init_new_tbo(
|
||||
layers: torch.nn.ModuleList,
|
||||
forward_mode: ForwardMode,
|
||||
) -> "OperationsStrategy":
|
||||
layer_name = layers[0].__class__.__name__
|
||||
if layer_name == "DeepseekV2DecoderLayer":
|
||||
return OperationsStrategy.concat(
|
||||
[
|
||||
_compute_moe_deepseek_layer_operations_strategy_tbo(
|
||||
layer, forward_mode
|
||||
)
|
||||
for layer in layers
|
||||
]
|
||||
)
|
||||
elif layer_name == "Qwen3MoeDecoderLayer":
|
||||
return OperationsStrategy.concat(
|
||||
[
|
||||
_compute_moe_qwen3_layer_operations_strategy_tbo(
|
||||
layer, forward_mode
|
||||
)
|
||||
for layer in layers
|
||||
]
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
def _assert_all_same(items: List):
|
||||
assert all(item == items[0] for item in items)
|
||||
return items[0]
|
||||
|
||||
|
||||
# -------------------------------- Strategy for DeepSeek ---------------------------------------
|
||||
|
||||
|
||||
# TODO can refactor to make it more fancy if we have more complex strategies
|
||||
def _compute_moe_deepseek_layer_operations_strategy_tbo(
|
||||
layer: torch.nn.Module,
|
||||
forward_mode: ForwardMode,
|
||||
) -> OperationsStrategy:
|
||||
assert layer.is_layer_sparse, "dense layer TBO not yet implemented"
|
||||
if forward_mode == ForwardMode.EXTEND:
|
||||
return _compute_moe_deepseek_blog_prefill(layer)
|
||||
elif (
|
||||
forward_mode == ForwardMode.DECODE or forward_mode == ForwardMode.TARGET_VERIFY
|
||||
):
|
||||
return _compute_moe_deepseek_blog_decode(layer)
|
||||
else:
|
||||
raise NotImplementedError(f"Unsupported {forward_mode=}")
|
||||
|
||||
|
||||
def _compute_moe_deepseek_blog_prefill(layer):
|
||||
device_properties = torch.cuda.get_device_properties(device="cuda")
|
||||
total_num_sms = device_properties.multi_processor_count
|
||||
deep_gemm_num_sms = total_num_sms - DeepEPConfig.get_instance().num_sms
|
||||
|
||||
return OperationsStrategy(
|
||||
deep_gemm_num_sms=deep_gemm_num_sms,
|
||||
tbo_delta_stages=0,
|
||||
operations=[
|
||||
layer.op_comm_prepare_attn,
|
||||
layer.self_attn.op_prepare,
|
||||
layer.self_attn.op_core,
|
||||
layer.op_comm_prepare_mlp,
|
||||
layer.mlp.op_gate,
|
||||
layer.mlp.op_select_experts,
|
||||
layer.mlp.op_dispatch_a,
|
||||
operations.YieldOperation(),
|
||||
layer.mlp.op_dispatch_b,
|
||||
layer.mlp.op_experts,
|
||||
layer.mlp.op_combine_a,
|
||||
operations.YieldOperation(),
|
||||
layer.mlp.op_shared_experts,
|
||||
layer.mlp.op_combine_b,
|
||||
layer.mlp.op_output,
|
||||
layer.op_comm_postprocess_layer,
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
def _compute_moe_deepseek_blog_decode(layer):
|
||||
return OperationsStrategy(
|
||||
deep_gemm_num_sms=None,
|
||||
tbo_delta_stages=2,
|
||||
operations=[
|
||||
layer.op_comm_prepare_attn,
|
||||
layer.self_attn.op_prepare,
|
||||
operations.YieldOperation(),
|
||||
layer.self_attn.op_core,
|
||||
layer.op_comm_prepare_mlp,
|
||||
layer.mlp.op_gate,
|
||||
layer.mlp.op_select_experts,
|
||||
operations.YieldOperation(),
|
||||
layer.mlp.op_dispatch_a,
|
||||
layer.mlp.op_shared_experts,
|
||||
operations.YieldOperation(),
|
||||
layer.mlp.op_dispatch_b,
|
||||
layer.mlp.op_experts,
|
||||
layer.mlp.op_combine_a,
|
||||
operations.YieldOperation(),
|
||||
layer.mlp.op_combine_b,
|
||||
operations.YieldOperation(),
|
||||
layer.mlp.op_output,
|
||||
layer.op_comm_postprocess_layer,
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
# -------------------------------- Strategy for Qwen3 ---------------------------------------
|
||||
|
||||
|
||||
# TODO: unstable, current strategy is almost the same as DeepSeek, keep redundant code here for
|
||||
# convenience to adjust strategy
|
||||
def _compute_moe_qwen3_layer_operations_strategy_tbo(
|
||||
layer: torch.nn.Module,
|
||||
forward_mode: ForwardMode,
|
||||
) -> OperationsStrategy:
|
||||
assert layer.is_layer_sparse, "qwen3 moe only support sparse layers"
|
||||
if forward_mode == ForwardMode.EXTEND:
|
||||
return _compute_moe_qwen3_prefill(layer)
|
||||
elif (
|
||||
forward_mode == ForwardMode.DECODE or forward_mode == ForwardMode.TARGET_VERIFY
|
||||
):
|
||||
return _compute_moe_qwen3_decode(layer)
|
||||
else:
|
||||
raise NotImplementedError(f"Unsupported {forward_mode=}")
|
||||
|
||||
|
||||
def _compute_moe_qwen3_prefill(layer):
|
||||
device_properties = torch.cuda.get_device_properties(device="cuda")
|
||||
total_num_sms = device_properties.multi_processor_count
|
||||
deep_gemm_num_sms = total_num_sms - DeepEPConfig.get_instance().num_sms
|
||||
|
||||
return OperationsStrategy(
|
||||
deep_gemm_num_sms=deep_gemm_num_sms,
|
||||
tbo_delta_stages=0,
|
||||
operations=[
|
||||
layer.op_comm_prepare_attn,
|
||||
layer.self_attn.op_prepare,
|
||||
layer.self_attn.op_core,
|
||||
layer.op_comm_prepare_mlp,
|
||||
layer.mlp.op_gate,
|
||||
layer.mlp.op_select_experts,
|
||||
layer.mlp.op_dispatch_a,
|
||||
operations.YieldOperation(),
|
||||
layer.mlp.op_dispatch_b,
|
||||
layer.mlp.op_experts,
|
||||
layer.mlp.op_combine_a,
|
||||
operations.YieldOperation(),
|
||||
layer.mlp.op_combine_b,
|
||||
layer.mlp.op_output,
|
||||
layer.op_comm_postprocess_layer,
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
def _compute_moe_qwen3_decode(layer):
|
||||
return OperationsStrategy(
|
||||
deep_gemm_num_sms=None,
|
||||
tbo_delta_stages=2,
|
||||
operations=[
|
||||
layer.op_comm_prepare_attn,
|
||||
layer.self_attn.op_prepare,
|
||||
operations.YieldOperation(),
|
||||
layer.self_attn.op_core,
|
||||
layer.op_comm_prepare_mlp,
|
||||
layer.mlp.op_gate,
|
||||
layer.mlp.op_select_experts,
|
||||
operations.YieldOperation(),
|
||||
layer.mlp.op_dispatch_a,
|
||||
operations.YieldOperation(),
|
||||
layer.mlp.op_dispatch_b,
|
||||
layer.mlp.op_experts,
|
||||
layer.mlp.op_combine_a,
|
||||
operations.YieldOperation(),
|
||||
layer.mlp.op_combine_b,
|
||||
layer.mlp.op_output,
|
||||
layer.op_comm_postprocess_layer,
|
||||
operations.YieldOperation(),
|
||||
],
|
||||
)
|
||||
116
python/sglang/srt/batch_overlap/single_batch_overlap.py
Normal file
116
python/sglang/srt/batch_overlap/single_batch_overlap.py
Normal file
@@ -0,0 +1,116 @@
|
||||
# Copyright 2025 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.moe import get_moe_runner_backend
|
||||
from sglang.srt.layers.moe.utils import is_sbo_enabled
|
||||
from sglang.srt.utils import get_int_env_var
|
||||
|
||||
|
||||
class SboFlags:
|
||||
# TODO may have: "enable_dispatch_shared_one_stream_overlap", "enable_dispatch_gateup_gemm_two_stream_overlap", ...
|
||||
|
||||
@classmethod
|
||||
def enable_combine_down_gemm_two_stream_overlap(cls):
|
||||
return (
|
||||
is_sbo_enabled()
|
||||
# currently only cutedsl backend supports it
|
||||
and get_moe_runner_backend().is_flashinfer_cutedsl()
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def enable_combine_shared_two_stream_overlap(cls):
|
||||
return is_sbo_enabled()
|
||||
|
||||
@classmethod
|
||||
def fuse_shared_experts_inside_sbo(cls):
|
||||
# TODO after antgroup's PR, should be `... or cls.enable_dispatch_shared_one_stream_overlap()`
|
||||
return cls.enable_combine_shared_two_stream_overlap()
|
||||
|
||||
|
||||
@dataclass
|
||||
class CombineOverlapArgs:
|
||||
# this "overlap" flag means overlapping with down gemm, not the general two-stream overlap
|
||||
overlap: bool
|
||||
stream: torch.cuda.Stream
|
||||
wait_event: torch.cuda.Event
|
||||
num_sms: int
|
||||
signal: Optional[torch.Tensor] = None
|
||||
threshold: int = 0
|
||||
|
||||
|
||||
@dataclass
|
||||
class DownGemmOverlapArgs:
|
||||
num_sms: int
|
||||
signal: torch.Tensor
|
||||
start_event: torch.cuda.Event
|
||||
|
||||
|
||||
def compute_overlap_args(dispatch_output, alt_stream):
|
||||
if not (
|
||||
SboFlags.enable_combine_down_gemm_two_stream_overlap()
|
||||
or SboFlags.enable_combine_shared_two_stream_overlap()
|
||||
):
|
||||
return None, None, {}
|
||||
|
||||
hidden_states = dispatch_output.hidden_states
|
||||
|
||||
num_local_experts, num_tokens_static, hidden_dim = hidden_states.shape
|
||||
|
||||
total_num_sms = torch.cuda.get_device_properties(
|
||||
device="cuda"
|
||||
).multi_processor_count
|
||||
communicate_num_sms = get_int_env_var("SGLANG_DEEPEP_LL_COMBINE_SEND_NUM_SMS", 32)
|
||||
compute_num_sms = total_num_sms - communicate_num_sms
|
||||
|
||||
assert alt_stream is not None
|
||||
combine_wait_event = torch.cuda.Event()
|
||||
combine_overlap_args = CombineOverlapArgs(
|
||||
overlap=False,
|
||||
num_sms=communicate_num_sms,
|
||||
stream=alt_stream,
|
||||
wait_event=combine_wait_event,
|
||||
)
|
||||
meta_overlap_args = dict(
|
||||
compute_num_sms=compute_num_sms,
|
||||
)
|
||||
down_gemm_overlap_args = None
|
||||
|
||||
if SboFlags.enable_combine_down_gemm_two_stream_overlap():
|
||||
# TODO use zero_allocator to remove this `torch.zeros` call
|
||||
# NOTE ours v2 use uint32 not int32 currently
|
||||
combine_signal = torch.zeros(
|
||||
num_local_experts, dtype=torch.uint32, device=hidden_states.device
|
||||
)
|
||||
|
||||
down_gemm_overlap_args = DownGemmOverlapArgs(
|
||||
signal=combine_signal,
|
||||
start_event=combine_wait_event,
|
||||
num_sms=compute_num_sms,
|
||||
)
|
||||
combine_overlap_args.overlap = True
|
||||
combine_overlap_args.signal = combine_signal
|
||||
combine_overlap_args.threshold = compute_num_sms
|
||||
else:
|
||||
meta_overlap_args |= dict(
|
||||
record_event_after_down=combine_wait_event,
|
||||
)
|
||||
|
||||
return combine_overlap_args, down_gemm_overlap_args, meta_overlap_args
|
||||
1027
python/sglang/srt/batch_overlap/two_batch_overlap.py
Normal file
1027
python/sglang/srt/batch_overlap/two_batch_overlap.py
Normal file
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user