[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:
Yuan Luo
2025-12-04 12:32:00 +08:00
committed by GitHub
co-authored by luoyuan.luo
parent 04df80a9a1
commit b2b09f5f24
3 changed files with 47 additions and 46 deletions
+4 -21
View File
@@ -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)