{
"cells": [
{
"cell_type": "code",
"execution_count": 1,
"id": "3a679ef8",
"metadata": {},
"outputs": [],
"source": [
"import pandas as pd\n",
"import pickle\n",
"import sys\n",
"sys.path.append('../src')\n",
"import seaborn as sns\n",
"import matplotlib.pyplot as plt\n",
" \n",
"sns.set(font_scale = 3) \n",
"sns.set_theme()\n",
"sns.set_style(\"whitegrid\") "
]
},
{
"cell_type": "code",
"execution_count": 2,
"id": "2f7a374e",
"metadata": {},
"outputs": [],
"source": [
"datasets = [\n",
" 'conference-1',\n",
" 'conference-2',\n",
" 'conference-3',\n",
" 'conference-4',\n",
" 'conference-5',\n",
" 'conference-6',\n",
" 'conference-7',\n",
" 'conference-8',\n",
" 'conference-9',\n",
" 'conference-10',\n",
" 'conference-11',\n",
" 'conference-12',\n",
" 'conference-13',\n",
" 'conference-14',\n",
" 'conference-15',\n",
" 'conference-16',\n",
" 'conference-17',\n",
" 'conference-18',\n",
" 'conference-19',\n",
" 'conference-20',\n",
" 'conference-21',\n",
" 'ai4eu-1',\n",
" 'ai4eu-2',\n",
" 'ai4eu-3',\n",
" 'nasa'\n",
"]\n",
"\n",
"# loading the results from each dataset\n",
"methods = [\n",
" 'WeSAL',\n",
" 'AL-RF',\n",
" 'DualLoop',\n",
"]\n",
"\n",
"ds_names = [\n",
" 'conference', \n",
" 'ai4eu', \n",
" 'nasa'\n",
"]\n",
"\n",
"titles = {\n",
" 'conference': 'Conference',\n",
" 'ai4eu': 'AI4EU',\n",
" 'nasa': 'AirTraffic'\n",
"}\n"
]
},
{
"cell_type": "code",
"execution_count": 3,
"id": "4859db0d",
"metadata": {},
"outputs": [],
"source": [
"def plot_recall_result(results, dataset, filename): \n",
" exp_result_df = pd.DataFrame(results)\n",
"\n",
"\n",
" colors = [\"gray\", \"blue\", \"green\"]\n",
" sns.set_palette(sns.color_palette(colors)) \n",
" \n",
" fig, ax = plt.subplots(figsize=(8, 6)) \n",
" plt.ylim(0, 100)\n",
" plt.xlim(0, 100) \n",
"\n",
" p = sns.lineplot(x='iteration', y='metric', hue='methods', ax=ax, data=exp_result_df, \n",
" style=\"methods\", markers=['D', 'o', 's'])\n",
"\n",
"# p.set_xlabel(\"percentage of query cost (%)\")\n",
"# p.set_ylabel(\"Recall (%)\") \n",
" \n",
"# ax.title.set_text(dataset)\n",
" \n",
" p.set_xlabel(\"% of queried candidates\", fontsize=20)\n",
" p.set_ylabel(\"Recall (%)\", fontsize=20) \n",
" plt.xticks(fontsize=16) \n",
" plt.yticks(fontsize=16) \n",
" \n",
" plt.savefig(filename, transparent=False)\n",
" plt.close(fig) \n",
" \n",
"def read_all_result(metric_name, result_folder):\n",
" all_results = {}\n",
"\n",
" for name in ds_names:\n",
" my_datasets = []\n",
" for ds in datasets:\n",
" if ds.startswith(name):\n",
" my_datasets.append(ds)\n",
"\n",
" results = []\n",
" for method in methods:\n",
" for ds in my_datasets:\n",
" result = pickle.load( open(result_folder + ds + '/result_' + method + '_' + ds + \".pk\", \"rb\" ) ) \n",
"\n",
" max_iteration = 1\n",
" for iteration in result:\n",
" if iteration > max_iteration:\n",
" max_iteration = iteration \n",
" \n",
" for iteration in result:\n",
" r = result[iteration]\n",
" metric = r[metric_name] * 100.0\n",
" \n",
" budget = r['human_effort'] \n",
" percentage = int(budget * 100.0/max_iteration // 2 * 2 + 2) \n",
" \n",
" if percentage >= 100.0:\n",
" metric = 100.0 \n",
" \n",
" measurement = {\n",
" 'methods': method,\n",
" 'iteration': percentage,\n",
" 'metric': metric\n",
" } \n",
" results.append(measurement)\n",
"\n",
" all_results[name] = results \n",
" \n",
" return all_results"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "77fbbe8a",
"metadata": {
"scrolled": true
},
"outputs": [],
"source": [
"metric = 'recall'\n",
"folder = '../results/'\n",
"\n",
"r = read_all_result(metric, folder)\n",
"for name in ds_names:\n",
" title = titles[name] \n",
" plot_recall_result(r[name], title, './figures/recall_' + name + '.pdf')"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "bd49e65c",
"metadata": {},
"outputs": [],
"source": []
},
{
"cell_type": "code",
"execution_count": null,
"id": "a332283d",
"metadata": {},
"outputs": [],
"source": []
}
],
"metadata": {
"kernelspec": {
"display_name": "Python 3 (ipykernel)",
"language": "python",
"name": "python3"
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 3
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.8.8"
}
},
"nbformat": 4,
"nbformat_minor": 5
}