[Build] Add srt_empty extra group for device-agnostic install (#31300)
Co-authored-by: yixiaodapeng <yixiaodapeng@example.com> Co-authored-by: Alex Nails <alex.nails@radixark.ai>
This commit is contained in:
co-authored by
yixiaodapeng
Alex Nails
parent
de34dd11e9
commit
5c4f72f92a
@@ -17,14 +17,16 @@ classifiers = [
|
|||||||
dependencies = ["aiohttp", "IPython", "numpy", "requests", "setproctitle", "tqdm"]
|
dependencies = ["aiohttp", "IPython", "numpy", "requests", "setproctitle", "tqdm"]
|
||||||
|
|
||||||
[project.optional-dependencies]
|
[project.optional-dependencies]
|
||||||
runtime_common = [
|
# runtime_base: torch-free subset of runtime_common.
|
||||||
|
# Used by srt_empty for device-agnostic install on non-NVIDIA platforms.
|
||||||
|
# NOTE: Do NOT add packages that depend on torch/triton here — they belong in runtime_common.
|
||||||
|
runtime_base = [
|
||||||
"aiohttp",
|
"aiohttp",
|
||||||
"anthropic>=0.20.0",
|
"anthropic>=0.20.0",
|
||||||
"apache-tvm-ffi",
|
"apache-tvm-ffi",
|
||||||
"av",
|
"av",
|
||||||
"blobfile==3.0.0",
|
"blobfile==3.0.0",
|
||||||
"build",
|
"build",
|
||||||
"compressed-tensors",
|
|
||||||
"datasets",
|
"datasets",
|
||||||
"easydict",
|
"easydict",
|
||||||
"einops",
|
"einops",
|
||||||
@@ -42,7 +44,6 @@ runtime_common = [
|
|||||||
"openai==2.6.1",
|
"openai==2.6.1",
|
||||||
"openai-harmony==0.0.4",
|
"openai-harmony==0.0.4",
|
||||||
"orjson",
|
"orjson",
|
||||||
"outlines==0.1.11",
|
|
||||||
"packaging",
|
"packaging",
|
||||||
"partial_json_parser",
|
"partial_json_parser",
|
||||||
"pillow",
|
"pillow",
|
||||||
@@ -60,16 +61,29 @@ runtime_common = [
|
|||||||
"smg-grpc-servicer>=0.5.0",
|
"smg-grpc-servicer>=0.5.0",
|
||||||
"soundfile==0.13.1",
|
"soundfile==0.13.1",
|
||||||
"tiktoken",
|
"tiktoken",
|
||||||
"timm==1.0.16",
|
|
||||||
"torchao==0.9.0",
|
|
||||||
"tqdm",
|
"tqdm",
|
||||||
"transformers==5.12.1",
|
"transformers==5.12.1",
|
||||||
"uvicorn",
|
"uvicorn",
|
||||||
"uvloop",
|
"uvloop",
|
||||||
"xxhash",
|
"xxhash",
|
||||||
|
]
|
||||||
|
|
||||||
|
# runtime_common: backward-compatible — same install result as before the split.
|
||||||
|
runtime_common = [
|
||||||
|
"sglang[runtime_base]",
|
||||||
|
"compressed-tensors",
|
||||||
|
"outlines==0.1.11",
|
||||||
|
"timm==1.0.16",
|
||||||
|
"torchao==0.9.0",
|
||||||
"xgrammar==0.2.1",
|
"xgrammar==0.2.1",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
# srt_empty: device-agnostic install — pure Python packages only, no torch dependency chain.
|
||||||
|
# Enables OOT plugins (e.g. sglang-plugin-FL) to install sglang without conflicting with
|
||||||
|
# vendor-specific PyTorch builds (torch_npu, torch_musa, etc.).
|
||||||
|
# Usage: cp pyproject_other.toml pyproject.toml && pip install -e ".[srt_empty]"
|
||||||
|
srt_empty = ["sglang[runtime_base]"]
|
||||||
|
|
||||||
diffusion_common = [
|
diffusion_common = [
|
||||||
"addict",
|
"addict",
|
||||||
"cloudpickle",
|
"cloudpickle",
|
||||||
|
|||||||
@@ -25,7 +25,7 @@ from collections import deque
|
|||||||
from contextlib import contextmanager, nullcontext
|
from contextlib import contextmanager, nullcontext
|
||||||
from functools import partial
|
from functools import partial
|
||||||
from http import HTTPStatus
|
from http import HTTPStatus
|
||||||
from typing import Any, Deque, Dict, List, Optional, Tuple, Union
|
from typing import TYPE_CHECKING, Any, Deque, Dict, List, Optional, Tuple, Union
|
||||||
|
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
get_device,
|
get_device,
|
||||||
@@ -50,12 +50,17 @@ import psutil # isort: skip
|
|||||||
import setproctitle
|
import setproctitle
|
||||||
import torch
|
import torch
|
||||||
import torch.distributed
|
import torch.distributed
|
||||||
from torch.cuda import Stream as CudaStream
|
|
||||||
from torch.distributed import barrier
|
from torch.distributed import barrier
|
||||||
|
|
||||||
from sglang.kernels.ops.mamba.triton_ops import (
|
if TYPE_CHECKING:
|
||||||
initialize_mamba_selective_state_update_backend,
|
from torch.cuda import Stream as CudaStream
|
||||||
)
|
|
||||||
|
try:
|
||||||
|
from sglang.kernels.ops.mamba.triton_ops import (
|
||||||
|
initialize_mamba_selective_state_update_backend,
|
||||||
|
)
|
||||||
|
except ImportError:
|
||||||
|
initialize_mamba_selective_state_update_backend = None
|
||||||
from sglang.srt.configs.model_config import (
|
from sglang.srt.configs.model_config import (
|
||||||
ModelConfig,
|
ModelConfig,
|
||||||
ModelImpl,
|
ModelImpl,
|
||||||
@@ -851,7 +856,8 @@ class Scheduler(
|
|||||||
)
|
)
|
||||||
|
|
||||||
def init_mamba_backend(self) -> None:
|
def init_mamba_backend(self) -> None:
|
||||||
initialize_mamba_selective_state_update_backend(self.server_args)
|
if initialize_mamba_selective_state_update_backend is not None:
|
||||||
|
initialize_mamba_selective_state_update_backend(self.server_args)
|
||||||
|
|
||||||
def init_moe_gemm_config(self):
|
def init_moe_gemm_config(self):
|
||||||
# For the MM models, check the text_config for MoE settings
|
# For the MM models, check the text_config for MoE settings
|
||||||
|
|||||||
@@ -31,7 +31,6 @@ import uuid
|
|||||||
from functools import cached_property
|
from functools import cached_property
|
||||||
from typing import Any, Callable, Dict, List, Literal, Optional, Union
|
from typing import Any, Callable, Dict, List, Literal, Optional, Union
|
||||||
|
|
||||||
from sglang.kernels.ops.attention.fla.chunk_delta_h import CHUNK_SIZE as FLA_CHUNK_SIZE
|
|
||||||
from sglang.kernels.ops.kv_canary.consts import RealKvHashMode
|
from sglang.kernels.ops.kv_canary.consts import RealKvHashMode
|
||||||
from sglang.srt.arg_groups.arg_utils import NS, A, Arg, add_cli_args_from_dataclass
|
from sglang.srt.arg_groups.arg_utils import NS, A, Arg, add_cli_args_from_dataclass
|
||||||
from sglang.srt.arg_groups.argparse_actions import (
|
from sglang.srt.arg_groups.argparse_actions import (
|
||||||
@@ -8675,6 +8674,14 @@ class ServerArgs:
|
|||||||
# It is used to determine the caching point in a sequence during prefill.
|
# It is used to determine the caching point in a sequence during prefill.
|
||||||
if not hasattr(self, "_mamba_cache_chunk_size"):
|
if not hasattr(self, "_mamba_cache_chunk_size"):
|
||||||
|
|
||||||
|
try:
|
||||||
|
from sglang.kernels.ops.attention.fla.chunk_delta_h import (
|
||||||
|
CHUNK_SIZE as FLA_CHUNK_SIZE,
|
||||||
|
)
|
||||||
|
except ImportError:
|
||||||
|
# Must match sglang.kernels.ops.attention.fla.chunk_delta_h.CHUNK_SIZE
|
||||||
|
FLA_CHUNK_SIZE = 64
|
||||||
|
|
||||||
hf_config = self.get_model_config().hf_config
|
hf_config = self.get_model_config().hf_config
|
||||||
chunk_size = getattr(hf_config, "mamba_chunk_size", FLA_CHUNK_SIZE)
|
chunk_size = getattr(hf_config, "mamba_chunk_size", FLA_CHUNK_SIZE)
|
||||||
page_size = resolved_view(self).page_size
|
page_size = resolved_view(self).page_size
|
||||||
|
|||||||
@@ -0,0 +1,91 @@
|
|||||||
|
# Copyright 2024 SGLang Team
|
||||||
|
# Licensed under the Apache License, Version 2.0
|
||||||
|
"""Test that runtime_base in pyproject_other.toml remains torch-free.
|
||||||
|
|
||||||
|
This prevents accidental introduction of packages that transitively pull
|
||||||
|
torch/triton into the srt_empty install target.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
# Packages known to transitively depend on torch or triton.
|
||||||
|
# If a new package is added to runtime_base and it pulls torch,
|
||||||
|
# add it here and move it to runtime_common instead.
|
||||||
|
TORCH_PULLING_PACKAGES = frozenset(
|
||||||
|
{
|
||||||
|
"torch",
|
||||||
|
"torchao",
|
||||||
|
"timm",
|
||||||
|
"xgrammar",
|
||||||
|
"compressed-tensors",
|
||||||
|
"outlines",
|
||||||
|
"flashinfer",
|
||||||
|
"sgl-kernel",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _parse_runtime_base() -> set:
|
||||||
|
"""Parse runtime_base package names from pyproject_other.toml."""
|
||||||
|
# Try tomllib (3.11+) or tomli
|
||||||
|
try:
|
||||||
|
import tomllib
|
||||||
|
except ModuleNotFoundError:
|
||||||
|
import tomli as tomllib # type: ignore[no-redef]
|
||||||
|
|
||||||
|
toml_path = Path(__file__).resolve().parents[3] / "python" / "pyproject_other.toml"
|
||||||
|
if not toml_path.exists():
|
||||||
|
pytest.skip(f"pyproject_other.toml not found at {toml_path}")
|
||||||
|
|
||||||
|
with open(toml_path, "rb") as f:
|
||||||
|
data = tomllib.load(f)
|
||||||
|
|
||||||
|
runtime_base = data["project"]["optional-dependencies"]["runtime_base"]
|
||||||
|
|
||||||
|
# Extract bare package names (strip version specifiers and extras)
|
||||||
|
pkg_names = set()
|
||||||
|
for dep in runtime_base:
|
||||||
|
# "package[extra]>=1.0,<2.0" -> "package"
|
||||||
|
name = (
|
||||||
|
dep.split("[")[0]
|
||||||
|
.split(">")[0]
|
||||||
|
.split("<")[0]
|
||||||
|
.split("=")[0]
|
||||||
|
.split("!")[0]
|
||||||
|
.split(";")[0]
|
||||||
|
.strip()
|
||||||
|
)
|
||||||
|
pkg_names.add(name.lower())
|
||||||
|
|
||||||
|
return pkg_names
|
||||||
|
|
||||||
|
|
||||||
|
def test_runtime_base_no_torch_deps():
|
||||||
|
"""runtime_base must not contain packages that pull in torch."""
|
||||||
|
pkg_names = _parse_runtime_base()
|
||||||
|
violations = pkg_names & TORCH_PULLING_PACKAGES
|
||||||
|
assert not violations, (
|
||||||
|
f"runtime_base contains torch-pulling packages: {sorted(violations)}. "
|
||||||
|
f"Move them to runtime_common to keep srt_empty torch-free."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_runtime_base_not_empty():
|
||||||
|
"""Sanity check: runtime_base should have a reasonable number of packages."""
|
||||||
|
pkg_names = _parse_runtime_base()
|
||||||
|
assert len(pkg_names) >= 20, (
|
||||||
|
f"runtime_base only has {len(pkg_names)} packages, expected >= 20. "
|
||||||
|
f"Did the toml structure change?"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
import sys
|
||||||
|
|
||||||
|
sys.exit(pytest.main([__file__, "-v"]))
|
||||||
Reference in New Issue
Block a user