ARotting's picture
Publish Interactive periodic versus conventional recurrence comparison
5111d13 verified
Raw
History Blame Contribute Delete
2.3 kB
from __future__ import annotations
import torch
from torch import nn
class ClockworkRNN(nn.Module):
def __init__(self, hidden_dimensions: int = 32) -> None:
super().__init__()
self.hidden_dimensions = hidden_dimensions
self.periods = [1, 2, 4, 8]
self.block_size = hidden_dimensions // len(self.periods)
self.input_projection = nn.Linear(1, hidden_dimensions)
self.recurrent_projection = nn.Linear(hidden_dimensions, hidden_dimensions)
self.output = nn.Linear(hidden_dimensions, 1)
def forward(self, sequence: torch.Tensor) -> torch.Tensor:
batch, length, _ = sequence.shape
hidden = torch.zeros(batch, self.hidden_dimensions, device=sequence.device)
outputs = []
for step in range(length):
candidate = torch.tanh(
self.input_projection(sequence[:, step])
+ self.recurrent_projection(hidden)
)
blocks = []
for index, period in enumerate(self.periods):
start = index * self.block_size
end = start + self.block_size
blocks.append(
candidate[:, start:end] if step % period == 0 else hidden[:, start:end]
)
hidden = torch.cat(blocks, dim=1)
outputs.append(self.output(hidden))
return torch.stack(outputs, dim=1)
class PlainRNN(nn.Module):
def __init__(self, hidden_dimensions: int = 32) -> None:
super().__init__()
self.recurrent = nn.RNN(1, hidden_dimensions, batch_first=True)
self.output = nn.Linear(hidden_dimensions, 1)
def forward(self, sequence: torch.Tensor) -> torch.Tensor:
hidden, _ = self.recurrent(sequence)
return self.output(hidden)
class MatchedGRU(nn.Module):
def __init__(self, hidden_dimensions: int = 18) -> None:
super().__init__()
self.recurrent = nn.GRU(1, hidden_dimensions, batch_first=True)
self.output = nn.Linear(hidden_dimensions, 1)
def forward(self, sequence: torch.Tensor) -> torch.Tensor:
hidden, _ = self.recurrent(sequence)
return self.output(hidden)
def parameter_count(model: nn.Module) -> int:
return sum(parameter.numel() for parameter in model.parameters())