Slightly improve the sampler to skip unnecessary steps (#6956)
This commit is contained in:
@@ -1,17 +1,15 @@
|
||||
import dataclasses
|
||||
import logging
|
||||
from typing import TYPE_CHECKING, Dict, List, Optional, Sequence
|
||||
from typing import Dict, List, Optional, Sequence
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||
from sglang.srt.layers.communicator import (
|
||||
CommunicateContext,
|
||||
CommunicateSimpleFn,
|
||||
CommunicateSummableTensorPairFn,
|
||||
ScatterMode,
|
||||
)
|
||||
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
||||
from sglang.srt.layers.moe.ep_moe.token_dispatcher import DeepEPDispatcher
|
||||
from sglang.srt.layers.quantization.deep_gemm import configure_deep_gemm_num_sms
|
||||
from sglang.srt.managers.schedule_batch import global_server_args_dict
|
||||
@@ -20,9 +18,6 @@ from sglang.srt.operations import execute_operations, execute_overlapped_operati
|
||||
from sglang.srt.operations_strategy import OperationsStrategy
|
||||
from sglang.srt.utils import BumpAllocator, DeepEPMode, get_bool_env_var
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.model_executor.cuda_graph_runner import CudaGraphRunner
|
||||
|
||||
_tbo_debug = get_bool_env_var("SGLANG_TBO_DEBUG")
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -46,7 +41,7 @@ def compute_split_seq_index(
|
||||
assert num_tokens == 0
|
||||
return 0
|
||||
else:
|
||||
raise NotImplementedError
|
||||
raise NotImplementedError()
|
||||
|
||||
|
||||
def _split_array_by_half_sum(arr: Sequence[int]) -> int:
|
||||
|
||||
Reference in New Issue
Block a user