(self)
| 400 | self.logger.log_info('Resume from {}'.format(path)) |
| 401 | |
| 402 | def train_epoch(self): |
| 403 | self.model.train() |
| 404 | self.last_epoch += 1 |
| 405 | |
| 406 | if self.args.distributed: |
| 407 | self.dataloader['train_loader'].sampler.set_epoch(self.last_epoch) |
| 408 | |
| 409 | epoch_start = time.time() |
| 410 | itr_start = time.time() |
| 411 | itr = -1 |
| 412 | for itr, batch in enumerate(self.dataloader['train_loader']): |
| 413 | if itr == 0: |
| 414 | print("time2 is " + str(time.time())) |
| 415 | data_time = time.time() - itr_start |
| 416 | step_start = time.time() |
| 417 | self.last_iter += 1 |
| 418 | loss = self.step(batch, phase='train') |
| 419 | # logging info |
| 420 | if self.logger is not None and self.last_iter % self.args.log_frequency == 0: |
| 421 | info = '{}: train'.format(self.args.name) |
| 422 | info = info + ': Epoch {}/{} iter {}/{}'.format(self.last_epoch, self.max_epochs, self.last_iter%self.dataloader['train_iterations'], self.dataloader['train_iterations']) |
| 423 | for loss_n, loss_dict in loss.items(): |
| 424 | info += ' ||' |
| 425 | loss_dict = reduce_dict(loss_dict) |
| 426 | info += '' if loss_n == 'none' else ' {}'.format(loss_n) |
| 427 | # info = info + ': Epoch {}/{} iter {}/{}'.format(self.last_epoch, self.max_epochs, self.last_iter%self.dataloader['train_iterations'], self.dataloader['train_iterations']) |
| 428 | for k in loss_dict: |
| 429 | info += ' | {}: {:.4f}'.format(k, float(loss_dict[k])) |
| 430 | self.logger.add_scalar(tag='train/{}/{}'.format(loss_n, k), scalar_value=float(loss_dict[k]), global_step=self.last_iter) |
| 431 | |
| 432 | # log lr |
| 433 | lrs = self._get_lr(return_type='dict') |
| 434 | for k in lrs.keys(): |
| 435 | lr = lrs[k] |
| 436 | self.logger.add_scalar(tag='train/{}_lr'.format(k), scalar_value=lrs[k], global_step=self.last_iter) |
| 437 | |
| 438 | # add lr to info |
| 439 | info += ' || {}'.format(self._get_lr()) |
| 440 | |
| 441 | # add time consumption to info |
| 442 | spend_time = time.time() - self.start_train_time |
| 443 | itr_time_avg = spend_time / (self.last_iter + 1) |
| 444 | info += ' || data_time: {dt}s | fbward_time: {fbt}s | iter_time: {it}s | iter_avg_time: {ita}s | epoch_time: {et} | spend_time: {st} | left_time: {lt}'.format( |
| 445 | dt=round(data_time, 1), |
| 446 | it=round(time.time() - itr_start, 1), |
| 447 | fbt=round(time.time() - step_start, 1), |
| 448 | ita=round(itr_time_avg, 1), |
| 449 | et=format_seconds(time.time() - epoch_start), |
| 450 | st=format_seconds(spend_time), |
| 451 | lt=format_seconds(itr_time_avg*self.max_epochs*self.dataloader['train_iterations']-spend_time) |
| 452 | ) |
| 453 | self.logger.log_info(info) |
| 454 | |
| 455 | itr_start = time.time() |
| 456 | |
| 457 | # sample |
| 458 | if self.sample_iterations > 0 and (self.last_iter + 1) % self.sample_iterations == 0: |
| 459 | # print("save model here") |
no test coverage detected