Skip to content
Kaveh Razavi edited this page Jun 28, 2026 · 8 revisions

NNx

Lightweight PyTorch training / eval / visualization toolkit. First-class support for graph neural networks (GCN / GraphSAGE / GAT). Originally extracted from thekaveh/ml to underpin training loops, checkpointing, and result visualization across notebook-based experiments; now standalone. NNx owns the boilerplate around supervised training so you can focus on the model: it builds the network from frozen-dataclass configs, runs the train / eval / predict loop, manages checkpoints under a content-addressed runs/<id>/ directory, dispatches a documented Callback lifecycle, and exposes pluggable extension points for fine-tuning, multi-optimizer training, diffusion, alternative training paradigms, and parameter-efficient fine-tuning.

Install

pip install thekaveh-nnx

Python 3.10+ · PyTorch 2.x. See Installation for optional extras (lm, embeddings, quantize, viz, tensorboard, wandb, hub, gguf-write, onnx-dynamo).

60-second quickstart

import torch
from torch.utils.data import DataLoader, TensorDataset

from nnx import (
    NNModel, NNParams, NNModelParams, NNTrainParams,
    NNOptimParams, NNSchedulerParams,
    Activations, Devices, Losses, Nets, Optims,
    EarlyStopping,
)

# 1. Data
X_train, y_train = torch.randn(256, 8), torch.randint(0, 3, (256,))
X_val,   y_val   = torch.randn(64,  8), torch.randint(0, 3, (64,))
train_loader = DataLoader(TensorDataset(X_train, y_train), batch_size=32, shuffle=True)
val_loader   = DataLoader(TensorDataset(X_val,   y_val),   batch_size=32)

# 2. Model
net_params   = NNParams(input_dim=8, output_dim=3, hidden_dims=[32, 16],
                        dropout_prob=0.1, activation=Activations.RELU)
model_params = NNModelParams(net=Nets.FEED_FWD, device=Devices.CPU,
                             loss=Losses.CROSS_ENTROPY)
model = NNModel(net_params=net_params, params=model_params)

# 3. Train
train_params = NNTrainParams(
    n_epochs=10,
    seed=42,
    train_loader=train_loader,
    val_loader=val_loader,
    optim=NNOptimParams(name=Optims.ADAM, max_lr=1e-2,
                        momentum=(0.9, 0.999), weight_decay=5e-5,
                        grad_clip_norm=1.0),
    scheduler=NNSchedulerParams(min_lr=1e-7, factor=0.5,
                                patience=3, cooldown=1, threshold=1e-3),
)
run = model.train(params=train_params, callbacks=[EarlyStopping(patience=5)])

# 4. Use it
print(f"trained {len(run.idps)} iterations; saved under runs/{run.id}/")
result = model.predict(X=X_val)
print(f"predicted {len(result.classes)} samples")

See Quickstart for GPU / MPS variants, warm-resume, custom metrics, scheduler choices, TensorBoard, LR finder, and more.

Capabilities at a glance

NNx covers supervised training, GNNs, Transformers, PEFT (LoRA / DoRA / IA³ / prefix / prompt / adapters), QAT + PTQ quantization, pruning, model surgery, contrastive embeddings, GGUF + Ollama export, HuggingFace Hub, and rich visualization — all wired through the same train_step_fn hook.

Capabilities-at-a-Glance — full capability matrix with one-liners per subsystem.

Architecture

NNx is organized around two public entry points (NNModel / Trainer), a single extensibility hook (train_step_fn), and content-addressed persistence in runs/<id>/. The diagram below shows the eight-layer top-to-bottom flow:

