summaryrefslogtreecommitdiff
path: root/spires_memristor_sim/spires_interface.py
diff options
context:
space:
mode:
authorTanner Robison <[email protected]>2026-06-25 15:27:45 -0700
committerTanner Robison <[email protected]>2026-06-25 15:49:24 -0700
commit3c9a3d7d92d155545ce4892524f00dcd092b52d3 (patch)
tree2b797ec584b7496da7764cb8fce8db4aad9e8ba4 /spires_memristor_sim/spires_interface.py
parentc53c13b873ea1e69ed260844268f11708d4b4a5b (diff)
using memTorch to simulate spires reservoir
focusing on trying to simulate a spires reservoir on a memristor crossbar using memTorch
Diffstat (limited to 'spires_memristor_sim/spires_interface.py')
-rw-r--r--spires_memristor_sim/spires_interface.py35
1 files changed, 35 insertions, 0 deletions
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