[diffusion] chore: reuse srt siglip vision model (#34988)
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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),
|
||||
)
|
||||
|
||||
@@ -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"
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user