[Runtime] Let out-of-tree platforms provide full graph backends (#37969)
Co-authored-by: Xiaoyu Zhang <1182563586@qq.com>
This commit is contained in:
co-authored by
Xiaoyu Zhang
parent
c46bf5e990
commit
1fdd6c8921
@@ -37,9 +37,8 @@ from sglang.srt.model_executor.runner_backend.full_cuda_graph_backend import (
|
|||||||
from sglang.srt.model_executor.runner_backend.tc_piecewise_cuda_graph_backend import (
|
from sglang.srt.model_executor.runner_backend.tc_piecewise_cuda_graph_backend import (
|
||||||
TcPiecewiseCudaGraphBackend,
|
TcPiecewiseCudaGraphBackend,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.platforms import current_platform
|
||||||
get_exec,
|
from sglang.srt.runtime_context import get_exec
|
||||||
)
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.model_executor.runner.base_cuda_graph_runner import (
|
from sglang.srt.model_executor.runner.base_cuda_graph_runner import (
|
||||||
@@ -99,8 +98,17 @@ def resolve_decode_backend(
|
|||||||
"falling back to 'full'."
|
"falling back to 'full'."
|
||||||
)
|
)
|
||||||
_TC_PIECEWISE_DECODE_FALLBACK_LOGGED = True
|
_TC_PIECEWISE_DECODE_FALLBACK_LOGGED = True
|
||||||
return FullCudaGraphBackend(
|
|
||||||
cuda_graph_runner, enable_memory_saver=enable_memory_saver
|
full_backend_cls = None
|
||||||
|
if current_platform.is_out_of_tree():
|
||||||
|
full_backend_cls = current_platform.get_full_graph_backend_cls()
|
||||||
|
|
||||||
|
if full_backend_cls is None:
|
||||||
|
full_backend_cls = FullCudaGraphBackend
|
||||||
|
|
||||||
|
return full_backend_cls(
|
||||||
|
cuda_graph_runner,
|
||||||
|
enable_memory_saver=enable_memory_saver,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -12,7 +12,7 @@ Out-of-tree platforms register via setuptools entry_points under the
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from typing import TYPE_CHECKING, Optional, Type
|
from typing import TYPE_CHECKING, Any, Optional, Type
|
||||||
|
|
||||||
from sglang.srt.platforms.device_mixin import DeviceMixin, PlatformEnum
|
from sglang.srt.platforms.device_mixin import DeviceMixin, PlatformEnum
|
||||||
|
|
||||||
@@ -60,6 +60,10 @@ class SRTPlatform(DeviceMixin):
|
|||||||
"""Return the graph runner class for this platform."""
|
"""Return the graph runner class for this platform."""
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def get_full_graph_backend_cls(self) -> type[Any]:
|
||||||
|
"""Return the full device-graph backend class for this platform."""
|
||||||
|
return None
|
||||||
|
|
||||||
def get_mha_kv_pool_cls(self) -> type:
|
def get_mha_kv_pool_cls(self) -> type:
|
||||||
"""Return the MHA KV pool class for this platform."""
|
"""Return the MHA KV pool class for this platform."""
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|||||||
@@ -0,0 +1,189 @@
|
|||||||
|
"""CPU-only tests for decode graph backend resolution."""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest import mock
|
||||||
|
|
||||||
|
from sglang.srt.model_executor.cuda_graph_config import Backend
|
||||||
|
from sglang.srt.model_executor.runner_backend import utils as backend_utils
|
||||||
|
from sglang.srt.platforms.device_mixin import PlatformEnum
|
||||||
|
from sglang.srt.platforms.interface import SRTPlatform
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
def _make_graph_runner(*, device="custom"):
|
||||||
|
return SimpleNamespace(model_runner=SimpleNamespace(device=device))
|
||||||
|
|
||||||
|
|
||||||
|
def _make_exec_config(
|
||||||
|
*,
|
||||||
|
backend=Backend.FULL,
|
||||||
|
enable_memory_saver=False,
|
||||||
|
debug_cuda_graph=False,
|
||||||
|
):
|
||||||
|
return SimpleNamespace(
|
||||||
|
graph=SimpleNamespace(
|
||||||
|
cuda_graph_config=SimpleNamespace(decode=SimpleNamespace(backend=backend)),
|
||||||
|
debug_cuda_graph=debug_cuda_graph,
|
||||||
|
),
|
||||||
|
features=SimpleNamespace(
|
||||||
|
enable_memory_saver=enable_memory_saver,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_platform(*, is_out_of_tree, full_backend_cls=None):
|
||||||
|
return SimpleNamespace(
|
||||||
|
is_out_of_tree=mock.Mock(return_value=is_out_of_tree),
|
||||||
|
get_full_graph_backend_cls=mock.Mock(return_value=full_backend_cls),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class _OutOfTreePlatformWithoutFullGraphBackend(SRTPlatform):
|
||||||
|
_enum = PlatformEnum.OOT
|
||||||
|
|
||||||
|
|
||||||
|
class TestResolveDecodeBackend(CustomTestCase):
|
||||||
|
def test_out_of_tree_platform_provides_full_backend(self):
|
||||||
|
runner = _make_graph_runner()
|
||||||
|
exec_config = _make_exec_config(enable_memory_saver=True)
|
||||||
|
expected_backend = object()
|
||||||
|
backend_cls = mock.Mock(return_value=expected_backend)
|
||||||
|
platform = _make_platform(
|
||||||
|
is_out_of_tree=True,
|
||||||
|
full_backend_cls=backend_cls,
|
||||||
|
)
|
||||||
|
|
||||||
|
with (
|
||||||
|
mock.patch.object(
|
||||||
|
backend_utils,
|
||||||
|
"get_exec",
|
||||||
|
return_value=exec_config,
|
||||||
|
),
|
||||||
|
mock.patch.object(
|
||||||
|
backend_utils,
|
||||||
|
"current_platform",
|
||||||
|
platform,
|
||||||
|
),
|
||||||
|
mock.patch.object(
|
||||||
|
backend_utils,
|
||||||
|
"FullCudaGraphBackend",
|
||||||
|
) as default_backend_cls,
|
||||||
|
):
|
||||||
|
actual_backend = backend_utils.resolve_decode_backend(runner)
|
||||||
|
|
||||||
|
self.assertIs(actual_backend, expected_backend)
|
||||||
|
platform.get_full_graph_backend_cls.assert_called_once_with()
|
||||||
|
backend_cls.assert_called_once_with(
|
||||||
|
runner,
|
||||||
|
enable_memory_saver=True,
|
||||||
|
)
|
||||||
|
default_backend_cls.assert_not_called()
|
||||||
|
|
||||||
|
def test_out_of_tree_platform_without_full_backend_uses_default(self):
|
||||||
|
runner = _make_graph_runner()
|
||||||
|
exec_config = _make_exec_config()
|
||||||
|
expected_backend = object()
|
||||||
|
platform = _OutOfTreePlatformWithoutFullGraphBackend()
|
||||||
|
|
||||||
|
with (
|
||||||
|
mock.patch.object(
|
||||||
|
backend_utils,
|
||||||
|
"get_exec",
|
||||||
|
return_value=exec_config,
|
||||||
|
),
|
||||||
|
mock.patch.object(
|
||||||
|
backend_utils,
|
||||||
|
"current_platform",
|
||||||
|
platform,
|
||||||
|
),
|
||||||
|
mock.patch.object(
|
||||||
|
backend_utils,
|
||||||
|
"FullCudaGraphBackend",
|
||||||
|
return_value=expected_backend,
|
||||||
|
) as default_backend_cls,
|
||||||
|
):
|
||||||
|
actual_backend = backend_utils.resolve_decode_backend(runner)
|
||||||
|
|
||||||
|
self.assertIs(actual_backend, expected_backend)
|
||||||
|
default_backend_cls.assert_called_once_with(
|
||||||
|
runner,
|
||||||
|
enable_memory_saver=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_in_tree_platform_uses_default_full_backend(self):
|
||||||
|
runner = _make_graph_runner(device="cuda")
|
||||||
|
exec_config = _make_exec_config()
|
||||||
|
expected_backend = object()
|
||||||
|
platform = _make_platform(is_out_of_tree=False)
|
||||||
|
|
||||||
|
with (
|
||||||
|
mock.patch.object(
|
||||||
|
backend_utils,
|
||||||
|
"get_exec",
|
||||||
|
return_value=exec_config,
|
||||||
|
),
|
||||||
|
mock.patch.object(
|
||||||
|
backend_utils,
|
||||||
|
"current_platform",
|
||||||
|
platform,
|
||||||
|
),
|
||||||
|
mock.patch.object(
|
||||||
|
backend_utils,
|
||||||
|
"FullCudaGraphBackend",
|
||||||
|
return_value=expected_backend,
|
||||||
|
) as default_backend_cls,
|
||||||
|
):
|
||||||
|
actual_backend = backend_utils.resolve_decode_backend(runner)
|
||||||
|
|
||||||
|
self.assertIs(actual_backend, expected_backend)
|
||||||
|
platform.get_full_graph_backend_cls.assert_not_called()
|
||||||
|
default_backend_cls.assert_called_once_with(
|
||||||
|
runner,
|
||||||
|
enable_memory_saver=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_breakable_backend_does_not_consult_full_backend_factory(self):
|
||||||
|
runner = _make_graph_runner()
|
||||||
|
exec_config = _make_exec_config(
|
||||||
|
backend=Backend.BREAKABLE,
|
||||||
|
enable_memory_saver=True,
|
||||||
|
debug_cuda_graph=True,
|
||||||
|
)
|
||||||
|
expected_backend = object()
|
||||||
|
platform = _make_platform(is_out_of_tree=True)
|
||||||
|
|
||||||
|
with (
|
||||||
|
mock.patch.object(
|
||||||
|
backend_utils,
|
||||||
|
"get_exec",
|
||||||
|
return_value=exec_config,
|
||||||
|
),
|
||||||
|
mock.patch.object(
|
||||||
|
backend_utils,
|
||||||
|
"current_platform",
|
||||||
|
platform,
|
||||||
|
),
|
||||||
|
mock.patch.object(
|
||||||
|
backend_utils,
|
||||||
|
"BreakableCudaGraphBackend",
|
||||||
|
return_value=expected_backend,
|
||||||
|
) as breakable_backend_cls,
|
||||||
|
):
|
||||||
|
actual_backend = backend_utils.resolve_decode_backend(runner)
|
||||||
|
|
||||||
|
self.assertIs(actual_backend, expected_backend)
|
||||||
|
platform.is_out_of_tree.assert_not_called()
|
||||||
|
platform.get_full_graph_backend_cls.assert_not_called()
|
||||||
|
breakable_backend_cls.assert_called_once_with(
|
||||||
|
runner,
|
||||||
|
enable_memory_saver=True,
|
||||||
|
debug_eager=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user