From 8a82c70297fdf1a24c4e32d9111f0299d389d9c0 Mon Sep 17 00:00:00 2001 From: Yuan Luo Date: Mon, 16 Feb 2026 11:19:44 +0800 Subject: [PATCH] [VLM] Optimize Ernie4.5-VL rotary embedding with fused triton kernel (#18856) Co-authored-by: luoyuan.luo --- python/sglang/srt/layers/rotary_embedding.py | 271 ++++++++++++++++++- 1 file changed, 268 insertions(+), 3 deletions(-) diff --git a/python/sglang/srt/layers/rotary_embedding.py b/python/sglang/srt/layers/rotary_embedding.py index 4ff07a6b0..4a881200f 100644 --- a/python/sglang/srt/layers/rotary_embedding.py +++ b/python/sglang/srt/layers/rotary_embedding.py @@ -2788,9 +2788,238 @@ class YaRNScalingMRotaryEmbedding(MRotaryEmbedding): return cache +@triton.jit +def _triton_ernie45_rope_qk_fused( + q_ptr, + k_ptr, + cos_sin_cache_ptr, + positions_ptr, # [3, num_tokens] (t/h/w) + q_stride0: tl.constexpr, + k_stride0: tl.constexpr, + pos_stride0: tl.constexpr, # positions.stride(0) + n_qh: tl.constexpr, + n_kh: tl.constexpr, + hd: tl.constexpr, + rd: tl.constexpr, # rotary_dim + pad_n_qh: tl.constexpr, + pad_n_kh: tl.constexpr, + pad_hd: tl.constexpr, + section_hw: tl.constexpr, # section_h + section_w (Ernie: 2*section_h) + is_neox_style: tl.constexpr, +): + pid = tl.program_id(0) # token id + q_ptr = q_ptr + pid * q_stride0 + k_ptr = k_ptr + pid * k_stride0 + + half_rd = rd // 2 + + # positions: [3, num_tokens] => (t, h, w) + tpos = tl.load(positions_ptr + 0 * pos_stride0 + pid).to(tl.int32) + hpos = tl.load(positions_ptr + 1 * pos_stride0 + pid).to(tl.int32) + wpos = tl.load(positions_ptr + 2 * pos_stride0 + pid).to(tl.int32) + + # rotary pair index vector [0 .. pad_hd/2) + ridx = tl.arange(0, pad_hd // 2) + rmask = ridx < half_rd + + # Choose which axis position to use for each ridx + # ridx < section_hw: even->hpos, odd->wpos ; else -> tpos + use_hw = ridx < section_hw + use_h = (ridx & 1) == 0 + pos = tl.where(use_hw, tl.where(use_h, hpos, wpos), tpos) + + # Load cos/sin for each ridx from cache[pos, :] + # cache row stride is rd (cos first half, sin second half) + cos = tl.load(cos_sin_cache_ptr + pos * rd + ridx, mask=rmask, other=0.0) + sin = tl.load( + cos_sin_cache_ptr + pos * rd + (ridx + half_rd), + mask=rmask, + other=0.0, + ) + + # Apply to Q/K in-place. + # Q: [n_qh, hd], K: [n_kh, hd], but stored flattened [num_heads * hd] + if is_neox_style: + # Load first half / second half of rotary dim (size half_rd) + q_head = tl.arange(0, pad_n_qh)[:, None] + k_head = tl.arange(0, pad_n_kh)[:, None] + + d = tl.arange(0, pad_hd // 2)[None, :] + + q_mask = (q_head < n_qh) & (d < half_rd) + k_mask = (k_head < n_kh) & (d < half_rd) + + # offsets for first half within each head + q_off0 = q_head * hd + d + k_off0 = k_head * hd + d + + # offsets for second half within each head (shift by half_rd within rotary dim) + q_off1 = q_off0 + half_rd + k_off1 = k_off0 + half_rd + + # Load + q0 = tl.load(q_ptr + q_off0, mask=q_mask, other=0.0).to(cos.dtype) + q1 = tl.load(q_ptr + q_off1, mask=q_mask, other=0.0).to(cos.dtype) + k0 = tl.load(k_ptr + k_off0, mask=k_mask, other=0.0).to(cos.dtype) + k1 = tl.load(k_ptr + k_off1, mask=k_mask, other=0.0).to(cos.dtype) + + # Broadcast cos/sin to [heads, half_rd] + cos_b = cos[None, :] + sin_b = sin[None, :] + + # Rotate + nq0 = q0 * cos_b - q1 * sin_b + nq1 = q1 * cos_b + q0 * sin_b + nk0 = k0 * cos_b - k1 * sin_b + nk1 = k1 * cos_b + k0 * sin_b + + # Store back + tl.store(q_ptr + q_off0, nq0, mask=q_mask) + tl.store(q_ptr + q_off1, nq1, mask=q_mask) + tl.store(k_ptr + k_off0, nk0, mask=k_mask) + tl.store(k_ptr + k_off1, nk1, mask=k_mask) + + else: + # GPT-J style: pairs are (even, odd) within rotary_dim + q_head = tl.arange(0, pad_n_qh)[:, None] + k_head = tl.arange(0, pad_n_kh)[:, None] + p = tl.arange(0, pad_hd // 2)[None, :] # pair index + + q_mask = (q_head < n_qh) & (p < half_rd) + k_mask = (k_head < n_kh) & (p < half_rd) + + even = 2 * p + odd = even + 1 + + q_even_off = q_head * hd + even + q_odd_off = q_head * hd + odd + k_even_off = k_head * hd + even + k_odd_off = k_head * hd + odd + + q_even = tl.load(q_ptr + q_even_off, mask=q_mask, other=0.0).to(cos.dtype) + q_odd = tl.load(q_ptr + q_odd_off, mask=q_mask, other=0.0).to(cos.dtype) + k_even = tl.load(k_ptr + k_even_off, mask=k_mask, other=0.0).to(cos.dtype) + k_odd = tl.load(k_ptr + k_odd_off, mask=k_mask, other=0.0).to(cos.dtype) + + cos_b = cos[None, :] + sin_b = sin[None, :] + + nq_even = q_even * cos_b - q_odd * sin_b + nq_odd = q_odd * cos_b + q_even * sin_b + nk_even = k_even * cos_b - k_odd * sin_b + nk_odd = k_odd * cos_b + k_even * sin_b + + tl.store(q_ptr + q_even_off, nq_even, mask=q_mask) + tl.store(q_ptr + q_odd_off, nq_odd, mask=q_mask) + tl.store(k_ptr + k_even_off, nk_even, mask=k_mask) + tl.store(k_ptr + k_odd_off, nk_odd, mask=k_mask) + + +def triton_ernie45_rope_fused_inplace( + q: torch.Tensor, # [num_tokens, n_qh*hd], contiguous + k: torch.Tensor, # [num_tokens, n_kh*hd], contiguous + cos_sin_cache: torch.Tensor, # [max_pos, rd], contiguous + positions: torch.Tensor, # [3, num_tokens], contiguous, stride(1)==1 + mrope_section: list[int], # [h, w, t] (Ernie expects h==w) + head_size: int, + rotary_dim: int, + is_neox_style: bool, +) -> None: + assert q.is_cuda and k.is_cuda and cos_sin_cache.is_cuda and positions.is_cuda + assert q.dim() == 2 and k.dim() == 2 + assert positions.dim() == 2 and positions.shape[0] == 3 + assert q.stride(1) == 1 and k.stride(1) == 1 + assert positions.stride(1) == 1 + assert cos_sin_cache.dim() == 2 and cos_sin_cache.is_contiguous() + + num_tokens = q.shape[0] + assert positions.shape[1] == num_tokens + assert q.shape[0] == k.shape[0] == num_tokens + + n_q_dim = q.shape[1] + n_k_dim = k.shape[1] + assert n_q_dim % head_size == 0 and n_k_dim % head_size == 0 + + n_qh = n_q_dim // head_size + n_kh = n_k_dim // head_size + + rd = rotary_dim + assert rd % 2 == 0 + assert rd <= head_size + + # Ernie section sanity + section_h, section_w, section_t = mrope_section + assert section_h == section_w, "Ernie4.5 layout assumes section_h == section_w" + assert section_h + section_w + section_t == ( + rd // 2 + ), "mrope_section must sum to rotary_dim//2" + + # Ensure cache dtype matches q/k dtype for best perf (avoid implicit casts) + if cos_sin_cache.dtype != q.dtype or cos_sin_cache.device != q.device: + cos_sin_cache = cos_sin_cache.to(device=q.device, dtype=q.dtype) + + pad_n_qh = triton.next_power_of_2(n_qh) + pad_n_kh = triton.next_power_of_2(n_kh) + pad_hd = triton.next_power_of_2(head_size) + + # Heuristic warps + num_warps = 4 if (pad_n_qh * pad_hd) <= 8192 else 8 + + _triton_ernie45_rope_qk_fused[(num_tokens,)]( + q, + k, + cos_sin_cache, + positions, + q.stride(0), + k.stride(0), + positions.stride(0), + n_qh=n_qh, + n_kh=n_kh, + hd=head_size, + rd=rd, + pad_n_qh=pad_n_qh, + pad_n_kh=pad_n_kh, + pad_hd=pad_hd, + section_hw=section_h + section_w, + is_neox_style=is_neox_style, + num_warps=num_warps, + ) + + class Ernie4_5_VLRotaryEmbedding(MRotaryEmbedding): """3D rotary positional embedding. [h w h w h w h w... t t t...]""" + def __init__( + self, + head_size: int, + rotary_dim: int, + max_position_embeddings: int, + base: int, + is_neox_style: bool, + dtype: torch.dtype, + mrope_section: Optional[List[int]] = None, + mrope_interleaved: bool = False, + ) -> None: + super().__init__( + head_size, + rotary_dim, + max_position_embeddings, + base, + is_neox_style, + dtype, + mrope_section=mrope_section, + mrope_interleaved=mrope_interleaved, + ) + self.head_size = head_size + self.rotary_dim = rotary_dim + self.max_position_embeddings = max_position_embeddings + self.base = base + self.is_neox_style = is_neox_style + self.dtype = dtype + self.mrope_section = mrope_section + self.mrope_interleaved = mrope_interleaved + self._apply_rotary_emb_wrapped = torch.compile(dynamic=True)(_apply_rotary_emb) + def forward_native( # type: ignore[override] self, positions: torch.Tensor, @@ -2835,14 +3064,16 @@ class Ernie4_5_VLRotaryEmbedding(MRotaryEmbedding): query = query.view(num_tokens, -1, self.head_size) query_rot = query[..., : self.rotary_dim] query_pass = query[..., self.rotary_dim :] - query_rot = _apply_rotary_emb(query_rot, cos, sin, self.is_neox_style) + query_rot = self._apply_rotary_emb_wrapped( + query_rot, cos, sin, self.is_neox_style + ) query = torch.cat((query_rot, query_pass), dim=-1).reshape(query_shape) key_shape = key.shape key = key.view(num_tokens, -1, self.head_size) key_rot = key[..., : self.rotary_dim] key_pass = key[..., self.rotary_dim :] - key_rot = _apply_rotary_emb(key_rot, cos, sin, self.is_neox_style) + key_rot = self._apply_rotary_emb_wrapped(key_rot, cos, sin, self.is_neox_style) key = torch.cat((key_rot, key_pass), dim=-1).reshape(key_shape) return query, key @@ -2852,6 +3083,40 @@ class Ernie4_5_VLRotaryEmbedding(MRotaryEmbedding): query: torch.Tensor, key: torch.Tensor | None = None, ) -> tuple[torch.Tensor, torch.Tensor | None]: + assert key is not None + assert positions.ndim in (1, 2) + + # Ensure cache dtype/device matches q/k to avoid extra casts + self._match_cos_sin_cache_dtype(query) + + if positions.ndim == 2: + assert self.mrope_section is not None + # positions: [3, num_tokens] + triton_ernie45_rope_fused_inplace( + q=query, + k=key, + cos_sin_cache=self.cos_sin_cache, + positions=positions, + mrope_section=self.mrope_section, # [h, w, t] + head_size=self.head_size, + rotary_dim=self.rotary_dim, + is_neox_style=self.is_neox_style, + ) + return query, key + + # positions.ndim == 1 (text-only): use existing fused kernel if available + if _is_cuda and (apply_rope_with_cos_sin_cache_inplace is not None): + apply_rope_with_cos_sin_cache_inplace( + positions=positions, + query=query, + key=key, + head_size=self.head_size, + cos_sin_cache=self.cos_sin_cache, + is_neox=self.is_neox_style, + ) + return query, key + + # fallback return self.forward_native(positions, query, key) def forward( @@ -2870,7 +3135,7 @@ class Ernie4_5_VLRotaryEmbedding(MRotaryEmbedding): key: [num_tokens, num_kv_heads * head_size] """ assert positions.ndim == 1 or positions.ndim == 2 - return self.forward_native(positions, query, key) + return self.forward_cuda(positions, query, key) class DualChunkRotaryEmbedding(MultiPlatformOp):