"""Vendored FLA chunk KDA forward, returning ``(o, ht)`` like the public op.""" from __future__ import annotations import torch from kda._fla.ops.kda.chunk import chunk_kda from kda._fla.ops.kda.chunk_fwd import chunk_kda_fwd as fla_chunk_kda_fwd __all__ = ["chunk_kda_fwd", "fla_chunk_kda_fwd"] def chunk_kda_fwd( q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, g: torch.Tensor, beta: torch.Tensor, scale: float | None = None, initial_state: torch.Tensor | None = None, output_final_state: bool = False, chunk_size: int = 64, **kwargs, ): """Chunked KDA forward with FLA kernels. Returns ``(o, ht)``.""" return chunk_kda( q, k, v, g, beta, scale=scale, initial_state=initial_state, output_final_state=output_final_state, chunk_size=chunk_size, **kwargs, )