From fd655128f59d9714910beb808ad7651967af4c81 Mon Sep 17 00:00:00 2001 From: zhangtaoshan Date: Fri, 11 Sep 2026 09:17:40 +0000 Subject: [PATCH] feat: add hardware platform and plugin support --- lightllm/platform/__init__.py | 34 +++++ lightllm/platform/backends/__init__.py | 1 + .../platform/backends/cuda_like/__init__.py | 14 ++ lightllm/platform/backends/cuda_like/graph.py | 23 ++++ .../platform/backends/cuda_like/runtime.py | 63 +++++++++ lightllm/platform/base/__init__.py | 0 lightllm/platform/base/backend.py | 24 ++++ lightllm/platform/base/graph.py | 28 ++++ lightllm/platform/base/registry.py | 37 +++++ lightllm/platform/base/runtime.py | 122 +++++++++++++++++ lightllm/platform/plugin/__init__.py | 12 ++ lightllm/platform/plugin/att.py | 63 +++++++++ lightllm/platform/plugin/common.py | 128 ++++++++++++++++++ lightllm/platform/plugin/ops.py | 43 ++++++ lightllm/server/api_cli.py | 14 ++ lightllm/server/core/objs/start_args_type.py | 3 + 16 files changed, 609 insertions(+) create mode 100644 lightllm/platform/__init__.py create mode 100644 lightllm/platform/backends/__init__.py create mode 100644 lightllm/platform/backends/cuda_like/__init__.py create mode 100644 lightllm/platform/backends/cuda_like/graph.py create mode 100644 lightllm/platform/backends/cuda_like/runtime.py create mode 100644 lightllm/platform/base/__init__.py create mode 100644 lightllm/platform/base/backend.py create mode 100644 lightllm/platform/base/graph.py create mode 100644 lightllm/platform/base/registry.py create mode 100644 lightllm/platform/base/runtime.py create mode 100644 lightllm/platform/plugin/__init__.py create mode 100644 lightllm/platform/plugin/att.py create mode 100644 lightllm/platform/plugin/common.py create mode 100644 lightllm/platform/plugin/ops.py diff --git a/lightllm/platform/__init__.py b/lightllm/platform/__init__.py new file mode 100644 index 0000000000..70fcd51c65 --- /dev/null +++ b/lightllm/platform/__init__.py @@ -0,0 +1,34 @@ +from __future__ import annotations + +from typing import Optional + +import lightllm.platform.backends # noqa: F401 +from lightllm.platform.base.backend import HardwareBackend +from lightllm.platform.base.registry import get_platform_spec +from lightllm.platform.plugin import configure_plugins +from lightllm.utils.envs_utils import get_env_start_args + +_backend: Optional[HardwareBackend] = None + + +def get_hardware_backend() -> HardwareBackend: + global _backend + + if _backend is not None: + return _backend + + configure_plugins() + + platform_name = get_env_start_args().hardware_platform + spec = get_platform_spec(platform_name) + + backend_cls = spec.backend_cls + _backend = backend_cls() + + if not _backend.runtime.is_available(): + raise RuntimeError(f"Hardware backend {backend_cls.__name__} is not available.") + + return _backend + + +__all__ = ["get_hardware_backend"] diff --git a/lightllm/platform/backends/__init__.py b/lightllm/platform/backends/__init__.py new file mode 100644 index 0000000000..25cb41dae8 --- /dev/null +++ b/lightllm/platform/backends/__init__.py @@ -0,0 +1 @@ +from lightllm.platform.backends import cuda_like # noqa: F401 diff --git a/lightllm/platform/backends/cuda_like/__init__.py b/lightllm/platform/backends/cuda_like/__init__.py new file mode 100644 index 0000000000..7ea2d3a81a --- /dev/null +++ b/lightllm/platform/backends/cuda_like/__init__.py @@ -0,0 +1,14 @@ +from lightllm.platform.base.backend import HardwareBackend +from lightllm.platform.base.registry import register_platform +from lightllm.platform.backends.cuda_like.runtime import CudaLikeRuntime +from lightllm.platform.backends.cuda_like.graph import CudaLikeGraph + + +class CudaLikeBackend(HardwareBackend): + def __init__(self) -> None: + super().__init__(CudaLikeRuntime(), CudaLikeGraph()) + + +@register_platform("cuda") +class CudaBackend(CudaLikeBackend): + pass diff --git a/lightllm/platform/backends/cuda_like/graph.py b/lightllm/platform/backends/cuda_like/graph.py new file mode 100644 index 0000000000..987013a81e --- /dev/null +++ b/lightllm/platform/backends/cuda_like/graph.py @@ -0,0 +1,23 @@ +from typing import Any, ContextManager, Optional + +import torch +from lightllm.platform.base.graph import HardwareBackendGraph + + +class CudaLikeGraph(HardwareBackendGraph): + def create_graph(self) -> Any: + return torch.cuda.CUDAGraph() + + def graph( + self, + graph_obj: Any, + pool: Optional[Any] = None, + stream: Optional[Any] = None, + ) -> ContextManager: + return torch.cuda.graph(graph_obj, pool=pool, stream=stream) + + def graph_pool_handle(self) -> Any: + return torch.cuda.graph_pool_handle() + + def is_capturing(self) -> bool: + return torch.cuda.is_current_stream_capturing() diff --git a/lightllm/platform/backends/cuda_like/runtime.py b/lightllm/platform/backends/cuda_like/runtime.py new file mode 100644 index 0000000000..88fa595cf5 --- /dev/null +++ b/lightllm/platform/backends/cuda_like/runtime.py @@ -0,0 +1,63 @@ +from typing import Any, ContextManager, Optional, Tuple + +import torch +from lightllm.platform.base.runtime import DeviceLike, HardwareBackendRuntime + + +class CudaLikeRuntime(HardwareBackendRuntime): + @property + def device_type(self) -> str: + return "cuda" + + @property + def dist_backend(self) -> str: + return "nccl" + + def mem_get_info(self, device: DeviceLike) -> Tuple[int, int]: + return torch.cuda.mem_get_info(self._parse(device)) + + def get_device_properties(self, device: DeviceLike) -> Any: + return torch.cuda.get_device_properties(self._parse(device)) + + def device_count(self) -> int: + return torch.cuda.device_count() + + def is_available(self) -> bool: + return torch.cuda.is_available() + + def current_device(self) -> int: + return torch.cuda.current_device() + + def get_device_name(self, device_id: Optional[int] = None) -> str: + if device_id is None: + device_id = self.current_device() + return torch.cuda.get_device_name(device_id) + + def set_device(self, device: DeviceLike) -> None: + torch.cuda.set_device(self._parse(device)) + + def create_stream(self, **kwargs) -> Any: + return torch.cuda.Stream(**kwargs) + + def stream(self, stream: Any) -> ContextManager[Any]: + return torch.cuda.stream(stream) + + def current_stream(self, device_id: Optional[int] = None) -> Any: + if device_id is None: + device_id = self.current_device() + return torch.cuda.current_stream(device_id) + + def create_event(self, **kwargs) -> torch.Event: + return torch.cuda.Event(**kwargs) + + def synchronize(self, device: Optional[DeviceLike] = None) -> None: + if device is None: + torch.cuda.synchronize() + return + torch.cuda.synchronize(self._parse(device)) + + def empty_cache(self) -> None: + torch.cuda.empty_cache() + + def manual_seed_all(self, seed: int) -> None: + torch.cuda.manual_seed_all(seed) diff --git a/lightllm/platform/base/__init__.py b/lightllm/platform/base/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/lightllm/platform/base/backend.py b/lightllm/platform/base/backend.py new file mode 100644 index 0000000000..3669cc0e8a --- /dev/null +++ b/lightllm/platform/base/backend.py @@ -0,0 +1,24 @@ +from abc import ABC + +from lightllm.platform.base.graph import HardwareBackendGraph +from lightllm.platform.base.runtime import HardwareBackendRuntime + + +class HardwareBackend(ABC): + platform_name: str + + def __init__(self, runtime: HardwareBackendRuntime, graph: HardwareBackendGraph) -> None: + self._runtime = runtime + self._graph = graph + + @property + def name(self) -> str: + return self.platform_name + + @property + def runtime(self) -> HardwareBackendRuntime: + return self._runtime + + @property + def graph(self) -> HardwareBackendGraph: + return self._graph diff --git a/lightllm/platform/base/graph.py b/lightllm/platform/base/graph.py new file mode 100644 index 0000000000..ea015cbd2a --- /dev/null +++ b/lightllm/platform/base/graph.py @@ -0,0 +1,28 @@ +from abc import ABC, abstractmethod +from typing import Any, ContextManager, Optional + + +class HardwareBackendGraph(ABC): + @abstractmethod + def create_graph(self) -> Any: + pass + + @abstractmethod + def graph( + self, + graph_obj: Any, + pool: Optional[Any] = None, + stream: Optional[Any] = None, + ) -> ContextManager: + pass + + def replay_graph(self, graph_obj: Any) -> Any: + graph_obj.replay() + + @abstractmethod + def graph_pool_handle(self) -> Any: + pass + + @abstractmethod + def is_capturing(self) -> bool: + pass diff --git a/lightllm/platform/base/registry.py b/lightllm/platform/base/registry.py new file mode 100644 index 0000000000..61f4bbef21 --- /dev/null +++ b/lightllm/platform/base/registry.py @@ -0,0 +1,37 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Callable, Type + +from lightllm.platform.base.backend import HardwareBackend + + +@dataclass(frozen=True) +class PlatformSpec: + name: str + backend_cls: Type[HardwareBackend] + + +PLATFORMS: dict[str, PlatformSpec] = {} + + +def register_platform(name: str) -> Callable[[Type[HardwareBackend]], Type[HardwareBackend]]: + def decorator(backend_cls: Type[HardwareBackend]) -> Type[HardwareBackend]: + if name in PLATFORMS: + raise ValueError(f"Platform {name!r} is already registered.") + + backend_cls.platform_name = name + PLATFORMS[name] = PlatformSpec( + name=name, + backend_cls=backend_cls, + ) + return backend_cls + + return decorator + + +def get_platform_spec(name: str) -> PlatformSpec: + spec = PLATFORMS.get(name) + if spec is None: + raise RuntimeError(f"Platform {name!r} is not registered, registered: {sorted(PLATFORMS.keys())}.") + return spec diff --git a/lightllm/platform/base/runtime.py b/lightllm/platform/base/runtime.py new file mode 100644 index 0000000000..d1f834ca3b --- /dev/null +++ b/lightllm/platform/base/runtime.py @@ -0,0 +1,122 @@ +from abc import ABC, abstractmethod +from typing import Any, ContextManager, Optional, Tuple, Union + +import torch +import torch.distributed as dist + +DeviceLike = Union[int, str, torch.device] + + +class HardwareBackendRuntime(ABC): + @property + @abstractmethod + def device_type(self) -> str: + pass + + @property + @abstractmethod + def dist_backend(self) -> str: + pass + + def dist_init_kwargs(self, target_device: torch.device) -> dict[str, Any]: + return {} + + def init_process_group( + self, + *, + host: str, + port: int, + rank: int, + world_size: int, + device_id: int, + ) -> None: + target_device = self.target_device(device_id) + self.set_device(target_device) + + kwargs: dict[str, Any] = { + "backend": self.dist_backend, + "init_method": f"tcp://{host}:{port}", + "rank": rank, + "world_size": world_size, + } + kwargs.update(self.dist_init_kwargs(target_device)) + dist.init_process_group(**kwargs) + + @abstractmethod + def mem_get_info(self, device: DeviceLike) -> Tuple[int, int]: + pass + + @abstractmethod + def get_device_properties(self, device: DeviceLike) -> Any: + pass + + def target_device(self, device_id: Optional[int] = None) -> torch.device: + if device_id is None: + device_id = self.current_device() + return torch.device(self.device_type, device_id) + + @abstractmethod + def device_count(self) -> int: + pass + + @abstractmethod + def is_available(self) -> bool: + pass + + @abstractmethod + def current_device(self) -> int: + pass + + @abstractmethod + def get_device_name(self, device_id: Optional[int] = None) -> str: + pass + + def _parse(self, device: DeviceLike) -> torch.device: + if isinstance(device, torch.device): + _device = device + elif isinstance(device, int): + _device = torch.device(self.device_type, device) + elif isinstance(device, str): + _device = torch.device(device) + else: + raise TypeError(f"Invalid device: {device}") + + if _device.type != self.device_type: + raise ValueError(f"Expected device type {self.device_type!r}, got {_device.type!r} ({_device})") + + if _device.index is None: + _device = torch.device(self.device_type, self.current_device()) + + return _device + + @abstractmethod + def set_device(self, device: DeviceLike) -> None: + pass + + @abstractmethod + def create_stream(self, **kwargs) -> Any: + pass + + @abstractmethod + def stream(self, stream: Any) -> ContextManager[Any]: + pass + + @abstractmethod + def current_stream(self, device_id: Optional[int] = None) -> Any: + pass + + @abstractmethod + def create_event(self, **kwargs) -> torch.Event: + pass + + @abstractmethod + def synchronize(self, device: Optional[DeviceLike] = None) -> None: + pass + + @abstractmethod + def empty_cache(self) -> None: + pass + + @abstractmethod + def manual_seed_all(self, seed: int) -> None: + pass diff --git a/lightllm/platform/plugin/__init__.py b/lightllm/platform/plugin/__init__.py new file mode 100644 index 0000000000..21c2ab8e6d --- /dev/null +++ b/lightllm/platform/plugin/__init__.py @@ -0,0 +1,12 @@ +from lightllm.platform.plugin.common import Plugin + + +OPS = Plugin(name="ops", entry_point_group="lightllm.ops_plugin") +ATT = Plugin(name="att", entry_point_group="lightllm.att_plugin") + +_PLUGINS = (OPS, ATT) + + +def configure_plugins() -> None: + for plugin in _PLUGINS: + plugin.load() diff --git a/lightllm/platform/plugin/att.py b/lightllm/platform/plugin/att.py new file mode 100644 index 0000000000..d113853b31 --- /dev/null +++ b/lightllm/platform/plugin/att.py @@ -0,0 +1,63 @@ +from dataclasses import dataclass + +ATT_TABLE: dict[tuple[str, str, str], list["AttBackendEntry"]] = {} + + +@dataclass(frozen=True) +class AttBackendEntry: + # The name of the att backend. e.g., "triton", "flashinfer", "fa3". + name: str + # The category of the att backend. e.g., "standard", "mla", "nsa". + category: str + # The type of the kv. e.g., "int8kv", "int4kv", "fp8kv_sph". + kv_type: str + # The class of the att backend. + backend_cls: type + # The platforms of the att backend. e.g., "cuda", "ascend". + platforms: tuple[str, ...] + + +def register_att_backend( + name: str, + *, + category: str, + kv_type: str = "None", + platforms: tuple[str, ...] = ("cuda",), +): + def decorator(backend_cls): + if not platforms: + raise ValueError("The value of platforms must be specified for att backend registration.") + + key = (name, category, kv_type) + # Check if the att backend is already registered. + for existing in ATT_TABLE.get(key, []): + if set(existing.platforms) & set(platforms): + raise ValueError(f"The att backend {name} is already registered for platforms {existing.platforms}.") + + ATT_TABLE.setdefault(key, []).append( + AttBackendEntry( + name=name, + category=category, + kv_type=kv_type, + backend_cls=backend_cls, + platforms=platforms, + ) + ) + return backend_cls + + return decorator + + +def get_att_backend_class( + name: str, + category: str, + kv_type: str, + platform: str, +) -> type: + key = (name, category, kv_type) + for entry in ATT_TABLE.get(key, []): + if platform in entry.platforms: + return entry.backend_cls + raise ValueError( + f"The att backend {name} is not registered for category {category}, kv type {kv_type}, and platform {platform}." + ) diff --git a/lightllm/platform/plugin/common.py b/lightllm/platform/plugin/common.py new file mode 100644 index 0000000000..c44b674565 --- /dev/null +++ b/lightllm/platform/plugin/common.py @@ -0,0 +1,128 @@ +import importlib +from dataclasses import dataclass +from importlib.metadata import entry_points, EntryPoint +from typing import Any, Iterable, List, Mapping + +from lightllm.utils.envs_utils import get_env_start_args + + +@dataclass(frozen=True) +class PluginConfig: + modules: tuple[str, ...] = () + + +class Plugin: + def __init__(self, name: str, entry_point_group: str) -> None: + """Loader for pip-installed plugins. + + Args: + name: Plugin kind. + entry_point_group: Entry point group name. Looked up with ``importlib.metadata``. + + pyproject.toml: + [project.entry-points."lightllm.ops_plugin"] + example_ops = "lightllm_example_plugin.register:register_ops" + + setup.py: + setup( + entry_points={ + "lightllm.ops_plugin": [ + "example_ops = lightllm_example_plugin.register:register_ops", + ], + }, + ) + + After ``pip install``, pass ``--extra-ops example_ops`` to load and register. + """ + self.name = name + self.entry_point_group = entry_point_group + + def load(self) -> None: + names = self._name_from_cli() + if not names: + return + + configs = self._load_entry_point_plugins(names) + config = PluginConfig(modules=merge_config_field(configs, "modules")) + + for module in config.modules: + importlib.import_module(module) + + def _name_from_cli(self) -> tuple[str, ...]: + start_args = get_env_start_args() + return _to_str_tuple(getattr(start_args, f"extra_{self.name}", None)) + + def _load_entry_point_plugins(self, plugin_names: tuple[str, ...]) -> List[PluginConfig]: + given_names = set(plugin_names) + configs: List[PluginConfig] = [] + loaded_names: set[str] = set() + available_entry_points = list(_iter_entry_points(self.entry_point_group)) + for entry_point in available_entry_points: + # Only load entry points from CLI. + if entry_point.name not in given_names: + continue + + # example_ops = "lightllm_example_plugin.register:register_ops" + # -> register_fn = lightllm_example_plugin.register.register_ops + register_fn = entry_point.load() + config = parse_plugin_config(register_fn(), plugin_kind=self.name) + + configs.append(config) + loaded_names.add(entry_point.name) + + # Check if all required plugins are loaded. + missing = given_names - loaded_names + if missing: + available = tuple(sorted(ep.name for ep in available_entry_points)) + message = ( + f"{self.name} plugin(s) not found in entry point group " + f"{self.entry_point_group!r}: {sorted(missing)}" + ) + if available: + message += f". Installed plugins: {available}" + else: + message += ( + f". No {self.name} plugins installed; register entry points in group " + f"{self.entry_point_group!r} and pip install -e your plugin package." + ) + raise RuntimeError(message) + + return configs + + +def parse_plugin_config(value: Any, plugin_kind: str) -> PluginConfig: + if not isinstance(value, Mapping): + raise TypeError(f"{plugin_kind} plugin config must be a mapping, got {type(value)}") + + return PluginConfig(modules=_to_str_tuple(value.get("modules"))) + + +def merge_config_field(configs: Iterable[Any], field_name: str) -> tuple[str, ...]: + merged: list[str] = [] + seen: set[str] = set() + for config in configs: + for item in getattr(config, field_name): + if item in seen: + continue + seen.add(item) + merged.append(item) + return tuple(merged) + + +def _iter_entry_points(entry_point_group: str) -> Iterable[EntryPoint]: + eps = entry_points() + if hasattr(eps, "select"): + yield from eps.select(group=entry_point_group) + else: + yield from eps.get(entry_point_group, []) + + +def _to_str_tuple(value: str | Iterable[str] | None) -> tuple[str, ...]: + if value is None: + return () + + if isinstance(value, str): + parts: Iterable[str] = value.split(",") + else: + parts = value + return tuple(item.strip() for item in parts if item and item.strip()) diff --git a/lightllm/platform/plugin/ops.py b/lightllm/platform/plugin/ops.py new file mode 100644 index 0000000000..771b8f925f --- /dev/null +++ b/lightllm/platform/plugin/ops.py @@ -0,0 +1,43 @@ +from dataclasses import dataclass, field +from typing import Callable + +# The default order of fallback implementations. +DEFAULT_ORDER = ("triton", "torch") + +OPS_TABLE: dict[str, "OpEntry"] = {} + + +@dataclass +class OpEntry: + fns: dict[str, Callable] = field(default_factory=dict) + extras: list[str] = field(default_factory=list) + + @property + def impls(self) -> tuple[str, ...]: + extras = tuple(name for name in self.extras if name in self.fns) + defaults = tuple(kind for kind in DEFAULT_ORDER if kind in self.fns) + return extras + defaults + + +def register_op(op_name: str, *, impl: str): + def decorator(fn: Callable) -> Callable: + entry = OPS_TABLE.setdefault(op_name, OpEntry()) + if impl in entry.fns: + raise RuntimeError(f"The {op_name} op has already been registered with the {impl!r} implementation.") + entry.fns[impl] = fn + if impl not in DEFAULT_ORDER: + entry.extras.insert(0, impl) + return fn + + return decorator + + +def get_op(op_name: str) -> Callable: + entry = OPS_TABLE.get(op_name) + if entry is None: + raise NotImplementedError(f"The {op_name!r} op is not registered.") + for impl in entry.impls: + fn = entry.fns.get(impl) + if fn is not None: + return fn + raise NotImplementedError(f"The {op_name!r} op has no usable implementation in {entry.impls}.") diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index e6c077dbe4..d278d6329e 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -973,6 +973,20 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: (you should set it up by yourself). A NVTX range named 'LIGHTLLM_PROFILE' will be added within the profiling range.""", ) + parser.add_argument( + "--extra_ops", + type=str, + default=None, + help="""Extra ops plugin(s) to load at startup, comma-separated names. + Each name must match a pip-installed plugin.""" + ) + parser.add_argument( + "--extra_att", + type=str, + default=None, + help="""Extra att plugin(s) to load at startup, comma-separated names. + Each name must match a pip-installed plugin.""" + ) return parser diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index 9c89975de7..d063f41bd5 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -240,3 +240,6 @@ class StartArgs: disable_linear_att_small_page_cpu_cache: bool = field(default=False) linear_att_cache_size: Optional[int] = field(default=None) linear_att_ssm_data_type: Optional[str] = field(default="float32", metadata={"choices": ["bfloat16", "float32"]}) + + extra_ops: Optional[str] = field(default=None) + extra_att: Optional[str] = field(default=None)