MaRN#

Manifold Regularized Networks (marn) is a model-agnostic PyTorch library for training target models through low-dimensional, generated parameter spaces. It implements the ideas from Mapping Networks as a reusable toolkit: frozen targets, explicit generation strategies, composite mapping losses, trainers, and versioned checkpoints.

Explore the documentation:

  • User guide — concepts, data flow, and subsystem walkthroughs.

  • Cookbook — runnable end-to-end training recipes in the repository.

  • API reference — autogenerated signatures and source links for every public type.

Installation#

From PyPI#

pip install marn

With Poetry:

poetry add marn

From source#

Clone the repository and install in editable mode (includes the cookbook scripts):

git clone https://github.com/arjunmnath/MaRN.git
cd MaRN
poetry install

Requirements#

  • Python: 3.12+ (see pyproject.toml for the exact supported range)

  • PyTorch: >=2.6, <3.0

  • Pydantic: >=2.0

  • PyYAML: ^6.0.3

Quickstart#

This example trains a small classifier through a layerwise latent mapping in a few lines. For a full CNN + SLVT workflow, see cookbook recipe 01.

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

from marn import ClassificationLoss, MappingLoss, MappingModel, MappingTrainer

# 1. Target model and synthetic data
target = nn.Sequential(nn.Linear(10, 16), nn.ReLU(), nn.Linear(16, 2))
features = torch.randn(128, 10)
labels = torch.randint(0, 2, (128,))
loader = DataLoader(TensorDataset(features, labels), batch_size=16, shuffle=True)

# 2. Wrap the target with a low-dimensional mapping strategy
model = MappingModel(
    target_model=target,
    latent_dim=32,
    strategy="layerwise",
)

# 3. Composite task + regularization loss and trainer
loss_fn = MappingLoss(task_loss=ClassificationLoss())
trainer = MappingTrainer(
    model=model,
    train_loader=loader,
    loss_fn=loss_fn,
    learning_rate=1e-3,
)

# 4. Optimize latent coordinates (target weights stay frozen)
trainer.fit(epochs=5)

Package overview#

Area

Role

Runtime

ParameterSpec / ParameterTree name and store target weights for generation.

Models

TargetModel (stateless forward) and MappingModel (latent → weights → prediction).

Mappers & modulation

Project latents and apply additive, affine, or low-rank updates to parameters.

Generators & strategies

SLVT, layerwise, grouped, LRD, and fine-tuning parameter layouts.

Losses

Task loss plus stability, smoothness, and alignment regularizers.

Training

MappingTrainer, batch adapters, callbacks, LR finder, checkpoint I/O.

Scale-out

DDP helpers, lazy generation, memory profiling, torch.compile compatibility.