import torch from torch import nn import torch.nn.functional as F class GeometricallyAttractingRandomRNN(nn.Module): """RNN with K candidate tanh maps and a time-dependent categorical gate. Forward uses either a sampled candidate (stochastic=True) or the exact probability-weighted output mixture. The latter is useful for low-variance training. Candidate spectral norms provide a cheap Jacobian upper bound. """ def __init__(self, input_size, hidden_size, candidates=2, target_rho=0.95, gate_input_size=None): super().__init__() self.input_size = input_size self.hidden_size = hidden_size self.candidates = candidates self.target_rho = target_rho self.W = nn.Parameter(torch.empty(candidates, hidden_size, hidden_size)) self.U = nn.Parameter(torch.empty(candidates, hidden_size, input_size)) self.bias = nn.Parameter(torch.zeros(candidates, hidden_size)) nn.init.orthogonal_(self.W[0]) for k in range(1, candidates): nn.init.orthogonal_(self.W[k]) nn.init.xavier_uniform_(self.U) gate_input_size = input_size if gate_input_size is None else gate_input_size self.gate = nn.Linear(gate_input_size, candidates) def gains(self): # Exact matrix-norm upper bound for tanh Jacobian (which is <= 1). return torch.linalg.matrix_norm(self.W, ord=2, dim=(-2, -1)) def contraction_penalty(self, probabilities, eps=1e-8): """Squared positive excess of log expected gain over log(target_rho).""" expected_gain = (probabilities * self.gains()).sum(dim=-1) excess = torch.log(expected_gain + eps) - torch.log( torch.as_tensor(self.target_rho, device=expected_gain.device)) return F.relu(excess).square().mean() def step(self, x, h, gate_features=None, stochastic=True): # x: [B,input], h: [B,hidden], gate_features: [B,gate_input]. z = x if gate_features is None else gate_features probabilities = torch.softmax(self.gate(z), dim=-1) candidates = torch.tanh( torch.einsum('kij,bj->bki', self.W, h) + torch.einsum('kij,bj->bki', self.U, x) + self.bias[None]) if stochastic: index = torch.multinomial(probabilities, 1).squeeze(-1) next_h = candidates[torch.arange(x.shape[0], device=x.device), index] else: next_h = (probabilities[..., None] * candidates).sum(dim=1) return next_h, probabilities def forward(self, x_sequence, h0=None, stochastic=True): # x_sequence: [T,B,input] T, B, _ = x_sequence.shape h = x_sequence.new_zeros(B, self.hidden_size) if h0 is None else h0 states, routes = [], [] for t in range(T): h, p = self.step(x_sequence[t], h, stochastic=stochastic) states.append(h) routes.append(p) return torch.stack(states), torch.stack(routes)