| 6 | |
| 7 | |
| 8 | class GeminiProEvaluator(Evaluator): |
| 9 | |
| 10 | def __init__(self, pretrained_model_name_or_path, api_key, temperature=0): |
| 11 | super(GeminiProEvaluator, self).__init__() |
| 12 | |
| 13 | gemini.configure( |
| 14 | api_key=api_key, |
| 15 | transport='rest' |
| 16 | ) |
| 17 | self.model = gemini.GenerativeModel(pretrained_model_name_or_path) |
| 18 | self.temperature = temperature |
| 19 | self.candidate_count = 1 |
| 20 | self.max_output_tokens = 2048 |
| 21 | |
| 22 | def format_prompt(self, prompt): |
| 23 | return prompt |
| 24 | |
| 25 | @retry.Retry() |
| 26 | def generate_text(self, prompt): |
| 27 | prompt = self.format_prompt(prompt) |
| 28 | response = self.model.generate_content( |
| 29 | contents=prompt, |
| 30 | generation_config=gemini.types.GenerationConfig( |
| 31 | temperature=self.temperature, |
| 32 | candidate_count=self.candidate_count, |
| 33 | max_output_tokens=self.max_output_tokens |
| 34 | ) |
| 35 | ) |
| 36 | |
| 37 | return response.text.strip() if response.text is not None else response.text |
| 38 | |
| 39 | @retry.Retry() |
| 40 | def count_tokens(self, string): |
| 41 | return gemini.count_message_tokens(prompt=string)['token_count'] |