MCPcopy Create free account
hub / github.com/FreedomIntelligence/CMB / init_generation_config

Method init_generation_config

workers/base.py:118–134  ·  view source on GitHub ↗
(self, config)

Source from the content-addressed store, hash-verified

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)

Callers 1

__post_init__Method · 0.95

Calls 1

print_in_mainMethod · 0.95

Tested by

no test coverage detected