lite_boost.ops.recurrent_gated_delta_rule
- lite_boost.ops.recurrent_gated_delta_rule(query, key, value, beta, state, actual_seq_lengths, ssm_state_indices, g, gk, num_accepted_tokens, scale_value=1.0)[源代码]
基于CANN aclnn后端的递推式线性注意力decode算子。
实现Gated Delta Rule的逐token递推前向计算,更新递推状态矩阵并输出注意力结果。主要用于混合线性注意力模型(如Qwen3.5)的decode阶段推理加速。
算法流程(对每个batch中的每个token依次执行)。状态衰减为
\[S = S * \exp(g) * \exp(gk)\]记忆检索为
\[kv\_mem = S^{\top} k\]Delta更新为
\[S = S + k^{\top} ((v - kv\_mem) * \beta)\]输出计算为
\[o = S^{\top} q\]其中 \(S\) 是shape为 \((N_v, D_k, D_v)\) 的递推状态矩阵,存储了线性注意力的key-value关联信息。
- 参数:
query (Tensor) - 查询张量,shape \((B, N_k, S, D_k)\) ,dtype=bfloat16。必须L2归一化(每个head向量的L2范数为1,值域[0, 1])。其中B=batch_size,N_k=查询头数,S=序列长度,D_k=key维度。
key (Tensor) - 键张量,shape \((B, N_k, S, D_k)\) ,dtype=bfloat16。必须L2归一化(同query)。
value (Tensor) - 值张量,shape \((B, N_v, S, D_v)\) ,dtype=bfloat16。N_v=值头数(须为N_k的整数倍),D_v=value维度。
beta (Tensor) - Delta更新步长,shape \((B, N_v, S)\) ,dtype=bfloat16。取值范围[0, 1]。控制每次delta更新的幅度,beta越大,新信息覆盖旧记忆的程度越强;beta越小,倾向于保留已有记忆。
state (Tensor) - 递推状态池,shape \((state\_slots, N_v, D_k, D_v)\) ,dtype=bfloat16。state_slots为池中状态槽的个数,每个槽独立存储一个序列累积的key-value关联,各token经 ssm_state_indices 映射到对应槽位(普通推理时每个batch占用一个槽,state_slots通常等于B)。D_k为key维度(行),D_v为value维度(列)。首次调用时可初始化为零张量。
actual_seq_lengths (Tensor) - 实际序列长度,shape \((B)\) ,dtype=int32。用于变长序列推理。每个元素表示对应batch中的有效token数。例如
[4, 3, 5]表示3个batch的序列长度分别为4、3、5。ssm_state_indices (Tensor) - 状态槽索引,shape \((T)\) ,dtype=int32,其中
T = B * S(每个展平后的token一个索引),每个token据此在全局状态池( state 的第0维,共state_slots个槽)中选择一个状态槽。g (Tensor) - 全局衰减门,shape \((B, N_v, S)\) ,dtype=float32。必须为负值。该取值范围不做校验,因为校验需引入额外的归约算子(对张量求min/max),在推理路径上带来明显开销。超出范围的输入不会报错,但结果无意义。
exp(g)作为状态衰减因子,值域(0, 1)。g越负,历史信息遗忘越快。例如g=-1时,每步保留约37%的历史状态。gk (Tensor) - key维度门控,shape \((B, N_v, S, D_k)\) ,dtype=float32。必须为负值(不做校验,理由同 g )。
exp(gk)对每个key维度独立施加衰减,实现更细粒度的记忆控制。与全局门g的区别在于gk在D_k维度上逐元素操作。num_accepted_tokens (Tensor) - 已接受token数,shape \((B)\) ,dtype=int32。在speculative decoding等场景中用于标记实际接受(非拒绝)的token数量。普通推理时与
actual_seq_lengths相同。scale_value (float, 可选) - 注意力缩放因子。通常设为
1.0 / sqrt(D_k),与标准注意力缩放一致。query在计算前会乘以此缩放因子。默认值:1.0。
- 返回:
tuple[Tensor, Tensor]
out (Tensor) - 注意力输出,shape \((B, N_v, S, D_v)\) ,dtype=bfloat16。每个token位置的线性注意力计算结果。
state_out (Tensor) - 更新后的递推状态池,shape与 state 一致,dtype=bfloat16。需在下一步递推时作为
state输入传入,形成状态传递链。
- 异常:
RuntimeError - 输入张量dtype或设备不符合要求,或CANN算子执行失败时抛出。
ValueError - 输入张量shape不符合要求时抛出,包括
N_v不是N_k的整数倍,或 ssm_state_indices 的元素个数不是B * S(每个展平后的token一个索引)。
说明
本算子仅支持 decode阶段 (逐token推理),序列长度S不应超过8。Prefill阶段的并行计算请使用chunk-level算子。
支持分组递推头(grouped recurrent heads),即N_v为N_k的整数倍。
所有输入张量必须在同一NPU设备上。
CANN算子内部状态存储为 \((state\_slots, N_v, D_v, D_k)\) 布局(value维度在前),本函数会自动进行布局转换。
- 支持平台:
Ascend
样例:
>>> import torch >>> import lite_boost.ops as lite_ops >>> device = torch.device("npu:0") >>> B, N, S, Dk, Dv = 1, 64, 1, 64, 512 >>> 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 >>> state = torch.zeros(B, N, Dk, Dv, device=device, dtype=torch.bfloat16) >>> g = -(torch.rand(B, N, S, device=device) + 0.01) >>> gk = -(torch.rand(B, N, S, Dk, device=device) + 0.01) >>> actual_seq_lengths = torch.tensor([S], dtype=torch.int32, device=device) >>> ssm_state_indices = torch.tensor([0], dtype=torch.int32, device=device) >>> num_accepted_tokens = torch.tensor([S], dtype=torch.int32, device=device) >>> output, state_out = lite_ops.recurrent_gated_delta_rule( ... query, key, value, beta, state, ... actual_seq_lengths, ssm_state_indices, ... g, gk, num_accepted_tokens, ... scale_value=1.0 / (Dk ** 0.5))