Reworked fast_pos_embed_interpolate() using torch (#10959)

This commit is contained in:
Vitaly Tuzov
2025-12-30 14:45:34 +08:00
committed by GitHub
parent 8a84b1e7e0
commit 1048803c1f
6 changed files with 143 additions and 91 deletions
+2 -3
View File
@@ -626,7 +626,7 @@ class VisionAttention(nn.Module):
prefix=add_prefix("proj", prefix), prefix=add_prefix("proj", prefix),
) )
self.aux_stream = aux_stream 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: def _determine_attention_backend(self, passed_backend: Optional[str]) -> str:
"""Decide the multimodal attention backend string. """Decide the multimodal attention backend string.
@@ -693,8 +693,7 @@ class VisionAttention(nn.Module):
q, k = maybe_execute_in_parallel( q, k = maybe_execute_in_parallel(
q_l2norm, q_l2norm,
k_l2norm, k_l2norm,
self.ln_events[0], self.ln_events,
self.ln_events[1],
self.aux_stream, self.aux_stream,
) )
return q, k return q, k
+25 -80
View File
@@ -19,7 +19,6 @@ import re
from functools import lru_cache, partial from functools import lru_cache, partial
from typing import Callable, Iterable, List, Optional, Tuple, Union from typing import Callable, Iterable, List, Optional, Tuple, Union
import numpy as np
import torch import torch
import torch.nn as nn import torch.nn as nn
from einops import rearrange from einops import rearrange
@@ -282,6 +281,11 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin):
self.hidden_size = vision_config.hidden_size self.hidden_size = vision_config.hidden_size
self.num_heads = vision_config.num_heads self.num_heads = vision_config.num_heads
self.num_position_embeddings = vision_config.num_position_embeddings 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.patch_size = vision_config.patch_size
self.spatial_merge_size = vision_config.spatial_merge_size self.spatial_merge_size = vision_config.spatial_merge_size
self.spatial_merge_unit = self.spatial_merge_size**2 self.spatial_merge_unit = self.spatial_merge_size**2
@@ -378,89 +382,30 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin):
return cos_combined, sin_combined return cos_combined, sin_combined
def fast_pos_embed_interpolate(self, grid_thw): 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 = [] patch_pos_embeds_permute = []
m_size = self.spatial_merge_size m_size = self.spatial_merge_size
for pos_embed, (t, h, w) in zip(patch_pos_embeds, grid_thw):
pos_embed = ( embeds = torch.arange(self.num_grid, device=self.pos_embed.weight.device)
pos_embed.view(t, h // m_size, m_size, w // m_size, m_size, -1) embeds = (
.permute(0, 1, 3, 2, 4, 5) self.pos_embed(embeds)
.flatten(0, 4) .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_permute.append(pos_embed)
patch_pos_embeds = torch.cat(patch_pos_embeds_permute) return torch.cat(patch_pos_embeds_permute)
return patch_pos_embeds
def forward( def forward(
self, self,
+6
View File
@@ -579,6 +579,7 @@ class ServerArgs:
# Context parallelism used in the long sequence prefill phase of DeepSeek v3.2 # Context parallelism used in the long sequence prefill phase of DeepSeek v3.2
enable_nsa_prefill_context_parallel: bool = False enable_nsa_prefill_context_parallel: bool = False
enable_fused_qk_norm_rope: bool = False enable_fused_qk_norm_rope: bool = False
enable_precise_embedding_interpolation: bool = False
# Dynamic batch tokenizer # Dynamic batch tokenizer
enable_dynamic_batch_tokenizer: bool = False enable_dynamic_batch_tokenizer: bool = False
@@ -4191,6 +4192,11 @@ class ServerArgs:
action="store_true", action="store_true",
help="Enable fused qk normalization and rope rotary embedding.", 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 # Dynamic batch tokenizer
parser.add_argument( parser.add_argument(
@@ -37,8 +37,7 @@ def with_multi_stream(enable: bool):
def maybe_execute_in_parallel( def maybe_execute_in_parallel(
fn0: Callable, fn0: Callable,
fn1: Callable, fn1: Callable,
event0: torch.cuda.Event, events: list[torch.cuda.Event],
event1: torch.cuda.Event,
aux_stream: Optional[torch.cuda.Stream] = None, aux_stream: Optional[torch.cuda.Stream] = None,
) -> tuple[Any, Any]: ) -> tuple[Any, Any]:
"""Utility function to run two functions in two cuda streams in parallel. Multi-stream is """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: Args:
fn0 (Callable): callable for the default stream fn0 (Callable): callable for the default stream
fn1 (Callable): callable for the second stream, aux_stream fn1 (Callable): callable for the second stream, aux_stream
event0 (torch.cuda.Event): cuda event for fn0 events (list[torch.cuda.Event]): cuda events for callables
event1 (torch.cuda.Event): cuda event for fn1
aux_stream (Optional[torch.cuda.Stream]): the second cuda stream for fn1. aux_stream (Optional[torch.cuda.Stream]): the second cuda stream for fn1.
Multi-stream is disabled when aux_stream is None. 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 multi_stream = do_multi_stream() and aux_stream is not None
if multi_stream: if multi_stream:
event0.record() events[0].record()
result0 = fn0() result0 = fn0()
with torch.cuda.stream(aux_stream): with torch.cuda.stream(aux_stream):
event0.wait() events[0].wait()
result1 = fn1() result1 = fn1()
event1.record() events[1].record()
event1.wait() events[1].wait()
else: else:
result0 = fn0() result0 = fn0()
result1 = fn1() result1 = fn1()
+1
View File
@@ -334,6 +334,7 @@ suite_ascend = {
TestFile("ascend/test_ascend_sampling_backend.py", 400), TestFile("ascend/test_ascend_sampling_backend.py", 400),
TestFile("ascend/test_ascend_tp1_bf16.py", 400), TestFile("ascend/test_ascend_tp1_bf16.py", 400),
TestFile("ascend/test_ascend_compile_graph_tp1_bf16.py", 400), TestFile("ascend/test_ascend_compile_graph_tp1_bf16.py", 400),
TestFile("test_embed_interpolate_unittest.py", 400),
], ],
"per-commit-2-npu-a2": [ "per-commit-2-npu-a2": [
TestFile("ascend/test_ascend_graph_tp2_bf16.py", 400), TestFile("ascend/test_ascend_graph_tp2_bf16.py", 400),
+103
View File
@@ -0,0 +1,103 @@
import unittest
import torch
from sglang.srt.configs.qwen3_vl import Qwen3VLConfig
from sglang.srt.distributed.parallel_state import (
init_distributed_environment,
initialize_model_parallel,
)
from sglang.srt.layers.dp_attention import initialize_dp_attention
from sglang.srt.layers.quantization.unquant import (
LinearMethodBase,
UnquantizedLinearMethod,
)
from sglang.srt.models.qwen3_vl import Qwen3VLMoeVisionModel
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
def unpack(tensor, dim_len, pack_len):
dim_part = dim_len // pack_len
ret_val = tensor.reshape(dim_part, dim_part, pack_len, pack_len, -1)
ret_val = ret_val.permute(4, 0, 2, 1, 3).reshape(1, -1, dim_len, dim_len)
return ret_val
class TestEmbedInterpolate(unittest.TestCase):
@classmethod
def setUpClass(cls):
cls.pDevice = torch.get_default_device()
torch.set_default_device("npu")
@classmethod
def tearDownClass(cls):
torch.set_default_device(cls.pDevice)
def test_embed_interpolate(self):
self.assertTrue(issubclass(UnquantizedLinearMethod, LinearMethodBase))
t_dim = [16, 32]
s_dim = [192, 574]
sarg = ServerArgs(model_path="dummy", device="npu")
mconf = Qwen3VLConfig(
hidden_size=64,
num_heads=1,
num_position_embeddings=2304,
patch_size=16,
spatial_merge_size=2,
temporal_patch_size=2,
deepstack_visual_indexes=[5, 11, 17],
in_channels=3,
depth=24,
intermediate_size=256,
hidden_act="gelu_pytorch_tanh",
out_hidden_size=2560,
)
set_global_server_args_for_scheduler(sarg)
init_distributed_environment(
backend="gloo",
world_size=1,
rank=0,
local_rank=0,
distributed_init_method="tcp://127.0.0.1:2646",
)
initialize_model_parallel()
initialize_dp_attention(
server_args=sarg,
model_config=mconf,
)
model = Qwen3VLMoeVisionModel(
mconf,
quant_config=None,
norm_eps=1e-6,
prefix="visual",
)
embeddings = model.fast_pos_embed_interpolate(
[(t, s, s) for t, s in zip(t_dim, s_dim)]
)
embeddings_s0 = embeddings[: s_dim[0] * s_dim[0], :]
embeddings_s1 = embeddings[s_dim[0] * s_dim[0] : 2 * s_dim[0] * s_dim[0], :]
self.assertTrue(torch.allclose(embeddings_s0, embeddings_s1, atol=5e-5))
embeddings_l = embeddings[
t_dim[0] * s_dim[0] * s_dim[0] : t_dim[0] * s_dim[0] * s_dim[0]
+ s_dim[1] * s_dim[1],
:,
]
embeddings_s0 = torch.nn.functional.interpolate(
unpack(embeddings_s0, s_dim[0], 2),
size=(48, 48),
mode="area",
)
embeddings_r = torch.nn.functional.interpolate(
unpack(embeddings_l, s_dim[1], 2),
size=(48, 48),
mode="area",
)
self.assertTrue(
torch.allclose(embeddings_s0, embeddings_r, atol=5e-1, rtol=5e-1)
)
if __name__ == "__main__":
unittest.main()