[diffusion] refactor: reuse srt qwen vision and text modules (#35006)
This commit is contained in:
@@ -182,10 +182,12 @@ class MiniMaxH3Qwen3VLEncoder(TextEncoder):
|
||||
for name, loaded_weight in weights:
|
||||
if not self.should_materialize_checkpoint_weight(name):
|
||||
continue
|
||||
param = params.get(name)
|
||||
param_name = name.replace(".attn.qkv.", ".attn.qkv_proj.")
|
||||
param = params.get(param_name)
|
||||
if param is None:
|
||||
raise KeyError(
|
||||
f"Unexpected MiniMax H3 Qwen3-VL checkpoint weight: {name}"
|
||||
"Unexpected MiniMax H3 Qwen3-VL checkpoint weight: "
|
||||
f"{name} (mapped to {param_name})"
|
||||
)
|
||||
weight_loader = getattr(param, "weight_loader", default_weight_loader)
|
||||
try:
|
||||
@@ -196,7 +198,7 @@ class MiniMaxH3Qwen3VLEncoder(TextEncoder):
|
||||
f"{name!r}: checkpoint={tuple(loaded_weight.shape)}, "
|
||||
f"parameter={tuple(param.shape)}"
|
||||
) from exc
|
||||
loaded.add(name)
|
||||
loaded.add(param_name)
|
||||
return loaded
|
||||
|
||||
|
||||
|
||||
@@ -5,7 +5,6 @@ from transformers import (
|
||||
DynamicCache,
|
||||
PretrainedConfig,
|
||||
Qwen2_5_VLTextConfig,
|
||||
Qwen2RMSNorm,
|
||||
)
|
||||
from transformers.masking_utils import (
|
||||
create_causal_mask,
|
||||
@@ -17,23 +16,30 @@ from transformers.utils import TransformersKwargs, is_torchdynamo_compiling
|
||||
|
||||
from sglang.multimodal_gen.configs.models.encoders.qwen_image import Qwen2_5VLConfig
|
||||
from sglang.multimodal_gen.runtime.distributed import (
|
||||
get_tp_rank,
|
||||
get_tp_world_size,
|
||||
model_parallel_is_initialized,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.layers.attention import LocalAttention
|
||||
from sglang.multimodal_gen.runtime.layers.linear import (
|
||||
ColumnParallelLinear,
|
||||
MergedColumnParallelLinear,
|
||||
RowParallelLinear,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.layers.quantization import QuantizationConfig
|
||||
from sglang.multimodal_gen.runtime.loader.weight_utils import default_weight_loader
|
||||
from sglang.multimodal_gen.runtime.models.encoders.base import TextEncoder
|
||||
from sglang.multimodal_gen.runtime.models.encoders.qwen2_5vl_vision import (
|
||||
Qwen2_5VLVisionTransformer,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.models.encoders.qwen_vl_rope import (
|
||||
apply_qwen_vl_text_rope,
|
||||
build_qwen_vl_text_rope,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
|
||||
from sglang.multimodal_gen.runtime.utils.common import add_prefix
|
||||
from sglang.srt.layers.layernorm import RMSNorm
|
||||
from sglang.srt.layers.linear import (
|
||||
ColumnParallelLinear,
|
||||
ReplicatedLinear,
|
||||
RowParallelLinear,
|
||||
)
|
||||
from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||
from sglang.srt.models.qwen2_5_vl import Qwen2_5_VLMLP
|
||||
|
||||
# coding=utf-8
|
||||
# Adapted from
|
||||
@@ -70,12 +76,9 @@ except ImportError:
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from transformers.activations import ACT2FN
|
||||
from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import (
|
||||
Qwen2_5_VLCausalLMOutputWithPast,
|
||||
Qwen2_5_VLModelOutputWithPast,
|
||||
Qwen2_5_VLRotaryEmbedding,
|
||||
apply_multimodal_rotary_pos_emb,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -134,6 +137,12 @@ def _tp_world_size() -> int:
|
||||
return get_tp_world_size()
|
||||
|
||||
|
||||
def _tp_rank() -> int:
|
||||
if not model_parallel_is_initialized():
|
||||
return 0
|
||||
return get_tp_rank()
|
||||
|
||||
|
||||
def _linear_output(linear: nn.Module, x: torch.Tensor) -> torch.Tensor:
|
||||
output = linear(x)
|
||||
return output[0] if isinstance(output, tuple) else output
|
||||
@@ -152,8 +161,10 @@ def _make_column_linear(
|
||||
out_features,
|
||||
bias=bias,
|
||||
gather_output=False,
|
||||
tp_size=_tp_world_size(),
|
||||
tp_rank=_tp_rank(),
|
||||
)
|
||||
return nn.Linear(in_features, out_features, bias=bias)
|
||||
return ReplicatedLinear(in_features, out_features, bias=bias)
|
||||
|
||||
|
||||
def _make_row_linear(
|
||||
@@ -168,8 +179,10 @@ def _make_row_linear(
|
||||
in_features,
|
||||
out_features,
|
||||
bias=bias,
|
||||
tp_size=_tp_world_size(),
|
||||
tp_rank=_tp_rank(),
|
||||
)
|
||||
return nn.Linear(in_features, out_features, bias=bias)
|
||||
return ReplicatedLinear(in_features, out_features, bias=bias)
|
||||
|
||||
|
||||
class Qwen2_5_VLAttention(nn.Module):
|
||||
@@ -183,10 +196,9 @@ class Qwen2_5_VLAttention(nn.Module):
|
||||
self.config = config
|
||||
self.layer_idx = layer_idx
|
||||
if layer_idx is None:
|
||||
logger.warn(
|
||||
f"Instantiating {self.__class__.__name__} without passing `layer_idx` is not recommended and will "
|
||||
"to errors during the forward call, if caching is used. Please make sure to provide a `layer_idx` "
|
||||
"when creating this class."
|
||||
logger.warning(
|
||||
"Instantiating %s without layer_idx disables correct cache updates",
|
||||
self.__class__.__name__,
|
||||
)
|
||||
|
||||
self.hidden_size = config.hidden_size
|
||||
@@ -221,7 +233,6 @@ class Qwen2_5_VLAttention(nn.Module):
|
||||
self.num_key_value_groups = self.num_heads // self.num_key_value_heads
|
||||
self.is_causal = True
|
||||
self.attention_dropout = config.attention_dropout
|
||||
self.rope_scaling = config.rope_scaling
|
||||
self.scaling = self.head_dim**-0.5
|
||||
|
||||
self.q_proj = _make_column_linear(
|
||||
@@ -254,7 +265,7 @@ class Qwen2_5_VLAttention(nn.Module):
|
||||
else None
|
||||
)
|
||||
|
||||
self.rotary_emb = Qwen2_5_VLRotaryEmbedding(config=config)
|
||||
self.rotary_emb = build_qwen_vl_text_rope(config)
|
||||
self.attn = LocalAttention(
|
||||
num_heads=self.num_heads,
|
||||
head_size=self.head_dim,
|
||||
@@ -276,9 +287,6 @@ class Qwen2_5_VLAttention(nn.Module):
|
||||
output_attentions: bool = False,
|
||||
use_cache: bool = False,
|
||||
cache_position: Optional[torch.LongTensor] = None,
|
||||
position_embeddings: Optional[
|
||||
tuple[torch.Tensor, torch.Tensor]
|
||||
] = None, # necessary, but kept here for BC
|
||||
**kwargs: Unpack[FlashAttentionKwargs],
|
||||
) -> tuple[torch.Tensor, Optional[torch.Tensor], Optional[tuple[torch.Tensor]]]:
|
||||
bsz, q_len, _ = hidden_states.size()
|
||||
@@ -291,17 +299,15 @@ class Qwen2_5_VLAttention(nn.Module):
|
||||
key_states = key_states.view(bsz, q_len, -1, self.head_dim).transpose(1, 2)
|
||||
value_states = value_states.view(bsz, q_len, -1, self.head_dim).transpose(1, 2)
|
||||
|
||||
cos, sin = position_embeddings
|
||||
query_states, key_states = apply_multimodal_rotary_pos_emb(
|
||||
query_states, key_states, cos, sin, self.rope_scaling["mrope_section"]
|
||||
query_states, key_states = apply_qwen_vl_text_rope(
|
||||
self.rotary_emb,
|
||||
position_ids,
|
||||
query_states,
|
||||
key_states,
|
||||
)
|
||||
|
||||
if past_key_values is not None:
|
||||
cache_kwargs = {
|
||||
"sin": sin,
|
||||
"cos": cos,
|
||||
"cache_position": cache_position,
|
||||
} # Specific to RoPE models
|
||||
cache_kwargs = {"cache_position": cache_position}
|
||||
key_states, value_states = past_key_values.update(
|
||||
key_states, value_states, self.layer_idx, cache_kwargs
|
||||
)
|
||||
@@ -324,38 +330,6 @@ class Qwen2_5_VLAttention(nn.Module):
|
||||
return attn_output
|
||||
|
||||
|
||||
class Qwen2_5_VLTextMLP(nn.Module):
|
||||
def __init__(self, config: Qwen2_5_VLTextConfig):
|
||||
super().__init__()
|
||||
tp_size = _tp_world_size()
|
||||
use_tensor_parallel = tp_size > 1 and config.intermediate_size % tp_size == 0
|
||||
self.gate_proj = _make_column_linear(
|
||||
config.hidden_size,
|
||||
config.intermediate_size,
|
||||
bias=False,
|
||||
use_tensor_parallel=use_tensor_parallel,
|
||||
)
|
||||
self.up_proj = _make_column_linear(
|
||||
config.hidden_size,
|
||||
config.intermediate_size,
|
||||
bias=False,
|
||||
use_tensor_parallel=use_tensor_parallel,
|
||||
)
|
||||
self.down_proj = _make_row_linear(
|
||||
config.intermediate_size,
|
||||
config.hidden_size,
|
||||
bias=False,
|
||||
use_tensor_parallel=use_tensor_parallel,
|
||||
)
|
||||
self.act_fn = ACT2FN[config.hidden_act]
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = self.act_fn(_linear_output(self.gate_proj, x)) * _linear_output(
|
||||
self.up_proj, x
|
||||
)
|
||||
return _linear_output(self.down_proj, x)
|
||||
|
||||
|
||||
class Qwen2_5_VLDecoderLayer(nn.Module):
|
||||
def __init__(self, config: Qwen2_5_VLTextConfig, layer_idx: int):
|
||||
super().__init__()
|
||||
@@ -371,11 +345,26 @@ class Qwen2_5_VLDecoderLayer(nn.Module):
|
||||
)
|
||||
self.self_attn = Qwen2_5_VLAttention(config, layer_idx)
|
||||
|
||||
self.mlp = Qwen2_5_VLTextMLP(config)
|
||||
self.input_layernorm = Qwen2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||
self.post_attention_layernorm = Qwen2RMSNorm(
|
||||
config.hidden_size, eps=config.rms_norm_eps
|
||||
mlp_tp_size = _tp_world_size()
|
||||
if config.intermediate_size % mlp_tp_size != 0:
|
||||
mlp_tp_size = 1
|
||||
self.mlp = Qwen2_5_VLMLP(
|
||||
config.hidden_size,
|
||||
config.intermediate_size,
|
||||
bias=False,
|
||||
hidden_act=config.hidden_act,
|
||||
prefix=f"model.language_model.layers.{layer_idx}.mlp",
|
||||
fuse_gate_up=False,
|
||||
tp_size=mlp_tp_size,
|
||||
tp_rank=_tp_rank() if mlp_tp_size > 1 else 0,
|
||||
)
|
||||
norm_kwargs = dict(
|
||||
eps=config.rms_norm_eps,
|
||||
cast_x_before_out_mul=True,
|
||||
force_native=True,
|
||||
)
|
||||
self.input_layernorm = RMSNorm(config.hidden_size, **norm_kwargs)
|
||||
self.post_attention_layernorm = RMSNorm(config.hidden_size, **norm_kwargs)
|
||||
self.attention_type = config.layer_types[layer_idx]
|
||||
|
||||
def forward(
|
||||
@@ -387,9 +376,6 @@ class Qwen2_5_VLDecoderLayer(nn.Module):
|
||||
output_attentions: Optional[bool] = False,
|
||||
use_cache: Optional[bool] = False,
|
||||
cache_position: Optional[torch.LongTensor] = None,
|
||||
position_embeddings: Optional[
|
||||
tuple[torch.Tensor, torch.Tensor]
|
||||
] = None, # necessary, but kept here for BC
|
||||
**kwargs: Unpack[FlashAttentionKwargs],
|
||||
) -> tuple[
|
||||
torch.FloatTensor, Optional[tuple[torch.FloatTensor, torch.FloatTensor]]
|
||||
@@ -408,9 +394,6 @@ class Qwen2_5_VLDecoderLayer(nn.Module):
|
||||
past_key_values (`Tuple(torch.FloatTensor)`, *optional*): cached past key and value projection states
|
||||
cache_position (`torch.LongTensor` of shape `(sequence_length)`, *optional*):
|
||||
Indices depicting the position of the input sequence tokens in the sequence.
|
||||
position_embeddings (`tuple[torch.FloatTensor, torch.FloatTensor]`, *optional*):
|
||||
Tuple containing the cosine and sine positional embeddings of shape `(batch_size, seq_len, head_dim)`,
|
||||
with `head_dim` being the embedding dimension of each attention head.
|
||||
kwargs (`dict`, *optional*):
|
||||
Arbitrary kwargs to be ignored, used for FSDP and other methods that injects code
|
||||
into the model
|
||||
@@ -429,7 +412,6 @@ class Qwen2_5_VLDecoderLayer(nn.Module):
|
||||
output_attentions=output_attentions,
|
||||
use_cache=use_cache,
|
||||
cache_position=cache_position,
|
||||
position_embeddings=position_embeddings,
|
||||
**kwargs,
|
||||
)
|
||||
hidden_states = residual + hidden_states
|
||||
@@ -443,41 +425,6 @@ class Qwen2_5_VLDecoderLayer(nn.Module):
|
||||
return hidden_states
|
||||
|
||||
|
||||
class Qwen2_5_VLMLP(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_features: int,
|
||||
hidden_features: int = None,
|
||||
bias: bool = True,
|
||||
hidden_act="silu",
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
self.gate_up_proj = MergedColumnParallelLinear(
|
||||
input_size=in_features,
|
||||
output_sizes=[hidden_features] * 2, # [gate_proj, up_proj]
|
||||
bias=bias,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("gate_up_proj", prefix),
|
||||
)
|
||||
self.down_proj = RowParallelLinear(
|
||||
hidden_features,
|
||||
in_features,
|
||||
bias=bias,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("down_proj", prefix),
|
||||
)
|
||||
self.act = ACT2FN[hidden_act]
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
gate_up, _ = self.gate_up_proj(x)
|
||||
gate, up = gate_up.chunk(2, dim=-1)
|
||||
x = self.act(gate) * up
|
||||
x_down, _ = self.down_proj(x)
|
||||
return x_down
|
||||
|
||||
|
||||
class Qwen2_5_VLTextModel(nn.Module):
|
||||
def __init__(self, config: PretrainedConfig):
|
||||
super().__init__()
|
||||
@@ -485,8 +432,11 @@ class Qwen2_5_VLTextModel(nn.Module):
|
||||
self.padding_idx = config.pad_token_id
|
||||
self.vocab_size = config.vocab_size
|
||||
|
||||
self.embed_tokens = nn.Embedding(
|
||||
config.vocab_size, config.hidden_size, self.padding_idx
|
||||
self.embed_tokens = VocabParallelEmbedding(
|
||||
config.vocab_size,
|
||||
config.hidden_size,
|
||||
org_num_embeddings=config.vocab_size,
|
||||
prefix="model.language_model.embed_tokens",
|
||||
)
|
||||
self.layers = nn.ModuleList(
|
||||
[
|
||||
@@ -495,8 +445,12 @@ class Qwen2_5_VLTextModel(nn.Module):
|
||||
]
|
||||
)
|
||||
self._attn_implementation = config._attn_implementation
|
||||
self.norm = Qwen2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||
self.rotary_emb = Qwen2_5_VLRotaryEmbedding(config=config)
|
||||
self.norm = RMSNorm(
|
||||
config.hidden_size,
|
||||
eps=config.rms_norm_eps,
|
||||
cast_x_before_out_mul=True,
|
||||
force_native=True,
|
||||
)
|
||||
self.has_sliding_layers = "sliding_attention" in self.config.layer_types
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
@@ -600,9 +554,6 @@ class Qwen2_5_VLTextModel(nn.Module):
|
||||
|
||||
hidden_states = inputs_embeds
|
||||
|
||||
# create position embeddings to be shared across the decoder layers
|
||||
position_embeddings = self.rotary_emb(hidden_states, position_ids)
|
||||
|
||||
# decoder layers
|
||||
all_hidden_states = () if output_hidden_states else None
|
||||
all_self_attns = () if output_attentions else None
|
||||
@@ -614,12 +565,11 @@ class Qwen2_5_VLTextModel(nn.Module):
|
||||
hidden_states = decoder_layer(
|
||||
hidden_states,
|
||||
attention_mask=causal_mask_mapping[decoder_layer.attention_type],
|
||||
position_ids=text_position_ids,
|
||||
position_ids=position_ids,
|
||||
past_key_values=past_key_values,
|
||||
output_attentions=output_attentions,
|
||||
use_cache=use_cache,
|
||||
cache_position=cache_position,
|
||||
position_embeddings=position_embeddings,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@@ -1426,6 +1376,26 @@ class Qwen2_5_VLForConditionalGeneration(TextEncoder):
|
||||
if not self.enable_image_understanding:
|
||||
continue
|
||||
name = name.replace("visual.", "model.visual.")
|
||||
name = name.replace(".attn.qkv.", ".attn.qkv_proj.")
|
||||
|
||||
loaded_stacked_param = False
|
||||
for weight_name, shard_id in (
|
||||
(".gate_proj.", 0),
|
||||
(".up_proj.", 1),
|
||||
):
|
||||
if weight_name not in name:
|
||||
continue
|
||||
fused_name = name.replace(weight_name, ".gate_up_proj.")
|
||||
if fused_name not in params_dict:
|
||||
continue
|
||||
param = params_dict[fused_name]
|
||||
loaded_weight = loaded_weight.to(param.dtype)
|
||||
param.weight_loader(param, loaded_weight, shard_id)
|
||||
loaded_params.add(fused_name)
|
||||
loaded_stacked_param = True
|
||||
break
|
||||
if loaded_stacked_param:
|
||||
continue
|
||||
try:
|
||||
# Skip loading extra bias for GPTQ models.
|
||||
if name.endswith(".bias") and name not in params_dict:
|
||||
|
||||
@@ -3,285 +3,75 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from sglang.multimodal_gen.runtime.layers.attention.selector import get_attn_backend
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _PackedSequenceMetadata:
|
||||
cu_seqlens: torch.Tensor
|
||||
cu_seqlens_host: tuple[int, ...]
|
||||
max_seqlen: int
|
||||
|
||||
@classmethod
|
||||
def from_cu_seqlens(cls, cu_seqlens: torch.Tensor) -> _PackedSequenceMetadata:
|
||||
bounds = tuple(int(value) for value in cu_seqlens.tolist())
|
||||
return cls(
|
||||
cu_seqlens=cu_seqlens,
|
||||
cu_seqlens_host=bounds,
|
||||
max_seqlen=max(
|
||||
stop - start for start, stop in zip(bounds[:-1], bounds[1:])
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class Qwen2_5VLVisionRMSNorm(nn.Module):
|
||||
def __init__(self, hidden_size: int, eps: float = 1e-6) -> None:
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.ones(hidden_size))
|
||||
self.variance_epsilon = eps
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
input_dtype = hidden_states.dtype
|
||||
hidden_states = hidden_states.float()
|
||||
variance = hidden_states.square().mean(dim=-1, keepdim=True)
|
||||
hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
|
||||
return self.weight * hidden_states.to(input_dtype)
|
||||
|
||||
|
||||
class Qwen2_5VLVisionPatchEmbed(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
patch_size: int,
|
||||
temporal_patch_size: int,
|
||||
in_channels: int,
|
||||
embed_dim: int,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.patch_size = patch_size
|
||||
self.temporal_patch_size = temporal_patch_size
|
||||
self.in_channels = in_channels
|
||||
self.embed_dim = embed_dim
|
||||
kernel_size = (temporal_patch_size, patch_size, patch_size)
|
||||
self.proj = nn.Conv3d(
|
||||
in_channels,
|
||||
embed_dim,
|
||||
kernel_size=kernel_size,
|
||||
stride=kernel_size,
|
||||
bias=False,
|
||||
)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states = hidden_states.view(
|
||||
-1,
|
||||
self.in_channels,
|
||||
self.temporal_patch_size,
|
||||
self.patch_size,
|
||||
self.patch_size,
|
||||
)
|
||||
return self.proj(hidden_states.to(self.proj.weight.dtype)).view(
|
||||
-1, self.embed_dim
|
||||
)
|
||||
|
||||
|
||||
class Qwen2_5VLVisionRotaryEmbedding(nn.Module):
|
||||
def __init__(self, dim: int, theta: float = 10000.0) -> None:
|
||||
super().__init__()
|
||||
inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim))
|
||||
self.register_buffer("inv_freq", inv_freq, persistent=False)
|
||||
|
||||
def forward(self, position_ids: torch.Tensor) -> torch.Tensor:
|
||||
return (position_ids.unsqueeze(-1) * self.inv_freq).flatten(1)
|
||||
|
||||
|
||||
def _rotate_half(hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
first, second = hidden_states.chunk(2, dim=-1)
|
||||
return torch.cat((-second, first), dim=-1)
|
||||
|
||||
|
||||
def _apply_vision_rotary_embedding(
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
cos: torch.Tensor,
|
||||
sin: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
query_dtype = query.dtype
|
||||
key_dtype = key.dtype
|
||||
query = query.float()
|
||||
key = key.float()
|
||||
cos = cos.unsqueeze(-2).float()
|
||||
sin = sin.unsqueeze(-2).float()
|
||||
query = query * cos + _rotate_half(query) * sin
|
||||
key = key * cos + _rotate_half(key) * sin
|
||||
return query.to(query_dtype), key.to(key_dtype)
|
||||
|
||||
|
||||
class Qwen2_5VLVisionAttention(nn.Module):
|
||||
def __init__(self, config: Any, prefix: str) -> None:
|
||||
super().__init__()
|
||||
self.num_heads = config.num_heads
|
||||
self.head_dim = config.hidden_size // config.num_heads
|
||||
self.scaling = self.head_dim**-0.5
|
||||
self.prefix = prefix
|
||||
self.qkv = nn.Linear(config.hidden_size, config.hidden_size * 3, bias=True)
|
||||
self.proj = nn.Linear(config.hidden_size, config.hidden_size)
|
||||
self._attention_impl = None
|
||||
self._initialize_attention(torch.get_default_dtype())
|
||||
|
||||
def _initialize_attention(self, dtype: torch.dtype) -> None:
|
||||
backend = get_attn_backend(self.head_dim, dtype)
|
||||
if backend.supports_packed_varlen():
|
||||
self._attention_impl = backend.get_impl_cls()(
|
||||
num_heads=self.num_heads,
|
||||
head_size=self.head_dim,
|
||||
num_kv_heads=self.num_heads,
|
||||
softmax_scale=self.scaling,
|
||||
causal=False,
|
||||
prefix=self.prefix,
|
||||
)
|
||||
else:
|
||||
logger.warning_once(
|
||||
"Qwen2.5-VL vision attention uses torch SDPA because "
|
||||
f"{backend.get_enum().name.lower()} does not support packed sequences"
|
||||
)
|
||||
|
||||
def _packed_attention(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
cu_seqlens: torch.Tensor,
|
||||
cu_seqlens_host: tuple[int, ...],
|
||||
max_seqlen: int,
|
||||
) -> torch.Tensor:
|
||||
if self._attention_impl is not None:
|
||||
return self._attention_impl.forward_varlen(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
cu_seqlens=cu_seqlens,
|
||||
cu_seqlens_host=cu_seqlens_host,
|
||||
max_seqlen=max_seqlen,
|
||||
)
|
||||
|
||||
output = torch.empty_like(query)
|
||||
for start, stop in zip(cu_seqlens_host[:-1], cu_seqlens_host[1:]):
|
||||
if start == stop:
|
||||
continue
|
||||
query_segment = query[start:stop].transpose(0, 1).unsqueeze(0)
|
||||
key_segment = key[start:stop].transpose(0, 1).unsqueeze(0)
|
||||
value_segment = value[start:stop].transpose(0, 1).unsqueeze(0)
|
||||
segment = F.scaled_dot_product_attention(
|
||||
query_segment,
|
||||
key_segment,
|
||||
value_segment,
|
||||
dropout_p=0.0,
|
||||
is_causal=False,
|
||||
scale=self.scaling,
|
||||
)
|
||||
output[start:stop] = segment.squeeze(0).transpose(0, 1)
|
||||
return output
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
*,
|
||||
cu_seqlens: torch.Tensor,
|
||||
cu_seqlens_host: tuple[int, ...],
|
||||
max_seqlen: int,
|
||||
position_embeddings: tuple[torch.Tensor, torch.Tensor],
|
||||
) -> torch.Tensor:
|
||||
seq_len = hidden_states.shape[0]
|
||||
query, key, value = (
|
||||
self.qkv(hidden_states)
|
||||
.reshape(seq_len, 3, self.num_heads, self.head_dim)
|
||||
.permute(1, 0, 2, 3)
|
||||
.unbind(0)
|
||||
)
|
||||
query, key = _apply_vision_rotary_embedding(query, key, *position_embeddings)
|
||||
output = self._packed_attention(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
cu_seqlens,
|
||||
cu_seqlens_host,
|
||||
max_seqlen,
|
||||
)
|
||||
return self.proj(output.reshape(seq_len, -1).contiguous())
|
||||
|
||||
|
||||
class Qwen2_5VLVisionMLP(nn.Module):
|
||||
def __init__(self, config: Any) -> None:
|
||||
super().__init__()
|
||||
if config.hidden_act != "silu":
|
||||
raise ValueError(
|
||||
f"Unsupported Qwen2.5-VL vision activation: {config.hidden_act}"
|
||||
)
|
||||
self.gate_proj = nn.Linear(
|
||||
config.hidden_size, config.intermediate_size, bias=True
|
||||
)
|
||||
self.up_proj = nn.Linear(
|
||||
config.hidden_size, config.intermediate_size, bias=True
|
||||
)
|
||||
self.down_proj = nn.Linear(
|
||||
config.intermediate_size, config.hidden_size, bias=True
|
||||
)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
return self.down_proj(
|
||||
F.silu(self.gate_proj(hidden_states)) * self.up_proj(hidden_states)
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.models.encoders.qwen_vl_vision import (
|
||||
PackedSequenceMetadata,
|
||||
QwenVLVisionAttention,
|
||||
)
|
||||
from sglang.srt.layers.layernorm import RMSNorm
|
||||
from sglang.srt.models.qwen2_5_vl import (
|
||||
Qwen2_5_VisionPatchEmbed as Qwen2_5VLVisionPatchEmbed,
|
||||
)
|
||||
from sglang.srt.models.qwen2_5_vl import (
|
||||
Qwen2_5_VisionPatchMerger as Qwen2_5VLVisionPatchMerger,
|
||||
)
|
||||
from sglang.srt.models.qwen2_5_vl import (
|
||||
Qwen2_5_VisionRotaryEmbedding as Qwen2_5VLVisionRotaryEmbedding,
|
||||
)
|
||||
from sglang.srt.models.qwen2_5_vl import (
|
||||
Qwen2_5_VLMLP,
|
||||
)
|
||||
|
||||
|
||||
class Qwen2_5VLVisionBlock(nn.Module):
|
||||
def __init__(self, config: Any, layer_idx: int) -> None:
|
||||
super().__init__()
|
||||
self.norm1 = Qwen2_5VLVisionRMSNorm(config.hidden_size)
|
||||
self.norm2 = Qwen2_5VLVisionRMSNorm(config.hidden_size)
|
||||
self.attn = Qwen2_5VLVisionAttention(
|
||||
config, prefix=f"visual.blocks.{layer_idx}.attn"
|
||||
self.norm1 = RMSNorm(
|
||||
config.hidden_size,
|
||||
eps=1e-6,
|
||||
cast_x_before_out_mul=True,
|
||||
force_native=True,
|
||||
)
|
||||
self.norm2 = RMSNorm(
|
||||
config.hidden_size,
|
||||
eps=1e-6,
|
||||
cast_x_before_out_mul=True,
|
||||
force_native=True,
|
||||
)
|
||||
self.attn = QwenVLVisionAttention(
|
||||
config,
|
||||
prefix=f"visual.blocks.{layer_idx}.attn",
|
||||
model_name="Qwen2.5-VL",
|
||||
)
|
||||
self.mlp = Qwen2_5_VLMLP(
|
||||
config.hidden_size,
|
||||
config.intermediate_size,
|
||||
bias=True,
|
||||
hidden_act=config.hidden_act,
|
||||
prefix=f"visual.blocks.{layer_idx}.mlp",
|
||||
fuse_gate_up=False,
|
||||
)
|
||||
self.mlp = Qwen2_5VLVisionMLP(config)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
*,
|
||||
cu_seqlens: torch.Tensor,
|
||||
cu_seqlens_host: tuple[int, ...],
|
||||
max_seqlen: int,
|
||||
metadata: PackedSequenceMetadata,
|
||||
position_embeddings: tuple[torch.Tensor, torch.Tensor],
|
||||
) -> torch.Tensor:
|
||||
hidden_states = hidden_states + self.attn(
|
||||
self.norm1(hidden_states),
|
||||
cu_seqlens=cu_seqlens,
|
||||
cu_seqlens_host=cu_seqlens_host,
|
||||
max_seqlen=max_seqlen,
|
||||
metadata=metadata,
|
||||
position_embeddings=position_embeddings,
|
||||
)
|
||||
return hidden_states + self.mlp(self.norm2(hidden_states))
|
||||
|
||||
|
||||
class Qwen2_5VLVisionPatchMerger(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
output_dim: int,
|
||||
context_dim: int,
|
||||
spatial_merge_size: int,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.hidden_size = context_dim * spatial_merge_size**2
|
||||
self.ln_q = Qwen2_5VLVisionRMSNorm(context_dim)
|
||||
self.mlp = nn.Sequential(
|
||||
nn.Linear(self.hidden_size, self.hidden_size),
|
||||
nn.GELU(),
|
||||
nn.Linear(self.hidden_size, output_dim),
|
||||
)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states = self.ln_q(hidden_states).view(-1, self.hidden_size)
|
||||
return self.mlp(hidden_states)
|
||||
|
||||
|
||||
def _vision_position_ids(
|
||||
grid_thw: torch.Tensor, spatial_merge_size: int
|
||||
) -> torch.Tensor:
|
||||
@@ -383,6 +173,7 @@ class Qwen2_5VLVisionTransformer(nn.Module):
|
||||
temporal_patch_size=config.temporal_patch_size,
|
||||
in_channels=config.in_channels,
|
||||
embed_dim=config.hidden_size,
|
||||
disable_linear=True,
|
||||
)
|
||||
head_dim = config.hidden_size // config.num_heads
|
||||
self.rotary_pos_emb = Qwen2_5VLVisionRotaryEmbedding(head_dim // 2)
|
||||
@@ -390,9 +181,13 @@ class Qwen2_5VLVisionTransformer(nn.Module):
|
||||
Qwen2_5VLVisionBlock(config, layer_idx) for layer_idx in range(config.depth)
|
||||
)
|
||||
self.merger = Qwen2_5VLVisionPatchMerger(
|
||||
output_dim=config.out_hidden_size,
|
||||
dim=config.out_hidden_size,
|
||||
context_dim=config.hidden_size,
|
||||
padded_context_dim=config.hidden_size,
|
||||
spatial_merge_size=config.spatial_merge_size,
|
||||
prefix="visual.merger",
|
||||
cast_x_before_out_mul=True,
|
||||
force_native_norm=True,
|
||||
)
|
||||
|
||||
@property
|
||||
@@ -440,8 +235,8 @@ class Qwen2_5VLVisionTransformer(nn.Module):
|
||||
).cumsum(dim=0, dtype=torch.int32)
|
||||
cu_seqlens = F.pad(cu_seqlens, (1, 0), value=0)
|
||||
|
||||
full_metadata = _PackedSequenceMetadata.from_cu_seqlens(cu_seqlens)
|
||||
window_metadata = _PackedSequenceMetadata.from_cu_seqlens(cu_window_seqlens)
|
||||
full_metadata = PackedSequenceMetadata.from_cu_seqlens(cu_seqlens)
|
||||
window_metadata = PackedSequenceMetadata.from_cu_seqlens(cu_window_seqlens)
|
||||
|
||||
for layer_idx, block in enumerate(self.blocks):
|
||||
metadata = (
|
||||
@@ -451,9 +246,7 @@ class Qwen2_5VLVisionTransformer(nn.Module):
|
||||
)
|
||||
hidden_states = block(
|
||||
hidden_states,
|
||||
cu_seqlens=metadata.cu_seqlens,
|
||||
cu_seqlens_host=metadata.cu_seqlens_host,
|
||||
max_seqlen=metadata.max_seqlen,
|
||||
metadata=metadata,
|
||||
position_embeddings=position_embeddings,
|
||||
)
|
||||
|
||||
|
||||
@@ -7,9 +7,8 @@ from torch import nn
|
||||
from sglang.multimodal_gen.configs.models.encoders import BaseEncoderOutput
|
||||
from sglang.multimodal_gen.configs.models.encoders.qwen3 import Qwen3TextConfig
|
||||
from sglang.multimodal_gen.runtime.distributed import get_tp_world_size
|
||||
from sglang.multimodal_gen.runtime.layers.activation import SiluAndMul
|
||||
from sglang.multimodal_gen.runtime.layers.attention import LocalAttention
|
||||
from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm
|
||||
from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm as MMGenRMSNorm
|
||||
from sglang.multimodal_gen.runtime.layers.linear import (
|
||||
MergedColumnParallelLinear,
|
||||
QKVParallelLinear,
|
||||
@@ -25,6 +24,8 @@ from sglang.multimodal_gen.runtime.loader.weight_utils import (
|
||||
maybe_remap_kv_scale_name,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.models.encoders.base import TextEncoder
|
||||
from sglang.srt.layers.activation import SiluAndMul
|
||||
from sglang.srt.layers.layernorm import RMSNorm
|
||||
|
||||
|
||||
class Qwen3MLP(nn.Module):
|
||||
@@ -131,8 +132,9 @@ class Qwen3Attention(nn.Module):
|
||||
|
||||
# QK-Norm: Key difference from LLaMA
|
||||
rms_norm_eps = getattr(config, "rms_norm_eps", 1e-6)
|
||||
self.q_norm = RMSNorm(self.head_dim, eps=rms_norm_eps)
|
||||
self.k_norm = RMSNorm(self.head_dim, eps=rms_norm_eps)
|
||||
# Keep the small-hidden one-pass kernel used by diffusion QK norm.
|
||||
self.q_norm = MMGenRMSNorm(self.head_dim, eps=rms_norm_eps)
|
||||
self.k_norm = MMGenRMSNorm(self.head_dim, eps=rms_norm_eps)
|
||||
|
||||
# Rotary embeddings
|
||||
self.rotary_emb = get_rope(
|
||||
|
||||
@@ -32,7 +32,12 @@ from sglang.multimodal_gen.runtime.models.encoders.base import TextEncoder
|
||||
from sglang.multimodal_gen.runtime.models.encoders.qwen3vl_vision import (
|
||||
Qwen3VLVisionTransformer,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.models.encoders.qwen_vl_rope import (
|
||||
apply_qwen_vl_text_rope,
|
||||
build_qwen_vl_text_rope,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
|
||||
from sglang.srt.layers.layernorm import RMSNorm
|
||||
|
||||
"""Inference-only Qwen3-VL model compatible with HuggingFace weights."""
|
||||
import logging
|
||||
@@ -57,12 +62,18 @@ from transformers.models.qwen3_vl.configuration_qwen3_vl import (
|
||||
from transformers.models.qwen3_vl.modeling_qwen3_vl import (
|
||||
Qwen3VLCausalLMOutputWithPast,
|
||||
Qwen3VLModelOutputWithPast,
|
||||
Qwen3VLTextRMSNorm,
|
||||
Qwen3VLTextRotaryEmbedding,
|
||||
apply_rotary_pos_emb,
|
||||
)
|
||||
|
||||
|
||||
def _make_text_rms_norm(hidden_size: int, eps: float) -> RMSNorm:
|
||||
return RMSNorm(
|
||||
hidden_size,
|
||||
eps=eps,
|
||||
cast_x_before_out_mul=True,
|
||||
force_native=True,
|
||||
)
|
||||
|
||||
|
||||
class Qwen3VLQuantizedLinear(ReplicatedLinear):
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return super().forward(x)[0]
|
||||
@@ -270,12 +281,9 @@ class Qwen3VLTextAttention(nn.Module):
|
||||
use_tensor_parallel=use_tensor_parallel,
|
||||
prefix=f"{prefix}.o_proj",
|
||||
)
|
||||
self.q_norm = Qwen3VLTextRMSNorm(
|
||||
self.head_dim, eps=config.rms_norm_eps
|
||||
) # unlike olmo, only on the head dim!
|
||||
self.k_norm = Qwen3VLTextRMSNorm(
|
||||
self.head_dim, eps=config.rms_norm_eps
|
||||
) # thus post q_norm does not need reshape
|
||||
self.q_norm = _make_text_rms_norm(self.head_dim, config.rms_norm_eps)
|
||||
self.k_norm = _make_text_rms_norm(self.head_dim, config.rms_norm_eps)
|
||||
self.rotary_emb = build_qwen_vl_text_rope(config, mrope_interleaved=True)
|
||||
|
||||
self.attn = LocalAttention(
|
||||
num_heads=self.num_heads,
|
||||
@@ -292,7 +300,7 @@ class Qwen3VLTextAttention(nn.Module):
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
position_embeddings: tuple[torch.Tensor, torch.Tensor],
|
||||
position_ids: torch.LongTensor,
|
||||
attention_mask: Optional[torch.Tensor],
|
||||
past_key_values: Optional[Cache] = None,
|
||||
cache_position: Optional[torch.LongTensor] = None,
|
||||
@@ -309,14 +317,15 @@ class Qwen3VLTextAttention(nn.Module):
|
||||
).transpose(1, 2)
|
||||
value_states = self.v_proj(hidden_states).view(hidden_shape).transpose(1, 2)
|
||||
|
||||
cos, sin = position_embeddings
|
||||
query_states, key_states = apply_rotary_pos_emb(
|
||||
query_states, key_states, cos, sin
|
||||
query_states, key_states = apply_qwen_vl_text_rope(
|
||||
self.rotary_emb,
|
||||
position_ids,
|
||||
query_states,
|
||||
key_states,
|
||||
)
|
||||
|
||||
if past_key_values is not None:
|
||||
# sin and cos are specific to RoPE models; cache_position needed for the static cache
|
||||
cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position}
|
||||
cache_kwargs = {"cache_position": cache_position}
|
||||
key_states, value_states = past_key_values.update(
|
||||
key_states, value_states, self.layer_idx, cache_kwargs
|
||||
)
|
||||
@@ -432,17 +441,16 @@ class Qwen3VLTextDecoderLayer(nn.Module):
|
||||
use_tensor_parallel=use_tensor_parallel,
|
||||
prefix=f"{prefix}.mlp",
|
||||
)
|
||||
self.input_layernorm = Qwen3VLTextRMSNorm(
|
||||
config.hidden_size, eps=config.rms_norm_eps
|
||||
self.input_layernorm = _make_text_rms_norm(
|
||||
config.hidden_size, config.rms_norm_eps
|
||||
)
|
||||
self.post_attention_layernorm = Qwen3VLTextRMSNorm(
|
||||
config.hidden_size, eps=config.rms_norm_eps
|
||||
self.post_attention_layernorm = _make_text_rms_norm(
|
||||
config.hidden_size, config.rms_norm_eps
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
position_embeddings: tuple[torch.Tensor, torch.Tensor],
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
position_ids: Optional[torch.LongTensor] = None,
|
||||
past_key_values: Optional[Cache] = None,
|
||||
@@ -460,7 +468,6 @@ class Qwen3VLTextDecoderLayer(nn.Module):
|
||||
past_key_values=past_key_values,
|
||||
use_cache=use_cache,
|
||||
cache_position=cache_position,
|
||||
position_embeddings=position_embeddings,
|
||||
**kwargs,
|
||||
)
|
||||
hidden_states = residual + hidden_states
|
||||
@@ -505,8 +512,7 @@ class Qwen3VLTextModel(nn.Module):
|
||||
for layer_idx in range(config.num_hidden_layers)
|
||||
]
|
||||
)
|
||||
self.norm = Qwen3VLTextRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||
self.rotary_emb = Qwen3VLTextRotaryEmbedding(config=config)
|
||||
self.norm = _make_text_rms_norm(config.hidden_size, config.rms_norm_eps)
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
# Initialize weights and apply final processing
|
||||
@@ -582,15 +588,10 @@ class Qwen3VLTextModel(nn.Module):
|
||||
position_ids = position_ids[None, ...].expand(3, position_ids.shape[0], -1)
|
||||
|
||||
if position_ids.ndim == 3 and position_ids.shape[0] == 4:
|
||||
text_position_ids = position_ids[0]
|
||||
position_ids = position_ids[1:]
|
||||
else:
|
||||
text_position_ids = position_ids[0]
|
||||
|
||||
hidden_states = inputs_embeds
|
||||
|
||||
# create position embeddings to be shared across the decoder layers
|
||||
position_embeddings = self.rotary_emb(hidden_states, position_ids)
|
||||
all_hidden_states = () if output_hidden_states else None
|
||||
all_self_attns = () if output_attentions else None
|
||||
# decoder layers
|
||||
@@ -598,11 +599,10 @@ class Qwen3VLTextModel(nn.Module):
|
||||
hidden_states = decoder_layer(
|
||||
hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
position_ids=text_position_ids,
|
||||
position_ids=position_ids,
|
||||
past_key_values=past_key_values,
|
||||
cache_position=cache_position,
|
||||
output_attentions=output_attentions,
|
||||
position_embeddings=position_embeddings,
|
||||
**kwargs,
|
||||
)
|
||||
# hidden_states = layer_outputs
|
||||
@@ -1269,6 +1269,8 @@ class Qwen3VLForConditionalGeneration(TextEncoder):
|
||||
for name, loaded_weight in weights:
|
||||
if "rotary_emb.inv_freq" in name:
|
||||
continue
|
||||
if "visual." in name:
|
||||
name = name.replace(".attn.qkv.", ".attn.qkv_proj.")
|
||||
|
||||
try:
|
||||
param = params_dict[name]
|
||||
|
||||
@@ -10,10 +10,16 @@ import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from sglang.multimodal_gen.runtime.layers.attention.selector import get_attn_backend
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
from sglang.multimodal_gen.runtime.models.encoders.qwen_vl_vision import (
|
||||
PackedSequenceMetadata,
|
||||
QwenVLVisionAttention,
|
||||
)
|
||||
from sglang.srt.models.qwen3_vl import (
|
||||
Qwen3_VisionMLP,
|
||||
Qwen3VLMoeVisionPatchMerger,
|
||||
Qwen3VLVisionPatchEmbed,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@@ -23,57 +29,6 @@ class Qwen3VLVisionOutput:
|
||||
deepstack_features: list[torch.Tensor]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _PackedSequenceMetadata:
|
||||
cu_seqlens: torch.Tensor
|
||||
cu_seqlens_host: tuple[int, ...]
|
||||
max_seqlen: int
|
||||
|
||||
@classmethod
|
||||
def from_cu_seqlens(cls, cu_seqlens: torch.Tensor) -> _PackedSequenceMetadata:
|
||||
bounds = tuple(int(value) for value in cu_seqlens.tolist())
|
||||
return cls(
|
||||
cu_seqlens=cu_seqlens,
|
||||
cu_seqlens_host=bounds,
|
||||
max_seqlen=max(
|
||||
stop - start for start, stop in zip(bounds[:-1], bounds[1:])
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class Qwen3VLVisionPatchEmbed(nn.Module):
|
||||
def __init__(self, config: Any) -> None:
|
||||
super().__init__()
|
||||
self.patch_size = config.patch_size
|
||||
self.temporal_patch_size = config.temporal_patch_size
|
||||
self.in_channels = config.in_channels
|
||||
self.embed_dim = config.hidden_size
|
||||
kernel_size = (
|
||||
config.temporal_patch_size,
|
||||
config.patch_size,
|
||||
config.patch_size,
|
||||
)
|
||||
self.proj = nn.Conv3d(
|
||||
config.in_channels,
|
||||
config.hidden_size,
|
||||
kernel_size=kernel_size,
|
||||
stride=kernel_size,
|
||||
bias=True,
|
||||
)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states = hidden_states.view(
|
||||
-1,
|
||||
self.in_channels,
|
||||
self.temporal_patch_size,
|
||||
self.patch_size,
|
||||
self.patch_size,
|
||||
)
|
||||
return self.proj(hidden_states.to(self.proj.weight.dtype)).view(
|
||||
-1, self.embed_dim
|
||||
)
|
||||
|
||||
|
||||
class Qwen3VLVisionRotaryEmbedding(nn.Module):
|
||||
def __init__(self, dim: int, theta: float = 10000.0) -> None:
|
||||
super().__init__()
|
||||
@@ -89,140 +44,32 @@ class Qwen3VLVisionRotaryEmbedding(nn.Module):
|
||||
return torch.outer(positions, self.inv_freq)
|
||||
|
||||
|
||||
def _rotate_half(hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
first, second = hidden_states.chunk(2, dim=-1)
|
||||
return torch.cat((-second, first), dim=-1)
|
||||
|
||||
|
||||
def _apply_vision_rotary_embedding(
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
cos: torch.Tensor,
|
||||
sin: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
query_dtype = query.dtype
|
||||
key_dtype = key.dtype
|
||||
query = query.float()
|
||||
key = key.float()
|
||||
cos = cos.unsqueeze(-2).float()
|
||||
sin = sin.unsqueeze(-2).float()
|
||||
query = query * cos + _rotate_half(query) * sin
|
||||
key = key * cos + _rotate_half(key) * sin
|
||||
return query.to(query_dtype), key.to(key_dtype)
|
||||
|
||||
|
||||
class Qwen3VLVisionAttention(nn.Module):
|
||||
def __init__(self, config: Any, prefix: str) -> None:
|
||||
super().__init__()
|
||||
self.num_heads = config.num_heads
|
||||
self.head_dim = config.hidden_size // config.num_heads
|
||||
self.scaling = self.head_dim**-0.5
|
||||
self.qkv = nn.Linear(config.hidden_size, config.hidden_size * 3, bias=True)
|
||||
self.proj = nn.Linear(config.hidden_size, config.hidden_size)
|
||||
backend = get_attn_backend(self.head_dim, torch.get_default_dtype())
|
||||
self._attention_impl = None
|
||||
if backend.supports_packed_varlen():
|
||||
self._attention_impl = backend.get_impl_cls()(
|
||||
num_heads=self.num_heads,
|
||||
head_size=self.head_dim,
|
||||
num_kv_heads=self.num_heads,
|
||||
softmax_scale=self.scaling,
|
||||
causal=False,
|
||||
prefix=prefix,
|
||||
)
|
||||
else:
|
||||
logger.warning_once(
|
||||
"Qwen3-VL vision attention uses torch SDPA because "
|
||||
f"{backend.get_enum().name.lower()} does not support packed sequences"
|
||||
)
|
||||
|
||||
def _packed_attention(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
metadata: _PackedSequenceMetadata,
|
||||
) -> torch.Tensor:
|
||||
if self._attention_impl is not None:
|
||||
return self._attention_impl.forward_varlen(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
cu_seqlens=metadata.cu_seqlens,
|
||||
cu_seqlens_host=metadata.cu_seqlens_host,
|
||||
max_seqlen=metadata.max_seqlen,
|
||||
)
|
||||
|
||||
output = torch.empty_like(query)
|
||||
for start, stop in zip(
|
||||
metadata.cu_seqlens_host[:-1], metadata.cu_seqlens_host[1:]
|
||||
):
|
||||
if start == stop:
|
||||
continue
|
||||
segment = F.scaled_dot_product_attention(
|
||||
query[start:stop].transpose(0, 1).unsqueeze(0),
|
||||
key[start:stop].transpose(0, 1).unsqueeze(0),
|
||||
value[start:stop].transpose(0, 1).unsqueeze(0),
|
||||
dropout_p=0.0,
|
||||
is_causal=False,
|
||||
scale=self.scaling,
|
||||
)
|
||||
output[start:stop] = segment.squeeze(0).transpose(0, 1)
|
||||
return output
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
*,
|
||||
metadata: _PackedSequenceMetadata,
|
||||
position_embeddings: tuple[torch.Tensor, torch.Tensor],
|
||||
) -> torch.Tensor:
|
||||
sequence_length = hidden_states.shape[0]
|
||||
query, key, value = (
|
||||
self.qkv(hidden_states)
|
||||
.reshape(sequence_length, 3, self.num_heads, self.head_dim)
|
||||
.permute(1, 0, 2, 3)
|
||||
.unbind(0)
|
||||
)
|
||||
query, key = _apply_vision_rotary_embedding(query, key, *position_embeddings)
|
||||
output = self._packed_attention(query, key, value, metadata)
|
||||
return self.proj(output.reshape(sequence_length, -1).contiguous())
|
||||
|
||||
|
||||
class Qwen3VLVisionMLP(nn.Module):
|
||||
def __init__(self, config: Any) -> None:
|
||||
super().__init__()
|
||||
if config.hidden_act != "gelu_pytorch_tanh":
|
||||
raise ValueError(
|
||||
f"Unsupported Qwen3-VL vision activation: {config.hidden_act}"
|
||||
)
|
||||
self.linear_fc1 = nn.Linear(
|
||||
config.hidden_size, config.intermediate_size, bias=True
|
||||
)
|
||||
self.linear_fc2 = nn.Linear(
|
||||
config.intermediate_size, config.hidden_size, bias=True
|
||||
)
|
||||
self.act_fn = nn.GELU(approximate="tanh")
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
return self.linear_fc2(self.act_fn(self.linear_fc1(hidden_states)))
|
||||
|
||||
|
||||
class Qwen3VLVisionBlock(nn.Module):
|
||||
def __init__(self, config: Any, layer_idx: int) -> None:
|
||||
super().__init__()
|
||||
parallel = get_parallel()
|
||||
self.norm1 = nn.LayerNorm(config.hidden_size, eps=1e-6)
|
||||
self.norm2 = nn.LayerNorm(config.hidden_size, eps=1e-6)
|
||||
self.attn = Qwen3VLVisionAttention(
|
||||
config, prefix=f"visual.blocks.{layer_idx}.attn"
|
||||
self.attn = QwenVLVisionAttention(
|
||||
config,
|
||||
prefix=f"visual.blocks.{layer_idx}.attn",
|
||||
model_name="Qwen3-VL",
|
||||
)
|
||||
self.mlp = Qwen3_VisionMLP(
|
||||
config.hidden_size,
|
||||
config.intermediate_size,
|
||||
bias=True,
|
||||
hidden_act=config.hidden_act,
|
||||
prefix=f"visual.blocks.{layer_idx}.mlp",
|
||||
tp_rank=parallel.tp_rank,
|
||||
tp_size=parallel.tp_size,
|
||||
)
|
||||
self.mlp = Qwen3VLVisionMLP(config)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
*,
|
||||
metadata: _PackedSequenceMetadata,
|
||||
metadata: PackedSequenceMetadata,
|
||||
position_embeddings: tuple[torch.Tensor, torch.Tensor],
|
||||
) -> torch.Tensor:
|
||||
hidden_states = hidden_states + self.attn(
|
||||
@@ -233,24 +80,6 @@ class Qwen3VLVisionBlock(nn.Module):
|
||||
return hidden_states + self.mlp(self.norm2(hidden_states))
|
||||
|
||||
|
||||
class Qwen3VLVisionPatchMerger(nn.Module):
|
||||
def __init__(self, config: Any, *, use_postshuffle_norm: bool) -> None:
|
||||
super().__init__()
|
||||
self.hidden_size = config.hidden_size * config.spatial_merge_size**2
|
||||
self.use_postshuffle_norm = use_postshuffle_norm
|
||||
norm_size = self.hidden_size if use_postshuffle_norm else config.hidden_size
|
||||
self.norm = nn.LayerNorm(norm_size, eps=1e-6)
|
||||
self.linear_fc1 = nn.Linear(self.hidden_size, self.hidden_size)
|
||||
self.act_fn = nn.GELU()
|
||||
self.linear_fc2 = nn.Linear(self.hidden_size, config.out_hidden_size)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
if self.use_postshuffle_norm:
|
||||
hidden_states = hidden_states.view(-1, self.hidden_size)
|
||||
hidden_states = self.norm(hidden_states).view(-1, self.hidden_size)
|
||||
return self.linear_fc2(self.act_fn(self.linear_fc1(hidden_states)))
|
||||
|
||||
|
||||
def _vision_position_ids(
|
||||
grid_thw: torch.Tensor, spatial_merge_size: int
|
||||
) -> torch.Tensor:
|
||||
@@ -349,11 +178,12 @@ def _vision_cu_seqlens(grid_thw: torch.Tensor) -> torch.Tensor:
|
||||
class Qwen3VLVisionTransformer(nn.Module):
|
||||
def __init__(self, config: Any) -> None:
|
||||
super().__init__()
|
||||
parallel = get_parallel()
|
||||
self.config = config
|
||||
self.spatial_merge_size = config.spatial_merge_size
|
||||
self.spatial_merge_unit = config.spatial_merge_size**2
|
||||
self.patch_size = config.patch_size
|
||||
self.patch_embed = Qwen3VLVisionPatchEmbed(config)
|
||||
self.patch_embed = Qwen3VLVisionPatchEmbed(config, disable_linear=True)
|
||||
self.pos_embed = nn.Embedding(
|
||||
config.num_position_embeddings, config.hidden_size
|
||||
)
|
||||
@@ -363,11 +193,29 @@ class Qwen3VLVisionTransformer(nn.Module):
|
||||
self.blocks = nn.ModuleList(
|
||||
Qwen3VLVisionBlock(config, layer_idx) for layer_idx in range(config.depth)
|
||||
)
|
||||
self.merger = Qwen3VLVisionPatchMerger(config, use_postshuffle_norm=False)
|
||||
self.merger = Qwen3VLMoeVisionPatchMerger(
|
||||
dim=config.out_hidden_size,
|
||||
context_dim=config.hidden_size,
|
||||
padded_context_dim=config.hidden_size,
|
||||
spatial_merge_size=config.spatial_merge_size,
|
||||
use_postshuffle_norm=False,
|
||||
prefix="visual.merger",
|
||||
tp_rank=parallel.tp_rank,
|
||||
tp_size=parallel.tp_size,
|
||||
)
|
||||
self.deepstack_visual_indexes = tuple(config.deepstack_visual_indexes)
|
||||
self.deepstack_merger_list = nn.ModuleList(
|
||||
Qwen3VLVisionPatchMerger(config, use_postshuffle_norm=True)
|
||||
for _ in self.deepstack_visual_indexes
|
||||
Qwen3VLMoeVisionPatchMerger(
|
||||
dim=config.out_hidden_size,
|
||||
context_dim=config.hidden_size,
|
||||
padded_context_dim=config.hidden_size,
|
||||
spatial_merge_size=config.spatial_merge_size,
|
||||
use_postshuffle_norm=True,
|
||||
prefix=f"visual.deepstack_merger_list.{merger_idx}",
|
||||
tp_rank=parallel.tp_rank,
|
||||
tp_size=parallel.tp_size,
|
||||
)
|
||||
for merger_idx, _ in enumerate(self.deepstack_visual_indexes)
|
||||
)
|
||||
self._deepstack_merger_by_layer = {
|
||||
layer_idx: merger_idx
|
||||
@@ -407,7 +255,7 @@ class Qwen3VLVisionTransformer(nn.Module):
|
||||
rotary = rotary.flatten(1)
|
||||
rotary = torch.cat((rotary, rotary), dim=-1)
|
||||
position_embeddings = (rotary.cos(), rotary.sin())
|
||||
metadata = _PackedSequenceMetadata.from_cu_seqlens(_vision_cu_seqlens(grid_thw))
|
||||
metadata = PackedSequenceMetadata.from_cu_seqlens(_vision_cu_seqlens(grid_thw))
|
||||
|
||||
deepstack_features = []
|
||||
for layer_idx, block in enumerate(self.blocks):
|
||||
|
||||
@@ -0,0 +1,65 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Shared SRT rotary embedding adapter for Qwen-VL text encoders."""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.rotary_embedding import get_rope
|
||||
from sglang.srt.layers.rotary_embedding.base import RotaryEmbedding
|
||||
from sglang.srt.utils.hf_transformers.common import get_rope_config
|
||||
|
||||
|
||||
def build_qwen_vl_text_rope(
|
||||
config: Any, *, mrope_interleaved: bool = False
|
||||
) -> RotaryEmbedding:
|
||||
head_dim = getattr(config, "head_dim", None) or (
|
||||
config.hidden_size // config.num_attention_heads
|
||||
)
|
||||
rope_theta, rope_scaling = get_rope_config(config)
|
||||
rope_scaling = dict(rope_scaling or {})
|
||||
rope_scaling["mrope_interleaved"] = mrope_interleaved
|
||||
return get_rope(
|
||||
head_size=head_dim,
|
||||
rotary_dim=head_dim,
|
||||
max_position=config.max_position_embeddings,
|
||||
base=rope_theta,
|
||||
is_neox_style=True,
|
||||
rope_scaling=rope_scaling,
|
||||
)
|
||||
|
||||
|
||||
def apply_qwen_vl_text_rope(
|
||||
rotary_emb: RotaryEmbedding,
|
||||
position_ids: torch.Tensor,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Apply three-axis MRoPE to batched attention tensors."""
|
||||
if query.ndim != 4 or key.ndim != 4:
|
||||
raise ValueError(
|
||||
"Qwen-VL query and key must have shape "
|
||||
"[batch, heads, sequence, head_dim]"
|
||||
)
|
||||
if position_ids.ndim != 3 or position_ids.shape[0] != 3:
|
||||
raise ValueError(
|
||||
"Qwen-VL text position_ids must have shape [3, batch, sequence]"
|
||||
)
|
||||
batch_size, num_query_heads, sequence_length, head_dim = query.shape
|
||||
key_batch_size, num_key_value_heads, key_sequence_length, key_head_dim = key.shape
|
||||
if (key_batch_size, key_sequence_length, key_head_dim) != (
|
||||
batch_size,
|
||||
sequence_length,
|
||||
head_dim,
|
||||
):
|
||||
raise ValueError("Qwen-VL query and key shapes are incompatible")
|
||||
if tuple(position_ids.shape[1:]) != (batch_size, sequence_length):
|
||||
raise ValueError("Qwen-VL position_ids do not match the attention input")
|
||||
|
||||
query = query.transpose(1, 2).reshape(-1, num_query_heads * head_dim)
|
||||
key = key.transpose(1, 2).reshape(-1, num_key_value_heads * head_dim)
|
||||
# Preserve HF's bf16 arithmetic order; fused MRoPE changes generated images.
|
||||
query, key = rotary_emb.forward_native(position_ids.reshape(3, -1), query, key)
|
||||
query = query.view(batch_size, sequence_length, num_query_heads, head_dim)
|
||||
key = key.view(batch_size, sequence_length, num_key_value_heads, head_dim)
|
||||
return query.transpose(1, 2), key.transpose(1, 2)
|
||||
@@ -0,0 +1,154 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Shared Qwen-VL vision attention for multimodal generation."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from sglang.multimodal_gen.runtime.layers.attention.selector import get_attn_backend
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.srt.layers.linear import QKVParallelLinear, RowParallelLinear
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PackedSequenceMetadata:
|
||||
cu_seqlens: torch.Tensor
|
||||
cu_seqlens_host: tuple[int, ...]
|
||||
max_seqlen: int
|
||||
|
||||
@classmethod
|
||||
def from_cu_seqlens(cls, cu_seqlens: torch.Tensor) -> PackedSequenceMetadata:
|
||||
bounds = tuple(int(value) for value in cu_seqlens.tolist())
|
||||
return cls(
|
||||
cu_seqlens=cu_seqlens,
|
||||
cu_seqlens_host=bounds,
|
||||
max_seqlen=max(
|
||||
stop - start for start, stop in zip(bounds[:-1], bounds[1:])
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _rotate_half(hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
first, second = hidden_states.chunk(2, dim=-1)
|
||||
return torch.cat((-second, first), dim=-1)
|
||||
|
||||
|
||||
def _apply_rotary_embedding(
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
cos: torch.Tensor,
|
||||
sin: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
query_dtype = query.dtype
|
||||
key_dtype = key.dtype
|
||||
query = query.float()
|
||||
key = key.float()
|
||||
cos = cos.unsqueeze(-2).float()
|
||||
sin = sin.unsqueeze(-2).float()
|
||||
query = query * cos + _rotate_half(query) * sin
|
||||
key = key * cos + _rotate_half(key) * sin
|
||||
return query.to(query_dtype), key.to(key_dtype)
|
||||
|
||||
|
||||
class QwenVLVisionAttention(nn.Module):
|
||||
def __init__(self, config: Any, *, prefix: str, model_name: str) -> None:
|
||||
super().__init__()
|
||||
parallel = get_parallel()
|
||||
self.num_heads = config.num_heads // parallel.tp_size
|
||||
self.head_dim = config.hidden_size // config.num_heads
|
||||
self.scaling = self.head_dim**-0.5
|
||||
self.qkv_proj = QKVParallelLinear(
|
||||
hidden_size=config.hidden_size,
|
||||
head_size=self.head_dim,
|
||||
total_num_heads=config.num_heads,
|
||||
bias=True,
|
||||
prefix=f"{prefix}.qkv_proj",
|
||||
tp_rank=parallel.tp_rank,
|
||||
tp_size=parallel.tp_size,
|
||||
)
|
||||
self.proj = RowParallelLinear(
|
||||
input_size=config.hidden_size,
|
||||
output_size=config.hidden_size,
|
||||
bias=True,
|
||||
prefix=f"{prefix}.proj",
|
||||
tp_rank=parallel.tp_rank,
|
||||
tp_size=parallel.tp_size,
|
||||
)
|
||||
|
||||
backend = get_attn_backend(self.head_dim, torch.get_default_dtype())
|
||||
self._attention_impl = None
|
||||
if backend.supports_packed_varlen():
|
||||
self._attention_impl = backend.get_impl_cls()(
|
||||
num_heads=self.num_heads,
|
||||
head_size=self.head_dim,
|
||||
num_kv_heads=self.num_heads,
|
||||
softmax_scale=self.scaling,
|
||||
causal=False,
|
||||
prefix=prefix,
|
||||
)
|
||||
else:
|
||||
logger.warning_once(
|
||||
f"{model_name} vision attention uses torch SDPA because "
|
||||
f"{backend.get_enum().name.lower()} does not support packed sequences"
|
||||
)
|
||||
|
||||
def _packed_attention(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
metadata: PackedSequenceMetadata,
|
||||
) -> torch.Tensor:
|
||||
if self._attention_impl is not None:
|
||||
return self._attention_impl.forward_varlen(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
cu_seqlens=metadata.cu_seqlens,
|
||||
cu_seqlens_host=metadata.cu_seqlens_host,
|
||||
max_seqlen=metadata.max_seqlen,
|
||||
)
|
||||
|
||||
output = torch.empty_like(query)
|
||||
for start, stop in zip(
|
||||
metadata.cu_seqlens_host[:-1], metadata.cu_seqlens_host[1:]
|
||||
):
|
||||
if start == stop:
|
||||
continue
|
||||
segment = F.scaled_dot_product_attention(
|
||||
query[start:stop].transpose(0, 1).unsqueeze(0),
|
||||
key[start:stop].transpose(0, 1).unsqueeze(0),
|
||||
value[start:stop].transpose(0, 1).unsqueeze(0),
|
||||
dropout_p=0.0,
|
||||
is_causal=False,
|
||||
scale=self.scaling,
|
||||
)
|
||||
output[start:stop] = segment.squeeze(0).transpose(0, 1)
|
||||
return output
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
*,
|
||||
metadata: PackedSequenceMetadata,
|
||||
position_embeddings: tuple[torch.Tensor, torch.Tensor],
|
||||
) -> torch.Tensor:
|
||||
sequence_length = hidden_states.shape[0]
|
||||
qkv, _ = self.qkv_proj(hidden_states)
|
||||
query, key, value = (
|
||||
qkv.reshape(sequence_length, 3, self.num_heads, self.head_dim)
|
||||
.permute(1, 0, 2, 3)
|
||||
.unbind(0)
|
||||
)
|
||||
query, key = _apply_rotary_embedding(query, key, *position_embeddings)
|
||||
output = self._packed_attention(query, key, value, metadata)
|
||||
output, _ = self.proj(output.reshape(sequence_length, -1).contiguous())
|
||||
return output
|
||||
@@ -1177,6 +1177,7 @@ STANDALONE_FILES = {
|
||||
"../single_test_file/test_disagg_server.py",
|
||||
"../single_test_file/test_ar_models.py",
|
||||
"../single_test_file/test_ipc_a2a_2_gpu.py",
|
||||
"../single_test_file/test_encoder_fold_srt_linear_2_gpu.py",
|
||||
"../single_test_file/test_encoder_fold_srt_2_gpu.py",
|
||||
"../single_test_file/test_diffusion_bcg_tp2_zimage_turbo.py",
|
||||
"../single_test_file/test_dp_serving_2_gpu.py",
|
||||
@@ -1216,6 +1217,7 @@ STANDALONE_FILE_EST_TIMES = {
|
||||
"../single_test_file/test_ar_models.py": 600.0,
|
||||
# no model load; the cost is the one-time JIT build of the sync kernels
|
||||
"../single_test_file/test_ipc_a2a_2_gpu.py": 240.0,
|
||||
"../single_test_file/test_encoder_fold_srt_linear_2_gpu.py": 120.0,
|
||||
"../single_test_file/test_encoder_fold_srt_2_gpu.py": 240.0,
|
||||
# ~60 s locally with a warm HF cache (load + one capture + 4 steps);
|
||||
# padded for cold-cache CI.
|
||||
|
||||
@@ -477,37 +477,40 @@ class AccuracyEngine:
|
||||
for name, tensor in target.named_parameters():
|
||||
total += 1
|
||||
src_tensor = None
|
||||
for cand in generate_name_candidates(name, reverse_mapping):
|
||||
candidates = generate_name_candidates(name, reverse_mapping)
|
||||
for cand in candidates:
|
||||
if cand in lookup:
|
||||
src_tensor = lookup[cand]
|
||||
break
|
||||
if src_tensor is None:
|
||||
for cand in generate_name_candidates(name, reverse_mapping):
|
||||
for cand in candidates:
|
||||
src_tensor = fuse_qkv(lookup, cand)
|
||||
if src_tensor is not None:
|
||||
break
|
||||
if src_tensor is None:
|
||||
for cand in generate_name_candidates(name, reverse_mapping):
|
||||
for cand in candidates:
|
||||
src_tensor = fuse_gate_up_proj(lookup, cand)
|
||||
if src_tensor is not None:
|
||||
break
|
||||
if src_tensor is None:
|
||||
unmatched_details.append(f"{name}: no matching source tensor")
|
||||
continue
|
||||
shard_context = shard_contexts.get(name)
|
||||
shard_world_size = (
|
||||
shard_context.world_size if shard_context is not None else tp_world
|
||||
)
|
||||
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.
|
||||
needs_weight_loader = (
|
||||
shard_world_size > 1 or tensor.shape != src_tensor.shape
|
||||
# Production loaders own fused projection sharding and alignment
|
||||
# padding. Use them for TP parameters and whenever a direct copy
|
||||
# cannot represent the source layout, including TP=1.
|
||||
requires_weight_loader = (
|
||||
shard_world_size > 1
|
||||
or src_tensor is None
|
||||
or src_tensor.shape != tensor.shape
|
||||
)
|
||||
if needs_weight_loader and load_param_with_weight_loader(
|
||||
if requires_weight_loader and load_param_with_weight_loader(
|
||||
tensor, name, lookup, reverse_mapping
|
||||
):
|
||||
matched += 1
|
||||
elif src_tensor is None:
|
||||
unmatched_details.append(f"{name}: no matching source tensor")
|
||||
elif copy_tensor(tensor, src_tensor, shard_world_size, shard_rank):
|
||||
matched += 1
|
||||
else:
|
||||
|
||||
@@ -915,10 +915,14 @@ def load_param_with_weight_loader(param, name, lookup, reverse_mapping) -> bool:
|
||||
loader(param, tensor.to(dtype=param.dtype), shard_id)
|
||||
return True
|
||||
for cand in candidates:
|
||||
src = lookup.get(cand)
|
||||
if src is not None:
|
||||
loader(param, src.to(dtype=param.dtype))
|
||||
return True
|
||||
source_names = [cand]
|
||||
if "qkv_proj" in cand:
|
||||
source_names.append(cand.replace("qkv_proj", "qkv"))
|
||||
for source_name in source_names:
|
||||
src = lookup.get(source_name)
|
||||
if src is not None:
|
||||
loader(param, src.to(dtype=param.dtype))
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
return False
|
||||
|
||||
+120
@@ -0,0 +1,120 @@
|
||||
"""A folded encoder must run SRT collectives on its bound TP group."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import unittest
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
_WORLD_SIZE = 2
|
||||
|
||||
|
||||
def _worker() -> int:
|
||||
from sglang.multimodal_gen.runtime.distributed import (
|
||||
cleanup_dist_env_and_memory,
|
||||
get_tp_group,
|
||||
get_world_group,
|
||||
init_distributed_environment,
|
||||
initialize_model_parallel,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.models.encoders.base import (
|
||||
EncoderTensorParallelMixin,
|
||||
)
|
||||
from sglang.srt.distributed import parallel_state as srt_parallel_state
|
||||
from sglang.srt.layers.linear import RowParallelLinear
|
||||
|
||||
rank = int(os.environ["RANK"])
|
||||
world_size = int(os.environ["WORLD_SIZE"])
|
||||
device = torch.device(f"cuda:{rank}")
|
||||
torch.cuda.set_device(device)
|
||||
init_distributed_environment(
|
||||
world_size=world_size,
|
||||
rank=rank,
|
||||
local_rank=rank,
|
||||
)
|
||||
initialize_model_parallel(
|
||||
tensor_parallel_degree=1,
|
||||
sequence_parallel_degree=world_size,
|
||||
ulysses_degree=world_size,
|
||||
ring_degree=1,
|
||||
)
|
||||
|
||||
class FoldedEncoder(EncoderTensorParallelMixin, nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.bind_encoder_tp_group(get_world_group())
|
||||
self.proj = RowParallelLinear(
|
||||
input_size=8,
|
||||
output_size=6,
|
||||
bias=False,
|
||||
tp_rank=rank,
|
||||
tp_size=world_size,
|
||||
params_dtype=torch.float32,
|
||||
)
|
||||
|
||||
def forward(self, inputs: torch.Tensor) -> torch.Tensor:
|
||||
local_inputs = inputs.chunk(world_size, dim=-1)[rank].contiguous()
|
||||
output, _ = self.proj(local_inputs)
|
||||
return output
|
||||
|
||||
full_weight = torch.arange(48, dtype=torch.float32, device=device).reshape(6, 8)
|
||||
full_weight = (full_weight - 23.5) / 32
|
||||
inputs = torch.arange(24, dtype=torch.float32, device=device).reshape(3, 8) / 8
|
||||
model = FoldedEncoder().to(device).eval()
|
||||
with torch.no_grad():
|
||||
model.proj.weight.copy_(full_weight[:, rank * 4 : (rank + 1) * 4].contiguous())
|
||||
expected = F.linear(inputs, full_weight)
|
||||
actual = model(inputs)
|
||||
|
||||
torch.testing.assert_close(actual, expected, rtol=1e-6, atol=1e-6)
|
||||
assert get_tp_group().world_size == 1
|
||||
assert srt_parallel_state.get_tp_group().world_size == 1
|
||||
assert srt_parallel_state.get_attn_tp_group().world_size == 1
|
||||
|
||||
if rank == 0:
|
||||
print("ENCODER_FOLD_SRT_LINEAR_PARITY PASS", flush=True)
|
||||
torch.distributed.barrier()
|
||||
cleanup_dist_env_and_memory()
|
||||
return 0
|
||||
|
||||
|
||||
class TestEncoderFoldSrtLinearTwoGpu(CustomTestCase):
|
||||
def test_folded_srt_linear_matches_unsharded_reference(self):
|
||||
if not current_platform.is_cuda():
|
||||
self.skipTest("CUDA-only test")
|
||||
if torch.cuda.device_count() < _WORLD_SIZE:
|
||||
self.skipTest(f"needs {_WORLD_SIZE} GPUs")
|
||||
|
||||
proc = subprocess.run(
|
||||
[
|
||||
sys.executable,
|
||||
"-m",
|
||||
"torch.distributed.run",
|
||||
f"--nproc-per-node={_WORLD_SIZE}",
|
||||
"--master-port=29618",
|
||||
__file__,
|
||||
"--worker",
|
||||
],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=600,
|
||||
)
|
||||
print(proc.stdout[-4000:])
|
||||
if proc.returncode != 0:
|
||||
print(proc.stderr[-4000:], file=sys.stderr)
|
||||
self.assertEqual(proc.returncode, 0, "folded SRT linear output diverged")
|
||||
self.assertIn("ENCODER_FOLD_SRT_LINEAR_PARITY PASS", proc.stdout)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
if "--worker" in sys.argv:
|
||||
raise SystemExit(_worker())
|
||||
unittest.main()
|
||||
@@ -0,0 +1,59 @@
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from sglang.multimodal_gen.test.single_test_file.component_accuracy.engine import (
|
||||
AccuracyEngine,
|
||||
)
|
||||
|
||||
|
||||
class _SourceProjectionSet(nn.Module):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.qkv = nn.Linear(2, 6, bias=False)
|
||||
self.gate_proj = nn.Linear(2, 3, bias=False)
|
||||
self.up_proj = nn.Linear(2, 3, bias=False)
|
||||
self.down_proj = nn.Linear(3, 2, bias=False)
|
||||
|
||||
|
||||
class _TargetProjectionSet(nn.Module):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.qkv_proj = nn.Linear(2, 6, bias=False)
|
||||
self.gate_up_proj = nn.Linear(2, 8, bias=False)
|
||||
self.down_proj = nn.Linear(4, 2, bias=False)
|
||||
|
||||
self.qkv_proj.weight.weight_loader = self._load_qkv
|
||||
self.gate_up_proj.weight.weight_loader = self._load_gate_up
|
||||
self.down_proj.weight.weight_loader = self._load_down
|
||||
|
||||
@staticmethod
|
||||
def _load_qkv(param: nn.Parameter, source: torch.Tensor) -> None:
|
||||
param.data.copy_(source)
|
||||
|
||||
@staticmethod
|
||||
def _load_gate_up(param: nn.Parameter, source: torch.Tensor, shard_id: int) -> None:
|
||||
offset = shard_id * 4
|
||||
param.data[offset : offset + source.shape[0]].copy_(source)
|
||||
|
||||
@staticmethod
|
||||
def _load_down(param: nn.Parameter, source: torch.Tensor) -> None:
|
||||
param.data[:, : source.shape[1]].copy_(source)
|
||||
|
||||
|
||||
def test_transfer_weights_uses_loaders_for_fused_aliases_and_padding() -> None:
|
||||
source = _SourceProjectionSet().to(dtype=torch.bfloat16)
|
||||
target = _TargetProjectionSet()
|
||||
with torch.no_grad():
|
||||
for index, parameter in enumerate(source.parameters(), start=1):
|
||||
parameter.fill_(index)
|
||||
for parameter in target.parameters():
|
||||
parameter.zero_()
|
||||
|
||||
AccuracyEngine.transfer_weights(source, target, target_device=torch.device("cpu"))
|
||||
|
||||
torch.testing.assert_close(target.qkv_proj.weight, source.qkv.weight)
|
||||
torch.testing.assert_close(target.gate_up_proj.weight[:3], source.gate_proj.weight)
|
||||
torch.testing.assert_close(target.gate_up_proj.weight[4:7], source.up_proj.weight)
|
||||
assert torch.count_nonzero(target.gate_up_proj.weight[[3, 7]]) == 0
|
||||
torch.testing.assert_close(target.down_proj.weight[:, :3], source.down_proj.weight)
|
||||
assert torch.count_nonzero(target.down_proj.weight[:, 3]) == 0
|
||||
@@ -15,6 +15,8 @@ from sglang.multimodal_gen.runtime.models.encoders.qwen2_5vl import (
|
||||
Qwen2_5_VLAttention,
|
||||
Qwen2_5_VLForConditionalGeneration,
|
||||
_apply_repetition_penalty,
|
||||
_make_column_linear,
|
||||
_make_row_linear,
|
||||
_select_next_token,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.models.encoders.qwen2_5vl_vision import (
|
||||
@@ -24,6 +26,17 @@ from sglang.multimodal_gen.runtime.models.encoders.qwen2_5vl_vision import (
|
||||
_vision_window_index,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines.longcat_image import LongCatImagePipeline
|
||||
from sglang.srt.layers.linear import (
|
||||
ColumnParallelLinear,
|
||||
ReplicatedLinear,
|
||||
RowParallelLinear,
|
||||
)
|
||||
from sglang.srt.models.qwen2_5_vl import (
|
||||
Qwen2_5_VisionPatchEmbed,
|
||||
Qwen2_5_VisionPatchMerger,
|
||||
Qwen2_5_VLMLP,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
|
||||
|
||||
class _StubQwen2_5VL(Qwen2_5_VLForConditionalGeneration):
|
||||
@@ -66,6 +79,75 @@ class _AttentionRecorder(nn.Module):
|
||||
return query
|
||||
|
||||
|
||||
def test_native_vision_reuses_srt_modules():
|
||||
config = SimpleNamespace(
|
||||
hidden_size=16,
|
||||
intermediate_size=24,
|
||||
hidden_act="silu",
|
||||
num_heads=2,
|
||||
depth=0,
|
||||
patch_size=2,
|
||||
temporal_patch_size=1,
|
||||
in_channels=3,
|
||||
spatial_merge_size=2,
|
||||
out_hidden_size=12,
|
||||
fullatt_block_indexes=[],
|
||||
window_size=8,
|
||||
)
|
||||
with get_parallel().override(tp_size=1, tp_rank=0):
|
||||
model = Qwen2_5VLVisionTransformer(config)
|
||||
mlp = Qwen2_5_VLMLP(
|
||||
16,
|
||||
24,
|
||||
fuse_gate_up=False,
|
||||
)
|
||||
fused_mlp = Qwen2_5_VLMLP(16, 24)
|
||||
|
||||
assert isinstance(model.patch_embed, Qwen2_5_VisionPatchEmbed)
|
||||
assert isinstance(model.merger, Qwen2_5_VisionPatchMerger)
|
||||
assert not mlp.fuse_gate_up
|
||||
assert isinstance(mlp.gate_proj, ColumnParallelLinear)
|
||||
assert isinstance(mlp.up_proj, ColumnParallelLinear)
|
||||
assert mlp.gate_proj.tp_size == mlp.up_proj.tp_size == 1
|
||||
assert isinstance(mlp.down_proj, ReplicatedLinear)
|
||||
assert isinstance(fused_mlp.down_proj, RowParallelLinear)
|
||||
assert mlp.act is not None
|
||||
assert isinstance(
|
||||
_make_column_linear(16, 24, bias=False, use_tensor_parallel=False),
|
||||
ReplicatedLinear,
|
||||
)
|
||||
assert isinstance(
|
||||
_make_row_linear(24, 16, bias=False, use_tensor_parallel=False),
|
||||
ReplicatedLinear,
|
||||
)
|
||||
|
||||
|
||||
def test_text_mlp_uses_single_rank_when_intermediate_size_is_not_tp_divisible(
|
||||
monkeypatch,
|
||||
):
|
||||
monkeypatch.setattr(qwen2_5vl, "Qwen2_5_VLAttention", lambda *_args: nn.Identity())
|
||||
monkeypatch.setattr(qwen2_5vl, "_tp_world_size", lambda: 3)
|
||||
monkeypatch.setattr(qwen2_5vl, "_tp_rank", lambda: 2)
|
||||
config = SimpleNamespace(
|
||||
hidden_size=16,
|
||||
intermediate_size=25,
|
||||
hidden_act="silu",
|
||||
rms_norm_eps=1e-6,
|
||||
use_sliding_window=False,
|
||||
_attn_implementation="flash_attention_2",
|
||||
layer_types=["full_attention"],
|
||||
)
|
||||
|
||||
layer = qwen2_5vl.Qwen2_5_VLDecoderLayer(config, layer_idx=0)
|
||||
|
||||
assert layer.mlp.tp_size == 1
|
||||
assert layer.mlp.tp_rank == 0
|
||||
assert isinstance(layer.mlp.gate_proj, ColumnParallelLinear)
|
||||
assert isinstance(layer.mlp.up_proj, ColumnParallelLinear)
|
||||
assert layer.mlp.gate_proj.tp_rank == layer.mlp.up_proj.tp_rank == 0
|
||||
assert isinstance(layer.mlp.down_proj, ReplicatedLinear)
|
||||
|
||||
|
||||
def test_explicit_attention_mask_is_limited_to_cached_generation(monkeypatch):
|
||||
attention = Qwen2_5_VLAttention.__new__(Qwen2_5_VLAttention)
|
||||
nn.Module.__init__(attention)
|
||||
@@ -76,12 +158,12 @@ def test_explicit_attention_mask_is_limited_to_cached_generation(monkeypatch):
|
||||
attention.num_heads = 1
|
||||
attention.num_key_value_heads = 1
|
||||
attention.head_dim = 4
|
||||
attention.rope_scaling = {"mrope_section": [1, 1, 0]}
|
||||
attention.rotary_emb = object()
|
||||
attention.attn = _AttentionRecorder()
|
||||
monkeypatch.setattr(
|
||||
qwen2_5vl,
|
||||
"apply_multimodal_rotary_pos_emb",
|
||||
lambda query, key, *_args: (query, key),
|
||||
"apply_qwen_vl_text_rope",
|
||||
lambda _rotary_emb, _position_ids, query, key: (query, key),
|
||||
)
|
||||
|
||||
hidden_states = torch.randn(1, 2, 4)
|
||||
@@ -89,7 +171,7 @@ def test_explicit_attention_mask_is_limited_to_cached_generation(monkeypatch):
|
||||
kwargs = {
|
||||
"hidden_states": hidden_states,
|
||||
"attention_mask": explicit_mask,
|
||||
"position_embeddings": (torch.empty(0), torch.empty(0)),
|
||||
"position_ids": torch.zeros(3, 1, 2, dtype=torch.long),
|
||||
}
|
||||
|
||||
attention(**kwargs, use_cache=False)
|
||||
|
||||
@@ -2,7 +2,10 @@ from types import SimpleNamespace
|
||||
|
||||
import torch
|
||||
|
||||
import sglang.multimodal_gen.runtime.models.encoders.qwen3 as qwen3
|
||||
import sglang.srt.layers.activation as srt_activation
|
||||
from sglang.multimodal_gen.runtime.models.encoders.qwen3 import Qwen3ForCausalLM
|
||||
from sglang.srt.layers.activation import SiluAndMul
|
||||
|
||||
|
||||
class _CaptureLayer(torch.nn.Module):
|
||||
@@ -26,6 +29,53 @@ class _IdentityNorm(torch.nn.Module):
|
||||
return hidden_states, None
|
||||
|
||||
|
||||
def test_mlp_reuses_srt_activation_without_server_context(monkeypatch):
|
||||
def fail_get_exec():
|
||||
raise AssertionError("SiluAndMul must not read an unpublished context")
|
||||
|
||||
monkeypatch.setattr(srt_activation, "publish_role", lambda: None)
|
||||
monkeypatch.setattr(srt_activation, "get_exec", fail_get_exec)
|
||||
|
||||
def make_linear(*_args, **_kwargs):
|
||||
return torch.nn.Identity()
|
||||
|
||||
monkeypatch.setattr(qwen3, "MergedColumnParallelLinear", make_linear)
|
||||
monkeypatch.setattr(qwen3, "RowParallelLinear", make_linear)
|
||||
|
||||
mlp = qwen3.Qwen3MLP(16, 24, "silu")
|
||||
|
||||
assert isinstance(mlp.act_fn, SiluAndMul)
|
||||
|
||||
|
||||
def test_attention_keeps_diffusion_one_pass_qk_norm(monkeypatch):
|
||||
monkeypatch.setattr(qwen3, "get_tp_world_size", lambda: 1)
|
||||
monkeypatch.setattr(
|
||||
qwen3, "QKVParallelLinear", lambda **kwargs: torch.nn.Identity()
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
qwen3, "RowParallelLinear", lambda **kwargs: torch.nn.Identity()
|
||||
)
|
||||
monkeypatch.setattr(qwen3, "get_rope", lambda *args, **kwargs: torch.nn.Identity())
|
||||
monkeypatch.setattr(
|
||||
qwen3, "LocalAttention", lambda *args, **kwargs: torch.nn.Identity()
|
||||
)
|
||||
config = SimpleNamespace(
|
||||
head_dim=128,
|
||||
rms_norm_eps=1e-6,
|
||||
_supported_attention_backends=(),
|
||||
)
|
||||
|
||||
attention = qwen3.Qwen3Attention(
|
||||
config,
|
||||
hidden_size=256,
|
||||
num_heads=2,
|
||||
num_kv_heads=1,
|
||||
)
|
||||
|
||||
assert isinstance(attention.q_norm, qwen3.MMGenRMSNorm)
|
||||
assert isinstance(attention.k_norm, qwen3.MMGenRMSNorm)
|
||||
|
||||
|
||||
def test_default_position_ids_batch_shape():
|
||||
model = Qwen3ForCausalLM.__new__(Qwen3ForCausalLM)
|
||||
torch.nn.Module.__init__(model)
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
from types import SimpleNamespace
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
import sglang.multimodal_gen.runtime.models.encoders.qwen3vl as qwen3vl
|
||||
|
||||
|
||||
class _IdentityAttention(nn.Module):
|
||||
def forward(self, query, key, value):
|
||||
return query
|
||||
|
||||
|
||||
def test_qwen3vl_attention_uses_interleaved_mrope(monkeypatch):
|
||||
captured_kwargs = {}
|
||||
|
||||
def build_rope(_config, **kwargs):
|
||||
captured_kwargs.update(kwargs)
|
||||
return object()
|
||||
|
||||
monkeypatch.setattr(qwen3vl, "build_qwen_vl_text_rope", build_rope)
|
||||
monkeypatch.setattr(
|
||||
qwen3vl, "_make_text_linear", lambda *args, **kwargs: nn.Identity()
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
qwen3vl, "_make_text_row_linear", lambda *args, **kwargs: nn.Identity()
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
qwen3vl, "_make_text_rms_norm", lambda *args, **kwargs: nn.Identity()
|
||||
)
|
||||
monkeypatch.setattr(qwen3vl, "LocalAttention", lambda **kwargs: nn.Identity())
|
||||
config = SimpleNamespace(
|
||||
head_dim=8,
|
||||
hidden_size=8,
|
||||
num_attention_heads=1,
|
||||
num_key_value_heads=1,
|
||||
attention_dropout=0.0,
|
||||
attention_bias=False,
|
||||
rms_norm_eps=1e-6,
|
||||
)
|
||||
|
||||
qwen3vl.Qwen3VLTextAttention(config, layer_idx=0)
|
||||
|
||||
assert captured_kwargs == {"mrope_interleaved": True}
|
||||
|
||||
|
||||
def test_qwen3vl_attention_passes_three_axis_positions_to_srt_rope(monkeypatch):
|
||||
attention = qwen3vl.Qwen3VLTextAttention.__new__(qwen3vl.Qwen3VLTextAttention)
|
||||
nn.Module.__init__(attention)
|
||||
attention.q_proj = nn.Identity()
|
||||
attention.k_proj = nn.Identity()
|
||||
attention.v_proj = nn.Identity()
|
||||
attention.o_proj = nn.Identity()
|
||||
attention.q_norm = nn.Identity()
|
||||
attention.k_norm = nn.Identity()
|
||||
attention.head_dim = 4
|
||||
attention.rotary_emb = object()
|
||||
attention.attn = _IdentityAttention()
|
||||
|
||||
captured_position_ids = None
|
||||
|
||||
def apply_rope(_rotary_emb, position_ids, query, key):
|
||||
nonlocal captured_position_ids
|
||||
captured_position_ids = position_ids
|
||||
return query, key
|
||||
|
||||
monkeypatch.setattr(qwen3vl, "apply_qwen_vl_text_rope", apply_rope)
|
||||
|
||||
hidden_states = torch.randn(1, 2, 4)
|
||||
position_ids = torch.arange(6).view(3, 1, 2)
|
||||
output = attention(
|
||||
hidden_states,
|
||||
position_ids=position_ids,
|
||||
attention_mask=None,
|
||||
)
|
||||
|
||||
assert captured_position_ids is position_ids
|
||||
torch.testing.assert_close(output, hidden_states)
|
||||
@@ -9,6 +9,7 @@ from sglang.multimodal_gen.runtime.models.encoders.minimax_h3_qwen3vl import (
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.models.encoders.qwen3vl import (
|
||||
Qwen3VLForConditionalGeneration,
|
||||
_make_text_rms_norm,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.models.encoders.qwen3vl_vision import (
|
||||
Qwen3VLVisionRotaryEmbedding,
|
||||
@@ -16,6 +17,12 @@ from sglang.multimodal_gen.runtime.models.encoders.qwen3vl_vision import (
|
||||
_vision_cu_seqlens,
|
||||
_vision_position_ids,
|
||||
)
|
||||
from sglang.srt.layers.layernorm import RMSNorm
|
||||
from sglang.srt.models.qwen3_vl import (
|
||||
Qwen3VLMoeVisionPatchMerger,
|
||||
Qwen3VLVisionPatchEmbed,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
|
||||
|
||||
def test_native_vision_layout_matches_qwen3_merge_order():
|
||||
@@ -38,6 +45,13 @@ def test_native_vision_layout_matches_qwen3_merge_order():
|
||||
assert cu_seqlens.tolist() == [0, 24, 32, 40]
|
||||
|
||||
|
||||
def test_qwen3vl_text_reuses_srt_rms_norm():
|
||||
norm = _make_text_rms_norm(16, 1e-6)
|
||||
|
||||
assert isinstance(norm, RMSNorm)
|
||||
assert norm.cast_x_before_out_mul
|
||||
|
||||
|
||||
def test_native_vision_keeps_checkpoint_parameter_names():
|
||||
config = SimpleNamespace(
|
||||
hidden_size=16,
|
||||
@@ -53,7 +67,11 @@ def test_native_vision_keeps_checkpoint_parameter_names():
|
||||
out_hidden_size=12,
|
||||
deepstack_visual_indexes=[],
|
||||
)
|
||||
model = Qwen3VLVisionTransformer(config)
|
||||
with get_parallel().override(tp_size=1, tp_rank=0):
|
||||
model = Qwen3VLVisionTransformer(config)
|
||||
|
||||
assert isinstance(model.patch_embed, Qwen3VLVisionPatchEmbed)
|
||||
assert isinstance(model.merger, Qwen3VLMoeVisionPatchMerger)
|
||||
|
||||
assert set(model.state_dict()) == {
|
||||
"patch_embed.proj.weight",
|
||||
|
||||
@@ -0,0 +1,150 @@
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
import sglang.multimodal_gen.runtime.models.encoders.qwen_vl_rope as qwen_vl_rope
|
||||
import sglang.srt.layers.rotary_embedding.base as rope_base
|
||||
import sglang.srt.layers.rotary_embedding.factory as rope_factory
|
||||
from sglang.multimodal_gen.runtime.models.encoders.qwen_vl_rope import (
|
||||
apply_qwen_vl_text_rope,
|
||||
build_qwen_vl_text_rope,
|
||||
)
|
||||
|
||||
|
||||
class _RecordingRotaryEmbedding(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.positions = None
|
||||
self.query_shape = None
|
||||
self.key_shape = None
|
||||
|
||||
def forward_native(self, positions, query, key):
|
||||
self.positions = positions
|
||||
self.query_shape = query.shape
|
||||
self.key_shape = key.shape
|
||||
return query + 1, key + 2
|
||||
|
||||
|
||||
def test_qwen_vl_rope_supports_transformers_v5_config(monkeypatch):
|
||||
rope_parameters = {
|
||||
"rope_type": "default",
|
||||
"rope_theta": 1_000_000.0,
|
||||
"mrope_section": [2, 1, 1],
|
||||
}
|
||||
config = SimpleNamespace(
|
||||
head_dim=None,
|
||||
hidden_size=32,
|
||||
num_attention_heads=4,
|
||||
max_position_embeddings=128,
|
||||
rope_parameters=rope_parameters,
|
||||
)
|
||||
captured_kwargs = {}
|
||||
rotary_emb = object()
|
||||
|
||||
def get_rope(**kwargs):
|
||||
captured_kwargs.update(kwargs)
|
||||
return rotary_emb
|
||||
|
||||
monkeypatch.setattr(qwen_vl_rope, "get_rope", get_rope)
|
||||
|
||||
assert build_qwen_vl_text_rope(config) is rotary_emb
|
||||
assert captured_kwargs == {
|
||||
"head_size": 8,
|
||||
"rotary_dim": 8,
|
||||
"max_position": 128,
|
||||
"base": 1_000_000.0,
|
||||
"is_neox_style": True,
|
||||
"rope_scaling": {**rope_parameters, "mrope_interleaved": False},
|
||||
}
|
||||
|
||||
|
||||
def test_qwen_vl_rope_enables_interleaved_layout_explicitly(monkeypatch):
|
||||
config = SimpleNamespace(
|
||||
head_dim=8,
|
||||
max_position_embeddings=128,
|
||||
rope_parameters={
|
||||
"rope_type": "default",
|
||||
"rope_theta": 1_000_000.0,
|
||||
"mrope_section": [2, 1, 1],
|
||||
},
|
||||
)
|
||||
captured_kwargs = {}
|
||||
|
||||
def get_rope(**kwargs):
|
||||
captured_kwargs.update(kwargs)
|
||||
return object()
|
||||
|
||||
monkeypatch.setattr(qwen_vl_rope, "get_rope", get_rope)
|
||||
|
||||
build_qwen_vl_text_rope(config, mrope_interleaved=True)
|
||||
|
||||
assert captured_kwargs["rope_scaling"] == {
|
||||
**config.rope_parameters,
|
||||
"mrope_interleaved": True,
|
||||
}
|
||||
|
||||
|
||||
def test_qwen_vl_rope_does_not_require_srt_runtime_context(monkeypatch):
|
||||
def fail_get_exec():
|
||||
raise AssertionError("Qwen-VL RoPE must not read an unpublished context")
|
||||
|
||||
monkeypatch.setattr(rope_base, "get_exec", fail_get_exec)
|
||||
monkeypatch.setattr(rope_base, "publish_role", lambda: None)
|
||||
monkeypatch.setattr(rope_factory, "_ROPE_DICT", {})
|
||||
config = SimpleNamespace(
|
||||
head_dim=None,
|
||||
hidden_size=32,
|
||||
num_attention_heads=4,
|
||||
max_position_embeddings=37,
|
||||
rope_parameters={
|
||||
"rope_type": "default",
|
||||
"rope_theta": 123_457.0,
|
||||
"mrope_section": [2, 1, 1],
|
||||
},
|
||||
)
|
||||
|
||||
rotary_emb = build_qwen_vl_text_rope(config)
|
||||
positions = torch.arange(9).view(3, 3)
|
||||
query = torch.randn(3, 16)
|
||||
key = torch.randn(3, 8)
|
||||
|
||||
rotated_query, rotated_key = rotary_emb.forward_native(positions, query, key)
|
||||
|
||||
assert rotated_query.shape == query.shape
|
||||
assert rotated_key.shape == key.shape
|
||||
|
||||
|
||||
def test_qwen_vl_rope_adapts_batched_gqa_layout():
|
||||
rotary_emb = _RecordingRotaryEmbedding()
|
||||
query = torch.randn(2, 4, 5, 8)
|
||||
key = torch.randn(2, 2, 5, 8)
|
||||
position_ids = torch.arange(30).view(3, 2, 5)
|
||||
|
||||
rotated_query, rotated_key = apply_qwen_vl_text_rope(
|
||||
rotary_emb, position_ids, query, key
|
||||
)
|
||||
|
||||
assert rotary_emb.query_shape == (10, 32)
|
||||
assert rotary_emb.key_shape == (10, 16)
|
||||
assert torch.equal(rotary_emb.positions, position_ids.reshape(3, -1))
|
||||
torch.testing.assert_close(rotated_query, query + 1)
|
||||
torch.testing.assert_close(rotated_key, key + 2)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("position_ids", "key"),
|
||||
[
|
||||
(torch.zeros(2, 1, 3, dtype=torch.long), torch.zeros(1, 1, 3, 4)),
|
||||
(torch.zeros(3, 1, 2, dtype=torch.long), torch.zeros(1, 1, 3, 4)),
|
||||
],
|
||||
)
|
||||
def test_qwen_vl_rope_rejects_incompatible_shapes(position_ids, key):
|
||||
with pytest.raises(ValueError):
|
||||
apply_qwen_vl_text_rope(
|
||||
_RecordingRotaryEmbedding(),
|
||||
position_ids,
|
||||
torch.zeros(1, 1, 3, 4),
|
||||
key,
|
||||
)
|
||||
@@ -108,6 +108,26 @@ class TestMiniMaxH3CheckpointFilter(unittest.TestCase):
|
||||
expected,
|
||||
)
|
||||
|
||||
def test_vision_qkv_checkpoint_name_maps_to_native_projection(self):
|
||||
encoder = MiniMaxH3Qwen3VLEncoder.__new__(MiniMaxH3Qwen3VLEncoder)
|
||||
torch.nn.Module.__init__(encoder)
|
||||
encoder.model = torch.nn.Module()
|
||||
encoder.model.visual = torch.nn.Module()
|
||||
block = torch.nn.Module()
|
||||
block.attn = torch.nn.Module()
|
||||
block.attn.qkv_proj = torch.nn.Linear(2, 2)
|
||||
encoder.model.visual.blocks = torch.nn.ModuleList([block])
|
||||
|
||||
loaded = encoder.load_weights(
|
||||
[("model.visual.blocks.0.attn.qkv.bias", torch.tensor([1.0, 2.0]))]
|
||||
)
|
||||
|
||||
self.assertEqual(loaded, {"model.visual.blocks.0.attn.qkv_proj.bias"})
|
||||
torch.testing.assert_close(
|
||||
encoder.model.visual.blocks[0].attn.qkv_proj.bias,
|
||||
torch.tensor([1.0, 2.0]),
|
||||
)
|
||||
|
||||
|
||||
class TestTextEncoderQuantization(unittest.TestCase):
|
||||
def setUp(self):
|
||||
|
||||
@@ -33,7 +33,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
|
||||
Phase,
|
||||
check_cuda_graph_backend,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel, publish_role
|
||||
from sglang.srt.utils import (
|
||||
cpu_has_amx_support,
|
||||
get_bool_env_var,
|
||||
@@ -130,7 +130,10 @@ logger = logging.getLogger(__name__)
|
||||
class SiluAndMul(BaseFusedOp):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
if get_exec().deterministic.rl_on_policy_target is not None:
|
||||
if (
|
||||
publish_role() is not None
|
||||
and get_exec().deterministic.rl_on_policy_target is not None
|
||||
):
|
||||
self._forward_method = self.forward_native
|
||||
elif _use_aiter and envs.SGLANG_OPT_USE_AITER_SILU_MUL.get():
|
||||
self._forward_method = self.forward_aiter
|
||||
|
||||
@@ -432,6 +432,7 @@ class RMSNorm(BaseFusedOp):
|
||||
weight_dtype: Optional = None,
|
||||
override_orig_dtype: Optional = None,
|
||||
x_pad_to_multiple: int = 0,
|
||||
force_native: bool = False,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.has_weight = has_weight
|
||||
@@ -467,6 +468,8 @@ class RMSNorm(BaseFusedOp):
|
||||
except ImportError:
|
||||
self._fused_pad_kernel = None
|
||||
self._forward_method = self.forward_aiter
|
||||
if force_native:
|
||||
self._forward_method = self.forward_native
|
||||
|
||||
def forward_cuda(
|
||||
self,
|
||||
@@ -481,11 +484,6 @@ class RMSNorm(BaseFusedOp):
|
||||
residual = residual + post_residual_addition
|
||||
return x, residual
|
||||
return x
|
||||
# sgl_kernel rmsnorm requires 2D input; reshape higher-rank tensors
|
||||
needs_reshape = x.dim() != 2 and residual is None
|
||||
if needs_reshape:
|
||||
original_shape = x.shape
|
||||
x = x.contiguous().reshape(-1, original_shape[-1])
|
||||
if self.variance_size_override is not None:
|
||||
return self.forward_native(x, residual, post_residual_addition)
|
||||
if is_batch_invariant_mode_enabled():
|
||||
@@ -495,6 +493,10 @@ class RMSNorm(BaseFusedOp):
|
||||
or get_exec().deterministic.rl_on_policy_target == "fsdp"
|
||||
):
|
||||
return self.forward_native(x, residual, post_residual_addition)
|
||||
original_shape = x.shape
|
||||
needs_reshape = x.dim() != 2
|
||||
if needs_reshape:
|
||||
x = x.contiguous().reshape(-1, original_shape[-1])
|
||||
out = rms_norm_batch_invariant(
|
||||
x,
|
||||
self.weight.data,
|
||||
@@ -517,6 +519,21 @@ class RMSNorm(BaseFusedOp):
|
||||
return self.forward_with_per_tensor_quant_fusion(
|
||||
x, scale, residual, post_residual_addition
|
||||
)
|
||||
|
||||
# CUDA RMSNorm kernels require 2D inputs. Flatten token dimensions for
|
||||
# the kernel call and restore each returned tensor to its input shape.
|
||||
original_shape = x.shape
|
||||
residual_shape = residual.shape if residual is not None else original_shape
|
||||
needs_reshape = x.dim() != 2
|
||||
if needs_reshape:
|
||||
x = x.contiguous().reshape(-1, original_shape[-1])
|
||||
if residual is not None:
|
||||
residual = residual.contiguous().reshape(-1, residual_shape[-1])
|
||||
if post_residual_addition is not None:
|
||||
post_residual_addition = post_residual_addition.contiguous().reshape(
|
||||
-1, post_residual_addition.shape[-1]
|
||||
)
|
||||
|
||||
if self.cast_x_before_out_mul and residual is None:
|
||||
# Use HF-semantics kernel (cast to dtype before weight multiply).
|
||||
if (
|
||||
@@ -531,10 +548,8 @@ class RMSNorm(BaseFusedOp):
|
||||
else:
|
||||
# Fallback: pure-Python HF semantics (already implemented in forward_native).
|
||||
out = self.forward_native(x, None, None)
|
||||
if needs_reshape:
|
||||
out = out.reshape(original_shape)
|
||||
return out
|
||||
if residual is not None:
|
||||
result = out
|
||||
elif residual is not None:
|
||||
if self.cast_x_before_out_mul:
|
||||
if (
|
||||
x.dtype in (torch.float16, torch.bfloat16)
|
||||
@@ -554,20 +569,28 @@ class RMSNorm(BaseFusedOp):
|
||||
self.variance_epsilon,
|
||||
cast_x_before_out_mul=self.cast_x_before_out_mul,
|
||||
)
|
||||
return x, residual
|
||||
return self.forward_native(x, residual, post_residual_addition)
|
||||
# TODO: Ideally we want to have (hidden_states+residual)+post_residual_addition.
|
||||
# but right now we can only have hidden_states+(residual+post_residual_addition).
|
||||
# (hidden_states+residual)+post_residual_addition != hidden_states+(residual+post_residual_addition),
|
||||
# we probably need to add another parameter to fused_add_rmsnorm
|
||||
if post_residual_addition is not None:
|
||||
residual = residual + post_residual_addition
|
||||
fused_add_rmsnorm(x, residual, self.weight.data, self.variance_epsilon)
|
||||
return x, residual
|
||||
out = rmsnorm(x, self.weight.data, self.variance_epsilon)
|
||||
result = x, residual
|
||||
else:
|
||||
result = self.forward_native(x, residual, post_residual_addition)
|
||||
else:
|
||||
# TODO: Ideally we want to have (hidden_states+residual)+post_residual_addition.
|
||||
# but right now we can only have hidden_states+(residual+post_residual_addition).
|
||||
# (hidden_states+residual)+post_residual_addition != hidden_states+(residual+post_residual_addition),
|
||||
# we probably need to add another parameter to fused_add_rmsnorm
|
||||
if post_residual_addition is not None:
|
||||
residual = residual + post_residual_addition
|
||||
fused_add_rmsnorm(x, residual, self.weight.data, self.variance_epsilon)
|
||||
result = x, residual
|
||||
else:
|
||||
result = rmsnorm(x, self.weight.data, self.variance_epsilon)
|
||||
|
||||
if needs_reshape:
|
||||
out = out.reshape(original_shape)
|
||||
return out
|
||||
if residual is not None:
|
||||
return result[0].reshape(original_shape), result[1].reshape(
|
||||
residual_shape
|
||||
)
|
||||
return result.reshape(original_shape)
|
||||
return result
|
||||
|
||||
def forward_npu(
|
||||
self,
|
||||
@@ -602,15 +625,6 @@ class RMSNorm(BaseFusedOp):
|
||||
# AITER's ROCm rmsnorm2d_fwd requires weight/activation dtypes to match;
|
||||
# FP32 weight + BF16 activation yields finite-but-corrupted output on gfx950.
|
||||
return self.forward_native(x, residual, post_residual_addition)
|
||||
# Aiter's RMSNorm kernels expect 2D contiguous inputs. Keep the
|
||||
# already-safe layout as a zero-copy path, and only normalize strided or
|
||||
# higher-rank views such as Q/K slices from packed QKV projections.
|
||||
needs_reshape = x.dim() != 2 and residual is None
|
||||
if needs_reshape:
|
||||
original_shape = x.shape
|
||||
x = x.contiguous().reshape(-1, original_shape[-1])
|
||||
elif not x.is_contiguous():
|
||||
x = x.contiguous()
|
||||
if is_batch_invariant_mode_enabled():
|
||||
if (
|
||||
residual is not None
|
||||
@@ -619,6 +633,10 @@ class RMSNorm(BaseFusedOp):
|
||||
or (self._fused_pad_kernel is not None and self.x_pad_to_multiple > 0)
|
||||
):
|
||||
return self.forward_native(x, residual, post_residual_addition)
|
||||
original_shape = x.shape
|
||||
needs_reshape = x.dim() != 2
|
||||
if needs_reshape:
|
||||
x = x.contiguous().reshape(-1, original_shape[-1])
|
||||
out = rms_norm_batch_invariant(
|
||||
x,
|
||||
self.weight.data,
|
||||
@@ -627,6 +645,25 @@ class RMSNorm(BaseFusedOp):
|
||||
if needs_reshape:
|
||||
out = out.reshape(original_shape)
|
||||
return out
|
||||
|
||||
# AITER's RMSNorm kernels require 2D contiguous inputs.
|
||||
original_shape = x.shape
|
||||
residual_shape = residual.shape if residual is not None else original_shape
|
||||
needs_reshape = x.dim() != 2
|
||||
if needs_reshape:
|
||||
x = x.contiguous().reshape(-1, original_shape[-1])
|
||||
if residual is not None:
|
||||
residual = residual.contiguous().reshape(-1, residual_shape[-1])
|
||||
if post_residual_addition is not None:
|
||||
post_residual_addition = post_residual_addition.contiguous().reshape(
|
||||
-1, post_residual_addition.shape[-1]
|
||||
)
|
||||
else:
|
||||
if not x.is_contiguous():
|
||||
x = x.contiguous()
|
||||
if residual is not None and not residual.is_contiguous():
|
||||
residual = residual.contiguous()
|
||||
|
||||
# Fused (add +) rmsnorm + zero-pad path. Triggered when caller
|
||||
# constructed RMSNorm with x_pad_to_multiple > 0. Output last
|
||||
# dim is padded up; residual_out stays at original width. Used
|
||||
@@ -636,13 +673,20 @@ class RMSNorm(BaseFusedOp):
|
||||
if self._fused_pad_kernel is not None and self.x_pad_to_multiple > 0:
|
||||
if post_residual_addition is not None and residual is not None:
|
||||
residual = residual + post_residual_addition
|
||||
return self._fused_pad_kernel(
|
||||
result = self._fused_pad_kernel(
|
||||
x,
|
||||
self.weight.data,
|
||||
self.variance_epsilon,
|
||||
residual,
|
||||
self.x_pad_to_multiple,
|
||||
)
|
||||
if needs_reshape and residual is not None:
|
||||
output, residual_out = result
|
||||
output_shape = (*original_shape[:-1], output.shape[-1])
|
||||
return output.reshape(output_shape), residual_out.reshape(
|
||||
residual_shape
|
||||
)
|
||||
return result
|
||||
if residual is not None:
|
||||
residual_out = torch.empty_like(x)
|
||||
output = torch.empty_like(x)
|
||||
@@ -656,6 +700,10 @@ class RMSNorm(BaseFusedOp):
|
||||
self.weight.data,
|
||||
self.variance_epsilon,
|
||||
)
|
||||
if needs_reshape:
|
||||
return output.reshape(original_shape), residual_out.reshape(
|
||||
residual_shape
|
||||
)
|
||||
return output, residual_out
|
||||
output = rms_norm(x, self.weight.data, self.variance_epsilon)
|
||||
if needs_reshape:
|
||||
|
||||
@@ -11,7 +11,7 @@ from sglang.kernels.fused_op import BaseFusedOp
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.rotary_embedding.utils import apply_rotary_emb
|
||||
from sglang.srt.platforms import current_platform
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
from sglang.srt.runtime_context import get_exec, publish_role
|
||||
from sglang.srt.utils import (
|
||||
cpu_has_amx_support,
|
||||
get_bool_env_var,
|
||||
@@ -94,6 +94,10 @@ class RotaryEmbedding(BaseFusedOp):
|
||||
self.base = base
|
||||
self.is_neox_style = is_neox_style
|
||||
self.dtype = dtype
|
||||
self._force_native = (
|
||||
publish_role() is not None
|
||||
and get_exec().deterministic.rl_on_policy_target is not None
|
||||
)
|
||||
|
||||
cache = self._compute_cos_sin_cache()
|
||||
# NOTE(ByronHsu): cache needs to be in FP32 for numerical stability.
|
||||
@@ -129,7 +133,7 @@ class RotaryEmbedding(BaseFusedOp):
|
||||
self._apply_rotary_emb_wrapped = apply_rotary_emb
|
||||
|
||||
# XXX (MUSA): Implement sgl_kernel.rotary_embedding support for MUSA backend
|
||||
if get_exec().deterministic.rl_on_policy_target is not None or _is_musa:
|
||||
if self._force_native or _is_musa:
|
||||
self._forward_method = self.forward_native
|
||||
self._apply_rotary_emb_wrapped = torch.compile(
|
||||
dynamic=True,
|
||||
@@ -152,9 +156,7 @@ class RotaryEmbedding(BaseFusedOp):
|
||||
# use CPU to compute the cache and then move it to GPU. However, we
|
||||
# create the cache on GPU for faster initialization. This may cause
|
||||
# a slight numerical difference between the HF implementation and ours.
|
||||
init_device = (
|
||||
"cpu" if get_exec().deterministic.rl_on_policy_target is not None else None
|
||||
)
|
||||
init_device = "cpu" if self._force_native else None
|
||||
inv_freq = 1.0 / (
|
||||
base
|
||||
** (
|
||||
@@ -164,7 +166,7 @@ class RotaryEmbedding(BaseFusedOp):
|
||||
/ self.rotary_dim
|
||||
)
|
||||
)
|
||||
if get_exec().deterministic.rl_on_policy_target is not None:
|
||||
if self._force_native:
|
||||
inv_freq = inv_freq.cuda()
|
||||
return inv_freq
|
||||
|
||||
|
||||
@@ -18,7 +18,7 @@ from sglang.srt.layers.rotary_embedding.yarn import (
|
||||
yarn_get_mscale_simple,
|
||||
yarn_linear_ramp_mask,
|
||||
)
|
||||
from sglang.srt.runtime_context import attention_backends, get_exec
|
||||
from sglang.srt.runtime_context import attention_backends
|
||||
from sglang.srt.utils import (
|
||||
cpu_has_amx_support,
|
||||
is_cuda,
|
||||
@@ -131,7 +131,7 @@ class MRotaryEmbedding(RotaryEmbedding):
|
||||
self.register_buffer("axis_map", axis_map, persistent=False)
|
||||
else:
|
||||
self.axis_map = None
|
||||
if get_exec().deterministic.rl_on_policy_target is not None:
|
||||
if self._force_native:
|
||||
self._forward_method = self.forward_native
|
||||
|
||||
def get_cos_sin_with_position(self, positions):
|
||||
|
||||
@@ -37,10 +37,6 @@ from transformers.models.qwen2_5_vl.configuration_qwen2_5_vl import (
|
||||
Qwen2_5_VLConfig,
|
||||
Qwen2_5_VLVisionConfig,
|
||||
)
|
||||
from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import (
|
||||
Qwen2_5_VisionPatchEmbed,
|
||||
Qwen2_5_VisionRotaryEmbedding,
|
||||
)
|
||||
|
||||
from sglang.srt.distributed.parallel_state import get_pp_group
|
||||
from sglang.srt.environ import envs
|
||||
@@ -50,10 +46,12 @@ from sglang.srt.layers.attention.vision import (
|
||||
VisionAttentionMetadata,
|
||||
prepare_vision_attention_metadata,
|
||||
)
|
||||
from sglang.srt.layers.conv import Conv3dLayer
|
||||
from sglang.srt.layers.layernorm import RMSNorm
|
||||
from sglang.srt.layers.linear import (
|
||||
ColumnParallelLinear,
|
||||
MergedColumnParallelLinear,
|
||||
ReplicatedLinear,
|
||||
RowParallelLinear,
|
||||
)
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessor
|
||||
@@ -85,6 +83,52 @@ _is_cpu = is_cpu()
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class Qwen2_5_VisionPatchEmbed(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
patch_size: int,
|
||||
temporal_patch_size: int,
|
||||
in_channels: int,
|
||||
embed_dim: int,
|
||||
disable_linear: bool = False,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.patch_size = patch_size
|
||||
self.temporal_patch_size = temporal_patch_size
|
||||
self.in_channels = in_channels
|
||||
self.embed_dim = embed_dim
|
||||
kernel_size = (temporal_patch_size, patch_size, patch_size)
|
||||
self.proj = Conv3dLayer(
|
||||
in_channels,
|
||||
embed_dim,
|
||||
kernel_size=kernel_size,
|
||||
stride=kernel_size,
|
||||
bias=False,
|
||||
disable_linear=disable_linear,
|
||||
)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states = hidden_states.view(
|
||||
-1,
|
||||
self.in_channels,
|
||||
self.temporal_patch_size,
|
||||
self.patch_size,
|
||||
self.patch_size,
|
||||
)
|
||||
hidden_states = self.proj(hidden_states.to(self.proj.weight.dtype))
|
||||
return hidden_states.view(-1, self.embed_dim)
|
||||
|
||||
|
||||
class Qwen2_5_VisionRotaryEmbedding(nn.Module):
|
||||
def __init__(self, dim: int, theta: float = 10000.0) -> None:
|
||||
super().__init__()
|
||||
inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim))
|
||||
self.register_buffer("inv_freq", inv_freq, persistent=False)
|
||||
|
||||
def forward(self, position_ids: torch.Tensor) -> torch.Tensor:
|
||||
return (position_ids.unsqueeze(-1) * self.inv_freq).flatten(1)
|
||||
|
||||
|
||||
class Qwen2_5_VLMLP(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
@@ -95,31 +139,73 @@ class Qwen2_5_VLMLP(nn.Module):
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
use_data_parallel: bool = False,
|
||||
fuse_gate_up: bool = True,
|
||||
tp_size: Optional[int] = None,
|
||||
tp_rank: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.tp_size = 1 if use_data_parallel else get_parallel().tp_size
|
||||
self.tp_rank = 0 if use_data_parallel else get_parallel().tp_rank
|
||||
self.gate_up_proj = MergedColumnParallelLinear(
|
||||
input_size=in_features,
|
||||
output_sizes=[hidden_features] * 2, # [gate_proj, up_proj]
|
||||
bias=bias,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("gate_up_proj", prefix),
|
||||
tp_size=self.tp_size,
|
||||
tp_rank=self.tp_rank,
|
||||
)
|
||||
self.down_proj = RowParallelLinear(
|
||||
hidden_features,
|
||||
in_features,
|
||||
bias=bias,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("down_proj", prefix),
|
||||
tp_size=self.tp_size,
|
||||
tp_rank=self.tp_rank,
|
||||
)
|
||||
if use_data_parallel:
|
||||
if tp_size is not None or tp_rank is not None:
|
||||
raise ValueError(
|
||||
"Explicit MLP TP cannot be combined with data parallel"
|
||||
)
|
||||
self.tp_size, self.tp_rank = 1, 0
|
||||
else:
|
||||
if (tp_size is None) != (tp_rank is None):
|
||||
raise ValueError("MLP tp_size and tp_rank must be set together")
|
||||
self.tp_size = get_parallel().tp_size if tp_size is None else tp_size
|
||||
self.tp_rank = get_parallel().tp_rank if tp_rank is None else tp_rank
|
||||
self.fuse_gate_up = fuse_gate_up
|
||||
if fuse_gate_up:
|
||||
self.gate_up_proj = MergedColumnParallelLinear(
|
||||
input_size=in_features,
|
||||
output_sizes=[hidden_features] * 2, # [gate_proj, up_proj]
|
||||
bias=bias,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("gate_up_proj", prefix),
|
||||
tp_size=self.tp_size,
|
||||
tp_rank=self.tp_rank,
|
||||
)
|
||||
else:
|
||||
projection_kwargs = dict(
|
||||
input_size=in_features,
|
||||
output_size=hidden_features,
|
||||
bias=bias,
|
||||
quant_config=quant_config,
|
||||
tp_size=self.tp_size,
|
||||
tp_rank=self.tp_rank,
|
||||
)
|
||||
self.gate_proj = ColumnParallelLinear(
|
||||
**projection_kwargs,
|
||||
prefix=add_prefix("gate_proj", prefix),
|
||||
)
|
||||
self.up_proj = ColumnParallelLinear(
|
||||
**projection_kwargs,
|
||||
prefix=add_prefix("up_proj", prefix),
|
||||
)
|
||||
if not self.fuse_gate_up and self.tp_size == 1:
|
||||
self.down_proj = ReplicatedLinear(
|
||||
hidden_features,
|
||||
in_features,
|
||||
bias=bias,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("down_proj", prefix),
|
||||
)
|
||||
else:
|
||||
self.down_proj = RowParallelLinear(
|
||||
hidden_features,
|
||||
in_features,
|
||||
bias=bias,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("down_proj", prefix),
|
||||
tp_size=self.tp_size,
|
||||
tp_rank=self.tp_rank,
|
||||
)
|
||||
self.hidden_act = hidden_act
|
||||
if self.hidden_act == "silu":
|
||||
if self.fuse_gate_up and self.hidden_act == "silu":
|
||||
self.act = SiluAndMul()
|
||||
elif not self.fuse_gate_up:
|
||||
self.act = ACT2FN[self.hidden_act]
|
||||
else:
|
||||
base_act = ACT2FN[self.hidden_act]
|
||||
|
||||
@@ -130,8 +216,13 @@ class Qwen2_5_VLMLP(nn.Module):
|
||||
self.act = _act_fn
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
gate_up, _ = self.gate_up_proj(x)
|
||||
x = self.act(gate_up)
|
||||
if self.fuse_gate_up:
|
||||
gate_up, _ = self.gate_up_proj(x)
|
||||
x = self.act(gate_up)
|
||||
else:
|
||||
gate, _ = self.gate_proj(x)
|
||||
up, _ = self.up_proj(x)
|
||||
x = self.act(gate) * up
|
||||
x_down, _ = self.down_proj(x)
|
||||
return x_down
|
||||
|
||||
@@ -225,11 +316,18 @@ class Qwen2_5_VisionPatchMerger(nn.Module):
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
use_data_parallel: bool = False,
|
||||
cast_x_before_out_mul: bool = False,
|
||||
force_native_norm: bool = False,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.hidden_size = context_dim * (spatial_merge_size**2)
|
||||
self.padded_context_dim = padded_context_dim * (spatial_merge_size**2)
|
||||
self.ln_q = RMSNorm(context_dim, eps=1e-6)
|
||||
self.ln_q = RMSNorm(
|
||||
context_dim,
|
||||
eps=1e-6,
|
||||
cast_x_before_out_mul=cast_x_before_out_mul,
|
||||
force_native=force_native_norm,
|
||||
)
|
||||
tp_size = 1 if use_data_parallel else get_parallel().tp_size
|
||||
tp_rank = 0 if use_data_parallel else get_parallel().tp_rank
|
||||
self.mlp = nn.ModuleList(
|
||||
@@ -257,10 +355,8 @@ class Qwen2_5_VisionPatchMerger(nn.Module):
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
# x expected shape: [S, B, context_dim]
|
||||
S, B, D = x.shape
|
||||
x2d = x.reshape(-1, D)
|
||||
x2d = self.ln_q(x2d) # RMSNorm expects 2D
|
||||
x2d = x.reshape(-1, x.shape[-1])
|
||||
x2d = self.ln_q(x2d)
|
||||
x2d = x2d.view(-1, self.hidden_size) # group into spatial_merge_unit
|
||||
mlp_fc1, mlp_act, mlp_fc2 = self.mlp
|
||||
x_parallel, _ = mlp_fc1(x2d)
|
||||
|
||||
@@ -102,6 +102,25 @@ _is_cpu = is_cpu()
|
||||
_VECTORIZED_VL_POS_EMBED_MIN_IMAGES = 6
|
||||
|
||||
|
||||
def _resolve_vision_tp(
|
||||
*,
|
||||
use_data_parallel: bool,
|
||||
tp_size: Optional[int],
|
||||
tp_rank: Optional[int],
|
||||
) -> tuple[int, int]:
|
||||
if use_data_parallel:
|
||||
if tp_size is not None or tp_rank is not None:
|
||||
raise ValueError("Explicit vision TP cannot be combined with data parallel")
|
||||
return 1, 0
|
||||
if (tp_size is None) != (tp_rank is None):
|
||||
raise ValueError("Vision tp_size and tp_rank must be set together")
|
||||
if tp_size is None:
|
||||
parallel = get_parallel()
|
||||
return parallel.attn_tp_size, parallel.attn_tp_rank
|
||||
assert tp_rank is not None
|
||||
return tp_size, tp_rank
|
||||
|
||||
|
||||
class Qwen3_VisionMLP(nn.Module):
|
||||
|
||||
def __init__(
|
||||
@@ -113,10 +132,15 @@ class Qwen3_VisionMLP(nn.Module):
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
use_data_parallel: bool = False,
|
||||
tp_size: Optional[int] = None,
|
||||
tp_rank: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.tp_size = 1 if use_data_parallel else get_parallel().attn_tp_size
|
||||
self.tp_rank = 0 if use_data_parallel else get_parallel().attn_tp_rank
|
||||
self.tp_size, self.tp_rank = _resolve_vision_tp(
|
||||
use_data_parallel=use_data_parallel,
|
||||
tp_size=tp_size,
|
||||
tp_rank=tp_rank,
|
||||
)
|
||||
self.linear_fc1 = ColumnParallelLinear(
|
||||
in_features,
|
||||
hidden_features,
|
||||
@@ -145,7 +169,7 @@ class Qwen3_VisionMLP(nn.Module):
|
||||
|
||||
|
||||
class Qwen3VLVisionPatchEmbed(nn.Module):
|
||||
def __init__(self, config) -> None:
|
||||
def __init__(self, config, disable_linear: bool = False) -> None:
|
||||
super().__init__()
|
||||
self.patch_size = config.patch_size
|
||||
self.temporal_patch_size = config.temporal_patch_size
|
||||
@@ -159,6 +183,7 @@ class Qwen3VLVisionPatchEmbed(nn.Module):
|
||||
kernel_size=kernel_size,
|
||||
stride=kernel_size,
|
||||
bias=True,
|
||||
disable_linear=disable_linear,
|
||||
)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
@@ -265,6 +290,8 @@ class Qwen3VLMoeVisionPatchMerger(nn.Module):
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
use_data_parallel: bool = False,
|
||||
tp_size: Optional[int] = None,
|
||||
tp_rank: Optional[int] = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.hidden_size = context_dim * (spatial_merge_size**2)
|
||||
@@ -277,8 +304,11 @@ class Qwen3VLMoeVisionPatchMerger(nn.Module):
|
||||
self.norm = norm_layer(
|
||||
self.hidden_size if use_postshuffle_norm else context_dim
|
||||
)
|
||||
self.tp_size = 1 if use_data_parallel else get_parallel().attn_tp_size
|
||||
self.tp_rank = 0 if use_data_parallel else get_parallel().attn_tp_rank
|
||||
self.tp_size, self.tp_rank = _resolve_vision_tp(
|
||||
use_data_parallel=use_data_parallel,
|
||||
tp_size=tp_size,
|
||||
tp_rank=tp_rank,
|
||||
)
|
||||
self.linear_fc1 = ColumnParallelLinear(
|
||||
self.hidden_size,
|
||||
self.padded_context_dim,
|
||||
|
||||
Reference in New Issue
Block a user