[diffusion] chore: reuse srt siglip vision model (#34988)

This commit is contained in:
Mick
2026-08-17 09:16:17 +08:00
committed by GitHub
parent 0aa09ab40d
commit 0e178c3d22
9 changed files with 298 additions and 274 deletions
@@ -159,11 +159,15 @@ def _sync_srt_tp_group() -> None:
if srt_parallel_state._TP is None: if srt_parallel_state._TP is None:
srt_parallel_state._TP = _TP 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: def _clear_srt_tp_group() -> None:
import sglang.srt.distributed.parallel_state as srt_parallel_state 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: if srt_parallel_state._TP is _TP:
srt_parallel_state._TP = None srt_parallel_state._TP = None
@@ -607,14 +611,26 @@ def patch_tensor_parallel_group(tp_group: GroupCoordinator):
_TP_STATE_PATCHED = True _TP_STATE_PATCHED = True
old_tp_group = get_tp_group() 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 global _TP
_TP = tp_group _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: try:
yield yield
finally: finally:
# restore the original state # restore the original state
_TP_STATE_PATCHED = False _TP_STATE_PATCHED = False
_TP = old_tp_group _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: def get_tp_world_size() -> int:
@@ -4,7 +4,7 @@
# Adapted from sglang: python/sglang/srt/models/gemma3_causal.py # Adapted from sglang: python/sglang/srt/models/gemma3_causal.py
import logging import logging
from functools import partial from contextlib import nullcontext
from typing import Any, Iterable, Optional, Set, Tuple from typing import Any, Iterable, Optional, Set, Tuple
import torch 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.base import BaseEncoderOutput
from sglang.multimodal_gen.configs.models.encoders.gemma_3 import Gemma3Config 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.activation import GeluAndMul
from sglang.multimodal_gen.runtime.layers.attention import LocalAttention
from sglang.multimodal_gen.runtime.layers.linear import ( from sglang.multimodal_gen.runtime.layers.linear import (
ColumnParallelLinear,
MergedColumnParallelLinear, MergedColumnParallelLinear,
QKVParallelLinear, QKVParallelLinear,
RowParallelLinear, RowParallelLinear,
@@ -28,6 +29,7 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im
LayerwiseOffloadableModuleMixin, LayerwiseOffloadableModuleMixin,
) )
from sglang.multimodal_gen.runtime.utils.common import add_prefix from sglang.multimodal_gen.runtime.utils.common import add_prefix
from sglang.srt.models.siglip import SiglipVisionModel
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -440,270 +442,6 @@ class Gemma3TextScaledWordEmbedding(nn.Embedding):
return super().forward(input_ids) * self.embed_scale 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): class Gemma3MultiModalProjector(nn.Module):
"""Projector for Gemma3 multimodal.""" """Projector for Gemma3 multimodal."""
@@ -949,9 +687,11 @@ class Gemma3ForConditionalGeneration(nn.Module, LayerwiseOffloadableModuleMixin)
param_names_mapping = { param_names_mapping = {
r"^(vision_tower\.)(embeddings|encoder|post_layernorm|head)\.": r"\1vision_model.\2.", 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 = { reverse_param_names_mapping = {
r"^(vision_tower\.)vision_model\.(embeddings|encoder|post_layernorm|head)\.": r"\1\2.", 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__( def __init__(
@@ -964,10 +704,12 @@ class Gemma3ForConditionalGeneration(nn.Module, LayerwiseOffloadableModuleMixin)
self.config = config self.config = config
self.quant_config = quant_config self.quant_config = quant_config
self.text_config = config.text_config self.text_config = config.text_config
self._vision_tensor_parallel_group = get_tp_group()
# Vision Tower # Vision Tower
self.vision_tower = SiglipVisionModel( self.vision_tower = SiglipVisionModel(
config=config.vision_config, config=config.vision_config,
qkv_backend="sdpa",
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("vision_tower", prefix), prefix=add_prefix("vision_tower", prefix),
) )
@@ -978,6 +720,11 @@ class Gemma3ForConditionalGeneration(nn.Module, LayerwiseOffloadableModuleMixin)
# Text Model # Text Model
self.language_model = Gemma3TextModel(config) 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( def get_placeholder_mask(
self, self,
input_ids: torch.LongTensor, input_ids: torch.LongTensor,
@@ -1030,7 +777,8 @@ class Gemma3ForConditionalGeneration(nn.Module, LayerwiseOffloadableModuleMixin)
elif pixel_values.dim() != 4: elif pixel_values.dim() != 4:
raise ValueError(f"Unexpected pixel_values shape: {pixel_values.shape}") 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 = self.multi_modal_projector(vision_outputs)
image_features = image_features.to( image_features = image_features.to(
device=inputs_embeds.device, dtype=inputs_embeds.dtype device=inputs_embeds.device, dtype=inputs_embeds.dtype
@@ -501,7 +501,10 @@ class AccuracyEngine:
shard_rank = shard_context.rank if shard_context is not None else rank shard_rank = shard_context.rank if shard_context is not None else rank
# TP-sharded params must load via their own weight_loader; the # TP-sharded params must load via their own weight_loader; the
# generic narrow mis-slices fused QKV/gate_up projections. # 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 tensor, name, lookup, reverse_mapping
): ):
matched += 1 matched += 1
@@ -2,14 +2,21 @@ from contextlib import ExitStack
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import call, patch 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 import parallel_state
from sglang.multimodal_gen.runtime.distributed.device_communicators.ipc_a2a import ( from sglang.multimodal_gen.runtime.distributed.device_communicators.ipc_a2a import (
IPC_A2A, IPC_A2A,
) )
from sglang.multimodal_gen.runtime.distributed.parallel_groups import PROCESS_GROUP 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 ( from sglang.multimodal_gen.test.single_test_file.component_accuracy.utils import (
initialize_parallel_runtime, 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" _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 destroy_group.call_args_list == [call(ulysses_group), call(ring_group)]
assert PROCESS_GROUP.ULYSSES_PG is None assert PROCESS_GROUP.ULYSSES_PG is None
assert PROCESS_GROUP.RING_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),
)
@@ -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"
+10 -3
View File
@@ -17,7 +17,7 @@ from sglang.kernels.ops.layernorm.norm import (
) )
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.models.utils import apply_qk_norm 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 ( from sglang.srt.utils import (
cpu_has_amx_support, cpu_has_amx_support,
get_bool_env_var, get_bool_env_var,
@@ -1104,7 +1104,7 @@ class VisionAttention(nn.Module):
# Select attention backend via a unified method # Select attention backend via a unified method
_passed_backend = qkv_backend _passed_backend = qkv_backend
qkv_backend = self._determine_attention_backend(_passed_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"Multimodal attention backend not set. Use {qkv_backend}.")
print_info_once(f"Using {qkv_backend} as multimodal attention backend.") print_info_once(f"Using {qkv_backend} as multimodal attention backend.")
@@ -1214,7 +1214,14 @@ class VisionAttention(nn.Module):
- Ascend NPU: "ascend_attn" - Ascend NPU: "ascend_attn"
- Other platforms: device-specific optimized backend or "sdpa" - 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: if override_backend is not None:
backend = override_backend backend = override_backend
elif passed_backend is not None: elif passed_backend is not None:
+11 -1
View File
@@ -100,6 +100,7 @@ class SiglipEncoderLayer(nn.Module):
norm_layer: Type[nn.Module] = None, norm_layer: Type[nn.Module] = None,
quant_config: Optional[QuantizationConfig] = None, quant_config: Optional[QuantizationConfig] = None,
prefix: str = "", prefix: str = "",
qkv_backend: Optional[str] = None,
) -> None: ) -> None:
super().__init__() super().__init__()
if norm_layer is None: if norm_layer is None:
@@ -112,6 +113,7 @@ class SiglipEncoderLayer(nn.Module):
projection_size=config.hidden_size, projection_size=config.hidden_size,
use_qkv_parallel=True, use_qkv_parallel=True,
flatten_batch=True, flatten_batch=True,
qkv_backend=qkv_backend,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("self_attn", prefix), prefix=add_prefix("self_attn", prefix),
) )
@@ -167,6 +169,7 @@ class SiglipEncoder(nn.Module):
config: SiglipVisionConfig, config: SiglipVisionConfig,
quant_config: Optional[QuantizationConfig] = None, quant_config: Optional[QuantizationConfig] = None,
prefix: str = "", prefix: str = "",
qkv_backend: Optional[str] = None,
) -> None: ) -> None:
super().__init__() super().__init__()
@@ -179,6 +182,7 @@ class SiglipEncoder(nn.Module):
SiglipEncoderLayer( SiglipEncoderLayer(
config=config, config=config,
norm_layer=norm_layer, norm_layer=norm_layer,
qkv_backend=qkv_backend,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix(f"layers.{layer_idx}", prefix), prefix=add_prefix(f"layers.{layer_idx}", prefix),
) )
@@ -215,6 +219,7 @@ class SiglipVisionTransformer(nn.Module):
config: SiglipVisionConfig, config: SiglipVisionConfig,
quant_config: Optional[QuantizationConfig] = None, quant_config: Optional[QuantizationConfig] = None,
prefix: str = "", prefix: str = "",
qkv_backend: Optional[str] = None,
) -> None: ) -> None:
super().__init__() super().__init__()
@@ -225,6 +230,7 @@ class SiglipVisionTransformer(nn.Module):
self.encoder = SiglipEncoder( self.encoder = SiglipEncoder(
config=config, config=config,
qkv_backend=qkv_backend,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("encoder", prefix), prefix=add_prefix("encoder", prefix),
) )
@@ -268,10 +274,14 @@ class SiglipVisionModel(nn.Module):
config: SiglipVisionConfig, config: SiglipVisionConfig,
quant_config: Optional[QuantizationConfig] = None, quant_config: Optional[QuantizationConfig] = None,
prefix: str = "", prefix: str = "",
qkv_backend: Optional[str] = None,
): ):
super().__init__() super().__init__()
self.vision_model = SiglipVisionTransformer( 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 @property
+5
View File
@@ -857,6 +857,11 @@ class RuntimeContext:
self._check_role_namespace(name) self._check_role_namespace(name)
return bags[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: def _check_role_namespace(self, name: str) -> None:
# Out of line so the mode gate above stays one dead-branch-prunable # 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 # check under dynamo in the default "off" mode (config_bag runs inside
@@ -52,6 +52,39 @@ def test_npu_backend_selection_priority(
assert backend == expected 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(): def test_sdpa_preserves_flattened_batch_layout():
torch.manual_seed(0) torch.manual_seed(0)
bsz, seq_len, num_heads, head_dim = 3, 5, 2, 8 bsz, seq_len, num_heads, head_dim = 3, 5, 2, 8