Reduce a tensor of losses across all GPUs.
(losses)
| 79 | |
| 80 | |
| 81 | def average_losses_across_data_parallel_group(losses): |
| 82 | """Reduce a tensor of losses across all GPUs.""" |
| 83 | averaged_losses = torch.cat([loss.clone().detach().view(1) for loss in losses]) |
| 84 | torch.distributed.all_reduce(averaged_losses, group=mpu.get_data_parallel_group()) |
| 85 | averaged_losses = averaged_losses / torch.distributed.get_world_size( |
| 86 | group=mpu.get_data_parallel_group() |
| 87 | ) |
| 88 | |
| 89 | return averaged_losses |
| 90 | |
| 91 | |
| 92 | def report_memory(name): |
no outgoing calls
no test coverage detected