"""MLX counterpart of the FP32 GDN2 trajectory recurrence.""" from __future__ import annotations import mlx.core as mx from mlx import nn def _l2(value: mx.array) -> mx.array: return value * mx.rsqrt(mx.maximum(mx.sum(mx.square(value), axis=-1, keepdims=True), 1e-12)) class _GateProjection(nn.Module): 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 __call__(self, source: mx.array) -> mx.array: 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 = mx.zeros((heads,), mx.float32) self.dt_bias = mx.full((heads, key_dim), -6.906255, mx.float32) def empty(self, *leading: int) -> mx.array: return mx.zeros((*leading, self.heads, self.key_dim, self.value_dim), mx.float32) def _projections(self, source: mx.array): normalized = source.astype(mx.float32) source = (normalized * mx.rsqrt(mx.mean(mx.square(normalized), axis=-1, keepdims=True) + 1e-6)).astype(source.dtype) shape = source.shape[:-1] key_shape = (*shape, self.heads, self.key_dim) value_shape = (*shape, self.heads, self.value_dim) k = _l2(nn.silu(self.k_proj(source).astype(mx.float32)).reshape(key_shape)) v = nn.silu(self.v_proj(source).astype(mx.float32)).reshape(value_shape) head_rate = mx.exp(self.a_log.astype(mx.float32)).reshape( *((1,) * (source.ndim - 1)), self.heads, 1) decay = mx.exp(-head_rate * nn.softplus( self.f_proj(source).astype(mx.float32).reshape(key_shape) + self.dt_bias.astype(mx.float32))) erase = mx.sigmoid(self.b_proj(source).astype(mx.float32).reshape(key_shape)) write = mx.sigmoid(self.w_proj(source).astype(mx.float32).reshape(value_shape)) return k, v, decay, erase, write def _query(self, source: mx.array) -> mx.array: return _l2(nn.silu(self.q_proj(source).astype(mx.float32)).reshape( *source.shape[:-1], self.heads, self.key_dim)) def _output(self, value: mx.array, source: mx.array) -> mx.array: gate = nn.silu(self.g_proj(source).astype(mx.float32).reshape(value.shape)) value = value * mx.rsqrt(mx.mean(mx.square(value), axis=-1, keepdims=True) + 1e-6) * gate return self.o_proj(value.reshape(*source.shape[:-1], -1).astype(source.dtype)) def read(self, state: mx.array, source: mx.array) -> mx.array: if state.shape != (*source.shape[:-1], self.heads, self.key_dim, self.value_dim): raise ValueError("GDN2 state and query leading dimensions differ.") value = mx.sum(self._query(source)[..., None] * state.astype(mx.float32), axis=-2) return self._output(value, source) def read_shared(self, state: mx.array, source: mx.array) -> mx.array: 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(0, 2, 1, 3) value = (query @ state.astype(mx.float32)).transpose(0, 2, 1, 3) return self._output(value, source) @staticmethod def _transition(state, k, v, decay, erase, write, valid): decayed = state.astype(mx.float32) * decay[..., None] old = mx.sum((erase * k)[..., None] * decayed, axis=-2) candidate = decayed + k[..., None] * (write * v - old)[..., None, :] return candidate if valid is None else mx.where( valid[..., None, None, None], candidate, state.astype(mx.float32)) def transition(self, state: mx.array, source: mx.array, valid: mx.array | None = None) -> mx.array: # Do not override nn.Module.update: it installs parameters for optimizers, # dtype conversion, and autodiff. State evolution is a separate operation. 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: mx.array, source: mx.array, valid: mx.array) -> mx.array: 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 __all__ = ["GDN2Memory"]