| 6 | |
| 7 | |
| 8 | class ChatGPTEvaluator(Evaluator): |
| 9 | |
| 10 | def __init__(self, pretrained_model_name_or_path, api_key, temperature=0): |
| 11 | super(ChatGPTEvaluator, self).__init__() |
| 12 | |
| 13 | openai.api_key = api_key |
| 14 | self.model = pretrained_model_name_or_path |
| 15 | self.temperature = temperature |
| 16 | |
| 17 | def format_prompt(self, prompt): |
| 18 | return [ |
| 19 | { |
| 20 | 'role': 'user', |
| 21 | 'content': prompt |
| 22 | } |
| 23 | ] |
| 24 | |
| 25 | @backoff.on_exception(backoff.expo, openai.error.RateLimitError) |
| 26 | def generate_text(self, prompt): |
| 27 | prompt = self.format_prompt(prompt) |
| 28 | response = openai.ChatCompletion.create( |
| 29 | model=self.model, |
| 30 | messages=prompt, |
| 31 | temperature=self.temperature |
| 32 | ) |
| 33 | |
| 34 | return response['choices'][0]['message']['content'].strip() |
| 35 | |
| 36 | @backoff.on_exception(backoff.expo, openai.error.RateLimitError) |
| 37 | def count_tokens(self, prompt): |
| 38 | try: |
| 39 | encoding = tiktoken.encoding_for_model(self.model) |
| 40 | except KeyError: |
| 41 | encoding = tiktoken.get_encoding('cl100k_base') |
| 42 | num_tokens = len(encoding.encode(prompt)) |
| 43 | |
| 44 | return num_tokens |