Skip to content

Commit 195bb32

Browse files
committed
Added json fix in factool and factcheckgpt
1 parent 0fe75ae commit 195bb32

2 files changed

Lines changed: 24 additions & 25 deletions

File tree

src/openfactcheck/solvers/webservice/factcheckgpt_utils/openai_api.py

Lines changed: 18 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -6,21 +6,25 @@
66
client = None
77

88

9+
def _json_fix(output):
10+
return output.replace("```json\n", "").replace("```", "")
11+
12+
913
def init_client():
1014
global client
1115
if client is None:
12-
if openai.api_key is None and 'OPENAI_API_KEY' not in os.environ:
16+
if openai.api_key is None and "OPENAI_API_KEY" not in os.environ:
1317
print("openai_key not presented, delay to initialize.")
1418
return
1519
client = OpenAI()
1620

1721

1822
def request(
19-
user_inputs,
20-
model,
21-
system_role,
22-
temperature=1.0,
23-
return_all=False,
23+
user_inputs,
24+
model,
25+
system_role,
26+
temperature=1.0,
27+
return_all=False,
2428
):
2529
init_client()
2630

@@ -29,41 +33,31 @@ def request(
2933
elif type(user_inputs) == list:
3034
if all([type(x) == str for x in user_inputs]):
3135
chat_histories = [
32-
{
33-
"role": "user" if i % 2 == 0 else "assistant", "content": x
34-
} for i, x in enumerate(user_inputs)
36+
{"role": "user" if i % 2 == 0 else "assistant", "content": x} for i, x in enumerate(user_inputs)
3537
]
3638
elif all([type(x) == dict for x in user_inputs]):
3739
chat_histories = user_inputs
3840
else:
3941
raise ValueError("Invalid input for OpenAI API calling")
4042
else:
4143
raise ValueError("Invalid input for OpenAI API calling")
42-
4344

4445
messages = [{"role": "system", "content": system_role}] + chat_histories
4546

46-
response = client.chat.completions.create(
47-
model=model,
48-
messages=messages,
49-
temperature=temperature
50-
)
47+
response = client.chat.completions.create(model=model, messages=messages, temperature=temperature)
48+
49+
# Fix the json format
50+
response = _json_fix(response)
51+
5152
if return_all:
5253
return response
53-
response_str = ''
54+
response_str = ""
5455
for choice in response.choices:
5556
response_str += choice.message.content
5657
return response_str
5758

5859

59-
def gpt(
60-
user_inputs,
61-
model,
62-
system_role,
63-
temperature=1.0,
64-
num_retries=3,
65-
waiting=1
66-
):
60+
def gpt(user_inputs, model, system_role, temperature=1.0, num_retries=3, waiting=1):
6761
response = None
6862
for _ in range(num_retries):
6963
try:

src/openfactcheck/solvers/webservice/factool_utils/chat_api.py

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -72,6 +72,9 @@ def extract_dict_from_string(self, input_string):
7272
else:
7373
return None
7474

75+
def _json_fix(self, output):
76+
return output.replace("```json\n", "").replace("```", "")
77+
7578
def _boolean_fix(self, output):
7679
return output.replace("true", "True").replace("false", "False")
7780

@@ -166,7 +169,9 @@ def run(self, messages_list, expected_type):
166169
)
167170

168171
preds = [
169-
self._type_check(self._boolean_fix(prediction.choices[0].message.content), expected_type)
172+
self._type_check(
173+
self._boolean_fix(self._json_fix(prediction.choices[0].message.content)), expected_type
174+
)
170175
if prediction is not None
171176
else None
172177
for prediction in predictions

0 commit comments

Comments
 (0)