diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 143054cd6..ef38e21de 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -203,6 +203,7 @@ from sglang.srt.observability.scheduler_metrics_mixin import ( ) from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info from sglang.srt.parser.reasoning_parser import ReasoningParser +from sglang.srt.platforms import current_platform from sglang.srt.plugins import load_plugins from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.server_args import PortArgs, ServerArgs, get_global_server_args @@ -214,7 +215,6 @@ from sglang.srt.utils import ( broadcast_pyobj, configure_gc_logger, configure_logger, - empty_device_cache, freeze_gc, get_available_gpu_memory, get_bool_env_var, @@ -3454,7 +3454,7 @@ class Scheduler( self.draft_worker.clear_cache_pool() if empty_cache: - empty_device_cache(self.device_module) + current_platform.empty_cache() logger.info("Cache flushed successfully!") success = True else: @@ -3852,7 +3852,7 @@ class IdleSleeper: and real_time() - self.last_empty_time > self.empty_cache_interval ): self.last_empty_time = real_time() - empty_device_cache() + current_platform.empty_cache() def is_health_check_generate_req(recv_req): diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index 9bf36c317..8d7ad26a4 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -406,7 +406,7 @@ class MambaPool: return dst_index def get_cpu_copy(self, indices): - torch.cuda.synchronize() + current_platform.synchronize() conv_cpu = [ conv[:, indices].to("cpu", non_blocking=True) for conv in self.mamba_cache.conv @@ -414,18 +414,18 @@ class MambaPool: temporal_cpu = self.mamba_cache.temporal[:, indices].to( "cpu", non_blocking=True ) - torch.cuda.synchronize() + current_platform.synchronize() return conv_cpu, temporal_cpu def load_cpu_copy(self, mamba_cache_cpu, indices): conv_cpu, temporal_cpu = mamba_cache_cpu - torch.cuda.synchronize() + current_platform.synchronize() for i, conv in enumerate(self.mamba_cache.conv): conv[:, indices] = conv_cpu[i].to(conv.device, non_blocking=True) self.mamba_cache.temporal[:, indices] = temporal_cpu.to( self.mamba_cache.temporal.device, non_blocking=True ) - torch.cuda.synchronize() + current_platform.synchronize() def get_contiguous_buf_infos(self): """ @@ -982,7 +982,7 @@ class MHATokenToKVPool(KVCache): return kv_data_ptrs, kv_data_lens, kv_item_lens def get_cpu_copy(self, indices, mamba_indices=None): - torch.cuda.synchronize() + current_platform.synchronize() kv_cache_cpu = [] chunk_size = self.cpu_offloading_chunk_size for layer_id in range(self.layer_num): @@ -996,11 +996,11 @@ class MHATokenToKVPool(KVCache): "cpu", non_blocking=True ) kv_cache_cpu[-1].append([k_cpu, v_cpu]) - torch.cuda.synchronize() + current_platform.synchronize() return kv_cache_cpu def load_cpu_copy(self, kv_cache_cpu, indices, mamba_indices=None): - torch.cuda.synchronize() + current_platform.synchronize() chunk_size = self.cpu_offloading_chunk_size for layer_id in range(self.layer_num): for i in range(0, len(indices), chunk_size): @@ -1014,7 +1014,7 @@ class MHATokenToKVPool(KVCache): v_chunk = v_cpu.to(self.v_buffer[0].device, non_blocking=True) self.k_buffer[layer_id][chunk_indices] = k_chunk self.v_buffer[layer_id][chunk_indices] = v_chunk - torch.cuda.synchronize() + current_platform.synchronize() def _get_key_buffer(self, layer_id: int): # for internal use of referencing @@ -1822,7 +1822,7 @@ class MLATokenToKVPool(KVCache): return cache_k_nope, cache_k_rope def get_cpu_copy(self, indices, mamba_indices=None): - torch.cuda.synchronize() + current_platform.synchronize() kv_cache_cpu = [] chunk_size = self.cpu_offloading_chunk_size for layer_id in range(self.layer_num): @@ -1833,11 +1833,11 @@ class MLATokenToKVPool(KVCache): "cpu", non_blocking=True ) kv_cache_cpu[-1].append(kv_cpu) - torch.cuda.synchronize() + current_platform.synchronize() return kv_cache_cpu def load_cpu_copy(self, kv_cache_cpu, indices, mamba_indices=None): - torch.cuda.synchronize() + current_platform.synchronize() chunk_size = self.cpu_offloading_chunk_size for layer_id in range(self.layer_num): for i in range(0, len(indices), chunk_size): @@ -1846,7 +1846,7 @@ class MLATokenToKVPool(KVCache): assert kv_cpu.shape[0] == len(chunk_indices) kv_chunk = kv_cpu.to(self.kv_buffer[0].device, non_blocking=True) self.kv_buffer[layer_id][chunk_indices] = kv_chunk - torch.cuda.synchronize() + current_platform.synchronize() class MLATokenToKVPoolFP4(MLATokenToKVPool): diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index ff937dd9c..53bb200c2 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -1288,7 +1288,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): # Single warmup all_reduce to initialize NCCL/RCCL communicator warmup_tensor = torch.zeros(1, device=torch.cuda.current_device()) dist.all_reduce(warmup_tensor, group=tp_group_handle) - torch.cuda.synchronize() + current_platform.synchronize() warmup_elapsed = time.perf_counter() - warmup_start logger.info( @@ -1847,7 +1847,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): f"group_rank={group_rank}, world_size={world_size}, group_name={group_name}, backend={backend}" ) - torch.cuda.empty_cache() + current_platform.empty_cache() success = False message = "" try: @@ -1867,7 +1867,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): message = f"Failed to init group: {e}." logger.error(message) - torch.cuda.empty_cache() + current_platform.empty_cache() return success, message def send_weights_to_remote_instance( @@ -1895,7 +1895,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): logger.error(message) return False, message - torch.cuda.empty_cache() + current_platform.empty_cache() success = False na = NetworkAddress(master_address, group_port) message = "" @@ -1915,7 +1915,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): # destroy the process group after sending weights del self._weights_send_group[group_name] torch.distributed.distributed_c10d.destroy_process_group(send_group) - torch.cuda.empty_cache() + current_platform.empty_cache() return success, message def init_weights_update_group( diff --git a/python/sglang/srt/model_loader/loader.py b/python/sglang/srt/model_loader/loader.py index 6dfcdb5d3..9869740c9 100644 --- a/python/sglang/srt/model_loader/loader.py +++ b/python/sglang/srt/model_loader/loader.py @@ -106,6 +106,7 @@ from sglang.srt.model_loader.weight_utils import ( safetensors_weights_iterator, set_runai_streamer_env, ) +from sglang.srt.platforms import current_platform from sglang.srt.utils import ( get_bool_env_var, get_device_capability, @@ -1221,7 +1222,7 @@ class QuantizedRLModelLoader(DefaultModelLoader): del current_param_data if is_last_update: gc.collect() - torch.cuda.empty_cache() + current_platform.empty_cache() logger.info("[QuantizedRL] Reload complete") return updated_param_names, is_last_update @@ -1912,7 +1913,7 @@ class BitsAndBytesModelLoader(BaseModelLoader): model.load_weights(qweight_iterator) - torch.cuda.empty_cache() + current_platform.empty_cache() param_dict = dict(model.named_parameters()) stacked_quant_state_dict: Dict[str, Dict[int, Any]] = {} @@ -2220,7 +2221,7 @@ class RemoteInstanceModelLoader(BaseModelLoader): tp_rank=load_config.tp_rank, instance_ip=instance_ip, ) - torch.cuda.synchronize() + current_platform.synchronize() end_build_group_tic = time.time() logger.debug( f"finish building group for remote instance, time used: {(end_build_group_tic - start_build_group_tic):.4f}s" @@ -2246,7 +2247,7 @@ class RemoteInstanceModelLoader(BaseModelLoader): src=0, group=client._model_update_group, ) - torch.cuda.synchronize() + current_platform.synchronize() _post_load_weights(model) end_get_weights_tic = time.time() @@ -2257,7 +2258,7 @@ class RemoteInstanceModelLoader(BaseModelLoader): torch.distributed.distributed_c10d.destroy_process_group( client._model_update_group ) - torch.cuda.empty_cache() + current_platform.empty_cache() def load_model_from_remote_instance_by_transfer_engine( self, model, transfer_engine, seed_url, tp_rank diff --git a/python/sglang/srt/platforms/__init__.py b/python/sglang/srt/platforms/__init__.py index 5282fe92f..9c68eefa5 100644 --- a/python/sglang/srt/platforms/__init__.py +++ b/python/sglang/srt/platforms/__init__.py @@ -14,8 +14,12 @@ import logging import pkgutil from importlib.metadata import entry_points +import torch + from sglang.srt.environ import envs +from sglang.srt.platforms.cuda import CudaSRTPlatform from sglang.srt.platforms.interface import SRTPlatform +from sglang.srt.platforms.rocm import RocmSRTPlatform from sglang.srt.plugins import PLATFORM_PLUGINS_GROUP, load_plugins_by_group logger = logging.getLogger(__name__) @@ -23,6 +27,14 @@ logger = logging.getLogger(__name__) _current_platform: SRTPlatform | None = None +def _is_cuda_available() -> bool: + return bool(torch.cuda.is_available() and torch.version.hip is None) + + +def _is_rocm_available() -> bool: + return bool(torch.cuda.is_available() and torch.version.hip is not None) + + def _resolve_platform() -> SRTPlatform: """ Discover and instantiate the active platform. @@ -39,7 +51,9 @@ def _resolve_platform() -> SRTPlatform: SGLANG_PLATFORM unset (auto-discover): - Import and activate all discovered plugins - - 0 activated → fallback base SRTPlatform + - 0 activated + CUDA available → fallback CudaSRTPlatform + - 0 activated + ROCm available → fallback RocmSRTPlatform + - 0 activated + neither → fallback base SRTPlatform - 1 activated → use it - N activated → RuntimeError (must set SGLANG_PLATFORM) @@ -90,7 +104,17 @@ def _resolve_platform() -> SRTPlatform: logger.exception("Failed to activate platform plugin: %s", name) if len(activated) == 0: - logger.debug("No platform detected. Using base SRTPlatform with defaults.") + if _is_cuda_available(): + logger.debug( + "No platform plugin detected. Using CUDA SRTPlatform defaults." + ) + return CudaSRTPlatform() + if _is_rocm_available(): + logger.debug( + "No platform plugin detected. Using ROCm SRTPlatform defaults." + ) + return RocmSRTPlatform() + logger.debug("No platform detected. Using base SRTPlatform.") return SRTPlatform() if len(activated) == 1: diff --git a/python/sglang/srt/platforms/cuda.py b/python/sglang/srt/platforms/cuda.py new file mode 100644 index 000000000..5800886d2 --- /dev/null +++ b/python/sglang/srt/platforms/cuda.py @@ -0,0 +1,75 @@ +"""CUDA device operations for the SRT platform layer.""" + +from typing import Optional + +import torch + +from sglang.srt.platforms.device_mixin import ( + DeviceCapability, + DeviceMixin, + PlatformEnum, +) +from sglang.srt.platforms.interface import SRTPlatform + + +class CudaDeviceMixin(DeviceMixin): + """CUDA implementation of the shared device operations.""" + + _enum: PlatformEnum = PlatformEnum.CUDA + device_name: str = "cuda" + device_type: str = "cuda" + + def get_device_total_memory(self, device_id: int = 0) -> int: + return int(torch.cuda.get_device_properties(device_id).total_memory) + + def get_current_memory_usage( + self, device: Optional["torch.device"] = None + ) -> float: + return float(torch.cuda.max_memory_allocated(device)) + + def get_device(self, local_rank: int) -> "torch.device": + return torch.device("cuda", local_rank) + + def set_device(self, device: "torch.device") -> None: + torch.cuda.set_device(device) + + def get_device_name(self, device_id: int = 0) -> str: + return str(torch.cuda.get_device_name(device_id)) + + def get_device_uuid(self, device_id: int = 0) -> str: + return str(torch.cuda.get_device_properties(device_id).uuid) + + def get_device_capability(self, device_id: int = 0) -> DeviceCapability: + major, minor = torch.cuda.get_device_capability(device_id) + return DeviceCapability(major, minor) + + def empty_cache(self) -> None: + torch.cuda.empty_cache() + + def synchronize(self) -> None: + torch.cuda.synchronize() + + def get_available_memory(self, device_id: int = 0) -> tuple[int, int]: + return torch.cuda.mem_get_info(device_id) + + def get_torch_distributed_backend_str(self) -> str: + return "nccl" + + @classmethod + def seed_everything(cls, seed: int | None = None) -> None: + if seed is not None: + super().seed_everything(seed) + torch.cuda.manual_seed_all(seed) + + +class CudaSRTPlatform(CudaDeviceMixin, SRTPlatform): + """Default in-tree CUDA SRT platform.""" + + def supports_fp8(self) -> bool: + return True + + def support_cuda_graph(self) -> bool: + return True + + def support_piecewise_cuda_graph(self) -> bool: + return True diff --git a/python/sglang/srt/platforms/device_mixin.py b/python/sglang/srt/platforms/device_mixin.py index 49bd2f0c5..c42c4b2bd 100644 --- a/python/sglang/srt/platforms/device_mixin.py +++ b/python/sglang/srt/platforms/device_mixin.py @@ -26,10 +26,11 @@ Method status annotations: """ import enum -from typing import TYPE_CHECKING, NamedTuple, Optional +import random +from typing import NamedTuple, Optional -if TYPE_CHECKING: - import torch +import numpy as np +import torch class PlatformEnum(enum.Enum): @@ -203,19 +204,12 @@ class DeviceMixin: @classmethod def inference_mode(cls): """[Planned] Return inference mode context manager.""" - import torch - return torch.inference_mode(mode=True) @classmethod def seed_everything(cls, seed: int | None = None) -> None: """[Planned] Set random seeds for reproducibility across all libraries.""" if seed is not None: - import random - - import numpy as np - import torch - random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) diff --git a/python/sglang/srt/platforms/interface.py b/python/sglang/srt/platforms/interface.py index b7d4bfdc7..eb6086d8f 100644 --- a/python/sglang/srt/platforms/interface.py +++ b/python/sglang/srt/platforms/interface.py @@ -7,16 +7,11 @@ and adds SRT-specific subsystem factory methods, capability flags, and configuration lifecycle hooks. Out-of-tree platforms register via setuptools entry_points under the -"sglang.platform_plugins" group and should subclass SRTPlatform. +"sglang.srt.platforms" group and should subclass SRTPlatform. """ -from typing import TYPE_CHECKING - from sglang.srt.platforms.device_mixin import DeviceMixin, PlatformEnum -if TYPE_CHECKING: - pass - # Re-export for convenience __all__ = ["SRTPlatform", "PlatformEnum"] diff --git a/python/sglang/srt/platforms/rocm.py b/python/sglang/srt/platforms/rocm.py new file mode 100644 index 000000000..c250c482e --- /dev/null +++ b/python/sglang/srt/platforms/rocm.py @@ -0,0 +1,31 @@ +"""ROCm device operations for the SRT platform layer. + +PyTorch exposes ROCm through the same ``torch.cuda.*`` API surface as CUDA +(HIP is a binary shim, and ``torch.device("rocm")`` does not exist). So +``RocmDeviceMixin`` inherits all device ops from ``CudaDeviceMixin`` and +only overrides identity (``_enum``, ``device_name``). +""" + +from sglang.srt.platforms.cuda import CudaDeviceMixin +from sglang.srt.platforms.device_mixin import PlatformEnum +from sglang.srt.platforms.interface import SRTPlatform + + +class RocmDeviceMixin(CudaDeviceMixin): + """ROCm device ops — identical surface to CUDA via torch.cuda's HIP shim.""" + + _enum: PlatformEnum = PlatformEnum.ROCM + device_name: str = "rocm" + # device_type stays "cuda" — torch.device("cuda") is the only valid + # device-type string for HIP devices in PyTorch. + + +class RocmSRTPlatform(RocmDeviceMixin, SRTPlatform): + """Default in-tree ROCm SRT platform. + + Capability flags (supports_fp8, support_cuda_graph, support_piecewise_cuda_graph) + keep the conservative SRTPlatform defaults rather than mirroring CudaSRTPlatform. + They are currently only consulted in OOT branches gated on is_out_of_tree(), + so the defaults are behaviorally inert for the in-tree ROCm path. A follow-up + that migrates AMD-specific gating off legacy is_hip() should set these here. + """ diff --git a/test/registered/unit/platforms/test_platform_interface.py b/test/registered/unit/platforms/test_platform_interface.py index f5bfcd337..8c8cec606 100644 --- a/test/registered/unit/platforms/test_platform_interface.py +++ b/test/registered/unit/platforms/test_platform_interface.py @@ -7,7 +7,10 @@ and the platform discovery / lazy initialization mechanism. from unittest.mock import MagicMock, patch +import torch + from sglang.srt.platforms import _load_platform_class, _resolve_platform +from sglang.srt.platforms.cuda import CudaDeviceMixin, CudaSRTPlatform from sglang.srt.platforms.device_mixin import ( CpuArchEnum, DeviceCapability, @@ -222,6 +225,111 @@ class TestSRTPlatform(CustomTestCase): base = SRTPlatform() self.assertEqual(base.get_compile_backend(mode="npugraph_ex"), "inductor") + def test_base_device_identity_stays_unspecified(self): + """The abstract SRT base should not claim any concrete in-tree device.""" + base = SRTPlatform() + self.assertFalse(base.is_cuda()) + self.assertFalse(base.is_cuda_alike()) + + +class TestCudaDeviceMixin(CustomTestCase): + """Tests for CUDA device operation defaults.""" + + def test_default_get_device_returns_cuda_device(self): + base = CudaSRTPlatform() + self.assertEqual(base.get_device(2), torch.device("cuda", 2)) + + def test_cuda_platform_identity(self): + base = CudaSRTPlatform() + self.assertTrue(base.is_cuda()) + self.assertTrue(base.is_cuda_alike()) + self.assertIsInstance(base, CudaDeviceMixin) + + @patch("torch.cuda.get_device_properties") + def test_default_get_device_total_memory_uses_cuda( + self, mock_get_device_properties + ): + mock_get_device_properties.return_value.total_memory = 123 + base = CudaSRTPlatform() + self.assertEqual(base.get_device_total_memory(1), 123) + mock_get_device_properties.assert_called_once_with(1) + + @patch("torch.cuda.max_memory_allocated", return_value=456) + def test_default_get_current_memory_usage_uses_cuda( + self, mock_max_memory_allocated + ): + base = CudaSRTPlatform() + device = torch.device("cuda", 1) + self.assertEqual(base.get_current_memory_usage(device), 456.0) + mock_max_memory_allocated.assert_called_once_with(device) + + @patch("torch.cuda.set_device") + def test_default_set_device_uses_cuda(self, mock_set_device): + base = CudaSRTPlatform() + device = torch.device("cuda", 1) + base.set_device(device) + mock_set_device.assert_called_once_with(device) + + @patch("torch.cuda.get_device_name", return_value="NVIDIA H100") + def test_default_get_device_name_uses_cuda(self, mock_get_device_name): + base = CudaSRTPlatform() + self.assertEqual(base.get_device_name(1), "NVIDIA H100") + mock_get_device_name.assert_called_once_with(1) + + @patch("torch.cuda.get_device_properties") + def test_default_get_device_uuid_uses_cuda(self, mock_get_device_properties): + mock_get_device_properties.return_value.uuid = "1234" + base = CudaSRTPlatform() + self.assertEqual(base.get_device_uuid(1), "1234") + mock_get_device_properties.assert_called_once_with(1) + + @patch("torch.cuda.get_device_capability", return_value=(9, 0)) + def test_default_get_device_capability_uses_cuda(self, mock_get_device_capability): + base = CudaSRTPlatform() + self.assertEqual(base.get_device_capability(1), DeviceCapability(9, 0)) + mock_get_device_capability.assert_called_once_with(1) + + @patch("torch.cuda.empty_cache") + def test_default_empty_cache_uses_cuda(self, mock_empty_cache): + base = CudaSRTPlatform() + base.empty_cache() + mock_empty_cache.assert_called_once_with() + + @patch("torch.cuda.synchronize") + def test_default_synchronize_uses_cuda(self, mock_synchronize): + base = CudaSRTPlatform() + base.synchronize() + mock_synchronize.assert_called_once_with() + + @patch("torch.cuda.mem_get_info", return_value=(123, 456), create=True) + def test_default_get_available_memory_uses_cuda(self, mock_mem_get_info): + base = CudaSRTPlatform() + self.assertEqual(base.get_available_memory(1), (123, 456)) + mock_mem_get_info.assert_called_once_with(1) + + def test_default_distributed_backend_is_nccl(self): + base = CudaSRTPlatform() + self.assertEqual(base.get_torch_distributed_backend_str(), "nccl") + + @patch("torch.cuda.manual_seed_all") + @patch("torch.manual_seed") + @patch("sglang.srt.platforms.device_mixin.np.random.seed") + @patch("sglang.srt.platforms.device_mixin.random.seed") + def test_default_seed_everything_seeds_cuda( + self, mock_random_seed, mock_np_seed, mock_torch_seed, mock_cuda_seed + ): + CudaSRTPlatform.seed_everything(123) + mock_random_seed.assert_called_once_with(123) + mock_np_seed.assert_called_once_with(123) + mock_torch_seed.assert_called_once_with(123) + mock_cuda_seed.assert_called_once_with(123) + + def test_cuda_srt_platform_capabilities(self): + base = CudaSRTPlatform() + self.assertTrue(base.supports_fp8()) + self.assertTrue(base.support_cuda_graph()) + self.assertTrue(base.support_piecewise_cuda_graph()) + class TestSRTPlatformOverrides(CustomTestCase): """Tests for SRTPlatform method overrides via plugins.""" @@ -320,6 +428,16 @@ class TestResolvePlatformWithEnv(CustomTestCase): class TestResolvePlatformAutoDiscover(CustomTestCase): """Tests for _resolve_platform auto-discovery when SGLANG_PLATFORM is not set.""" + @patch("sglang.srt.platforms.torch") + def test_is_cuda_available_excludes_rocm(self, mock_torch): + """ROCm exposes torch.cuda, but should not use the CUDA platform identity.""" + mock_torch.cuda.is_available.return_value = True + mock_torch.version.hip = "6.0" + + import sglang.srt.platforms as plat_mod + + self.assertFalse(plat_mod._is_cuda_available()) + @patch("sglang.srt.platforms.load_plugins_by_group") @patch("sglang.srt.platforms.envs") def test_single_plugin_activates(self, mock_envs, mock_load): @@ -335,13 +453,48 @@ class TestResolvePlatformAutoDiscover(CustomTestCase): self.assertEqual(result, mock_instance) @patch("sglang.srt.platforms.load_plugins_by_group") + @patch("sglang.srt.platforms._is_cuda_available") @patch("sglang.srt.platforms.envs") - def test_no_plugin_activates_fallback(self, mock_envs, mock_load): - """When no plugin activates, return base SRTPlatform with warning.""" + def test_no_plugin_activates_cuda_fallback( + self, mock_envs, mock_is_cuda_available, mock_load + ): + """When CUDA is available and no plugin activates, return CUDA defaults.""" mock_envs.SGLANG_PLATFORM.get.return_value = "" + mock_is_cuda_available.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._is_cuda_available") + @patch("sglang.srt.platforms.envs") + def test_no_plugin_no_cuda_activates_base_fallback( + self, mock_envs, mock_is_cuda_available, mock_load + ): + """When no plugin or CUDA is available, return the abstract base platform.""" + mock_envs.SGLANG_PLATFORM.get.return_value = "" + mock_is_cuda_available.return_value = False mock_load.return_value = {} result = _resolve_platform() self.assertIsInstance(result, SRTPlatform) + self.assertNotIsInstance(result, CudaSRTPlatform) + + @patch("sglang.srt.platforms.load_plugins_by_group") + @patch("sglang.srt.platforms.torch") + @patch("sglang.srt.platforms.envs") + def test_no_plugin_rocm_does_not_activate_cuda_fallback( + self, mock_envs, mock_torch, mock_load + ): + """ROCm exposes torch.cuda but must not use the CUDA fallback platform.""" + mock_envs.SGLANG_PLATFORM.get.return_value = "" + mock_torch.cuda.is_available.return_value = True + mock_torch.version.hip = "6.0" + mock_load.return_value = {} + + result = _resolve_platform() + + self.assertIsInstance(result, SRTPlatform) + self.assertNotIsInstance(result, CudaSRTPlatform) @patch("sglang.srt.platforms.load_plugins_by_group") @patch("sglang.srt.platforms.envs")