[VLM] Tiny: Unify VLM environment variables (#15572)
Co-authored-by: luoyuan.luo <luoyuan.luo@antgroup.com>
This commit is contained in:
@@ -342,6 +342,7 @@ class Envs:
|
|||||||
SGLANG_RESIZE_RESAMPLE = EnvStr("")
|
SGLANG_RESIZE_RESAMPLE = EnvStr("")
|
||||||
SGLANG_MM_BUFFER_SIZE_MB = EnvInt(0)
|
SGLANG_MM_BUFFER_SIZE_MB = EnvInt(0)
|
||||||
SGLANG_MM_PRECOMPUTE_HASH = EnvBool(False)
|
SGLANG_MM_PRECOMPUTE_HASH = EnvBool(False)
|
||||||
|
SGLANG_VIT_ENABLE_CUDA_GRAPH = EnvBool(False)
|
||||||
|
|
||||||
# VLM Item CUDA IPC Transport
|
# VLM Item CUDA IPC Transport
|
||||||
SGLANG_USE_CUDA_IPC_TRANSPORT = EnvBool(False)
|
SGLANG_USE_CUDA_IPC_TRANSPORT = EnvBool(False)
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ import torch.nn as nn
|
|||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
from einops import rearrange
|
from einops import rearrange
|
||||||
|
|
||||||
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.dp_attention import get_attention_tp_rank, get_attention_tp_size
|
from sglang.srt.layers.dp_attention import get_attention_tp_rank, get_attention_tp_size
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
get_bool_env_var,
|
get_bool_env_var,
|
||||||
@@ -294,7 +295,7 @@ class VisionTritonAttention(nn.Module):
|
|||||||
Returns:
|
Returns:
|
||||||
[b * s, h, head_size]
|
[b * s, h, head_size]
|
||||||
"""
|
"""
|
||||||
if get_bool_env_var("SGLANG_VIT_ENABLE_CUDA_GRAPH"):
|
if envs.SGLANG_VIT_ENABLE_CUDA_GRAPH.get():
|
||||||
if "output_ws" not in kwargs:
|
if "output_ws" not in kwargs:
|
||||||
raise RuntimeError("output_ws should be prepared for cuda-graph mode")
|
raise RuntimeError("output_ws should be prepared for cuda-graph mode")
|
||||||
|
|
||||||
@@ -363,7 +364,7 @@ class VisionFlash3Attention(nn.Module):
|
|||||||
Returns:
|
Returns:
|
||||||
[b * s, h, head_size]
|
[b * s, h, head_size]
|
||||||
"""
|
"""
|
||||||
if get_bool_env_var("SGLANG_VIT_ENABLE_CUDA_GRAPH"):
|
if envs.SGLANG_VIT_ENABLE_CUDA_GRAPH.get():
|
||||||
max_seqlen = cu_seqlens[1]
|
max_seqlen = cu_seqlens[1]
|
||||||
output = flash_attn_varlen_func(
|
output = flash_attn_varlen_func(
|
||||||
q,
|
q,
|
||||||
|
|||||||
@@ -46,6 +46,7 @@ from sglang.srt.distributed import (
|
|||||||
get_tensor_model_parallel_world_size,
|
get_tensor_model_parallel_world_size,
|
||||||
)
|
)
|
||||||
from sglang.srt.distributed.parallel_state import get_pp_group
|
from sglang.srt.distributed.parallel_state import get_pp_group
|
||||||
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.attention.vision import VisionAttention
|
from sglang.srt.layers.attention.vision import VisionAttention
|
||||||
from sglang.srt.layers.layernorm import RMSNorm
|
from sglang.srt.layers.layernorm import RMSNorm
|
||||||
from sglang.srt.layers.linear import (
|
from sglang.srt.layers.linear import (
|
||||||
@@ -74,7 +75,7 @@ from sglang.srt.models.utils import RotaryPosMixin, WeightsMapper, permute_inv
|
|||||||
from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model
|
from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model
|
||||||
from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner
|
from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.utils import add_prefix, get_bool_env_var
|
from sglang.srt.utils import add_prefix
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -389,7 +390,7 @@ class Qwen2_5_VisionTransformer(nn.Module, RotaryPosMixin):
|
|||||||
x: torch.Tensor,
|
x: torch.Tensor,
|
||||||
grid_thw: torch.Tensor,
|
grid_thw: torch.Tensor,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
if get_bool_env_var("SGLANG_VIT_ENABLE_CUDA_GRAPH"):
|
if envs.SGLANG_VIT_ENABLE_CUDA_GRAPH.get():
|
||||||
return self.forward_with_cuda_graph(x, grid_thw)
|
return self.forward_with_cuda_graph(x, grid_thw)
|
||||||
|
|
||||||
# patchify
|
# patchify
|
||||||
|
|||||||
@@ -31,6 +31,7 @@ from sglang.srt.distributed import (
|
|||||||
get_tensor_model_parallel_world_size,
|
get_tensor_model_parallel_world_size,
|
||||||
)
|
)
|
||||||
from sglang.srt.distributed.parallel_state import get_pp_group
|
from sglang.srt.distributed.parallel_state import get_pp_group
|
||||||
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.attention.vision import VisionAttention
|
from sglang.srt.layers.attention.vision import VisionAttention
|
||||||
from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear
|
from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear
|
||||||
from sglang.srt.layers.logits_processor import LogitsProcessor
|
from sglang.srt.layers.logits_processor import LogitsProcessor
|
||||||
@@ -59,7 +60,7 @@ from sglang.srt.models.utils import (
|
|||||||
from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model
|
from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model
|
||||||
from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner
|
from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.utils import add_prefix, get_bool_env_var, get_int_env_var
|
from sglang.srt.utils import add_prefix, get_int_env_var
|
||||||
from sglang.srt.utils.hf_transformers_utils import get_processor
|
from sglang.srt.utils.hf_transformers_utils import get_processor
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -466,7 +467,7 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin):
|
|||||||
x: torch.Tensor,
|
x: torch.Tensor,
|
||||||
grid_thw: torch.Tensor,
|
grid_thw: torch.Tensor,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
if get_bool_env_var("SGLANG_VIT_ENABLE_CUDA_GRAPH"):
|
if envs.SGLANG_VIT_ENABLE_CUDA_GRAPH.get():
|
||||||
return self.forward_with_cuda_graph(x, grid_thw)
|
return self.forward_with_cuda_graph(x, grid_thw)
|
||||||
|
|
||||||
x = x.to(device=self.device, dtype=self.dtype)
|
x = x.to(device=self.device, dtype=self.dtype)
|
||||||
|
|||||||
Reference in New Issue
Block a user