-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy path05_get_downstream_utility.py
More file actions
347 lines (321 loc) · 12.2 KB
/
Copy path05_get_downstream_utility.py
File metadata and controls
347 lines (321 loc) · 12.2 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
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
"""
Get mean downstream utility for each split for each model
For non-learned (baseline and basic routing methods), load from full-recall data
For learned methods, load from learned model's report
Also supports downstream utility of RAG with a single retriever
"""
import os
import json
import argparse
import random
import statistics
random.seed(42)
def filter_data(fr_data, dataset_type: str, split: str) -> dict:
if dataset_type == "balanced":
# no filtering
return fr_data
elif dataset_type == "multi-aspect":
if split == "train":
# filter out qids that do have "multi-aspect" in question_categories
fr_data = {
qid: q_data
for qid, q_data in fr_data.items()
if "multi-aspect"
not in [
cat["category_name"]
for cat in q_data["question_categories"]
if cat["categorization_name"] == "answer-type"
]
}
elif split == "test":
# only include qids that do have "multi-aspect" in question_categories
fr_data = {
qid: q_data
for qid, q_data in fr_data.items()
if "multi-aspect"
in [
cat["category_name"]
for cat in q_data["question_categories"]
if cat["categorization_name"] == "answer-type"
]
}
elif dataset_type == "comparison":
if split == "train":
# filter out qids that do have "comparison" in question_categories
fr_data = {
qid: q_data
for qid, q_data in fr_data.items()
if "comparison"
not in [
cat["category_name"]
for cat in q_data["question_categories"]
if cat["categorization_name"] == "answer-type"
]
}
elif split == "test":
# only include qids that do have "comparison" in question_categories
fr_data = {
qid: q_data
for qid, q_data in fr_data.items()
if "comparison"
in [
cat["category_name"]
for cat in q_data["question_categories"]
if cat["categorization_name"] == "answer-type"
]
}
elif dataset_type == "complex":
if split == "train":
# filter out qids that do have "expert" in user_categories
fr_data = {
qid: q_data
for qid, q_data in fr_data.items()
if "expert"
not in [
cat["category_name"]
for cat in q_data["user_categories"]
if cat["categorization_name"] == "user-expertise"
]
}
elif split == "test":
# only include qids that do have "expert" in user_categories
fr_data = {
qid: q_data
for qid, q_data in fr_data.items()
if "expert"
in [
cat["category_name"]
for cat in q_data["user_categories"]
if cat["categorization_name"] == "user-expertise"
]
}
elif dataset_type == "open-ended":
if split == "train":
# filter out qids that do have "open-ended" in question_categories
fr_data = {
qid: q_data
for qid, q_data in fr_data.items()
if "open-ended"
not in [
cat["category_name"]
for cat in q_data["question_categories"]
if cat["categorization_name"] == "answer-type"
]
}
elif split == "test":
# only include qids that do have "open-ended" in question_categories
fr_data = {
qid: q_data
for qid, q_data in fr_data.items()
if "open-ended"
in [
cat["category_name"]
for cat in q_data["question_categories"]
if cat["categorization_name"] == "answer-type"
]
}
return fr_data
def get_args():
parser = argparse.ArgumentParser()
parser.add_argument("--base_metric", type=str, default="bem", choices=["bem", "ac"])
parser.add_argument(
"--dataset_type",
type=str,
default="balanced",
choices=["balanced", "multi-aspect", "comparison", "complex", "open-ended"],
)
parser.add_argument(
"--single_retriever",
type=str,
required=False,
choices=[
"bm25",
"bm25_stochastic",
"bm25_regularize",
"e5base",
"e5base_stochastic",
"e5base_regularize",
],
)
parser.add_argument(
"--non_learned_algo",
type=str,
required=False,
choices=[
"noRetrieval",
"random",
"overall_sim",
"avg_sim",
"max_sim",
"var_sim",
"moran",
],
)
parser.add_argument(
"--ltrr_algo",
type=str,
required=False,
choices=[
"pointwise-xgboost",
"pointwise-svm",
"pointwise-neural",
"pointwise-deberta",
"pairwise-xgboost",
"pairwise-svm",
"pairwise-neural",
"pairwise-deberta",
"listwise-listnet",
"listwise-lambdamart",
"listwise-deberta",
],
)
args = parser.parse_args()
return args
def get_utility_scores_non_learned_models(
fr_data: dict, model: str, higher_the_better: bool
) -> list:
utilities = []
for qid, q_data in fr_data.items():
action_value_map = {}
for action_key, action_data in q_data["routing_data"].items():
action_num = int(action_key.split("_")[-1])
model_value = action_data["post_retrieval_features"][model]
metric_value = action_data[f"{BASE_METRIC}_score"]
action_value_map.update({action_num: (model_value, metric_value)})
# Sort actions by model value (first element of tuple) in descending order
sorted_actions = sorted(
action_value_map.items(),
key=lambda x: x[1][0], # Sort by first element of value tuple (model_value)
reverse=True if higher_the_better else False, # Descending order
)
# Get the metric value (second element of tuple) of the highest-ranked action
utility_score = sorted_actions[0][1][1]
utilities.append(utility_score)
return utilities
if __name__ == "__main__":
args = get_args()
BASE_METRIC = str(args.base_metric)
DATASET_TYPE = str(args.dataset_type)
CUR_DIR_PATH = os.path.dirname(os.path.realpath(__file__))
FULL_RECALL_TEST_FP = os.path.join(CUR_DIR_PATH, "data", "full-recall", "test.json")
if args.ltrr_algo is not None:
LTRR_LEARNING_METHOD = str(args.ltrr_algo.split("-")[0])
LTRR_MODEL_NAME = str(args.ltrr_algo.split("-")[1])
LTRR_MODEL_DIR = os.path.join(
CUR_DIR_PATH,
"trained-ltrr-models",
f"{BASE_METRIC}-based",
f"{DATASET_TYPE}",
f"{LTRR_LEARNING_METHOD}",
)
LTRR_REPORT_FP = os.path.join(LTRR_MODEL_DIR, f"{LTRR_MODEL_NAME}-report.json")
DS_RESULT_DIR = os.path.join(
CUR_DIR_PATH,
"downstream_results",
f"{BASE_METRIC}-based",
f"{DATASET_TYPE}",
)
os.makedirs(DS_RESULT_DIR, exist_ok=True)
with open(FULL_RECALL_TEST_FP, "r") as f:
fr_test = json.load(f)
f.close()
# Filter test data based on dataset type
fr_test = filter_data(fr_test, DATASET_TYPE, split="test")
if True:
# get oracle utility
utilities = []
for qid, q_data in fr_test.items():
baseline_utility = q_data[f"baseline_{BASE_METRIC}_score"]
oracle_utility = max(
q_data["routing_data"][f"action_{i}"][f"{BASE_METRIC}_score"]
for i in range(1, 7)
)
oracle_utility = max(baseline_utility, oracle_utility)
utilities.append(oracle_utility)
avg_utility = statistics.mean(utilities)
std_utility = statistics.stdev(utilities)
with open(os.path.join(DS_RESULT_DIR, f"oracle.txt"), "w") as f:
f.write(f"avg_rag_utility: {round(avg_utility, 4)}\n")
f.write(f"std_rag_utility: {round(std_utility, 4)}")
f.close()
if args.single_retriever is not None:
model = args.single_retriever
retriever_to_action_num = {
"bm25": 1,
"bm25_stochastic": 2,
"bm25_regularize": 3,
"e5base": 4,
"e5base_stochastic": 5,
"e5base_regularize": 6,
}
action_num = retriever_to_action_num[model]
utilities = []
for qid, q_data in fr_test.items():
utility_score = q_data["routing_data"][f"action_{action_num}"][
f"{BASE_METRIC}_score"
]
utilities.append(utility_score)
avg_utility = statistics.mean(utilities)
std_utility = statistics.stdev(utilities)
with open(os.path.join(DS_RESULT_DIR, f"{model}.txt"), "w") as f:
f.write(f"avg_rag_utility: {round(avg_utility, 4)}\n")
f.write(f"std_rag_utility: {round(std_utility, 4)}")
f.close()
# getting downstream utility for non-learned models
if args.non_learned_algo is not None:
model = args.non_learned_algo
if model == "noRetrieval":
utilities = []
for qid, q_data in fr_test.items():
utilities.append(q_data[f"baseline_{BASE_METRIC}_score"])
elif model == "random":
utilities = []
for qid, q_data in fr_test.items():
random_num = random.randint(0, 6)
if random_num == 0:
utilities.append(q_data[f"baseline_{BASE_METRIC}_score"])
else:
utility_score = q_data["routing_data"][f"action_{random_num}"][
f"{BASE_METRIC}_score"
]
utilities.append(utility_score)
# for post-retrieval-feature-based models
elif model in ["overall_sim", "avg_sim", "max_sim", "moran"]:
utilities = get_utility_scores_non_learned_models(
fr_test, model=model, higher_the_better=True
)
elif model in ["var_sim"]:
utilities = get_utility_scores_non_learned_models(
fr_test, model=model, higher_the_better=False
)
else:
raise ValueError(f"Invalid non-learned model: {model}")
avg_utility = statistics.mean(utilities)
std_utility = statistics.stdev(utilities)
with open(os.path.join(DS_RESULT_DIR, f"{model}.txt"), "w") as f:
f.write(f"avg_rag_utility: {round(avg_utility, 4)}\n")
f.write(f"std_rag_utility: {round(std_utility, 4)}")
f.close()
# getting downstream utility for learned models
if args.ltrr_algo is not None:
model = args.ltrr_algo
with open(LTRR_REPORT_FP, "r") as f:
report_dict = json.load(f)
f.close()
report_dict: dict = report_dict["per_query_predicted_rankings"]
utilities = []
for qid, pred_dict in report_dict.items():
top_choice = int(pred_dict["predicted_rankings"][1])
if top_choice != 0:
utility = fr_test[qid]["routing_data"][f"action_{top_choice}"][
f"{BASE_METRIC}_score"
]
else:
utility = fr_test[qid][f"baseline_{BASE_METRIC}_score"]
utilities.append(utility)
avg_utility = statistics.mean(utilities)
std_utility = statistics.stdev(utilities)
with open(os.path.join(DS_RESULT_DIR, f"{model}.txt"), "w") as f:
f.write(f"avg_rag_utility: {round(avg_utility, 4)}\n")
f.write(f"std_rag_utility: {round(std_utility, 4)}")
f.close()