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))