(self, step_idx: int = 0)
| 597 | self.is_valid = False |
| 598 | |
| 599 | def step(self, step_idx: int = 0): |
| 600 | assert self.is_valid, self.error_msg |
| 601 | for optimizer_idx in range(0, len(self._opts)): |
| 602 | batch = self._cfg.data(step_idx, optimizer_idx) |
| 603 | outputs = self._model.training_step( |
| 604 | batch=batch, optimizer_idx=optimizer_idx |
| 605 | ) |
| 606 | loss = None |
| 607 | if isinstance(outputs, tuple) and len(outputs) > 0: |
| 608 | loss = outputs[0] |
| 609 | else: |
| 610 | loss = outputs |
| 611 | loss.backward() |
| 612 | opt = self._opts[optimizer_idx] |
| 613 | opt.step() |
| 614 | opt.zero_grad() |
| 615 | self._method_callback( |
| 616 | "on_training_step_end", |
| 617 | outputs=outputs, |
| 618 | step_idx=step_idx, |
| 619 | optimizer_idx=optimizer_idx, |
| 620 | ) |
| 621 | |
| 622 | def _get_and_check_step(self): |
| 623 | if not self._model.method_overrided("training_step"): |
nothing calls this directly
no test coverage detected