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

Class DataIterator

code/BertModel/data.py:154–191  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

152
153
154class 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

Callers 2

trainFunction · 0.90
testFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected