From 1048803c1fc502079046b0fd28ecef8213cda430 Mon Sep 17 00:00:00 2001 From: Vitaly Tuzov <13571155+terfendail@users.noreply.github.com> Date: Tue, 30 Dec 2025 09:45:34 +0300 Subject: [PATCH] Reworked fast_pos_embed_interpolate() using torch (#10959) --- python/sglang/srt/layers/attention/vision.py | 5 +- python/sglang/srt/models/qwen3_vl.py | 105 +++++------------- python/sglang/srt/server_args.py | 6 + python/sglang/srt/utils/multi_stream_utils.py | 14 +-- test/srt/run_suite.py | 1 + test/srt/test_embed_interpolate_unittest.py | 103 +++++++++++++++++ 6 files changed, 143 insertions(+), 91 deletions(-) create mode 100644 test/srt/test_embed_interpolate_unittest.py diff --git a/python/sglang/srt/layers/attention/vision.py b/python/sglang/srt/layers/attention/vision.py index 1bec0fe63..7ad15e7ec 100644 --- a/python/sglang/srt/layers/attention/vision.py +++ b/python/sglang/srt/layers/attention/vision.py @@ -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 diff --git a/python/sglang/srt/models/qwen3_vl.py b/python/sglang/srt/models/qwen3_vl.py index 1c3b46d2f..bdf4c8312 100644 --- a/python/sglang/srt/models/qwen3_vl.py +++ b/python/sglang/srt/models/qwen3_vl.py @@ -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, diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index d506110ac..3022bae40 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -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( diff --git a/python/sglang/srt/utils/multi_stream_utils.py b/python/sglang/srt/utils/multi_stream_utils.py index fa5c55837..fb3d2cf7e 100644 --- a/python/sglang/srt/utils/multi_stream_utils.py +++ b/python/sglang/srt/utils/multi_stream_utils.py @@ -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() diff --git a/test/srt/run_suite.py b/test/srt/run_suite.py index 4f286703b..ffe8db9b7 100644 --- a/test/srt/run_suite.py +++ b/test/srt/run_suite.py @@ -334,6 +334,7 @@ suite_ascend = { TestFile("ascend/test_ascend_sampling_backend.py", 400), TestFile("ascend/test_ascend_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": [ TestFile("ascend/test_ascend_graph_tp2_bf16.py", 400), diff --git a/test/srt/test_embed_interpolate_unittest.py b/test/srt/test_embed_interpolate_unittest.py new file mode 100644 index 000000000..cb09935bc --- /dev/null +++ b/test/srt/test_embed_interpolate_unittest.py @@ -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()