Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions AUTHORS
Original file line number Diff line number Diff line change
Expand Up @@ -16,3 +16,5 @@ Maxbeth2 (Ohas)
pagrawal-psu
pulinagrawal
antonvice
Jack Foreback

25 changes: 23 additions & 2 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,27 @@ Matplotlib (>=3.8.0) and imageio (>=2.31.5) and both plotting and density estima
tools (routines within ``ngclearn.utils.density``) will require Scikit-learn (>=0.24.2).
Many of the tutorials will require Matplotlib (>=3.8.0), imageio (>=2.31.5), and Scikit-learn (>=0.24.2).

<i>Note</i>: If you are working with Cuda 12 and want to use jax/jaxlib versions > 0.4.28, you might need to
check that you are working with the right version of Cudnn (e.g., `nvidia-cudnn-cu12==9.10.2.21`) to ensure
that all of ngc-learn's internal supported tools, like in-built convolution/deconvolution, compile
correctly onto the GPU (if using an architecture based on Pascal GPUs, i.e., Compute Capability 6.1,
combined with NVIDIA Driver 580+).

**Important Note for Legacy GPU Users (Pascal Architecture)**
> If you are running JAX (`> 0.4.28`) on **CUDA 12** using an older
> **Pascal-generation GPU** (Compute Capability 6.1, e.g., GTX 1080/1080Ti, Titan X)
> combined with **NVIDIA Driver 580+**, you might encounter compilation crashes during
> convolution/deconvolution operations (such as `unknown cudnn status: 5003`).
>
> Newer versions of `nvidia-cudnn-cu12` have dropped critical hardware support for
> these legacy architectures. To fix this and ensure `ngclearn` compiles correctly
> on your GPU, you will need to explicitly "pin" your cuDNN library version using
> this command (after installing Cuda-12 JAX):
>
> ```bash
> pip install --force-reinstall "nvidia-cudnn-cu12==9.10.2.21"
> ```

### User Installation

<i>Setup</i>: The easiest way to install ngc-learn is through <code>pip</code>:
Expand All @@ -68,7 +89,7 @@ and complete the following sequence of steps as depicted in the screenshot below
right major and minor version of ngc-learn):

```console
Python 3.11.4 (main, MONTH DAY YEAR, TIME) [GCC XX.X.X] on linux
Python 3.12.13 (main, MONTH DAY YEAR, TIME) [GCC XX.X.X] on linux
Type "help", "copyright", "credits" or "license" for more information.
>>> import ngclearn
>>> ngclearn.__version__
Expand Down Expand Up @@ -119,7 +140,7 @@ $ python install -e .
</pre>

**Version:**<br>
3.2.1 <!--1.2.3-Beta--> <!-- -Alpha -->
3.2.2

Author:
Alexander G. Ororbia II<br>
Expand Down
18 changes: 5 additions & 13 deletions docs/installation.md
Original file line number Diff line number Diff line change
Expand Up @@ -5,10 +5,10 @@
<i>Setup:</i> <a href="https://github.com/NACLab/ngc-learn">NGC-Learn</a>, in its entirety (including its supporting utility sub-packages), requires that you ensure that you have installed the following base dependencies in your system. Note that this library was developed and tested on Ubuntu 22.04 (with much earlier versions on Ubuntu 18.04/20.04).
Specifically, NGC-Learn requires:
* Python (>=3.10)
* ngcsimlib (>=3.0.0), (<a href="https://github.com/NACLab/ngc-sim-lib">official page</a>)
* ngcsimlib (>=3.1.1), (<a href="https://github.com/NACLab/ngc-sim-lib">official page</a>)
* NumPy (>=1.22.0)
* SciPy (>=1.7.0)
* JAX (>= 0.4.28; and jaxlib>=0.4.28) <!--(tested for cuda 11.8)-->
* JAX (>= 0.11.1; and jaxlib>=0.11.1) <!--(tested for cuda 12)-->
* Matplotlib (>=3.8.0), (for `ngclearn.utils.viz`)
* Scikit-learn (>=1.6.1), (for `ngclearn.utils.patch_utils` and `ngclearn.utils.density`)

Expand All @@ -33,7 +33,7 @@ $ git clone https://github.com/NACLab/ngc-learn.git
$ cd ngc-learn
```

