| 7 | |
| 8 | class PromptDatasetProcessor(object): |
| 9 | def __init__( |
| 10 | self, |
| 11 | tokenize: Callable, |
| 12 | pad_token: int, |
| 13 | keep_order: bool = False, |
| 14 | max_seq_len: int = 2048, |
| 15 | sliding_stride: int = 200, |
| 16 | discard_overlong: bool = True, |
| 17 | eod_token: int = None, |
| 18 | preprocess: Callable = None, |
| 19 | ): |
| 20 | super(PromptDatasetProcessor, self).__init__() |
| 21 | self._keep_order = keep_order |
| 22 | self._max_seq_len = max_seq_len |
| 23 | self._sliding_stride = sliding_stride |
| 24 | self._tokenize = tokenize |
| 25 | self._pad_token = pad_token |
| 26 | self._discard_overlong = discard_overlong |
| 27 | self._eod_token = eod_token |
| 28 | self._preprocess = preprocess |
| 29 | |
| 30 | self.doc_processed = 0 |
| 31 | self.doc_generated = 0 |
| 32 | self.start_time = 0 |
| 33 | |
| 34 | def pad_seq(self, prompt_tokens: List[int], code_tokens: List[int], extra: dict = None) -> Dict[str, List[int]]: |
| 35 | total_length = len(prompt_tokens) + len(code_tokens) |