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