MCPcopy Create free account
hub / github.com/AkaliKong/MiniOneRec / __init__

Method __init__

LogitProcessor.py:26–41  ·  view source on GitHub ↗
(
        self,
        prefix_allowed_tokens_fn: Callable[[int, torch.Tensor], List[int]],
        num_beams: int,
        base_model: str = None,
        eos_token_id: int = None
    )

Source from the content-addressed store, hash-verified

24class 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)

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected