Repository navigation
Expand file tree
/
Copy pathrephraser.py
More file actions
133 lines (108 loc) · 4.86 KB
/
Copy pathrephraser.py
File metadata and controls
133 lines (108 loc) · 4.86 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
import json
from agentic_base import *
from agentic_utils import *
class RephraserAgent(BaseClientAgent):
def __init__(self, introduction=""):
super().__init__('rephraser')
self.introduction = introduction
def validate(self, response, file_handler=None):
if file_handler:
file_handler.write("#### Rephraser Response ####\n")
file_handler.write(response + "\n")
match = CustomJSONValidator.validate(response)
if not match:
return False, "No JSON section found in response."
try:
# second pass
CustomJSONValidator.pydantic_validate(response, RephraserResponse)
json_obj = json.loads(match.group(1))
options = json_obj.get("option")
if isinstance(options, int):
if options == 2:
return True, {'option': 2, 'subquestions': json_obj.get('subquestions')}
else:
return True, {'option': 1, 'answer': json_obj.get('answer')}
else:
return False, "Expected 'option' to be a list of integers: " + str(json_obj)
except json.JSONDecodeError:
return False, "Invalid JSON format: " + match.group(1)
def act(self, **kwargs):
print('Trying to resolve question based on existing walks...')
question = kwargs.get('question')
walks = kwargs.get('walks')
extra_help = kwargs.get('extra_help')
qas = ""
for q in extra_help:
if extra_help[q] is not None and 'answer' in extra_help[q]:
answer = extra_help[q]["answer"]
if isinstance(answer, list):
if all(isinstance(el, str) for el in answer):
answer = ', '.join(answer)
else:
try:
# try to convert...
answer = [str(el) for el in answer]
answer = ', '.join(answer)
except Exception:
answer = ''
elif not isinstance(answer, str):
try:
# try to convert...
answer = str(answer)
except Exception:
answer = ''
qas += "Question: " + q + " | Answer: " + answer + '\n'
# add extra information about subquestions that have already been answered
intro_text = self.introduction
if qas != "":
intro_text = f"""
{self.introduction}
Note that the following questions and answers are already known:
{qas}
"""
prompt = generate_rephraser_prompt(introduction=intro_text,
question=question,
walks=walks)
file_handler = kwargs.get('file_handler')
if file_handler:
file_handler.write("#### Rephraser prompt ####\n")
file_handler.write(prompt + "\n")
print('Asking to see if question can be resolved...')
iter_count = 0
augmented_prompt = prompt
while iter_count < MAX_ITER:
# output = self.ask(augmented_prompt, RephraserResponse)
output = self.ask(augmented_prompt)
valid, outcome = self.validate(output, file_handler=file_handler)
if valid:
return outcome
augmented_prompt = generate_failed_prompt(output, outcome, self.name)
iter_count += 1
# will lead to trying the whole loop again...
return {"option": 2}
def refine(self, **kwargs):
print('Trying to refine subquestion...')
question = kwargs.get('question')
subquestion = kwargs.get('subquestion')
entities = kwargs.get('entities')
prompt = refine_subquestion_prompt(introduction=self.introduction,
question=question,
subquestion=subquestion,
entities=entities)
file_handler = kwargs.get('file_handler')
if file_handler:
file_handler.write("#### Refinement prompt ####\n")
file_handler.write(prompt + "\n")
print('Asking to see if subquestion can be refined...')
# output = self.ask(prompt, RefinementResponse)
output = self.ask(prompt)
match = CustomJSONValidator.validate(output)
if not match:
return "No JSON section found in response."
try:
# second pass
CustomJSONValidator.pydantic_validate(output, RefinementResponse)
json_obj = json.loads(match.group(1))
return json_obj.get('subquestion')
except json.JSONDecodeError:
return "Invalid JSON format: " + match.group(1)