-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathutils.py
More file actions
63 lines (49 loc) · 1.49 KB
/
Copy pathutils.py
File metadata and controls
63 lines (49 loc) · 1.49 KB
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
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
from __future__ import print_function
import numpy as np
import math
np.set_printoptions(suppress=True)
import os
import time
import torch
from importlib import reload
def makePath(path):
if not os.path.isdir(path):
os.makedirs(path)
return path
def monitor(process, multiple, second):
while True:
sum = 0
for ps in process:
if ps.is_alive():
sum += 1
if sum < multiple:
break
else:
time.sleep(second)
def save_load_name(args, name=''):
name = name if len(name) > 0 else 'default_model'
return name
def save_model(args, model, name=''):
name = save_load_name(args, name)
torch.save(model, f'./pre_trained_models/{name}.pt')
def load_model(args, name=''):
name = save_load_name(args, name)
model = torch.load(f'./pre_trained_models/{name}.pt')
return model
def get_logger(name, log_path, length):
import logging
reload(logging)
logger = logging.getLogger()
logger.setLevel(logging.INFO)
logfile = makePath(log_path) + "/Train"+str(length)+"s_" + name + ".log"
fh = logging.FileHandler(logfile, mode='w')
fh.setLevel(logging.DEBUG)
formatter = logging.Formatter("%(asctime)s - %(levelname)s: %(message)s")
fh.setFormatter(formatter)
logger.addHandler(fh)
if log_path == "./result/test":
ch = logging.StreamHandler()
ch.setLevel(logging.INFO)
ch.setFormatter(formatter)
logger.addHandler(ch)
return logger