blob: 67e9cb6af00a8a5163eee83ee8fec5000a5e525d (
plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
|
import torch
import memtorch
# 1. Standard memTorch setup (Static crossbar weights)
ann_layer = torch.nn.Linear(100, 10)
patched_layer = memtorch.mn.Module.patch_model(ann_layer, memristor_model)
# 2. Your custom SNN wrapper loop
def forward_snn(input_spikes_over_time):
# input_spikes_over_time shape: (time_steps, batch_size, input_dim)
time_steps = input_spikes_over_time.shape[0]
v_mem = torch.zeros(batch_size, 10) # Hidden neuron membrane potentials
output_spikes = []
for t in range(time_steps):
# Pass binary spikes through memTorch's physical crossbar simulation
current_in = patched_layer(input_spikes_over_time[t])
# Leaky Integrate-and-Fire (LIF) logic (Written by you!)
v_mem = 0.9 * v_mem + current_in # Leak & Integrate
# Fire threshold
spike = (v_mem >= 1.0).float()
v_mem[v_mem >= 1.0] = 0.0 # Reset
output_spikes.append(spike)
return torch.stack(output_spikes)
|