[diffusion] refactor: reuse srt qwen vision and text modules (#35006)

This commit is contained in:
Mick
2026-08-19 10:12:44 +08:00
committed by GitHub
parent 58c5bee3ac
commit 4cef72faee
26 changed files with 1347 additions and 713 deletions
@@ -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
@@ -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):
+5 -2
View File
@@ -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
+80 -32
View File
@@ -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):
+128 -32
View File
@@ -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)
+35 -5
View File
@@ -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,