[fix] moss-vl: use Conv3dLayer and remove no-op flat_encoder_result (#23932)

This commit is contained in:
Shaojun Zhou
2026-04-30 14:19:45 +08:00
committed by GitHub
parent 936c9c2355
commit 4f0b44c5c6
+12 -60
View File
@@ -20,6 +20,7 @@ from sglang.srt.distributed import (
from sglang.srt.layers.activation import SiluAndMul
from sglang.srt.layers.attention.vision import VisionAttention
from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes
from sglang.srt.layers.conv import Conv3dLayer
from sglang.srt.layers.dp_attention import get_attention_tp_rank, get_attention_tp_size
from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import (
@@ -96,7 +97,7 @@ class MossVLVisionPatchEmbed(nn.Module):
self.embed_dim = config.hidden_size
kernel_size = [self.temporal_patch_size, self.patch_size, self.patch_size]
self.proj = nn.Conv3d(
self.proj = Conv3dLayer(
self.in_channels,
self.embed_dim,
kernel_size=kernel_size,
@@ -1142,11 +1143,10 @@ class MossVLForConditionalGeneration(nn.Module):
def _collect_mm_data(self, forward_batch: ForwardBatch):
"""Collect pixel_values, grid_thw, and vision_position_ids from uncached requests."""
if forward_batch.forward_mode.is_decode() or all(forward_batch.encoder_cached):
return None, None, None, None
return None, None, None
pixel_values_list = []
grid_thw_list = []
encoder_lens_need = []
vision_pos_ids_list = []
for i, mm_input in enumerate(forward_batch.mm_inputs):
@@ -1161,14 +1161,13 @@ class MossVLForConditionalGeneration(nn.Module):
if grid_thw is not None:
grid_thw_list.append(torch.as_tensor(grid_thw, dtype=torch.long))
encoder_len = forward_batch.encoder_lens_cpu[i]
encoder_lens_need.append(encoder_len)
vp = mm_input.vision_position_ids
if vp is not None:
vision_pos_ids_list.append(vp[:, :encoder_len])
if not pixel_values_list:
return None, None, None, None
return None, None, None
pixel_values = torch.cat(pixel_values_list, dim=0)
grid_thw = torch.cat(grid_thw_list, dim=0) if grid_thw_list else None
@@ -1176,7 +1175,7 @@ class MossVLForConditionalGeneration(nn.Module):
torch.cat(vision_pos_ids_list, dim=1) if vision_pos_ids_list else None
)
return pixel_values, grid_thw, encoder_lens_need, packed_vision_pos_ids
return pixel_values, grid_thw, packed_vision_pos_ids
def _get_vision_features(
self,
@@ -1225,51 +1224,6 @@ class MossVLForConditionalGeneration(nn.Module):
return torch.cat(output_parts, dim=0)
def flat_encoder_result(
self,
cross_attention_states: torch.Tensor,
encoder_lens_need: List[int],
) -> torch.Tensor:
"""Copy vision states into a flat packed tensor, trimmed to encoder_lens."""
total_encoder_len = sum(encoder_lens_need)
head_dim = cross_attention_states.shape[-1]
if cross_attention_states.dim() == 1:
return cross_attention_states
# cross_attention_states is already packed (total_tokens, hidden_size)
# We need to split it according to encoder_lens_need
result = torch.zeros(
total_encoder_len,
head_dim,
device=cross_attention_states.device,
dtype=cross_attention_states.dtype,
)
src_offset = 0
dst_offset = 0
for encoder_len in encoder_lens_need:
if encoder_len > 0:
if src_offset + encoder_len > cross_attention_states.shape[0]:
raise RuntimeError(
"Encoder length mismatch: expected "
f"{encoder_len} tokens, but only "
f"{cross_attention_states.shape[0] - src_offset} remaining."
)
result[dst_offset : dst_offset + encoder_len] = cross_attention_states[
src_offset : src_offset + encoder_len
]
src_offset += encoder_len
dst_offset += encoder_len
if src_offset != cross_attention_states.shape[0]:
raise RuntimeError(
"Encoder length mismatch: produced "
f"{cross_attention_states.shape[0]} tokens, expected {src_offset}."
)
return result
# ---- prepare_forward_batch (called before attn backend init) ----
def prepare_forward_batch(self, forward_batch: ForwardBatch):
@@ -1509,8 +1463,8 @@ class MossVLForConditionalGeneration(nn.Module):
positions = forward_batch.mrope_positions
# 1. Collect vision inputs for uncached requests
pixel_values, grid_thw, encoder_lens_need, vision_position_ids = (
self._collect_mm_data(forward_batch)
pixel_values, grid_thw, vision_position_ids = self._collect_mm_data(
forward_batch
)
cross_attention_mask = None
@@ -1534,14 +1488,12 @@ class MossVLForConditionalGeneration(nn.Module):
if pixel_values is not None and grid_thw is not None:
# Run ViT
vision_hidden_states = self._get_vision_features(pixel_values, grid_thw)
# Insert separator tokens after each frame
vision_with_sep = self._insert_separator_tokens(
# Insert separator tokens after each frame. The result is already
# packed (total_tokens, hidden_size) matching encoder_lens, so it
# can be passed directly into the cross-attention path.
cross_attention_states = self._insert_separator_tokens(
vision_hidden_states, grid_thw
)
# Flatten to match encoder_lens
cross_attention_states = self.flat_encoder_result(
vision_with_sep, encoder_lens_need
)
# Drop heavy per-request vision tensors now that the encoder KV
# has been produced and will be cached. Otherwise pixel_values and
# vision_position_ids stay pinned on req.multimodal_inputs across
@@ -1551,7 +1503,7 @@ class MossVLForConditionalGeneration(nn.Module):
# Note: the local `vision_position_ids` is still needed by the LM
# cross-attention below, so we keep it; but we drop the per-request
# copy on mm_input, which we won't read again.
del pixel_values, vision_hidden_states, vision_with_sep
del pixel_values, vision_hidden_states
for i, mm_input in enumerate(forward_batch.mm_inputs):
if forward_batch.encoder_cached[i] or mm_input is None:
continue