import os import torch import torch.distributed as dist from torch import nn from torch.nn.parallel import DistributedDataParallel from torch.utils.data import DataLoader, TensorDataset from torch.utils.data.distributed import DistributedSampler def build_training_state(rank, world_size): torch.manual_seed(17) model = DistributedDataParallel(nn.Linear(2, 1)) features = torch.tensor( [[0.0, 0.0], [1.0, 1.0], [2.0, 2.0], [3.0, 3.0]], dtype=torch.float32, ) targets = torch.tensor([[0.0], [2.0], [4.0], [6.0]], dtype=torch.float32) dataset = TensorDataset(features, targets) sampler = DistributedSampler( dataset, num_replicas=world_size, rank=rank, shuffle=False, ) loader = DataLoader(dataset, batch_size=2, sampler=sampler) optimizer = torch.optim.SGD(model.parameters(), lr=0.1) return model, loader, optimizer