How do I calculate the overall loss for an epoch with distributed training via ray?

Viewed 39

I'm training a model using Ray Train across 2 workers. Prior to using ray, I would have used a function like this:

train_loader = trcdata.DataLoader(train_set, batch_size, num_workers=4)

# Other code here

def get_overall_mse(loader):
    all_outputs = []
    all_labels = []

    for x, labels in loader:
        outputs = model.forward(x, True)
        outputs = torch.squeeze(outputs)

        all_outputs.append(outputs.detach())
        all_labels.append(labels)

    all_outputs = torch.cat(all_outputs)
    all_labels = torch.cat(all_labels)

    return torchmetrics.functional.mean_squared_error(all_outputs, all_labels).item()

train_mse = get_overall_mse(train_loader)

However, if I do this with Ray I end up with two separate sets of metrics for each worker in any callback triggered ray.train.report().

To note, I am preparing my datasets with ray as follows:

model = ray.train.torch.prepare_model(model)


train_loader = trcdata.DataLoader(train_set, batch_size, num_workers=4)
train_loader = ray.train.torch.prepare_data_loader(train_loader)

Is there anyway to ensure I get a single statistic for the entire dataset? I want this so I can properly checkpoint models with their evaluated statistics. I had thought about implementing a callback to calculate the final stats for the epoch - however I do not believe this will work properly with checkpointing the model.

0 Answers
Related