Repository navigation
Expand file tree
/
Copy pathevaluate_v2.py
More file actions
106 lines (86 loc) · 4.13 KB
/
Copy pathevaluate_v2.py
File metadata and controls
106 lines (86 loc) · 4.13 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
from util import generate_output, schema_induction, show_token_usage
from approach_3 import instruction_generation, run_dialogues, compare_dialogue_states
import json
from functools import partial
import os
from huggingface_hub import login
if __name__ == "__main__":
instruction_model = "gemma-3-27b-it"
system_model = "gemma-3-27b-it"
user_model = "gemma-3-27b-it"
eval_model = "gemma-3-27b-it"
train_size = 100
test_size = 10
note = "ge"
regenerate = False
regenerate_dialogues = False
with open (".config/huggingface_token.txt", "r", encoding="utf-8") as f:
token = f.read().strip()
login(token)
train_prefix = f"{note}_{train_size}"
a1_instruction, a2_instruction = instruction_generation(train_size, note=train_prefix, instruction_model = instruction_model, system_model = system_model, user_model = user_model, regenerate = regenerate)
if os.path.exists("outputs/"+train_prefix+"/a1_"+str(test_size)+"_dialogues.json") and not regenerate_dialogues:
print("Loading dialogues...")
with open("outputs/"+train_prefix+"/a1_"+str(test_size)+"_dialogues.json", "r", encoding="utf-8") as f:
a1_dialogues = json.load(f)
with open("outputs/"+train_prefix+"/a2_"+str(test_size)+"_dialogues.json", "r", encoding="utf-8") as f:
a2_dialogues = json.load(f)
print("Dialogues loaded.")
else:
print("Generating dialogues...")
a1_agent = partial(generate_output, instruction=a1_instruction, model = system_model)
a1_dialogues = run_dialogues(a1_agent, test_size, user_model=user_model)
a2_agent = partial(generate_output, instruction=a2_instruction, model = system_model)
a2_dialogues = run_dialogues(a2_agent, test_size, user_model=user_model)
with open("outputs/"+train_prefix+"/a1_"+str(test_size)+"_dialogues.json", "w", encoding="utf-8") as f:
json.dump(a1_dialogues, f, ensure_ascii=False, indent=4)
with open("outputs/"+train_prefix+"/a2_"+str(test_size)+"_dialogues.json", "w", encoding="utf-8") as f:
json.dump(a2_dialogues, f, ensure_ascii=False, indent=4)
print("Dialogues generated and saved.")
with open("data/multiwoz_test.json", "r", encoding="utf-8") as f:
true_dialogues = json.load(f)
inform_1, inform_correct_1, request_1 = 0, 0, 0
inform_2, inform_correct_2, request_2 = 0, 0, 0
for dial_id, dialog_a1 in a1_dialogues.items():
dialogue_state_true = true_dialogues[dial_id]
information_types = dialogue_state_true["information_types"]
dialog_a2 = a2_dialogues[dial_id]
schema_a1 = schema_induction(dialog_a1, information_types)
schema_a2 = schema_induction(dialog_a2, information_types)
print("dialogue:", dialogue_state_true)
print("true state: ")
for slot, value in dialogue_state_true["slot_values"].items():
print(f"* {slot}: {','.join(value)}")
print("schema_a1:")
for goal, slots in schema_a1.items():
for slot, value in slots.items():
print(f"* {slot}: {value}")
print("schema_a2:", schema_a2)
for goal, slots in schema_a2.items():
for slot, value in slots.items():
print(f"* {slot}: {value}")
i_1, ic_1, r_1 = compare_dialogue_states(dialogue_state_true, schema_a1)
i_2, ic_2, r_2 = compare_dialogue_states(dialogue_state_true, schema_a2)
print(i_1, ic_1, r_1)
print(i_2, ic_2, r_2)
inform_1 += i_1
inform_correct_1 += ic_1
request_1 += r_1
inform_2 += i_2
inform_correct_2 += ic_2
request_2 += r_2
inform_1 /= test_size
inform_correct_1 /= test_size
request_1 /= test_size
inform_2 /= test_size
inform_correct_2 /= test_size
request_2 /= test_size
print("A1:")
print(f"Inform:\t{inform_1:.4f}")
print(f"Inform correct\t:{inform_correct_1:.4f}")
print(f"Request:{request_1:.4f}")
print("A2:")
print(f"Inform:{inform_2:.4f}")
print(f"Inform correct:{inform_correct_2:.4f}")
print(f"Request:{request_2:.4f}")
show_token_usage()