added dec and sym exp loops

This commit is contained in:
Kashif Rasul committed 2022-05-20 14:50:14 +02:00
1 parent 8628813bae
commit af6b4ea765
1 file changed
+142 -44
+142 -44
View File
@@ -3,7 +3,7 @@
{
"cell_type": "code",
"execution_count": 55,
"id": "98b88925",
"id": "5a62c0ff",
"metadata": {},
"outputs": [],
"source": [
@@ -18,7 +18,7 @@
{
"cell_type": "code",
"execution_count": 50,
"id": "042a0773",
"id": "c75f5208",
"metadata": {},
"outputs": [],
"source": [
@@ -32,7 +32,7 @@
{
"cell_type": "code",
"execution_count": 8,
"id": "35c300a9",
"id": "aa342b02",
"metadata": {},
"outputs": [
{
@@ -169,7 +169,7 @@
{
"cell_type": "code",
"execution_count": 16,
"id": "c7b82213",
"id": "5aefbb4a",
"metadata": {},
"outputs": [],
"source": [
@@ -180,7 +180,7 @@
{
"cell_type": "code",
"execution_count": 11,
"id": "053971b7",
"id": "047f9d32",
"metadata": {},
"outputs": [],
"source": [
@@ -190,7 +190,7 @@
{
"cell_type": "code",
"execution_count": 12,
"id": "260573cc",
"id": "513bbe1f",
"metadata": {},
"outputs": [],
"source": [
@@ -200,7 +200,7 @@
{
"cell_type": "code",
"execution_count": 13,
"id": "491f8c85",
"id": "2e484da6",
"metadata": {},
"outputs": [],
"source": [
@@ -209,7 +209,7 @@
},
{
"cell_type": "markdown",
"id": "ed03133a",
"id": "b5630b79",
"metadata": {},
"source": [
"## Sclaing Experiments\n",
@@ -228,7 +228,7 @@
{
"cell_type": "code",
"execution_count": 51,
"id": "0a65866c",
"id": "7945eff4",
"metadata": {},
"outputs": [],
"source": [
@@ -237,7 +237,7 @@
},
{
"cell_type": "markdown",
"id": "d51bfbfc",
"id": "b0b59bb2",
"metadata": {},
"source": [
"### Encoder Scaling"
@@ -246,7 +246,7 @@
{
"cell_type": "code",
"execution_count": null,
"id": "af9dc0ae",
"id": "681199ae",
"metadata": {},
"outputs": [],
"source": [
@@ -296,37 +296,19 @@
{
"cell_type": "code",
"execution_count": 53,
"id": "ef47755d",
"metadata": {},
"outputs": [],
"source": [
"enc_metrics_out = open(\"elec_enc_metrics.pkl\", \"wb\")"
]
},
{
"cell_type": "code",
"execution_count": 56,
"id": "855b130e",
"metadata": {},
"outputs": [],
"source": [
"pickle.dump(enc_metrics, enc_metrics_out)"
]
},
{
"cell_type": "code",
"execution_count": 57,
"id": "869ccbec",
"id": "c6cb9386",
"metadata": {},
"outputs": [],
"source": [
"enc_metrics_out = open(\"elec_enc_metrics.pkl\", \"wb\")\n",
"pickle.dump(enc_metrics, enc_metrics_out)\n",
"enc_metrics_out.close()"
]
},
{
"cell_type": "code",
"execution_count": 66,
"id": "02df0064",
"id": "513cbbbf",
"metadata": {
"scrolled": true
},
@@ -367,7 +349,7 @@
{
"cell_type": "code",
"execution_count": 68,
"id": "95290703",
"id": "af5715b2",
"metadata": {},
"outputs": [
{
@@ -405,8 +387,8 @@
},
{
"cell_type": "code",
"execution_count": 72,
"id": "61cf3e7d",
"execution_count": 79,
"id": "1c64e02c",
"metadata": {},
"outputs": [
{
@@ -415,13 +397,13 @@
"Text(0.5, 1.0, 'Encoder Scaling')"
]
},
"execution_count": 72,
"execution_count": 79,
"metadata": {},
"output_type": "execute_result"
},
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYoAAAEWCAYAAAB42tAoAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjUuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8qNh9FAAAACXBIWXMAAAsTAAALEwEAmpwYAABDyklEQVR4nO3deXxcZfX48c/J3mZrs3VP06Q7pbSldJEdCpSyiYiC6BcUQf2pgIKI+nVXBEQQ0C984SvigiyKbFKWUnZoCyVdaJt0S5N0zZ5mbZrl/P6YO2Uaksk2+5z36zWvztx7587JJJ0zz32e5zyiqhhjjDG9iQl2AMYYY0KbJQpjjDFeWaIwxhjjlSUKY4wxXlmiMMYY45UlCmOMMV5ZojBmAETkDRH5arDj8EZEHhGRXzn3TxaRrcGOyYQ3SxQmbIhIqYi0ikiTx+0PwY5rKETkahEpFpFGEakQkeUikuqr86vq26o6zVfnM9EpLtgBGDNAF6jqq8EOYqBEJE5VO7ptOxW4FViqqutEJAO4ICgBGuOFtShMRBCRq0TkHRG5U0TqRGSXiJzrsT9DRP4sIvuc/c947LtGRHaISK2IPCciYz32neV84z/otF6k2+t+RUSKnHO+LCITPfapiHxTRLYD23sI+wRglaquA1DVWlX9i6o2Os8fJiK/E5Ey5/XfEZFhzr5/isgBZ/tbInJML+/LaSKyx+NxqYjcJCIbnec+ISJJHvtvFpH9zvv0VednmNzPX4OJUJYoTCRZCGwFsoA7gD+JiPuD/W/AcOAYIAe4G0BEzgB+A3wOGAOUAY87+7KAfwP/7ZxzJ3Ci+8VE5CLgh8BngGzgbeCxbjF92olrZg/xrgHOEZGfi8iJIpLYbf+dwPHAp4AM4Gagy9n3IjDF+VkKgUf7eG88fQ5YCkwCZgNXOT/PUuC7wBJgMnDaAM5pIpmq2s1uYXEDSoEmoN7jdo2z7ypgh8exwwEFRuNKAF3AyB7O+SfgDo/HKUA7kAf8F7DaY58Ae4CvOo9fBK722B8DtAATnccKnNHHz3Qu8LzzszQBdwGxzrlageP68b6McF4r3Xn8CPAr5/5pwJ5u7+EXPR7fATzg3H8Y+I3HvsnOeScH+3dvt+DerEVhws2nVXWEx+0hj30H3HdUtcW5mwJMAGpVta6H843F1YpwP68JqAHGOft2e+xTz8fAROAeEakXkXqgFlcyGedxjOfxn6CqL6rqBbhaDBfhSnhfxdWCScLVijmKiMSKyG0islNEGnB9+OM8pz8OeNxvwfUeQbeft6/YTfSwRGGiwW4gQ0RG9LBvH64PfABEJBnIBPYC+3ElGfc+8XzsnPdr3RLXMFV9z+OYfpVnVtUuVV0JvAbMAqqBQ0BBD4d/AVdSWQKk42r9QLf+k0HYD4z3eDyhtwNNdLFEYSKequ7HdZnof0RkpIjEi8gpzu7HgC+LyBynj+BWYI2qlgIvAMeIyGdEJA64DtelLLcHgB+4O5JFJF1ELu1vXCJykYhc5sQkIrIAOBXX5a4uXJeC7hKRsU4rYrETYyrQhqvlM9yJ2ReexPVezBCR4cCPfXReE+YsUZhw83y3eRRP9/N5X8LV91AMVAI3AKhrqO2PgadwfaMuAC5z9lUDlwK34fpQngK86z6hqj4N3A487lwC2oSrz6G/6oBrcI2IagD+DvxWVd0d0zcBHwEf4LqsdTuu/7N/xXW5bC+wBVg9gNfslaq+CNwLvA7s8Dhvmy/Ob8KXuC67GmPM0URkBq7kl6jd5oCY6GItCmPMESJysYgkishIXC2Y5y1JGEsUxhhPX8N1aW4n0Al8I7jhmFBgl56MMcZ4ZS0KY4wxXkVUUcCsrCzNy8sLdhjGGBM2Pvzww2pVzfZ2TEQliry8PNauXRvsMIwxJmyISFlfx9ilJ2OMMV5ZojDGGOOVJQpjjDFeWaIwxhjjlSUKY4wxXvktUYjIBBF5XUS2iMhmEbne2f4zEdkrIuud27Jenr9URLY6S1Te4q84jTHGeOfP4bEdwI2qWigiqcCHIrLC2Xe3qt7Z2xNFJBb4I3AWrhXFPhCR51R1ix/jNcYY0wO/tShUdb+qFjr3G4Eijl75y5sFuJa1LFHVw7jWML7IP5Ga7g61d/LomjLaOjqDHYoxJgQEpI9CRPKAubgWkwf4lohsFJGHnSqV3Y3j6GUY99BLkhGRa0VkrYisraqq8mXYUeuht0r40dObeL3Y3k9jTAAShYik4FoU5gZVbQDux7U4zBxcC8X8bijnV9UHVXW+qs7PzvY6C930Q1VjGw+86VqmeWdVU5CjMcaEAr8mChGJx5UkHlXVfwOoaoWqdjpLPT6E6zJTd3s5er3e8c4242e/f3UbbR1dpCbGUVLVHOxwjDEhwJ+jngT4E1Ckqnd5bB/jcdjFuFbQ6u4DYIqITBKRBFxLUz7nr1iNy47KRh7/YDdXLMxl1rh0a1EYYwD/tihOxLVO8RndhsLeISIfichG4HTgOwDOAvLLAZwVtb4FvIyrE/xJVd3sx1gNcNuLxQyPj+W6M6dQkJNMSVUTtl6JMcZvw2NV9R1Aeti1vJfj9wHLPB4v7+1Y43vv7azm1aJKvr90OpkpieRnpdBwqIOa5sNkpSQGOzxjTBDZzGxDV5dy6/IixqYn8eUT8wAoyEkBYGelXX4yJtpZojA8u2Evm/Y28L2l00iKjwUgPysZgJJq69A2JtpZoohyh9o7ufPlbcwal8ZFx308VWXciGEkxsVQYh3axkQ9SxRR7s/vlrK3vpUfLptBTMzHXUoxMcKkrGR22hBZY6KeJYooVtPUxv+8voMzp+fwqYKsT+wvyE6xFoUxxhJFNLvvtR20tHfyg2XTe9yfn53M7rpWq/lkTJSzRBGlSqqa+PvqMj5/wgQm56T2eExBdgqdXUp5TUuAozPGhBJLFFHq9peKSYyL4YYlU3o9Jj/bNfLJ+imMiW6WKKLQB6W1vLy5gq+fWkBOalKvx03KcicK66cwJppZoogyqsqvXihiVFoiXz053+uxqUnxjEpLtOKAxkQ5SxRR5j8b97Nhdz03nj2NYQmxfR6fn5VCSbW1KIyJZpYookhbRye3v1TM9NGpXDJvfL+eU5CTzM5KKw5oTDSzRBFF/raqjD11rfzovBnExvRUr/GTPIsDGmOikyWKKFHfcph7V27nlKnZnDyl/ysBukc+WT+FMdHLEkWUuO+1HTS1dfDDXibX9aYg26kiayOfjIlaliiiQHlNC39dVcqlx09g+ui0AT3XigMaYyxRRIHbXy4mLiaG7549dcDPteKAxhhLFBGusLyOFzbu55pT8hmV1vvkOm+sOKAx0c0SRQRTVW59oYislES+dor3yXXeWHFAY6Kb3xKFiEwQkddFZIuIbBaR653tvxWRYhHZKCJPi8iIXp5fKiIfich6EVnrrzgj2cubD7C2rI4bz55KcuLgl0fPz0624oDGRDF/tig6gBtVdSawCPimiMwEVgCzVHU2sA34gZdznK6qc1R1vh/jjEiHO7q47cVipuSkcOnx/Ztc15uPRz5ZP4Ux0chviUJV96tqoXO/ESgCxqnqK6ra4Ry2Ghjap5jp0aNryiitaeGHy2YQFzu0X/OkI+tnWz+FMdEoIH0UIpIHzAXWdNv1FeDFXp6mwCsi8qGIXOvl3NeKyFoRWVtVVeWTeMPdwdZ27l25nRMnZ3LatP5PrutNalI8OamJ7Ky0FoUx0cjviUJEUoCngBtUtcFj+49wXZ56tJennqSq84BzcV22OqWng1T1QVWdr6rzs7OH/qEYCf7njR3Ut7bzg3NnINK/Uh19Kci24oDGRCu/JgoRiceVJB5V1X97bL8KOB+4QnupNqeqe51/K4GngQX+jDVS7Klr4c/vlnLx3HHMGpfus/PmZydTUtVsxQGNiUL+HPUkwJ+AIlW9y2P7UuBm4EJV7XEYjYgki0iq+z5wNrDJX7FGkjtf3ooAN509zafnLchO4WBruxUHNCYK+bNFcSLwJeAMZ4jrehFZBvwBSAVWONseABCRsSKy3HnuKOAdEdkAvA+8oKov+THWiLBxTz3PrN/H1SdNYuyIYT49txUHNCZ6DX5wfR9U9R2gpwvky3vYhqruA5Y590uA4/wVWyRSVX79QhGZyQl847QCn5/fszjggkkZPj+/MSZ02czsCLGyqJI1u2q5YckUUpPifX7+sVYc0JioZYkiQtyzcjv5WclctiDXL+ePdYoD2qUnY6KPJYoI0NmlFB9o4OxjRhM/xMl13uRnJ9u6FMZEIUsUEWBffSvtncrEzOF+fZ2C7BR217VyuKPLr69jjAktligiQJlTrM/fieJIccBau/xkTDSxRBEBypwP7rzMZL++Tn6Wa+TTDivlYUxUsUQRAcpqWkiIi2H0IBcm6q8jcymslIcxUcUSRQQoq2kmN2M4MTG+qevUGysOaEx0skQRAcpqWsjzc/+EmxUHNCb6WKIIc6pKWU0LuRn+7Z9ws+KAxkQfSxRhrqqxjdb2TvKyAtOiyLfigMZEHUsUYa7UGRqbmxGoS09WHNCYaGOJIsyV1gRmaKybuzig1XwyJnpYoghz5TUtxMYI40b6tqx4b8aOGEZCXIyV8jAmiliiCHOlNc2MGzHMrzWePMXGCPlWHNCYqGKJIsyV1bT4vXRHd/nZyZRUW6IwJlpYoghjqkppTXPA+ifc8rNSKK9tseKAxkQJSxRhrL6lncZDHQFvURTkWHFAY6KJJYow5h7xNDEILQqw4oDGRAu/JQoRmSAir4vIFhHZLCLXO9szRGSFiGx3/h3Zy/OvdI7ZLiJX+ivOcFZe65pDEajyHW5WHNCY6OLPFkUHcKOqzgQWAd8UkZnALcBKVZ0CrHQeH0VEMoCfAguBBcBPe0so0ay02pUoJgRosp2buzigjXwyJjr4LVGo6n5VLXTuNwJFwDjgIuAvzmF/AT7dw9PPAVaoaq2q1gErgKX+ijVcldU2MyY9iaTLine truncated
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYMAAAEWCAYAAACEz/viAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjUuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8qNh9FAAAACXBIWXMAAAsTAAALEwEAmpwYAAA+oElEQVR4nO3dd3xU55X4/88ZdQlJgxogJCGBRAdRhQu4xDYuseM4No6dZBN77bTNpmySzWZ3Uzf5btqud5OfU9ZJvE7srIlbihPbYMcNHGOqKAKDBEIVNdQbKvP8/pg7eBCSRmXKHem8X695MTP3zr1HgzRn7lPOI8YYlFJKTW+OUAeglFIq9DQZKKWU0mSglFJKk4FSSik0GSillEKTgVJKKTQZKHUREXlVRO4PdRyjEZFHROTb1v1NInI81DGp8KbJQNmKiJwWkR4R6fS6PRjquCZDRO4TkbdFpENE6kXkORFJ9NfxjTE7jDGL/HU8NT1FhjoApYZxizHmpVAHMV4iEmmMGRjy3JXAvwM3GGMOiEgKcEtIAlRqFHploMKGiNwjIjtF5D9EpEVEykXkRq/tKSLyvyJSa23/vde2j4pImYg0i8gfRSTTa9t11jf3NusqRIac929F5Jh1zG0iMs9rmxGRT4lIKVA6TNjrgTeNMQcAjDHNxphfGWM6rNfHich/ikiFdf6dIhJnbXtSROqs518XkWUjvC9XiUi11+PTIvJFETlkvfa3IhLrtf1LInLGep/ut36G/DH+N6gpSpOBCjcbgONAGvB94Jci4vnwfhSIB5YBGcB/AYjIu4DvAHcCc4AKYKu1LQ14BviKdcyTwOWek4nIrcC/AO8D0oEdwONDYnqvFdfSYeJ9C7heRL4pIpeLSMyQ7f8BrAUuA1KALwEua9vzQIH1s+wHfuPjvfF2J3ADkAesBO6xfp4bgM8D1wL5wFXjOKaayowxetObbW7AaaATaPW6fdTadg9Q5rVvPGCA2bg/5F3AzGGO+Uvg+16PZwD9QC7wYWCX1zYBqoH7rcfPA/d5bXcA3cA867EB3uXjZ7oReNb6WTqBB4AI61g9QOEY3henda5k6/EjwLet+1cB1UPeww95Pf4+8DPr/sPAd7y25VvHzQ/1/73eQnvTKwNlR+81xji9bj/32lbnuWOM6bbuzgCygWZjTMswx8vEfTXgeV0ncBaYa22r8tpmvB8D84AfikiriLQCzbgTxlyvfbz3v4gx5nljzC24v/nfijup3Y/7SiQW99XIBUQkQkS+KyInRaQd9wc81mvGos7rfjfu9wiG/Ly+YlfThyYDNVVUASki4hxmWy3uD3UARCQBSAVqgDO4E4lnm3g/to778SHJKc4Y81evfcZU+tcY4zLG/AV4GVgONAG9wIJhdv8A7sRxLZCM+yoGhvRnTMAZIMvrcfZIO6rpRZOBmhKMMWdwN+n8RERmikiUiFxhbX4cuFdEVllt9v8OvGWMOQ38GVgmIu8TkUjgM7ibnTx+Bvyzp/NWRJJFZMtY4xKRW0XkLismEZEi4ErcTVMu3M02D4hIpnU1cKkVYyJwDvcVTLwVsz88gfu9WCIi8cBX/XRcFeY0GSg7enbIPIPfjfF1f4O7L+BtoAH4HIBxD1P9KvA07m/GC4C7rG1NwBbgu7g/eAuANzwHNMb8DvgesNVqrjmCuw9grFqAj+IeadQOPAb8wBjj6Qz+InAY2IO7Cep7uP8uf427aasGOArsGsc5R2SMeR74EfAKUOZ13HP+OL4KX+JuIlVKTUcisgR3gosxQ+ZIqOlFrwyUmmZE5DYRiRGRmbivRJ7VRKA0GSg1/XwcdzPaSWAQ+GRow1F2oM1ESiml9MpAKaVUGBaqS0tLM7m5uaEOQymlwsq+ffuajDHpI20Pu2SQm5vL3r17Qx2GUkqFFRGpGG27NhMppZTSZKCUUkqTgVJKKTQZKKWUQpOBUkopNBkopZRCk4FSSik0GSg/ev7wGaqau33vqJSyHU0Gyi9K6zv45G/285NXy0IdilJqAjQZKL94+I1yAIqr2kIciVJqIjQZqEk723mOp/fXEBPp4ER9B919WhpfqXCjyUBN2qO7KugbcPH56xYy6DKU1LaHOiSl1DhpMlCT0ts/yKNvVnDN4gxuWzMXgINVraENSik1bpoM1KT8obiGs1193Lcpj4zEWOY64yjWZKBU2NFkoCbMGMMvdpSzdE4Sl85PBaAwO5mD1a2hDUwpNW6aDNSEvXaikdKGTu7flIeIAFCY5aSquYeznedCHJ1Sajw0GagJ++XOcmYlxXDzyszzzxVmOwE4VK1DTJUKJ5oM1IS8XdfOjtImPnJZLtGR7/warZibjEPQfgOlwowmAzUhv9hRTlxUBB8oyrng+YSYSAoyErXfQKkwo8lAjVtDey9/KK7hznVZOOOjL9pemJ3MwapWjDEhiE4pNRGaDNS4PbqrggGX4d7L84bdXpjtpKW7n6rmniBHppSaKE0Galx6+gZ5bFcF1y2ZRW5awrD7FGY5ASjWpiKlwoYmAzUuT++vpqW7n/s3zR9xn0WzE4mNclBc2Rq8wJRSk6LJQI2Zy2V4eGc5hVnJrM+dOeJ+UREOlmfq5DOlwokmAzVmrxxv4FRTF/dtmn9+ktlICrOdHKlpo3/QFaTolFKTEbBkICIPi0iDiBwZYXuyiDwrIgdFpERE7g1ULMo/fr7jFJnJsdy4fLbPfQuznZwbcHG8riMIkSmlJiuQVwaPADeMsv1TwFFjTCFwFfCfInLxOEVlC0dq2th1qpl7L88jKsL3r80qqxNZm4qUCg8BSwbGmNeB5tF2ARLF3d4ww9pXV0WxqV/uLCchOoL3F2WPaf/slDhmxkdpOWulwkQo+wweBJYAtcBh4LPGmGEbmEXkYyKyV0T2NjY2BjNGBZxp6+HZg7W8f30OSbFRY3qNiFCY7eSgLoOpVFgIZTK4HigGMoFVwIMikjTcjsaYh4wx64wx69LT04MXoQLgV3+twGUM916eO67XFWY5OdHQQec5veBTyu5CmQzuBZ4xbmVAObA4hPGoYXSdG+D/3qrgxuVzyE6JH9drV2U7Mcbd36CUsrdQJoNK4BoAEZkFLAJOhTAeNYyn9lXT3jvAfZuGLz0xmpVZyYAug6lUOIgM1IFF5HHco4TSRKQa+DoQBWCM+RnwLeARETkMCPBPxpimQMWjxm/QZfjlznLW5DhZkzPyJLORpM6IITslTkcUKRUGApYMjDF3+9heC2wO1PnV5L14tJ7K5m6+fOPEW+8Ks5wc0LIUStmezkBWI/rlzlNkp8Rx/TLfk8xGsirbSU1rDw0dvX6MTCnlb5oM1LCKq1rZc7qFey/LI8IxeumJ0ZxfBlOHmCpla5oM1LB+seMUibGR3Ll+bJPMRrIsM4kIh2i/gVI2p8lAXaS6pZvnj9TxgaIcZsRMrlspPjqShbMSdU1kpWxOk4G6yK/+ehqAj1yW65fjrcp2crCqFZdLl8FUyq40GagLdPT2s3V3Fe9eMYdMZ5xfjrkqO5n23gFOn+3yy/GUUv6nyUBd4Ld7qug4N8D9E5hkNhJPJ7L2GyhlX5oM1HkDgy7+943TFOWlsNIqQe0PBRmJxEdHaNE6pWxMk4E6b1tJPTWtPdy/0X9XBQARDmH53GTtRFbKxjQZKACMMfx8xylyU+O5Zsksvx9/VbaTo7Xt9A3oMphK2ZEmAwXA/soWiqtauW/j5CaZjaQwy0nfoIu369r9fmyl1ORpMlAA/GJHOclxUdy+Nisgxy/M1gqmStmZJgNF5dlutpXU8cENOcRHB6Z24VxnHGkzoinWTmSlbEmTgeLhN8qJcIjfJpkNR0QozHLq8FKlbEqTwTTX1tPPE3uruKUwk1lJsQE9V2G2k5ONnbT39gf0PEqp8dNkMM1t3V1Jd98g9/l5OOlwCj3LYFZrU5FSdqPJYBrrH3TxyF9Pc9mCVJZlJgf8fIXWMpjF2lSklO1oMpjGnjt8hjNtvXx00/ygnM8ZH01uaryOKFLKhjQZTGNP7atmXmo8Vy5MD9o5C7OdU6IsRXFVK739g6EOQym/0WQwTbV19/PmybPcuHwOjgBMMhtJYZaTuvZe6trCdxnMps5zvO8nb/DLneWhDkUpv9FkME29fLyeAZfh+mX+Lz0xGk8F03CuU3SivgOXgTdPng11KEr5jSaDaWp7ST0ZiTEU+rE66Vgsy0wiMsyXwTzZ0Am4S3gMDGqtJTU1aDKYhnr7B3n1eCObl80KahMRQGxUBEvmJIV1J3KplQy6+wYpqdVaS2pq0GQwDe0sbaKnf5DNS2eH5PyF2ckcqm4L22UwS+s7yZrpXgVud3lziKNRyj8ClgxE5GERaRCRI6Psc5WIFItIiYi8FqhY1IW2ldSRGBvJJfNTQ3L+wiwnnecGONXUGZLzT1ZpQyeXzk9lXmo8u09rMlBTQyCvDB4Bbhhpo4g4gZ8A7zHGLAO2BDAWZRkYdPHSsXquWZxBdGRoLgxXne9EDr8hpq3dfTR1niM/YwZFuSnsOd0ctlc4SnkL2KeBMeZ1YLSvTR8AnjHGVFr7NwQqFvWOvRUttHT3s3lZaJqIAOanz2BGTGRY9huUWf0FBbNmsD4vhdbufsoaw/MKRylvoewzWAjMFJFXRWSfiHx4pB1F5GMisldE9jY2NgYxxKlnW0kd0ZGOoE40GyrCIayYmxyWI4o8nccFGYlsyEsB4C3tN1BTQCiTQSSwFng3cD3wVRFZONyOxpiHjDHrjDHr0tND9yEW7owxbC+pZ1N+GgkxgVm3YKwKs50cO9MedrN4yxo6iY1yMNcZR05KPBmJMezRZKCmgFAmg2pgmzGmyxjTBLwOFIYwnimvpLadmtYerg9hE5HHquxk+gcNx86E19DM0oZOFqTPwOEQRISivBR2lzdjjPYbqPAWymTwB2CjiESKSDywATgWwnimvO1H63EIXLMkI9ShnJ+JHG79BmX1HRRkzDj/uCgvhbrLine truncated
"text/plain": [
"<Figure size 432x288 with 1 Axes>"
]
@@ -435,24 +417,86 @@
"source": [
"plt.plot(\n",
" [metrics[\"trainable_parameters\"] for metrics in enc_metrics],\n",
" [metrics[\"MSIS\"] for metrics in enc_metrics], \n",
" [metrics[\"NRMSE\"] for metrics in enc_metrics], \n",
")\n",
"plt.xlabel(\"trainable parameters\")\n",
"plt.ylabel(\"MSIS\")\n",
"plt.ylabel(\"NRMSE\")\n",
"plt.title(\"Encoder Scaling\")"
]
},
{
"cell_type": "markdown",
"id": "4dedbf76",
"id": "13356916",
"metadata": {},
"source": [
"### Decoder Scaling"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "a175e127",
"metadata": {},
"outputs": [],
"source": [
"dec_metrics = []\n",
"for layer in layers:\n",
" estimator = TransformerEstimator(\n",
" freq=freq,\n",
" prediction_length=prediction_length,\n",
" context_length=prediction_length*7,\n",
"\n",
" nhead=2,\n",
" num_encoder_layers=6,\n",
" num_decoder_layers=layer,\n",
" dim_feedforward=16,\n",
" activation=\"gelu\",\n",
"\n",
" num_feat_static_cat=1,\n",
" cardinality=[320],\n",
" embedding_dimension=[5],\n",
"\n",
" batch_size=128,\n",
" num_batches_per_epoch=100,\n",
" trainer_kwargs=dict(max_epochs=50, accelerator='auto', gpus=1),\n",
" )\n",
" \n",
" predictor = estimator.train(\n",
" training_data=train_ds,\n",
" validation_data=val_ds,\n",
" num_workers=8,\n",
" shuffle_buffer_length=1024\n",
" )\n",
" \n",
" forecast_it, ts_it = make_evaluation_predictions(\n",
" dataset=test_ds, \n",
" predictor=predictor\n",
" )\n",
" forecasts = list(forecast_it)\n",
" if layer == layers[0]:\n",
" tss = list(ts_it)\n",
" \n",
" evaluator = Evaluator()\n",
" agg_metrics, _ = evaluator(iter(tss), iter(forecasts))\n",
" agg_metrics[\"trainable_parameters\"] = summarize(estimator.create_lightning_module()).trainable_parameters\n",
" dec_metrics.append(agg_metrics.copy())"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "c439510e",
"metadata": {},
"outputs": [],
"source": [
"dec_metrics_out = open(\"elec_dec_metrics.pkl\", \"wb\")\n",
"pickle.dump(dec_metrics, dec_metrics_out)\n",
"dec_metrics_out.close()"
]
},
{
"cell_type": "markdown",
"id": "abeca54b",
"id": "78be6dd9",
"metadata": {},
"source": [
"### Symmetric Scaling"
@@ -461,10 +505,64 @@
{
"cell_type": "code",
"execution_count": null,
"id": "972f3216",
"id": "069b20d9",
"metadata": {},
"outputs": [],
"source": []
"source": [
"sym_metrics = []\n",
"for layer in layers:\n",
" estimator = TransformerEstimator(\n",
" freq=freq,\n",
" prediction_length=prediction_length,\n",
" context_length=prediction_length*7,\n",
"\n",
" nhead=2,\n",
" num_encoder_layers=layer,\n",
" num_decoder_layers=layer,\n",
" dim_feedforward=16,\n",
" activation=\"gelu\",\n",
"\n",
" num_feat_static_cat=1,\n",
" cardinality=[320],\n",
" embedding_dimension=[5],\n",
"\n",
" batch_size=128,\n",
" num_batches_per_epoch=100,\n",
" trainer_kwargs=dict(max_epochs=50, accelerator='auto', gpus=1),\n",
" )\n",
" \n",
" predictor = estimator.train(\n",
" training_data=train_ds,\n",
" validation_data=val_ds,\n",
" num_workers=8,\n",
" shuffle_buffer_length=1024\n",
" )\n",
" \n",
" forecast_it, ts_it = make_evaluation_predictions(\n",
" dataset=test_ds, \n",
" predictor=predictor\n",
" )\n",
" forecasts = list(forecast_it)\n",
" if layer == layers[0]:\n",
" tss = list(ts_it)\n",
" \n",
" evaluator = Evaluator()\n",
" agg_metrics, _ = evaluator(iter(tss), iter(forecasts))\n",
" agg_metrics[\"trainable_parameters\"] = summarize(estimator.create_lightning_module()).trainable_parameters\n",
" sym_metrics.append(agg_metrics.copy())"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "9d48d029",
"metadata": {},
"outputs": [],
"source": [
"sym_metrics_out = open(\"elec_sym_metrics.pkl\", \"wb\")\n",
"pickle.dump(sym_metrics, sym_metrics_out)\n",
"sym_metrics_out.close()"
]
}
],
"metadata": {