Args: input_dict (dict): all the values will be reduced average (bool): whether to do average or sum Reduce the values in the dictionary from all processes so that all processes have the averaged results. Returns a dict with the same fields as input_dict, after reduc
(input_dict, average=True)
| 129 | |
| 130 | |
| 131 | def reduce_dict(input_dict, average=True): |
| 132 | """ |
| 133 | Args: |
| 134 | input_dict (dict): all the values will be reduced |
| 135 | average (bool): whether to do average or sum |
| 136 | Reduce the values in the dictionary from all processes so that all processes |
| 137 | have the averaged results. Returns a dict with the same fields as |
| 138 | input_dict, after reduction. |
| 139 | """ |
| 140 | world_size = get_world_size() |
| 141 | if world_size < 2: |
| 142 | return input_dict |
| 143 | with torch.no_grad(): |
| 144 | names = [] |
| 145 | values = [] |
| 146 | # sort the keys so that they are consistent across processes |
| 147 | for k in sorted(input_dict.keys()): |
| 148 | names.append(k) |
| 149 | values.append(input_dict[k]) |
| 150 | values = torch.stack(values, dim=0) |
| 151 | dist.all_reduce(values) |
| 152 | if average: |
| 153 | values /= world_size |
| 154 | reduced_dict = {k: v for k, v in zip(names, values)} |
| 155 | return reduced_dict |
| 156 | |
| 157 | |
| 158 | class MetricLogger(object): |
nothing calls this directly
no test coverage detected