Repository navigation
Expand file tree
/
Copy pathAutoTOD_eval.py
More file actions
76 lines (63 loc) · 2.6 KB
/
Copy pathAutoTOD_eval.py
File metadata and controls
76 lines (63 loc) · 2.6 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
import os
import json
from util import prompt_LLM, generate_output
from functools import partial
from user_simulator import user_simulator
def run(user_model, sys_model, max_iter=20):
for f in os.listdir("data/multiwoz/test"):
if f.endswith(".json"):
data = json.load(open("data/multiwoz/test/"+f, "r", encoding="utf-8"))
for dialog in data:
dial_id = dialog["dialogue_id"][:-5].lower()
logs = []
turn_idx = 1
sys_utter = None
finish_status = None
context = []
for turn_idx in range(1, max_iter + 1):
# print(HEADER_COLOR + '=' * HEADER_WIDTH + f' Turn {turn_idx} ' + '=' * HEADER_WIDTH + RESET_COLOR, end='\n\n')
user_utter = user_model(dialog, context, sys_utter, model=user_model)
context.append(user_utter)
# print(USER_COLOR + f'User: {user_utter}' + RESET_COLOR, end='\n')
if 'dialogue ends' in user_utter.lower():
finish_status = 'dialogue ends'
break
sys_utter = sys_model(instruction, context=context, model=sys_model)
context.append(sys_utter)
# print()
# print(AGENT_COLOR + f'AI Assistant: {sys_utter}' + RESET_COLOR, end='\n\n')
logs.append({
'turn_idx': turn_idx,
'user': user_utter,
'agent': sys_utter,
})
goals, goal_messages, dialog_refer = transform_dialog(dialog)
result = {
'cost': cost_handler.cost,
'dialog_pred': logs,
'goals': goals,
'goal_messages': goal_messages,
'dialog_refer': dialog_refer,
'finish_status': finish_status,
}
return result
# from AutoTOD
def transform_dialog(dialog):
goals = {d: dialog['goal'][d] for d in DOMAINS if dialog['goal'].get(d)}
dialog_refer = []
turn_idx = 0
for i, turn in enumerate(dialog['log']):
if i % 2 == 0:
turn_idx += 1
log_turn = {
'turn_idx': turn_idx,
'user': turn['text'],
'agent': None,
}
else:
log_turn['agent'] = turn['text']
dialog_refer.append(log_turn)
return goals, dialog['goal']['message'], dialog_refer
if __name__ == "__main__":
generate_output_fn = partial(generate_output, instruction="")
print(run(user_model=user_simulator, sys_model=generate_output_fn))