"""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