Skip to main content

GSMC-Torch: Gated Spiking Memory Cell Plugin

Python 3.9+ PyTorch 2.0+ License: MIT

A standalone, production-grade PyTorch plugin library for Gated Spiking Memory Cells (GSMC).

GSMC provides a fundamental solution to the vanishing gradient problem in Spiking Neural Networks (SNNs) and Spiking Transformers via a learnable Constant-Error Carousel (CEC), while maintaining strict binary inter-neuron communication and multiplier-free operations (~41 pJ/step/neuron on 45 nm, 46× cheaper than ANN-LSTM).


Key Features

  • Drop-in PyTorch Plugin (nn.Module): Seamlessly integrates into any PyTorch model architecture (SNN-Transformers, RNNs, ConvNets, hybrid models).
  • Dual Execution Modes:
    • GSMCv2 / GSMCLayer: Vectorized sequence-to-sequence layer for fast BPTT over multi-timestep sequence tensors (batch, time, features).
    • GSMCCell: Low-level single-timestep stateful cell for step-by-step unrolling, streaming inference, or custom Transformer attention blocks.
  • Vanishing-Gradient Immunity: Preserves temporal gradients over $T=784$ steps 34+ orders of magnitude above LIF baselines.
  • Split-Gamma Reset ($\gamma_r = 0.1$): Eliminates the "reset tax" on temporal gradients while maintaining negative feedback stabilization.
  • Hardware Energy Model: Built-in 45nm CMOS energy metrics generator (compute_energy_per_step).

Installation

Install directly in editable mode:

cd gsmc-plugin
pip install -e .

Or install with development dependencies:

pip install -e ".[dev]"

Quickstart

import torch
from gsmc_torch import GSMCv2, GSMCCell

# 1. High-level sequence layer (batch_first=True)
layer = GSMCv2(input_size=1, hidden_size=128)

# Input binary spikes: (batch=32, time=784, features=1)
x = (torch.rand(32, 784, 1) > 0.5).float()

# Forward pass -> returns output binary spikes (32, 784, 128)
spikes = layer(x)
print(f"Output spikes shape: {spikes.shape}")

# 2. Low-level stateful step cell
cell = GSMCCell(input_size=1, hidden_size=128)
state = cell.init_state(batch_size=32)

x_t = (torch.rand(32, 1) > 0.5).float()
s_next, state = cell(x_t, state)
print(f"Step spike shape: {s_next.shape}")

Integration into Spiking Transformers

GSMC can be used directly inside Spiking Attention blocks to provide long-horizon temporal memory:

import torch
import torch.nn as nn
from gsmc_torch import GSMCv2

class SpikingAttentionBlock(nn.Module):
    def __init__(self, embed_dim=64, hidden_dim=128):
        super().__init__()
        self.q_proj = nn.Linear(embed_dim, embed_dim)
        self.k_proj = nn.Linear(embed_dim, embed_dim)
        self.v_proj = nn.Linear(embed_dim, embed_dim)
        
        # GSMC Temporal Memory Cell replacing standard attention decay
        self.gsmc_memory = GSMCv2(input_size=embed_dim, hidden_size=hidden_dim, batch_first=True)
        self.out_proj = nn.Linear(hidden_dim, embed_dim)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        attn_features = (self.q_proj(x) * self.k_proj(x)) + self.v_proj(x)
        spiking_attn = (attn_features > 0.0).float()
        memory_spikes = self.gsmc_memory(spiking_attn)
        return self.out_proj(memory_spikes)

Mathematical Formulation

Forward Dynamics

All affine operations act on binary vectors $X[t], S[t-1] \in {0, 1}$, eliminating dense multiplications on neuromorphic hardware:

$$\begin{aligned} f[t] &= \sigma(W_f X[t] + U_f S[t-1] + b_f) && \text{(Forget gate: initial } b_f=8.0) \ i[t] &= \sigma(W_i X[t] + U_i S[t-1]) && \text{(Input write gate)} \ o[t] &= \sigma(W_o X[t] + U_o S[t-1] + b_o) && \text{(Output exposure gate: initial } b_o=-2.0) \ g[t] &= \tanh(W_g X[t] + U_g S[t-1]) && \text{(Candidate state)} \ A[t] &= f[t] \odot M[t-1] + i[t] \odot g[t] && \text{(Memory-bus accumulator)} \ V[t] &= o[t] \odot \text{Norm}(A[t]) + W_d X[t] && \text{(Exposed membrane voltage)} \ S[t] &= \Theta(V[t] - \theta_t) && \text{(Spike generation)} \ M[t] &= A[t] - v_{th} \tilde{S}[t] && \text{(Refractory reset via split } \gamma_r) \end{aligned}$$

BPTT Temporal Jacobian

The temporal Jacobian decomposes into:

$$J_t = \frac{\partial M[t]}{\partial M[t-1]} = \operatorname{diag}(f[t]) + \mathcal{B}_t$$

Holding $S$ constant gives $\prod_{t=1}^T \operatorname{diag}(f[t])$, a learnable constant-error carousel immune to exponential decay.


Hardware Energy Footprint (45 nm CMOS)

Model Dense MACs Energy ($\text{pJ}/\text{step}/\text{neuron}$) Energy vs ANN-LSTM
VanillaLIF 0 16.9 pJ 112× cheaper
GSMC v2 (gsmc_torch) 0 41.2 pJ 46× cheaper
SpikingLSTM 0 40.8 pJ 46× cheaper
ANN-LSTM Dense ($32 \times 32$) 1900.8 pJ 1.0× (Baseline)

Running Tests

Execute the comprehensive Pytest suite:

pytest

License

MIT License. See LICENSE for details.

Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

gsmc_torch-0.1.0.tar.gz (17.9 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

gsmc_torch-0.1.0-py3-none-any.whl (16.9 kB view details)

Uploaded Python 3

File details

Details for the file gsmc_torch-0.1.0.tar.gz.

File metadata

  • Download URL: gsmc_torch-0.1.0.tar.gz
  • Upload date:
  • Size: 17.9 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/7.0.0 CPython/3.10.11

File hashes

Hashes for gsmc_torch-0.1.0.tar.gz
Algorithm Hash digest
SHA256 73ce91fed17a76d71ac46c0e516e68c26a9c6783f3ba7c5a76691ef6deb60be3
MD5 a250e368cc8c7394c8a76bb43519edcc
BLAKE2b-256 86159e3742fc432037602dbab1283a67a69ecf170f83f180de71a017b859bc8d

See more details on using hashes here.

File details

Details for the file gsmc_torch-0.1.0-py3-none-any.whl.

File metadata

  • Download URL: gsmc_torch-0.1.0-py3-none-any.whl
  • Upload date:
  • Size: 16.9 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/7.0.0 CPython/3.10.11

File hashes

Hashes for gsmc_torch-0.1.0-py3-none-any.whl
Algorithm Hash digest
SHA256 ad5af2fabaef96098ed7d600ee647c65cd5722b7361efb7897f4a0d02c3b404a
MD5 3d40bb07a3a98a49b40604d06db74041
BLAKE2b-256 4646b4f2b7f1cc868fc11e1cb194ecb80e68c146ff5db657b01fc472dd951edf

See more details on using hashes here.

Release history Release notifications | RSS feed

0.1.1

2 files

This release

0.1.0 This release

2 files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page