| 139 | |
| 140 | @torch.no_grad() |
| 141 | def generate(self, data_dict: dict=None, num_return_sequences: int=8, generation_config: dict=dict()) -> dict: |
| 142 | |
| 143 | net_device = next(self.parameters()).device |
| 144 | max_length = 8191 |
| 145 | output_ids = torch.ones(num_return_sequences, max_length).long().to(net_device) * self.eos_token_id |
| 146 | |
| 147 | # batch x ntokens |
| 148 | results = self.transformer.generate( |
| 149 | inputs_embeds=data_dict['prefix_embeds'], |
| 150 | max_length=max_length-1, |
| 151 | num_return_sequences=num_return_sequences, |
| 152 | bos_token_id=self.bos_token_id, |
| 153 | eos_token_id=self.eos_token_id, |
| 154 | pad_token_id=self.eos_token_id, |
| 155 | **generation_config |
| 156 | ) |
| 157 | output_ids[:, :results.shape[1]] = results |
| 158 | |
| 159 | # discard <bos> and <eos> tokens to pad tokens |
| 160 | output_ids = output_ids[:, :-1] |
| 161 | output_ids[output_ids == self.eos_token_id] = self.tokenizer.pad_id |
| 162 | |
| 163 | decoder_output = self.tokenizer.detokenize(input_ids=output_ids) |
| 164 | |
| 165 | return decoder_output |
| 166 | |
| 167 | |
| 168 | |