flashinfer.rope.apply_llama31_rope_pos_ids_inplace¶
- flashinfer.rope.apply_llama31_rope_pos_ids_inplace(q: Tensor, k: Tensor, pos_ids: Tensor, rotary_dim: int | None = None, interleave: bool = False, rope_scale: float = 8, rope_theta: float = 500000.0, low_freq_factor: float = 1, high_freq_factor: float = 4, old_context_len: int = 8192) None¶
将 Llama 3.1 风格的旋转嵌入应用于一批查询/键(存储为 RaggedTensor)就地。 cos/sin 值在内核内部实时计算。
我们使用
indptr表示批处理中每个片段的起始指针,第 i 个片段的查询是q[indptr[i]:indptr[i+1]],第 i 个片段的键是k[indptr[i]:indptr[i+1]],indptr的第一个元素始终为 0,indptr的最后一个元素是批处理中的查询/键总数。有关 Ragged Tensor 的更多详细信息,请参阅 Ragged Tensor 教程。- 参数:
q (torch.Tensor) – 查询稀疏张量,形状:
(nnz, num_q_heads, head_dim),其中nnz是indptr的最后一个元素。k (torch.Tensor) – 键 Ragged Tensor,形状:
(nnz, num_k_heads, head_dim),其中nnz是indptr的最后一个元素。pos_ids (torch.Tensor) – 位置索引,形状:
(nnz)。rotary_dim (Optional[int]) – 应用 RoPE 的维度,如果为
None,则将 RoPE 应用于整个 head 维度,否则将 RoPE 应用于前rotary_dim个维度,默认值:None。interleave (bool) –
是否在最后一个维度中使用交错布局,默认值:
False。如果为
True,则查询/键张量的最后一个维度是交错的,即我们旋转偶数维度([..., ::2])和奇数维度([..., 1::2])。如果为
False,则查询/键张量的最后一个维度不交错,即我们旋转前一半维度([..., :head_dim//2])和后一半维度([..., head_dim//2:])。
rope_scale (float) – 在 rope 嵌入中使用的缩放因子,默认值:
8。rope_theta (float) – 在 rope 嵌入中使用的 theta 值,默认值:
5e5。low_freq_factor (float) – 在 Llama 3.1 RoPE 中使用的低频因子,默认值:
1。high_freq_factor (float) – 在 Llama 3.1 RoPE 中使用的低频因子,默认值:
4。old_context_len (int) – 在 Llama 3.1 RoPE 中使用的旧上下文长度,默认值:
8192。