CoolFace
Apppublic

Linhz/ViMNer

sourceHugging Faceupdated 2y agoView on Hugging Face
1likes
predict.py354 linesDownload Raw Back to MultimodelNER
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