Modilify-Mk2-preview-mlx / modilify_mk2 /mlx_gdn2_memory.py
ydy9038074's picture
Publish Modilify Mk2 Preview MLX
e4f7326 verified
Raw History Blame Contribute Delete
6.35 kB
"""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"]