MCPcopy Create free account
hub / github.com/OpenBitSys/BitDistiller / MultiTokenEOSCriteria

Class MultiTokenEOSCriteria

test/general/lm_eval/models/huggingface.py:723–751  ·  view source on GitHub ↗

Criteria to stop on the specified multi-token sequence.

Source from the content-addressed store, hash-verified

721
722
723class 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
754def stop_sequences_criteria(

Callers 1

stop_sequences_criteriaFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected