-
Notifications
You must be signed in to change notification settings - Fork 28
Expand file tree
/
Copy pathmagazine.py
More file actions
81 lines (65 loc) · 2.49 KB
/
Copy pathmagazine.py
File metadata and controls
81 lines (65 loc) · 2.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
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
from pathlib import Path
import xml.etree.ElementTree as ET
import torch
from torch_geometric.data import Data
from data.base import BaseDataset
class Magazine(BaseDataset):
labels = [
'text',
'image',
'headline',
'text-over-image',
'headline-over-image',
]
def __init__(self, split='train', transform=None):
super().__init__('magazine', split, transform)
def download(self):
super().download()
def process(self):
data_list = []
ann_dir = Path(self.raw_dir) / 'layoutdata' / 'annotations'
for xml_path in sorted(ann_dir.glob('*.xml')):
with xml_path.open() as f:
root = ET.parse(f).getroot()
W = float(root.find('size/width').text)
H = float(root.find('size/height').text)
name = root.find('filename').text
elements = root.findall('layout/element')
boxes = []
labels = []
for element in elements:
# bbox
px = list(map(float, element.get('polygon_x').split()))
py = list(map(float, element.get('polygon_y').split()))
x1, x2 = min(px), max(px)
y1, y2 = min(py), max(py)
xc = (x1 + x2) / 2.
yc = (y1 + y2) / 2.
width = x2 - x1
height = y2 - y1
b = [xc / W, yc / H,
width / W, height / H]
boxes.append(b)
# label
l = element.get('label')
labels.append(self.label2index[l])
boxes = torch.tensor(boxes, dtype=torch.float)
labels = torch.tensor(labels, dtype=torch.long)
data = Data(x=boxes, y=labels)
data.attr = {
'name': name,
'width': W,
'height': H,
'has_canvas_element': False,
}
data_list.append(data)
# shuffle with seed
generator = torch.Generator().manual_seed(0)
indices = torch.randperm(len(data_list), generator=generator)
data_list = [data_list[i] for i in indices]
# train 85% / val 5% / test 10%
N = len(data_list)
s = [int(N * .85), int(N * .90)]
torch.save(self.collate(data_list[:s[0]]), self.processed_paths[0])
torch.save(self.collate(data_list[s[0]:s[1]]), self.processed_paths[1])
torch.save(self.collate(data_list[s[1]:]), self.processed_paths[2])