diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 435c30a5c..9b84337e3 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -509,9 +509,6 @@ class Envs: # None = standard attention. See https://arxiv.org/abs/2512.12087 SGLANG_SKIP_SOFTMAX_PREFILL_THRESHOLD_SCALE_FACTOR = EnvFloat(None) SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR = EnvFloat(None) - # TODO(mmangkad): Remove this once the FlashInfer unified allreduce-fusion - # transport issue on GB200/GB300 platforms is fixed and verified resolved. - SGLANG_FLASHINFER_FORCE_POSIX_FD_TRANSPORT = EnvBool(None) # Triton SGLANG_TRITON_DECODE_ATTN_STATIC_KV_SPLITS = EnvBool(False) diff --git a/python/sglang/srt/layers/flashinfer_comm_fusion.py b/python/sglang/srt/layers/flashinfer_comm_fusion.py index d441276ed..8bb26d512 100644 --- a/python/sglang/srt/layers/flashinfer_comm_fusion.py +++ b/python/sglang/srt/layers/flashinfer_comm_fusion.py @@ -1,7 +1,5 @@ -import contextlib import inspect import logging -import platform from typing import Optional, Tuple import torch @@ -18,7 +16,6 @@ from sglang.srt.distributed import ( get_moe_tp_group, get_tp_group, ) -from sglang.srt.environ import envs from sglang.srt.utils import ( ceil_align, get_cuda_driver_bindings, @@ -34,72 +31,6 @@ _flashinfer_allreduce_unavailable = False _flashinfer_create_workspace_supports_group = False _flashinfer_create_workspace_supports_comm_backend = False _flashinfer_allreduce_supports_trigger_completion = False -_posix_transport_override_logged = False - - -def _should_force_posix_fd_transport() -> bool: - force_posix_env = envs.SGLANG_FLASHINFER_FORCE_POSIX_FD_TRANSPORT.get() - if force_posix_env is not None: - return force_posix_env - - machine = platform.machine().lower() - if machine not in ("aarch64", "arm64"): - return False - - if not torch.cuda.is_available(): - return False - - try: - major, _minor = torch.cuda.get_device_capability(torch.cuda.current_device()) - except Exception as e: - logger.debug("Failed to get CUDA device capability: %s", e) - return False - - return major == 10 - - -@contextlib.contextmanager -def _flashinfer_posix_fd_transport_override_if_needed(): - # TODO(mmangkad): Remove this temporary override once the - # FlashInfer unified allreduce-fusion transport issue on - # GB200/GB300 platforms is fixed and verified resolved. - global _posix_transport_override_logged - - if not _should_force_posix_fd_transport(): - yield - return - - try: - import flashinfer.comm.mnnvl as flashinfer_mnnvl - except Exception as e: - logger.debug( - "Failed to import flashinfer.comm.mnnvl for transport override: %s", e - ) - yield - return - - original_checker = getattr(flashinfer_mnnvl, "is_mnnvl_fabric_supported", None) - if original_checker is None: - yield - return - - if not _posix_transport_override_logged: - logger.warning( - "Applying FlashInfer transport workaround: forcing PosixFD " - "symmetric-memory handle exchange on aarch64 + sm10x to avoid " - "known data corruption with Fabric handle exchange on GB systems. " - "Set SGLANG_FLASHINFER_FORCE_POSIX_FD_TRANSPORT=0 to disable." - ) - _posix_transport_override_logged = True - - def _always_disable_fabric(_device_idx: int) -> bool: - return False - - flashinfer_mnnvl.is_mnnvl_fabric_supported = _always_disable_fabric - try: - yield - finally: - flashinfer_mnnvl.is_mnnvl_fabric_supported = original_checker if is_flashinfer_available(): @@ -177,21 +108,13 @@ def is_flashinfer_allreduce_unavailable() -> bool: def _make_flashinfer_workspace_allocation_prop(cuda_driver): - if _should_force_posix_fd_transport(): - handle_type = ( - cuda_driver.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_POSIX_FILE_DESCRIPTOR - ) - else: - from flashinfer.comm.mnnvl import is_mnnvl_fabric_supported + from flashinfer.comm.mnnvl import is_mnnvl_fabric_supported - if is_mnnvl_fabric_supported(torch.cuda.current_device()): - handle_type = ( - cuda_driver.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_FABRIC - ) - else: - handle_type = ( - cuda_driver.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_POSIX_FILE_DESCRIPTOR - ) + handle_types = cuda_driver.CUmemAllocationHandleType + if is_mnnvl_fabric_supported(torch.cuda.current_device()): + handle_type = handle_types.CU_MEM_HANDLE_TYPE_FABRIC + else: + handle_type = handle_types.CU_MEM_HANDLE_TYPE_POSIX_FILE_DESCRIPTOR prop = cuda_driver.CUmemAllocationProp() prop.requestedHandleTypes = handle_type @@ -413,8 +336,7 @@ class FlashInferWorkspaceManager: kwargs["comm_backend"] = _TorchDistBackend( device_group=device_group, cpu_group=cpu_group ) - with _flashinfer_posix_fd_transport_override_if_needed(): - self.workspace = create_workspace(**kwargs) + self.workspace = create_workspace(**kwargs) except Exception as e: _flashinfer_allreduce_unavailable = True logger.warning(