{
 "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
}
