summaryrefslogtreecommitdiff
path: root/neurobench_testing
diff options
context:
space:
mode:
authorTanner Robison <[email protected]>2026-06-19 00:02:30 -0700
committerTanner Robison <[email protected]>2026-06-19 00:02:30 -0700
commit30dc9553e8a7a7d603077a80238782298d0b9349 (patch)
tree5fcce0b3632f597a53ecafbecb1b16def93d31a5 /neurobench_testing
parent00ba3dd475eadaedc0984401eb0272554439fb66 (diff)
customizeable memristor model
started trying to implement a customizeable memristor model to pass to memTorch, porting a model I got from kaevin too fit the pyTorch framework
Diffstat (limited to 'neurobench_testing')
-rw-r--r--neurobench_testing/custom_memristor_model.py42
-rw-r--r--neurobench_testing/memTorch_testing.py81
-rw-r--r--neurobench_testing/memTorch_testing_custom_memristor.py160
3 files changed, 269 insertions, 14 deletions
diff --git a/neurobench_testing/custom_memristor_model.py b/neurobench_testing/custom_memristor_model.py
new file mode 100644
index 0000000..8bdcde0
--- /dev/null
+++ b/neurobench_testing/custom_memristor_model.py
@@ -0,0 +1,42 @@
+import torch
+from memtorch.bh.memristor.Memristor import Memristor
+
+class MemtorchMemristor(Memristor):
+ def __init__(
+ self,
+ k_off = 1.0, # switching rate for off state
+ k_on = -1.0, # switching rate for on state
+ alpha_off = 5, # exponent controlling nonlinearity
+ alpha_on = 5, # exponent controlling nonlinearity
+ i_off = 0.5e-3, # threshhold current to trigger off state
+ i_on = 0.5e-3, # threshold current to trigger on state
+ r_on = 1e3, # maximum resistance
+ r_off = 10e3, # minimum resistance
+ p = 2, # window function exponent
+ **kwargs
+ ):
+ #initializing base memristor class
+ super(MemtorchMemristor, self).__init__(r_off=r_off, r_on=r_on, **kwargs)
+
+ # hyper parameters
+ self.k_off = k_off
+ self.k_on = k_on
+ self.alpha_off = alpha_off
+ self.alpha_on = alpha_on
+ self.i_on = i_on
+ self.i_off = i_off
+ self.p = p
+
+ # makes sure w starts in valid state
+ if not hasattr(self, 'w'):
+ self.w = torch.tensor(0.5)
+
+
+ """
+ Updates w and computes new resistance
+
+ """
+ def step(self, v, dt):
+ i = v / self.r_curr
+
+
diff --git a/neurobench_testing/memTorch_testing.py b/neurobench_testing/memTorch_testing.py
index a44fd98..5c3d828 100644
--- a/neurobench_testing/memTorch_testing.py
+++ b/neurobench_testing/memTorch_testing.py
@@ -1,14 +1,19 @@
-from memtorch import memristor
+"""
+This is a program to test out using memTorch with neurobench.
+June 17th, 2026
+Author: Tanner Robison,
+Teuscher Lab
+"""
import torch
import torch.nn as nn
import snntorch as snn
+from snntorch import surrogate
from torch.utils.data import DataLoader
from neurobench.models import SNNTorchModel
from neurobench.benchmarks import Benchmark, benchmark
from neurobench.datasets import SpeechCommands
-
from neurobench.metrics.workload import (
ActivationSparsity,
SynapticOperations,
@@ -29,8 +34,7 @@ from memtorch.map.Parameter import naive_map
from memtorch.bh.memristor import VTEAM
beta = 0.9
-
-class SimpleSNN(nn.Module):
+class SNN(nn.Module):
def __init__(self):
super().__init__()
@@ -54,36 +58,46 @@ class SimpleSNN(nn.Module):
return spk2, mem2
device = torch.device("cpu")
-net = SimpleSNN().to(device)
+spike_grad = surrogate.fast_sigmoid()
+net = nn.Sequential(
+ nn.Flatten(),
+ nn.Linear(20, 256),
+ snn.Leaky(beta=beta, spike_grad=spike_grad, init_hidden=True),
+ nn.Linear(256, 256),
+ snn.Leaky(beta=beta, spike_grad=spike_grad, init_hidden=True),
+ nn.Linear(256, 256),
+ snn.Leaky(beta=beta, spike_grad=spike_grad, init_hidden=True),
+ nn.Linear(256, 35),
+ snn.Leaky(beta=beta, spike_grad=spike_grad, init_hidden=True, output=True),
+)
#memristor patch
reference_memristor = VTEAM()
-print("Patching model to memristive crossbar")
+net.load_state_dict(torch.load("examples/gsc/model_data/s2s_gsc_snntorch", map_location=device))
+
patched_net = patch_model(
copy.deepcopy(net),
memristor_model=reference_memristor,
- memristor_model_params={},
+ memristor_model_params={'time_series_resolution': 1e-8},
mapping_routine=naive_map,
transistor=True,
+ tile_shape=(128, 128),
ADC_resolution=8,
- use_bindings=False
+ use_bindings=True
)
-model = SNNTorchModel(patched_net)
-
static_metrics = [Footprint, ConnectionSparsity]
workload_metrics = [ActivationSparsity, SynapticOperations, ClassificationAccuracy]
# data loader here maybe??
test_set = SpeechCommands(path="data/SpeechCommands/", subset="testing")
-test_set_loader = DataLoader(test_set, batch_size=50, shuffle=True)
+test_set_loader = DataLoader(test_set, batch_size=500, shuffle=True)
pre_processor = [S2SPreProcessor(device=device)]
post_processor = [ChooseMaxCount()]
-# print("Layer 1 Beta:", net[2].beta)
-
+model = SNNTorchModel(net)
benchmark = Benchmark(
model,
test_set_loader,
@@ -93,7 +107,46 @@ benchmark = Benchmark(
)
results = benchmark.run()
-print(results)
+print("\n\n----- IDEAL BENCHMARK -----")
+print(f"Footprint: {results['Footprint']}")
+print(f"Connection Sparsity: {results['ConnectionSparsity']}")
+print(f"Activation Sparsity: {results['ActivationSparsity']}")
+print(f"Synaptic Operations: {results['SynapticOperations']}")
+print(f"Classification Accuracy: {results['ClassificationAccuracy']}")
+
+#energy calculations
+ENERGY_PER_MAC = 0.9e-12
+ENERGY_PER_AC = 0.1e-12
+
+macs = results['SynapticOperations']['Effective_MACs']
+acs = results['SynapticOperations']['Effective_ACs']
+
+total_energy = (macs * ENERGY_PER_MAC) + (acs * ENERGY_PER_AC)
+
+print("Energy Report:")
+print(f"Total Operations: {macs} MACS, {acs} acs ")
+print(f"Calculated Energy cost: {total_energy} joules per batch")
+
+
+model = SNNTorchModel(patched_net)
+benchmark = Benchmark(
+ model,
+ test_set_loader,
+ pre_processor,
+ post_processor,
+ [static_metrics, workload_metrics]
+)
+
+with torch.no_grad(): #makes sure its in inference mode
+ #otherwise you get memory leaks
+ results = benchmark.run()
+
+print("\n\n----- MEMRISTOR BENCHMARK -----")
+print(f"Footprint: {results['Footprint']}")
+print(f"Connection Sparsity: {results['ConnectionSparsity']}")
+print(f"Activation Sparsity: {results['ActivationSparsity']}")
+print(f"Synaptic Operations: {results['SynapticOperations']}")
+print(f"Classification Accuracy: {results['ClassificationAccuracy']}")
#energy calculations
ENERGY_PER_MAC = 0.9e-12
diff --git a/neurobench_testing/memTorch_testing_custom_memristor.py b/neurobench_testing/memTorch_testing_custom_memristor.py
new file mode 100644
index 0000000..64f338e
--- /dev/null
+++ b/neurobench_testing/memTorch_testing_custom_memristor.py
@@ -0,0 +1,160 @@
+from memtorch import memristor
+import torch
+import torch.nn as nn
+import snntorch as snn
+from snntorch import surrogate
+
+from torch.utils.data import DataLoader
+
+from neurobench.models import SNNTorchModel
+from neurobench.benchmarks import Benchmark, benchmark
+from neurobench.datasets import SpeechCommands
+
+from neurobench.metrics.workload import (
+ ActivationSparsity,
+ SynapticOperations,
+ ClassificationAccuracy,
+)
+
+from neurobench.metrics.static import (
+ Footprint,
+ ConnectionSparsity,
+)
+
+from neurobench.processors.preprocessors import S2SPreProcessor
+from neurobench.processors.postprocessors import ChooseMaxCount
+
+import copy
+from memtorch.mn.Module import patch_model
+from memtorch.map.Parameter import naive_map
+from memtorch.bh.memristor import VTEAM
+
+from kaevin_memristor import TEAMMemristor
+
+beta = 0.9
+class SNN(nn.Module):
+ def __init__(self):
+ super().__init__()
+
+ #standard layers
+ self.fc1 = nn.Linear(20, 128)
+ self.fc2 = nn.Linear(128, 35)
+
+ #Spiking neurons
+ self.lif1 = snn.Leaky(beta=beta, init_hidden=True)
+ self.lif2 = snn.Leaky(beta=beta, init_hidden=True, output=True)
+
+ def forward(self, x):
+ x = x.view(x.size(0), -1)
+
+ cur1 = self.fc1(x)
+ spk1 = self.lif1(cur1)
+
+ cur2 = self.fc2(spk1)
+ spk2, mem2 = self.lif2(cur2)
+
+ return spk2, mem2
+
+device = torch.device("cpu")
+spike_grad = surrogate.fast_sigmoid()
+net = nn.Sequential(
+ nn.Flatten(),
+ nn.Linear(20, 256),
+ snn.Leaky(beta=beta, spike_grad=spike_grad, init_hidden=True),
+ nn.Linear(256, 256),
+ snn.Leaky(beta=beta, spike_grad=spike_grad, init_hidden=True),
+ nn.Linear(256, 256),
+ snn.Leaky(beta=beta, spike_grad=spike_grad, init_hidden=True),
+ nn.Linear(256, 35),
+ snn.Leaky(beta=beta, spike_grad=spike_grad, init_hidden=True, output=True),
+)
+
+#memristor patch
+reference_memristor = TEAMMemristor
+
+net.load_state_dict(torch.load("examples/gsc/model_data/s2s_gsc_snntorch", map_location=device))
+
+patched_net = patch_model(
+ copy.deepcopy(net),
+ memristor_model=reference_memristor,
+ memristor_model_params={'time_series_resolution': 1e-8},
+ mapping_routine=naive_map,
+ transistor=True,
+ tile_shape=(128, 128),
+ ADC_resolution=8,
+ use_bindings=True
+)
+
+static_metrics = [Footprint, ConnectionSparsity]
+workload_metrics = [ActivationSparsity, SynapticOperations, ClassificationAccuracy]
+
+# data loader here maybe??
+test_set = SpeechCommands(path="data/SpeechCommands/", subset="testing")
+test_set_loader = DataLoader(test_set, batch_size=500, shuffle=True)
+
+pre_processor = [S2SPreProcessor(device=device)]
+post_processor = [ChooseMaxCount()]
+
+model = SNNTorchModel(net)
+benchmark = Benchmark(
+ model,
+ test_set_loader,
+ pre_processor,
+ post_processor,
+ [static_metrics, workload_metrics]
+)
+
+results = benchmark.run()
+print("\n\n----- IDEAL BENCHMARK -----")
+print(f"Footprint: {results['Footprint']}")
+print(f"Connection Sparsity: {results['ConnectionSparsity']}")
+print(f"Activation Sparsity: {results['ActivationSparsity']}")
+print(f"Synaptic Operations: {results['SynapticOperations']}")
+print(f"Classification Accuracy: {results['ClassificationAccuracy']}")
+
+#energy calculations
+ENERGY_PER_MAC = 0.9e-12
+ENERGY_PER_AC = 0.1e-12
+
+macs = results['SynapticOperations']['Effective_MACs']
+acs = results['SynapticOperations']['Effective_ACs']
+
+total_energy = (macs * ENERGY_PER_MAC) + (acs * ENERGY_PER_AC)
+
+print("Energy Report:")
+print(f"Total Operations: {macs} MACS, {acs} acs ")
+print(f"Calculated Energy cost: {total_energy} joules per batch")
+
+
+model = SNNTorchModel(patched_net)
+benchmark = Benchmark(
+ model,
+ test_set_loader,
+ pre_processor,
+ post_processor,
+ [static_metrics, workload_metrics]
+)
+
+with torch.no_grad(): #makes sure its in inference mode
+ #otherwise you get memory leaks
+ results = benchmark.run()
+
+print("\n\n----- MEMRISTOR BENCHMARK -----")
+print(f"Footprint: {results['Footprint']}")
+print(f"Connection Sparsity: {results['ConnectionSparsity']}")
+print(f"Activation Sparsity: {results['ActivationSparsity']}")
+print(f"Synaptic Operations: {results['SynapticOperations']}")
+print(f"Classification Accuracy: {results['ClassificationAccuracy']}")
+
+#energy calculations
+ENERGY_PER_MAC = 0.9e-12
+ENERGY_PER_AC = 0.1e-12
+
+macs = results['SynapticOperations']['Effective_MACs']
+acs = results['SynapticOperations']['Effective_ACs']
+
+total_energy = (macs * ENERGY_PER_MAC) + (acs * ENERGY_PER_AC)
+
+print("Energy Report:")
+print(f"Total Operations: {macs} MACS, {acs} acs ")
+print(f"Calculated Energy cost: {total_energy} joules per batch")