CoTyle / models /lakonlab /runner /timer.py
liuhuijie
update
619344d
Raw
History Blame
1.83 kB
# modified from
# https://github.com/tjiiv-cprg/EPro-PnP/blob/42412220b641aef9e8943ceba516b3175631d370/EPro-PnP-Det/epropnp_det/utils/timer.py
"""
Copyright (C) 2010-2022 Alibaba Group Holding Limited.
"""
import numpy as np
import torch
import mmcv
from mmcv import Timer
from mmgen.utils import get_root_logger
class IterTimer:
def __init__(self, name='time', sync=True, enabled=True):
self.name = name
self.times = []
self.timer = Timer(start=False)
self.sync = sync
self.enabled = enabled
def __enter__(self):
if not self.enabled:
return
if self.sync:
torch.cuda.synchronize()
self.timer.start()
return self
def __exit__(self, type, value, traceback):
if not self.enabled:
return
if self.sync:
torch.cuda.synchronize()
self.timer_record()
self.timer._is_running = False
def timer_start(self):
self.timer.start()
def timer_record(self):
self.times.append(self.timer.since_last_check())
def print_time(self):
if not self.enabled:
return
logger = get_root_logger()
mmcv.print_log(f'Average {self.name} = {np.average(self.times):.4f}', logger=logger)
def reset(self):
self.times = []
class IterTimers(dict):
def __init__(self, *args, **kwargs):
super(IterTimers, self).__init__(*args, **kwargs)
def disable_all(self):
for timer in self.values():
timer.enabled = False
def enable_all(self):
for timer in self.values():
timer.enabled = True
def add_timer(self, name='time', sync=True, enabled=False):
self[name] = IterTimer(
name, sync=sync, enabled=enabled)
default_timers = IterTimers()