Get a suffix of prompt which is no longer than num_token tokens. Args: prompt (str): Input string. num_token (int): The upper bound of token numbers. mode (str): The method of input truncation ('front', 'mid', or 'rear') Returns:
(self, prompt: str, num_token: int, mode: str)
| 432 | return len(enc.encode(prompt, disallowed_special=())) |
| 433 | |
| 434 | def _bin_trim(self, prompt: str, num_token: int, mode: str) -> str: |
| 435 | """Get a suffix of prompt which is no longer than num_token tokens. |
| 436 | |
| 437 | Args: |
| 438 | prompt (str): Input string. |
| 439 | num_token (int): The upper bound of token numbers. |
| 440 | mode (str): The method of input truncation |
| 441 | ('front', 'mid', or 'rear') |
| 442 | |
| 443 | Returns: |
| 444 | str: The trimmed prompt. |
| 445 | """ |
| 446 | token_len = self.get_token_len(prompt) |
| 447 | if token_len <= num_token: |
| 448 | return prompt |
| 449 | pattern = re.compile(r'[\u4e00-\u9fa5]') |
| 450 | if pattern.search(prompt): |
| 451 | words = list(jieba.cut(prompt, cut_all=False)) |
| 452 | sep = '' |
| 453 | else: |
| 454 | words = prompt.split(' ') |
| 455 | sep = ' ' |
| 456 | |
| 457 | l, r = 1, len(words) |
| 458 | while l + 2 < r: |
| 459 | mid = (l + r) // 2 |
| 460 | if mode == 'front': |
| 461 | cur_prompt = sep.join(words[-mid:]) |
| 462 | elif mode == 'mid': |
| 463 | cur_prompt = sep.join(words[:mid]) + sep.join(words[-mid:]) |
| 464 | elif mode == 'rear': |
| 465 | cur_prompt = sep.join(words[:mid]) |
| 466 | |
| 467 | if self.get_token_len(cur_prompt) <= num_token: |
| 468 | l = mid # noqa: E741 |
| 469 | else: |
| 470 | r = mid |
| 471 | |
| 472 | if mode == 'front': |
| 473 | prompt = sep.join(words[-l:]) |
| 474 | elif mode == 'mid': |
| 475 | prompt = sep.join(words[:l]) + sep.join(words[-l:]) |
| 476 | elif mode == 'rear': |
| 477 | prompt = sep.join(words[:l]) |
| 478 | return prompt |
| 479 | |
| 480 | def _preprocess_messages( |
| 481 | self, |
no test coverage detected