spike-pocket-lab / model.py
ARotting's picture
Publish Interactive Poisson spike raster and LIF classifier
b63a6a2 verified
Raw
History Blame Contribute Delete
2.61 kB
from __future__ import annotations
import torch
from torch import nn
from torch.nn import functional as F
def surrogate_spike(membrane: torch.Tensor, threshold: float = 1.0) -> torch.Tensor:
hard = (membrane >= threshold).float()
soft = torch.sigmoid((membrane - threshold) * 10.0)
return hard + soft - soft.detach()
class LIFSpikingClassifier(nn.Module):
def __init__(self, hidden_dimensions: int = 64, decay: float = 0.85) -> None:
super().__init__()
self.hidden_dimensions = hidden_dimensions
self.decay = decay
self.input = nn.Linear(64, hidden_dimensions)
self.output = nn.Linear(hidden_dimensions, 10)
def forward(
self,
pixels: torch.Tensor,
*,
timesteps: int = 24,
generator: torch.Generator | None = None,
return_raster: bool = False,
) -> tuple[torch.Tensor, torch.Tensor] | tuple[
torch.Tensor, torch.Tensor, torch.Tensor
]:
membrane = torch.zeros(
len(pixels),
self.hidden_dimensions,
device=pixels.device,
)
logits = torch.zeros(len(pixels), 10, device=pixels.device)
spike_total = torch.zeros((), device=pixels.device)
input_raster = []
for _ in range(timesteps):
random_values = torch.rand(
pixels.shape,
generator=generator,
device=pixels.device,
)
input_spikes = (random_values < pixels).float()
membrane = self.decay * membrane + self.input(input_spikes)
hidden_spikes = surrogate_spike(membrane)
membrane = membrane - hidden_spikes.detach()
logits = logits + self.output(hidden_spikes)
spike_total = spike_total + hidden_spikes.sum()
if return_raster:
input_raster.append(input_spikes)
spike_rate = spike_total / (len(pixels) * self.hidden_dimensions * timesteps)
if return_raster:
return logits / timesteps, spike_rate, torch.stack(input_raster, dim=1)
return logits / timesteps, spike_rate
class MatchedDenseClassifier(nn.Module):
def __init__(self, hidden_dimensions: int = 64) -> None:
super().__init__()
self.first = nn.Linear(64, hidden_dimensions)
self.output = nn.Linear(hidden_dimensions, 10)
def forward(self, pixels: torch.Tensor) -> torch.Tensor:
return self.output(F.gelu(self.first(pixels)))
def parameter_count(model: nn.Module) -> int:
return sum(parameter.numel() for parameter in model.parameters())