Reworked fast_pos_embed_interpolate() using torch (#10959)

This commit is contained in:
Vitaly Tuzov
2025-12-30 09:45:34 +03:00
committed by GitHub
parent 8a84b1e7e0
commit 1048803c1f
6 changed files with 143 additions and 91 deletions

View File

@@ -626,7 +626,7 @@ class VisionAttention(nn.Module):
prefix=add_prefix("proj", prefix),
)
self.aux_stream = aux_stream
self.ln_events = [torch.cuda.Event(), torch.cuda.Event()]
self.ln_events = [torch.cuda.Event(), torch.cuda.Event()] if aux_stream else []
def _determine_attention_backend(self, passed_backend: Optional[str]) -> str:
"""Decide the multimodal attention backend string.
@@ -693,8 +693,7 @@ class VisionAttention(nn.Module):
q, k = maybe_execute_in_parallel(
q_l2norm,
k_l2norm,
self.ln_events[0],
self.ln_events[1],
self.ln_events,
self.aux_stream,
)
return q, k

View File

@@ -19,7 +19,6 @@ import re
from functools import lru_cache, partial
from typing import Callable, Iterable, List, Optional, Tuple, Union
import numpy as np
import torch
import torch.nn as nn
from einops import rearrange
@@ -282,6 +281,11 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin):
self.hidden_size = vision_config.hidden_size
self.num_heads = vision_config.num_heads
self.num_position_embeddings = vision_config.num_position_embeddings
self.num_grid_per_side = int(self.num_position_embeddings**0.5)
self.num_grid = self.num_grid_per_side * self.num_grid_per_side
self.align_corners = (
get_global_server_args().enable_precise_embedding_interpolation
)
self.patch_size = vision_config.patch_size
self.spatial_merge_size = vision_config.spatial_merge_size
self.spatial_merge_unit = self.spatial_merge_size**2
@@ -378,89 +382,30 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin):
return cos_combined, sin_combined
def fast_pos_embed_interpolate(self, grid_thw):
num_grid_per_side = int(self.num_position_embeddings**0.5)
idx_list = [[] for _ in range(4)]
weight_list = [[] for _ in range(4)]
# TODO: use torch instand of np
for t, h, w in grid_thw:
h_idxs = np.linspace(0, num_grid_per_side - 1, h)
w_idxs = np.linspace(0, num_grid_per_side - 1, w)
h_idxs_floor = h_idxs.astype(int)
w_idxs_floor = w_idxs.astype(int)
h_idxs_ceil = (h_idxs.astype(int) + 1).clip(max=num_grid_per_side - 1)
w_idxs_ceil = (w_idxs.astype(int) + 1).clip(max=num_grid_per_side - 1)
dh = h_idxs - h_idxs_floor
dw = w_idxs - w_idxs_floor
idx_list[0].extend(
((h_idxs_floor * num_grid_per_side)[None].T + w_idxs_floor[None])
.flatten()
.tolist()
* t
)
idx_list[1].extend(
((h_idxs_floor * num_grid_per_side)[None].T + w_idxs_ceil[None])
.flatten()
.tolist()
* t
)
idx_list[2].extend(
((h_idxs_ceil * num_grid_per_side)[None].T + w_idxs_floor[None])
.flatten()
.tolist()
* t
)
idx_list[3].extend(
((h_idxs_ceil * num_grid_per_side)[None].T + w_idxs_ceil[None])
.flatten()
.tolist()
* t
)
weight_list[0].extend(
((1 - dh)[None].T * (1 - dw)[None]).flatten().tolist() * t
)
weight_list[1].extend(((1 - dh)[None].T * dw[None]).flatten().tolist() * t)
weight_list[2].extend((dh[None].T * (1 - dw)[None]).flatten().tolist() * t)
weight_list[3].extend((dh[None].T * dw[None]).flatten().tolist() * t)
device = self.pos_embed.weight.device
dtype = self.pos_embed.weight.dtype
p0 = (
self.pos_embed(torch.tensor(idx_list[0], dtype=torch.long, device=device))
* torch.tensor(weight_list[0], dtype=dtype, device=device)[:, None]
)
p1 = (
self.pos_embed(torch.tensor(idx_list[1], dtype=torch.long, device=device))
* torch.tensor(weight_list[1], dtype=dtype, device=device)[:, None]
)
p2 = (
self.pos_embed(torch.tensor(idx_list[2], dtype=torch.long, device=device))
* torch.tensor(weight_list[2], dtype=dtype, device=device)[:, None]
)
p3 = (
self.pos_embed(torch.tensor(idx_list[3], dtype=torch.long, device=device))
* torch.tensor(weight_list[3], dtype=dtype, device=device)[:, None]
)
patch_pos_embeds = p0 + p1 + p2 + p3
patch_pos_embeds = patch_pos_embeds.split([t * h * w for t, h, w in grid_thw])
patch_pos_embeds_permute = []
m_size = self.spatial_merge_size
for pos_embed, (t, h, w) in zip(patch_pos_embeds, grid_thw):
pos_embed = (
pos_embed.view(t, h // m_size, m_size, w // m_size, m_size, -1)
.permute(0, 1, 3, 2, 4, 5)
.flatten(0, 4)
embeds = torch.arange(self.num_grid, device=self.pos_embed.weight.device)
embeds = (
self.pos_embed(embeds)
.permute(1, 0)
.reshape(1, -1, self.num_grid_per_side, self.num_grid_per_side)
)
for t, h, w in grid_thw:
pos_embed = torch.nn.functional.interpolate(
embeds, size=(h, w), mode="bilinear", align_corners=self.align_corners
)
pos_embed = pos_embed.reshape(
-1,
h // self.spatial_merge_size,
self.spatial_merge_size,
w // self.spatial_merge_size,
self.spatial_merge_size,
)
pos_embed = pos_embed.permute(1, 3, 2, 4, 0)
pos_embed = pos_embed.flatten(0, 3).repeat(t, 1)
patch_pos_embeds_permute.append(pos_embed)
patch_pos_embeds = torch.cat(patch_pos_embeds_permute)
return patch_pos_embeds
return torch.cat(patch_pos_embeds_permute)
def forward(
self,

View File

@@ -579,6 +579,7 @@ class ServerArgs:
# Context parallelism used in the long sequence prefill phase of DeepSeek v3.2
enable_nsa_prefill_context_parallel: bool = False
enable_fused_qk_norm_rope: bool = False
enable_precise_embedding_interpolation: bool = False
# Dynamic batch tokenizer
enable_dynamic_batch_tokenizer: bool = False
@@ -4191,6 +4192,11 @@ class ServerArgs:
action="store_true",
help="Enable fused qk normalization and rope rotary embedding.",
)
parser.add_argument(
"--enable-precise-embedding-interpolation",
action="store_true",
help="Enable corner alignment for resize of embeddings grid to ensure more accurate(but slower) evaluation of interpolated embedding values.",
)
# Dynamic batch tokenizer
parser.add_argument(

View File

@@ -37,8 +37,7 @@ def with_multi_stream(enable: bool):
def maybe_execute_in_parallel(
fn0: Callable,
fn1: Callable,
event0: torch.cuda.Event,
event1: torch.cuda.Event,
events: list[torch.cuda.Event],
aux_stream: Optional[torch.cuda.Stream] = None,
) -> tuple[Any, Any]:
"""Utility function to run two functions in two cuda streams in parallel. Multi-stream is
@@ -51,8 +50,7 @@ def maybe_execute_in_parallel(
Args:
fn0 (Callable): callable for the default stream
fn1 (Callable): callable for the second stream, aux_stream
event0 (torch.cuda.Event): cuda event for fn0
event1 (torch.cuda.Event): cuda event for fn1
events (list[torch.cuda.Event]): cuda events for callables
aux_stream (Optional[torch.cuda.Stream]): the second cuda stream for fn1.
Multi-stream is disabled when aux_stream is None.
@@ -63,14 +61,14 @@ def maybe_execute_in_parallel(
multi_stream = do_multi_stream() and aux_stream is not None
if multi_stream:
event0.record()
events[0].record()
result0 = fn0()
with torch.cuda.stream(aux_stream):
event0.wait()
events[0].wait()
result1 = fn1()
event1.record()
event1.wait()
events[1].record()
events[1].wait()
else:
result0 = fn0()
result1 = fn1()