naver/PUMP
1
1# Copyright 2022-present NAVER Corp.2# CC BY-NC-SA 4.03# Available only for non-commercial use4 5import pdb; bb = pdb.set_trace6from tqdm import tqdm7from collections import defaultdict8 9import torch10import torch.nn as nn11from torch.nn import DataParallel12 13from .common import todevice14 15 16class Trainer (nn.Module):17 """ Helper class to train a deep network.18 Overload this class `forward_backward` for your actual needs.19 20 Usage: 21 train = Trainer(net, loss, optimizer)22 for epoch in range(n_epochs):23 train()24 """25 def __init__(self, net, loss, optimizer, epoch=0):26 super().__init__()27 self.net = net28 self.loss = loss29 self.optimizer = optimizer30 self.epoch = epoch31 32 @property33 def device(self):34 return next(self.net.parameters()).device35 36 @property37 def model(self):38 return self.net.module if isinstance(self.net, DataParallel) else self.net39 40 def distribute(self):41 self.net = DataParallel(self.net) # DataDistributed not implemented yet42 43 def __call__(self, data_loader):44 print(f'>> Training (epoch {self.epoch} --> {self.epoch+1})')45 self.net.train()46 47 stats = defaultdict(list)48 49 for batch in tqdm(data_loader):50 batch = todevice(batch, self.device)51 52 # compute gradient and do model update53 self.optimizer.zero_grad()54 details = self.forward_backward(batch)55 self.optimizer.step()56 57 for key, val in details.items():58 stats[key].append( val )59 60 self.epoch += 161 62 print(" Summary of losses during this epoch:")63 for loss_name, vals in stats.items():64 N = 1 + len(vals)//1065 print(f" - {loss_name:10}: {avg(vals[:N]):.3f} --> {avg(vals[-N:]):.3f} (avg: {avg(vals):.3f})")66 67 def forward_backward(self, inputs):68 raise NotImplementedError()69 70 def save(self, path):71 print(f"\n>> Saving model to {path}")72 73 data = {'model': self.model.state_dict(),74 'optimizer': self.optimizer.state_dict(),75 'loss': self.loss.state_dict(),76 'epoch': self.epoch}77 78 torch.save(data, open(path,'wb'))79 80 def load(self, path, resume=True):81 print(f">> Loading weights from {path} ...")82 checkpoint = torch.load(path, map_location='cpu')83 assert isinstance(checkpoint, dict)84 85 self.net.load_state_dict(checkpoint['model'])86 if resume:87 self.optimizer.load_state_dict(checkpoint['optimizer'])88 self.loss.load_state_dict(checkpoint['optimizer'])89 self.epoch = checkpoint['epoch']90 print(f" Resuming training at Epoch {self.epoch}!")91 92 93def get_loss( loss ):94 """ returns a tuple (loss, dictionary of loss details)95 """96 assert isinstance(loss, dict)97 grads = None98 99 k,l = next(iter(loss.items())) # first item is assumed to be the main loss100 if isinstance(l, tuple):101 l, grads = l102 loss[k] = l103 104 return (l, grads), {k:float(v) for k,v in loss.items()}105 106 107def backward( loss ):108 if isinstance(loss, tuple):109 loss, grads = loss110 else:111 loss, grads = (loss, None)112 113 assert loss == loss, 'loss is NaN'114 115 if grads is None:116 loss.backward()117 else:118 # dictionary of separate subgraphs119 for var,grad in grads:120 var.backward(grad)121 return float(loss)122 123 124def avg( lis ):125 return sum(lis) / len(lis)126 