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