| 7 | |
| 8 | |
| 9 | class QwenChatEvaluator(Evaluator): |
| 10 | |
| 11 | def __init__(self, pretrained_model_name_or_path, api_key, temperature=0): |
| 12 | super(QwenChatEvaluator, self).__init__() |
| 13 | |
| 14 | dashscope.api_key = api_key |
| 15 | self.model = pretrained_model_name_or_path |
| 16 | self.temperature = temperature |
| 17 | |
| 18 | def format_prompt(self, prompt): |
| 19 | return [ |
| 20 | { |
| 21 | 'role': 'user', |
| 22 | 'content': prompt |
| 23 | } |
| 24 | ] |
| 25 | |
| 26 | @retry(wait=wait_random_exponential(min=1, max=60), stop=stop_after_attempt(6)) |
| 27 | def generate_text(self, prompt): |
| 28 | prompt = self.format_prompt(prompt) |
| 29 | response = dashscope.Generation.call( |
| 30 | self.model, |
| 31 | messages=prompt, |
| 32 | temperature=self.temperature, |
| 33 | result_format='message' |
| 34 | ) |
| 35 | if response.status_code == HTTPStatus.OK: |
| 36 | return response['output']['choices'][0]['message']['content'].strip() |
| 37 | else: |
| 38 | print('Request id: %s, Status code: %s, error code: %s, error message: %s' % (response.request_id, response.status_code, response.code, response.message)) |
| 39 | return '' |
| 40 | |
| 41 | @retry(wait=wait_random_exponential(min=1, max=60), stop=stop_after_attempt(6)) |
| 42 | def count_tokens(self, prompt): |
| 43 | response = dashscope.Tokenization.call( |
| 44 | 'qwen-14b-chat', |
| 45 | prompt=prompt, |
| 46 | ) |
| 47 | if response.status_code == HTTPStatus.OK: |
| 48 | return response['usage']['input_tokens'] |
| 49 | else: |
| 50 | print('Failed request_id: %s, status_code: %s, code: %s, message:%s' % (response.request_id, response.status_code, response.code, response.message)) |
| 51 | return 0 |