VLM: add Conv2dLayer/Conv3dLayer to fix PyTorch 2.9.1 CuDNN Conv3d (#20282)

Co-authored-by: wili-65535 <wili-65535@users.noreply.github.com>
This commit is contained in:
Yuhao Yang
2026-03-15 19:17:44 +08:00
committed by GitHub
parent f07529b947
commit 1c456a0af5
18 changed files with 704 additions and 90 deletions

View File

@@ -0,0 +1,300 @@
"""
Conv2d/Conv3d layers with unfold+linear optimization for patch embeddings.
When kernel_size == stride, padding == 0, dilation == 1, groups == 1, the conv
is equivalent to unfold + F.linear, which is significantly faster on CUDA and
also avoids the PyTorch 2.9.1 + CuDNN < 9.15 Conv3d bug
(https://github.com/pytorch/pytorch/issues/168167).
"""
import math
from typing import Tuple, Union
import torch
import torch.nn as nn
import torch.nn.functional as F
from sglang.srt.layers.utils.multi_platform import MultiPlatformOp
_VALID_PADDING_STRINGS = {"same", "valid"}
_VALID_PADDING_MODES = {"zeros", "reflect", "replicate", "circular"}
def _tuplify(val, n: int) -> tuple:
if isinstance(val, (list, tuple)):
assert len(val) == n
return tuple(val)
return (val,) * n
def _check_enable_linear(
kernel_size: tuple,
stride: tuple,
padding: tuple,
dilation: tuple,
groups: int,
) -> bool:
"""Check if conv can be replaced with unfold + F.linear."""
return (
kernel_size == stride
and all(p == 0 for p in padding)
and all(d == 1 for d in dilation)
and groups == 1
)
def _reverse_repeat_tuple(t: tuple) -> tuple:
"""(1, 2, 3) -> (3, 3, 2, 2, 1, 1). Used for F.pad with non-zeros padding_mode."""
return tuple(x for x in reversed(t) for _ in range(2))
def _compute_same_padding_for_pad(kernel_size: tuple, dilation: tuple) -> tuple:
"""Compute _reversed_padding_repeated_twice for padding='same'.
This mirrors PyTorch's nn.Conv*d behavior: pre-compute the exact pad
amounts so that F.pad can be called before F.conv*d(padding=0).
"""
pad = []
for k, d in zip(reversed(kernel_size), reversed(dilation)):
total = d * (k - 1)
pad.append(total // 2)
pad.append(total - total // 2)
return tuple(pad)
def _validate_conv_args(
in_channels: int,
out_channels: int,
groups: int,
padding,
padding_mode: str,
stride: tuple,
) -> None:
if in_channels % groups != 0:
raise ValueError(
f"in_channels ({in_channels}) must be divisible by groups ({groups})"
)
if out_channels % groups != 0:
raise ValueError(
f"out_channels ({out_channels}) must be divisible by groups ({groups})"
)
if padding_mode not in _VALID_PADDING_MODES:
raise ValueError(
f"padding_mode must be one of {_VALID_PADDING_MODES}, got '{padding_mode}'"
)
if isinstance(padding, str):
if padding not in _VALID_PADDING_STRINGS:
raise ValueError(
f"padding must be one of {_VALID_PADDING_STRINGS}, got '{padding}'"
)
if padding == "same" and any(s != 1 for s in stride):
raise ValueError("padding='same' is not supported for strided convolutions")
class Conv2dLayer(MultiPlatformOp):
"""Drop-in replacement for nn.Conv2d. Linear optimization disabled by default."""
def __init__(
self,
in_channels: int,
out_channels: int,
kernel_size: Union[int, Tuple[int, int]],
stride: Union[int, Tuple[int, int]] = 1,
padding: Union[int, Tuple[int, int], str] = 0,
dilation: Union[int, Tuple[int, int]] = 1,
groups: int = 1,
bias: bool = True,
padding_mode: str = "zeros",
disable_linear: bool = True,
):
super().__init__()
self.in_channels = in_channels
self.out_channels = out_channels
self.kernel_size = _tuplify(kernel_size, 2)
self.stride = _tuplify(stride, 2)
self.dilation = _tuplify(dilation, 2)
self.groups = groups
self.padding_mode = padding_mode
_validate_conv_args(
in_channels, out_channels, groups, padding, padding_mode, self.stride
)
if isinstance(padding, str):
self.padding = (0, 0) if padding == "valid" else padding
else:
self.padding = _tuplify(padding, 2)
# Pre-compute pad tuple for padding_mode != "zeros" (mirrors nn.Conv2d).
# When padding="same", we need numeric values for F.pad;
# when padding is already numeric, _reverse_repeat_tuple handles it.
if isinstance(self.padding, str):
self._reversed_padding_repeated_twice = _compute_same_padding_for_pad(
self.kernel_size, self.dilation
)
else:
self._reversed_padding_repeated_twice = _reverse_repeat_tuple(self.padding)
padding_tuple = self.padding if isinstance(self.padding, tuple) else (1, 1)
self.enable_linear = not disable_linear and _check_enable_linear(
self.kernel_size, self.stride, padding_tuple, self.dilation, groups
)
self.weight = nn.Parameter(
torch.empty(out_channels, in_channels // groups, *self.kernel_size)
)
if bias:
self.bias = nn.Parameter(torch.empty(out_channels))
else:
self.register_parameter("bias", None)
self._reset_parameters()
def _reset_parameters(self):
nn.init.kaiming_uniform_(self.weight, a=math.sqrt(5))
if self.bias is not None:
fan_in = nn.init._calculate_correct_fan(self.weight, "fan_in")
bound = 1 / math.sqrt(fan_in) if fan_in > 0 else 0
nn.init.uniform_(self.bias, -bound, bound)
def _forward_mulmat(self, x: torch.Tensor) -> torch.Tensor:
K1, K2 = self.kernel_size
x = x.unfold(2, K1, K1).unfold(3, K2, K2)
N, _, Hp, Wp = x.shape[:4]
x = x.permute(0, 2, 3, 1, 4, 5).reshape(N, Hp, Wp, -1)
x = F.linear(x, self.weight.reshape(self.out_channels, -1), self.bias)
return x.permute(0, 3, 1, 2)
def _forward_conv(self, x: torch.Tensor) -> torch.Tensor:
if self.padding_mode != "zeros":
return F.conv2d(
F.pad(x, self._reversed_padding_repeated_twice, mode=self.padding_mode),
self.weight,
self.bias,
self.stride,
(0, 0),
self.dilation,
self.groups,
)
return F.conv2d(
x,
self.weight,
self.bias,
self.stride,
self.padding,
self.dilation,
self.groups,
)
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
if self.enable_linear:
return self._forward_mulmat(x)
return self._forward_conv(x)
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
if self.enable_linear:
return self._forward_mulmat(x)
return self._forward_conv(x)
class Conv3dLayer(MultiPlatformOp):
"""Drop-in replacement for nn.Conv3d with automatic linear optimization."""
def __init__(
self,
in_channels: int,
out_channels: int,
kernel_size: Union[int, Tuple[int, int, int]],
stride: Union[int, Tuple[int, int, int]] = 1,
padding: Union[int, Tuple[int, int, int], str] = 0,
dilation: Union[int, Tuple[int, int, int]] = 1,
groups: int = 1,
bias: bool = True,
padding_mode: str = "zeros",
disable_linear: bool = False,
):
super().__init__()
self.in_channels = in_channels
self.out_channels = out_channels
self.kernel_size = _tuplify(kernel_size, 3)
self.stride = _tuplify(stride, 3)
self.dilation = _tuplify(dilation, 3)
self.groups = groups
self.padding_mode = padding_mode
_validate_conv_args(
in_channels, out_channels, groups, padding, padding_mode, self.stride
)
if isinstance(padding, str):
self.padding = (0, 0, 0) if padding == "valid" else padding
else:
self.padding = _tuplify(padding, 3)
if isinstance(self.padding, str):
self._reversed_padding_repeated_twice = _compute_same_padding_for_pad(
self.kernel_size, self.dilation
)
else:
self._reversed_padding_repeated_twice = _reverse_repeat_tuple(self.padding)
padding_tuple = self.padding if isinstance(self.padding, tuple) else (1, 1, 1)
self.enable_linear = not disable_linear and _check_enable_linear(
self.kernel_size, self.stride, padding_tuple, self.dilation, groups
)
self.weight = nn.Parameter(
torch.empty(out_channels, in_channels // groups, *self.kernel_size)
)
if bias:
self.bias = nn.Parameter(torch.empty(out_channels))
else:
self.register_parameter("bias", None)
self._reset_parameters()
def _reset_parameters(self):
nn.init.kaiming_uniform_(self.weight, a=math.sqrt(5))
if self.bias is not None:
fan_in = nn.init._calculate_correct_fan(self.weight, "fan_in")
bound = 1 / math.sqrt(fan_in) if fan_in > 0 else 0
nn.init.uniform_(self.bias, -bound, bound)
def _forward_mulmat(self, x: torch.Tensor) -> torch.Tensor:
K1, K2, K3 = self.kernel_size
x = x.unfold(2, K1, K1).unfold(3, K2, K2).unfold(4, K3, K3)
N, Dp, Hp, Wp = x.shape[0], x.shape[2], x.shape[3], x.shape[4]
x = x.permute(0, 2, 3, 4, 1, 5, 6, 7).reshape(N, Dp, Hp, Wp, -1)
x = F.linear(x, self.weight.reshape(self.out_channels, -1), self.bias)
return x.permute(0, 4, 1, 2, 3)
def _forward_conv(self, x: torch.Tensor) -> torch.Tensor:
if self.padding_mode != "zeros":
return F.conv3d(
F.pad(x, self._reversed_padding_repeated_twice, mode=self.padding_mode),
self.weight,
self.bias,
self.stride,
(0, 0, 0),
self.dilation,
self.groups,
)
return F.conv3d(
x,
self.weight,
self.bias,
self.stride,
self.padding,
self.dilation,
self.groups,
)
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
if self.enable_linear:
return self._forward_mulmat(x)
return self._forward_conv(x)
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
if self.enable_linear:
return self._forward_mulmat(x)
return self._forward_conv(x)

View File

@@ -11,6 +11,7 @@ from transformers.modeling_attn_mask_utils import _create_4d_causal_attention_ma
from sglang.srt.layers.activation import QuickGELU
from sglang.srt.layers.attention.vision import VisionAttention
from sglang.srt.layers.conv import Conv2dLayer
from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear
from sglang.srt.layers.pooler import EmbeddingPoolerOutput, Pooler, PoolingType
from sglang.srt.layers.quantization.base_config import QuantizationConfig
@@ -32,7 +33,7 @@ class CLIPVisionEmbeddings(nn.Module):
self.class_embedding = nn.Parameter(torch.randn(self.embed_dim))
self.patch_embedding = nn.Conv2d(
self.patch_embedding = Conv2dLayer(
in_channels=config.num_channels,
out_channels=self.embed_dim,
kernel_size=self.patch_size,

View File

@@ -11,6 +11,7 @@ from transformers.modeling_utils import PreTrainedModel
from sglang.srt.configs.dots_vlm import DotsVisionConfig
from sglang.srt.distributed import parallel_state
from sglang.srt.layers.attention.vision import VisionAttention
from sglang.srt.layers.conv import Conv2dLayer
from sglang.srt.layers.quantization import QuantizationConfig
from sglang.srt.utils import add_prefix, is_npu
@@ -113,7 +114,7 @@ class DotsPatchEmbed(nn.Module):
self.temporal_patch_size = config.temporal_patch_size
self.embed_dim = config.embed_dim
self.config = config
self.proj = nn.Conv2d(
self.proj = Conv2dLayer(
config.num_channels,
config.embed_dim,
kernel_size=(config.patch_size, config.patch_size),

View File

@@ -35,6 +35,7 @@ from sglang.srt.distributed.parallel_state import get_pp_group
from sglang.srt.layers.activation import SiluAndMul
from sglang.srt.layers.attention import vision_utils
from sglang.srt.layers.attention.vision import VisionAttention
from sglang.srt.layers.conv import Conv3dLayer
from sglang.srt.layers.layernorm import LayerNorm, RMSNorm
from sglang.srt.layers.linear import (
MergedColumnParallelLinear,
@@ -203,7 +204,7 @@ class Glm4vVisionPatchEmbed(nn.Module):
self.in_channels = in_channels
kernel_size = (temporal_patch_size, patch_size, patch_size)
self.proj = nn.Conv3d(
self.proj = Conv3dLayer(
in_channels,
hidden_size,
kernel_size=kernel_size,
@@ -211,26 +212,17 @@ class Glm4vVisionPatchEmbed(nn.Module):
bias=True,
)
k = self.in_channels * self.temporal_patch_size * self.patch_size**2
self.linear = nn.Linear(
in_features=k,
out_features=self.hidden_size,
bias=True,
dtype=self.proj.weight.dtype,
)
def copy_conv3d_weight_to_linear(self):
# Call this after weight loading
with torch.no_grad():
self.linear.weight.copy_(self.proj.weight.view(self.hidden_size, -1))
self.linear.bias.copy_(self.proj.bias)
del self.proj
def forward(self, x: torch.Tensor) -> torch.Tensor:
# After copy_conv3d_weight_to_linear(), self.linear exists and
# self.proj has been deleted. Input x is already 2-D:
# (num_patches, C * T * P * P)
return self.linear(x)
# Input x is 2-D: (num_patches, C * T * P * P)
# Reshape to 5-D for Conv3dLayer, then flatten back.
x = x.view(
-1,
self.in_channels,
self.temporal_patch_size,
self.patch_size,
self.patch_size,
)
return self.proj(x).view(-1, self.hidden_size)
class Glm4vPatchMerger(nn.Module):
@@ -456,16 +448,10 @@ class Glm4vVisionModel(nn.Module):
@property
def dtype(self) -> torch.dtype:
# After Conv3d to Linear conversion, self.proj is deleted and
# self.linear takes its place.
if hasattr(self.patch_embed, "linear"):
return self.patch_embed.linear.weight.dtype
return self.patch_embed.proj.weight.dtype
@property
def device(self) -> torch.device:
if hasattr(self.patch_embed, "linear"):
return self.patch_embed.linear.weight.device
return self.patch_embed.proj.weight.device
def rot_pos_emb(
@@ -815,7 +801,6 @@ class Glm4vForConditionalGeneration(nn.Module):
self.config, name, loaded_weight
)
weight_loader(param, loaded_weight)
self.visual.patch_embed.copy_conv3d_weight_to_linear()
def get_embed_and_head(self):
return self.model.embed_tokens.weight, self.lm_head.weight

View File

@@ -26,6 +26,7 @@ from transformers import PretrainedConfig
from sglang.srt.layers.activation import get_act_fn
from sglang.srt.layers.attention.vision import VisionAttention
from sglang.srt.layers.conv import Conv2dLayer
from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear
from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.utils import add_prefix, is_npu
@@ -193,7 +194,7 @@ class Idefics2VisionEmbeddings(nn.Module):
self.embed_dim = config.hidden_size
self.image_size = config.image_size
self.patch_size = config.patch_size
self.patch_embedding = nn.Conv2d(
self.patch_embedding = Conv2dLayer(
in_channels=config.num_channels,
out_channels=self.embed_dim,
kernel_size=self.patch_size,

View File

@@ -17,6 +17,7 @@ from sglang.srt.environ import envs
from sglang.srt.layers.activation import get_act_fn
from sglang.srt.layers.attention import vision_utils
from sglang.srt.layers.attention.vision import SingletonCache, VisionAttention
from sglang.srt.layers.conv import Conv2dLayer
from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
from sglang.srt.layers.quantization.base_config import QuantizationConfig
@@ -113,7 +114,7 @@ class InternVisionEmbeddings(nn.Module):
torch.randn(1, 1, self.embed_dim),
)
self.patch_embedding = nn.Conv2d(
self.patch_embedding = Conv2dLayer(
in_channels=3,
out_channels=self.embed_dim,
kernel_size=self.patch_size,

View File

@@ -10,6 +10,7 @@ from transformers import activations
from sglang.srt.configs.kimi_k25 import KimiK25Config, KimiK25VisionConfig
from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation
from sglang.srt.layers.conv import Conv2dLayer
from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.managers.mm_utils import (
MultiModalityDataPaddingPatternMultimodalTokens,
@@ -401,7 +402,7 @@ class MoonVision3dPatchEmbed(nn.Module):
), f"Expected patch_size to be a tuple of 2, got {patch_size}"
self.patch_size = patch_size
self.proj = nn.Conv2d(
self.proj = Conv2dLayer(
in_dim, out_dim, kernel_size=patch_size, stride=patch_size
)

View File

@@ -58,6 +58,7 @@ except ImportError:
flash_attn_varlen_func = None
from sglang.srt.configs import MoonViTConfig
from sglang.srt.layers.conv import Conv2dLayer
from sglang.srt.layers.linear import ReplicatedLinear
from sglang.srt.layers.quantization import QuantizationConfig
from sglang.srt.layers.quantization.modelslim.modelslim import ModelSlimConfig
@@ -250,7 +251,7 @@ class MoonVisionPatchEmbed(nn.Module):
), f"Expected patch_size to be a tuple of 2, got {patch_size}"
self.patch_size = patch_size
self.proj = nn.Conv2d(
self.proj = Conv2dLayer(
in_dim, out_dim, kernel_size=patch_size, stride=patch_size
)

View File

@@ -10,6 +10,7 @@ import torchaudio.functional as F
from transformers import PretrainedConfig
from sglang.srt.layers.attention.vision import VisionAttention
from sglang.srt.layers.conv import Conv2dLayer
from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear
from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.managers.mm_utils import (
@@ -79,7 +80,7 @@ class AudioPatchEmbed(nn.Module):
)
self.num_patches = self.grid_size[0] * self.grid_size[1]
self.flatten = flatten
self.proj = nn.Conv2d(
self.proj = Conv2dLayer(
in_chans,
embed_dim,
kernel_size=self.patch_size,

View File

@@ -26,6 +26,7 @@ from transformers.utils import torch_int
from sglang.srt.layers.activation import get_act_fn
from sglang.srt.layers.attention.vision import VisionAttention
from sglang.srt.layers.conv import Conv2dLayer
from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear
from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.managers.mm_utils import (
@@ -113,7 +114,7 @@ class SiglipVisionEmbeddings(nn.Module):
self.image_size = config.image_size
self.patch_size = config.patch_size
self.patch_embedding = nn.Conv2d(
self.patch_embedding = Conv2dLayer(
in_channels=config.num_channels,
out_channels=self.embed_dim,
kernel_size=self.patch_size,

View File

@@ -35,6 +35,7 @@ from transformers.models.pixtral.modeling_pixtral import (
from sglang.srt.layers.activation import SiluAndMul
from sglang.srt.layers.attention.vision import VisionAttention
from sglang.srt.layers.conv import Conv2dLayer
from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import MergedColumnParallelLinear, RowParallelLinear
from sglang.srt.layers.quantization.base_config import QuantizationConfig
@@ -328,7 +329,7 @@ class VisionTransformer(nn.Module):
def __init__(self, args: VisionEncoderArgs):
super().__init__()
self.args = args
self.patch_conv = nn.Conv2d(
self.patch_conv = Conv2dLayer(
in_channels=args.num_channels,
out_channels=args.hidden_size,
kernel_size=args.patch_size,
@@ -850,7 +851,7 @@ class PixtralHFVisionModel(nn.Module):
self.image_size = config.image_size
self.patch_size = config.patch_size
self.patch_conv = nn.Conv2d(
self.patch_conv = Conv2dLayer(
in_channels=config.num_channels,
out_channels=config.hidden_size,
kernel_size=config.patch_size,

View File

@@ -35,6 +35,7 @@ from transformers.models.qwen2_vl.configuration_qwen2_vl import Qwen2VLVisionCon
from sglang.srt.layers.activation import QuickGELU
from sglang.srt.layers.attention.vision import VisionAttention
from sglang.srt.layers.conv import Conv3dLayer
from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear
from sglang.srt.layers.logits_processor import LogitsProcessor
from sglang.srt.layers.pooler import Pooler, PoolingType
@@ -190,7 +191,7 @@ class Qwen2VisionPatchEmbed(nn.Module):
self.embed_dim = embed_dim
kernel_size = [temporal_patch_size, patch_size, patch_size]
self.proj = nn.Conv3d(
self.proj = Conv3dLayer(
in_chans, embed_dim, kernel_size=kernel_size, stride=kernel_size, bias=False
)

View File

@@ -37,6 +37,7 @@ from sglang.srt.layers.attention.vision import (
FLASHINFER_WORKSPACE_SIZE_BYTES,
VisionAttention,
)
from sglang.srt.layers.conv import Conv3dLayer
from sglang.srt.layers.dp_attention import (
get_attention_tp_rank,
get_attention_tp_size,
@@ -139,7 +140,7 @@ class Qwen3VLVisionPatchEmbed(nn.Module):
self.embed_dim = config.hidden_size
kernel_size = [self.temporal_patch_size, self.patch_size, self.patch_size]
self.proj = nn.Conv3d(
self.proj = Conv3dLayer(
self.in_channels,
self.embed_dim,
kernel_size=kernel_size,

View File

@@ -10,6 +10,7 @@ from transformers import SiglipVisionConfig
from sglang.srt.layers.activation import QuickGELU
from sglang.srt.layers.attention.vision import VisionAttention
from sglang.srt.layers.conv import Conv2dLayer
from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear
from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding
@@ -26,7 +27,7 @@ class SiglipVisionEmbeddings(nn.Module):
self.image_size = config.image_size
self.patch_size = config.patch_size
self.patch_embedding = nn.Conv2d(
self.patch_embedding = Conv2dLayer(
in_channels=config.num_channels,
out_channels=self.embed_dim,
kernel_size=self.patch_size,

View File

@@ -24,6 +24,7 @@ from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation
from sglang.srt.layers.activation import SiluAndMul
from sglang.srt.layers.attention.vision import VisionAttention
from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes
from sglang.srt.layers.conv import Conv2dLayer
from sglang.srt.layers.dp_attention import (
get_attention_tp_rank,
get_attention_tp_size,
@@ -616,7 +617,7 @@ class Step3VisionEmbeddings(nn.Module):
self.class_embedding = nn.Parameter(torch.randn(1, self.embed_dim))
self.patch_embedding = nn.Conv2d(
self.patch_embedding = Conv2dLayer(
in_channels=config.num_channels,
out_channels=self.embed_dim,
kernel_size=self.patch_size,

View File

@@ -13,6 +13,7 @@ from transformers.activations import ACT2FN
from sglang.srt.configs.step3_vl import Step3VLConfig
from sglang.srt.layers.attention.vision import VisionAttention
from sglang.srt.layers.conv import Conv2dLayer
from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear
from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.managers.mm_utils import (
@@ -316,7 +317,7 @@ class PerceptionEncoder(nn.Module):
raise ValueError("use_rope2d must be True")
self.image_size = config.image_size
self.conv1 = nn.Conv2d(
self.conv1 = Conv2dLayer(
in_channels=3,
out_channels=config.width,
kernel_size=config.patch_size,

View File

@@ -5721,11 +5721,6 @@ class ServerArgs:
# Check LoRA
self.check_lora_server_args()
# torch 2.9.1 has compatibility issues with cuDNN 9.14 and below,
# causing extremely slow nn.Conv3d performance.
# TODO(yhyang201): Remove this check when sglang no longer uses torch 2.9.1.
self.check_torch_2_9_1_cudnn_compatibility()
# Check speculative decoding
if self.speculative_algorithm is not None:
assert (
@@ -5834,49 +5829,6 @@ class ServerArgs:
"When enabling two batch overlap, moe_a2a_backend cannot be 'none'."
)
def check_torch_2_9_1_cudnn_compatibility(self):
if get_bool_env_var("SGLANG_DISABLE_CUDNN_CHECK"):
return
if self.get_model_config().is_multimodal:
import torch
if torch_release[:3] == (2, 9, 1):
cudnn_version = None
try:
cudnn_version = torch.backends.cudnn.version()
except Exception:
cudnn_version = None
if cudnn_version is not None:
version_float = float(str(cudnn_version)[:3]) / 100
if version_float < 9.15:
RED = "\033[91m"
BOLD = "\033[1m"
RESET = "\033[0m"
msg = (
f"{RED}{BOLD}"
"CRITICAL WARNING: PyTorch 2.9.1 & CuDNN Compatibility Issue Detected\n"
"--------------------------------------------------------------------------------\n"
f"Current Environment: PyTorch {torch.__version__} | CuDNN {version_float:.2f}\n\n"
"Issue: There is a KNOWN BUG in PyTorch 2.9.1's `nn.Conv3d` implementation\n"
" when used with CuDNN versions older than 9.15. This can cause\n"
" SEVERE PERFORMANCE DEGRADATION and EXCESSIVE MEMORY USAGE.\n\n"
"Reference: https://github.com/pytorch/pytorch/issues/168167\n\n"
"Solution: You MUST upgrade CuDNN to version 9.15+ to ensure correctness.\n\n"
"Run the following command immediately to fix:\n"
" pip install nvidia-cudnn-cu12==9.16.0.29\n\n"
"Or you can disable this check by setting env var SGLANG_DISABLE_CUDNN_CHECK=1\n"
"--------------------------------------------------------------------------------\n"
f"{RESET}"
)
raise RuntimeError(msg)
else:
RED = "\033[91m"
RESET = "\033[0m"
logger.warning(
f"{RED}WARNING: Could not determine CuDNN version for torch==2.9.1. Please ensure CuDNN >= 9.15 to avoid nn.Conv3d bugs.{RESET}"
)
def check_lora_server_args(self):
assert self.max_loras_per_batch > 0, "max_loras_per_batch must be positive"