CoolFace
Apppublic

naver/PUMP

sourceHugging Faceupdated 4y agoView on Hugging Face
1likes
trainer.py126 linesDownload Raw Back to tools
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