"""C/R-conditioned KDA recurrent linear attention expert. Keeps exact Q/K/V/C/R attention as the global path; this expert carries long-running sequence state. Forget/write gates are conditioned on intent context ``C`` and relation ``R`` glyphs. Output blend is zero-init so existing checkpoints remain identity-compatible. Uses ``fla.ops.kda.chunk_kda`` for CUDA execution and a tensor-native reference recurrence for CPU execution. A CUDA kernel/import failure is surfaced instead of silently changing the production algorithm. """ from __future__ import annotations from typing import cast import torch import torch.nn as nn import torch.nn.functional as F from resynthesis.config import GLYPH_DIM RESYNTHESIS_KDA_LOG_DECAY_FLOOR = -5.0 class CRConditionedKDAExpert(nn.Module): """Channel-decayed recurrent attention conditioned on C/R planes.""" def __init__( self, hidden_size: int, num_heads: int, *, glyph_dim: int = GLYPH_DIM, head_dim: int | None = None, ) -> None: super().__init__() heads = max(1, int(num_heads)) width = int(hidden_size) if width % heads != 0 and head_dim is None: while heads > 1 and width % heads != 0: heads -= 1 self.hidden_size = width self.num_heads = heads self.head_dim = int(head_dim) if head_dim is not None else width // heads self.glyph_dim = int(glyph_dim) inner = self.num_heads * self.head_dim # Head-grouped projections retain the full recurrent state width without # duplicating four dense hidden-by-hidden matrices per science layer. self.q_proj = nn.Conv1d( width, inner, kernel_size=1, groups=self.num_heads, bias=False, ) self.k_proj = nn.Conv1d( width, inner, kernel_size=1, groups=self.num_heads, bias=False, ) self.v_proj = nn.Conv1d( width, inner, kernel_size=1, groups=self.num_heads, bias=False, ) self.out_proj = nn.Conv1d( inner, width, kernel_size=1, groups=self.num_heads, bias=False, ) # Keep the recurrent output head-normalized and apply an # input-dependent full-rank gate before the output projection. This is # intentionally a dense projection: grouping it by recurrent head would # prevent the gate from coordinating channels learned by different # heads. self.output_gate_proj = nn.Linear(width, inner, bias=False) self.output_norm = nn.RMSNorm(self.head_dim) # Channel-wise forget (log-space raw) + scalar write, C/R conditioned. self.forget_hidden_proj = nn.Conv1d( width, inner, kernel_size=1, groups=self.num_heads, bias=True, ) self.forget_intent_proj = nn.Linear( self.glyph_dim, inner, bias=False, ) self.forget_relation_proj = nn.Linear( self.glyph_dim, inner, bias=False, ) self.write_hidden_proj = nn.Conv1d( width, self.num_heads, kernel_size=1, groups=self.num_heads, bias=True, ) self.write_intent_proj = nn.Linear( self.glyph_dim, self.num_heads, bias=False, ) self.write_relation_proj = nn.Linear( self.glyph_dim, self.num_heads, bias=False, ) self.short_conv = nn.Conv1d( width, width, kernel_size=3, padding=2, groups=width, bias=False, ) # Parameterize the finite log decay as # ``g_min * sigmoid(exp(A_h) * z_h)``. ``A_h`` is one learned scalar per # recurrent head, matching FLA's fused safe-gate contract. Its # one-dimensional zero initialization is also important for checkpoint # adoption: native graph migration preserves constructor values for # vector parameters, so an older checkpoint resumes with exp(A_h)=1. self.decay_log_scale = nn.Parameter(torch.zeros(self.num_heads)) self.blend_scale = nn.Parameter(torch.zeros(())) self._reset() def _reset(self) -> None: for module in ( self.q_proj, self.k_proj, self.v_proj, self.out_proj, ): nn.init.xavier_uniform_(module.weight) nn.init.xavier_uniform_(self.output_gate_proj.weight) nn.init.ones_(self.output_norm.weight) nn.init.zeros_(self.decay_log_scale) nn.init.xavier_uniform_(self.forget_hidden_proj.weight) forget_bias_t = self.forget_hidden_proj.bias if forget_bias_t is None: raise RuntimeError("KDA forget projection has no trained bias") nn.init.zeros_(forget_bias_t) nn.init.xavier_uniform_(self.forget_intent_proj.weight) nn.init.xavier_uniform_(self.forget_relation_proj.weight) nn.init.xavier_uniform_(self.write_hidden_proj.weight) write_bias_t = self.write_hidden_proj.bias if write_bias_t is None: raise RuntimeError("KDA write projection has no trained bias") nn.init.zeros_(write_bias_t) nn.init.xavier_uniform_(self.write_intent_proj.weight) nn.init.xavier_uniform_(self.write_relation_proj.weight) nn.init.dirac_(self.short_conv.weight) nn.init.zeros_(self.blend_scale) @staticmethod def _project_sequence( projection: nn.Conv1d, tensor: torch.Tensor, ) -> torch.Tensor: projected_t = cast(torch.Tensor, projection(tensor.transpose(1, 2))) return projected_t.transpose(1, 2) def _reshape_heads(self, tensor: torch.Tensor) -> torch.Tensor: batch, seq, _ = tensor.shape return tensor.view(batch, seq, self.num_heads, self.head_dim) def _bounded_log_decay( self, forget_raw: torch.Tensor, ) -> torch.Tensor: """Map learned decay logits into the finite recurrent log range.""" if ( forget_raw.ndim != 4 or forget_raw.shape[-2:] != ( self.num_heads, self.head_dim, ) ): raise ValueError("KDA forget-logit geometry differs") decay_scale_t = self.decay_log_scale.exp().view( 1, 1, self.num_heads, 1, ) return torch.sigmoid(decay_scale_t * forget_raw).mul( RESYNTHESIS_KDA_LOG_DECAY_FLOOR ) def forward( self, hidden: torch.Tensor, *, intent_glyph_context: torch.Tensor, relation_glyph_context: torch.Tensor, ) -> torch.Tensor: if hidden.ndim != 3 or hidden.shape[-1] != self.hidden_size: raise ValueError("KDA expert hidden geometry differs") if intent_glyph_context.shape[:2] != hidden.shape[:2]: raise ValueError("KDA intent context geometry differs") if relation_glyph_context.shape[:2] != hidden.shape[:2]: raise ValueError("KDA relation context geometry differs") # Depthwise short conv on sequence (B,H,T) — causal via left pad trim. conv_in = hidden.transpose(1, 2) conv_out = self.short_conv(conv_in)[..., : hidden.shape[1]].transpose(1, 2) x = F.silu(conv_out) intent_t = intent_glyph_context.to(dtype=x.dtype) relation_t = relation_glyph_context.to(dtype=x.dtype) q = self._reshape_heads(self._project_sequence(self.q_proj, x)) k = self._reshape_heads(self._project_sequence(self.k_proj, x)) v = self._reshape_heads(self._project_sequence(self.v_proj, x)) q = F.normalize(q, dim=-1) k = F.normalize(k, dim=-1) # fla chunk_kda: g is forget in log space [B,T,H,K]; beta write [B,T,H] forget_raw = ( self._project_sequence(self.forget_hidden_proj, x) + self.forget_intent_proj(intent_t) + self.forget_relation_proj(relation_t) ).view( hidden.shape[0], hidden.shape[1], self.num_heads, self.head_dim, ) write_logits_t = ( self._project_sequence(self.write_hidden_proj, x) + self.write_intent_proj(intent_t) + self.write_relation_proj(relation_t) ) out = self._run_kda_from_logits( q, k, v, forget_raw, write_logits_t, ) normalized_out = self.output_norm(out) flat = normalized_out.reshape(hidden.shape[0], hidden.shape[1], -1) output_gate_t = torch.sigmoid(self.output_gate_proj(x)) gated_flat = output_gate_t * flat projected = self._project_sequence( self.out_proj, gated_flat.to(dtype=hidden.dtype), ) return torch.tanh(self.blend_scale) * projected def _run_kda_from_logits( self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, forget_raw: torch.Tensor, write_logits: torch.Tensor, ) -> torch.Tensor: """Run KDA while retaining the native fused safe-gate CUDA path. FLA's KDA backend owns the same lower-bounded recurrent activation. Supplying the logits and per-head ``A_h`` directly lets it keep gate activation, 16-token safe rescaling, and write sigmoid inside the kernel. CPU and the one-token zero-state shortcut materialize the identical equations explicitly. """ if q.is_cuda and q.shape[1] > 1: from fla.ops.kda import chunk_kda # type: ignore[import-untyped] out, _state = chunk_kda( q.contiguous(), k.contiguous(), v.contiguous(), forget_raw.contiguous(), write_logits.contiguous(), use_gate_in_kernel=True, use_beta_sigmoid_in_kernel=True, safe_gate=True, lower_bound=RESYNTHESIS_KDA_LOG_DECAY_FLOOR, A_log=self.decay_log_scale.contiguous(), ) if not isinstance(out, torch.Tensor): raise RuntimeError("FLA KDA returned a non-tensor output") return out return self._run_kda( q, k, v, self._bounded_log_decay(forget_raw), torch.sigmoid(write_logits), ) def _run_kda( self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, g: torch.Tensor, beta: torch.Tensor, ) -> torch.Tensor: if q.shape[1] == 1: # CAS bulk science attempts are isolated single-token recurrent # sessions. With the exact zero initial state, decay and delta # correction multiply zero, so KDA reduces to beta*v*(k·q). # Keep a zero dependency on g because the reference recurrence # produces a zero (not absent) forget-gate gradient. q_t = q[:, 0] k_t = k[:, 0] v_t = v[:, 0] beta_t = beta[:, 0].unsqueeze(-1).float() scale = q_t.shape[-1] ** -0.5 alignment_t = ( (k_t.float() * q_t.float()).sum(dim=-1, keepdim=True) * scale ) forget_zero_t = g[:, :1].sum(dim=-1, keepdim=True).mul(0) return ( beta_t.mul(v_t.float()) .mul(alignment_t) .to(dtype=q.dtype) .unsqueeze(1) + forget_zero_t.to(dtype=q.dtype) ) # Triton/FLA kernels require CUDA tensors; CPU has an explicit reference # implementation. The historical sequence-parallel canary split this # recurrence and carried state in a Python loop. That was sequential, # not context parallel. A full native launch is both faster and exact; # real KDA context-parallel owners use the associative tensor packet in # ``sequence_parallel.py`` rather than a host-flag simulation. if q.is_cuda: from fla.ops.kda import chunk_kda # type: ignore[import-untyped] out, _state = chunk_kda( q.contiguous(), k.contiguous(), v.contiguous(), g.contiguous(), beta.contiguous(), ) if not isinstance(out, torch.Tensor): raise RuntimeError("FLA KDA returned a non-tensor output") return out return self._reference_kda(q, k, v, g, beta) @staticmethod def _reference_kda( q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, g: torch.Tensor, beta: torch.Tensor, ) -> torch.Tensor: """O(T) reference KDA / gated delta-rule recurrence for CPU.""" batch, seq, heads, dim = q.shape value_dim = v.shape[-1] state = q.new_zeros(batch, heads, dim, value_dim) outputs = q.new_empty(batch, seq, heads, value_dim) scale = dim ** -0.5 for t in range(seq): alpha = torch.exp(g[:, t]).clamp(0.0, 1.0) # [B,H,K] bt = beta[:, t].unsqueeze(-1).unsqueeze(-1) # [B,H,1,1] kt = k[:, t].unsqueeze(-1) # [B,H,K,1] vt = v[:, t].unsqueeze(-2) # [B,H,1,V] state = state * alpha.unsqueeze(-1) # Delta-rule correction then write. read = torch.matmul(state.transpose(-1, -2), k[:, t].unsqueeze(-1)) state = state - bt * torch.matmul(kt, read.transpose(-1, -2)) state = state + bt * torch.matmul(kt, vt) ot = torch.matmul( state.transpose(-1, -2), q[:, t].mul(scale).unsqueeze(-1), ).squeeze(-1) outputs[:, t] = ot return outputs