ngclearn 1.0b0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- ngclearn/__init__.py +23 -0
- ngclearn/commands/__init__.py +1 -0
- ngclearn/components/__init__.py +19 -0
- ngclearn/components/baseComponentTemplate.py +47 -0
- ngclearn/components/input_encoders/__init__.py +2 -0
- ngclearn/components/input_encoders/bernoulliCell.py +123 -0
- ngclearn/components/input_encoders/poissonCell.py +132 -0
- ngclearn/components/neurons/__init__.py +9 -0
- ngclearn/components/neurons/graded/__init__.py +4 -0
- ngclearn/components/neurons/graded/gaussianErrorCell.py +174 -0
- ngclearn/components/neurons/graded/laplacianErrorCell.py +174 -0
- ngclearn/components/neurons/graded/rateCell.py +201 -0
- ngclearn/components/neurons/spiking/LIFCell.py +305 -0
- ngclearn/components/neurons/spiking/__init__.py +5 -0
- ngclearn/components/neurons/spiking/izhikevichCell.py +150 -0
- ngclearn/components/neurons/spiking/quadLIFCell.py +327 -0
- ngclearn/components/neurons/spiking/sLIFCell.py +378 -0
- ngclearn/components/other/__init__.py +2 -0
- ngclearn/components/other/expKernel.py +129 -0
- ngclearn/components/other/varTrace.py +150 -0
- ngclearn/components/synapses/__init__.py +3 -0
- ngclearn/components/synapses/hebbian/__init__.py +3 -0
- ngclearn/components/synapses/hebbian/expSTDPSynapse.py +234 -0
- ngclearn/components/synapses/hebbian/hebbianSynapse.py +334 -0
- ngclearn/components/synapses/hebbian/traceSTDPSynapse.py +264 -0
- ngclearn/components/wrappers.py +8 -0
- ngclearn/utils/__init__.py +0 -0
- ngclearn/utils/density/__init__.py +0 -0
- ngclearn/utils/density/gmm.py +82 -0
- ngclearn/utils/io_utils.py +67 -0
- ngclearn/utils/model_utils.py +274 -0
- ngclearn/utils/optim/__init__.py +2 -0
- ngclearn/utils/optim/adam.py +88 -0
- ngclearn/utils/optim/opt.py +25 -0
- ngclearn/utils/optim/sgd.py +38 -0
- ngclearn/utils/patch_utils.py +79 -0
- ngclearn/utils/viz/__init__.py +0 -0
- ngclearn/utils/viz/dim_reduce.py +91 -0
- ngclearn/utils/viz/raster.py +177 -0
- ngclearn/utils/viz/synapse_plot.py +151 -0
- ngclearn-1.0b0.dist-info/AUTHORS +15 -0
- ngclearn-1.0b0.dist-info/LICENSE +29 -0
- ngclearn-1.0b0.dist-info/METADATA +188 -0
- ngclearn-1.0b0.dist-info/RECORD +46 -0
- ngclearn-1.0b0.dist-info/WHEEL +5 -0
- ngclearn-1.0b0.dist-info/top_level.txt +1 -0
ngclearn/__init__.py
ADDED
|
@@ -0,0 +1,23 @@
|
|
|
1
|
+
import sys
|
|
2
|
+
import subprocess
|
|
3
|
+
import pkg_resources
|
|
4
|
+
from pkg_resources import get_distribution
|
|
5
|
+
#from pathlib import Path
|
|
6
|
+
#from sys import argv
|
|
7
|
+
|
|
8
|
+
__version__ = get_distribution('ngclearn').version
|
|
9
|
+
|
|
10
|
+
#required = {'ngcsimlib', 'jax', 'jaxlib'} ## list of core ngclearn dependencies
|
|
11
|
+
required = {'ngcsimlib', 'jax', 'jaxlib'}
|
|
12
|
+
installed = {pkg.key for pkg in pkg_resources.working_set}
|
|
13
|
+
missing = required - installed
|
|
14
|
+
|
|
15
|
+
for key in required:
|
|
16
|
+
if key in missing:
|
|
17
|
+
raise ImportError(str(key) + ", a core dependency of ngclearn, is not " \
|
|
18
|
+
"currently installed!")
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
## Needed to preload is called before anything in ngclearn
|
|
22
|
+
import ngcsimlib
|
|
23
|
+
from ngcsimlib.controller import Controller
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
from ngcsimlib.commands import *
|
|
@@ -0,0 +1,19 @@
|
|
|
1
|
+
## point to rate-coded cell componet types
|
|
2
|
+
from .neurons.graded.rateCell import RateCell
|
|
3
|
+
from .neurons.graded.gaussianErrorCell import GaussianErrorCell
|
|
4
|
+
from .neurons.graded.laplacianErrorCell import LaplacianErrorCell
|
|
5
|
+
## point to standard spiking cell component types
|
|
6
|
+
from .neurons.spiking.sLIFCell import SLIFCell
|
|
7
|
+
from .neurons.spiking.LIFCell import LIFCell
|
|
8
|
+
from .neurons.spiking.quadLIFCell import QuadLIFCell
|
|
9
|
+
from .neurons.spiking.izhikevichCell import IzhikevichCell
|
|
10
|
+
## point to transformer/operater component types
|
|
11
|
+
from .other.varTrace import VarTrace
|
|
12
|
+
from .other.expKernel import ExpKernel
|
|
13
|
+
## point to input encoder component types
|
|
14
|
+
from .input_encoders.bernoulliCell import BernoulliCell
|
|
15
|
+
from .input_encoders.poissonCell import PoissonCell
|
|
16
|
+
## point to synapse component types
|
|
17
|
+
from .synapses.hebbian.hebbianSynapse import HebbianSynapse
|
|
18
|
+
from .synapses.hebbian.traceSTDPSynapse import TraceSTDPSynapse
|
|
19
|
+
from .synapses.hebbian.expSTDPSynapse import ExpSTDPSynapse
|
|
@@ -0,0 +1,47 @@
|
|
|
1
|
+
from ngcsimlib.component import Component
|
|
2
|
+
from jax import random
|
|
3
|
+
import time
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class COMPONENT_TEMPLATE(Component):
|
|
7
|
+
## Class Methods for Compartment Names
|
|
8
|
+
@classmethod
|
|
9
|
+
def DEFAULTCompartmentName(cls):
|
|
10
|
+
return 'DEFAULT'
|
|
11
|
+
|
|
12
|
+
## Bind Properties to Compartments for ease of use
|
|
13
|
+
@property
|
|
14
|
+
def DEFAULTCompartment(self):
|
|
15
|
+
return self.compartments.get(self.DEFAULTCompartmentName(), None)
|
|
16
|
+
|
|
17
|
+
@DEFAULTCompartment.setter
|
|
18
|
+
def DEFAULTCompartment(self, x):
|
|
19
|
+
if x is not None:
|
|
20
|
+
if True:
|
|
21
|
+
raise RuntimeError("")
|
|
22
|
+
self.compartments[self.DEFAULTCompartmentName()] = x
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
# Define Functions
|
|
26
|
+
def __init__(self, name, key=None, useVerboseDict=False, **kwargs):
|
|
27
|
+
super().__init__(name, useVerboseDict, **kwargs)
|
|
28
|
+
|
|
29
|
+
##Random Number Set up
|
|
30
|
+
self.key = key
|
|
31
|
+
if self.key is None:
|
|
32
|
+
self.key = random.PRNGKey(time.time_ns())
|
|
33
|
+
|
|
34
|
+
##Reset to initialize stuff
|
|
35
|
+
self.reset()
|
|
36
|
+
|
|
37
|
+
def verify_connections(self):
|
|
38
|
+
pass
|
|
39
|
+
|
|
40
|
+
def advance_state(self, **kwargs):
|
|
41
|
+
pass
|
|
42
|
+
|
|
43
|
+
def reset(self, **kwargs):
|
|
44
|
+
pass
|
|
45
|
+
|
|
46
|
+
def save(self, directory, **kwargs):
|
|
47
|
+
pass
|
|
@@ -0,0 +1,123 @@
|
|
|
1
|
+
from ngcsimlib.component import Component
|
|
2
|
+
from jax import numpy as jnp, random, jit
|
|
3
|
+
from functools import partial
|
|
4
|
+
import time
|
|
5
|
+
|
|
6
|
+
@jit
|
|
7
|
+
def update_times(t, s, tols):
|
|
8
|
+
"""
|
|
9
|
+
Updates time-of-last-spike (tols) variable.
|
|
10
|
+
|
|
11
|
+
Args:
|
|
12
|
+
t: current time (a scalar/int value)
|
|
13
|
+
|
|
14
|
+
s: binary spike vector
|
|
15
|
+
|
|
16
|
+
tols: current time-of-last-spike variable
|
|
17
|
+
|
|
18
|
+
Returns:
|
|
19
|
+
updated tols variable
|
|
20
|
+
"""
|
|
21
|
+
_tols = (1. - s) * tols + (s * t)
|
|
22
|
+
return _tols
|
|
23
|
+
|
|
24
|
+
@jit
|
|
25
|
+
def sample_bernoulli(dkey, data):
|
|
26
|
+
"""
|
|
27
|
+
Samples a Bernoulli spike train on-the-fly
|
|
28
|
+
|
|
29
|
+
Args:
|
|
30
|
+
data: sensory data (vector/matrix)
|
|
31
|
+
|
|
32
|
+
dt: integration time constant
|
|
33
|
+
|
|
34
|
+
Returns:
|
|
35
|
+
binary spikes
|
|
36
|
+
"""
|
|
37
|
+
s_t = random.bernoulli(dkey, p=data).astype(jnp.float32)
|
|
38
|
+
return s_t
|
|
39
|
+
|
|
40
|
+
class BernoulliCell(Component):
|
|
41
|
+
"""
|
|
42
|
+
A Bernoulli cell that produces Bernoulli-distributed spikes on-the-fly.
|
|
43
|
+
|
|
44
|
+
Args:
|
|
45
|
+
name: the string name of this cell
|
|
46
|
+
|
|
47
|
+
n_units: number of cellular entities (neural population size)
|
|
48
|
+
|
|
49
|
+
key: PRNG key to control determinism of any underlying synapses
|
|
50
|
+
associated with this cell
|
|
51
|
+
|
|
52
|
+
useVerboseDict: triggers slower, verbose dictionary mode (Default: False)
|
|
53
|
+
"""
|
|
54
|
+
|
|
55
|
+
## Class Methods for Compartment Names
|
|
56
|
+
@classmethod
|
|
57
|
+
def inputCompartmentName(cls):
|
|
58
|
+
return 'in'
|
|
59
|
+
|
|
60
|
+
@classmethod
|
|
61
|
+
def outputCompartmentName(cls):
|
|
62
|
+
return 'out'
|
|
63
|
+
|
|
64
|
+
@classmethod
|
|
65
|
+
def timeOfLastSpikeCompartmentName(cls):
|
|
66
|
+
return 'tols'
|
|
67
|
+
|
|
68
|
+
## Bind Properties to Compartments for ease of use
|
|
69
|
+
@property
|
|
70
|
+
def inputCompartment(self):
|
|
71
|
+
return self.compartments.get(self.inputCompartmentName(), None)
|
|
72
|
+
|
|
73
|
+
@inputCompartment.setter
|
|
74
|
+
def inputCompartment(self, inp):
|
|
75
|
+
self.compartments[self.inputCompartmentName()] = inp
|
|
76
|
+
|
|
77
|
+
@property
|
|
78
|
+
def outputCompartment(self):
|
|
79
|
+
return self.compartments.get(self.outputCompartmentName(), None)
|
|
80
|
+
|
|
81
|
+
@outputCompartment.setter
|
|
82
|
+
def outputCompartment(self, out):
|
|
83
|
+
self.compartments[self.outputCompartmentName()] = out
|
|
84
|
+
|
|
85
|
+
@property
|
|
86
|
+
def timeOfLastSpike(self):
|
|
87
|
+
return self.compartments.get(self.timeOfLastSpikeCompartmentName(), None)
|
|
88
|
+
|
|
89
|
+
@timeOfLastSpike.setter
|
|
90
|
+
def timeOfLastSpike(self, t):
|
|
91
|
+
self.compartments[self.timeOfLastSpikeCompartmentName()] = t
|
|
92
|
+
|
|
93
|
+
# Define Functions
|
|
94
|
+
def __init__(self, name, n_units, key=None, useVerboseDict=False, **kwargs):
|
|
95
|
+
super().__init__(name, useVerboseDict, **kwargs)
|
|
96
|
+
|
|
97
|
+
##Random Number Set up
|
|
98
|
+
self.key = key
|
|
99
|
+
if self.key is None:
|
|
100
|
+
self.key = random.PRNGKey(time.time_ns())
|
|
101
|
+
|
|
102
|
+
##Layer Size Setup
|
|
103
|
+
self.batch_size = 1
|
|
104
|
+
self.n_units = n_units
|
|
105
|
+
self.reset()
|
|
106
|
+
|
|
107
|
+
def verify_connections(self):
|
|
108
|
+
pass
|
|
109
|
+
|
|
110
|
+
def advance_state(self, t, dt, **kwargs):
|
|
111
|
+
self.key, *subkeys = random.split(self.key, 2)
|
|
112
|
+
|
|
113
|
+
self.outputCompartment = sample_bernoulli(subkeys[0], data=self.inputCompartment)
|
|
114
|
+
#self.timeOfLastSpike = (1 - self.outputCompartment) * self.timeOfLastSpike + (self.outputCompartment * t)
|
|
115
|
+
self.timeOfLastSpike = update_times(t, self.outputCompartment, self.timeOfLastSpike)
|
|
116
|
+
|
|
117
|
+
def reset(self, **kwargs):
|
|
118
|
+
self.inputCompartment = None
|
|
119
|
+
self.outputCompartment = jnp.zeros((self.batch_size, self.n_units)) #None
|
|
120
|
+
self.timeOfLastSpike = jnp.zeros((self.batch_size, self.n_units))
|
|
121
|
+
|
|
122
|
+
def save(self, **kwargs):
|
|
123
|
+
pass
|
|
@@ -0,0 +1,132 @@
|
|
|
1
|
+
from ngcsimlib.component import Component
|
|
2
|
+
from jax import numpy as jnp, random, jit
|
|
3
|
+
from functools import partial
|
|
4
|
+
import time
|
|
5
|
+
|
|
6
|
+
@jit
|
|
7
|
+
def update_times(t, s, tols):
|
|
8
|
+
"""
|
|
9
|
+
Updates time-of-last-spike (tols) variable.
|
|
10
|
+
|
|
11
|
+
Args:
|
|
12
|
+
t: current time (a scalar/int value)
|
|
13
|
+
|
|
14
|
+
s: binary spike vector
|
|
15
|
+
|
|
16
|
+
tols: current time-of-last-spike variable
|
|
17
|
+
|
|
18
|
+
Returns:
|
|
19
|
+
updated tols variable
|
|
20
|
+
"""
|
|
21
|
+
_tols = (1. - s) * tols + (s * t)
|
|
22
|
+
return _tols
|
|
23
|
+
|
|
24
|
+
@partial(jit, static_argnums=[3])
|
|
25
|
+
def sample_poisson(dkey, data, dt, fmax=63.75):
|
|
26
|
+
"""
|
|
27
|
+
Samples a Poisson spike train on-the-fly.
|
|
28
|
+
|
|
29
|
+
Args:
|
|
30
|
+
data: sensory data (vector/matrix)
|
|
31
|
+
|
|
32
|
+
dt: integration time constant
|
|
33
|
+
|
|
34
|
+
fmax: maximum frequency (Hz)
|
|
35
|
+
|
|
36
|
+
Returns:
|
|
37
|
+
binary spikes
|
|
38
|
+
"""
|
|
39
|
+
pspike = data * (dt/1000.) * fmax
|
|
40
|
+
eps = random.uniform(dkey, data.shape, minval=0., maxval=1., dtype=jnp.float32)
|
|
41
|
+
s_t = (eps < pspike).astype(jnp.float32)
|
|
42
|
+
return s_t
|
|
43
|
+
|
|
44
|
+
class PoissonCell(Component):
|
|
45
|
+
"""
|
|
46
|
+
A Poisson cell that produces approximately Poisson-distributed spikes on-the-fly.
|
|
47
|
+
|
|
48
|
+
Args:
|
|
49
|
+
name: the string name of this cell
|
|
50
|
+
|
|
51
|
+
n_units: number of cellular entities (neural population size)
|
|
52
|
+
|
|
53
|
+
max_freq: maximum frequency (in Hertz) of this Poisson spike train (must be > 0.)
|
|
54
|
+
|
|
55
|
+
key: PRNG key to control determinism of any underlying synapses
|
|
56
|
+
associated with this cell
|
|
57
|
+
|
|
58
|
+
useVerboseDict: triggers slower, verbose dictionary mode (Default: False)
|
|
59
|
+
"""
|
|
60
|
+
|
|
61
|
+
## Class Methods for Compartment Names
|
|
62
|
+
@classmethod
|
|
63
|
+
def inputCompartmentName(cls):
|
|
64
|
+
return 'in'
|
|
65
|
+
|
|
66
|
+
@classmethod
|
|
67
|
+
def outputCompartmentName(cls):
|
|
68
|
+
return 'out'
|
|
69
|
+
|
|
70
|
+
@classmethod
|
|
71
|
+
def timeOfLastSpikeCompartmentName(cls):
|
|
72
|
+
return 'tols'
|
|
73
|
+
|
|
74
|
+
## Bind Properties to Compartments for ease of use
|
|
75
|
+
@property
|
|
76
|
+
def inputCompartment(self):
|
|
77
|
+
return self.compartments.get(self.inputCompartmentName(), None)
|
|
78
|
+
|
|
79
|
+
@inputCompartment.setter
|
|
80
|
+
def inputCompartment(self, inp):
|
|
81
|
+
self.compartments[self.inputCompartmentName()] = inp
|
|
82
|
+
|
|
83
|
+
@property
|
|
84
|
+
def outputCompartment(self):
|
|
85
|
+
return self.compartments.get(self.outputCompartmentName(), None)
|
|
86
|
+
|
|
87
|
+
@outputCompartment.setter
|
|
88
|
+
def outputCompartment(self, out):
|
|
89
|
+
self.compartments[self.outputCompartmentName()] = out
|
|
90
|
+
|
|
91
|
+
@property
|
|
92
|
+
def timeOfLastSpike(self):
|
|
93
|
+
return self.compartments.get(self.timeOfLastSpikeCompartmentName(), None)
|
|
94
|
+
|
|
95
|
+
@timeOfLastSpike.setter
|
|
96
|
+
def timeOfLastSpike(self, t):
|
|
97
|
+
self.compartments[self.timeOfLastSpikeCompartmentName()] = t
|
|
98
|
+
|
|
99
|
+
# Define Functions
|
|
100
|
+
def __init__(self, name, n_units, max_freq=63.75, key=None,
|
|
101
|
+
useVerboseDict=False, **kwargs):
|
|
102
|
+
super().__init__(name, useVerboseDict, **kwargs)
|
|
103
|
+
|
|
104
|
+
##Random Number Set up
|
|
105
|
+
self.key = key
|
|
106
|
+
if self.key is None:
|
|
107
|
+
self.key = random.PRNGKey(time.time_ns())
|
|
108
|
+
|
|
109
|
+
## Poisson parameters
|
|
110
|
+
self.max_freq = max_freq ## maximum frequency (in Hertz/Hz)
|
|
111
|
+
|
|
112
|
+
##Layer Size Setup
|
|
113
|
+
self.batch_size = 1
|
|
114
|
+
self.n_units = n_units
|
|
115
|
+
self.reset()
|
|
116
|
+
|
|
117
|
+
def verify_connections(self):
|
|
118
|
+
pass
|
|
119
|
+
|
|
120
|
+
def advance_state(self, t, dt, **kwargs):
|
|
121
|
+
self.key, *subkeys = random.split(self.key, 2)
|
|
122
|
+
self.outputCompartment = sample_poisson(subkeys[0], data=self.inputCompartment, dt=dt, fmax=self.max_freq)
|
|
123
|
+
#self.timeOfLastSpike = (1 - self.outputCompartment) * self.timeOfLastSpike + (self.outputCompartment * t)
|
|
124
|
+
self.timeOfLastSpike = update_times(t, self.outputCompartment, self.timeOfLastSpike)
|
|
125
|
+
|
|
126
|
+
def reset(self, **kwargs):
|
|
127
|
+
self.inputCompartment = None
|
|
128
|
+
self.outputCompartment = jnp.zeros((self.batch_size, self.n_units)) #None
|
|
129
|
+
self.timeOfLastSpike = jnp.zeros((self.batch_size, self.n_units))
|
|
130
|
+
|
|
131
|
+
def save(self, **kwargs):
|
|
132
|
+
pass
|
|
@@ -0,0 +1,9 @@
|
|
|
1
|
+
## point to rate-coded cell componet types
|
|
2
|
+
from .graded.rateCell import RateCell
|
|
3
|
+
from .graded.gaussianErrorCell import GaussianErrorCell
|
|
4
|
+
from .graded.laplacianErrorCell import LaplacianErrorCell
|
|
5
|
+
## point to standard spiking cell component types
|
|
6
|
+
from .spiking.sLIFCell import SLIFCell
|
|
7
|
+
from .spiking.LIFCell import LIFCell
|
|
8
|
+
from .spiking.quadLIFCell import QuadLIFCell
|
|
9
|
+
from .spiking.izhikevichCell import IzhikevichCell
|
|
@@ -0,0 +1,174 @@
|
|
|
1
|
+
from ngcsimlib.component import Component
|
|
2
|
+
from jax import numpy as jnp, random, jit
|
|
3
|
+
from functools import partial
|
|
4
|
+
import time, sys
|
|
5
|
+
|
|
6
|
+
#@partial(jit, static_argnums=[3])
|
|
7
|
+
def run_cell(dt, targ, mu, eType="gaussian"):
|
|
8
|
+
"""
|
|
9
|
+
Moves cell dynamics one step forward.
|
|
10
|
+
|
|
11
|
+
Args:
|
|
12
|
+
dt: integration time constant
|
|
13
|
+
|
|
14
|
+
targ: target pattern value
|
|
15
|
+
|
|
16
|
+
mu: prediction value
|
|
17
|
+
|
|
18
|
+
Returns:
|
|
19
|
+
derivative w.r.t. mean "dmu", derivative w.r.t. target dtarg
|
|
20
|
+
"""
|
|
21
|
+
return run_gaussian_cell(dt, targ, mu)
|
|
22
|
+
|
|
23
|
+
@jit
|
|
24
|
+
def run_gaussian_cell(dt, targ, mu):
|
|
25
|
+
"""
|
|
26
|
+
Moves Gaussian cell dynamics one step forward. Specifically, this
|
|
27
|
+
routine emulates the error unit behavior of the local cost functional:
|
|
28
|
+
|
|
29
|
+
| L(targ, mu) = (1/2) * ||targ - mu||^2_2
|
|
30
|
+
| or log likelihood of the multivariate Gaussian with identity covariance
|
|
31
|
+
|
|
32
|
+
Args:
|
|
33
|
+
dt: integration time constant
|
|
34
|
+
|
|
35
|
+
targ: target pattern value
|
|
36
|
+
|
|
37
|
+
mu: prediction value
|
|
38
|
+
|
|
39
|
+
Returns:
|
|
40
|
+
derivative w.r.t. mean "dmu", derivative w.r.t. target dtarg
|
|
41
|
+
"""
|
|
42
|
+
dmu = (targ - mu) # e (error unit)
|
|
43
|
+
dtarg = -dmu # reverse of e
|
|
44
|
+
return dmu, dtarg
|
|
45
|
+
|
|
46
|
+
class GaussianErrorCell(Component): ## Rate-coded/real-valued error unit/cell
|
|
47
|
+
"""
|
|
48
|
+
A simple (non-spiking) Gaussian error cell - this is a fixed-point solution
|
|
49
|
+
of a mismatch signal.
|
|
50
|
+
|
|
51
|
+
Args:
|
|
52
|
+
name: the string name of this cell
|
|
53
|
+
|
|
54
|
+
n_units: number of cellular entities (neural population size)
|
|
55
|
+
|
|
56
|
+
tau_m: (Unused -- currently cell is a fixed-point model)
|
|
57
|
+
|
|
58
|
+
leakRate: (Unused -- currently cell is a fixed-point model)
|
|
59
|
+
|
|
60
|
+
key: PRNG Key to control determinism of any underlying synapses
|
|
61
|
+
associated with this cell
|
|
62
|
+
|
|
63
|
+
useVerboseDict: triggers slower, verbose dictionary mode (Default: False)
|
|
64
|
+
"""
|
|
65
|
+
|
|
66
|
+
## Class Methods for Compartment Names
|
|
67
|
+
@classmethod
|
|
68
|
+
def inputCompartmentName(cls):
|
|
69
|
+
return 'j' ## electrical current
|
|
70
|
+
|
|
71
|
+
@classmethod
|
|
72
|
+
def outputCompartmentName(cls):
|
|
73
|
+
return 'e' ## rate-coded output
|
|
74
|
+
|
|
75
|
+
@classmethod
|
|
76
|
+
def meanName(cls):
|
|
77
|
+
return 'mu'
|
|
78
|
+
|
|
79
|
+
@classmethod
|
|
80
|
+
def derivMeanName(cls):
|
|
81
|
+
return 'dmu'
|
|
82
|
+
|
|
83
|
+
@classmethod
|
|
84
|
+
def targetName(cls):
|
|
85
|
+
return 'target'
|
|
86
|
+
|
|
87
|
+
@classmethod
|
|
88
|
+
def derivTargetName(cls):
|
|
89
|
+
return 'dtarget'
|
|
90
|
+
|
|
91
|
+
@classmethod
|
|
92
|
+
def modulatorName(cls):
|
|
93
|
+
return 'modulator'
|
|
94
|
+
|
|
95
|
+
## Bind Properties to Compartments for ease of use
|
|
96
|
+
@property
|
|
97
|
+
def mean(self):
|
|
98
|
+
return self.compartments.get(self.meanName(), None)
|
|
99
|
+
|
|
100
|
+
@mean.setter
|
|
101
|
+
def mean(self, inp):
|
|
102
|
+
self.compartments[self.meanName()] = inp
|
|
103
|
+
|
|
104
|
+
@property
|
|
105
|
+
def derivMean(self):
|
|
106
|
+
return self.compartments.get(self.derivMeanName(), None)
|
|
107
|
+
|
|
108
|
+
@derivMean.setter
|
|
109
|
+
def derivMean(self, inp):
|
|
110
|
+
self.compartments[self.derivMeanName()] = inp
|
|
111
|
+
|
|
112
|
+
@property
|
|
113
|
+
def target(self):
|
|
114
|
+
return self.compartments.get(self.targetName(), None)
|
|
115
|
+
|
|
116
|
+
@target.setter
|
|
117
|
+
def target(self, inp):
|
|
118
|
+
self.compartments[self.targetName()] = inp
|
|
119
|
+
|
|
120
|
+
@property
|
|
121
|
+
def derivTarget(self):
|
|
122
|
+
return self.compartments.get(self.derivTargetName(), None)
|
|
123
|
+
|
|
124
|
+
@derivTarget.setter
|
|
125
|
+
def derivTarget(self, inp):
|
|
126
|
+
self.compartments[self.derivTargetName()] = inp
|
|
127
|
+
|
|
128
|
+
@property
|
|
129
|
+
def modulator(self):
|
|
130
|
+
return self.compartments.get(self.modulatorName(), None)
|
|
131
|
+
|
|
132
|
+
@modulator.setter
|
|
133
|
+
def modulator(self, inp):
|
|
134
|
+
self.compartments[self.modulatorName()] = inp
|
|
135
|
+
|
|
136
|
+
# Define Functions
|
|
137
|
+
def __init__(self, name, n_units, tau_m=0., leakRate=0., key=None,
|
|
138
|
+
useVerboseDict=False, **kwargs):
|
|
139
|
+
super().__init__(name, useVerboseDict, **kwargs)
|
|
140
|
+
|
|
141
|
+
##Random Number Set up
|
|
142
|
+
self.key = key
|
|
143
|
+
if self.key is None:
|
|
144
|
+
self.key = random.PRNGKey(time.time_ns())
|
|
145
|
+
|
|
146
|
+
##Layer Size Setup
|
|
147
|
+
self.n_units = n_units
|
|
148
|
+
self.batch_size = 1
|
|
149
|
+
|
|
150
|
+
## Set up bundle for multiple inputs of current
|
|
151
|
+
self.reset()
|
|
152
|
+
|
|
153
|
+
def verify_connections(self):
|
|
154
|
+
#self.metadata.check_incoming_connections(self.inputCompartmentName(), min_connections=1)
|
|
155
|
+
self.metadata.check_incoming_connections(self.meanName(), min_connections=1)
|
|
156
|
+
self.metadata.check_incoming_connections(self.targetName(), min_connections=1)
|
|
157
|
+
|
|
158
|
+
def advance_state(self, t, dt, **kwargs):
|
|
159
|
+
## currently only Gaussian error cells supported
|
|
160
|
+
self.derivMean, self.derivTarget = run_cell(dt, self.target, self.mean)
|
|
161
|
+
if self.modulator is not None:
|
|
162
|
+
self.derivMean = self.derivMean * self.modulator
|
|
163
|
+
self.derivTarget = self.derivTarget * self.modulator
|
|
164
|
+
self.modulator = None ## use and consume modulator
|
|
165
|
+
|
|
166
|
+
def reset(self, **kwargs):
|
|
167
|
+
self.derivMean = jnp.zeros((self.batch_size, self.n_units))
|
|
168
|
+
self.derivTarget = jnp.zeros((self.batch_size, self.n_units))
|
|
169
|
+
self.target = jnp.zeros((self.batch_size, self.n_units)) #None
|
|
170
|
+
self.mean = jnp.zeros((self.batch_size, self.n_units)) #None
|
|
171
|
+
self.modulator = None
|
|
172
|
+
|
|
173
|
+
def save(self, **kwargs):
|
|
174
|
+
pass
|