(self, path, max_padding_length, split='train')
| 61 | """A class for processing a LLama text dataset""" |
| 62 | |
| 63 | def __init__(self, path, max_padding_length, split='train'): |
| 64 | args = get_args() |
| 65 | self.tokenizer = get_tokenizer() |
| 66 | self.IGNORE_INDEX = self.tokenizer.pad_token_id |
| 67 | if "-Pretrain" in args.dataset: |
| 68 | self.max_padding_length = max_padding_length + 1 |
| 69 | else: |
| 70 | self.max_padding_length = max_padding_length |
| 71 | |
| 72 | list_data_dict = load_dataset( |
| 73 | 'json', |
| 74 | data_files=path[0], |
| 75 | split=split, |
| 76 | ) |
| 77 | |
| 78 | train_dataset = list_data_dict.map( |
| 79 | self.preprocess, |
| 80 | batched=True, |
| 81 | batch_size=3000, |
| 82 | num_proc=16, |
| 83 | remove_columns=list_data_dict.column_names, |
| 84 | load_from_cache_file=False, |
| 85 | desc="Running Encoding" |
| 86 | ) |
| 87 | |
| 88 | self.input_ids = np.array(train_dataset['input_ids']) |
| 89 | self.labels = np.array(train_dataset['labels']) |
| 90 | self.samples = [] |
| 91 | |
| 92 | for inputs, labels in tqdm(zip(self.input_ids, self.labels)): |
| 93 | if self.tokenizer.eos_token_id not in inputs: continue |
| 94 | self.samples.append([inputs, labels]) |
| 95 | |
| 96 | print(' >> total number of samples: {}'.format(len(self.samples))) |
| 97 | |
| 98 | def _make_r_io_base(self, f, mode: str): |
| 99 | if not isinstance(f, io.IOBase): |
nothing calls this directly
no test coverage detected