3. (<i>Optional</i>; only for GPU version) Install JAX for either CUDA 12 , depending on your system setup. Follow the <a href="https://jax.readthedocs.io/en/latest/installation.html">installation instructions</a> on the official JAX page to properly install the CUDA 11 or 12 version.
3. (<i>Optional</i>; only for GPU version) Install JAX for either CUDA 12 or 13, depending on your system setup. Follow the <a href="https://jax.readthedocs.io/en/latest/installation.html">installation instructions</a> on the official JAX page to properly install the CUDA 12 or 13 version.

4. Install the NGC-Learn package via:
```console
Expand All @@ -47,18 +47,10 @@ $ pip install -e .
If the installation was successful, you should see the following if you test it against your Python interpreter, i.e., run the <code>$ python</code> command and complete the following sequence of steps as depicted in the screenshot below:<br>

```console
Python 3.11.4 (main, MONTH DAY YEAR, TIME) [GCC XX.X.X] on linux
Python 3.12.13 (main, MONTH DAY YEAR, TIME) [GCC XX.X.X] on linux
Type "help", "copyright", "credits" or "license" for more information.
>>> import ngclearn
>>> ngclearn.__version__
'3.0.1'
'3.2.2'
```

<!--
<i>Note</i>: If you do not have a JSON configuration file in place (see tutorials
for details) locally where you call the import to ngc-learn, a warning will pop
up containing within it "<i>UserWarning: Missing file to preload modules from.</i>";
this still means that ngc-learn installed successfully but you will need to
point to a JSON configuration when building projects with ngc-learn.
-->

21 changes: 10 additions & 11 deletions docs/requirements.txt
Original file line number Diff line number Diff line change
@@ -1,11 +1,10 @@
sphinx>=4.5.0
sphinx_rtd_theme>=0.5.2
myst-parser>=0.17.2
numpy>=1.26.0
scikit-learn>=0.24.2
scipy>=1.7.0
matplotlib>=3.8.0
jax>=0.4.28
jaxlib>=0.4.28
imageio>=2.31.5
ngcsimlib>=1.0.1
numpy>=2.5.2
scikit-learn>=1.9.0
scipy>=1.18.1
matplotlib>=3.11.1
jax>=0.11.1
jaxlib>=0.11.1
ngcsimlib>=3.1.1
imageio>=2.37.4
pandas>=3.0.5
typing_extensions>=4.15.0
8 changes: 8 additions & 0 deletions docs/source/ngclearn.components.synapses.hebbian.rst
Original file line number Diff line number Diff line change
Expand Up @@ -60,6 +60,14 @@ ngclearn.components.synapses.hebbian.inhibitorySTDPSynapse module
:undoc-members:
:show-inheritance:

ngclearn.components.synapses.hebbian.ojaTensorSynapse module
------------------------------------------------------------

.. automodule:: ngclearn.components.synapses.hebbian.ojaTensorSynapse
:members:
:undoc-members:
:show-inheritance:

ngclearn.components.synapses.hebbian.traceSTDPSynapse module
------------------------------------------------------------

Expand Down
7 changes: 7 additions & 0 deletions history.txt
Original file line number Diff line number Diff line change
Expand Up @@ -115,3 +115,10 @@ History
* integration of additional visualization tools
* integration of sparse-tensor synaptic cable (locally-connected/unshared-convolutional structure)
* additional component integration/revisions, including updates to patched-synaptic cable components

3.2.2
— — — — — — — — -
* upgrades to utils/effective dimension toolset
* minor patches
* adjustments to requirements to nudge to modern >=Python 3.12 and >=Jax 0.11.1 (for cuda12)

1 change: 0 additions & 1 deletion ngclearn/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,6 @@
"with python 3.8 is maintained to allow for lava-nc components and should only be used with those")

## Following obtains installed package names (as normalized keys) for ngc-learn
#required = {'ngcsimlib', 'jax', 'jaxlib'} ## list of core ngclearn dependencies
required = {'ngcsimlib'} #, 'jax', 'jaxlib'}
#installed = {pkg.key for pkg in pkg_resources.working_set}
#missing = required - installed
Expand Down
72 changes: 64 additions & 8 deletions ngclearn/utils/analysis/effective_dim.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,13 +2,38 @@
import jax
from jax import numpy as jnp, jit

'''
Some useful notes on effective dimensional analysis:

* Participation ratio (PR), which measures the general usage of a vector space (how many
features are being used), can be easily "fooled" by a "bully" dimension, specifically yielding
cases where, say all D dimensions are all active but one of them holds 99% of variance while
the other D-1 dims share the remaining 1%; in this case, PR would yield a rather high, seemingly
healthy-looking score yet it is not accounting for the fact that other low-variance dims are
participating yet are far too "quiet"

* Stable rank (SR; which is a function of the Rayleigh coefficient) is good with detecting
if a single feature/dimension is
completely drowning out the rest of the vector space - if stable rank goes close to 1, then
model has collapsed to a 1-dim case even though the PR is high; this metric is useful to
examine to check if a vector space is multi-dimensional and balanced (and not just a single
massive eigenvector surrounded by insignificant/low-contributing dimensions)

PR, SR, and Rankme are metrics along a spectral analysis metric spectrum:
* Rankme is the exponential Shannon entropy of spectrum,
* PR is the Renyi-2 "effective dimension", and,
* SR focuses on the single largest eigenvalue of the dimensional space
'''

@partial(jit, static_argnums=[1])
def participation_ratio(
latent_codes, use_NaN_fallback=False
):
"""
Calculates the participation ratio coefficient (also known as the Gini effective
dimension) for a set of latent codes.
Calculates the participation ratio (PR) coefficient (also known as the Gini effective
dimension) for a set of latent codes. PR is useful for detecting "total dimensional
collapse", where the data/vector-space essentially flattens into a line (or only
make use of too few or even just a single dimension of the space).

Args:
latent_codes: a set of (N x D) latent code vectors (one row per vector code)
Expand Down Expand Up @@ -36,6 +61,35 @@ def participation_ratio(
##else, use ML-oriented NaN return value fallback
return tr2_cov / cov2_tr if cov2_tr > 0 else float("nan")

@jit
def covariance_error(latent_codes):
"""
Calculates the off-diagonal covariance error of a set of latent codes. This dimensional metric is useful for
quantifying informational redundancy. If the error/score is high, units/dimensions are highly correlated, which
means the vector code space is wasting its dimensional capacity by having different dimensions/features model
the exact same piece of information.

Args:
latent_codes: a set of (N x D) latent code vectors (one row per vector code)

Returns:
scalar measurement of the off-diagonal covariance error
"""
Z = latent_codes
Zc = Z - Z.mean(axis=0, keepdims=True)
cov = (Zc.T @ Zc) / (Zc.shape[0] - 1)
## normalize covariance to get correlation matrix
d = jnp.diag(cov)
std_dev = jnp.sqrt(jnp.clip(d, a_min=1e-8))
denominator = std_dev[:, None] * std_dev[None, :]
corr = cov / jnp.clip(denominator, a_min=1e-6)
## zero out diagonal elements
diag_mask = jnp.eye(corr.shape[0])
off_diag = corr * (1.0 - diag_mask)
## calc mean squared off-diagonal error
off_diagonal_error = jnp.sum(off_diag ** 2) / (corr.shape[0] * (corr.shape[0] - 1))
return off_diagonal_error

@partial(jit, static_argnums=[1])
def rankme(latent_codes, eps=1e-7):
"""
Expand Down Expand Up @@ -71,8 +125,11 @@ def rankme(latent_codes, eps=1e-7):
@partial(jit, static_argnums=[1])
def stable_rank(latent_codes, num_iters=10): ## power-iterator method
"""
Computes the stable rank via the power iteration method in order to find the
top singular value.
Computes the "stable rank} via the power iteration method in order to find the
top singular value (this metric is a function of the Rayleigh coefficient). Note that
this metric is useful for detecting a case of dimensional collapse known as "dominant
component collapse", where a single feature "hogs" up all
the power of the representational vector space while ignoring everything else.

Args:
latent_codes: a set of (N x D) latent code vectors (one row per vector code)
Expand All @@ -91,14 +148,13 @@ def stable_rank(latent_codes, num_iters=10): ## power-iterator method
key = jax.random.PRNGKey(0)
v = jax.random.normal(key, (Zc.shape[1], 1))
v = v / jnp.linalg.norm(v)
## apply standard power iteration loop
## run power iteration loop
for _ in range(num_iters):
## v = (Zc.T @ (Zc @ v))
v = Zc.T @ (Zc @ v)
v = v / jnp.linalg.norm(v)
## compute largest singular value squared (i.e., the Rayleigh quotient):
### sigma_max^2 = ||Zc @ v||^2
sigma_max_sq = jnp.sum(jnp.square(Zc @ v))
## compute largest singular value squared => sigma_max^2 = ||Zc @ v||^2
sigma_max_sq = jnp.sum(jnp.square(Zc @ v)) ## Rayleigh coefficient/quotient
return jnp.where(sigma_max_sq > 0.0, frobenius_norm_sq / sigma_max_sq, 1.0) # stable-rank score


13 changes: 9 additions & 4 deletions ngclearn/utils/viz/synapse_plot.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,6 @@
import imageio.v3 as iio
import jax.numpy as jnp


def visualize_macro_grid( ## more complex filter visualization co-routine
thetas,
sizes,
Expand Down Expand Up @@ -101,7 +100,8 @@ def visualize(
sizes,
prefix,
order=None,
suffix='.jpg'
suffix='.jpg',
contrast_by_data=False
):
"""

Expand All @@ -113,6 +113,8 @@ def visualize(
prefix:

suffix:

contrast_by_data:
"""
if order is None:
order = ['C' for _ in range(len(thetas))]
Expand Down Expand Up @@ -143,8 +145,11 @@ def visualize(
point = start + 1 + i + (r * extra)
plt.subplot(n_rows_total, n_cols_total, point)
_filter = T[i, :]
max_val = float(jnp.max(jnp.abs(_filter)))
min_val = float(jnp.min(jnp.abs(_filter)))
max_val = None # 1.
min_val = None # -1.
if contrast_by_data:
max_val = float(jnp.max(jnp.abs(_filter)))
min_val = float(jnp.min(jnp.abs(_filter)))
plt.imshow(
np.reshape(_filter, (sizes[idx][0], sizes[idx][1]), order=order[idx]),
cmap=plt.cm.bone, interpolation='nearest', vmin=min_val, vmax=max_val
Expand Down
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@ build-backend = "setuptools.build_meta" # using setuptool building engine

[project]
name = "ngclearn"
version = "3.2.1"
version = "3.2.2"
description = "Simulation software for building and analyzing computational neuroscience models, brain-inspired computing systems, and NeuroAI agents."
authors = [
{name = "Alexander Ororbia", email = "ago@cs.rit.edu"},
Expand Down
16 changes: 8 additions & 8 deletions requirements.txt
Original file line number Diff line number Diff line change
@@ -1,10 +1,10 @@
numpy>=1.26.4
scikit-learn>=1.6.1
scipy>=1.14.1
matplotlib>=3.9.4
jax>=0.4.28
jaxlib>=0.4.28
numpy>=2.5.2
scikit-learn>=1.9.0
scipy>=1.18.1
matplotlib>=3.11.1
jax>=0.11.1
jaxlib>=0.11.1
ngcsimlib>=3.1.1
imageio>=2.37.0
pandas>=2.2.3
imageio>=2.37.4
pandas>=3.0.5
typing_extensions>=4.15.0
Original file line number Diff line number Diff line change
Expand Up @@ -31,16 +31,17 @@ def test_HebbianConvSynapse1():
stride=stride, padding=padding_style, batch_size=batch_size, key=subkeys[0]
)

evolve_process = (MethodProcess("evolve_process")
use_jit = True #False
evolve_process = (MethodProcess("evolve_process", use_jit=use_jit)
>> a.evolve)

backtransmit_process = (MethodProcess("backtransmit_process")
backtransmit_process = (MethodProcess("backtransmit_process", use_jit=use_jit)
>> a.backtransmit)

advance_process = (MethodProcess("advance_proc")
advance_process = (MethodProcess("advance_proc", use_jit=use_jit)
>> a.advance_state)

reset_process = (MethodProcess("reset_proc")
reset_process = (MethodProcess("reset_proc", use_jit=use_jit)
>> a.reset)

x = jnp.ones(x_shape)
Expand Down
Loading