"""Vendored FLA KDA gate fusion (standard + safe gate + chunk cumsum).""" from __future__ import annotations import torch from kda._fla.ops.kda.gate import ( kda_gate_bwd, kda_gate_chunk_cumsum, kda_gate_fwd as _kda_gate_fwd, ) DEFAULT_LOWER_BOUND = -5.0 def kda_gate_fwd( g: torch.Tensor, A_log: torch.Tensor, dt_bias: torch.Tensor | None = None, lower_bound: float | None = DEFAULT_LOWER_BOUND, ): return _kda_gate_fwd( g, A_log=A_log, dt_bias=dt_bias, lower_bound=lower_bound, output_dtype=g.dtype, ) __all__ = [ "DEFAULT_LOWER_BOUND", "kda_gate_bwd", "kda_gate_chunk_cumsum", "kda_gate_fwd", ]