From 0e178c3d2246b0fb373fcbaa84a157f67234f45a Mon Sep 17 00:00:00 2001 From: Mick Date: Mon, 17 Aug 2026 09:16:17 +0800 Subject: [PATCH] [diffusion] chore: reuse srt siglip vision model (#34988) --- .../runtime/distributed/parallel_state.py | 16 + .../runtime/models/encoders/gemma_3.py | 286 ++---------------- .../component_accuracy/engine.py | 5 +- ...est_component_accuracy_parallel_runtime.py | 91 ++++++ .../test/unit/test_srt_siglip_reuse.py | 111 +++++++ python/sglang/srt/layers/attention/vision.py | 13 +- python/sglang/srt/models/siglip.py | 12 +- python/sglang/srt/runtime_context.py | 5 + .../test_vision_backend_selection.py | 33 ++ 9 files changed, 298 insertions(+), 274 deletions(-) create mode 100644 python/sglang/multimodal_gen/test/unit/test_srt_siglip_reuse.py diff --git a/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py b/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py index 2200c28bd..8a56a94f8 100644 --- a/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py +++ b/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py @@ -159,11 +159,15 @@ def _sync_srt_tp_group() -> None: if srt_parallel_state._TP is None: srt_parallel_state._TP = _TP + if srt_parallel_state._ATTN_TP is None: + srt_parallel_state._ATTN_TP = _TP def _clear_srt_tp_group() -> None: import sglang.srt.distributed.parallel_state as srt_parallel_state + if srt_parallel_state._ATTN_TP is _TP: + srt_parallel_state._ATTN_TP = None if srt_parallel_state._TP is _TP: srt_parallel_state._TP = None @@ -607,14 +611,26 @@ def patch_tensor_parallel_group(tp_group: GroupCoordinator): _TP_STATE_PATCHED = True old_tp_group = get_tp_group() + import sglang.srt.distributed.parallel_state as srt_parallel_state + + patch_srt_tp = srt_parallel_state._TP is old_tp_group + patch_srt_attention_tp = srt_parallel_state._ATTN_TP is old_tp_group global _TP _TP = tp_group + if patch_srt_tp: + srt_parallel_state._TP = tp_group + if patch_srt_attention_tp: + srt_parallel_state._ATTN_TP = tp_group try: yield finally: # restore the original state _TP_STATE_PATCHED = False _TP = old_tp_group + if patch_srt_tp and srt_parallel_state._TP is tp_group: + srt_parallel_state._TP = old_tp_group + if patch_srt_attention_tp and srt_parallel_state._ATTN_TP is tp_group: + srt_parallel_state._ATTN_TP = old_tp_group def get_tp_world_size() -> int: diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/gemma_3.py b/python/sglang/multimodal_gen/runtime/models/encoders/gemma_3.py index 7e6f67b70..ffd3a7d87 100644 --- a/python/sglang/multimodal_gen/runtime/models/encoders/gemma_3.py +++ b/python/sglang/multimodal_gen/runtime/models/encoders/gemma_3.py @@ -4,7 +4,7 @@ # Adapted from sglang: python/sglang/srt/models/gemma3_causal.py import logging -from functools import partial +from contextlib import nullcontext from typing import Any, Iterable, Optional, Set, Tuple import torch @@ -12,11 +12,12 @@ from torch import nn from sglang.multimodal_gen.configs.models.encoders.base import BaseEncoderOutput from sglang.multimodal_gen.configs.models.encoders.gemma_3 import Gemma3Config -from sglang.multimodal_gen.runtime.distributed import get_tp_world_size +from sglang.multimodal_gen.runtime.distributed import get_tp_group, get_tp_world_size +from sglang.multimodal_gen.runtime.distributed.parallel_state import ( + patch_tensor_parallel_group, +) from sglang.multimodal_gen.runtime.layers.activation import GeluAndMul -from sglang.multimodal_gen.runtime.layers.attention import LocalAttention from sglang.multimodal_gen.runtime.layers.linear import ( - ColumnParallelLinear, MergedColumnParallelLinear, QKVParallelLinear, RowParallelLinear, @@ -28,6 +29,7 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im LayerwiseOffloadableModuleMixin, ) from sglang.multimodal_gen.runtime.utils.common import add_prefix +from sglang.srt.models.siglip import SiglipVisionModel logger = logging.getLogger(__name__) @@ -440,270 +442,6 @@ class Gemma3TextScaledWordEmbedding(nn.Embedding): return super().forward(input_ids) * self.embed_scale -# --- Siglip Vision Model Implementation --- - - -class QuickGELU(nn.Module): - def forward(self, x: torch.Tensor) -> torch.Tensor: - return x * torch.sigmoid(1.702 * x) - - -class SiglipVisionEmbeddings(nn.Module): - def __init__(self, config): - super().__init__() - self.config = config - self.embed_dim = config.hidden_size - self.image_size = config.image_size - self.patch_size = config.patch_size - - self.patch_embedding = nn.Conv2d( - in_channels=config.num_channels, - out_channels=self.embed_dim, - kernel_size=self.patch_size, - stride=self.patch_size, - padding="valid", - ) - - self.num_patches = (self.image_size // self.patch_size) ** 2 - self.num_positions = self.num_patches - # Use simple Embedding for position embeddings (usually small enough) - self.position_embedding = nn.Embedding(self.num_positions, self.embed_dim) - self.register_buffer( - "position_ids", - torch.arange(self.num_positions).expand((1, -1)), - persistent=False, - ) - - def forward(self, pixel_values: torch.Tensor) -> torch.Tensor: - target_dtype = self.patch_embedding.weight.dtype - patch_embeds = self.patch_embedding( - pixel_values.to(dtype=target_dtype) - ) # shape = [*, width, grid, grid] - embeddings = patch_embeds.flatten(2).transpose(1, 2) - embeddings = embeddings + self.position_embedding(self.position_ids) - - return embeddings - - -class SiglipMLP(nn.Module): - def __init__( - self, - config, - act_layer: type[nn.Module] = QuickGELU, - quant_config: Optional[QuantizationConfig] = None, - prefix: str = "", - ): - super().__init__() - self.fc1 = ColumnParallelLinear( - config.hidden_size, - config.intermediate_size, - quant_config=quant_config, - prefix=add_prefix("fc1", prefix), - ) - self.act = act_layer() - self.fc2 = RowParallelLinear( - config.intermediate_size, - config.hidden_size, - quant_config=quant_config, - prefix=add_prefix("fc2", prefix), - ) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - x_parallel, _ = self.fc1(x) - x_parallel = self.act(x_parallel) - x, _ = self.fc2(x_parallel) - return x - - -class SiglipAttention(nn.Module): - def __init__( - self, - hidden_size: int, - num_heads: int, - quant_config: Optional[QuantizationConfig] = None, - prefix: str = "", - ): - super().__init__() - self.hidden_size = hidden_size - self.num_heads = num_heads - tp_size = get_tp_world_size() - self.head_dim = hidden_size // num_heads - self.num_heads_per_partition = num_heads // tp_size - # Cache the per-rank projection width so forward() does not re-read the - # global TP size (which is not patched to the folding group at run time). - self.embed_dim_per_partition = self.num_heads_per_partition * self.head_dim - self.scaling = self.head_dim**-0.5 - - self.qkv_proj = QKVParallelLinear( - hidden_size=hidden_size, - head_size=self.head_dim, - total_num_heads=num_heads, - total_num_kv_heads=num_heads, - bias=True, - quant_config=quant_config, - prefix=add_prefix("qkv_proj", prefix), - ) - - self.out_proj = RowParallelLinear( - input_size=hidden_size, - output_size=hidden_size, - bias=True, - quant_config=quant_config, - prefix=add_prefix("out_proj", prefix), - ) - - self.attn = LocalAttention( - num_heads=self.num_heads_per_partition, - head_size=self.head_dim, - num_kv_heads=self.num_heads_per_partition, - softmax_scale=self.scaling, - causal=False, # Bidirectional for Vision - ) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - qkv, _ = self.qkv_proj(hidden_states) - q, k, v = qkv.split([self.embed_dim_per_partition] * 3, dim=-1) - - batch_size, seq_len, _ = q.shape - q = q.view(batch_size, seq_len, self.num_heads_per_partition, self.head_dim) - k = k.view(batch_size, seq_len, self.num_heads_per_partition, self.head_dim) - v = v.view(batch_size, seq_len, self.num_heads_per_partition, self.head_dim) - - attn_output = self.attn(q, k, v) - - attn_output = attn_output.reshape( - batch_size, seq_len, self.embed_dim_per_partition - ) - - output, _ = self.out_proj(attn_output) - return output - - -class SiglipEncoderLayer(nn.Module): - def __init__( - self, - config, - act_layer: type[nn.Module] = QuickGELU, - norm_layer: type[nn.Module] = None, - quant_config: Optional[QuantizationConfig] = None, - prefix: str = "", - ) -> None: - super().__init__() - if norm_layer is None: - norm_layer = partial(nn.LayerNorm, eps=config.layer_norm_eps) - self.layer_norm1 = norm_layer(config.hidden_size) - self.layer_norm2 = norm_layer(config.hidden_size) - self.self_attn = SiglipAttention( - hidden_size=config.hidden_size, - num_heads=config.num_attention_heads, - quant_config=quant_config, - prefix=add_prefix("self_attn", prefix), - ) - self.mlp = SiglipMLP( - config, - act_layer=act_layer, - quant_config=quant_config, - prefix=add_prefix("mlp", prefix), - ) - - def forward( - self, - hidden_states: torch.Tensor, - ) -> torch.Tensor: - residual = hidden_states - hidden_states = self.layer_norm1(hidden_states) - hidden_states = self.self_attn(hidden_states) - hidden_states = residual + hidden_states - - residual = hidden_states - hidden_states = self.layer_norm2(hidden_states) - hidden_states = self.mlp(hidden_states) - hidden_states = residual + hidden_states - return hidden_states - - -class SiglipEncoder(nn.Module): - def __init__( - self, - config, - quant_config: Optional[QuantizationConfig] = None, - prefix: str = "", - ) -> None: - super().__init__() - self.config = config - num_hidden_layers = config.num_hidden_layers - norm_layer = partial(nn.LayerNorm, eps=config.layer_norm_eps) - self.layers = nn.ModuleList( - [ - SiglipEncoderLayer( - config=config, - norm_layer=norm_layer, - quant_config=quant_config, - prefix=add_prefix(f"layers.{layer_idx}", prefix), - ) - for layer_idx in range(num_hidden_layers) - ] - ) - - def forward( - self, - inputs_embeds: torch.Tensor, - ) -> torch.Tensor: - hidden_states = inputs_embeds - for encoder_layer in self.layers: - hidden_states = encoder_layer(hidden_states) - return hidden_states - - -class SiglipVisionTransformer(nn.Module): - def __init__( - self, - config, - quant_config: Optional[QuantizationConfig] = None, - prefix: str = "", - ) -> None: - super().__init__() - self.config = config - embed_dim = config.hidden_size - self.embeddings = SiglipVisionEmbeddings(config) - self.encoder = SiglipEncoder( - config=config, - quant_config=quant_config, - prefix=add_prefix("encoder", prefix), - ) - self.post_layernorm = nn.LayerNorm(embed_dim, eps=config.layer_norm_eps) - - @property - def device(self) -> torch.device: - return self.encoder.layers[0].layer_norm1.weight.device - - def forward(self, pixel_values: torch.Tensor) -> torch.Tensor: - hidden_states = self.embeddings(pixel_values.to(self.device)) - last_hidden_state = self.encoder(inputs_embeds=hidden_states) - last_hidden_state = self.post_layernorm(last_hidden_state) - return last_hidden_state - - -class SiglipVisionModel(nn.Module): - def __init__( - self, - config, - quant_config: Optional[QuantizationConfig] = None, - prefix: str = "", - ): - super().__init__() - self.vision_model = SiglipVisionTransformer( - config, quant_config, prefix=add_prefix("vision_model", prefix) - ) - - @property - def device(self) -> torch.device: - return self.vision_model.device - - def forward(self, pixel_values: torch.Tensor): - return self.vision_model(pixel_values) - - class Gemma3MultiModalProjector(nn.Module): """Projector for Gemma3 multimodal.""" @@ -949,9 +687,11 @@ class Gemma3ForConditionalGeneration(nn.Module, LayerwiseOffloadableModuleMixin) param_names_mapping = { r"^(vision_tower\.)(embeddings|encoder|post_layernorm|head)\.": r"\1vision_model.\2.", + r"^(vision_tower\.vision_model\.encoder\.layers\.\d+\.self_attn\.)out_proj\.": r"\1proj.", } reverse_param_names_mapping = { r"^(vision_tower\.)vision_model\.(embeddings|encoder|post_layernorm|head)\.": r"\1\2.", + r"^(vision_tower\.vision_model\.encoder\.layers\.\d+\.self_attn\.)proj\.": r"\1out_proj.", } def __init__( @@ -964,10 +704,12 @@ class Gemma3ForConditionalGeneration(nn.Module, LayerwiseOffloadableModuleMixin) self.config = config self.quant_config = quant_config self.text_config = config.text_config + self._vision_tensor_parallel_group = get_tp_group() # Vision Tower self.vision_tower = SiglipVisionModel( config=config.vision_config, + qkv_backend="sdpa", quant_config=quant_config, prefix=add_prefix("vision_tower", prefix), ) @@ -978,6 +720,11 @@ class Gemma3ForConditionalGeneration(nn.Module, LayerwiseOffloadableModuleMixin) # Text Model self.language_model = Gemma3TextModel(config) + def _vision_parallel_context(self): + if get_tp_group() is self._vision_tensor_parallel_group: + return nullcontext() + return patch_tensor_parallel_group(self._vision_tensor_parallel_group) + def get_placeholder_mask( self, input_ids: torch.LongTensor, @@ -1030,7 +777,8 @@ class Gemma3ForConditionalGeneration(nn.Module, LayerwiseOffloadableModuleMixin) elif pixel_values.dim() != 4: raise ValueError(f"Unexpected pixel_values shape: {pixel_values.shape}") - vision_outputs = self.vision_tower(pixel_values) + with self._vision_parallel_context(): + vision_outputs = self.vision_tower(pixel_values) image_features = self.multi_modal_projector(vision_outputs) image_features = image_features.to( device=inputs_embeds.device, dtype=inputs_embeds.dtype diff --git a/python/sglang/multimodal_gen/test/single_test_file/component_accuracy/engine.py b/python/sglang/multimodal_gen/test/single_test_file/component_accuracy/engine.py index c724dc36a..b500b21d1 100644 --- a/python/sglang/multimodal_gen/test/single_test_file/component_accuracy/engine.py +++ b/python/sglang/multimodal_gen/test/single_test_file/component_accuracy/engine.py @@ -501,7 +501,10 @@ class AccuracyEngine: shard_rank = shard_context.rank if shard_context is not None else rank # TP-sharded params must load via their own weight_loader; the # generic narrow mis-slices fused QKV/gate_up projections. - if shard_world_size > 1 and load_param_with_weight_loader( + needs_weight_loader = ( + shard_world_size > 1 or tensor.shape != src_tensor.shape + ) + if needs_weight_loader and load_param_with_weight_loader( tensor, name, lookup, reverse_mapping ): matched += 1 diff --git a/python/sglang/multimodal_gen/test/unit/test_component_accuracy_parallel_runtime.py b/python/sglang/multimodal_gen/test/unit/test_component_accuracy_parallel_runtime.py index 5ea5b96fd..50aa50e86 100644 --- a/python/sglang/multimodal_gen/test/unit/test_component_accuracy_parallel_runtime.py +++ b/python/sglang/multimodal_gen/test/unit/test_component_accuracy_parallel_runtime.py @@ -2,14 +2,21 @@ from contextlib import ExitStack from types import SimpleNamespace from unittest.mock import call, patch +import torch +from torch import nn + from sglang.multimodal_gen.runtime.distributed import parallel_state from sglang.multimodal_gen.runtime.distributed.device_communicators.ipc_a2a import ( IPC_A2A, ) from sglang.multimodal_gen.runtime.distributed.parallel_groups import PROCESS_GROUP +from sglang.multimodal_gen.test.single_test_file.component_accuracy.engine import ( + AccuracyEngine, +) from sglang.multimodal_gen.test.single_test_file.component_accuracy.utils import ( initialize_parallel_runtime, ) +from sglang.srt.distributed import parallel_state as srt_parallel_state _UTILS = "sglang.multimodal_gen.test.single_test_file.component_accuracy.utils" @@ -113,3 +120,87 @@ def test_destroy_releases_sequence_parallel_subgroups_after_partial_init(): assert destroy_group.call_args_list == [call(ulysses_group), call(ring_group)] assert PROCESS_GROUP.ULYSSES_PG is None assert PROCESS_GROUP.RING_PG is None + + +def test_srt_attention_tp_group_tracks_diffusion_tp_group(): + tp_group = object() + + with ( + patch.object(parallel_state, "_TP", tp_group), + patch.object(srt_parallel_state, "_TP", None), + patch.object(srt_parallel_state, "_ATTN_TP", None), + ): + parallel_state._sync_srt_tp_group() + + assert srt_parallel_state._TP is tp_group + assert srt_parallel_state._ATTN_TP is tp_group + + parallel_state._clear_srt_tp_group() + + assert srt_parallel_state._TP is None + assert srt_parallel_state._ATTN_TP is None + + +def test_srt_owned_groups_are_not_overwritten_or_cleared(): + diffusion_tp_group = object() + srt_tp_group = object() + srt_attention_tp_group = object() + + with ( + patch.object(parallel_state, "_TP", diffusion_tp_group), + patch.object(srt_parallel_state, "_TP", srt_tp_group), + patch.object(srt_parallel_state, "_ATTN_TP", srt_attention_tp_group), + ): + parallel_state._sync_srt_tp_group() + parallel_state._clear_srt_tp_group() + + assert srt_parallel_state._TP is srt_tp_group + assert srt_parallel_state._ATTN_TP is srt_attention_tp_group + + +def test_srt_tp_groups_follow_encoder_folding_context(): + original_tp_group = object() + folding_tp_group = object() + + with ( + patch.object(parallel_state, "_TP", original_tp_group), + patch.object(parallel_state, "_TP_STATE_PATCHED", False), + patch.object(srt_parallel_state, "_TP", original_tp_group), + patch.object(srt_parallel_state, "_ATTN_TP", original_tp_group), + ): + with parallel_state.patch_tensor_parallel_group(folding_tp_group): + assert parallel_state._TP is folding_tp_group + assert srt_parallel_state._TP is folding_tp_group + assert srt_parallel_state._ATTN_TP is folding_tp_group + + assert parallel_state._TP is original_tp_group + assert srt_parallel_state._TP is original_tp_group + assert srt_parallel_state._ATTN_TP is original_tp_group + + +def test_weight_transfer_uses_loader_for_implicit_srt_shard(): + source = nn.Module() + source.weight = nn.Parameter(torch.arange(8, dtype=torch.float32).reshape(4, 2)) + target = nn.Module() + target.weight = nn.Parameter(torch.empty(2, 2)) + + def load_first_shard(param, loaded_weight): + param.data.copy_(loaded_weight[:2]) + + target.weight.weight_loader = load_first_shard + + with patch( + "sglang.multimodal_gen.test.single_test_file.component_accuracy.engine.model_parallel_is_initialized", + return_value=False, + ): + AccuracyEngine.transfer_weights( + source, + target, + min_match_ratio=1.0, + target_device=torch.device("cpu"), + ) + + torch.testing.assert_close( + target.weight, + source.weight[:2].to(dtype=torch.bfloat16), + ) diff --git a/python/sglang/multimodal_gen/test/unit/test_srt_siglip_reuse.py b/python/sglang/multimodal_gen/test/unit/test_srt_siglip_reuse.py new file mode 100644 index 000000000..45d26b720 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_srt_siglip_reuse.py @@ -0,0 +1,111 @@ +from contextlib import contextmanager +from types import SimpleNamespace +from unittest.mock import patch + +from torch import nn + +from sglang.multimodal_gen.runtime.loader.utils import get_param_names_mapping +from sglang.multimodal_gen.runtime.models.encoders import gemma_3 +from sglang.srt.models import siglip + + +def _vision_config(): + return SimpleNamespace( + hidden_size=16, + intermediate_size=32, + num_attention_heads=2, + num_hidden_layers=1, + layer_norm_eps=1e-6, + ) + + +def test_siglip_encoder_propagates_attention_backend(): + with ( + patch.object(siglip, "VisionAttention", return_value=nn.Identity()) as attn, + patch.object(siglip, "SiglipMLP", return_value=nn.Identity()), + ): + siglip.SiglipEncoder( + _vision_config(), + qkv_backend="sdpa", + ) + + assert attn.call_count == 1 + assert attn.call_args.kwargs["qkv_backend"] == "sdpa" + + +def test_gemma3_uses_srt_siglip_with_stable_backend(): + config = SimpleNamespace(vision_config=object(), text_config=object()) + folding_group = object() + + with ( + patch.object(gemma_3, "get_tp_group", return_value=folding_group), + patch.object( + gemma_3, + "SiglipVisionModel", + return_value=nn.Identity(), + ) as vision_model, + patch.object( + gemma_3, + "Gemma3MultiModalProjector", + return_value=nn.Identity(), + ), + patch.object( + gemma_3, + "Gemma3TextModel", + return_value=nn.Identity(), + ), + ): + model = gemma_3.Gemma3ForConditionalGeneration(config) + + vision_model.assert_called_once_with( + config=config.vision_config, + qkv_backend="sdpa", + quant_config=None, + prefix="vision_tower", + ) + assert model._vision_tensor_parallel_group is folding_group + + +def test_gemma3_restores_vision_tensor_parallel_group(): + model = gemma_3.Gemma3ForConditionalGeneration.__new__( + gemma_3.Gemma3ForConditionalGeneration + ) + nn.Module.__init__(model) + folding_group = object() + active_group = object() + model._vision_tensor_parallel_group = folding_group + events = [] + + @contextmanager + def use_group(group): + events.append(("enter", group)) + yield + events.append(("exit", group)) + + with ( + patch.object(gemma_3, "get_tp_group", return_value=active_group), + patch.object( + gemma_3, + "patch_tensor_parallel_group", + side_effect=use_group, + ) as patch_group, + ): + with model._vision_parallel_context(): + events.append(("forward", folding_group)) + + patch_group.assert_called_once_with(folding_group) + assert events == [ + ("enter", folding_group), + ("forward", folding_group), + ("exit", folding_group), + ] + + +def test_gemma3_maps_hf_siglip_projection_name(): + map_name = get_param_names_mapping( + gemma_3.Gemma3ForConditionalGeneration.param_names_mapping + ) + + mapped, _, _ = map_name("vision_tower.encoder.layers.0.self_attn.out_proj.weight") + + assert mapped == "vision_tower.vision_model.encoder.layers.0.self_attn.proj.weight" diff --git a/python/sglang/srt/layers/attention/vision.py b/python/sglang/srt/layers/attention/vision.py index 5008fad4e..11a90a93e 100644 --- a/python/sglang/srt/layers/attention/vision.py +++ b/python/sglang/srt/layers/attention/vision.py @@ -17,7 +17,7 @@ from sglang.kernels.ops.layernorm.norm import ( ) from sglang.srt.environ import envs from sglang.srt.models.utils import apply_qk_norm -from sglang.srt.runtime_context import get_exec, get_mm, get_parallel +from sglang.srt.runtime_context import get_context, get_exec, get_mm, get_parallel from sglang.srt.utils import ( cpu_has_amx_support, get_bool_env_var, @@ -1104,7 +1104,7 @@ class VisionAttention(nn.Module): # Select attention backend via a unified method _passed_backend = qkv_backend qkv_backend = self._determine_attention_backend(_passed_backend) - if get_mm().mm_attention_backend is None and _passed_backend is None: + if _passed_backend is None and get_mm().mm_attention_backend is None: print_info_once(f"Multimodal attention backend not set. Use {qkv_backend}.") print_info_once(f"Using {qkv_backend} as multimodal attention backend.") @@ -1214,7 +1214,14 @@ class VisionAttention(nn.Module): - Ascend NPU: "ascend_attn" - Other platforms: device-specific optimized backend or "sdpa" """ - override_backend = get_mm().mm_attention_backend + try: + override_backend = get_mm().mm_attention_backend + except ValueError: + if passed_backend is None or get_context().is_config_namespace_published( + "mm" + ): + raise + override_backend = None if override_backend is not None: backend = override_backend elif passed_backend is not None: diff --git a/python/sglang/srt/models/siglip.py b/python/sglang/srt/models/siglip.py index bd57dc581..60443b76c 100644 --- a/python/sglang/srt/models/siglip.py +++ b/python/sglang/srt/models/siglip.py @@ -100,6 +100,7 @@ class SiglipEncoderLayer(nn.Module): norm_layer: Type[nn.Module] = None, quant_config: Optional[QuantizationConfig] = None, prefix: str = "", + qkv_backend: Optional[str] = None, ) -> None: super().__init__() if norm_layer is None: @@ -112,6 +113,7 @@ class SiglipEncoderLayer(nn.Module): projection_size=config.hidden_size, use_qkv_parallel=True, flatten_batch=True, + qkv_backend=qkv_backend, quant_config=quant_config, prefix=add_prefix("self_attn", prefix), ) @@ -167,6 +169,7 @@ class SiglipEncoder(nn.Module): config: SiglipVisionConfig, quant_config: Optional[QuantizationConfig] = None, prefix: str = "", + qkv_backend: Optional[str] = None, ) -> None: super().__init__() @@ -179,6 +182,7 @@ class SiglipEncoder(nn.Module): SiglipEncoderLayer( config=config, norm_layer=norm_layer, + qkv_backend=qkv_backend, quant_config=quant_config, prefix=add_prefix(f"layers.{layer_idx}", prefix), ) @@ -215,6 +219,7 @@ class SiglipVisionTransformer(nn.Module): config: SiglipVisionConfig, quant_config: Optional[QuantizationConfig] = None, prefix: str = "", + qkv_backend: Optional[str] = None, ) -> None: super().__init__() @@ -225,6 +230,7 @@ class SiglipVisionTransformer(nn.Module): self.encoder = SiglipEncoder( config=config, + qkv_backend=qkv_backend, quant_config=quant_config, prefix=add_prefix("encoder", prefix), ) @@ -268,10 +274,14 @@ class SiglipVisionModel(nn.Module): config: SiglipVisionConfig, quant_config: Optional[QuantizationConfig] = None, prefix: str = "", + qkv_backend: Optional[str] = None, ): super().__init__() self.vision_model = SiglipVisionTransformer( - config, quant_config, prefix=add_prefix("vision_model", prefix) + config, + qkv_backend=qkv_backend, + quant_config=quant_config, + prefix=add_prefix("vision_model", prefix), ) @property diff --git a/python/sglang/srt/runtime_context.py b/python/sglang/srt/runtime_context.py index 01940988e..a52a3a942 100644 --- a/python/sglang/srt/runtime_context.py +++ b/python/sglang/srt/runtime_context.py @@ -857,6 +857,11 @@ class RuntimeContext: self._check_role_namespace(name) return bags[name] + def is_config_namespace_published(self, name: str) -> bool: + """Return whether a config namespace exists in the current context.""" + bags = self._config_bags + return bags is not None and name in bags + def _check_role_namespace(self, name: str) -> None: # Out of line so the mode gate above stays one dead-branch-prunable # check under dynamo in the default "off" mode (config_bag runs inside diff --git a/test/registered/unit/layers/attention/test_vision_backend_selection.py b/test/registered/unit/layers/attention/test_vision_backend_selection.py index 1b894f548..4ecf3cee7 100644 --- a/test/registered/unit/layers/attention/test_vision_backend_selection.py +++ b/test/registered/unit/layers/attention/test_vision_backend_selection.py @@ -52,6 +52,39 @@ def test_npu_backend_selection_priority( assert backend == expected +def test_explicit_backend_without_published_mm_context(monkeypatch, npu_platform): + monkeypatch.setattr( + vision, + "get_mm", + Mock(side_effect=ValueError("config namespace 'mm' not published")), + ) + monkeypatch.setattr( + vision, + "get_context", + lambda: SimpleNamespace(is_config_namespace_published=lambda namespace: False), + ) + + backend = vision.VisionAttention._determine_attention_backend(None, "sdpa") + + assert backend == "sdpa" + + +def test_explicit_backend_keeps_published_context_errors(monkeypatch, npu_platform): + monkeypatch.setattr( + vision, + "get_mm", + Mock(side_effect=ValueError("mm namespace is not available for this role")), + ) + monkeypatch.setattr( + vision, + "get_context", + lambda: SimpleNamespace(is_config_namespace_published=lambda namespace: True), + ) + + with pytest.raises(ValueError, match="not available for this role"): + vision.VisionAttention._determine_attention_backend(None, "sdpa") + + def test_sdpa_preserves_flattened_batch_layout(): torch.manual_seed(0) bsz, seq_len, num_heads, head_dim = 3, 5, 2, 8