Fix torch.__version__ for PEP440 (#15682)
This commit is contained in:
@@ -5,13 +5,13 @@ from typing import TYPE_CHECKING
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.distributed._functional_collectives as ft_c
|
import torch.distributed._functional_collectives as ft_c
|
||||||
from packaging.version import parse
|
|
||||||
from torch.distributed.tensor.experimental._attention import _cp_options
|
from torch.distributed.tensor.experimental._attention import _cp_options
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||||
get_sp_group,
|
get_sp_group,
|
||||||
get_ulysses_parallel_world_size,
|
get_ulysses_parallel_world_size,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.utils.common import torch_release
|
||||||
|
|
||||||
_cp_options.enable_load_balance = False
|
_cp_options.enable_load_balance = False
|
||||||
|
|
||||||
@@ -226,7 +226,7 @@ def ring_attn(
|
|||||||
|
|
||||||
# Starting from torch 2.6.0, _templated_ring_attention expects an integer
|
# Starting from torch 2.6.0, _templated_ring_attention expects an integer
|
||||||
# segment_id for the attention function.
|
# segment_id for the attention function.
|
||||||
use_segment_id = parse(torch.__version__).release >= parse("2.6.0").release
|
use_segment_id = torch_release >= (2, 6)
|
||||||
|
|
||||||
attn_kwargs = dict(
|
attn_kwargs = dict(
|
||||||
op=attn_callable_adapter,
|
op=attn_callable_adapter,
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ import torch.fx as fx
|
|||||||
|
|
||||||
from sglang.srt.compilation.compilation_counter import compilation_counter
|
from sglang.srt.compilation.compilation_counter import compilation_counter
|
||||||
from sglang.srt.compilation.inductor_pass import pass_context
|
from sglang.srt.compilation.inductor_pass import pass_context
|
||||||
|
from sglang.srt.utils.common import torch_release
|
||||||
|
|
||||||
|
|
||||||
class CompilerInterface:
|
class CompilerInterface:
|
||||||
@@ -226,7 +227,7 @@ class InductorAdaptor(CompilerInterface):
|
|||||||
hash_str, file_path = None, None
|
hash_str, file_path = None, None
|
||||||
from torch._inductor.codecache import FxGraphCache, compiled_fx_graph_hash
|
from torch._inductor.codecache import FxGraphCache, compiled_fx_graph_hash
|
||||||
|
|
||||||
if torch.__version__.startswith("2.5"):
|
if torch_release[:2] == (2, 5):
|
||||||
original_load = FxGraphCache.load
|
original_load = FxGraphCache.load
|
||||||
original_load_name = "torch._inductor.codecache.FxGraphCache.load"
|
original_load_name = "torch._inductor.codecache.FxGraphCache.load"
|
||||||
|
|
||||||
@@ -252,7 +253,7 @@ class InductorAdaptor(CompilerInterface):
|
|||||||
hijacked_compile_fx_inner = (
|
hijacked_compile_fx_inner = (
|
||||||
torch._inductor.compile_fx.compile_fx_inner
|
torch._inductor.compile_fx.compile_fx_inner
|
||||||
) # noqa
|
) # noqa
|
||||||
elif torch.__version__ >= "2.6":
|
elif torch_release >= (2, 6):
|
||||||
# function renamed in 2.6
|
# function renamed in 2.6
|
||||||
original_load_name = None
|
original_load_name = None
|
||||||
|
|
||||||
@@ -405,7 +406,7 @@ class InductorAdaptor(CompilerInterface):
|
|||||||
# Dynamo metrics context, see method for more details.
|
# Dynamo metrics context, see method for more details.
|
||||||
exit_stack.enter_context(self.metrics_context())
|
exit_stack.enter_context(self.metrics_context())
|
||||||
|
|
||||||
if torch.__version__.startswith("2.5"):
|
if torch_release[:2] == (2, 5):
|
||||||
inductor_compiled_graph = FxGraphCache._lookup_graph(
|
inductor_compiled_graph = FxGraphCache._lookup_graph(
|
||||||
hash_str, example_inputs, True, False
|
hash_str, example_inputs, True, False
|
||||||
)
|
)
|
||||||
@@ -413,7 +414,7 @@ class InductorAdaptor(CompilerInterface):
|
|||||||
"Inductor cache lookup failed. Please remove"
|
"Inductor cache lookup failed. Please remove"
|
||||||
f"the cache directory and try again." # noqa
|
f"the cache directory and try again." # noqa
|
||||||
)
|
)
|
||||||
elif torch.__version__ >= "2.6":
|
elif torch_release >= (2, 6):
|
||||||
from torch._inductor.output_code import CompiledFxGraphConstantsWithGm
|
from torch._inductor.output_code import CompiledFxGraphConstantsWithGm
|
||||||
|
|
||||||
constants = CompiledFxGraphConstantsWithGm(graph)
|
constants = CompiledFxGraphConstantsWithGm(graph)
|
||||||
|
|||||||
@@ -3,13 +3,13 @@ import tempfile
|
|||||||
from contextlib import nullcontext
|
from contextlib import nullcontext
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from packaging import version
|
|
||||||
from torch.cuda.memory import CUDAPluggableAllocator
|
from torch.cuda.memory import CUDAPluggableAllocator
|
||||||
|
|
||||||
from sglang.srt.distributed.parallel_state import GroupCoordinator
|
from sglang.srt.distributed.parallel_state import GroupCoordinator
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
|
from sglang.srt.utils.common import torch_release
|
||||||
|
|
||||||
after_2_8_0 = version.parse(torch.__version__) >= version.parse("2.8.0")
|
after_2_8_0 = torch_release >= (2, 8)
|
||||||
|
|
||||||
nccl_allocator_source = """
|
nccl_allocator_source = """
|
||||||
|
|
||||||
|
|||||||
@@ -14,6 +14,8 @@ import torch
|
|||||||
import triton
|
import triton
|
||||||
from packaging import version
|
from packaging import version
|
||||||
|
|
||||||
|
from sglang.srt.utils.common import torch_release
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
COMPILER_MODE = os.getenv("FLA_COMPILER_MODE") == "1"
|
COMPILER_MODE = os.getenv("FLA_COMPILER_MODE") == "1"
|
||||||
@@ -204,11 +206,6 @@ def checkpoint(fn):
|
|||||||
return wrapper
|
return wrapper
|
||||||
|
|
||||||
|
|
||||||
@lru_cache(maxsize=None)
|
|
||||||
def check_pytorch_version(version_s: str = "2.4") -> bool:
|
|
||||||
return version.parse(torch.__version__) >= version.parse(version_s)
|
|
||||||
|
|
||||||
|
|
||||||
def _cpu_device_warning():
|
def _cpu_device_warning():
|
||||||
import warnings
|
import warnings
|
||||||
|
|
||||||
@@ -309,7 +306,7 @@ def check_shared_mem(arch: str = "none", tensor_idx: int = 0) -> bool:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
if check_pytorch_version("2.4"):
|
if torch_release >= (2, 4):
|
||||||
device = "cuda" if device == "cpu" else device
|
device = "cuda" if device == "cpu" else device
|
||||||
autocast_custom_fwd = functools.partial(torch.amp.custom_fwd, device_type=device)
|
autocast_custom_fwd = functools.partial(torch.amp.custom_fwd, device_type=device)
|
||||||
autocast_custom_bwd = functools.partial(torch.amp.custom_bwd, device_type=device)
|
autocast_custom_bwd = functools.partial(torch.amp.custom_bwd, device_type=device)
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ from sglang.srt.environ import envs
|
|||||||
from sglang.srt.layers import deep_gemm_wrapper
|
from sglang.srt.layers import deep_gemm_wrapper
|
||||||
from sglang.srt.layers.quantization.fp8_kernel import sglang_per_token_group_quant_fp8
|
from sglang.srt.layers.quantization.fp8_kernel import sglang_per_token_group_quant_fp8
|
||||||
from sglang.srt.layers.quantization.mxfp4_tensor import MXFP4QuantizeUtil
|
from sglang.srt.layers.quantization.mxfp4_tensor import MXFP4QuantizeUtil
|
||||||
|
from sglang.srt.utils.common import torch_release
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
@@ -80,17 +81,12 @@ TORCH_DEVICE_IDENTITY = None
|
|||||||
|
|
||||||
|
|
||||||
def use_rowwise_torch_scaled_mm():
|
def use_rowwise_torch_scaled_mm():
|
||||||
_TORCH_VERSION = torch.__version__.split("+")[0]
|
|
||||||
try:
|
|
||||||
_TORCH_VERSION_TUPLE = tuple(map(int, _TORCH_VERSION.split(".")[:3]))
|
|
||||||
except ValueError:
|
|
||||||
_TORCH_VERSION_TUPLE = (0, 0, 0)
|
|
||||||
if _is_hip:
|
if _is_hip:
|
||||||
# The condition to determine if it is on a platform that supports
|
# The condition to determine if it is on a platform that supports
|
||||||
# torch._scaled_mm rowwise feature.
|
# torch._scaled_mm rowwise feature.
|
||||||
# The condition is determined once as the operations
|
# The condition is determined once as the operations
|
||||||
# are time consuming.
|
# are time consuming.
|
||||||
return get_device_capability() >= (9, 4) and _TORCH_VERSION_TUPLE >= (2, 7, 0)
|
return get_device_capability() >= (9, 4) and torch_release >= (2, 7)
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -62,6 +62,7 @@ from sglang.srt.utils.common import (
|
|||||||
json_list_type,
|
json_list_type,
|
||||||
nullable_str,
|
nullable_str,
|
||||||
parse_connector_type,
|
parse_connector_type,
|
||||||
|
torch_release,
|
||||||
wait_port_available,
|
wait_port_available,
|
||||||
xpu_has_xmx_support,
|
xpu_has_xmx_support,
|
||||||
)
|
)
|
||||||
@@ -4978,10 +4979,7 @@ class ServerArgs:
|
|||||||
# NOTE: CUDA Green Context may encounter potential issues with CudaGraph on torch 2.7.x – 2.8.x, leading to performance degradation.
|
# NOTE: CUDA Green Context may encounter potential issues with CudaGraph on torch 2.7.x – 2.8.x, leading to performance degradation.
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
parts = torch.__version__.split("+", 1)[0].split(".")
|
if torch_release >= (2, 7):
|
||||||
major = int(parts[0]) if len(parts) > 0 and parts[0].isdigit() else 0
|
|
||||||
minor = int(parts[1]) if len(parts) > 1 and parts[1].isdigit() else 0
|
|
||||||
if (major, minor) > (2, 6):
|
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"WARNING: PD-Multiplexing may experience performance degradation with torch versions > 2.6.x.\n"
|
"WARNING: PD-Multiplexing may experience performance degradation with torch versions > 2.6.x.\n"
|
||||||
f" Current torch version is {torch.__version__}.\n"
|
f" Current torch version is {torch.__version__}.\n"
|
||||||
@@ -5049,8 +5047,7 @@ class ServerArgs:
|
|||||||
if self.get_model_config().is_multimodal:
|
if self.get_model_config().is_multimodal:
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
torch_version = torch.__version__.split("+", 1)[0]
|
if torch_release[:3] == (2, 9, 1):
|
||||||
if torch_version == "2.9.1":
|
|
||||||
cudnn_version = None
|
cudnn_version = None
|
||||||
try:
|
try:
|
||||||
cudnn_version = torch.backends.cudnn.version()
|
cudnn_version = torch.backends.cudnn.version()
|
||||||
|
|||||||
@@ -101,6 +101,7 @@ if TYPE_CHECKING:
|
|||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
torch_release = pkg_version.parse(torch.__version__).release
|
||||||
|
|
||||||
|
|
||||||
# https://pytorch.org/docs/stable/notes/hip.html#checking-for-hip
|
# https://pytorch.org/docs/stable/notes/hip.html#checking-for-hip
|
||||||
@@ -1927,7 +1928,7 @@ def init_custom_process_group(
|
|||||||
# https://github.com/pytorch/pytorch/commit/a0c7029a75628cd5fa8df83c0de0ea98ee7fd844
|
# https://github.com/pytorch/pytorch/commit/a0c7029a75628cd5fa8df83c0de0ea98ee7fd844
|
||||||
# We need to determine the appropriate parameter name based on PyTorch version
|
# We need to determine the appropriate parameter name based on PyTorch version
|
||||||
pg_options_param_name = (
|
pg_options_param_name = (
|
||||||
"backend_options" if str(torch.__version__) >= "2.6" else "pg_options"
|
"backend_options" if torch_release >= (2, 6) else "pg_options"
|
||||||
)
|
)
|
||||||
pg, _ = _new_process_group_helper(
|
pg, _ = _new_process_group_helper(
|
||||||
world_size,
|
world_size,
|
||||||
|
|||||||
@@ -14,10 +14,9 @@
|
|||||||
from typing import Callable, Union
|
from typing import Callable, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from packaging import version
|
|
||||||
from torch.multiprocessing import reductions
|
from torch.multiprocessing import reductions
|
||||||
|
|
||||||
from sglang.srt.utils.common import is_npu
|
from sglang.srt.utils.common import is_npu, torch_release
|
||||||
|
|
||||||
_is_npu = is_npu()
|
_is_npu = is_npu()
|
||||||
|
|
||||||
@@ -104,7 +103,7 @@ def _modify_tuple(t, index: int, modifier: Callable):
|
|||||||
|
|
||||||
|
|
||||||
def monkey_patch_torch_compile():
|
def monkey_patch_torch_compile():
|
||||||
if version.parse(torch.__version__) < version.parse("2.8.0"):
|
if torch_release < (2, 8):
|
||||||
# These things are cacheable by torch.compile. torch.compile just doesn't know it.
|
# These things are cacheable by torch.compile. torch.compile just doesn't know it.
|
||||||
# This was fixed in PyTorch 2.8, but until then, we monkey patch.
|
# This was fixed in PyTorch 2.8, but until then, we monkey patch.
|
||||||
import torch._higher_order_ops.auto_functionalize as af
|
import torch._higher_order_ops.auto_functionalize as af
|
||||||
|
|||||||
Reference in New Issue
Block a user