{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 1,
   "id": "f00a55d1",
   "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": "markdown",
   "id": "1f9c1f52",
   "metadata": {},
   "source": [
    "# switch off different elements in DualLoop"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 2,
   "id": "6a08c229",
   "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"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "b0e156e1",
   "metadata": {},
   "outputs": [],
   "source": [
    "for "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "664ed149",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "e048f9f3",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "fe71d33b",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "8ecb1012",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": 3,
   "id": "77810621",
   "metadata": {},
   "outputs": [],
   "source": [
    "def plot_f1_result(results, dataset_name, filename):          \n",
    "    exp_result_df = pd.DataFrame(results)\n",
    "\n",
    "    colors = [\"green\", \"blue\", \"dimgray\", \"purple\"]\n",
    "    sns.set_palette(sns.color_palette(colors))    \n",
    "        \n",
    "    \n",
    "    fig, ax = plt.subplots(figsize=(8, 6))    \n",
    "    plt.xlim(0, 100)\n",
    "    plt.ylim(60, 100)   \n",
    "\n",
    "    p = sns.lineplot(x='iteration', y='metric', hue='methods', ax=ax, data=exp_result_df, \n",
    "                 style=\"methods\", markers=['s', 'D', 'o', 'X'])\n",
    "\n",
    "    p.set_xlabel(\"% of queried candidates\", fontsize=20)\n",
    "    p.set_ylabel(\"F1 (%)\", fontsize=20)   \n",
    "    plt.xticks(fontsize=16)    \n",
    "    plt.yticks(fontsize=16)    \n",
    "    \n",
    "    plt.legend(loc='lower right')\n",
    "    \n",
    "    plt.savefig(filename, transparent=False)\n",
    "    plt.close(fig)        \n",
    "    \n",
    "def plot_recall_result(results, dataset, filename):            \n",
    "    exp_result_df = pd.DataFrame(results)\n",
    "\n",
    "    colors = [\"green\", \"blue\", \"dimgray\", \"purple\"]\n",
    "    sns.set_palette(sns.color_palette(colors))    \n",
    "    \n",
    "    fig, ax = plt.subplots(figsize=(8, 6))    \n",
    "    plt.xlim(0, 100)\n",
    "    plt.ylim(50, 100)      \n",
    "    \n",
    "    marker_list = [\"\"]\n",
    "    \n",
    "    p = sns.lineplot(x='iteration', y='metric', hue='methods', ax=ax, data=exp_result_df, \n",
    "                 style=\"methods\", markers=['s', 'D', 'o', 'X'])\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.legend(loc='lower right')\n",
    "            \n",
    "    plt.savefig(filename, transparent=False)\n",
    "    plt.close(fig)         "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 4,
   "id": "3d9a8e21",
   "metadata": {},
   "outputs": [],
   "source": [
    "result_folder = '../results/'\n",
    "\n",
    "measurements = []\n",
    "for dataset in datasets:\n",
    "    for method in methods:\n",
    "        result = pickle.load( open(result_folder + dataset + '/result_' + method + '_' + dataset  + \".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",
    "            m = result[iteration]\n",
    "            f1= m['f1'] * 100.0\n",
    "            \n",
    "            budget = m['human_effort'] \n",
    "            percentage = int(budget * 100.0/max_iteration // 5 * 5 + 5)               \n",
    "            \n",
    "#             if iteration <=0:\n",
    "#                 percentage = 0\n",
    "#             else:\n",
    "#                 percentage = int((iteration-1) * 100.0/max_iteration // 5 * 5 + 5)               \n",
    "            \n",
    "            if percentage >= 100.0:\n",
    "                metric = 100.0            \n",
    "\n",
    "            measurement = {\n",
    "                'methods': method,\n",
    "                'iteration': percentage,\n",
    "                'metric': f1\n",
    "            }   \n",
    "            measurements.append(measurement)\n",
    "    "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 5,
   "id": "3040bf94",
   "metadata": {},
   "outputs": [],
   "source": [
    "plot_f1_result(measurements, 'for all datasets',  './figures/f1_breakdown.pdf')"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 6,
   "id": "2ad454b5",
   "metadata": {},
   "outputs": [],
   "source": [
    "result_folder = '../results/'\n",
    "\n",
    "measurements = []\n",
    "for dataset in datasets:\n",
    "    for method in methods:\n",
    "        result = pickle.load( open(result_folder + dataset + '/result_' + method + '_' + dataset  + \".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",
    "            m = result[iteration]\n",
    "            recall = m['recall'] * 100.0\n",
    "            \n",
    "            budget = m['human_effort'] \n",
    "            percentage = int(budget * 100.0/max_iteration // 5 * 5 + 5)         \n",
    "            \n",
    "            if percentage >= 100.0:\n",
    "                metric = 100.0                   \n",
    "\n",
    "            measurement = {\n",
    "                'methods': method,\n",
    "                'iteration': percentage,\n",
    "                'metric': recall\n",
    "            }   \n",
    "            measurements.append(measurement)\n",
    "    "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 7,
   "id": "4d6ded38",
   "metadata": {},
   "outputs": [],
   "source": [
    "plot_recall_result(measurements, 'for all datasets', './figures/recall_breakdown.pdf')"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "42b96635",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "da19d828",
   "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
}
