File size: 6,169 Bytes
b88f761
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Reference GDN2 matrix memory for trajectory-time recurrence.

The state is FP32.  Leading dimensions are independent rows or canvas cells;
only the last three dimensions, [heads, key, value], belong to the rule.
"""

from __future__ import annotations

import torch
from torch import nn
from torch.nn import functional as F


class _GateProjection(nn.Module):
    """GDN2 gates need channel control, not a second full-width content map."""

    def __init__(self, source: int, target: int, rank: int) -> None:
        super().__init__()
        self.down = nn.Linear(source, rank, bias=False)
        self.up = nn.Linear(rank, target, bias=False)

    def forward(self, source: torch.Tensor) -> torch.Tensor:
        return self.up(self.down(source))


class GDN2Memory(nn.Module):
    def __init__(self, input_dim: int, heads: int, key_dim: int, value_dim: int,
                 *, observation_dim: int | None = None) -> None:
        super().__init__()
        if min(input_dim, heads, key_dim, value_dim) <= 0:
            raise ValueError("GDN2 dimensions must be positive.")
        self.heads, self.key_dim, self.value_dim = heads, key_dim, value_dim
        observation_dim = input_dim if observation_dim is None else observation_dim
        if observation_dim <= 0:
            raise ValueError("GDN2 observation width must be positive.")
        self.q_proj = nn.Linear(input_dim, heads * key_dim, bias=False)
        self.k_proj = nn.Linear(observation_dim, heads * key_dim, bias=False)
        self.v_proj = nn.Linear(observation_dim, heads * value_dim, bias=False)
        self.f_proj = _GateProjection(observation_dim, heads * key_dim, min(observation_dim, key_dim))
        self.b_proj = nn.Linear(observation_dim, heads * key_dim, bias=False)
        self.w_proj = nn.Linear(observation_dim, heads * value_dim, bias=False)
        self.g_proj = _GateProjection(input_dim, heads * value_dim, min(input_dim, value_dim))
        self.o_proj = nn.Linear(heads * value_dim, input_dim, bias=False)
        self.a_log = nn.Parameter(torch.zeros(heads))
        self.dt_bias = nn.Parameter(torch.full((heads, key_dim), -6.906255))

    def _projections(self, source: torch.Tensor):
        source = F.rms_norm(source.float(), (source.shape[-1],), eps=1e-6).to(source.dtype)
        shape = source.shape[:-1]
        key_shape = (*shape, self.heads, self.key_dim)
        value_shape = (*shape, self.heads, self.value_dim)
        k = F.normalize(F.silu(self.k_proj(source).float()).reshape(key_shape), dim=-1)
        v = F.silu(self.v_proj(source).float()).reshape(value_shape)
        head_rate = self.a_log.float().exp().reshape(*((1,) * (source.ndim - 1)), self.heads, 1)
        decay = torch.exp(-head_rate * F.softplus(
            self.f_proj(source).float().reshape(key_shape) + self.dt_bias.float()))
        erase = torch.sigmoid(self.b_proj(source).float().reshape(key_shape))
        write = torch.sigmoid(self.w_proj(source).float().reshape(value_shape))
        return k, v, decay, erase, write

    def _query(self, source: torch.Tensor) -> torch.Tensor:
        return F.normalize(F.silu(self.q_proj(source).float()).reshape(
            *source.shape[:-1], self.heads, self.key_dim), dim=-1)

    def _output(self, value: torch.Tensor, source: torch.Tensor) -> torch.Tensor:
        gate = F.silu(self.g_proj(source).float().reshape(value.shape))
        value = F.rms_norm(value, (self.value_dim,), eps=1e-6) * gate
        return self.o_proj(value.flatten(-2).to(source.dtype))

    def read(self, state: torch.Tensor, source: torch.Tensor) -> torch.Tensor:
        if state.shape != (*source.shape[:-1], self.heads, self.key_dim, self.value_dim):
            raise ValueError("GDN2 state and query leading dimensions differ.")
        value = torch.einsum("...hk,...hkv->...hv", self._query(source), state.float())
        return self._output(value, source)

    def read_shared(self, state: torch.Tensor, source: torch.Tensor) -> torch.Tensor:
        """Read one row matrix for all queries without a canvas-sized matrix product."""
        if source.ndim != 3 or state.shape != (source.shape[0], self.heads, self.key_dim, self.value_dim):
            raise ValueError("Shared GDN2 read requires [batch, canvas, width] queries.")
        query = self._query(source).transpose(1, 2)
        value = (query @ state.float()).transpose(1, 2)
        return self._output(value, source)

    @staticmethod
    def _transition(state, k, v, decay, erase, write, valid):
        decayed = state.float() * decay.unsqueeze(-1)
        old = torch.einsum("...hk,...hkv->...hv", erase * k, decayed)
        candidate = decayed + k.unsqueeze(-1) * (write * v - old).unsqueeze(-2)
        return candidate if valid is None else torch.where(
            valid[..., None, None, None], candidate, state.float())

    def transition(self, state: torch.Tensor, source: torch.Tensor,
                   valid: torch.Tensor | None = None) -> torch.Tensor:
        if state.shape != (*source.shape[:-1], self.heads, self.key_dim, self.value_dim):
            raise ValueError("GDN2 state and observation leading dimensions differ.")
        k, v, decay, erase, write = self._projections(source)
        if valid is not None:
            if valid.shape != source.shape[:-1]:
                raise ValueError("GDN2 valid mask must match observation rows.")
        return self._transition(state, k, v, decay, erase, write, valid)

    def write_sequence(self, state: torch.Tensor, source: torch.Tensor,
                       valid: torch.Tensor) -> torch.Tensor:
        """Project a commit packet once, then preserve ordered GDN2 recurrence."""
        if source.ndim != 3 or valid.shape != source.shape[:2]:
            raise ValueError("GDN2 sequence and mask must share [batch, length].")
        if state.shape != (source.shape[0], self.heads, self.key_dim, self.value_dim):
            raise ValueError("GDN2 sequence state shape differs.")
        projected = self._projections(source)
        for index in range(source.shape[1]):
            state = self._transition(state, *(part[:, index] for part in projected), valid[:, index])
        return state