diff options
Diffstat (limited to 'neurobench_testing')
| -rw-r--r-- | neurobench_testing/custom_memristor_model.py | 42 | ||||
| -rw-r--r-- | neurobench_testing/memTorch_testing.py | 81 | ||||
| -rw-r--r-- | neurobench_testing/memTorch_testing_custom_memristor.py | 160 |
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") |
