-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy patheval_detr.py
More file actions
84 lines (60 loc) · 2.94 KB
/
Copy patheval_detr.py
File metadata and controls
84 lines (60 loc) · 2.94 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
import argparse
import torch
import torchvision
from tqdm import tqdm
from transformers import DetrForObjectDetection, DetrImageProcessor
from datasets.coco_eval import CocoEvaluator
class CocoDetection(torchvision.datasets.CocoDetection):
def __init__(self, img_folder, ann_file, feature_extractor):
super(CocoDetection, self).__init__(img_folder, ann_file)
self.feature_extractor = feature_extractor
def __getitem__(self, idx):
img, target = super(CocoDetection, self).__getitem__(idx)
image_id = self.ids[idx]
target = {'image_id': image_id, 'annotations': target}
encoding = self.feature_extractor(images=img, annotations=target, return_tensors="pt")
pixel_values = encoding["pixel_values"].squeeze()
target = encoding["labels"][0]
return pixel_values, target
def collate_fn(batch):
pixel_values = [item[0] for item in batch]
encoding = feature_extractor.pad(pixel_values, return_tensors="pt")
labels = [item[1] for item in batch]
batch = {}
batch['pixel_values'] = encoding['pixel_values']
batch['pixel_mask'] = encoding['pixel_mask']
batch['labels'] = labels
return batch
if __name__ == '__main__':
parser = argparse.ArgumentParser()
parser.add_argument('--coco', type=str)
parser.add_argument('--model', type=str)
args = parser.parse_args()
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = DetrForObjectDetection.from_pretrained("facebook/detr-resnet-50", revision="no_timm").to(device)
print(model)
print(args.model)
model.load_state_dict(torch.load(args.model))
model.eval()
feature_extractor = DetrImageProcessor()
dataset_val = CocoDetection(img_folder=args.coco + 'val2017',
ann_file=args.coco + 'annotations/instances_val2017.json',
feature_extractor=feature_extractor)
dataloader = torch.utils.data.DataLoader(dataset_val, 2, collate_fn=collate_fn)
base_ds = dataset_val.coco
iou_types = ['bbox']
coco_evaluator = CocoEvaluator(base_ds, iou_types)
print("Running evaluation...")
with torch.no_grad():
for idx, batch in enumerate(tqdm(dataloader)):
pixel_values = batch["pixel_values"].to(device)
pixel_mask = batch["pixel_mask"].to(device)
labels = [{k: v.to(device) for k, v in t.items()} for t in batch["labels"]]
outputs = model(pixel_values=pixel_values, pixel_mask=pixel_mask)
orig_target_sizes = torch.stack([target["orig_size"] for target in labels], dim=0)
results = feature_extractor.post_process_object_detection(outputs, 0, orig_target_sizes)
res = {target['image_id'].item(): output for target, output in zip(labels, results)}
coco_evaluator.update(res)
coco_evaluator.synchronize_between_processes()
coco_evaluator.accumulate()
print(coco_evaluator.summarize())