(self, prompt)
| 34 | |
| 35 | @retry(wait=wait_random_exponential(min=1, max=60), stop=stop_after_attempt(6)) |
| 36 | def generate_text(self, prompt): |
| 37 | prompt = self.format_prompt(prompt) |
| 38 | inputs = self.tokenizer(prompt, add_special_tokens=False, return_tensors='pt').to(self.model.device) |
| 39 | outputs = self.model.generate( |
| 40 | inputs['input_ids'], |
| 41 | do_sample=self.do_sample, |
| 42 | max_length=self.max_length, |
| 43 | eos_token_id=self.tokenizer.convert_tokens_to_ids('</s>') |
| 44 | ).to('cpu') |
| 45 | response = self.tokenizer.decode(outputs[0], skip_special_tokens=True) |
| 46 | |
| 47 | return response.split('Helper:')[-1].strip() |
| 48 | |
| 49 | @retry(wait=wait_random_exponential(min=1, max=60), stop=stop_after_attempt(6)) |
| 50 | def count_tokens(self, prompt): |
nothing calls this directly
no test coverage detected