| 87 | |
| 88 | class LabelDatasetProcessor(object): |
| 89 | def __init__( |
| 90 | self, |
| 91 | tokenize: Callable, |
| 92 | pad_token: int, |
| 93 | keep_order: bool = False, |
| 94 | max_seq_len: int = 2048, |
| 95 | sliding_stride: int = 200, |
| 96 | discard_overlong: bool = True, |
| 97 | eod_token: int = None, |
| 98 | preprocess: Callable = None, |
| 99 | ): |
| 100 | super(LabelDatasetProcessor, self).__init__() |
| 101 | self._keep_order = keep_order |
| 102 | self._max_seq_len = max_seq_len |
| 103 | self._sliding_stride = sliding_stride |
| 104 | self._tokenize = tokenize |
| 105 | self._pad_token = pad_token |
| 106 | self._discard_overlong = discard_overlong |
| 107 | self._eod_token = eod_token |
| 108 | self._preprocess = preprocess |
| 109 | |
| 110 | self.doc_processed = 0 |
| 111 | self.doc_generated = 0 |
| 112 | self.start_time = 0 |
| 113 | |
| 114 | def pad_seq(self, prompt_tokens: List[int], label: int, extra: dict = None) -> Dict[str, List[int]]: |
| 115 | total_length = len(prompt_tokens) |