[VLM] Introduce Cache for positional embedding ids for Qwen-VL family (#14292)
Co-authored-by: luoyuan.luo <luoyuan.luo@antgroup.com>
This commit is contained in:
@@ -50,7 +50,7 @@ from sglang.srt.managers.schedule_batch import (
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||
from sglang.srt.models.qwen3 import Qwen3Model
|
||||
from sglang.srt.models.utils import compute_cu_seqlens_from_grid_numpy
|
||||
from sglang.srt.models.utils import RotaryPosMixin, compute_cu_seqlens_from_grid_numpy
|
||||
from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model
|
||||
from sglang.srt.server_args import get_global_server_args
|
||||
from sglang.srt.utils import add_prefix
|
||||
@@ -257,7 +257,7 @@ class Qwen3VLMoeVisionPatchMerger(nn.Module):
|
||||
return out
|
||||
|
||||
|
||||
class Qwen3VLMoeVisionModel(nn.Module):
|
||||
class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -339,26 +339,9 @@ class Qwen3VLMoeVisionModel(nn.Module):
|
||||
def rot_pos_emb(self, grid_thw):
|
||||
pos_ids = []
|
||||
for t, h, w in grid_thw:
|
||||
hpos_ids = torch.arange(h).unsqueeze(1).expand(-1, w)
|
||||
hpos_ids = hpos_ids.reshape(
|
||||
h // self.spatial_merge_size,
|
||||
self.spatial_merge_size,
|
||||
w // self.spatial_merge_size,
|
||||
self.spatial_merge_size,
|
||||
)
|
||||
hpos_ids = hpos_ids.permute(0, 2, 1, 3)
|
||||
hpos_ids = hpos_ids.flatten()
|
||||
base = self.rot_pos_ids(h, w, self.spatial_merge_size)
|
||||
pos_ids.append(base if t == 1 else base.repeat(t, 1))
|
||||
|
||||
wpos_ids = torch.arange(w).unsqueeze(0).expand(h, -1)
|
||||
wpos_ids = wpos_ids.reshape(
|
||||
h // self.spatial_merge_size,
|
||||
self.spatial_merge_size,
|
||||
w // self.spatial_merge_size,
|
||||
self.spatial_merge_size,
|
||||
)
|
||||
wpos_ids = wpos_ids.permute(0, 2, 1, 3)
|
||||
wpos_ids = wpos_ids.flatten()
|
||||
pos_ids.append(torch.stack([hpos_ids, wpos_ids], dim=-1).repeat(t, 1))
|
||||
pos_ids = torch.cat(pos_ids, dim=0)
|
||||
max_grid_size = grid_thw[:, 1:].max()
|
||||
rotary_pos_emb_full = self.rotary_pos_emb(max_grid_size)
|
||||
|
||||
Reference in New Issue
Block a user