MCPcopy Create free account
hub / github.com/PaddlePaddle/Research / ClassifyReader

Class ClassifyReader

NLP/UNIMO/src/reader/visual_entailment_reader.py:30–237  ·  view source on GitHub ↗

ClassifyReader

Source from the content-addressed store, hash-verified

28
29
30class ClassifyReader(object):
31 """ClassifyReader"""
32 def __init__(self,
33 filelist,
34 max_seq_len,
35 tokenizer):
36
37 self.files = open(filelist).readlines()
38 self.current_file_index = 0
39 self.total_file = len(self.files)
40 self.current_file = None
41 self.tot_examples_nums = 0
42
43 self.max_seq_len = max_seq_len
44 self.pad_id = tokenizer.pad_token_id
45 self.sep_id = tokenizer.sep_token_id
46
47 self.trainer_id = int(os.getenv("PADDLE_TRAINER_ID", "0"))
48 self.trainer_nums = int(os.getenv("PADDLE_TRAINERS_NUM", "1"))
49
50 def get_num_examples(self):
51 """get_num_examples"""
52 for index, file_ in enumerate(self.files):
53 self.tot_examples_nums += int(os.popen('wc -l '+file_.strip()).read().split()[0])
54 return self.tot_examples_nums
55
56 def get_progress(self):
57 """return current progress of traning data
58 """
59 return self.current_epoch, self.current_example, self.current_file_index, self.total_file, self.current_file
60
61 def parse_line(self, line, max_seq_len=512):
62 """ parse one line to token_ids, sentence_ids, pos_ids, label
63 """
64 line = line.strip('\r\n').split(";")
65
66 if len(line) == 14:
67 (image_id, data_id, label, token_ids, sent_ids, pos_ids, _, image_w, image_h, \
68 number_box, boxes, image_embeddings, _, _) = line
69 else:
70 raise ValueError("One sample have %d fields!" % len(line))
71
72 def decode_feature(base64_str, size):
73 fea_base64 = base64.b64decode(base64_str)
74 fea_decode = np.frombuffer(fea_base64, dtype=np.float32)
75 shape = size, int(fea_decode.shape[0] / size)
76 features = np.resize(fea_decode, shape)
77 return features
78
79 token_ids = [int(token) for token in token_ids.split(" ")]
80 sent_ids = [int(token) for token in sent_ids.split(" ")]
81 pos_ids = [int(token) for token in pos_ids.split(" ")]
82 assert len(token_ids) == len(sent_ids) == len(pos_ids), \
83 "[Must be true]len(token_ids) == len(sent_ids) == len(pos_ids)"
84
85 number_box = int(number_box)
86 boxes = decode_feature(boxes, number_box)
87

Callers 1

mainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected