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.