reduce loss dict. In distributed training, it averages the losses among different GPUs . Args: loss_dict (OrderedDict): Loss dict.
(self, loss_dict)
| 369 | self.schedulers[i].load_state_dict(s) |
| 370 | |
| 371 | def reduce_loss_dict(self, loss_dict): |
| 372 | """reduce loss dict. |
| 373 | |
| 374 | In distributed training, it averages the losses among different GPUs . |
| 375 | |
| 376 | Args: |
| 377 | loss_dict (OrderedDict): Loss dict. |
| 378 | """ |
| 379 | with torch.no_grad(): |
| 380 | if self.opt['dist']: |
| 381 | keys = [] |
| 382 | losses = [] |
| 383 | for name, value in loss_dict.items(): |
| 384 | keys.append(name) |
| 385 | losses.append(value) |
| 386 | losses = torch.stack(losses, 0) |
| 387 | torch.distributed.reduce(losses, dst=0) |
| 388 | if self.opt['rank'] == 0: |
| 389 | losses /= self.opt['world_size'] |
| 390 | loss_dict = {key: loss for key, loss in zip(keys, losses)} |
| 391 | |
| 392 | log_dict = OrderedDict() |
| 393 | for name, value in loss_dict.items(): |
| 394 | log_dict[name] = value.mean().item() |
| 395 | |
| 396 | return log_dict |
no outgoing calls
no test coverage detected