-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy patheval_dino.py
More file actions
73 lines (51 loc) · 2.15 KB
/
Copy patheval_dino.py
File metadata and controls
73 lines (51 loc) · 2.15 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
import argparse
import torch
from tqdm import tqdm
from datasets.coco import build
from datasets.coco_eval import CocoEvaluator
from dino.dino import PostProcess
from util.build_dino import build_dino_model
from util.misc import collate_fn
def to_device(item, device):
if isinstance(item, torch.Tensor):
return item.to(device)
elif isinstance(item, list):
return [to_device(i, device) for i in item]
elif isinstance(item, dict):
return {k: to_device(v, device) for k, v in item.items()}
else:
raise NotImplementedError("Call Shilong if you use other containers! type: {}".format(type(item)))
if __name__ == '__main__':
parser = argparse.ArgumentParser()
parser.add_argument('--coco', type=str)
parser.add_argument('--model', type=str)
parser.add_argument('--root', type=str)
args = parser.parse_args()
dev = torch.device("cuda" if torch.cuda.is_available() else "cpu")
if "swin-L" in args.model:
backbone_model = "swin-L"
else:
backbone_model = "resnet50"
model = build_dino_model(args.root, backbone=backbone_model).to(dev)
print(model)
print(args.model)
model.load_state_dict(torch.load(args.model))
model.eval()
dataset_val = build("val", args.coco)
dataloader = torch.utils.data.DataLoader(dataset_val, 1, collate_fn=collate_fn)
base_ds = dataset_val.coco
iou_types = ['bbox']
coco_evaluator = CocoEvaluator(base_ds, iou_types)
postprocessors = {'bbox': PostProcess(num_select=300, nms_iou_threshold=-1)}
print("Running evaluation...")
for samples, targets in tqdm(dataloader):
samples = samples.to(dev)
targets = [{k: to_device(v, dev) for k, v in t.items()} for t in targets]
outputs = model(samples)
orig_target_sizes = torch.stack([t["orig_size"] for t in targets], dim=0)
results = postprocessors['bbox'](outputs, orig_target_sizes)
res = {target['image_id'].item(): output for target, output in zip(targets, results)}
coco_evaluator.update(res)
coco_evaluator.synchronize_between_processes()
coco_evaluator.accumulate()
print(coco_evaluator.summarize())