ClassifyReader
| 28 | |
| 29 | |
| 30 | class 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 |