MCPcopy Create free account
hub / github.com/InternScience/SciReason / DatasetEncoder

Class DatasetEncoder

opencompass/openicl/icl_dataset_reader.py:247–289  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

245
246
247class DatasetEncoder(torch.utils.data.Dataset):
248
249 def __init__(self,
250 datalist: List,
251 model_name=None,
252 tokenizer=None) -> None:
253 self.datalist = datalist
254 if model_name is None and tokenizer is None:
255 raise ValueError('model_name and tokenizer could not both be None')
256 if tokenizer is not None:
257 self.tokenizer = tokenizer
258 else:
259 self.tokenizer = AutoTokenizer.from_pretrained(model_name)
260 self.tokenizer.pad_token = self.tokenizer.eos_token
261 self.tokenizer.pad_token_id = self.tokenizer.eos_token_id
262 self.tokenizer.padding_side = 'left'
263 self.encode_dataset = []
264 self.init_dataset()
265 self.datalist_length = len(self.encode_dataset)
266
267 def init_dataset(self):
268 for idx, data in enumerate(self.datalist):
269 tokenized_data = self.tokenizer.encode_plus(data,
270 truncation=True,
271 return_tensors='pt',
272 verbose=False)
273 self.encode_dataset.append({
274 'input_ids':
275 tokenized_data.input_ids[0],
276 'attention_mask':
277 tokenized_data.attention_mask[0],
278 'metadata': {
279 'id': idx,
280 'len': len(tokenized_data.input_ids[0]),
281 'text': data
282 }
283 })
284
285 def __len__(self):
286 return self.datalist_length
287
288 def __getitem__(self, idx):
289 return self.encode_dataset[idx]

Callers 2

__init__Method · 0.90
create_indexMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected