Clean up imports and move files (#14317)
This commit is contained in:
@@ -18,6 +18,7 @@ dependencies = [
|
|||||||
"IPython",
|
"IPython",
|
||||||
"aiohttp",
|
"aiohttp",
|
||||||
"anthropic>=0.20.0",
|
"anthropic>=0.20.0",
|
||||||
|
"av ; sys_platform == 'linux' and (platform_machine == 'aarch64' or platform_machine == 'arm64' and platform_machine == 'armv7l')",
|
||||||
"blobfile==3.0.0",
|
"blobfile==3.0.0",
|
||||||
"build",
|
"build",
|
||||||
"compressed-tensors",
|
"compressed-tensors",
|
||||||
@@ -38,6 +39,7 @@ dependencies = [
|
|||||||
"ninja",
|
"ninja",
|
||||||
"numpy",
|
"numpy",
|
||||||
"nvidia-cutlass-dsl==4.2.1",
|
"nvidia-cutlass-dsl==4.2.1",
|
||||||
|
"nvidia-ml-py",
|
||||||
"openai-harmony==0.0.4",
|
"openai-harmony==0.0.4",
|
||||||
"openai==2.6.1",
|
"openai==2.6.1",
|
||||||
"orjson",
|
"orjson",
|
||||||
@@ -50,7 +52,6 @@ dependencies = [
|
|||||||
"py-spy",
|
"py-spy",
|
||||||
"pybase64",
|
"pybase64",
|
||||||
"pydantic",
|
"pydantic",
|
||||||
"nvidia-ml-py",
|
|
||||||
"python-multipart",
|
"python-multipart",
|
||||||
"pyzmq>=25.1.2",
|
"pyzmq>=25.1.2",
|
||||||
"requests",
|
"requests",
|
||||||
@@ -63,9 +64,8 @@ dependencies = [
|
|||||||
"timm==1.0.16",
|
"timm==1.0.16",
|
||||||
"torch_memory_saver==0.0.9",
|
"torch_memory_saver==0.0.9",
|
||||||
"torch==2.9.1",
|
"torch==2.9.1",
|
||||||
"torchcodec==0.7.0 ; sys_platform != 'linux' or (sys_platform == 'linux' and platform_machine != 'aarch64' and platform_machine != 'arm64' and platform_machine != 'armv7l')", # torchcodec does not exist in those systems. If not provided, transformer will use torchvision instead by default.
|
|
||||||
"av ; sys_platform == 'linux' and (platform_machine == 'aarch64' or platform_machine == 'arm64' and platform_machine == 'armv7l')",
|
|
||||||
"torchaudio==2.9.1",
|
"torchaudio==2.9.1",
|
||||||
|
"torchcodec==0.7.0 ; sys_platform != 'linux' or (sys_platform == 'linux' and platform_machine != 'aarch64' and platform_machine != 'arm64' and platform_machine != 'armv7l')", # torchcodec does not exist in those systems. If not provided, transformer will use torchvision instead by default.
|
||||||
"torchvision",
|
"torchvision",
|
||||||
"torchao==0.9.0",
|
"torchao==0.9.0",
|
||||||
"tqdm",
|
"tqdm",
|
||||||
@@ -73,6 +73,7 @@ dependencies = [
|
|||||||
"uvicorn",
|
"uvicorn",
|
||||||
"uvloop",
|
"uvloop",
|
||||||
"xgrammar==0.1.27",
|
"xgrammar==0.1.27",
|
||||||
|
|
||||||
"grpcio==1.75.1", # keep it align with compile_proto.py
|
"grpcio==1.75.1", # keep it align with compile_proto.py
|
||||||
"grpcio-tools==1.75.1", # keep it align with compile_proto.py
|
"grpcio-tools==1.75.1", # keep it align with compile_proto.py
|
||||||
"grpcio-reflection==1.75.1", # required by srt/entrypoints/grpc_server.py
|
"grpcio-reflection==1.75.1", # required by srt/entrypoints/grpc_server.py
|
||||||
@@ -82,17 +83,17 @@ dependencies = [
|
|||||||
[project.optional-dependencies]
|
[project.optional-dependencies]
|
||||||
checkpoint-engine = ["checkpoint-engine==0.1.2"]
|
checkpoint-engine = ["checkpoint-engine==0.1.2"]
|
||||||
diffusion = [
|
diffusion = [
|
||||||
|
"PyYAML==6.0.1",
|
||||||
|
"cloudpickle",
|
||||||
"diffusers==0.35.2",
|
"diffusers==0.35.2",
|
||||||
"yunchang==0.6.3.post1",
|
|
||||||
"opencv-python==4.10.0.84",
|
|
||||||
"imageio==2.36.0",
|
"imageio==2.36.0",
|
||||||
"imageio-ffmpeg==0.5.1",
|
"imageio-ffmpeg==0.5.1",
|
||||||
"PyYAML==6.0.1",
|
|
||||||
"moviepy>=2.0.0",
|
"moviepy>=2.0.0",
|
||||||
"cloudpickle",
|
"opencv-python==4.10.0.84",
|
||||||
"remote-pdb",
|
"remote-pdb",
|
||||||
"st_attn ==0.0.7",
|
"st_attn ==0.0.7",
|
||||||
"vsa==0.0.4",
|
"vsa==0.0.4",
|
||||||
|
"yunchang==0.6.3.post1",
|
||||||
]
|
]
|
||||||
|
|
||||||
[tool.uv.extra-build-dependencies]
|
[tool.uv.extra-build-dependencies]
|
||||||
|
|||||||
@@ -19,12 +19,12 @@ import requests
|
|||||||
|
|
||||||
from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST
|
from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST
|
||||||
from sglang.srt.entrypoints.http_server import launch_server
|
from sglang.srt.entrypoints.http_server import launch_server
|
||||||
|
from sglang.srt.entrypoints.warmup import warmup
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.managers.io_struct import GenerateReqInput
|
from sglang.srt.managers.io_struct import GenerateReqInput
|
||||||
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.srt.utils import kill_process_tree
|
from sglang.srt.utils import kill_process_tree
|
||||||
from sglang.srt.warmup import warmup
|
|
||||||
|
|
||||||
multiprocessing.set_start_method("spawn", force=True)
|
multiprocessing.set_start_method("spawn", force=True)
|
||||||
|
|
||||||
|
|||||||
@@ -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(),
|
||||||
|
],
|
||||||
|
)
|
||||||
+6
-3
@@ -8,6 +8,11 @@ from typing import TYPE_CHECKING, Dict, List, Optional, Sequence
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.batch_overlap.operations import (
|
||||||
|
execute_operations,
|
||||||
|
execute_overlapped_operations,
|
||||||
|
)
|
||||||
|
from sglang.srt.batch_overlap.operations_strategy import OperationsStrategy
|
||||||
from sglang.srt.layers import deep_gemm_wrapper
|
from sglang.srt.layers import deep_gemm_wrapper
|
||||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||||
from sglang.srt.layers.communicator import (
|
from sglang.srt.layers.communicator import (
|
||||||
@@ -32,15 +37,13 @@ from sglang.srt.model_executor.forward_batch_info import (
|
|||||||
ForwardMode,
|
ForwardMode,
|
||||||
compute_position,
|
compute_position,
|
||||||
)
|
)
|
||||||
from sglang.srt.operations import execute_operations, execute_overlapped_operations
|
|
||||||
from sglang.srt.operations_strategy import OperationsStrategy
|
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.speculative.spec_info import SpecInput
|
from sglang.srt.speculative.spec_info import SpecInput
|
||||||
from sglang.srt.utils import BumpAllocator, empty_context, get_bool_env_var, is_hip
|
from sglang.srt.utils import BumpAllocator, empty_context, get_bool_env_var, is_hip
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.batch_overlap.single_batch_overlap import CombineOverlapArgs
|
||||||
from sglang.srt.layers.moe.token_dispatcher import DispatchOutput
|
from sglang.srt.layers.moe.token_dispatcher import DispatchOutput
|
||||||
from sglang.srt.single_batch_overlap import CombineOverlapArgs
|
|
||||||
from sglang.srt.speculative.eagle_info import EagleVerifyInput
|
from sglang.srt.speculative.eagle_info import EagleVerifyInput
|
||||||
|
|
||||||
_is_hip = is_hip()
|
_is_hip = is_hip()
|
||||||
@@ -72,6 +72,7 @@ from sglang.srt.entrypoints.openai.serving_tokenize import (
|
|||||||
OpenAIServingDetokenize,
|
OpenAIServingDetokenize,
|
||||||
OpenAIServingTokenize,
|
OpenAIServingTokenize,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.entrypoints.warmup import execute_warmups
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.function_call.function_call_parser import FunctionCallParser
|
from sglang.srt.function_call.function_call_parser import FunctionCallParser
|
||||||
from sglang.srt.managers.io_struct import (
|
from sglang.srt.managers.io_struct import (
|
||||||
@@ -126,7 +127,6 @@ from sglang.srt.utils import (
|
|||||||
kill_process_tree,
|
kill_process_tree,
|
||||||
set_uvicorn_logging_configs,
|
set_uvicorn_logging_configs,
|
||||||
)
|
)
|
||||||
from sglang.srt.warmup import execute_warmups
|
|
||||||
from sglang.utils import get_exception_traceback
|
from sglang.utils import get_exception_traceback
|
||||||
from sglang.version import __version__
|
from sglang.version import __version__
|
||||||
|
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ from typing import TYPE_CHECKING, Callable, List, Optional
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt import two_batch_overlap
|
from sglang.srt.batch_overlap import two_batch_overlap
|
||||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||||
from sglang.srt.speculative.spec_info import SpecInput
|
from sglang.srt.speculative.spec_info import SpecInput
|
||||||
|
|
||||||
|
|||||||
@@ -6,6 +6,8 @@ from typing import List, Optional, Tuple
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.batch_overlap.single_batch_overlap import DownGemmOverlapArgs
|
||||||
|
from sglang.srt.batch_overlap.two_batch_overlap import MaybeTboDeepEPDispatcher
|
||||||
from sglang.srt.distributed import (
|
from sglang.srt.distributed import (
|
||||||
get_moe_expert_parallel_rank,
|
get_moe_expert_parallel_rank,
|
||||||
get_moe_expert_parallel_world_size,
|
get_moe_expert_parallel_world_size,
|
||||||
@@ -46,8 +48,6 @@ from sglang.srt.layers.quantization.modelopt_quant import ModelOptNvFp4FusedMoEM
|
|||||||
from sglang.srt.layers.quantization.unquant import UnquantizedFusedMoEMethod
|
from sglang.srt.layers.quantization.unquant import UnquantizedFusedMoEMethod
|
||||||
from sglang.srt.model_loader.weight_utils import narrow_padded_param_and_loaded_weight
|
from sglang.srt.model_loader.weight_utils import narrow_padded_param_and_loaded_weight
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.single_batch_overlap import DownGemmOverlapArgs
|
|
||||||
from sglang.srt.two_batch_overlap import MaybeTboDeepEPDispatcher
|
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
get_bool_env_var,
|
get_bool_env_var,
|
||||||
|
|||||||
@@ -19,6 +19,7 @@ from typing import (
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.batch_overlap.single_batch_overlap import CombineOverlapArgs
|
||||||
from sglang.srt.layers.moe.token_dispatcher import (
|
from sglang.srt.layers.moe.token_dispatcher import (
|
||||||
DeepEPLLCombineInput,
|
DeepEPLLCombineInput,
|
||||||
DeepEPLLDispatchOutput,
|
DeepEPLLDispatchOutput,
|
||||||
@@ -28,7 +29,6 @@ if TYPE_CHECKING:
|
|||||||
StandardDispatchOutput,
|
StandardDispatchOutput,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.moe.topk import TopKOutput
|
from sglang.srt.layers.moe.topk import TopKOutput
|
||||||
from sglang.srt.single_batch_overlap import CombineOverlapArgs
|
|
||||||
|
|
||||||
|
|
||||||
# ------------------------------ Dispatcher Hook -------------------------------------
|
# ------------------------------ Dispatcher Hook -------------------------------------
|
||||||
|
|||||||
@@ -34,7 +34,7 @@ from sglang.srt.utils import (
|
|||||||
_is_npu = is_npu()
|
_is_npu = is_npu()
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.single_batch_overlap import CombineOverlapArgs
|
from sglang.srt.batch_overlap.single_batch_overlap import CombineOverlapArgs
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from deep_ep import Buffer, Config
|
from deep_ep import Buffer, Config
|
||||||
|
|||||||
@@ -19,7 +19,7 @@ from sglang.srt.layers.moe.utils import DeepEPMode
|
|||||||
from sglang.srt.utils import get_int_env_var
|
from sglang.srt.utils import get_int_env_var
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.single_batch_overlap import CombineOverlapArgs
|
from sglang.srt.batch_overlap.single_batch_overlap import CombineOverlapArgs
|
||||||
|
|
||||||
from enum import Enum, auto
|
from enum import Enum, auto
|
||||||
|
|
||||||
|
|||||||
@@ -30,6 +30,11 @@ from typing import (
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
try:
|
||||||
|
from triton_kernels.routing import GatherIndx, RoutingData, ScatterIndx, routing
|
||||||
|
except ImportError:
|
||||||
|
pass
|
||||||
|
|
||||||
from sglang.srt.custom_op import CustomOp
|
from sglang.srt.custom_op import CustomOp
|
||||||
from sglang.srt.distributed import get_tp_group
|
from sglang.srt.distributed import get_tp_group
|
||||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||||
@@ -57,13 +62,8 @@ from sglang.srt.utils.patch_torch import register_fake_if_exists
|
|||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.layers.quantization import QuantizationConfig
|
from sglang.srt.layers.quantization import QuantizationConfig
|
||||||
|
|
||||||
try:
|
|
||||||
from triton_kernels.routing import GatherIndx, RoutingData, ScatterIndx, routing
|
|
||||||
except ImportError:
|
|
||||||
pass
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
_is_hip = is_hip()
|
_is_hip = is_hip()
|
||||||
_is_cpu = is_cpu()
|
_is_cpu = is_cpu()
|
||||||
@@ -80,7 +80,12 @@ if _is_cuda:
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
if _is_cuda or _is_hip:
|
if _is_cuda or _is_hip:
|
||||||
from sgl_kernel import topk_sigmoid, topk_softmax
|
from sgl_kernel import topk_softmax
|
||||||
|
|
||||||
|
try:
|
||||||
|
from sgl_kernel import topk_sigmoid
|
||||||
|
except ImportError:
|
||||||
|
pass
|
||||||
if _use_aiter:
|
if _use_aiter:
|
||||||
try:
|
try:
|
||||||
from aiter import biased_grouped_topk as aiter_biased_grouped_topk
|
from aiter import biased_grouped_topk as aiter_biased_grouped_topk
|
||||||
|
|||||||
@@ -88,13 +88,11 @@ _is_cuda = is_cuda()
|
|||||||
_is_npu = is_npu()
|
_is_npu = is_npu()
|
||||||
_is_cpu_amx_available = cpu_has_amx_support()
|
_is_cpu_amx_available = cpu_has_amx_support()
|
||||||
_is_cpu = is_cpu()
|
_is_cpu = is_cpu()
|
||||||
|
|
||||||
_is_fp8_fnuz = is_fp8_fnuz()
|
_is_fp8_fnuz = is_fp8_fnuz()
|
||||||
|
_use_hip_int4 = get_bool_env_var("SGLANG_INT4_WEIGHT") and _is_hip
|
||||||
_use_hip_int4 = get_bool_env_var("SGLANG_INT4_WEIGHT")
|
|
||||||
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
||||||
|
|
||||||
if _is_hip and (_use_aiter or _use_hip_int4):
|
if _use_aiter or _use_hip_int4:
|
||||||
from aiter import ActivationType, QuantType
|
from aiter import ActivationType, QuantType
|
||||||
from aiter.fused_moe import fused_moe
|
from aiter.fused_moe import fused_moe
|
||||||
from aiter.ops.shuffle import shuffle_weight
|
from aiter.ops.shuffle import shuffle_weight
|
||||||
|
|||||||
@@ -45,10 +45,6 @@ _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
|||||||
if _is_cuda:
|
if _is_cuda:
|
||||||
from sgl_kernel import sgl_per_tensor_quant_fp8, sgl_per_token_quant_fp8
|
from sgl_kernel import sgl_per_tensor_quant_fp8, sgl_per_token_quant_fp8
|
||||||
|
|
||||||
@torch.library.register_fake("sgl_kernel::sgl_per_tensor_quant_fp8")
|
|
||||||
def _sgl_per_tensor_quant_fp8(input, output_q, output_s, is_static):
|
|
||||||
return
|
|
||||||
|
|
||||||
# Temporary
|
# Temporary
|
||||||
try:
|
try:
|
||||||
from sgl_kernel import sgl_per_token_group_quant_8bit
|
from sgl_kernel import sgl_per_token_group_quant_8bit
|
||||||
@@ -1861,3 +1857,7 @@ if _is_cuda:
|
|||||||
@torch.library.register_fake("sgl_kernel::sgl_per_token_quant_fp8")
|
@torch.library.register_fake("sgl_kernel::sgl_per_token_quant_fp8")
|
||||||
def _(input, output_q, output_s):
|
def _(input, output_q, output_s):
|
||||||
return
|
return
|
||||||
|
|
||||||
|
@torch.library.register_fake("sgl_kernel::sgl_per_tensor_quant_fp8")
|
||||||
|
def _sgl_per_tensor_quant_fp8(input, output_q, output_s, is_static):
|
||||||
|
return
|
||||||
|
|||||||
@@ -53,12 +53,12 @@ from sglang.srt.utils.common import (
|
|||||||
from sglang.srt.utils.patch_torch import register_fake_if_exists
|
from sglang.srt.utils.patch_torch import register_fake_if_exists
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.batch_overlap.single_batch_overlap import DownGemmOverlapArgs
|
||||||
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
|
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
|
||||||
from sglang.srt.layers.moe.token_dispatcher import (
|
from sglang.srt.layers.moe.token_dispatcher import (
|
||||||
CombineInput,
|
CombineInput,
|
||||||
StandardDispatchOutput,
|
StandardDispatchOutput,
|
||||||
)
|
)
|
||||||
from sglang.srt.single_batch_overlap import DownGemmOverlapArgs
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
if is_sm120_supported():
|
if is_sm120_supported():
|
||||||
|
|||||||
@@ -21,7 +21,6 @@ from sglang.srt.managers.schedule_batch import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.multimodal_cache import MultiModalStaticCache
|
from sglang.srt.mem_cache.multimodal_cache import MultiModalStaticCache
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.model_executor.piecewise_cuda_graph_runner import use_original_ca_comm
|
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.utils import flatten_nested_list, is_npu, print_warning_once
|
from sglang.srt.utils import flatten_nested_list, is_npu, print_warning_once
|
||||||
from sglang.utils import logger
|
from sglang.utils import logger
|
||||||
@@ -660,6 +659,10 @@ def general_mm_embed_routine(
|
|||||||
Returns:
|
Returns:
|
||||||
Hidden states from language model forward pass
|
Hidden states from language model forward pass
|
||||||
"""
|
"""
|
||||||
|
# Lazy import to allow some monkey patch of piecewise_cuda_graph_runner
|
||||||
|
from sglang.srt.model_executor.piecewise_cuda_graph_runner import (
|
||||||
|
use_original_ca_comm,
|
||||||
|
)
|
||||||
|
|
||||||
tp_group = get_tp_group()
|
tp_group = get_tp_group()
|
||||||
|
|
||||||
|
|||||||
@@ -5,8 +5,8 @@ from typing import TYPE_CHECKING, Callable
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.batch_overlap.two_batch_overlap import TboDPAttentionPreparer
|
||||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||||
from sglang.srt.two_batch_overlap import TboDPAttentionPreparer
|
|
||||||
from sglang.srt.utils.common import require_mlp_tp_gather
|
from sglang.srt.utils.common import require_mlp_tp_gather
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
|
|||||||
@@ -28,6 +28,7 @@ import torch
|
|||||||
import tqdm
|
import tqdm
|
||||||
from torch.profiler import ProfilerActivity, profile
|
from torch.profiler import ProfilerActivity, profile
|
||||||
|
|
||||||
|
from sglang.srt.batch_overlap.two_batch_overlap import TboCudaGraphRunnerPlugin
|
||||||
from sglang.srt.constants import GPU_MEMORY_TYPE_CUDA_GRAPH
|
from sglang.srt.constants import GPU_MEMORY_TYPE_CUDA_GRAPH
|
||||||
from sglang.srt.custom_op import CustomOp
|
from sglang.srt.custom_op import CustomOp
|
||||||
from sglang.srt.distributed import get_tensor_model_parallel_rank
|
from sglang.srt.distributed import get_tensor_model_parallel_rank
|
||||||
@@ -60,7 +61,6 @@ from sglang.srt.model_executor.forward_batch_info import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.model_executor.input_buffers import GraphInputBuffers
|
from sglang.srt.model_executor.input_buffers import GraphInputBuffers
|
||||||
from sglang.srt.multiplex.pdmux_context import get_current_stream_idx, get_stream_groups
|
from sglang.srt.multiplex.pdmux_context import get_current_stream_idx, get_stream_groups
|
||||||
from sglang.srt.two_batch_overlap import TboCudaGraphRunnerPlugin
|
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
empty_context,
|
empty_context,
|
||||||
get_available_gpu_memory,
|
get_available_gpu_memory,
|
||||||
|
|||||||
@@ -719,7 +719,7 @@ class ForwardBatch:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def prepare_mlp_sync_batch(self, model_runner: ModelRunner):
|
def prepare_mlp_sync_batch(self, model_runner: ModelRunner):
|
||||||
from sglang.srt.two_batch_overlap import TboForwardBatchPreparer
|
from sglang.srt.batch_overlap.two_batch_overlap import TboForwardBatchPreparer
|
||||||
|
|
||||||
assert self.global_num_tokens_cpu is not None
|
assert self.global_num_tokens_cpu is not None
|
||||||
assert self.global_num_tokens_for_logprob_cpu is not None
|
assert self.global_num_tokens_for_logprob_cpu is not None
|
||||||
|
|||||||
@@ -24,6 +24,7 @@ from typing import TYPE_CHECKING, Union
|
|||||||
import torch
|
import torch
|
||||||
import tqdm
|
import tqdm
|
||||||
|
|
||||||
|
from sglang.srt.batch_overlap.two_batch_overlap import TboCudaGraphRunnerPlugin
|
||||||
from sglang.srt.compilation.compilation_config import CompilationConfig
|
from sglang.srt.compilation.compilation_config import CompilationConfig
|
||||||
from sglang.srt.compilation.compile import install_torch_compiled, set_compiled
|
from sglang.srt.compilation.compile import install_torch_compiled, set_compiled
|
||||||
from sglang.srt.compilation.piecewise_context_manager import (
|
from sglang.srt.compilation.piecewise_context_manager import (
|
||||||
@@ -51,7 +52,6 @@ from sglang.srt.model_executor.forward_batch_info import (
|
|||||||
ForwardMode,
|
ForwardMode,
|
||||||
PPProxyTensors,
|
PPProxyTensors,
|
||||||
)
|
)
|
||||||
from sglang.srt.two_batch_overlap import TboCudaGraphRunnerPlugin
|
|
||||||
from sglang.srt.utils import get_available_gpu_memory, log_info_on_rank0
|
from sglang.srt.utils import get_available_gpu_memory, log_info_on_rank0
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|||||||
@@ -29,6 +29,8 @@ import tqdm
|
|||||||
from torch import nn
|
from torch import nn
|
||||||
from transformers import PretrainedConfig
|
from transformers import PretrainedConfig
|
||||||
|
|
||||||
|
from sglang.srt.batch_overlap.single_batch_overlap import SboFlags, compute_overlap_args
|
||||||
|
from sglang.srt.batch_overlap.two_batch_overlap import model_forward_maybe_tbo
|
||||||
from sglang.srt.compilation.piecewise_context_manager import is_in_piecewise_cuda_graph
|
from sglang.srt.compilation.piecewise_context_manager import is_in_piecewise_cuda_graph
|
||||||
from sglang.srt.configs.model_config import (
|
from sglang.srt.configs.model_config import (
|
||||||
get_nsa_index_head_dim,
|
get_nsa_index_head_dim,
|
||||||
@@ -134,9 +136,7 @@ from sglang.srt.model_loader.utils import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.single_batch_overlap import SboFlags, compute_overlap_args
|
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
from sglang.srt.two_batch_overlap import model_forward_maybe_tbo
|
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
BumpAllocator,
|
BumpAllocator,
|
||||||
LazyValue,
|
LazyValue,
|
||||||
|
|||||||
@@ -22,6 +22,7 @@ import torch.nn.functional as F
|
|||||||
from torch import nn
|
from torch import nn
|
||||||
from transformers import PretrainedConfig
|
from transformers import PretrainedConfig
|
||||||
|
|
||||||
|
from sglang.srt.batch_overlap.two_batch_overlap import model_forward_maybe_tbo
|
||||||
from sglang.srt.distributed import (
|
from sglang.srt.distributed import (
|
||||||
get_moe_expert_parallel_world_size,
|
get_moe_expert_parallel_world_size,
|
||||||
get_pp_group,
|
get_pp_group,
|
||||||
@@ -75,7 +76,6 @@ from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode
|
|||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
||||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.two_batch_overlap import model_forward_maybe_tbo
|
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
add_prefix,
|
add_prefix,
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
|
|||||||
@@ -22,6 +22,7 @@ import torch
|
|||||||
from torch import nn
|
from torch import nn
|
||||||
from transformers import PretrainedConfig
|
from transformers import PretrainedConfig
|
||||||
|
|
||||||
|
from sglang.srt.batch_overlap.two_batch_overlap import model_forward_maybe_tbo
|
||||||
from sglang.srt.distributed import (
|
from sglang.srt.distributed import (
|
||||||
get_moe_expert_parallel_world_size,
|
get_moe_expert_parallel_world_size,
|
||||||
get_pp_group,
|
get_pp_group,
|
||||||
@@ -61,7 +62,6 @@ from sglang.srt.model_loader.weight_utils import (
|
|||||||
maybe_remap_kv_scale_name,
|
maybe_remap_kv_scale_name,
|
||||||
)
|
)
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.two_batch_overlap import model_forward_maybe_tbo
|
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
BumpAllocator,
|
BumpAllocator,
|
||||||
add_prefix,
|
add_prefix,
|
||||||
|
|||||||
@@ -25,6 +25,7 @@ import torch.nn.functional as F
|
|||||||
from torch import nn
|
from torch import nn
|
||||||
from transformers import PretrainedConfig
|
from transformers import PretrainedConfig
|
||||||
|
|
||||||
|
from sglang.srt.batch_overlap.two_batch_overlap import model_forward_maybe_tbo
|
||||||
from sglang.srt.distributed import (
|
from sglang.srt.distributed import (
|
||||||
get_moe_expert_parallel_world_size,
|
get_moe_expert_parallel_world_size,
|
||||||
get_pp_group,
|
get_pp_group,
|
||||||
@@ -70,7 +71,6 @@ from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode
|
|||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
||||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.two_batch_overlap import model_forward_maybe_tbo
|
|
||||||
from sglang.srt.utils import add_prefix, is_cuda, make_layers
|
from sglang.srt.utils import add_prefix, is_cuda, make_layers
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|||||||
@@ -3,10 +3,10 @@ from typing import List, Optional
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt import operations
|
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.layers.moe.token_dispatcher import DeepEPConfig
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||||
from sglang.srt.operations import Operation
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
|
|||||||
@@ -137,18 +137,22 @@ builtins.FP8_E4M3_MAX = FP8_E4M3_MAX
|
|||||||
builtins.FP8_E4M3_MIN = FP8_E4M3_MIN
|
builtins.FP8_E4M3_MIN = FP8_E4M3_MIN
|
||||||
|
|
||||||
|
|
||||||
|
@lru_cache(maxsize=1)
|
||||||
def is_cuda():
|
def is_cuda():
|
||||||
return torch.cuda.is_available() and torch.version.cuda
|
return torch.cuda.is_available() and torch.version.cuda
|
||||||
|
|
||||||
|
|
||||||
|
@lru_cache(maxsize=1)
|
||||||
def is_cuda_alike():
|
def is_cuda_alike():
|
||||||
return is_cuda() or is_hip()
|
return is_cuda() or is_hip()
|
||||||
|
|
||||||
|
|
||||||
|
@lru_cache(maxsize=1)
|
||||||
def is_hpu() -> bool:
|
def is_hpu() -> bool:
|
||||||
return hasattr(torch, "hpu") and torch.hpu.is_available()
|
return hasattr(torch, "hpu") and torch.hpu.is_available()
|
||||||
|
|
||||||
|
|
||||||
|
@lru_cache(maxsize=1)
|
||||||
def is_xpu() -> bool:
|
def is_xpu() -> bool:
|
||||||
return hasattr(torch, "xpu") and torch.xpu.is_available()
|
return hasattr(torch, "xpu") and torch.xpu.is_available()
|
||||||
|
|
||||||
@@ -158,6 +162,7 @@ def is_npu() -> bool:
|
|||||||
return hasattr(torch, "npu") and torch.npu.is_available()
|
return hasattr(torch, "npu") and torch.npu.is_available()
|
||||||
|
|
||||||
|
|
||||||
|
@lru_cache(maxsize=1)
|
||||||
def is_host_cpu_x86() -> bool:
|
def is_host_cpu_x86() -> bool:
|
||||||
machine = platform.machine().lower()
|
machine = platform.machine().lower()
|
||||||
return (
|
return (
|
||||||
@@ -167,6 +172,7 @@ def is_host_cpu_x86() -> bool:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@lru_cache(maxsize=1)
|
||||||
def is_cpu() -> bool:
|
def is_cpu() -> bool:
|
||||||
return os.getenv("SGLANG_USE_CPU_ENGINE", "0") == "1" and is_host_cpu_x86()
|
return os.getenv("SGLANG_USE_CPU_ENGINE", "0") == "1" and is_host_cpu_x86()
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user