flashinfer.rope.apply_llama31_rope_inplace

flashinfer.rope.apply_llama31_rope_inplace(q: Tensor, k: Tensor, indptr: Tensor, offsets: 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),其中 nnzindptr 的最后一个元素。

  • k (torch.Tensor) – 键 Ragged Tensor,形状:(nnz, num_k_heads, head_dim),其中 nnzindptr 的最后一个元素。

  • indptr (torch.Tensor) – Indptr Tensor,形状:(batch_size + 1)

  • offsets (torch.Tensor) – 批处理中每个查询的相对位置偏移量,形状:(batch_size)

  • 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

示例

>>> import torch
>>> import flashinfer
>>> batch_size = 128
>>> qkv_len = 1024
>>> num_qo_heads = 32
>>> num_kv_heads = 32
>>> head_dim = 128
>>> nnz = batch_size * qkv_len
>>> qkv_packed = torch.randn(
>>>    nnz,
>>>    (num_qo_heads + 2 * num_kv_heads) * head_dim,
>>>    dtype=torch.float16,
>>>    device="cuda:0",
>>> )
>>> q = qkv_packed[:, : num_qo_heads * head_dim].reshape(nnz, num_qo_heads, head_dim)
>>> k = qkv_packed[
...    :, num_qo_heads * head_dim : (num_qo_heads + num_kv_heads) * head_dim
... ].reshape(nnz, num_kv_heads, head_dim)
>>> indptr = torch.tensor(
...    [i * qkv_len for i in range(batch_size + 1)], dtype=torch.int32, device="cuda:0"
>>> )
>>> offsets = torch.full((batch_size,), 10, dtype=torch.int32, device="cuda:0")
>>> flashinfer.apply_llama31_rope_inplace(q, k, indptr, offsets)