diff --git a/flexeval/core/language_model/openai_api.py b/flexeval/core/language_model/openai_api.py index 5aa2594f..c6aa2b97 100644 --- a/flexeval/core/language_model/openai_api.py +++ b/flexeval/core/language_model/openai_api.py @@ -226,7 +226,7 @@ def _batch_complete_text( LMOutput( text=res.choices[0].message.content, reasoning_text=get_reasoning_text(res.choices[0].message), - finish_reason=res.choices[0].finish_reason, + finish_reason="empty" if res is self.empty_response else res.choices[0].finish_reason, ) for res in api_responses ] @@ -246,7 +246,7 @@ def _batch_generate_chat_response( LMOutput( text=res.choices[0].message.content, reasoning_text=get_reasoning_text(res.choices[0].message), - finish_reason=res.choices[0].finish_reason, + finish_reason="empty" if res is self.empty_response else res.choices[0].finish_reason, tool_calls=[tool_call.to_dict() for tool_call in res.choices[0].message.tool_calls] if res.choices[0].message.tool_calls else None, @@ -295,17 +295,21 @@ def _batch_compute_chat_log_probs( ) log_probs = [] - top_logprobs_list = [res.choices[0].logprobs.content[0].top_logprobs for res in api_responses] + top_logprobs_list = [ + None if res is self.empty_response else res.choices[0].logprobs.content[0].top_logprobs + for res in api_responses + ] for index, prompt in enumerate(prompt_list): target_token = response_contents[index] index_in_unique = unique_prompt_list.index(prompt) - log_prob = None # if target token not in top_logprobs, return None for log_prob of the token + log_prob = None # if target token not in top_logprobs, or the request errored, return None top_logprobs = top_logprobs_list[index_in_unique] - for token_logprob in top_logprobs: - if token_logprob.token == target_token: - log_prob = token_logprob.logprob - break + if top_logprobs is not None: + for token_logprob in top_logprobs: + if token_logprob.token == target_token: + log_prob = token_logprob.logprob + break log_probs.append(log_prob) return log_probs @@ -450,7 +454,12 @@ def _batch_complete_text( **kwargs, ) - return [LMOutput(text=res.choices[0].text, finish_reason=res.choices[0].finish_reason) for res in api_responses] + return [ + LMOutput(text="", finish_reason="empty") + if res is self.empty_response + else LMOutput(text=res.choices[0].text, finish_reason=res.choices[0].finish_reason) + for res in api_responses + ] def __repr__(self) -> str: return f"{self.__class__.__name__}(model={self.model})" diff --git a/flexeval/core/language_model/openai_batch_api.py b/flexeval/core/language_model/openai_batch_api.py index 0ca208b0..6f868edc 100644 --- a/flexeval/core/language_model/openai_batch_api.py +++ b/flexeval/core/language_model/openai_batch_api.py @@ -257,7 +257,12 @@ def _batch_complete_text( **kwargs, ) return [ - LMOutput(text=res["choices"][0]["message"]["content"], finish_reason=res["choices"][0]["finish_reason"]) + LMOutput(text="", finish_reason="empty") + if isinstance(res, str) + else LMOutput( + text=res["choices"][0]["message"]["content"], + finish_reason=res["choices"][0]["finish_reason"], + ) for res in api_responses ] @@ -273,7 +278,9 @@ def _batch_generate_chat_response( **kwargs, ) return [ - LMOutput( + LMOutput(text="", finish_reason="empty") + if isinstance(res, str) + else LMOutput( text=res["choices"][0]["message"]["content"], finish_reason=res["choices"][0]["finish_reason"], tool_calls=res["choices"][0]["message"].get("tool_calls", None), @@ -319,17 +326,21 @@ def _batch_compute_chat_log_probs( ) log_probs = [] - top_logprobs_list = [res["choices"][0]["logprobs"]["content"][0]["top_logprobs"] for res in api_responses] + top_logprobs_list = [ + None if isinstance(res, str) else res["choices"][0]["logprobs"]["content"][0]["top_logprobs"] + for res in api_responses + ] for index, prompt in enumerate(prompt_list): target_token = response_contents[index] index_in_unique = unique_prompt_list.index(prompt) - log_prob = None # if target token not in top_logprobs, return None for log_prob of the token + log_prob = None # if target token not in top_logprobs, or the request errored, return None top_logprobs = top_logprobs_list[index_in_unique] - for token_logprob in top_logprobs: - if token_logprob["token"] == target_token: - log_prob = token_logprob["logprob"] - break + if top_logprobs is not None: + for token_logprob in top_logprobs: + if token_logprob["token"] == target_token: + log_prob = token_logprob["logprob"] + break log_probs.append(log_prob) return log_probs