lite_boost.ops.chunk_gated_delta_rule

查看源文件
lite_boost.ops.chunk_gated_delta_rule(query, key, value, beta, initial_state, actual_seq_lengths, g=None, scale_value=1.0)[源代码]

分块(prefill)Gated Delta Rule算子(A2)。

包装 torch.ops.lite_boost.chunk_gated_delta_rule,对应A2上CANN的AscendC算子 aclnnChunkGatedDeltaRule。将用户友好的BNSD布局转换为CANN算子要求的TND布局,query/key/value/beta/initial_state 转换为低精度dtype(bf16或fp16,跟随输入dtype;fp32/其他默认bf16),可选门 g 保持float32,actual_seq_lengths 直接透传(T = sum(actual_seq_lengths))。

在每个时间步 \(t\),Gated Delta Rule按如下公式计算新的递推状态和注意力输出:

\[S_t = \alpha_t S_{t-1} + \beta_t (v_t - \alpha_t S_{t-1} k_t) k_t^{\top}\]
\[o_t = S_t q_t \cdot scale\]

其中 \(\alpha_t = \exp(g_t)\) 为衰减因子(省略 g 时禁用衰减,即 \(\alpha_t = 1\)),\(\beta_t\) 为Delta更新步长。本算子是上述递推的分块(按块并行)实现,在长序列场景比逐token形式计算效率更高,适用于prefill阶段;输出每一步的结果以及最终状态。

本接口仅支持A2,不支持300I Duo。

参数:
  • query (Tensor) - 查询张量,shape \((B, N_k, S, D_k)\) 。计算前转换为低精度dtype。

  • key (Tensor) - 键张量,shape \((B, N_k, S, D_k)\) 。计算前转换为低精度dtype。

  • value (Tensor) - 值张量,shape \((B, N_v, S, D_v)\)N_v 须为 N_k 的整数倍(GQA模式下每个key头对应 N_v / N_k 个value头)。计算前转换为低精度dtype。

  • beta (Tensor) - Delta更新步长,shape \((B, N_v, S)\) ,取值范围[0, 1]。计算前转换为低精度dtype。

  • initial_state (Tensor) - 输入递推状态,shape \((B, N_v, D_k, D_v)\) ,内部转换为算子的value在前布局 [B, N_v, D_v, D_k]

  • actual_seq_lengths (Tensor) - 每batch的有效token数,shape \((B)\) ,dtype=int32。每个元素须在 [0, S] 范围内,BNSD布局的 S 维为每个batch padding后的长度,仅前 actual_seq_lengths[b] 个token有效。支持非等长,有效token会打包成 T = sum(actual_seq_lengths) 的TND张量,输出中超出每个batch长度的位置补零。

  • g (Tensor, 可选) - 全局衰减门,shape \((B, N_v, S)\) ,dtype=float32,必须为负值。该取值范围不做校验,因为校验需引入额外的归约算子(对张量求min/max),在推理路径上带来明显开销。超出范围的输入不会报错,但结果无意义。 None 表示禁用衰减门(hasGamma=0路径)。默认值: None

  • scale_value (float, 可选) - 施加在 query 上的注意力缩放因子。默认值: 1.0

返回:

tuple[Tensor, Tensor]

  • out (Tensor) - 注意力输出,shape \((B, N_v, S, D_v)\) ,dtype与输入低精度转换结果一致(默认bfloat16)。每个batch仅前 actual_seq_lengths[b] 个位置有效,padding位置补零。

  • final_state (Tensor) - 更新后的递推状态,shape \((B, N_v, D_k, D_v)\) ,dtype与 out 一致。

异常:
  • RuntimeError - 输入张量形状、dtype或设备不符合要求,或CANN算子执行失败时抛出。

说明

  • 本接口仅支持A2。300I Duo上注册的是另一签名的 ascend_300iduo 算子,本绑定不兼容,请勿在300I Duo上使用。

  • 所有输入张量必须在同一NPU设备上。

  • N_v``(值头数)须为 ``N_k``(key/query头数)的整数倍,每个key头对应 ``N_v / N_k 个value头。

  • BNSD布局的 S 维为每个batch padding后的序列长度,仅前 actual_seq_lengths[b] 个token有效;非等长 actual_seq_lengths 会自动打包/解包,输出中padding位置补零。

  • CANN算子通过DataTypeList同时接受bf16和fp16(q/k/v/beta/state),跟随输入dtype,fp32/其他输入默认bf16;可选门 g 始终为float32。

支持平台:

Ascend

样例:

>>> import torch
>>> import lite_boost.ops as lite_ops
>>> device = torch.device("npu:0")
>>> B, N, S, Dk, Dv = 1, 8, 16, 32, 64
>>> query = torch.randn(B, N, S, Dk, device=device, dtype=torch.bfloat16)
>>> key = torch.randn(B, N, S, Dk, device=device, dtype=torch.bfloat16)
>>> value = torch.randn(B, N, S, Dv, device=device, dtype=torch.bfloat16)
>>> beta = torch.rand(B, N, S, device=device, dtype=torch.bfloat16) * 0.9 + 0.05
>>> initial_state = torch.zeros(B, N, Dk, Dv, device=device, dtype=torch.bfloat16)
>>> actual_seq_lengths = torch.tensor([S], dtype=torch.int32, device=device)
>>> out, final_state = lite_ops.chunk_gated_delta_rule(
...     query, key, value, beta, initial_state, actual_seq_lengths)
>>> print(out.shape)
torch.Size([1, 8, 16, 64])
>>> print(final_state.shape)
torch.Size([1, 8, 32, 64])
>>> print(out.dtype)
torch.bfloat16