Criteria to stop on the specified multi-token sequence.
| 721 | |
| 722 | |
| 723 | class MultiTokenEOSCriteria(transformers.StoppingCriteria): |
| 724 | """Criteria to stop on the specified multi-token sequence.""" |
| 725 | |
| 726 | def __init__( |
| 727 | self, |
| 728 | sequence: str, |
| 729 | tokenizer: transformers.PreTrainedTokenizer, |
| 730 | initial_decoder_input_length: int, |
| 731 | batch_size: int, |
| 732 | ): |
| 733 | self.initial_decoder_input_length = initial_decoder_input_length |
| 734 | self.done_tracker = [False] * batch_size |
| 735 | self.sequence = sequence |
| 736 | self.sequence_ids = tokenizer.encode(sequence, add_special_tokens=False) |
| 737 | self.sequence_id_len = len(self.sequence_ids) |
| 738 | self.tokenizer = tokenizer |
| 739 | |
| 740 | def __call__(self, input_ids, scores, **kwargs) -> bool: |
| 741 | # For efficiency, we compare the last n tokens where n is the number of tokens in the stop_sequence |
| 742 | lookback_ids_batch = input_ids[:, self.initial_decoder_input_length :][ |
| 743 | :, -self.sequence_id_len : |
| 744 | ] |
| 745 | |
| 746 | lookback_tokens_batch = self.tokenizer.batch_decode(lookback_ids_batch) |
| 747 | |
| 748 | for i, done in enumerate(self.done_tracker): |
| 749 | if not done: |
| 750 | self.done_tracker[i] = self.sequence in lookback_tokens_batch[i] |
| 751 | return False not in self.done_tracker |
| 752 | |
| 753 | |
| 754 | def stop_sequences_criteria( |
no outgoing calls
no test coverage detected