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.

@@ -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
+ [![Python 3.9+](https://img.shields.io/badge/python-3.9%2B-blue.svg)](https://www.python.org/downloads/)
32
+ [![PyTorch 2.0+](https://img.shields.io/badge/pytorch-2.0%2B-ee4c2c.svg)](https://pytorch.org/)
33
+ [![License: MIT](https://img.shields.io/badge/license-MIT-green.svg)](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
+ [![Python 3.9+](https://img.shields.io/badge/python-3.9%2B-blue.svg)](https://www.python.org/downloads/)
4
+ [![PyTorch 2.0+](https://img.shields.io/badge/pytorch-2.0%2B-ee4c2c.svg)](https://pytorch.org/)
5
+ [![License: MIT](https://img.shields.io/badge/license-MIT-green.svg)](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