-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathllm_generate.py
More file actions
399 lines (359 loc) · 16.3 KB
/
Copy pathllm_generate.py
File metadata and controls
399 lines (359 loc) · 16.3 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
# from litellm import completion
import anthropic
import together
import openai
# from litellm import completion
import asyncio
import json
import logging
from functools import wraps
import anthropic
import together
import openai
from enum import Enum
import random
from typing import Optional, List, Dict
class EvalErrorCode(Enum):
string_above_max_length = "string_above_max_length"
context_length_exceeded = "context_length_exceeded"
output_parse_error = "output_parse_error"
other_error = "other_error"
class EvalError(Exception):
"""Custom error type for evaluation."""
def __init__(self, code: EvalErrorCode, content: Optional[str] = None, inner_error: Optional[Exception] = None):
if inner_error is not None:
super().__init__(code, content, *inner_error.args)
else:
super().__init__(code, content)
self.code = code
self.content = content
self.inner_error = inner_error
def __repr__(self):
return \
f"""EvalError(
code = {self.code}
content = {self.content}
inner_error = (
type={type(self.inner_error)}
__str__={self.inner_error}
args=(
)
)"""
class EvalInput:
def __init__(self, messages: List[Dict], model_with_platform: str):
self._messages = messages
self.model_with_platform = model_with_platform
@property
def model_names(self):
model_names = self.model_with_platform.split("/")
assert len(model_names) == 2
return model_names
@property
def model_platform(self):
return self.model_names[0]
@property
def model_name(self):
return self.model_names[1]
@property
def messages(self):
return self._messages
def __repr__(self):
return \
f"""EvalInput(
messages = {self._messages}
model_with_platform = {self.model_with_platform}
)"""
def log_error_wrapper(func):
@wraps(func)
async def wrapper(*args, **kwargs):
try:
return await func(*args, **kwargs)
except Exception as error:
input = args[0] if args else None
logging.error(f"""
[Error on Eval] EvalInput = {input}
[Error Type]
{type(error)}
[Error]
{error}
[Error.args]
""")
pass
return wrapper
def handle_error_openai(func):
@wraps(func)
async def wrapper(*args, **kwargs):
while True:
try:
return await func(*args, **kwargs)
# sleep and cotinue errors
except (openai.RateLimitError, openai.InternalServerError, openai.APIConnectionError) as error:
match error.code:
case "rate_limit_exceeded":
if error.message.find("tokens per min (TPM)") != -1:
await asyncio.sleep(60)
else:
await asyncio.sleep(1)
continue
case "insufficient_quota":
print("Error: Insufficient quota", flush=True)
return "[ERROR] Insufficient quota"
case _:
print("Error: Unknown error", flush=True)
return "[ERROR] Unknown error"
# write EvalError into the pair.eval_error record
except openai.BadRequestError as error:
match error.code:
case "context_length_exceeded":
print("Error: Context length exceeded", flush=True)
return "[ERROR] Context length exceeded"
case "string_above_max_length":
print("Error: String above max length", flush=True)
return "[ERROR] String above max length"
case _:
print("Error: Unknown error", flush=True)
return "[ERROR] Unknown error"
# directly raise
except openai.AuthenticationError as error:
raise error
except Exception as error:
print("Error: Unknown error", flush=True)
return "[ERROR] Unknown error"
return wrapper
"""
def handle_error_deepseek(func):
@wraps(func)
async def wrapper(*args, **kwargs):
while True:
try:
return await func(*args, **kwargs)
# sleep and cotinue errors
except (openai.RateLimitError, openai.InternalServerError, openai.APIConnectionError, openai.BadRequestError) as error:
match error.code:
case "rate_limit_exceeded":
if error.message.find("tokens per min (TPM)") != -1:
await asyncio.sleep(60)
else:
await asyncio.sleep(1)
continue
case "insufficient_quota":
raise error
case "invalid_request_error":
# NOTE: Deepseek API only support 64k tokens
if error.message.find("This model's maximum context length is 65536 tokens") != -1:
print("Error: This model's maximum context length is 65536 tokens", flush=True)
return "Sorry, the model's maximum context length is 65536 tokens."
case _:
raise error
# write EvalError into the pair.eval_error record
except openai.BadRequestError as error:
match error.code:
case "context_length_exceeded":
raise EvalError(code=EvalErrorCode.context_length_exceeded, inner_error=error)
case "string_above_max_length":
raise EvalError(code=EvalErrorCode.string_above_max_length, inner_error=error)
case _:
raise error
# directly raise
# except openai.AuthenticationError as error:
# raise error
# except Exception as error:
# raise error
return wrapper
"""
def handle_error_anthropic(func):
@wraps(func)
async def wrapper(*args, **kwargs):
while True:
try:
return await func(*args, **kwargs)
except anthropic.RateLimitError as error:
# Error code: 429 - {'type': 'error', 'error': {'type': 'rate_limit_error', 'message': 'This request would exceed your organization’s rate limit of 400,000 input tokens per minute. For details, refer to: https://docs.anthropic.com/en/api/rate-limits; see the response headers for current usage. Please reduce the prompt length or the maximum tokens requested, or try again later. You may also contact sales at https://www.anthropic.com/contact-sales to discuss your options for a rate limit increase.'}}
match error.body["error"]["type"]:
case "rate_limit_error":
# sleep and cotinue errors
if error.body["error"]["message"].find("tokens per minute") != -1:
await asyncio.sleep(60)
else:
await asyncio.sleep(1)
continue
case _:
print("Error: Unknown error", flush=True)
return "[ERROR] Unknown error"
# directly raise
# except openai.AuthenticationError as error:
# raise error
except Exception as error:
print("Error: Unknown error", flush=True)
return "[ERROR] Unknown error"
return wrapper
def handle_error_together(func):
@wraps(func)
async def wrapper(*args, **kwargs):
while True:
try:
return await func(*args, **kwargs)
except together.error.ServiceUnavailableError as error:
error_message = error._message[18:]
match (type(error), error.http_status, error_message):
case (together.error.ServiceUnavailableError, 503, "The server is overloaded or not ready yet."):
await asyncio.sleep(60)
continue
case _:
print("Error: Unknown error", flush=True)
return "[ERROR] Unknown error"
except (together.error.APIError, together.error.RateLimitError) as error:
try:
error_obj = json.loads(error._message[18:])
error_type = error_obj.get("type_")
error_msg = error_obj.get("message", "")
#if error_obj["type_"] is None and error_obj["message"] == "Internal Server Error":
if error_type is None and "Internal Server Error" in error_msg:
await asyncio.sleep(60)
continue
match type(error), error.http_status, error_type:
case (together.error.RateLimitError, 429, "model_rate_limit") if "rate limit specific to this model" in error_msg:
await asyncio.sleep(60)
continue
case (together.error.APIError, 500, "server_error") if "Internal server error" in error_msg:
await asyncio.sleep(60)
continue
case (together.error.APIError, 413, "invalid_request_error") if "Request entity too large" in error_msg:
return "[ERROR] Request entity too large"
# raise EvalError(code=EvalErrorCode.string_above_max_length, inner_error=error)
# await asyncio.sleep(60)
# continue
case _:
print("Error: Unknown error", flush=True)
return "[ERROR] Unknown error"
except json.JSONDecodeError as json_error:
match type(error), error.http_status:
case (together.error.APIError, 502):
await asyncio.sleep(60)
continue
case _:
print("Error: Unknown error", flush=True)
return "[ERROR] Unknown error"
except together.error.InvalidRequestError as error:
error_obj = json.loads(error._message[18:])
match type(error), error.http_status, error_obj["type_"]:
case (together.error.InvalidRequestError, 422, "invalid_request_error") if error_obj["message"].startswith("Input validation error: `inputs` tokens + `max_new_tokens` must be <="):
# write EvalError into the pair.eval_error record
# raise EvalError(code=EvalErrorCode.context_length_exceeded, inner_error=error)
return "[ERROR] Input validation error: `inputs` tokens + `max_new_tokens` must be <="
case (together.error.InvalidRequestError, 400, "invalid_request_error") if error_obj["message"] == "Input validation error":
# raise EvalError(code=EvalErrorCode.string_above_max_length, inner_error=error)
return "[ERROR] Input validation error"
case (together.error.InvalidRequestError, 400, "invalid_request_error") if error_obj["message"].startswith("This model\'s maximum context length is "):
# raise EvalError(code=EvalErrorCode.context_length_exceeded, inner_error=error)
return "[ERROR] Exceeded maximum context length"
case (together.error.InvalidRequestError, 400, "invalid_request_error") if error_obj["message"] == "All connection attempts failed":
# sleep and cotinue errors
await asyncio.sleep(60)
continue
case (together.error.InvalidRequestError, 400, "invalid_request_error") if error_obj["message"] == "Input validation error":
# sleep and cotinue errors
# raise EvalError(code=EvalErrorCode.context_length_exceeded, inner_error=error)
return "[ERROR] Input validation error"
case _:
print("Error: Unknown error", flush=True)
return "[ERROR] Unknown error"
except together.error.Timeout as error:
await asyncio.sleep(1)
continue
except together.error.APIConnectionError as error:
print("Error: API connection error", flush=True)
await asyncio.sleep(60)
return "[ERROR] API connection error"
# directly raise
except together.error.AuthenticationError as error:
raise error
except Exception as error:
print("Error: Unknown error", flush=True)
return "[ERROR] Unknown error"
return wrapper
async def llm_generate(input: EvalInput) -> str:
match input.model_platform:
case "openai":
return await generate_openai(input=input)
case "anthropic":
return await generate_anthropic(input=input)
case _:
return await generate_together(input=input)
@handle_error_openai
@log_error_wrapper
async def generate_openai(input: EvalInput) -> str:
client = openai.AsyncClient()
if input.model_name == "o3-mini-high":
chat_completion = await client.chat.completions.create(
messages = input.messages,
model = "o3-mini",
reasoning_effort= "high",
# timeout = 60 * 5
)
else:
chat_completion = await client.chat.completions.create(
messages = input.messages,
model = input.model_name,
# timeout = 60 * 5
)
await client.close()
return chat_completion.choices[0].message.content
@handle_error_openai
@log_error_wrapper
async def generate_openai_multisample(input: EvalInput, n: int, temperature: float) -> List[str]:
client = openai.AsyncClient()
chat_completion = await client.chat.completions.create(
messages = input.messages,
model = input.model_name,
n = n,
temperature = temperature,
# timeout = 60 * 5
)
await client.close()
return [choice.message.content for choice in chat_completion.choices]
@handle_error_anthropic
@log_error_wrapper
async def generate_anthropic(input: EvalInput) -> str:
client = anthropic.AsyncClient()
message = await client.messages.create(
max_tokens = 8192,
model = input.model_name,
messages = input.messages
)
await client.close()
return "".join([block.text for block in message.content])
@handle_error_together
@log_error_wrapper
async def generate_together(input: EvalInput) -> str:
client = together.AsyncClient()
chat_completion = await client.chat.completions.create(
messages = input.messages,
model = f"{input.model_platform}/{input.model_name}"
)
# await client.close()
return chat_completion.choices[0].message.content
"""
@handle_error_deepseek
@log_error_wrapper
async def generate_deepseek(input: EvalInput) -> str:
deepseek_api_key = os.environ["DEEPSEEK_API_KEY"]
deepseek_base_url = os.environ.get("DEEPSEEK_BASE_URL", "https://api.deepseek.com")
client = openai.AsyncClient(api_key=deepseek_api_key, base_url=deepseek_base_url)
chat_completion = await client.chat.completions.create(
messages = input.messages,
model = input.model_name,
# timeout = 60 * 5
)
return chat_completion.choices[0].message.content
#@handle_aisuide_error
@log_error_wrapper
@async_wrap
def create_completion_aisuite(input: EvalInput) -> str:
client = aisuite.Client()
chat_completion = client.chat.completions.create(
messages = input.messages,
model = f"{input.model_platform}:{input.model_name}"
)
return chat_completion.choices[0].message.content
"""