| 54 | model_type: str = "" |
| 55 | |
| 56 | def __init__(self, **kwargs): |
| 57 | # Attributes with defaults |
| 58 | self.output_hidden_states = kwargs.pop("output_hidden_states", False) |
| 59 | self.output_attentions = kwargs.pop("output_attentions", False) |
| 60 | self.use_cache = kwargs.pop("use_cache", True) # Not used by all models |
| 61 | self.torchscript = kwargs.pop("torchscript", False) # Only used by PyTorch models |
| 62 | self.use_bfloat16 = kwargs.pop("use_bfloat16", False) |
| 63 | self.pruned_heads = kwargs.pop("pruned_heads", {}) |
| 64 | |
| 65 | # Is decoder is used in encoder-decoder models to differentiate encoder from decoder |
| 66 | self.is_encoder_decoder = kwargs.pop("is_encoder_decoder", False) |
| 67 | self.is_decoder = kwargs.pop("is_decoder", False) |
| 68 | |
| 69 | # Parameters for sequence generation |
| 70 | self.max_length = kwargs.pop("max_length", 20) |
| 71 | self.min_length = kwargs.pop("min_length", 0) |
| 72 | self.do_sample = kwargs.pop("do_sample", False) |
| 73 | self.early_stopping = kwargs.pop("early_stopping", False) |
| 74 | self.num_beams = kwargs.pop("num_beams", 1) |
| 75 | self.temperature = kwargs.pop("temperature", 1.0) |
| 76 | self.top_k = kwargs.pop("top_k", 50) |
| 77 | self.top_p = kwargs.pop("top_p", 1.0) |
| 78 | self.repetition_penalty = kwargs.pop("repetition_penalty", 1.0) |
| 79 | self.length_penalty = kwargs.pop("length_penalty", 1.0) |
| 80 | self.no_repeat_ngram_size = kwargs.pop("no_repeat_ngram_size", 0) |
| 81 | self.bad_words_ids = kwargs.pop("bad_words_ids", None) |
| 82 | self.num_return_sequences = kwargs.pop("num_return_sequences", 1) |
| 83 | |
| 84 | # Fine-tuning task arguments |
| 85 | self.architectures = kwargs.pop("architectures", None) |
| 86 | self.finetuning_task = kwargs.pop("finetuning_task", None) |
| 87 | self.id2label = kwargs.pop("id2label", None) |
| 88 | self.label2id = kwargs.pop("label2id", None) |
| 89 | if self.id2label is not None: |
| 90 | kwargs.pop("num_labels", None) |
| 91 | self.id2label = dict((int(key), value) for key, value in self.id2label.items()) |
| 92 | # Keys are always strings in JSON so convert ids to int here. |
| 93 | else: |
| 94 | self.num_labels = kwargs.pop("num_labels", 2) |
| 95 | |
| 96 | # Tokenizer arguments TODO: eventually tokenizer and models should share the same config |
| 97 | self.prefix = kwargs.pop("prefix", None) |
| 98 | self.bos_token_id = kwargs.pop("bos_token_id", None) |
| 99 | self.pad_token_id = kwargs.pop("pad_token_id", None) |
| 100 | self.eos_token_id = kwargs.pop("eos_token_id", None) |
| 101 | self.decoder_start_token_id = kwargs.pop("decoder_start_token_id", None) |
| 102 | |
| 103 | # task specific arguments |
| 104 | self.task_specific_params = kwargs.pop("task_specific_params", None) |
| 105 | |
| 106 | # TPU arguments |
| 107 | self.xla_device = kwargs.pop("xla_device", None) |
| 108 | |
| 109 | # Additional attributes without default values |
| 110 | for key, value in kwargs.items(): |
| 111 | try: |
| 112 | setattr(self, key, value) |
| 113 | except AttributeError as err: |