summaryrefslogtreecommitdiff
path: root/custom_benchmark/memTorch_SNN.py
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)