-
Notifications
You must be signed in to change notification settings - Fork 27
/
monitor.py
38 lines (31 loc) · 1.21 KB
/
monitor.py
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
# monitor.py
from collections import OrderedDict
class Monitor:
def __init__(self, smoothing=True, smoothness=0.7):
self.keys = []
self.losses = {}
self.smoothing = smoothing
self.smoothness = smoothness
self.num = 0
def register(self, modules):
for m in modules:
self.keys.append(m)
self.losses[m] = 0
def reset(self):
self.num = 0
for key, value in self.losses.items():
value = 0
def update(self, modules, batch_size):
if self.smoothing == False:
for key, value in modules.items():
self.losses[key] = (self.losses[key]*self.num + value*batch_size)/(self.num + batch_size)
if self.smoothing == True:
for key, value in modules.items():
temp = (self.losses[key]*self.num + value*batch_size)/(self.num + batch_size)
self.losses[key] = self.losses[key]*self.smoothness + value*(1-self.smoothness)
self.num += batch_size
def getvalues(self, key=None):
if key != None:
return self.losses[key]
if key == None:
return OrderedDict([(key,self.losses[key]) for key in self.keys])