Run a training step.
(self, data_list)
| 339 | print('Done.') |
| 340 | |
| 341 | def run_step(self, data_list): |
| 342 | """ |
| 343 | Run a training step. |
| 344 | """ |
| 345 | step_log = {'loss': {}, 'status': {}} |
| 346 | amp_context = partial(torch.autocast, device_type='cuda') if self.fp16_mode == 'amp' else nullcontext |
| 347 | elastic_controller_context = self.elastic_controller.record if self.elastic_controller_config is not None else nullcontext |
| 348 | |
| 349 | # Train |
| 350 | losses = [] |
| 351 | statuses = [] |
| 352 | elastic_controller_logs = [] |
| 353 | zero_grad(self.model_params) |
| 354 | for i, mb_data in enumerate(data_list): |
| 355 | ## sync at the end of each batch split |
| 356 | sync_contexts = [self.training_models[name].no_sync for name in self.training_models] if i != len(data_list) - 1 and self.world_size > 1 else [nullcontext] |
| 357 | with nested_contexts(*sync_contexts), elastic_controller_context(): |
| 358 | with amp_context(): |
| 359 | loss, status = self.training_losses(**mb_data) |
| 360 | l = loss['loss'] / len(data_list) |
| 361 | ## backward |
| 362 | if self.fp16_mode == 'amp': |
| 363 | self.scaler.scale(l).backward() |
| 364 | elif self.fp16_mode == 'inflat_all': |
| 365 | scaled_l = l * (2 ** self.log_scale) |
| 366 | scaled_l.backward() |
| 367 | else: |
| 368 | l.backward() |
| 369 | ## log |
| 370 | losses.append(dict_foreach(loss, lambda x: x.item() if isinstance(x, torch.Tensor) else x)) |
| 371 | statuses.append(dict_foreach(status, lambda x: x.item() if isinstance(x, torch.Tensor) else x)) |
| 372 | if self.elastic_controller_config is not None: |
| 373 | elastic_controller_logs.append(self.elastic_controller.log()) |
| 374 | ## gradient clip |
| 375 | if self.grad_clip is not None: |
| 376 | if self.fp16_mode == 'amp': |
| 377 | self.scaler.unscale_(self.optimizer) |
| 378 | elif self.fp16_mode == 'inflat_all': |
| 379 | model_grads_to_master_grads(self.model_params, self.master_params) |
| 380 | self.master_params[0].grad.mul_(1.0 / (2 ** self.log_scale)) |
| 381 | if isinstance(self.grad_clip, float): |
| 382 | grad_norm = torch.nn.utils.clip_grad_norm_(self.master_params, self.grad_clip) |
| 383 | else: |
| 384 | grad_norm = self.grad_clip(self.master_params) |
| 385 | if torch.isfinite(grad_norm): |
| 386 | statuses[-1]['grad_norm'] = grad_norm.item() |
| 387 | ## step |
| 388 | if self.fp16_mode == 'amp': |
| 389 | prev_scale = self.scaler.get_scale() |
| 390 | self.scaler.step(self.optimizer) |
| 391 | self.scaler.update() |
| 392 | elif self.fp16_mode == 'inflat_all': |
| 393 | prev_scale = 2 ** self.log_scale |
| 394 | if not any(not p.grad.isfinite().all() for p in self.model_params): |
| 395 | if self.grad_clip is None: |
| 396 | model_grads_to_master_grads(self.model_params, self.master_params) |
| 397 | self.master_params[0].grad.mul_(1.0 / (2 ** self.log_scale)) |
| 398 | self.optimizer.step() |
nothing calls this directly
no test coverage detected