Skip to main content

NN-Library

Test coverage Docstring coverage

We in the BONSAI Lab do research on neural networks, among other things, that requires loading/training/reconfiguring neural network models. This library is a work-in-progress suite of in-house tools to address some pain-points we've encountered in our research workflow.

We make no guarantees about the stability or usability of this library, but we hope that it can be useful to others in the research community. If you have any questions or suggestions, please feel free to reach out to us or open an issue on the GitHub repository.

Installation

Using pip:

pip install bonsai-nn-library[cpu]  # specifies torch-on-cpu
pip install bonsai-nn-library[cu130]  # specifies torch-on-cuda-13.0

Using uv (in a project):

uv add bonsai-nn-library --extra cpu  # for cpu
uv add bonsai-nn-library --extra cu130  # for cuda 13.0

CUDA versioning note: configuring CUDA dependencies for pytorch is notoriously tricky. More generally, pip and uv were not designed to handle switching between different dependencies on host systems with different hardware or driver support. At time of writing this, there are some open PyTorch issues and open PEPs that will someday make life easier. For now, our pyproject.toml file supports cpu and cu130. For other If you have a different CUDA version, you could update our pyproject.toml. See the disclaimer above: we don't guarantee our library will work for everyone on all systems, but if you have a way to improve it we welcome contributions.

Usage

The top-level import is nn_lib. Say you want to use one of our "fancy layers" like a low-rank convolution. You can do so like this:

from torch import nn
from nn_lib.models.fancy_layers import LowRankConv2d

model = nn.Sequential(
    LowRankConv2d(in_channels=3, out_channels=64, kernel_size=3, rank=8),
    nn.ReLU(),
    nn.Conv2d(in_channels=64, out_channels=64, kernel_size=3),
    nn.ReLU(),
    nn.Flatten(),
    nn.LazyLinear(10)
)

Useful thing #1: improved GraphModules.

PyTorch was not originally designed to handle explicit computation graphs, but it was added somewhat later in the torch.fx module. Others might use tensorflow or jax for this, but we like PyTorch. The torch.fx.GraphModule class is the built-in way to handle computation graphs in PyTorch, but it lacks some features that we find useful. We have extended the GraphModule class in our GraphModulePlus class, which inherits from GraphModule and adds some further functionality.

A motivating use-case is that we want to be able to "stitch" models together or extract out hidden layer activity. This is a little tricky to get right using GraphModule alone, but we've added some utilities like

  • GraphModulePlus.set_output(layer_name): use this to chop off the head of a model and make it output from a specific layer.
  • GraphModulePlus.new_from_merge(...): use this to merge or "stitch" existing models together. See demos/demo_stitching.py for a worked out example.

We've also done some metaprogramming trickery so that if you import GraphModulePlus anywhere in your code, it will automatically inject itself into the torch.fx module. The surprising but convenient behavior is:

from torch import nn
from torch.fx import symbolic_trace
from nn_lib.models import GraphModulePlus

my_regular_torch_model = nn.Sequential(
    nn.Conv2d(3, 64, 3),
    nn.ReLU(),
    nn.Conv2d(64, 64, 3),
    nn.ReLU(),
    nn.Flatten(),
    nn.LazyLinear(10)
)

# Natively, symbolic_trace is expected to return a GraphModule, but we've injected GraphModulePlus
graphified_model = symbolic_trace(my_regular_torch_model)
assert isinstance(graphified_model, GraphModulePlus)

Useful thing #2: Fancy layers.

We have implemented a few "fancy" layers, available via nn_lib.models or nn_lib.models.fancy_layers that we find useful in our research. These include:

  • Regressable linear layers: a Protocol that allows linear layers to be initialized by least squares regression. This is useful for initializing a linear layer to approximate a function learned by a different model.
  • RegressableLinear: a regressable version of nn.Linear
  • LowRankLinear: a regressable linear layer with a low-rank factorization.
  • ProcrustesLinear: a regressable linear layer constrained to rotation, with optional shift (bias) and optional scaling.
  • A conv2d version of each of the above.

