Reworked fast_pos_embed_interpolate() using torch (#10959)
This commit is contained in:
@@ -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
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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),
|
||||||
|
|||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user