diff --git a/AUTHORS b/AUTHORS
index 8523ca5e..4ce215fd 100644
--- a/AUTHORS
+++ b/AUTHORS
@@ -16,3 +16,5 @@ Maxbeth2 (Ohas)
pagrawal-psu
pulinagrawal
antonvice
+Jack Foreback
+
diff --git a/README.md b/README.md
index f7312883..16104817 100644
--- a/README.md
+++ b/README.md
@@ -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).
+Note: 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
Setup: The easiest way to install ngc-learn is through pip:
@@ -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__
@@ -119,7 +140,7 @@ $ python install -e .
**Version:**
-3.2.1
+3.2.2
Author:
Alexander G. Ororbia II
diff --git a/docs/installation.md b/docs/installation.md
index 963490f8..99517813 100644
--- a/docs/installation.md
+++ b/docs/installation.md
@@ -5,10 +5,10 @@
Setup: NGC-Learn, 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), (official page)
+* ngcsimlib (>=3.1.1), (official page)
* NumPy (>=1.22.0)
* SciPy (>=1.7.0)
-* JAX (>= 0.4.28; and jaxlib>=0.4.28)
+* JAX (>= 0.11.1; and jaxlib>=0.11.1)
* Matplotlib (>=3.8.0), (for `ngclearn.utils.viz`)
* Scikit-learn (>=1.6.1), (for `ngclearn.utils.patch_utils` and `ngclearn.utils.density`)
@@ -33,7 +33,7 @@ $ git clone https://github.com/NACLab/ngc-learn.git
$ cd ngc-learn
```
-3. (Optional; only for GPU version) Install JAX for either CUDA 12 , depending on your system setup. Follow the installation instructions on the official JAX page to properly install the CUDA 11 or 12 version.
+3. (Optional; only for GPU version) Install JAX for either CUDA 12 or 13, depending on your system setup. Follow the installation instructions on the official JAX page to properly install the CUDA 12 or 13 version.
4. Install the NGC-Learn package via:
```console
@@ -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 $ python command and complete the following sequence of steps as depicted in the screenshot below:
```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'
```
-
-
diff --git a/docs/requirements.txt b/docs/requirements.txt
index 0ebb3dc3..a61e3100 100644
--- a/docs/requirements.txt
+++ b/docs/requirements.txt
@@ -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
diff --git a/docs/source/ngclearn.components.synapses.hebbian.rst b/docs/source/ngclearn.components.synapses.hebbian.rst
index d42770a6..a972e440 100644
--- a/docs/source/ngclearn.components.synapses.hebbian.rst
+++ b/docs/source/ngclearn.components.synapses.hebbian.rst
@@ -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
------------------------------------------------------------
diff --git a/history.txt b/history.txt
index 13364b41..4f1897e4 100644
--- a/history.txt
+++ b/history.txt
@@ -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)
+
diff --git a/ngclearn/__init__.py b/ngclearn/__init__.py
index 5670cfca..e21f48ae 100644
--- a/ngclearn/__init__.py
+++ b/ngclearn/__init__.py
@@ -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
diff --git a/ngclearn/utils/analysis/effective_dim.py b/ngclearn/utils/analysis/effective_dim.py
index c427c47e..252fadbc 100644
--- a/ngclearn/utils/analysis/effective_dim.py
+++ b/ngclearn/utils/analysis/effective_dim.py
@@ -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)
@@ -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):
"""
@@ -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)
@@ -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
diff --git a/ngclearn/utils/viz/synapse_plot.py b/ngclearn/utils/viz/synapse_plot.py
index b60bc71c..20307914 100644
--- a/ngclearn/utils/viz/synapse_plot.py
+++ b/ngclearn/utils/viz/synapse_plot.py
@@ -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,
@@ -101,7 +100,8 @@ def visualize(
sizes,
prefix,
order=None,
- suffix='.jpg'
+ suffix='.jpg',
+ contrast_by_data=False
):
"""
@@ -113,6 +113,8 @@ def visualize(
prefix:
suffix:
+
+ contrast_by_data:
"""
if order is None:
order = ['C' for _ in range(len(thetas))]
@@ -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
diff --git a/pyproject.toml b/pyproject.toml
index ecd1b495..5cb289b2 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -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"},
diff --git a/requirements.txt b/requirements.txt
index ede6cb3e..a61e3100 100644
--- a/requirements.txt
+++ b/requirements.txt
@@ -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
diff --git a/tests/components/synapses/convolution/test_hebbianConvSynapse.py b/tests/components/synapses/convolution/test_hebbianConvSynapse.py
index e5ee3d74..7ae71c95 100644
--- a/tests/components/synapses/convolution/test_hebbianConvSynapse.py
+++ b/tests/components/synapses/convolution/test_hebbianConvSynapse.py
@@ -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)