MCPcopy Create free account
hub / github.com/cientgu/VQ-Diffusion / train_epoch

Method train_epoch

image_synthesis/engine/solver.py:402–468  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

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")

Callers 1

trainMethod · 0.95

Calls 10

stepMethod · 0.95
_get_lrMethod · 0.95
sampleMethod · 0.95
reduce_dictFunction · 0.90
format_secondsFunction · 0.90
add_scalarMethod · 0.80
log_infoMethod · 0.80
evalMethod · 0.80
trainMethod · 0.45
keysMethod · 0.45

Tested by

no test coverage detected