| 152 | |
| 153 | |
| 154 | class DataIterator(object): |
| 155 | def __init__(self, instances, args): |
| 156 | self.instances = instances |
| 157 | self.args = args |
| 158 | self.batch_count = math.ceil(len(instances)/args.batch_size) |
| 159 | |
| 160 | def get_batch(self, index): |
| 161 | sentence_ids = [] |
| 162 | sentences = [] |
| 163 | sens_lens = [] |
| 164 | token_ranges = [] |
| 165 | bert_tokens = [] |
| 166 | lengths = [] |
| 167 | masks = [] |
| 168 | aspect_tags = [] |
| 169 | opinion_tags = [] |
| 170 | tags = [] |
| 171 | |
| 172 | for i in range(index * self.args.batch_size, |
| 173 | min((index + 1) * self.args.batch_size, len(self.instances))): |
| 174 | sentence_ids.append(self.instances[i].id) |
| 175 | sentences.append(self.instances[i].sentence) |
| 176 | sens_lens.append(self.instances[i].sen_length) |
| 177 | token_ranges.append(self.instances[i].token_range) |
| 178 | bert_tokens.append(self.instances[i].bert_tokens_padding) |
| 179 | lengths.append(self.instances[i].length) |
| 180 | masks.append(self.instances[i].mask) |
| 181 | aspect_tags.append(self.instances[i].aspect_tags) |
| 182 | opinion_tags.append(self.instances[i].opinion_tags) |
| 183 | tags.append(self.instances[i].tags) |
| 184 | |
| 185 | bert_tokens = torch.stack(bert_tokens).to(self.args.device) |
| 186 | lengths = torch.tensor(lengths).to(self.args.device) |
| 187 | masks = torch.stack(masks).to(self.args.device) |
| 188 | aspect_tags = torch.stack(aspect_tags).to(self.args.device) |
| 189 | opinion_tags = torch.stack(opinion_tags).to(self.args.device) |
| 190 | tags = torch.stack(tags).to(self.args.device) |
| 191 | return sentence_ids, bert_tokens, lengths, masks, sens_lens, token_ranges, aspect_tags, tags |