Repository navigation
Expand file tree
/
Copy pathevaluate_graphrag_v3.py
More file actions
246 lines (186 loc) · 5.7 KB
/
Copy pathevaluate_graphrag_v3.py
File metadata and controls
246 lines (186 loc) · 5.7 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
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
"""
作用:
评测 GraphRAG 在 benchmark.json 上的表现
修改is_correct函数
作者:
Young + ChatGPT
注意修改benchmark和保存result地址
nohup python -u evaluate_graphrag_v3.py > ./logs/eval_GraphRAG_v3.log 2>&1 &
"""
import json
import time
from tqdm import tqdm
# 导入GraphRAG系统
from main_GraphRagRetriever_v2 import ask_question
from Models.LLM_Models import build_model
# ===================================
# 加载Benchmark
# ===================================
def load_benchmark(path="./benchmark/benchmark_v4.json"):
with open(
path,
"r",
encoding="utf-8"
) as f:
benchmark = json.load(f)
return benchmark
# ===================================
# 判断答案是否正确
# ===================================
def is_correct(prediction, ground_truth):
if not prediction:
return False
if not ground_truth:
return False
keywords = (
ground_truth
.replace(",", "、")
.replace(",", "、")
.split("、")
)
keywords = [
kw.strip()
for kw in keywords
if kw.strip()
]
if len(keywords) == 0:
return False
hit_count = 0
for kw in keywords:
if kw.lower() in prediction.lower():
hit_count += 1
# 单关键词
if len(keywords) == 1:
return hit_count == 1
# 多关键词覆盖率
coverage = hit_count / len(keywords)
return coverage >= 0.6
# ===================================
# 主评测函数
# ===================================
def evaluate():
evaluate_start = time.time()
benchmark = load_benchmark()
total = len(benchmark)
correct = 0
results = []
print(
f"\n开始评测,共 {total} 题...\n"
)
for idx, sample in enumerate(tqdm(benchmark), start=1):
question = sample["question"]
gt = sample["answer"]
start_time = time.time()
try:
result = ask_question(
question,
return_detail=True
)
prediction = result["answer"]
except Exception as e:
prediction = f"ERROR: {str(e)}"
result = {
"entity_extract_time": 0,
"kg_retrieve_time": 0,
"kg_totaltime":0,
"vector_retrieve_time": 0,
"answer_time": 0,
"total_time": 0
}
end_time = time.time()
cost_time = end_time - start_time
hit = is_correct(
prediction,
gt
)
if hit:
correct += 1
current_acc = correct / idx
results.append(
{
"question": question,
"ground_truth": gt,
"prediction": prediction,
"correct": hit,
"time": round(cost_time, 2),
"entity_extract_time":result["entity_extract_time"],
"kg_retrieve_time":result["kg_retrieve_time"],
"kg_totaltime": result["kg_totaltime"],
"vector_retrieve_time":result["vector_retrieve_time"],
"answer_time":result["answer_time"],
"total_time":result["total_time"]
}
)
print("\n" + "="*60)
print(f"题目 {idx}/{total}")
print("问题:", question)
print("标准答案:", gt)
print("模型答案:", prediction)
print("本题是否正确:", hit)
print( f"累计正确数: {correct}/{idx}")
print(f"当前准确率: {current_acc:.2%}")
print(f"本题耗时: {cost_time:.2f}s")
print("entity_extract_time:",result["entity_extract_time"])
print("kg_retrieve_time:",result["kg_retrieve_time"])
print("vector_retrieve_time:",result["vector_retrieve_time"])
print("answer_time:",result["answer_time"])
print("total_time:",result["total_time"])
print("")
print("="*60)
accuracy = correct / total
avg_time = (
sum(
item["time"]
for item in results
)
/ total
)
avg_entity_extract = sum(x["entity_extract_time"] for x in results) / total
avg_kg_retrieve = sum(x["kg_retrieve_time"] for x in results) / total
avg_vector_retrieve = sum(x["vector_retrieve_time"] for x in results) / total
avg_answer = sum(x["answer_time"] for x in results) / total
avg_total = sum(x["total_time"] for x in results) / total
avg_kg_total = sum(x["kg_totaltime"] for x in results) / total
print("\n====================")
print("评测完成")
print("====================")
print(f"总题数: {total}")
print(f"正确数: {correct}")
print(f"准确率: {accuracy:.2%}")
print(f"平均耗时: {avg_time:.2f}s")
print(f"avg_entity_extract: {avg_entity_extract:.2f}")
print(f"avg_kg_retrieve: {avg_kg_retrieve:.2f}")
print(f"avg_kg_totaltime: {avg_kg_total:.2f}")
print(f"avg_vector_retrieve: {avg_vector_retrieve:.2f}")
print(f"avg_answer: {avg_answer:.2f}")
print(f"avg_total: {avg_total:.2f}")
# 保存结果
with open(
"./eval_graphrag_result_v4.json",
"w",
encoding="utf-8"
) as f:
json.dump(
results,
f,
ensure_ascii=False,
indent=2
)
print(
"\n详细结果已保存到:"
"./eval_graphrag_result_v4.json"
)
evaluate_end = time.time()
total_time = (
evaluate_end
- evaluate_start
)
print(
f"总评测耗时: "
f"{total_time:.2f}s"
)
# ===================================
# 启动
# ===================================
if __name__ == "__main__":
evaluate()