| 693 | |
| 694 | @dataclass |
| 695 | class SamplingConfig: |
| 696 | end_id: int |
| 697 | pad_id: int |
| 698 | |
| 699 | max_new_tokens: int = field(default=20) |
| 700 | num_beams: int = field(default=1) |
| 701 | num_return_sequences: Optional[int] = field(default=None) |
| 702 | max_attention_window_size: Optional[int] = field(default=None) |
| 703 | sink_token_length: Optional[int] = field(default=None) |
| 704 | output_sequence_lengths: bool = field(default=False) |
| 705 | return_dict: bool = field(default=False) |
| 706 | stop_words_list: Optional[Union[list, np.ndarray, |
| 707 | torch.Tensor]] = field(default=None) |
| 708 | bad_words_list: Optional[Union[list, np.ndarray, |
| 709 | torch.Tensor]] = field(default=None) |
| 710 | |
| 711 | temperature: Union[float, torch.Tensor] = field(default=1.0) |
| 712 | top_k: Union[int, torch.Tensor] = field(default=1) |
| 713 | top_p: Union[float, torch.Tensor] = field(default=0.0) |
| 714 | top_p_decay: Optional[torch.Tensor] = field(default=None) # float |
| 715 | top_p_min: Optional[torch.Tensor] = field(default=None) # float |
| 716 | top_p_reset_ids: Optional[torch.Tensor] = field(default=None) # int |
| 717 | random_seed: Union[int, torch.Tensor] = field(default=None) |
| 718 | |
| 719 | length_penalty: Union[float, torch.Tensor] = field(default=1.0) |
| 720 | early_stopping: Union[int, torch.Tensor] = field(default=1) |
| 721 | repetition_penalty: Union[float, torch.Tensor] = field(default=1.0) |
| 722 | min_length: Union[int, torch.Tensor] = field(default=1) |
| 723 | presence_penalty: Union[float, torch.Tensor] = field(default=0.0) |
| 724 | frequency_penalty: Union[float, torch.Tensor] = field(default=0.0) |
| 725 | prompt_ignore_length: Union[int, torch.Tensor] = field(default=0) |
| 726 | use_beam_hyps: bool = field(default=True) |
| 727 | |
| 728 | # None here means user didn't set it, and dynamicDecodeOp.cpp take optional value |
| 729 | # The real default value is set in dynamicDecodeOp.cpp when it's None |
| 730 | beam_search_diversity_rate: Union[float, torch.Tensor] = field(init=False, |
| 731 | default=0.0) |
| 732 | output_cum_log_probs: bool = field(init=False, default=False) |
| 733 | output_log_probs: bool = field(init=False, default=False) |
| 734 | no_repeat_ngram_size: Union[int, torch.Tensor] = field(init=False, |
| 735 | default=None) |
| 736 | min_p: Union[float, torch.Tensor] = field(default=0.0) |
| 737 | |
| 738 | def update(self, **kwargs): |
| 739 | unused_kwargs = dict() |
| 740 | for key, value in kwargs.items(): |
| 741 | if hasattr(self, key): |
| 742 | setattr(self, key, value) |
| 743 | else: |
| 744 | unused_kwargs[key] = value |
| 745 | return unused_kwargs |
| 746 | |
| 747 | |
| 748 | class LogitsProcessor: |
no outgoing calls