| 245 | |
| 246 | |
| 247 | class 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] |
no outgoing calls
no test coverage detected