Useful thing #3: MLFLOW utilities.

We use MLFlow to track our experiments. We have a few utilities in nn_lib.utils.mlfow_cli that remove a bit of boilerplate from our code. The biggest contribution here is the run_registry which makes it relatively easy to manage experiments where you want to submit a singleton mlflow run per unique set of parameters.

Useful thing #4: Linear algebra and regression helpers

See nn_lib.utils.pca for a variety of tools for analyzing linear subspaces such as effective dimensionality and calculating principal components from data that may have missing values.

See nn_lib.utils.stats for helpers calculating variances and covariances, such as Welford's algorithm for numerically-stable batch-wise streaming updates of means and variances.

See nn_lib.utils.xval_nuc_norm for some novel methods we're developing to calculate cross-validated nuclear norms of cross-covariance matrices. It's useful for neural (dis)similarity analyses.

See nn_lib.analysis.regression for linear regression utilities such as regressing from x to y from streamed/batched data. This is used extensively in fancy_layers where we support initializing Linear or Conv2d layers by regressing to expected outputs. demos/demo_stitching.py shows off this functionality.

Useful thing #5: NTK utilities.

See nn_lib.analysis.ntk for some neural tangent kernel utilities.

Forthcoming/Planned features

  • More fancy layers
  • Vector Quantization utilities (but see nn_lib.models.sparse_auto_encoder which has some already)
  • Further analysis utilities especially focused on calculating neural similarity measures.

Obsolete/deprecated features

  • lightning training and overly-complex CLI utilities. Some straggler files might still need to be cleaned up.

Test and documentation coverage

We track test coverage (via coverage.py) and docstring coverage (via interrogate) for the badges/coverage.svg and badges/interrogate_badge.svg badges at the top of this file. Since the test suite requires CUDA, these aren't run in GitHub CI — instead, regenerate them locally (e.g. on the lab server) after making changes and commit the updated SVGs:

uv sync --extra dev
scripts/update_coverage_badges.sh

Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

bonsai_nn_library-0.7.1.tar.gz (111.6 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

bonsai_nn_library-0.7.1-py3-none-any.whl (92.0 kB view details)

Uploaded Python 3

File details

Details for the file bonsai_nn_library-0.7.1.tar.gz.

File metadata

  • Download URL: bonsai_nn_library-0.7.1.tar.gz
  • Upload date:
  • Size: 111.6 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/7.0.0 CPython/3.12.11

File hashes

Hashes for bonsai_nn_library-0.7.1.tar.gz
Algorithm Hash digest
SHA256 34faef196960f8a2ff36b542bf5d98865cd86629c650dc71e30d5637b0da1d80
MD5 3b6581dffec71ccfa7fec3432c1dbd23
BLAKE2b-256 8547170a8d5dc761a6e0067cfa354f7544bd98e739ac1c1168660aa64226974e

See more details on using hashes here.

File details

Details for the file bonsai_nn_library-0.7.1-py3-none-any.whl.

File metadata

File hashes

Hashes for bonsai_nn_library-0.7.1-py3-none-any.whl
Algorithm Hash digest
SHA256 e89204809f12cf218a71c347f5488b90eeb22ad3995804a654a0935c2ef83271
MD5 27d4b049f2da887243ae014473f17026
BLAKE2b-256 1448b3e8bebad7a652468316f694fd7c352bb5e9fef8c3f2a74963f7cac9480a

See more details on using hashes here.

Release history Release notifications | RSS feed

This release

0.7.1 This release

2 files

0.7.0

2 files

0.6.5

2 files

0.6.0

2 files

0.5.4

2 files

0.5.3

2 files

0.5.1

2 files

0.5.0

2 files

0.4.9

2 files

0.4.8

2 files

0.4.7

2 files

0.4.6

2 files

0.4.5

2 files

0.4.4

2 files

0.4.3

2 files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page