Skip to content

Repository files navigation

EXO Gym

Open source framework for simulated distributed training methods. Instead of training with multiple ranks, we simulate the distributed training process by running multiple nodes on a single machine.

Supported Devices

  • CPU
  • CUDA
  • MPS (CPU-bound for copy operations, see here)

Supported Methods

Installation

Basic Installation

Install with core dependencies only:

pip install exogym

Installation with Optional Features

Optional feature flags allowed are:

wandb,gpt,demo,examples,all,dev

For example, pip install exogym[demo]

Development Installation

To install for development:

git clone https://github.com/exo-explore/gym.git exogym
cd exogym
pip install -e ".[dev]"

Usage

Example Scripts

MNIST comparison of DDP, DiLoCo, and SPARTA:

python run/mnist.py

NanoGPT Shakespeare DiLoCo:

python run/nanogpt_diloco.py --dataset shakespeare

Custom Training

from exogym import LocalTrainer
from exogym.strategy import DiLoCoStrategy

train_dataset, val_dataset = ...
model = ... # model.forward() expects a batch, and returns a scalar loss

trainer = LocalTrainer(model, train_dataset, val_dataset)

# Strategy for optimization & communication
strategy = DiLoCoStrategy(
  inner_optim='adam',
  H=100
)

trainer.fit(
  strategy=strategy,
  num_nodes=4,
  device='mps'
)

Codebase Structure

  • Trainer: Builds simulation environment. Trainer will spawn multiple TrainNode instances, connect them together, and starts the training run.
  • TrainNode: A single node (rank) running its own training loop. At each train step, instead of calling optim.step(), it calls strategy.step().
  • Strategy: Abstract class for an optimization strategy, which both defines how the nodes communicate with each other and how model weights are updated. Typically, a gradient strategy will include an optimizer as well as a communication step. Sometimes (eg. DeMo), the optimizer step is comingled with the communication.

Technical Details

EXO Gym uses pytorch multiprocessing to a subprocess per-node, which are able to communicate with each other using regular operations such as all_reduce.

Model

The model is expected in a form that takes a batch (the same format as dataset outputs), and returns a scalar loss over the entire batch. This ensures the model is agnostic to the format of the data (eg. masked LM training doesn't have a clear x/y split).

Dataset

Recall that when we call trainer.fit(), $K$ subprocesses are spawned to handle each of the virtual workers. The dataset object is passed to every subprocess, and a DistributedSampler will be used to select indices per-node. If the dataset is entirely loaded into memory, this memory will be duplicated per-node - be careful not to run out of memory! If the dataset is larger, it should be lazily loaded.

About

No description, website, or topics provided.

Resources

Stars

1 star

Watchers

2 watching

Forks

Releases

Packages

Used by

Contributors

Languages