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.
Files changed (46) hide show
  1. ngclearn/__init__.py +23 -0
  2. ngclearn/commands/__init__.py +1 -0
  3. ngclearn/components/__init__.py +19 -0
  4. ngclearn/components/baseComponentTemplate.py +47 -0
  5. ngclearn/components/input_encoders/__init__.py +2 -0
  6. ngclearn/components/input_encoders/bernoulliCell.py +123 -0
  7. ngclearn/components/input_encoders/poissonCell.py +132 -0
  8. ngclearn/components/neurons/__init__.py +9 -0
  9. ngclearn/components/neurons/graded/__init__.py +4 -0
  10. ngclearn/components/neurons/graded/gaussianErrorCell.py +174 -0
  11. ngclearn/components/neurons/graded/laplacianErrorCell.py +174 -0
  12. ngclearn/components/neurons/graded/rateCell.py +201 -0
  13. ngclearn/components/neurons/spiking/LIFCell.py +305 -0
  14. ngclearn/components/neurons/spiking/__init__.py +5 -0
  15. ngclearn/components/neurons/spiking/izhikevichCell.py +150 -0
  16. ngclearn/components/neurons/spiking/quadLIFCell.py +327 -0
  17. ngclearn/components/neurons/spiking/sLIFCell.py +378 -0
  18. ngclearn/components/other/__init__.py +2 -0
  19. ngclearn/components/other/expKernel.py +129 -0
  20. ngclearn/components/other/varTrace.py +150 -0
  21. ngclearn/components/synapses/__init__.py +3 -0
  22. ngclearn/components/synapses/hebbian/__init__.py +3 -0
  23. ngclearn/components/synapses/hebbian/expSTDPSynapse.py +234 -0
  24. ngclearn/components/synapses/hebbian/hebbianSynapse.py +334 -0
  25. ngclearn/components/synapses/hebbian/traceSTDPSynapse.py +264 -0
  26. ngclearn/components/wrappers.py +8 -0
  27. ngclearn/utils/__init__.py +0 -0
  28. ngclearn/utils/density/__init__.py +0 -0
  29. ngclearn/utils/density/gmm.py +82 -0
  30. ngclearn/utils/io_utils.py +67 -0
  31. ngclearn/utils/model_utils.py +274 -0
  32. ngclearn/utils/optim/__init__.py +2 -0
  33. ngclearn/utils/optim/adam.py +88 -0
  34. ngclearn/utils/optim/opt.py +25 -0
  35. ngclearn/utils/optim/sgd.py +38 -0
  36. ngclearn/utils/patch_utils.py +79 -0
  37. ngclearn/utils/viz/__init__.py +0 -0
  38. ngclearn/utils/viz/dim_reduce.py +91 -0
  39. ngclearn/utils/viz/raster.py +177 -0
  40. ngclearn/utils/viz/synapse_plot.py +151 -0
  41. ngclearn-1.0b0.dist-info/AUTHORS +15 -0
  42. ngclearn-1.0b0.dist-info/LICENSE +29 -0
  43. ngclearn-1.0b0.dist-info/METADATA +188 -0
  44. ngclearn-1.0b0.dist-info/RECORD +46 -0
  45. ngclearn-1.0b0.dist-info/WHEEL +5 -0
  46. 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,2 @@
1
+ from .bernoulliCell import BernoulliCell
2
+ from .poissonCell import PoissonCell
@@ -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,4 @@
1
+ ## point to rate-coded cell componet types
2
+ from .rateCell import RateCell
3
+ from .gaussianErrorCell import GaussianErrorCell
4
+ from .laplacianErrorCell import LaplacianErrorCell
@@ -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