flowchart TD
    A["1 · User code + PyTorch"]
    B["2 · NNModel / Trainer\n(two public entry points)"]
    C["3 · train_step_fn / trainer_step_fn\n(orange hook — the specialization bus)"]
    D["4 · Specialization subpackages\nfinetune · peft · prune · surgery · quantize\ndiffusion · paradigms · trainer\nembeddings · interop · generation · viz"]
    E["5 · Training-loop internals\nepoch × batch dispatch · _step_scheduler\n_save_checkpoints · NaN guard · grad-clip"]
    F["6 · Callback bus\non_train_begin · on_epoch_begin\non_epoch_end · on_train_end"]
    G["7 · Callback listeners\nEarlyStopping · LRMonitor · ModelCheckpoint\nTensorBoardCallback · WandbCallback"]
    H["8 · Persistence\nNNRun + NNCheckpoint → runs/<id>/"]

    A --> B
    B --> C
    C --> D
    B --> E
    E --> F
    F --> G
    E --> H
Loading

Full prose walkthrough: Core-Concepts-and-Architecture.

Start here

Getting Started

Installation Python version support, optional extras, dev setup
Quickstart End-to-end CPU example; GPU / MPS / resume / metrics variants
Core-Concepts-and-Architecture Two-class surface, frozen params, hook pattern, persistence

Core Subsystems

Networks FeedFwdNN, GraphConvNN/GraphSageNN/GraphAttNN, TransformerNN, ViTNN
Training-Loop-and-Callbacks Epoch/batch loop, Callback lifecycle, EarlyStopping, LRMonitor
Multi-Optimizer-Trainer GAN / actor-critic multi-optimizer with disjoint param groups
Persistence-Runs-and-Checkpoints Content-addressed runs/<id>/, six checkpoint tags, warm-resume
Datasets NNDataset, NNGraphDataset, NNTabularDataset
Reproducibility-and-Diagnostics set_seed, dataloader_worker_init_fn, LR finder

Specialization Guides

Fine-Tuning Layer freezing, param groups, pretrained state-dict loading
PEFT LoRA, DoRA, IA³, prefix tuning, prompt tuning, adapters
Training-Paradigms KD, feature-KD, SimCLR, Mixup, CutMix, MoE, Born-Again, DPO
I-JEPA Self-supervised pretraining with masked latent prediction
DPO Preference-pair fine-tuning (Rafailov et al. 2023)
Diffusion DDPM denoiser, noise schedules, reverse sampler
Quantization PTQ INT8 weight-only + QAT 8da4w (torchao)
Pruning Magnitude unstructured + 2:4 semi-structured
Model-Surgery Net2Net widen/deepen, drop, low-rank factorize, embedding expand
Embeddings-and-FAISS Contrastive trainer + FAISS index export
Language-Modeling TransformerNN, KV-cache, tokenizer
Text-Generation generate(), LogitsProcessor chain, sampling strategies
GGUF-and-Ollama llama.cpp-compatible export + Ollama Modelfile
HuggingFace-Hub push_to_hub, from_pretrained, safetensors checkpoints
Visualization VisUtils run-output plots + model-internals viz (Captum, Netron)

Reference

Capabilities-at-a-Glance One-line summary of every subsystem
Params-Reference-Core NNParams, NNModelParams, NNTrainParams, NNOptimParams, NNSchedulerParams
Params-Reference-Trainer NNTrainerParams, multi-optim shape
Params-Reference-Transformer-and-Tokenizer NNTransformerParams, NNTokenizerParams
Enums-Reference Nets, Losses, Optims, Schedulers, Activations, Devices, Checkpoints, NoiseSchedulers
Fluent-Builders .builder() API for all params classes
Train-Step-Factories Every *_train_step_factory symbol, signatures
Public-API-Surface Full nnx.__all__ listing
Params-and-Back-Compat state() / from_state() contract, omit-when-default rule
Examples-Catalog All 24 examples in examples/, grouped by topic
Comparison NNx vs Lightning vs HF Transformers vs fastai vs Composer

Project

Contributing Dev setup, branching, ruff / mypy / pytest, PR checklist
FAQ-and-Troubleshooting Common errors, device selection, import issues
Status-and-Roadmap Current release, known gaps, planned features
Releases-and-Changelog Version history

See also

Clone this wiki locally