PyTorch#

The ray.train.torch module runs distributed PyTorch training. TorchXLAConfig lives in ray.train.torch.xla.

Trainer and configs#

TorchTrainer

A Trainer for data parallel PyTorch training.

TorchConfig

Configuration for torch process group setup.

TorchXLAConfig

Configuration for torch XLA setup.

Training loop utilities#

get_device

Gets the correct torch device configured for the current worker.

get_devices

Gets the list of torch devices configured for the current worker.

prepare_model

Prepares the model for distributed execution.

prepare_data_loader

Prepares DataLoader for distributed execution.

enable_reproducibility

Limits sources of nondeterministic behavior.