MCPcopy Create free account
hub / github.com/NJUNLP/GTS / get_batch

Method get_batch

code/NNModel/data.py:104–134  ·  view source on GitHub ↗
(self, index)

Source from the content-addressed store, hash-verified

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

Callers 2

trainFunction · 0.95
evalFunction · 0.45

Calls

no outgoing calls

Tested by

no test coverage detected