(self, bucket, dp_group)
| 2222 | log_dist(f"step={step}, skipped={self.skipped_steps}, lr={lr}, mom={mom}", ranks=[0]) |
| 2223 | |
| 2224 | def allreduce_bucket(self, bucket, dp_group): |
| 2225 | tensor = self.flatten(bucket) |
| 2226 | |
| 2227 | tensor_to_allreduce = tensor |
| 2228 | |
| 2229 | if self.communication_data_type != tensor.dtype: |
| 2230 | tensor_to_allreduce = tensor.to(self.communication_data_type) |
| 2231 | |
| 2232 | if self.postscale_gradients(): |
| 2233 | if self.gradient_predivide_factor() != 1.0: |
| 2234 | tensor_to_allreduce.mul_(1.0 / self.gradient_predivide_factor()) |
| 2235 | |
| 2236 | dist.all_reduce(tensor_to_allreduce, group=dp_group) |
| 2237 | if self.gradient_average: |
| 2238 | if self.gradient_predivide_factor() != dist.get_world_size(group=dp_group): |
| 2239 | tensor_to_allreduce.mul_(self.gradient_predivide_factor() / dist.get_world_size(group=dp_group)) |
| 2240 | else: |
| 2241 | tensor_to_allreduce.mul_(1. / dist.get_world_size(group=dp_group)) |
| 2242 | dist.all_reduce(tensor_to_allreduce, group=dp_group) |
| 2243 | |
| 2244 | if self.communication_data_type != tensor.dtype and tensor is not tensor_to_allreduce: |
| 2245 | tensor.copy_(tensor_to_allreduce) |
| 2246 | |
| 2247 | return tensor |
| 2248 | |
| 2249 | def allreduce_and_copy(self, small_bucket, dp_group): |
| 2250 | allreduced = self.allreduce_bucket(small_bucket, dp_group) |
no test coverage detected