[fix] moss-vl: use Conv3dLayer and remove no-op flat_encoder_result (#23932)
This commit is contained in:
@@ -20,6 +20,7 @@ from sglang.srt.distributed import (
|
|||||||
from sglang.srt.layers.activation import SiluAndMul
|
from sglang.srt.layers.activation import SiluAndMul
|
||||||
from sglang.srt.layers.attention.vision import VisionAttention
|
from sglang.srt.layers.attention.vision import VisionAttention
|
||||||
from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes
|
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.dp_attention import get_attention_tp_rank, get_attention_tp_size
|
||||||
from sglang.srt.layers.layernorm import RMSNorm
|
from sglang.srt.layers.layernorm import RMSNorm
|
||||||
from sglang.srt.layers.linear import (
|
from sglang.srt.layers.linear import (
|
||||||
@@ -96,7 +97,7 @@ class MossVLVisionPatchEmbed(nn.Module):
|
|||||||
self.embed_dim = config.hidden_size
|
self.embed_dim = config.hidden_size
|
||||||
|
|
||||||
kernel_size = [self.temporal_patch_size, self.patch_size, self.patch_size]
|
kernel_size = [self.temporal_patch_size, self.patch_size, self.patch_size]
|
||||||
self.proj = nn.Conv3d(
|
self.proj = Conv3dLayer(
|
||||||
self.in_channels,
|
self.in_channels,
|
||||||
self.embed_dim,
|
self.embed_dim,
|
||||||
kernel_size=kernel_size,
|
kernel_size=kernel_size,
|
||||||
@@ -1142,11 +1143,10 @@ class MossVLForConditionalGeneration(nn.Module):
|
|||||||
def _collect_mm_data(self, forward_batch: ForwardBatch):
|
def _collect_mm_data(self, forward_batch: ForwardBatch):
|
||||||
"""Collect pixel_values, grid_thw, and vision_position_ids from uncached requests."""
|
"""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):
|
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 = []
|
pixel_values_list = []
|
||||||
grid_thw_list = []
|
grid_thw_list = []
|
||||||
encoder_lens_need = []
|
|
||||||
vision_pos_ids_list = []
|
vision_pos_ids_list = []
|
||||||
|
|
||||||
for i, mm_input in enumerate(forward_batch.mm_inputs):
|
for i, mm_input in enumerate(forward_batch.mm_inputs):
|
||||||
@@ -1161,14 +1161,13 @@ class MossVLForConditionalGeneration(nn.Module):
|
|||||||
if grid_thw is not None:
|
if grid_thw is not None:
|
||||||
grid_thw_list.append(torch.as_tensor(grid_thw, dtype=torch.long))
|
grid_thw_list.append(torch.as_tensor(grid_thw, dtype=torch.long))
|
||||||
encoder_len = forward_batch.encoder_lens_cpu[i]
|
encoder_len = forward_batch.encoder_lens_cpu[i]
|
||||||
encoder_lens_need.append(encoder_len)
|
|
||||||
|
|
||||||
vp = mm_input.vision_position_ids
|
vp = mm_input.vision_position_ids
|
||||||
if vp is not None:
|
if vp is not None:
|
||||||
vision_pos_ids_list.append(vp[:, :encoder_len])
|
vision_pos_ids_list.append(vp[:, :encoder_len])
|
||||||
|
|
||||||
if not pixel_values_list:
|
if not pixel_values_list:
|
||||||
return None, None, None, None
|
return None, None, None
|
||||||
|
|
||||||
pixel_values = torch.cat(pixel_values_list, dim=0)
|
pixel_values = torch.cat(pixel_values_list, dim=0)
|
||||||
grid_thw = torch.cat(grid_thw_list, dim=0) if grid_thw_list else None
|
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
|
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(
|
def _get_vision_features(
|
||||||
self,
|
self,
|
||||||
@@ -1225,51 +1224,6 @@ class MossVLForConditionalGeneration(nn.Module):
|
|||||||
|
|
||||||
return torch.cat(output_parts, dim=0)
|
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) ----
|
# ---- prepare_forward_batch (called before attn backend init) ----
|
||||||
|
|
||||||
def prepare_forward_batch(self, forward_batch: ForwardBatch):
|
def prepare_forward_batch(self, forward_batch: ForwardBatch):
|
||||||
@@ -1509,8 +1463,8 @@ class MossVLForConditionalGeneration(nn.Module):
|
|||||||
positions = forward_batch.mrope_positions
|
positions = forward_batch.mrope_positions
|
||||||
|
|
||||||
# 1. Collect vision inputs for uncached requests
|
# 1. Collect vision inputs for uncached requests
|
||||||
pixel_values, grid_thw, encoder_lens_need, vision_position_ids = (
|
pixel_values, grid_thw, vision_position_ids = self._collect_mm_data(
|
||||||
self._collect_mm_data(forward_batch)
|
forward_batch
|
||||||
)
|
)
|
||||||
|
|
||||||
cross_attention_mask = None
|
cross_attention_mask = None
|
||||||
@@ -1534,14 +1488,12 @@ class MossVLForConditionalGeneration(nn.Module):
|
|||||||
if pixel_values is not None and grid_thw is not None:
|
if pixel_values is not None and grid_thw is not None:
|
||||||
# Run ViT
|
# Run ViT
|
||||||
vision_hidden_states = self._get_vision_features(pixel_values, grid_thw)
|
vision_hidden_states = self._get_vision_features(pixel_values, grid_thw)
|
||||||
# Insert separator tokens after each frame
|
# Insert separator tokens after each frame. The result is already
|
||||||
vision_with_sep = self._insert_separator_tokens(
|
# 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
|
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
|
# Drop heavy per-request vision tensors now that the encoder KV
|
||||||
# has been produced and will be cached. Otherwise pixel_values and
|
# has been produced and will be cached. Otherwise pixel_values and
|
||||||
# vision_position_ids stay pinned on req.multimodal_inputs across
|
# 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
|
# 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
|
# cross-attention below, so we keep it; but we drop the per-request
|
||||||
# copy on mm_input, which we won't read again.
|
# 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):
|
for i, mm_input in enumerate(forward_batch.mm_inputs):
|
||||||
if forward_batch.encoder_cached[i] or mm_input is None:
|
if forward_batch.encoder_cached[i] or mm_input is None:
|
||||||
continue
|
continue
|
||||||
|
|||||||
Reference in New Issue
Block a user