ttnn.transformer.gated_delta_attn_seq
- ttnn.transformer.gated_delta_attn_seq(L_unit: ttnn.Tensor, v_beta_sc: ttnn.Tensor, k_bd_sc: ttnn.Tensor, intra_attn: ttnn.Tensor, q_decay: ttnn.Tensor, k_decay_t: ttnn.Tensor, dl_exp: ttnn.Tensor, L_inv: ttnn.Tensor, *, initial_state: ttnn.Tensor | None, memory_config: ttnn.MemoryConfig | None) tuple[ttnn.Tensor, ttnn.Tensor]
-
Gated DeltaNet attention — sequential inter-chunk scan (Path A).
All inputs must be float32. Python pre-normalises L_unit to unit-diagonal form and precomputes L_inv (diagonal block inverses). The C++ kernel performs blocked forward substitution and the sequential inter-chunk state update.
- Parameters:
-
L_unit (ttnn.Tensor) – [BH, NC, C, C] unit-diagonal lower-tri
v_beta_sc (ttnn.Tensor) – [BH, NC, C, Dv] D^{-1} @ v_beta
k_bd_sc (ttnn.Tensor) – [BH, NC, C, Dk] D^{-1} @ k_beta_decay
intra_attn (ttnn.Tensor) – [BH, NC, C, C] intra-chunk attention
q_decay (ttnn.Tensor) – [BH, NC, C, Dk] queries with decay
k_decay_t (ttnn.Tensor) – [BH, NC, Dk, C] transposed keys with decay
dl_exp (ttnn.Tensor) – [BH, NC, 1, 1] state decay scalar (fp32)
L_inv (ttnn.Tensor) – [BH, NC, C, 32] 4 diagonal block inverses per chunk
- Keyword Arguments:
-
initial_state (ttnn.Tensor, optional) – [BH, Dk, Dv] initial state (zeros if absent).
memory_config (ttnn.MemoryConfig, optional) – output memory config.
compute_kernel_config (ttnn.DeviceComputeKernelConfig, optional) –
- Returns:
-
tuple[ttnn.Tensor, ttnn.Tensor] – output [BH, NC, C, Dv], final_state [BH, Dk, Dv]