Introduce CpuDeviceMixin and CpuSRTPlatform (#26385)
This commit is contained in:
@@ -10,6 +10,7 @@ from unittest.mock import MagicMock, patch
|
||||
import torch
|
||||
|
||||
from sglang.srt.platforms import _load_platform_class, _resolve_platform
|
||||
from sglang.srt.platforms.cpu import CpuDeviceMixin, CpuSRTPlatform
|
||||
from sglang.srt.platforms.cuda import CudaDeviceMixin, CudaSRTPlatform
|
||||
from sglang.srt.platforms.device_mixin import (
|
||||
CpuArchEnum,
|
||||
@@ -331,6 +332,123 @@ class TestCudaDeviceMixin(CustomTestCase):
|
||||
self.assertTrue(base.support_piecewise_cuda_graph())
|
||||
|
||||
|
||||
class TestCpuDeviceMixin(CustomTestCase):
|
||||
"""Tests for CPU device operation defaults (covers both x86 and ARM)."""
|
||||
|
||||
def test_cpu_platform_identity(self):
|
||||
base = CpuSRTPlatform()
|
||||
self.assertTrue(base.is_cpu())
|
||||
self.assertFalse(base.is_cuda())
|
||||
self.assertFalse(base.is_cuda_alike())
|
||||
self.assertIsInstance(base, CpuDeviceMixin)
|
||||
|
||||
def test_default_get_device_returns_cpu_device(self):
|
||||
base = CpuSRTPlatform()
|
||||
# ``local_rank`` is ignored — CPU has no per-rank device.
|
||||
self.assertEqual(base.get_device(0), torch.device("cpu"))
|
||||
self.assertEqual(base.get_device(7), torch.device("cpu"))
|
||||
|
||||
@patch("sglang.srt.platforms.cpu.psutil.virtual_memory")
|
||||
def test_default_get_device_total_memory_uses_psutil(self, mock_vm):
|
||||
mock_vm.return_value.total = 12345
|
||||
base = CpuSRTPlatform()
|
||||
self.assertEqual(base.get_device_total_memory(), 12345)
|
||||
|
||||
@patch("sglang.srt.platforms.cpu.psutil.virtual_memory")
|
||||
def test_default_get_available_memory_uses_psutil(self, mock_vm):
|
||||
mock_vm.return_value.available = 100
|
||||
mock_vm.return_value.total = 200
|
||||
base = CpuSRTPlatform()
|
||||
self.assertEqual(base.get_available_memory(), (100, 200))
|
||||
|
||||
@patch("sglang.srt.platforms.cpu.psutil.virtual_memory")
|
||||
def test_default_get_current_memory_usage_is_system_used(self, mock_vm):
|
||||
mock_vm.return_value.total = 1000
|
||||
mock_vm.return_value.available = 300
|
||||
base = CpuSRTPlatform()
|
||||
# system-used == total - available (not per-process RSS)
|
||||
self.assertEqual(base.get_current_memory_usage(), 700.0)
|
||||
|
||||
@patch("sglang.srt.platforms.cpu.psutil.virtual_memory")
|
||||
def test_memory_free_contract_yields_available(self, mock_vm):
|
||||
# The [Active] contract free = total - used must yield psutil.available.
|
||||
mock_vm.return_value.total = 1000
|
||||
mock_vm.return_value.available = 300
|
||||
base = CpuSRTPlatform()
|
||||
free = base.get_device_total_memory() - base.get_current_memory_usage()
|
||||
self.assertEqual(free, 300)
|
||||
|
||||
@patch("torch.cpu.set_device")
|
||||
def test_default_set_device_uses_torch_cpu(self, mock_set_device):
|
||||
base = CpuSRTPlatform()
|
||||
device = torch.device("cpu")
|
||||
base.set_device(device)
|
||||
# Documented CPU no-op, but called for symmetry with CudaDeviceMixin.
|
||||
mock_set_device.assert_called_once_with(device)
|
||||
|
||||
def test_default_set_device_does_not_flip_default(self):
|
||||
base = CpuSRTPlatform()
|
||||
# Must not call torch.set_default_device — process-wide default stays put.
|
||||
before = torch.empty(0).device
|
||||
base.set_device(torch.device("cpu"))
|
||||
after = torch.empty(0).device
|
||||
self.assertEqual(before, after)
|
||||
|
||||
@patch("sglang.srt.platforms.cpu.gc.collect")
|
||||
def test_default_empty_cache_calls_gc_collect(self, mock_collect):
|
||||
base = CpuSRTPlatform()
|
||||
base.empty_cache()
|
||||
mock_collect.assert_called_once_with()
|
||||
|
||||
@patch("torch.cpu.synchronize")
|
||||
def test_default_synchronize_uses_torch_cpu(self, mock_synchronize):
|
||||
base = CpuSRTPlatform()
|
||||
base.synchronize()
|
||||
mock_synchronize.assert_called_once_with()
|
||||
|
||||
def test_default_distributed_backend_is_gloo(self):
|
||||
base = CpuSRTPlatform()
|
||||
self.assertEqual(base.get_torch_distributed_backend_str(), "gloo")
|
||||
|
||||
@patch("platform.machine", return_value="aarch64")
|
||||
def test_cpu_arch_property_resolves_and_caches(self, mock_machine):
|
||||
base = CpuSRTPlatform()
|
||||
self.assertEqual(base.cpu_arch, CpuArchEnum.ARM)
|
||||
# cached_property: second access must not re-query platform.machine
|
||||
call_count = mock_machine.call_count
|
||||
self.assertEqual(base.cpu_arch, CpuArchEnum.ARM)
|
||||
self.assertEqual(mock_machine.call_count, call_count)
|
||||
|
||||
@patch("platform.machine", return_value="aarch64")
|
||||
def test_get_device_name_arm_branch(self, _mock_machine):
|
||||
base = CpuSRTPlatform()
|
||||
name = base.get_device_name()
|
||||
self.assertIn("aarch64", name)
|
||||
|
||||
@patch("platform.machine", return_value="x86_64")
|
||||
def test_get_device_name_x86_branch(self, _mock_machine):
|
||||
base = CpuSRTPlatform()
|
||||
name = base.get_device_name()
|
||||
self.assertIn("x86_64", name)
|
||||
|
||||
@patch("platform.machine", return_value="aarch64")
|
||||
def test_get_device_uuid_returns_machine(self, _mock_machine):
|
||||
base = CpuSRTPlatform()
|
||||
self.assertEqual(base.get_device_uuid(), "aarch64")
|
||||
|
||||
def test_get_device_capability_returns_none(self):
|
||||
base = CpuSRTPlatform()
|
||||
self.assertIsNone(base.get_device_capability())
|
||||
|
||||
def test_cpu_srt_platform_capabilities(self):
|
||||
base = CpuSRTPlatform()
|
||||
self.assertFalse(base.supports_fp8())
|
||||
self.assertFalse(base.support_cuda_graph())
|
||||
self.assertFalse(base.support_piecewise_cuda_graph())
|
||||
# Override of the SRTPlatform default (True) — no GPU to pin to.
|
||||
self.assertFalse(base.is_pin_memory_available())
|
||||
|
||||
|
||||
class TestSRTPlatformOverrides(CustomTestCase):
|
||||
"""Tests for SRTPlatform method overrides via plugins."""
|
||||
|
||||
@@ -496,6 +614,51 @@ class TestResolvePlatformAutoDiscover(CustomTestCase):
|
||||
self.assertIsInstance(result, SRTPlatform)
|
||||
self.assertNotIsInstance(result, CudaSRTPlatform)
|
||||
|
||||
@patch("sglang.srt.platforms.load_plugins_by_group")
|
||||
@patch("sglang.srt.platforms._is_cuda_available")
|
||||
@patch("sglang.srt.platforms._is_cpu_available")
|
||||
@patch("sglang.srt.platforms.envs")
|
||||
def test_no_plugin_cpu_engine_enabled_activates_cpu_fallback(
|
||||
self, mock_envs, mock_is_cpu, mock_is_cuda, mock_load
|
||||
):
|
||||
"""SGLANG_USE_CPU_ENGINE=1 + no plugins → CpuSRTPlatform."""
|
||||
mock_envs.SGLANG_PLATFORM.get.return_value = ""
|
||||
mock_is_cpu.return_value = True
|
||||
mock_is_cuda.return_value = False
|
||||
mock_load.return_value = {}
|
||||
result = _resolve_platform()
|
||||
self.assertIsInstance(result, CpuSRTPlatform)
|
||||
|
||||
@patch("sglang.srt.platforms.load_plugins_by_group")
|
||||
@patch("sglang.srt.platforms._is_cuda_available")
|
||||
@patch("sglang.srt.platforms._is_cpu_available")
|
||||
@patch("sglang.srt.platforms.envs")
|
||||
def test_cpu_engine_wins_over_cuda(
|
||||
self, mock_envs, mock_is_cpu, mock_is_cuda, mock_load
|
||||
):
|
||||
"""When both CPU engine and CUDA are available, explicit opt-in wins."""
|
||||
mock_envs.SGLANG_PLATFORM.get.return_value = ""
|
||||
mock_is_cpu.return_value = True
|
||||
mock_is_cuda.return_value = True
|
||||
mock_load.return_value = {}
|
||||
result = _resolve_platform()
|
||||
self.assertIsInstance(result, CpuSRTPlatform)
|
||||
|
||||
@patch("sglang.srt.platforms.load_plugins_by_group")
|
||||
@patch("sglang.srt.platforms._is_cuda_available")
|
||||
@patch("sglang.srt.platforms._is_cpu_available")
|
||||
@patch("sglang.srt.platforms.envs")
|
||||
def test_no_plugin_cpu_engine_disabled_prefers_cuda(
|
||||
self, mock_envs, mock_is_cpu, mock_is_cuda, mock_load
|
||||
):
|
||||
"""Regression: CPU opt-out leaves the existing CUDA fallback path intact."""
|
||||
mock_envs.SGLANG_PLATFORM.get.return_value = ""
|
||||
mock_is_cpu.return_value = False
|
||||
mock_is_cuda.return_value = True
|
||||
mock_load.return_value = {}
|
||||
result = _resolve_platform()
|
||||
self.assertIsInstance(result, CudaSRTPlatform)
|
||||
|
||||
@patch("sglang.srt.platforms.load_plugins_by_group")
|
||||
@patch("sglang.srt.platforms.envs")
|
||||
def test_multiple_plugins_activate_raises(self, mock_envs, mock_load):
|
||||
|
||||
Reference in New Issue
Block a user