changh95's picture
Add files using upload-large-folder tool
0190e6b verified
Raw History Blame Contribute Delete
2.08 kB
"""Host-side rotary embedding helpers for the Qwen3-VL text model (interleaved mrope).
Mirrors transformers 4.57 `Qwen3VLTextRotaryEmbedding` + `apply_interleaved_mrope`:
freqs[axis] = inv_freq * pos[axis] -> [3, L, head_dim/2]
out = freqs[t]; out[..., 1::3][:mrope_section[1]] = freqs[h]; out[..., 2::3][:mrope_section[2]] = freqs[w]
emb = cat(out, out) ; cos/sin -> [L, head_dim] (HF rotate_half convention)
All math in float32 on host; the device consumes bf16 cos/sin.
"""
from __future__ import annotations
import torch
def inv_freq(head_dim: int, theta: float) -> torch.Tensor:
return 1.0 / (theta ** (torch.arange(0, head_dim, 2, dtype=torch.float32) / head_dim))
def mrope_cos_sin(position_ids3: torch.Tensor, head_dim: int, theta: float,
mrope_section: tuple[int, int, int]) -> tuple[torch.Tensor, torch.Tensor]:
"""position_ids3: [3, L] or [3, 1, L] (t, h, w) integer positions. Returns cos, sin [L, head_dim] float32."""
pos = position_ids3.reshape(3, -1).to(torch.float32) # [3, L]
inv = inv_freq(head_dim, theta) # [head_dim/2]
freqs = pos[:, :, None] * inv[None, None, :] # [3, L, head_dim/2]
out = freqs[0].clone()
for axis, offset in ((1, 1), (2, 2)):
length = mrope_section[axis] * 3
idx = slice(offset, length, 3)
out[..., idx] = freqs[axis, ..., idx]
emb = torch.cat([out, out], dim=-1) # [L, head_dim]
return emb.cos(), emb.sin()
def text_cos_sin(positions: torch.Tensor, head_dim: int, theta: float) -> tuple[torch.Tensor, torch.Tensor]:
"""1-D rope (t == h == w) for text tokens after the prompt: positions [L] (already offset by rope_delta)."""
pos3 = positions.reshape(1, -1).expand(3, -1)
return mrope_cos_sin(pos3, head_dim, theta, (head_dim // 2, 0, 0))
def rope_delta(position_ids3: torch.Tensor, seq_len: int) -> int:
"""HF rope_deltas: max multimodal position + 1 - number of tokens (so decode position = index + delta)."""
return int(position_ids3.max().item()) + 1 - seq_len