Training loop used by `EpochBasedTrainer.train()`
(self, data_loader)
| 1220 | return data_loader |
| 1221 | |
| 1222 | def train_loop(self, data_loader): |
| 1223 | """ Training loop used by `EpochBasedTrainer.train()` |
| 1224 | """ |
| 1225 | self.invoke_hook(TrainerStages.before_run) |
| 1226 | self.model.train() |
| 1227 | for _ in range(self._epoch, self._max_epochs): |
| 1228 | self.invoke_hook(TrainerStages.before_train_epoch) |
| 1229 | for i, data_batch in enumerate(data_loader): |
| 1230 | if i < self.inner_iter: |
| 1231 | # inner_iter may be read out from the checkpoint file, so skip the trained iters in the epoch. |
| 1232 | continue |
| 1233 | data_batch = to_device(data_batch, self.device) |
| 1234 | self.data_batch = data_batch |
| 1235 | self._inner_iter = i |
| 1236 | self.invoke_hook(TrainerStages.before_train_iter) |
| 1237 | self.train_step(self.model, data_batch) |
| 1238 | self.invoke_hook(TrainerStages.after_train_iter) |
| 1239 | # Value changed after the hooks are invoked, do not move them above the invoke_hook code. |
| 1240 | del self.data_batch |
| 1241 | self._iter += 1 |
| 1242 | self._mode = ModeKeys.TRAIN |
| 1243 | |
| 1244 | if i + 1 >= self.iters_per_epoch: |
| 1245 | break |
| 1246 | |
| 1247 | self.invoke_hook(TrainerStages.after_train_epoch) |
| 1248 | # Value changed after the hooks are invoked, do not move them above the invoke_hook code. |
| 1249 | self._inner_iter = 0 |
| 1250 | self._epoch += 1 |
| 1251 | if self._stop_training: |
| 1252 | break |
| 1253 | |
| 1254 | self.invoke_hook(TrainerStages.after_run) |
| 1255 | |
| 1256 | def evaluation_step(self, data): |
| 1257 | """Perform a training step on a batch of inputs. |
no test coverage detected