-
Notifications
You must be signed in to change notification settings - Fork 1
Home
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.
pip install thekaveh-nnxPython 3.10+ · PyTorch 2.x. See Installation for optional extras (lm, embeddings, quantize, viz, tensorboard, wandb, hub, gguf-write, onnx-dynamo).
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.
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.
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
Full prose walkthrough: Core-Concepts-and-Architecture.
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 |
Apache-2.0 licensed.