| 102 | self.batch_count = math.ceil(len(instances)/args.batch_size) |
| 103 | |
| 104 | def get_batch(self, index): |
| 105 | sentence_ids = [] |
| 106 | sentence_tokens = [] |
| 107 | lengths = [] |
| 108 | masks = [] |
| 109 | aspect_tags = [] |
| 110 | opinion_tags = [] |
| 111 | tags = [] |
| 112 | |
| 113 | for i in range(index * self.args.batch_size, |
| 114 | min((index + 1) * self.args.batch_size, len(self.instances))): |
| 115 | sentence_ids.append(self.instances[i].id) |
| 116 | sentence_tokens.append(self.instances[i].sentence_tokens) |
| 117 | lengths.append(self.instances[i].length) |
| 118 | masks.append(self.instances[i].mask) |
| 119 | aspect_tags.append(self.instances[i].aspect_tags) |
| 120 | opinion_tags.append(self.instances[i].opinion_tags) |
| 121 | tags.append(self.instances[i].tags) |
| 122 | |
| 123 | indexes = list(range(len(sentence_tokens))) |
| 124 | indexes = sorted(indexes, key=lambda x: lengths[x], reverse=True) |
| 125 | |
| 126 | sentence_ids = [sentence_ids[i] for i in indexes] |
| 127 | sentence_tokens = torch.stack(sentence_tokens).to(self.args.device)[indexes] |
| 128 | lengths = torch.tensor(lengths).to(self.args.device)[indexes] |
| 129 | masks = torch.stack(masks).to(self.args.device)[indexes] |
| 130 | aspect_tags = torch.stack(aspect_tags).to(self.args.device)[indexes] |
| 131 | opinion_tags = torch.stack(opinion_tags).to(self.args.device)[indexes] |
| 132 | tags = torch.stack(tags).to(self.args.device)[indexes] |
| 133 | |
| 134 | return sentence_ids, sentence_tokens, lengths, masks, aspect_tags, opinion_tags, tags |