(
self,
prefix_allowed_tokens_fn: Callable[[int, torch.Tensor], List[int]],
num_beams: int,
base_model: str = None,
eos_token_id: int = None
)
| 24 | class ConstrainedLogitsProcessor(LogitsProcessor): |
| 25 | |
| 26 | def __init__( |
| 27 | self, |
| 28 | prefix_allowed_tokens_fn: Callable[[int, torch.Tensor], List[int]], |
| 29 | num_beams: int, |
| 30 | base_model: str = None, |
| 31 | eos_token_id: int = None |
| 32 | ): |
| 33 | self._prefix_allowed_tokens_fn = prefix_allowed_tokens_fn |
| 34 | self._num_beams = num_beams |
| 35 | self.count=0 |
| 36 | self.base_model = base_model |
| 37 | self.eos_token_id = eos_token_id |
| 38 | if self.base_model.lower().find("gpt2") > -1: |
| 39 | self.prefix_index = 4 |
| 40 | else: |
| 41 | self.prefix_index = 3 |
| 42 | |
| 43 | |
| 44 | @add_start_docstrings(LOGITS_PROCESSOR_INPUTS_DOCSTRING) |
nothing calls this directly
no outgoing calls
no test coverage detected