(self, config)
| 116 | raise NotImplementedError |
| 117 | |
| 118 | def init_generation_config(self, config): |
| 119 | if config.generation_config.get('num_return_sequences', None) is None: |
| 120 | config.generation_config['num_return_sequences'] = 1 |
| 121 | elif config.generation_config.get('num_return_sequences', 1) > 1 and config.generation_config.get('do_sample', False): |
| 122 | self.print_in_main('`num_return_sequences` must be 1 when using `do_sample=True`. Setting `num_return_sequences=1`') |
| 123 | config.generation_config['num_return_sequences'] = 1 |
| 124 | |
| 125 | if self.use_qa: |
| 126 | config.generation_config['repetition_penalty'] = 1.1 |
| 127 | |
| 128 | self.generation_config = config.generation_config |
| 129 | |
| 130 | if (self.tokenizer.pad_token_id is None) and (self.tokenizer.eos_token_id is not None): |
| 131 | self.print_in_main('warning: No pad_token in the config file. Setting pad_token_id to eos_token_id') |
| 132 | self.tokenizer.pad_token_id = self.tokenizer.eos_token_id |
| 133 | assert self.tokenizer.pad_token_id == self.tokenizer.eos_token_id |
| 134 | self.print_in_main(f'Generation config: {self.generation_config}') |
| 135 | |
| 136 | def init_dataloader(self, input_pth, batch_size): |
| 137 | dataset = MyDataset(input_pth) |
no test coverage detected