# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang import torch from fla.ops.linear_attn.utils import normalize_output from fla.ops.simple_gla import fused_chunk_simple_gla @torch.compiler.disable def fused_chunk_linear_attn( q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, scale: float | None = None, initial_state: torch.Tensor | None = None, output_final_state: bool = False, normalize: bool = True, cu_seqlens: torch.LongTensor | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: r""" Args: q (torch.Tensor): queries of shape `[B, T, H, K]`. k (torch.Tensor): keys of shape `[B, T, H, K]`. v (torch.Tensor): values of shape `[B, T, H, V]`. scale (Optional[float]): Scale factor for linear attention scores. If not provided, it will default to `1 / sqrt(K)`. Default: `None`. initial_state (Optional[torch.Tensor]): Initial state of shape `[B, H, K, V]`. Default: `None`. output_final_state (Optional[bool]): Whether to output the final state of shape `[B, H, K, V]`. Default: `False`. normalize (bool): Whether to normalize the output. Default: `True`. cu_seqlens (torch.LongTensor): Cumulative sequence lengths of shape `[N+1]` used for variable-length training, consistent with the FlashAttention API. Returns: o (torch.Tensor): Outputs of shape `[B, T, H, V]`. final_state (torch.Tensor): Final state of shape `[B, H, K, V]` if `output_final_state=True` else `None` """ o, final_state = fused_chunk_simple_gla( q=q, k=k, v=v, scale=scale, initial_state=initial_state, output_final_state=output_final_state, cu_seqlens=cu_seqlens, ) if normalize: o = normalize_output(q * scale, k, o) return o, final_state