"""PyTorch references for the two KDA gate activations.""" from __future__ import annotations import torch import torch.nn.functional as F def kda_gate_reference( g: torch.Tensor, A_log: torch.Tensor, dt_bias: torch.Tensor | None = None, *, safe_gate: bool = False, lower_bound: float | None = None, ) -> torch.Tensor: """Compute the official KDA gate semantics in PyTorch. ``A_log`` is head-wise with shape ``[HV]`` and ``dt_bias`` is per-dimension with shape ``[HV, K]`` (or flattened to ``[HV*K]``). """ HV, K = g.shape[-2:] gate_input = g if dt_bias is None else g + dt_bias.view(HV, K) rate = A_log.view(HV, 1).exp() if safe_gate: if lower_bound is None: raise ValueError("lower_bound is required when safe_gate=True") return lower_bound * torch.sigmoid(rate * gate_input) return -rate * F.softplus(gate_input) def kda_gate_naive( g: torch.Tensor, A_log: torch.Tensor, dt_bias: torch.Tensor | None = None, lower_bound: float | None = None, ) -> torch.Tensor: """Compatibility name matching FLA's reference gate convention.""" return kda_gate_reference( g, A_log, dt_bias, safe_gate=lower_bound is not None, lower_bound=lower_bound, ) __all__ = ["kda_gate_naive", "kda_gate_reference"]