class SelectiveSSM(nnx.Module):
"""A diagonal SSM whose step size, input matrix, and read-out are
functions of the input (Gu & Dao, 2023)."""
def __init__(self, num_hiddens, num_states=4, dt_min=0.001, dt_max=0.1,
rngs=None):
rngs = nnx.Rngs(0) if rngs is None else rngs
H, N, R = num_hiddens, num_states, max(2, num_hiddens // 16)
self.log_a = nnx.Param(jnp.tile(jnp.log(jnp.arange(1., N + 1)),
(H, 1)))
self.W_dt = nnx.Sequential(
nnx.Linear(H, R, rngs=rngs),
nnx.Linear(R, H, use_bias=False, rngs=rngs))
dt = jnp.exp(rngs.params.uniform((H,)) * math.log(dt_max / dt_min)
+ math.log(dt_min))
self.b_dt = nnx.Param(dt + jnp.log(-jnp.expm1(-dt)))
self.W_B = nnx.Linear(H, N, use_bias=False, rngs=rngs)
self.W_C = nnx.Linear(H, N, use_bias=False, rngs=rngs)
self.D = nnx.Param(jnp.ones(H))
def __call__(self, u): # (num_steps, batch, num_hiddens)
a = -jnp.exp(self.log_a[...]) # (H, N), Re(a) < 0
dt = jax.nn.softplus(self.W_dt(u) + self.b_dt) # (T, batch, H)
B, C = self.W_B(u), self.W_C(u) # (T, batch, N)
a_bar = jnp.exp(dt[..., None] * a) # (T, batch, H, N)
b_bar = (dt * u)[..., None] * B[..., None, :]
x = associative_scan(a_bar, b_bar)
return (x * C[..., None, :]).sum(-1) + self.D * u