-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy patheval.py
More file actions
90 lines (72 loc) · 2.33 KB
/
Copy patheval.py
File metadata and controls
90 lines (72 loc) · 2.33 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
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
import argparse
import torch
import os
from evaluator.lifeguard_evaluator import Lifeguard_Evaluator
from evaluator.lifeguard_evaluator_2Dseq import Lifeguard_Evaluator_2Dseq
from evaluator.lifeguard_evaluator_3Dseq import Lifeguard_Evaluator_3Dseq
from utils.misc import load_weight, CollateFunc_seq, CollateFunc
from config import build_default_config, build_model_config
from models import build_model
from utils.misc import parse_args
def lifeguard_eval(args, d_cfg, model):
# CHANGES ########################################################################
if d_cfg['mode']=='reference':
evaluator = Lifeguard_Evaluator(
d_cfg=d_cfg,
iou_thresh=0.5, # default 0.5
collate_fn=CollateFunc(),
)
elif d_cfg['mode']=='2D3Dseq':
evaluator = Lifeguard_Evaluator_3Dseq(
d_cfg=d_cfg,
iou_thresh=0.5, # default 0.5
collate_fn=CollateFunc_seq(),
)
elif d_cfg['mode']=='2Dseq':
evaluator = Lifeguard_Evaluator_2Dseq(
d_cfg=d_cfg,
iou_thresh=0.5, # default 0.5
collate_fn=CollateFunc_seq(),
)
elif d_cfg['mode']=='3Dseq':
evaluator = Lifeguard_Evaluator_3Dseq(
d_cfg=d_cfg,
iou_thresh=0.5, # default 0.5
collate_fn=CollateFunc_seq(),
)
else:
raise NotImplementedError
# CHANGES ########################################################################
# evaluate
evaluator.evaluate_frame_map(model, show_pr_curve=False) # default: True
if __name__ == '__main__':
args = parse_args()
num_classes = 2
# config
d_cfg = build_default_config(args)
m_cfg = build_model_config(args)
# cuda
if d_cfg['cuda']:
print('use cuda')
device = torch.device("cuda")
else:
device = torch.device("cpu")
# build model
model, _ = build_model(
args=args,
d_cfg=d_cfg,
m_cfg=m_cfg,
device=device,
num_classes=num_classes,
trainable=False
)
# load trained weight
model = load_weight(model=model, path_to_ckpt=os.path.join(d_cfg['save_folder'],d_cfg['weight']))
# to eval
model = model.to(device).eval()
# run
lifeguard_eval(
args=args,
d_cfg=d_cfg,
model=model,
)