po03087's picture
SRA: MID/LED/MoFlow code + RUNNING.md instructions (code only, no data/ckpts)
d4cbafd verified
Raw
History Blame Contribute Delete
4.4 kB
import os
import torch
import numpy as np
import random
import time
import logging
import logging.handlers
THOUSAND = 1000
MILLION = 1000000
class BlackHole(object):
def __setattr__(self, name, value):
pass
def __call__(self, *args, **kwargs):
return self
def __getattr__(self, name):
return self
class CheckpointManager(object):
def __init__(self, save_dir, logger=BlackHole()):
super().__init__()
os.makedirs(save_dir, exist_ok=True)
self.save_dir = save_dir
self.ckpts = []
self.logger = logger
for f in os.listdir(self.save_dir):
if f[:4] != 'ckpt':
continue
_, score, it = f.split('_')
it = it.split('.')[0]
self.ckpts.append({
'score': float(score),
'file': f,
'iteration': int(it),
})
def get_worst_ckpt_idx(self):
idx = -1
worst = float('-inf')
for i, ckpt in enumerate(self.ckpts):
if ckpt['score'] >= worst:
idx = i
worst = ckpt['score']
return idx if idx >= 0 else None
def get_best_ckpt_idx(self):
idx = -1
best = float('inf')
for i, ckpt in enumerate(self.ckpts):
if ckpt['score'] <= best:
idx = i
best = ckpt['score']
return idx if idx >= 0 else None
def get_latest_ckpt_idx(self):
idx = -1
latest_it = -1
for i, ckpt in enumerate(self.ckpts):
if ckpt['iteration'] > latest_it:
idx = i
latest_it = ckpt['iteration']
return idx if idx >= 0 else None
def save(self, model, args, score, others=None, step=None):
if step is None:
fname = 'ckpt_%.6f_.pt' % float(score)
else:
fname = 'ckpt_%.6f_%d.pt' % (float(score), int(step))
path = os.path.join(self.save_dir, fname)
torch.save({
'args': args,
'state_dict': model.state_dict(),
'others': others
}, path)
self.ckpts.append({
'score': score,
'file': fname
})
return True
def load_best(self):
idx = self.get_best_ckpt_idx()
if idx is None:
raise IOError('No checkpoints found.')
ckpt = torch.load(os.path.join(self.save_dir, self.ckpts[idx]['file']))
return ckpt
def load_latest(self):
idx = self.get_latest_ckpt_idx()
if idx is None:
raise IOError('No checkpoints found.')
ckpt = torch.load(os.path.join(self.save_dir, self.ckpts[idx]['file']))
return ckpt
def load_selected(self, file):
ckpt = torch.load(os.path.join(self.save_dir, file))
return ckpt
def seed_all(seed):
torch.manual_seed(seed)
np.random.seed(seed)
random.seed(seed)
def get_logger(name, log_dir=None):
logger = logging.getLogger(name)
logger.setLevel(logging.DEBUG)
formatter = logging.Formatter('[%(asctime)s::%(name)s::%(levelname)s] %(message)s')
stream_handler = logging.StreamHandler()
stream_handler.setLevel(logging.DEBUG)
stream_handler.setFormatter(formatter)
logger.addHandler(stream_handler)
if log_dir is not None:
file_handler = logging.FileHandler(os.path.join(log_dir, 'log.txt'))
file_handler.setLevel(logging.INFO)
file_handler.setFormatter(formatter)
logger.addHandler(file_handler)
return logger
def get_new_log_dir(root='./logs', postfix='', prefix=''):
log_dir = os.path.join(root, prefix + time.strftime('%Y_%m_%d__%H_%M_%S', time.localtime()) + postfix)
os.makedirs(log_dir)
return log_dir
def int_tuple(argstr):
return tuple(map(int, argstr.split(',')))
def str_tuple(argstr):
return tuple(argstr.split(','))
def int_list(argstr):
return list(map(int, argstr.split(',')))
def str_list(argstr):
return list(argstr.split(','))
def log_hyperparams(writer, args):
from torch.utils.tensorboard.summary import hparams
vars_args = {k:v if isinstance(v, str) else repr(v) for k, v in vars(args).items()}
exp, ssi, sei = hparams(vars_args, {})
writer.file_writer.add_summary(exp)
writer.file_writer.add_summary(ssi)
writer.file_writer.add_summary(sei)