| 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) |
| 36 | assert total_length <= self._max_seq_len, f"padding sequence: {total_length} > {self._max_seq_len}" |
| 37 | pad_len = self._max_seq_len - total_length |
| 38 | input_ids = prompt_tokens + code_tokens + [self._pad_token] * pad_len |
| 39 | attention_mask = [1] * len(prompt_tokens) + [1] * len(code_tokens) + [0] * pad_len |
| 40 | labels = [-100] * len(prompt_tokens) + code_tokens + [-100] * pad_len |
| 41 | |
| 42 | return { |
| 43 | "input_ids": input_ids, |
| 44 | "attention_mask": attention_mask, |
| 45 | "labels": labels, |
| 46 | } |
| 47 | |
| 48 | def process_sample(self, sample: PromptSample) -> Iterable[Dict[str, List[int]]]: |
| 49 | """ |