diff options
Diffstat (limited to 'neurobench_testing/custom_memristor_model.py')
| -rw-r--r-- | neurobench_testing/custom_memristor_model.py | 42 |
1 files changed, 42 insertions, 0 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 + + |
