summaryrefslogtreecommitdiff
path: root/spires_memristor_sim
diff options
context:
space:
mode:
Diffstat (limited to 'spires_memristor_sim')
-rw-r--r--spires_memristor_sim/README.md12
-rw-r--r--spires_memristor_sim/memTorch_SNN.py28
-rw-r--r--spires_memristor_sim/spires_interface.py35
-rw-r--r--spires_memristor_sim/torch_reservoir.py43
4 files changed, 118 insertions, 0 deletions
diff --git a/spires_memristor_sim/README.md b/spires_memristor_sim/README.md
new file mode 100644
index 0000000..bc89498
--- /dev/null
+++ b/spires_memristor_sim/README.md
@@ -0,0 +1,12 @@
+# Spires reservoir simulated on memristor crossbar using memTorch
+
+## Goal
+The goal is to able to simulate a spires reservoir on a memristor crossbar using
+a custom memristor spice model(HIL), thus you can run multiple simulations with
+multiple models to compare how each performs using benchmark results.
+
+## task
+The plan right now is to give the cart-pole task for benchmarking.
+
+## metrics
+
diff --git a/spires_memristor_sim/memTorch_SNN.py b/spires_memristor_sim/memTorch_SNN.py
new file mode 100644
index 0000000..67e9cb6
--- /dev/null
+++ b/spires_memristor_sim/memTorch_SNN.py
@@ -0,0 +1,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)
diff --git a/spires_memristor_sim/spires_interface.py b/spires_memristor_sim/spires_interface.py
new file mode 100644
index 0000000..3f1e302
--- /dev/null
+++ b/spires_memristor_sim/spires_interface.py
@@ -0,0 +1,35 @@
+import ctypes
+import numpy
+import torch
+
+spires_lib = ctypes.CDLL("../spires/build/.libspires.so")
+
+spires_lib.spires_reservoir_step.argtypes = [ctypes.POINTER(ctypes.c_float)]
+spires_lib.spires_reservoir_step.restype = None
+
+#change currents from tensor to a flat C pointer array for spires to read
+def send_currents_to_spires(currents):
+ # ----- extract from pytorhc graph -----
+ # .detach() removes it from auto gradient tracking
+ # .cpu() make sure data is in RAM, not VRAM
+ # .numpy() maps it to numpy array
+ numpy_array = torch_tensor.detach().cpu().numpy()
+
+ # ----- makes suren layout matches 32 bit float C-array -----
+ # .astype(npfloat32) forces standard CC float precistion
+ # .flatten() makes sure the memory is a 1D block
+ contiguous_array = numpy_array.astype(np.float32).flatten()
+
+ # ----- Ensure raw memory pointer -----
+ c_float_ptr = contiguous_array.ctypes.data_as(ctypes.POINTER(ctypes.c_float))
+ array_size = contiguous_array.size
+
+ # ----- Call the C library -----
+ spires_lib.spires_reservoir_step(c_float_ptr, array_size)
+ return contiguous_array
+
+ # NEED TO MODIFY SPIRES FOR THIS?? :((
+
+def read_spikes_from_spires():
+ placeholder = 0
+ return 0
diff --git a/spires_memristor_sim/torch_reservoir.py b/spires_memristor_sim/torch_reservoir.py
new file mode 100644
index 0000000..c3017c5
--- /dev/null
+++ b/spires_memristor_sim/torch_reservoir.py
@@ -0,0 +1,43 @@
+import torch
+import memtorch
+from spires_interface import (
+ send_currents_to_spires,
+ read_spikes_from_spires,
+)
+
+NUM_INPUTS = 4
+NUM_NEURONS = 800
+
+input_layer = torch.nn.Linear(NUM_INPUTS, NUM_NEURONS, bias=False)
+#this is a recurrent layer sort of???
+reservoir_layer = torch.nn.Linear(NUM_NEURONS, NUM_NEURONS, bias=False)
+
+#make the reservoir sparse and random
+with torch.no_grad():
+ reservoir_layer.weight.data.normal_(0.0, 0.5) #random weight
+
+ #mask so 10% of connections exist
+ mask = (torch.rand(NUM_NEURONS, NUM_NEURONS) < 0.10).float()
+ reservoir_layer.weight.data *= mask
+
+# patch layers into memristor crossbars with memTorch
+memristor_model = memtorch.bh.memristor.VTEAM
+mem_input_layer = memtorch.mn.Module.patch_model(input_layer, memristor_model)
+mem_reservoir_layer = memtorch.mn.Module.patch_model(reservoir_layer, memristor_model)
+
+#run spiking loop with spires
+previous_spikes = torch.zeros(1, NUM_NEURONS)
+for step in range(500):
+ input_tensor = torch.tensor(cartpole_state).float().unsqueeze(0)
+ currents_in = mem_input_layer(input_tensor)
+
+ currents_recv = mem_reservoir_layer(previous_spikes)
+
+ total_currents = currents_in + currents_recv
+
+ send_currents_to_spires(total_currents)
+
+ current_spikes_np = read_spikes_from_spires()
+ previous_spikes = torch.from_numpy(current_spikes_np).float().unsqueeze(0)
+
+#pass current spikes np to readout layer