gsmc-torch 0.1.0__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Potentially problematic release.
This version of gsmc-torch might be problematic. Click here for more details.
- gsmc_torch-0.1.0/LICENSE +20 -0
- gsmc_torch-0.1.0/PKG-INFO +174 -0
- gsmc_torch-0.1.0/README.md +146 -0
- gsmc_torch-0.1.0/gsmc_torch/__init__.py +31 -0
- gsmc_torch-0.1.0/gsmc_torch/baselines.py +132 -0
- gsmc_torch-0.1.0/gsmc_torch/cell.py +296 -0
- gsmc_torch-0.1.0/gsmc_torch/csrc/__init__.py +36 -0
- gsmc_torch-0.1.0/gsmc_torch/fused_cell.py +255 -0
- gsmc_torch-0.1.0/gsmc_torch/surrogate.py +53 -0
- gsmc_torch-0.1.0/gsmc_torch/transformer.py +86 -0
- gsmc_torch-0.1.0/gsmc_torch/utils.py +73 -0
- gsmc_torch-0.1.0/gsmc_torch.egg-info/PKG-INFO +174 -0
- gsmc_torch-0.1.0/gsmc_torch.egg-info/SOURCES.txt +22 -0
- gsmc_torch-0.1.0/gsmc_torch.egg-info/dependency_links.txt +1 -0
- gsmc_torch-0.1.0/gsmc_torch.egg-info/requires.txt +5 -0
- gsmc_torch-0.1.0/gsmc_torch.egg-info/top_level.txt +1 -0
- gsmc_torch-0.1.0/pyproject.toml +56 -0
- gsmc_torch-0.1.0/setup.cfg +4 -0
- gsmc_torch-0.1.0/tests/test_cell.py +44 -0
- gsmc_torch-0.1.0/tests/test_fused_cell.py +56 -0
- gsmc_torch-0.1.0/tests/test_jit.py +27 -0
- gsmc_torch-0.1.0/tests/test_layer.py +60 -0
- gsmc_torch-0.1.0/tests/test_surrogate.py +23 -0
- gsmc_torch-0.1.0/tests/test_transformer.py +23 -0
gsmc_torch-0.1.0/LICENSE
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
1
|
+
MIT License
|
|
2
|
+
|
|
3
|
+
Copyright (c) 2026 Sumit
|
|
4
|
+
|
|
5
|
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
|
6
|
+
of this software and associated documentation files (the "Software"), to deal
|
|
7
|
+
in the Software without restriction, including without limitation the rights
|
|
8
|
+
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
|
9
|
+
copies of the Software, and to permit persons to whom the Software is
|
|
10
|
+
furnished to do so, subject to the following conditions:
|
|
11
|
+
|
|
12
|
+
The above copyright notice and this permission notice shall be included in all
|
|
13
|
+
copies or substantial portions of the Software.
|
|
14
|
+
|
|
15
|
+
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
|
16
|
+
IMPLIED, INCLUDING BUT NOT LIMITED TO LICENSE OF MERCHANTABILITY, FITNESS FOR
|
|
17
|
+
A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR
|
|
18
|
+
COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER
|
|
19
|
+
IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN
|
|
20
|
+
CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
|
|
@@ -0,0 +1,174 @@
|
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
|
+
Name: gsmc-torch
|
|
3
|
+
Version: 0.1.0
|
|
4
|
+
Summary: Gated Spiking Memory Cell (GSMC) -- Standalone PyTorch plugin for vanishing-gradient-free spiking neural networks
|
|
5
|
+
Author: Sumit
|
|
6
|
+
License: MIT
|
|
7
|
+
Project-URL: Documentation, https://github.com/Griffith-7/GSMC-SNN#readme
|
|
8
|
+
Project-URL: Repository, https://github.com/Griffith-7/GSMC-SNN
|
|
9
|
+
Project-URL: Issues, https://github.com/Griffith-7/GSMC-SNN/issues
|
|
10
|
+
Keywords: spiking-neural-networks,snn,neuromorphic,vanishing-gradients,constant-error-carousel,pytorch,snn-transformer,plugin
|
|
11
|
+
Classifier: Development Status :: 4 - Beta
|
|
12
|
+
Classifier: Intended Audience :: Science/Research
|
|
13
|
+
Classifier: License :: OSI Approved :: MIT License
|
|
14
|
+
Classifier: Programming Language :: Python :: 3
|
|
15
|
+
Classifier: Programming Language :: Python :: 3.9
|
|
16
|
+
Classifier: Programming Language :: Python :: 3.10
|
|
17
|
+
Classifier: Programming Language :: Python :: 3.11
|
|
18
|
+
Classifier: Programming Language :: Python :: 3.12
|
|
19
|
+
Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
|
|
20
|
+
Requires-Python: >=3.9
|
|
21
|
+
Description-Content-Type: text/markdown
|
|
22
|
+
License-File: LICENSE
|
|
23
|
+
Requires-Dist: torch>=2.0
|
|
24
|
+
Provides-Extra: dev
|
|
25
|
+
Requires-Dist: pytest>=7.0; extra == "dev"
|
|
26
|
+
Requires-Dist: torchvision>=0.15; extra == "dev"
|
|
27
|
+
Dynamic: license-file
|
|
28
|
+
|
|
29
|
+
# GSMC-Torch: Gated Spiking Memory Cell Plugin
|
|
30
|
+
|
|
31
|
+
[](https://www.python.org/downloads/)
|
|
32
|
+
[](https://pytorch.org/)
|
|
33
|
+
[](LICENSE)
|
|
34
|
+
|
|
35
|
+
A standalone, production-grade PyTorch plugin library for **Gated Spiking Memory Cells (GSMC)**.
|
|
36
|
+
|
|
37
|
+
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)**.
|
|
38
|
+
|
|
39
|
+
---
|
|
40
|
+
|
|
41
|
+
## Key Features
|
|
42
|
+
|
|
43
|
+
- **Drop-in PyTorch Plugin (`nn.Module`)**: Seamlessly integrates into any PyTorch model architecture (SNN-Transformers, RNNs, ConvNets, hybrid models).
|
|
44
|
+
- **Dual Execution Modes**:
|
|
45
|
+
- `GSMCv2` / `GSMCLayer`: Vectorized sequence-to-sequence layer for fast BPTT over multi-timestep sequence tensors `(batch, time, features)`.
|
|
46
|
+
- `GSMCCell`: Low-level single-timestep stateful cell for step-by-step unrolling, streaming inference, or custom Transformer attention blocks.
|
|
47
|
+
- **Vanishing-Gradient Immunity**: Preserves temporal gradients over $T=784$ steps **34+ orders of magnitude above LIF baselines**.
|
|
48
|
+
- **Split-Gamma Reset ($\gamma_r = 0.1$)**: Eliminates the "reset tax" on temporal gradients while maintaining negative feedback stabilization.
|
|
49
|
+
- **Hardware Energy Model**: Built-in 45nm CMOS energy metrics generator (`compute_energy_per_step`).
|
|
50
|
+
|
|
51
|
+
---
|
|
52
|
+
|
|
53
|
+
## Installation
|
|
54
|
+
|
|
55
|
+
Install directly in editable mode:
|
|
56
|
+
|
|
57
|
+
```bash
|
|
58
|
+
cd gsmc-plugin
|
|
59
|
+
pip install -e .
|
|
60
|
+
```
|
|
61
|
+
|
|
62
|
+
Or install with development dependencies:
|
|
63
|
+
|
|
64
|
+
```bash
|
|
65
|
+
pip install -e ".[dev]"
|
|
66
|
+
```
|
|
67
|
+
|
|
68
|
+
---
|
|
69
|
+
|
|
70
|
+
## Quickstart
|
|
71
|
+
|
|
72
|
+
```python
|
|
73
|
+
import torch
|
|
74
|
+
from gsmc_torch import GSMCv2, GSMCCell
|
|
75
|
+
|
|
76
|
+
# 1. High-level sequence layer (batch_first=True)
|
|
77
|
+
layer = GSMCv2(input_size=1, hidden_size=128)
|
|
78
|
+
|
|
79
|
+
# Input binary spikes: (batch=32, time=784, features=1)
|
|
80
|
+
x = (torch.rand(32, 784, 1) > 0.5).float()
|
|
81
|
+
|
|
82
|
+
# Forward pass -> returns output binary spikes (32, 784, 128)
|
|
83
|
+
spikes = layer(x)
|
|
84
|
+
print(f"Output spikes shape: {spikes.shape}")
|
|
85
|
+
|
|
86
|
+
# 2. Low-level stateful step cell
|
|
87
|
+
cell = GSMCCell(input_size=1, hidden_size=128)
|
|
88
|
+
state = cell.init_state(batch_size=32)
|
|
89
|
+
|
|
90
|
+
x_t = (torch.rand(32, 1) > 0.5).float()
|
|
91
|
+
s_next, state = cell(x_t, state)
|
|
92
|
+
print(f"Step spike shape: {s_next.shape}")
|
|
93
|
+
```
|
|
94
|
+
|
|
95
|
+
---
|
|
96
|
+
|
|
97
|
+
## Integration into Spiking Transformers
|
|
98
|
+
|
|
99
|
+
GSMC can be used directly inside Spiking Attention blocks to provide long-horizon temporal memory:
|
|
100
|
+
|
|
101
|
+
```python
|
|
102
|
+
import torch
|
|
103
|
+
import torch.nn as nn
|
|
104
|
+
from gsmc_torch import GSMCv2
|
|
105
|
+
|
|
106
|
+
class SpikingAttentionBlock(nn.Module):
|
|
107
|
+
def __init__(self, embed_dim=64, hidden_dim=128):
|
|
108
|
+
super().__init__()
|
|
109
|
+
self.q_proj = nn.Linear(embed_dim, embed_dim)
|
|
110
|
+
self.k_proj = nn.Linear(embed_dim, embed_dim)
|
|
111
|
+
self.v_proj = nn.Linear(embed_dim, embed_dim)
|
|
112
|
+
|
|
113
|
+
# GSMC Temporal Memory Cell replacing standard attention decay
|
|
114
|
+
self.gsmc_memory = GSMCv2(input_size=embed_dim, hidden_size=hidden_dim, batch_first=True)
|
|
115
|
+
self.out_proj = nn.Linear(hidden_dim, embed_dim)
|
|
116
|
+
|
|
117
|
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
118
|
+
attn_features = (self.q_proj(x) * self.k_proj(x)) + self.v_proj(x)
|
|
119
|
+
spiking_attn = (attn_features > 0.0).float()
|
|
120
|
+
memory_spikes = self.gsmc_memory(spiking_attn)
|
|
121
|
+
return self.out_proj(memory_spikes)
|
|
122
|
+
```
|
|
123
|
+
|
|
124
|
+
---
|
|
125
|
+
|
|
126
|
+
## Mathematical Formulation
|
|
127
|
+
|
|
128
|
+
### Forward Dynamics
|
|
129
|
+
All affine operations act on binary vectors $X[t], S[t-1] \in \{0, 1\}$, eliminating dense multiplications on neuromorphic hardware:
|
|
130
|
+
|
|
131
|
+
$$\begin{aligned}
|
|
132
|
+
f[t] &= \sigma(W_f X[t] + U_f S[t-1] + b_f) && \text{(Forget gate: initial } b_f=8.0) \\
|
|
133
|
+
i[t] &= \sigma(W_i X[t] + U_i S[t-1]) && \text{(Input write gate)} \\
|
|
134
|
+
o[t] &= \sigma(W_o X[t] + U_o S[t-1] + b_o) && \text{(Output exposure gate: initial } b_o=-2.0) \\
|
|
135
|
+
g[t] &= \tanh(W_g X[t] + U_g S[t-1]) && \text{(Candidate state)} \\
|
|
136
|
+
A[t] &= f[t] \odot M[t-1] + i[t] \odot g[t] && \text{(Memory-bus accumulator)} \\
|
|
137
|
+
V[t] &= o[t] \odot \text{Norm}(A[t]) + W_d X[t] && \text{(Exposed membrane voltage)} \\
|
|
138
|
+
S[t] &= \Theta(V[t] - \theta_t) && \text{(Spike generation)} \\
|
|
139
|
+
M[t] &= A[t] - v_{th} \tilde{S}[t] && \text{(Refractory reset via split } \gamma_r)
|
|
140
|
+
\end{aligned}$$
|
|
141
|
+
|
|
142
|
+
### BPTT Temporal Jacobian
|
|
143
|
+
The temporal Jacobian decomposes into:
|
|
144
|
+
|
|
145
|
+
$$J_t = \frac{\partial M[t]}{\partial M[t-1]} = \operatorname{diag}(f[t]) + \mathcal{B}_t$$
|
|
146
|
+
|
|
147
|
+
Holding $S$ constant gives $\prod_{t=1}^T \operatorname{diag}(f[t])$, a learnable constant-error carousel immune to exponential decay.
|
|
148
|
+
|
|
149
|
+
---
|
|
150
|
+
|
|
151
|
+
## Hardware Energy Footprint (45 nm CMOS)
|
|
152
|
+
|
|
153
|
+
| Model | Dense MACs | Energy ($\text{pJ}/\text{step}/\text{neuron}$) | Energy vs ANN-LSTM |
|
|
154
|
+
| :--- | :--- | :--- | :--- |
|
|
155
|
+
| **VanillaLIF** | 0 | 16.9 pJ | 112× cheaper |
|
|
156
|
+
| **GSMC v2 (`gsmc_torch`)** | **0** | **41.2 pJ** | **46× cheaper** |
|
|
157
|
+
| **SpikingLSTM** | 0 | 40.8 pJ | 46× cheaper |
|
|
158
|
+
| **ANN-LSTM** | Dense ($32 \times 32$) | 1900.8 pJ | 1.0× (Baseline) |
|
|
159
|
+
|
|
160
|
+
---
|
|
161
|
+
|
|
162
|
+
## Running Tests
|
|
163
|
+
|
|
164
|
+
Execute the comprehensive Pytest suite:
|
|
165
|
+
|
|
166
|
+
```bash
|
|
167
|
+
pytest
|
|
168
|
+
```
|
|
169
|
+
|
|
170
|
+
---
|
|
171
|
+
|
|
172
|
+
## License
|
|
173
|
+
|
|
174
|
+
MIT License. See [LICENSE](LICENSE) for details.
|
|
@@ -0,0 +1,146 @@
|
|
|
1
|
+
# GSMC-Torch: Gated Spiking Memory Cell Plugin
|
|
2
|
+
|
|
3
|
+
[](https://www.python.org/downloads/)
|
|
4
|
+
[](https://pytorch.org/)
|
|
5
|
+
[](LICENSE)
|
|
6
|
+
|
|
7
|
+
A standalone, production-grade PyTorch plugin library for **Gated Spiking Memory Cells (GSMC)**.
|
|
8
|
+
|
|
9
|
+
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)**.
|
|
10
|
+
|
|
11
|
+
---
|
|
12
|
+
|
|
13
|
+
## Key Features
|
|
14
|
+
|
|
15
|
+
- **Drop-in PyTorch Plugin (`nn.Module`)**: Seamlessly integrates into any PyTorch model architecture (SNN-Transformers, RNNs, ConvNets, hybrid models).
|
|
16
|
+
- **Dual Execution Modes**:
|
|
17
|
+
- `GSMCv2` / `GSMCLayer`: Vectorized sequence-to-sequence layer for fast BPTT over multi-timestep sequence tensors `(batch, time, features)`.
|
|
18
|
+
- `GSMCCell`: Low-level single-timestep stateful cell for step-by-step unrolling, streaming inference, or custom Transformer attention blocks.
|
|
19
|
+
- **Vanishing-Gradient Immunity**: Preserves temporal gradients over $T=784$ steps **34+ orders of magnitude above LIF baselines**.
|
|
20
|
+
- **Split-Gamma Reset ($\gamma_r = 0.1$)**: Eliminates the "reset tax" on temporal gradients while maintaining negative feedback stabilization.
|
|
21
|
+
- **Hardware Energy Model**: Built-in 45nm CMOS energy metrics generator (`compute_energy_per_step`).
|
|
22
|
+
|
|
23
|
+
---
|
|
24
|
+
|
|
25
|
+
## Installation
|
|
26
|
+
|
|
27
|
+
Install directly in editable mode:
|
|
28
|
+
|
|
29
|
+
```bash
|
|
30
|
+
cd gsmc-plugin
|
|
31
|
+
pip install -e .
|
|
32
|
+
```
|
|
33
|
+
|
|
34
|
+
Or install with development dependencies:
|
|
35
|
+
|
|
36
|
+
```bash
|
|
37
|
+
pip install -e ".[dev]"
|
|
38
|
+
```
|
|
39
|
+
|
|
40
|
+
---
|
|
41
|
+
|
|
42
|
+
## Quickstart
|
|
43
|
+
|
|
44
|
+
```python
|
|
45
|
+
import torch
|
|
46
|
+
from gsmc_torch import GSMCv2, GSMCCell
|
|
47
|
+
|
|
48
|
+
# 1. High-level sequence layer (batch_first=True)
|
|
49
|
+
layer = GSMCv2(input_size=1, hidden_size=128)
|
|
50
|
+
|
|
51
|
+
# Input binary spikes: (batch=32, time=784, features=1)
|
|
52
|
+
x = (torch.rand(32, 784, 1) > 0.5).float()
|
|
53
|
+
|
|
54
|
+
# Forward pass -> returns output binary spikes (32, 784, 128)
|
|
55
|
+
spikes = layer(x)
|
|
56
|
+
print(f"Output spikes shape: {spikes.shape}")
|
|
57
|
+
|
|
58
|
+
# 2. Low-level stateful step cell
|
|
59
|
+
cell = GSMCCell(input_size=1, hidden_size=128)
|
|
60
|
+
state = cell.init_state(batch_size=32)
|
|
61
|
+
|
|
62
|
+
x_t = (torch.rand(32, 1) > 0.5).float()
|
|
63
|
+
s_next, state = cell(x_t, state)
|
|
64
|
+
print(f"Step spike shape: {s_next.shape}")
|
|
65
|
+
```
|
|
66
|
+
|
|
67
|
+
---
|
|
68
|
+
|
|
69
|
+
## Integration into Spiking Transformers
|
|
70
|
+
|
|
71
|
+
GSMC can be used directly inside Spiking Attention blocks to provide long-horizon temporal memory:
|
|
72
|
+
|
|
73
|
+
```python
|
|
74
|
+
import torch
|
|
75
|
+
import torch.nn as nn
|
|
76
|
+
from gsmc_torch import GSMCv2
|
|
77
|
+
|
|
78
|
+
class SpikingAttentionBlock(nn.Module):
|
|
79
|
+
def __init__(self, embed_dim=64, hidden_dim=128):
|
|
80
|
+
super().__init__()
|
|
81
|
+
self.q_proj = nn.Linear(embed_dim, embed_dim)
|
|
82
|
+
self.k_proj = nn.Linear(embed_dim, embed_dim)
|
|
83
|
+
self.v_proj = nn.Linear(embed_dim, embed_dim)
|
|
84
|
+
|
|
85
|
+
# GSMC Temporal Memory Cell replacing standard attention decay
|
|
86
|
+
self.gsmc_memory = GSMCv2(input_size=embed_dim, hidden_size=hidden_dim, batch_first=True)
|
|
87
|
+
self.out_proj = nn.Linear(hidden_dim, embed_dim)
|
|
88
|
+
|
|
89
|
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
90
|
+
attn_features = (self.q_proj(x) * self.k_proj(x)) + self.v_proj(x)
|
|
91
|
+
spiking_attn = (attn_features > 0.0).float()
|
|
92
|
+
memory_spikes = self.gsmc_memory(spiking_attn)
|
|
93
|
+
return self.out_proj(memory_spikes)
|
|
94
|
+
```
|
|
95
|
+
|
|
96
|
+
---
|
|
97
|
+
|
|
98
|
+
## Mathematical Formulation
|
|
99
|
+
|
|
100
|
+
### Forward Dynamics
|
|
101
|
+
All affine operations act on binary vectors $X[t], S[t-1] \in \{0, 1\}$, eliminating dense multiplications on neuromorphic hardware:
|
|
102
|
+
|
|
103
|
+
$$\begin{aligned}
|
|
104
|
+
f[t] &= \sigma(W_f X[t] + U_f S[t-1] + b_f) && \text{(Forget gate: initial } b_f=8.0) \\
|
|
105
|
+
i[t] &= \sigma(W_i X[t] + U_i S[t-1]) && \text{(Input write gate)} \\
|
|
106
|
+
o[t] &= \sigma(W_o X[t] + U_o S[t-1] + b_o) && \text{(Output exposure gate: initial } b_o=-2.0) \\
|
|
107
|
+
g[t] &= \tanh(W_g X[t] + U_g S[t-1]) && \text{(Candidate state)} \\
|
|
108
|
+
A[t] &= f[t] \odot M[t-1] + i[t] \odot g[t] && \text{(Memory-bus accumulator)} \\
|
|
109
|
+
V[t] &= o[t] \odot \text{Norm}(A[t]) + W_d X[t] && \text{(Exposed membrane voltage)} \\
|
|
110
|
+
S[t] &= \Theta(V[t] - \theta_t) && \text{(Spike generation)} \\
|
|
111
|
+
M[t] &= A[t] - v_{th} \tilde{S}[t] && \text{(Refractory reset via split } \gamma_r)
|
|
112
|
+
\end{aligned}$$
|
|
113
|
+
|
|
114
|
+
### BPTT Temporal Jacobian
|
|
115
|
+
The temporal Jacobian decomposes into:
|
|
116
|
+
|
|
117
|
+
$$J_t = \frac{\partial M[t]}{\partial M[t-1]} = \operatorname{diag}(f[t]) + \mathcal{B}_t$$
|
|
118
|
+
|
|
119
|
+
Holding $S$ constant gives $\prod_{t=1}^T \operatorname{diag}(f[t])$, a learnable constant-error carousel immune to exponential decay.
|
|
120
|
+
|
|
121
|
+
---
|
|
122
|
+
|
|
123
|
+
## Hardware Energy Footprint (45 nm CMOS)
|
|
124
|
+
|
|
125
|
+
| Model | Dense MACs | Energy ($\text{pJ}/\text{step}/\text{neuron}$) | Energy vs ANN-LSTM |
|
|
126
|
+
| :--- | :--- | :--- | :--- |
|
|
127
|
+
| **VanillaLIF** | 0 | 16.9 pJ | 112× cheaper |
|
|
128
|
+
| **GSMC v2 (`gsmc_torch`)** | **0** | **41.2 pJ** | **46× cheaper** |
|
|
129
|
+
| **SpikingLSTM** | 0 | 40.8 pJ | 46× cheaper |
|
|
130
|
+
| **ANN-LSTM** | Dense ($32 \times 32$) | 1900.8 pJ | 1.0× (Baseline) |
|
|
131
|
+
|
|
132
|
+
---
|
|
133
|
+
|
|
134
|
+
## Running Tests
|
|
135
|
+
|
|
136
|
+
Execute the comprehensive Pytest suite:
|
|
137
|
+
|
|
138
|
+
```bash
|
|
139
|
+
pytest
|
|
140
|
+
```
|
|
141
|
+
|
|
142
|
+
---
|
|
143
|
+
|
|
144
|
+
## License
|
|
145
|
+
|
|
146
|
+
MIT License. See [LICENSE](LICENSE) for details.
|
|
@@ -0,0 +1,31 @@
|
|
|
1
|
+
"""GSMC-Torch: Gated Spiking Memory Cell Plugin for PyTorch.
|
|
2
|
+
|
|
3
|
+
A drop-in PyTorch library implementing vanishing-gradient-free spiking neurons
|
|
4
|
+
with constant-error carousels (CEC) for SNNs, Spiking Transformers, and Neuromorphic computing.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from .cell import GSMCCell, GSMCv2, GSMCLayer, GSMCState
|
|
8
|
+
from .fused_cell import FusedGSMCv2, FusedGSMCLayer
|
|
9
|
+
from .transformer import GSMCSpikingAttention
|
|
10
|
+
from .surrogate import ATanSurrogate, surrogate_spike
|
|
11
|
+
from .baselines import VanillaLIF, LSNNCell
|
|
12
|
+
from .utils import compute_energy_per_step, compute_energy_table
|
|
13
|
+
|
|
14
|
+
__version__ = "0.1.0"
|
|
15
|
+
|
|
16
|
+
__all__ = [
|
|
17
|
+
"GSMCCell",
|
|
18
|
+
"GSMCv2",
|
|
19
|
+
"GSMCLayer",
|
|
20
|
+
"FusedGSMCv2",
|
|
21
|
+
"FusedGSMCLayer",
|
|
22
|
+
"GSMCSpikingAttention",
|
|
23
|
+
"GSMCState",
|
|
24
|
+
"ATanSurrogate",
|
|
25
|
+
"surrogate_spike",
|
|
26
|
+
"VanillaLIF",
|
|
27
|
+
"LSNNCell",
|
|
28
|
+
"compute_energy_per_step",
|
|
29
|
+
"compute_energy_table",
|
|
30
|
+
"__version__",
|
|
31
|
+
]
|
|
@@ -0,0 +1,132 @@
|
|
|
1
|
+
"""Baseline SNN cells (VanillaLIF, LSNNCell) for drop-in comparative analysis."""
|
|
2
|
+
|
|
3
|
+
import torch
|
|
4
|
+
import torch.nn as nn
|
|
5
|
+
from typing import Optional, Tuple
|
|
6
|
+
from .surrogate import surrogate_spike
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class VanillaLIF(nn.Module):
|
|
10
|
+
"""Standard Leaky Integrate-and-Fire (LIF) neuron layer with fixed decay beta < 1.
|
|
11
|
+
|
|
12
|
+
Args:
|
|
13
|
+
input_size: Feature dimension of inputs.
|
|
14
|
+
hidden_size: Hidden state dimension.
|
|
15
|
+
beta: Membrane potential leak factor (default: 0.9).
|
|
16
|
+
v_th: Spike threshold (default: 1.0).
|
|
17
|
+
batch_first: If True, sequence tensor is (batch, time, features). Default: True.
|
|
18
|
+
"""
|
|
19
|
+
|
|
20
|
+
def __init__(
|
|
21
|
+
self,
|
|
22
|
+
input_size: int,
|
|
23
|
+
hidden_size: int,
|
|
24
|
+
beta: float = 0.9,
|
|
25
|
+
v_th: float = 1.0,
|
|
26
|
+
batch_first: bool = True,
|
|
27
|
+
):
|
|
28
|
+
super().__init__()
|
|
29
|
+
self.input_size = input_size
|
|
30
|
+
self.hidden_size = hidden_size
|
|
31
|
+
self.beta = beta
|
|
32
|
+
self.v_th = v_th
|
|
33
|
+
self.batch_first = batch_first
|
|
34
|
+
|
|
35
|
+
self.W_x = nn.Linear(input_size, hidden_size, bias=True)
|
|
36
|
+
self.V_x = nn.Linear(hidden_size, hidden_size, bias=False)
|
|
37
|
+
nn.init.xavier_uniform_(self.W_x.weight)
|
|
38
|
+
nn.init.orthogonal_(self.V_x.weight)
|
|
39
|
+
|
|
40
|
+
def forward(self, x_seq: torch.Tensor) -> torch.Tensor:
|
|
41
|
+
if not self.batch_first:
|
|
42
|
+
x_seq = x_seq.transpose(0, 1)
|
|
43
|
+
|
|
44
|
+
batch_size, seq_len, _ = x_seq.size()
|
|
45
|
+
device = x_seq.device
|
|
46
|
+
dtype = x_seq.dtype
|
|
47
|
+
|
|
48
|
+
v_t = torch.zeros(batch_size, self.hidden_size, device=device, dtype=dtype)
|
|
49
|
+
s_t = torch.zeros(batch_size, self.hidden_size, device=device, dtype=dtype)
|
|
50
|
+
spikes = []
|
|
51
|
+
|
|
52
|
+
for t in range(seq_len):
|
|
53
|
+
x_t = x_seq[:, t, :]
|
|
54
|
+
i_t = self.W_x(x_t) + self.V_x(s_t)
|
|
55
|
+
v_t = self.beta * v_t + i_t
|
|
56
|
+
s_t = surrogate_spike(v_t - self.v_th)
|
|
57
|
+
v_t = v_t - s_t * self.v_th
|
|
58
|
+
spikes.append(s_t)
|
|
59
|
+
|
|
60
|
+
spikes = torch.stack(spikes, dim=1)
|
|
61
|
+
if not self.batch_first:
|
|
62
|
+
spikes = spikes.transpose(0, 1)
|
|
63
|
+
return spikes
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
class LSNNCell(nn.Module):
|
|
67
|
+
"""Long Short-Term Memory Spiking Neural Network Cell (LSNN) with adaptive threshold.
|
|
68
|
+
|
|
69
|
+
Based on Bellec et al. (2018). Combines leaky integration with refractory threshold adaptation.
|
|
70
|
+
|
|
71
|
+
Args:
|
|
72
|
+
input_size: Feature dimension of inputs.
|
|
73
|
+
hidden_size: Hidden state dimension.
|
|
74
|
+
beta: Membrane leak factor (default: 0.93).
|
|
75
|
+
beta_a: Threshold adaptation decay factor (default: 0.98).
|
|
76
|
+
v_th: Base threshold voltage (default: 1.0).
|
|
77
|
+
beta_scale: Adaptation scaling factor (default: 1.8).
|
|
78
|
+
batch_first: If True, input tensor shape is (batch, time, features).
|
|
79
|
+
"""
|
|
80
|
+
|
|
81
|
+
def __init__(
|
|
82
|
+
self,
|
|
83
|
+
input_size: int,
|
|
84
|
+
hidden_size: int,
|
|
85
|
+
beta: float = 0.93,
|
|
86
|
+
beta_a: float = 0.98,
|
|
87
|
+
v_th: float = 1.0,
|
|
88
|
+
beta_scale: float = 1.8,
|
|
89
|
+
batch_first: bool = True,
|
|
90
|
+
):
|
|
91
|
+
super().__init__()
|
|
92
|
+
self.input_size = input_size
|
|
93
|
+
self.hidden_size = hidden_size
|
|
94
|
+
self.beta = beta
|
|
95
|
+
self.beta_a = beta_a
|
|
96
|
+
self.v_th = v_th
|
|
97
|
+
self.beta_scale = beta_scale
|
|
98
|
+
self.batch_first = batch_first
|
|
99
|
+
|
|
100
|
+
self.W_x = nn.Linear(input_size, hidden_size, bias=True)
|
|
101
|
+
self.V_x = nn.Linear(hidden_size, hidden_size, bias=False)
|
|
102
|
+
nn.init.xavier_uniform_(self.W_x.weight)
|
|
103
|
+
nn.init.orthogonal_(self.V_x.weight)
|
|
104
|
+
|
|
105
|
+
def forward(self, x_seq: torch.Tensor) -> torch.Tensor:
|
|
106
|
+
if not self.batch_first:
|
|
107
|
+
x_seq = x_seq.transpose(0, 1)
|
|
108
|
+
|
|
109
|
+
batch_size, seq_len, _ = x_seq.size()
|
|
110
|
+
device = x_seq.device
|
|
111
|
+
dtype = x_seq.dtype
|
|
112
|
+
|
|
113
|
+
v_t = torch.zeros(batch_size, self.hidden_size, device=device, dtype=dtype)
|
|
114
|
+
s_t = torch.zeros(batch_size, self.hidden_size, device=device, dtype=dtype)
|
|
115
|
+
b_t = torch.zeros(batch_size, self.hidden_size, device=device, dtype=dtype)
|
|
116
|
+
spikes = []
|
|
117
|
+
|
|
118
|
+
for t in range(seq_len):
|
|
119
|
+
x_t = x_seq[:, t, :]
|
|
120
|
+
i_t = self.W_x(x_t) + self.V_x(s_t)
|
|
121
|
+
v_t = self.beta * v_t + i_t
|
|
122
|
+
|
|
123
|
+
theta_eff = self.v_th + self.beta_scale * b_t
|
|
124
|
+
s_t = surrogate_spike(v_t - theta_eff)
|
|
125
|
+
v_t = v_t - s_t * self.v_th
|
|
126
|
+
b_t = self.beta_a * b_t + s_t
|
|
127
|
+
spikes.append(s_t)
|
|
128
|
+
|
|
129
|
+
spikes = torch.stack(spikes, dim=1)
|
|
130
|
+
if not self.batch_first:
|
|
131
|
+
spikes = spikes.transpose(0, 1)
|
|
132
|
+
return spikes
|