Skip to main content

Tiny DDP training toolkit for quick-launch distributed training loops.

Project description

torchlitex

Tiny DDP launcher + trainer that exists because PyTorch 2.x still thinks we enjoy 400 lines of torchrun boilerplate and cryptic NCCL errors. This trims it to ~20 lines and keeps fork happy on Vast.

Why not just torchrun?

  • You like your code more than the 17 environment variables PyTorch asks you to memorize.
  • torchrun still feels like a 2010 MPI cosplay.
  • You want fork-based spawn that doesn’t randomly faceplant on Vast.
  • You want a one-function launcher + trainer, not a CLI maze.

Install

pip install -e .
# or with wandb logging
pip install -e .[wandb]

Quick start

from torchlitex.launcher import launch, DistributedConfig
from torchlitex.trainer import Trainer
from torch import nn, optim
import torch

def train_fn(rank, world_size, batch_size, epochs):
    dataset = MyDataset(...)
    model = MyModel(...)
    loss_fn = nn.CrossEntropyLoss()
    optimizer = optim.AdamW(model.parameters(), lr=3e-4)

    trainer = Trainer(
        model=model,
        dataset=dataset,
        loss_fn=loss_fn,
        optimizer=optimizer,
        grad_clip_norm=1.0,
        log_every=10,
    )
    trainer.ddp_train_loop(rank, world_size, batch_size=batch_size, epochs=epochs, ckpt_path="ckpt.pt")

if __name__ == "__main__":
    cfg = DistributedConfig(gpus=8)  # backend auto-switches nccl/gloo
    launch(train_fn, cfg, batch_size=64, epochs=20)

Features

  • Fork-first DDP launcher (no torchrun, no elastic).
  • Auto backend: nccl when CUDA exists, gloo when you're on a laptop/CI.
  • Trainer: AMP toggle, grad clipping, gradient accumulation, microbatching, optional schedulers, eval hook, callbacks, and EMA support.
  • Rich callback system with access to model outputs, inputs, targets for custom progress bars and visualizations.
  • Built-in validation loop with per-epoch val loss and sample outputs for reconstruction visualization.
  • Optional wandb init + logging on rank0 (install with .[wandb]).
  • DistributedSampler + DataLoader defaults that just work.
  • Start method auto-picks spawn when CUDA is present (avoids forked-CUDA init errors); override to fork if you really want it.
  • Checkpoint utilities that handle optimizer/scaler safely.
  • Rank-aware logging that doesn't spam.

Callbacks

The trainer fires rich callbacks with access to model outputs, inputs, and targets. This enables custom progress bars, reconstruction visualization, and more.

Callback hooks

Hook Parameters
on_train_start trainer, rank
on_epoch_start epoch, num_batches, trainer, rank
on_batch_start epoch, batch_idx, num_batches, trainer, rank
on_step_end epoch, step, batch_idx, num_batches, trainer, rank, loss, logits, inputs, targets
on_batch_end epoch, batch_idx, num_batches, trainer, rank, loss
on_val_end epoch, val_loss, trainer, rank, sample_logits, sample_inputs, sample_targets
on_epoch_end epoch, loss, val_loss, best, trainer, rank
on_train_end trainer, rank

Example: tqdm progress bar

from tqdm import tqdm

class TQDMCallback:
    def __init__(self):
        self.pbar = None

    def on_epoch_start(self, epoch, num_batches, rank, **_):
        if rank == 0:
            self.pbar = tqdm(total=num_batches, desc=f"Epoch {epoch}")

    def on_batch_end(self, rank, **_):
        if rank == 0 and self.pbar:
            self.pbar.update(1)

    def on_epoch_end(self, rank, **_):
        if rank == 0 and self.pbar:
            self.pbar.close()

Example: reconstruction visualization (autoencoders)

import matplotlib.pyplot as plt

class ReconCallback:
    def on_val_end(self, epoch, sample_inputs, sample_logits, rank, **_):
        if rank != 0 or sample_inputs is None:
            return
        # For an autoencoder: inputs are original, logits are reconstructed
        orig = sample_inputs[0].cpu()  # first sample
        recon = sample_logits[0].cpu()
        fig, axes = plt.subplots(1, 2)
        axes[0].imshow(orig.permute(1, 2, 0)); axes[0].set_title("Original")
        axes[1].imshow(recon.permute(1, 2, 0)); axes[1].set_title("Reconstructed")
        plt.savefig(f"recon_epoch_{epoch}.png")
        plt.close()

Example: validation dataset

trainer = Trainer(
    model=model,
    dataset=train_dataset,
    val_dataset=val_dataset,  # built-in val loop
    loss_fn=loss_fn,
    optimizer=optimizer,
    callbacks=[TQDMCallback(), ReconCallback()],
)

Testing levels

  • Level 1: CPU unit tests (no dist).
  • Level 2: CPU DDP (backend="gloo", world_size=2) to validate spawn/env/sampler.
  • Level 3: Single GPU (gpus=1) for end-to-end DDP path.
  • Level 4: Real multi-GPU (same code, just crank gpus).

PyTorch 2.x roast (lightly toasted)

  • DDP config still feels like “choose your own adventure” but every page ends with NCCL complaining.
  • torch.distributed docs read like a treasure map; the treasure is another flag.
  • “Just use torchrun” is 2020’s “have you tried turning it off and on again?”

torchlitex keeps the good bits of torch 2.x (SDPA, compile, etc.) and sidesteps the distributed busywork. Use it, ship models, spend less time appeasing the NCCL spirits.

Project details


Download files

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

Source Distribution

torchlitex-0.1.6.tar.gz (12.5 kB view details)

Uploaded Source

Built Distribution

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

torchlitex-0.1.6-py3-none-any.whl (11.4 kB view details)

Uploaded Python 3

File details

Details for the file torchlitex-0.1.6.tar.gz.

File metadata

  • Download URL: torchlitex-0.1.6.tar.gz
  • Upload date:
  • Size: 12.5 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/6.2.0 CPython/3.11.13

File hashes

Hashes for torchlitex-0.1.6.tar.gz
Algorithm Hash digest
SHA256 a80675dbe94b2c237e04cff0db7b28d900f65b58eab5f5b742c47cf30c65d431
MD5 9bf1aaee51aa7f2f4b4379c627423c24
BLAKE2b-256 a428a5247c133628c52764b659ce59b40e8e3e1309dcc03f4716c4f95b3707ee

See more details on using hashes here.

File details

Details for the file torchlitex-0.1.6-py3-none-any.whl.

File metadata

  • Download URL: torchlitex-0.1.6-py3-none-any.whl
  • Upload date:
  • Size: 11.4 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/6.2.0 CPython/3.11.13

File hashes

Hashes for torchlitex-0.1.6-py3-none-any.whl
Algorithm Hash digest
SHA256 366bb6135e3c9aaaf10158e1ae28b55f938bcdc589c69a43296a5da2cc69b05f
MD5 9b9a18e59ddd3a0548c928166216dfb4
BLAKE2b-256 9e1d9250653b25ae5d672095696daddd471ac5099ac9beeef3dc3291bdb1313d

See more details on using hashes here.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Pingdom Monitoring Sentry Error logging StatusPage Status page