[VLM] Optimize Ernie4.5-VL rotary embedding with fused triton kernel (#18856)

Co-authored-by: luoyuan.luo <luoyuan.luo@antgroup.com>
This commit is contained in:
Yuan Luo
2026-02-16 11:19:44 +08:00
committed by GitHub
parent 45715af50c
commit 8a82c70297

View File

@@ -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):