File size: 14,227 Bytes
919fd68
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
"""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