Linhz/ViMNer
1
1import torch
2import logging
3import os
4
5logger = logging.getLogger(__name__)
6from torchvision import transforms
7from PIL import Image
8
9
10class SBInputExample(object):
11 """A single training/test example for simple sequence classification."""
12
13 def __init__(self, guid, text_a, text_b, img_id, label=None, auxlabel=None):
14 """Constructs a InputExample.
15
16 Args:
17 guid: Unique id for the example.
18 text_a: string. The untokenized text of the first sequence. For single
19 sequence tasks, only this sequence must be specified.
20 text_b: (Optional) string. The untokenized text of the second sequence.
21 Only must be specified for sequence pair tasks.
22 label: (Optional) string. The label of the example. This should be
23 specified for train and dev examples, but not for test examples.
24 """
25 self.guid = guid
26 self.text_a = text_a
27 self.text_b = text_b
28 self.img_id = img_id
29 # Please note that the auxlabel is not used in SB
30 # it is just kept in order not to modify the original code
31
32
33class SBInputFeatures(object):
34 """A single set of features of data"""
35
36 def __init__(self, input_ids, input_mask, added_input_mask, segment_ids, img_feat):
37 self.input_ids = input_ids
38 self.input_mask = input_mask
39 self.added_input_mask = added_input_mask
40 self.segment_ids = segment_ids
41 self.img_feat = img_feat
42
43
44def sbreadfile(filename):
45 '''
46 Đọc dữ liệu từ tệp và trả về dưới dạng danh sách các từ và danh sách hình ảnh.
47 '''
48 print("Chuẩn bị dữ liệu từ", filename)
49 with open(filename, encoding='utf8') as f:
50 data = []
51 imgs = []
52 sentence = []
53 imgid = ''
54 for line in f:
55 line = line.strip() # Loại bỏ các dấu cách thừa ở đầu và cuối dòng
56 if line.startswith('IMGID:'):
57 imgid = line.split('IMGID:')[1] + '.jpg'
58 continue
59 if line == '':
60 if len(sentence) > 0:
61 data.append(sentence)
62 imgs.append(imgid)
63 sentence = []
64 imgid = ''
65 continue
66 word = line.split('\t')[0] # Chỉ lấy từ (không lấy nhãn)
67 sentence.append(word)
68
69 if len(sentence) > 0: # Xử lý dữ liệu cuối cùng trong tệp
70 data.append(sentence)
71 imgs.append(imgid)
72
73 print("Số lượng mẫu: " + str(len(data)))
74 print("Số lượng hình ảnh: " + str(len(imgs)))
75 return data, imgs
76
77
78class DataProcessor(object):
79 """Base class for data converters for sequence classification data sets."""
80
81 def get_train_examples(self, data_dir):
82 """Gets a collection of `InputExample`s for the train set."""
83 raise NotImplementedError()
84
85 def get_dev_examples(self, data_dir):
86 """Gets a collection of `InputExample`s for the dev set."""
87 raise NotImplementedError()
88
89 def get_labels(self):
90 """Gets the list of labels for this data set."""
91 raise NotImplementedError()
92
93 @classmethod
94 def _read_sbtsv(cls, input_file, quotechar=None):
95 """Reads a tab separated value file."""
96 return sbreadfile(input_file)
97
98
99class MNERProcessor(DataProcessor):
100 """Processor for the CoNLL-2003 data set."""
101
102 def get_train_examples(self, data_dir):
103 """See base class."""
104 data, imgs = self._read_sbtsv(os.path.join(data_dir, "train.txt"))
105 return self._create_examples(data, imgs, "train")
106
107 def get_dev_examples(self, data_dir):
108 """See base class."""
109 data, imgs = self._read_sbtsv(os.path.join(data_dir, "dev.txt"))
110 return self._create_examples(data, imgs, "dev")
111
112 def get_test_examples(self, data_dir):
113 """See base class."""
114 data, imgs = self._read_sbtsv(os.path.join(data_dir, "test.txt"))
115 return self._create_examples(data, imgs, "test")
116
117 def get_labels(self):
118 # return [
119 # "O","I-PRODUCT-AWARD",
120 # "B-MISCELLANEOUS",
121 # "B-QUANTITY-NUM",
122 # "B-ORGANIZATION-SPORTS",
123 # "B-DATETIME",
124 # "I-ADDRESS",
125 # "I-PERSON",
126 # "I-EVENT-SPORT",
127 # "B-ADDRESS",
128 # "B-EVENT-NATURAL",
129 # "I-LOCATION-GPE",
130 # "B-EVENT-GAMESHOW",
131 # "B-DATETIME-TIMERANGE",
132 # "I-QUANTITY-NUM",
133 # "I-QUANTITY-AGE",
134 # "B-EVENT-CUL",
135 # "I-QUANTITY-TEM",
136 # "I-PRODUCT-LEGAL",
137 # "I-LOCATION-STRUC",
138 # "I-ORGANIZATION",
139 # "B-PHONENUMBER",
140 # "B-IP",
141 # "B-QUANTITY-AGE",
142 # "I-DATETIME-TIME",
143 # "I-DATETIME",
144 # "B-ORGANIZATION-MED",
145 # "B-DATETIME-SET",
146 # "I-EVENT-CUL",
147 # "B-QUANTITY-DIM",
148 # "I-QUANTITY-DIM",
149 # "B-EVENT",
150 # "B-DATETIME-DATERANGE",
151 # "I-EVENT-GAMESHOW",
152 # "B-PRODUCT-AWARD",
153 # "B-LOCATION-STRUC",
154 # "B-LOCATION",
155 # "B-PRODUCT",
156 # "I-MISCELLANEOUS",
157 # "B-SKILL",
158 # "I-QUANTITY-ORD",
159 # "I-ORGANIZATION-STOCK",
160 # "I-LOCATION-GEO",
161 # "B-PERSON",
162 # "B-PRODUCT-COM",
163 # "B-PRODUCT-LEGAL",
164 # "I-LOCATION",
165 # "B-QUANTITY-TEM",
166 # "I-PRODUCT",
167 # "B-QUANTITY-CUR",
168 # "I-QUANTITY-CUR",
169 # "B-LOCATION-GPE",
170 # "I-PHONENUMBER",
171 # "I-ORGANIZATION-MED",
172 # "I-EVENT-NATURAL",
173 # "I-EMAIL",
174 # "B-ORGANIZATION",
175 # "B-URL",
176 # "I-DATETIME-TIMERANGE",
177 # "I-QUANTITY",
178 # "I-IP",
179 # "B-EVENT-SPORT",
180 # "B-PERSONTYPE",
181 # "B-QUANTITY-PER",
182 # "I-QUANTITY-PER",
183 # "I-PRODUCT-COM",
184 # "I-DATETIME-DURATION",
185 # "B-LOCATION-GPE-GEO",
186 # "B-QUANTITY-ORD",
187 # "I-EVENT",
188 # "B-DATETIME-TIME",
189 # "B-QUANTITY",
190 # "I-DATETIME-SET",
191 # "I-LOCATION-GPE-GEO",
192 # "B-ORGANIZATION-STOCK",
193 # "I-ORGANIZATION-SPORTS",
194 # "I-SKILL",
195 # "I-URL",
196 # "B-DATETIME-DURATION",
197 # "I-DATETIME-DATE",
198 # "I-PERSONTYPE",
199 # "B-DATETIME-DATE",
200 # "I-DATETIME-DATERANGE",
201 # "B-LOCATION-GEO",
202 # "B-EMAIL","X","<s>", "</s>"]
203
204 # vlsp2016
205 return [
206 "I-LOC", "B-MISC",
207 "I-PER",
208 "I-ORG",
209 "B-LOC",
210 "I-MISC",
211 "B-ORG",
212 "O",
213 "B-PER",
214 "X",
215 "<s>",
216 "</s>"]
217
218 # vlsp2018
219 # return [
220 # "O","I-ORGANIZATION",
221 # "B-ORGANIZATION",
222 # "I-LOCATION",
223 # "B-MISCELLANEOUS",
224 # "I-PERSON",
225 # "B-PERSON",
226 # "I-MISCELLANEOUS",
227 # "B-LOCATION",
228 # "X",
229 # "<s>",
230 # "</s>"]
231
232 def get_auxlabels(self):
233 return ["O", "B", "I", "X", "<s>", "</s>"]
234
235 def get_start_label_id(self):
236 label_list = self.get_labels()
237 label_map = {label: i for i, label in enumerate(label_list, 1)}
238 return label_map['<s>']
239
240 def get_stop_label_id(self):
241 label_list = self.get_labels()
242 label_map = {label: i for i, label in enumerate(label_list, 1)}
243 return label_map['</s>']
244
245 def _create_examples(self, lines, imgs, set_type):
246 examples = []
247 for i, (sentence) in enumerate(lines):
248 guid = "%s-%s" % (set_type, i)
249 text_a = ' '.join(sentence)
250 text_b = None
251 img_id = imgs[i]
252 examples.append(
253 SBInputExample(guid=guid, text_a=text_a, text_b=text_b, img_id=img_id))
254 return examples
255
256
257def create_examples(lines, imgs, set_type):
258 examples = []
259 for i, (sentence) in enumerate(lines):
260 guid = "%s-%s" % (set_type, i)
261 text_a = ' '.join(sentence)
262 text_b = None
263 img_id = imgs[i]
264 examples.append(
265 SBInputExample(guid=guid, text_a=text_a, text_b=text_b, img_id=img_id))
266 return examples
267
268
269def get_test_examples_predict(data_dir):
270 """See base class."""
271 data, imgs = sbreadfile(os.path.join(data_dir, "test.txt"))
272 return create_examples(data, imgs, "test")
273
274
275def image_process(image_path, transform):
276 image = Image.open(image_path).convert('RGB')
277 image = transform(image)
278 return image
279
280
281def convert_mm_examples_to_features_predict(examples,
282 max_seq_length, tokenizer, crop_size, path_img):
283 features = []
284 count = 0
285
286 transform = transforms.Compose([
287 transforms.Resize([256, 256]),
288 transforms.RandomCrop(crop_size), # args.crop_size, by default it is set to be 224
289 transforms.RandomHorizontalFlip(),
290 transforms.ToTensor(),
291 transforms.Normalize((0.485, 0.456, 0.406),
292 (0.229, 0.224, 0.225))])
293
294 for (ex_index, example) in enumerate(examples):
295 textlist = example.text_a.split(' ')
296 tokens = []
297 for i, word in enumerate(textlist):
298 token = tokenizer.tokenize(word)
299 tokens.extend(token)
300 if len(tokens) >= max_seq_length - 1:
301 tokens = tokens[0:(max_seq_length - 2)]
302 ntokens = []
303 segment_ids = []
304 ntokens.append("<s>")
305 segment_ids.append(0)
306 for i, token in enumerate(tokens):
307 ntokens.append(token)
308 segment_ids.append(0)
309 ntokens.append("</s>")
310 segment_ids.append(0)
311 input_ids = tokenizer.convert_tokens_to_ids(ntokens)
312 input_mask = [1] * len(input_ids)
313 added_input_mask = [1] * (len(input_ids) + 49) # 1 or 49 is for encoding regional image representations
314
315 while len(input_ids) < max_seq_length:
316 input_ids.append(0)
317 input_mask.append(0)
318 added_input_mask.append(0)
319 segment_ids.append(0)
320
321 assert len(input_ids) == max_seq_length
322 assert len(input_mask) == max_seq_length
323 assert len(segment_ids) == max_seq_length
324
325 image_name = example.img_id
326 image_path = os.path.join(path_img, image_name)
327
328 if not os.path.exists(image_path):
329 if 'NaN' not in image_path:
330 print(image_path)
331 try:
332 image = image_process(image_path, transform)
333 except:
334 count += 1
335 image_path_fail = os.path.join(path_img, 'background.jpg')
336 image = image_process(image_path_fail, transform)
337
338 else:
339 if ex_index < 1:
340 logger.info("*** Example ***")
341 logger.info("guid: %s" % (example.guid))
342 logger.info("tokens: %s" % " ".join(
343 [str(x) for x in tokens]))
344 logger.info("input_ids: %s" % " ".join([str(x) for x in input_ids]))
345 logger.info("input_mask: %s" % " ".join([str(x) for x in input_mask]))
346 logger.info(
347 "segment_ids: %s" % " ".join([str(x) for x in segment_ids]))
348
349 features.append(
350 SBInputFeatures(input_ids=input_ids, input_mask=input_mask, added_input_mask=added_input_mask,
351 segment_ids=segment_ids, img_feat=image))
352
353 print('the number of problematic samples: ' + str(count))
354 return features