flashinfer.rope.apply_rope_with_cos_sin_cache

flashinfer.rope.apply_rope_with_cos_sin_cache(positions: Tensor, query: Tensor, key: Tensor, head_size: int, cos_sin_cache: Tensor, is_neox: bool = True) Tuple[Tensor, Tensor]

使用预计算的 cos/sin 值将旋转嵌入应用于 key 和 query。此设计旨在与 SGL/vLLM 实现兼容。

参数:
  • positions (torch.Tensor) – 位置索引,形状:(nnz)

  • query (torch.Tensor) – 查询张量,形状:(nnz, num_q_heads * head_size)

  • key (torch.Tensor) – 键张量,形状:(nnz, num_k_heads * head_size)

  • cos_sin_cache (torch.Tensor) – Cosine 和 Sine 缓存张量,形状:(max_seq_len, rotary_dim)。Cosine 是 rotary_dim 的前半部分,Sine 是后半部分。

  • is_neox (bool) –

    是否使用 Neox 风格 RoPE,默认值:True

    • 如果为 True,则 query/key 张量的最后一个维度不交错,即我们旋转前一半维度 ([..., :head_dim//2]) 和后一半维度 ([..., head_dim//2:])

    • 如果为 False,则 query/key 张量的最后一个维度是交错的,即我们旋转偶数维度 ([..., ::2]) 和奇数维度 ([..., 1::2])

返回值:

  • query_out (torch.Tensor) – 旋转后的查询张量,形状:(nnz, num_q_heads * head_size)

  • key_out (torch.Tensor) – 旋转后的键张量,形状:(nnz, num_k_heads * head_size)

注意

旋转维度由 cosine 缓存和 sine 缓存确定。