mirror of
https://github.com/wassname/attentive-neural-processes.git
synced 2026-10-04 12:20:51 +08:00
test set, inputnorm, lstm before encoder
This commit is contained in:
1 parent
38ab6cac23
commit
62c377b05f
13 files changed
+1395
-2392
No files matched your search
+196
-1579
File diff suppressed because it is too large.
Load diff
+369
-287
@@ -20,8 +20,8 @@
|
||||
"execution_count": 1,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:41:29.841115Z",
|
||||
"start_time": "2020-03-01T03:41:29.484738Z"
|
||||
"end_time": "2020-03-15T01:04:18.458822Z",
|
||||
"start_time": "2020-03-15T01:04:18.001753Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
@@ -36,8 +36,8 @@
|
||||
"execution_count": 2,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:41:31.805982Z",
|
||||
"start_time": "2020-03-01T03:41:29.843485Z"
|
||||
"end_time": "2020-03-15T01:04:20.666366Z",
|
||||
"start_time": "2020-03-15T01:04:18.463835Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
@@ -63,8 +63,8 @@
|
||||
"execution_count": 3,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:41:31.857758Z",
|
||||
"start_time": "2020-03-01T03:41:31.809871Z"
|
||||
"end_time": "2020-03-15T01:04:20.722137Z",
|
||||
"start_time": "2020-03-15T01:04:20.669364Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
@@ -79,17 +79,30 @@
|
||||
"execution_count": 4,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:41:31.973473Z",
|
||||
"start_time": "2020-03-01T03:41:31.860963Z"
|
||||
"end_time": "2020-03-15T01:04:20.857479Z",
|
||||
"start_time": "2020-03-15T01:04:20.725266Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/home/wassname/.pyenv/versions/jup3.7.3/lib/python3.7/site-packages/pytorch_lightning/core/decorators.py:13: UserWarning:\n",
|
||||
"\n",
|
||||
"data_loader decorator deprecated in 0.7.0. Will remove 0.9.0\n",
|
||||
"\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"from src.models.model import LatentModel\n",
|
||||
"from src.data.smart_meter import collate_fns, SmartMeterDataSet, get_smartmeter_df\n",
|
||||
"from src.plot import plot_from_loader\n",
|
||||
"from src.models.lstm_seqseq import LSTMSeq2Seq_PL\n",
|
||||
"from src.dict_logger import DictLogger"
|
||||
"from src.dict_logger import DictLogger\n",
|
||||
"from src.utils import PyTorchLightningPruningCallback\n",
|
||||
"from src.train import main, objective, add_number, run_trial"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -97,8 +110,8 @@
|
||||
"execution_count": 5,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:41:32.025844Z",
|
||||
"start_time": "2020-03-01T03:41:31.977741Z"
|
||||
"end_time": "2020-03-15T01:04:20.918561Z",
|
||||
"start_time": "2020-03-15T01:04:20.865187Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
@@ -113,8 +126,8 @@
|
||||
"execution_count": 6,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:41:32.075520Z",
|
||||
"start_time": "2020-03-01T03:41:32.028741Z"
|
||||
"end_time": "2020-03-15T01:04:20.972745Z",
|
||||
"start_time": "2020-03-15T01:04:20.922359Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
@@ -136,13 +149,13 @@
|
||||
"execution_count": 7,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:41:40.783929Z",
|
||||
"start_time": "2020-03-01T03:41:32.077822Z"
|
||||
"end_time": "2020-03-15T01:04:30.344555Z",
|
||||
"start_time": "2020-03-15T01:04:20.976031Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df_train, df_test = get_smartmeter_df()"
|
||||
"df_train, df_val, df_test = get_smartmeter_df()"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -150,15 +163,15 @@
|
||||
"execution_count": 8,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:41:41.497026Z",
|
||||
"start_time": "2020-03-01T03:41:40.786905Z"
|
||||
"end_time": "2020-03-15T01:04:31.128105Z",
|
||||
"start_time": "2020-03-15T01:04:30.346646Z"
|
||||
}
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"<matplotlib.legend.Legend at 0x7fb05e5d2668>"
|
||||
"<matplotlib.legend.Legend at 0x7f8bde1390f0>"
|
||||
]
|
||||
},
|
||||
"execution_count": 8,
|
||||
@@ -167,7 +180,7 @@
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAXoAAAEGCAYAAABrQF4qAAAABHNCSVQICAgIfAhkiAAAAAlwSFlzAAALEgAACxIB0t1+/AAAADh0RVh0U29mdHdhcmUAbWF0cGxvdGxpYiB2ZXJzaW9uMy4xLjEsIGh0dHA6Ly9tYXRwbG90bGliLm9yZy8QZhcZAAAgAElEQVR4nO2dd5gUVdaH30OOIlEF1EFERVBREWVxTRhAVl0/V9fsurrornGDirsqCAZWd13XHDGwiiIGQJCkICpxiJJzGIacGWZgwvn+qOqhZqZnumemuru657zP089033ur6tc11adunXvuuaKqGIZhGKlLtUQLMAzDMGKLGXrDMIwUxwy9YRhGimOG3jAMI8UxQ28YhpHimKE3DMNIcczQG0YpiMhdIvKi+z5NRFREaviw3wtEJKMc7d8XkafKqFcROb6UuitE5NOK6DRSBzP0hhEGEakFPAY8H0XbsSLyiOdzK9f4his7shwalorICeXV7kVVRwIdROTUyuzHSG7M0BspiQ8976uAJaq6IYq2k4HzPJ/PA5aEKVuuqpuiObiItAWqq+qyKPWWxRCgtw/7MZIUM/RG3BGRliLyuYhsFZHVInK/W95PRIaKyIcisldEFopI50jbebYdJiL/E5E9wO9EpK6IfCAiO0VksYg8HHKZiMhDIvJ5MV0vich/3Y89ge/L+A7XiMgaEemIY+i7iUjo9/RL4EWgc7GyycX28VcR2SIiG0Xk9mKH6AWM9nxuLCKj3PMy3b0ReLlYRJaLyC4ReVVExFM3yd2fUUUxQ2/EFdfwjQTmAa2A7sCDInKZ2+RK4BPgcGAE8EqU24HTCx/mbvsR0BdIA44DLgFu9rT9H9BDRA53918DuB740K0/BVhayne4HfgncLGqLgBmALWB09wm5wHjgRXFyryG/kigkftd7gBeFZHGnvrLgVGez9cDTwKN3f0+XUzWr4CzgFOB6wDveVkMpInIYeG+j5H6mKE34s1ZQHNV7a+qB1V1FfA2jiED+FFVR6tqPjCYQ4Yy0nYAU1X1K1UtUNVsHIP3jKruVNUM4KVQQ1XdiGN4r3WLegDbVHWW+/lwYG8Y/Q8CDwEXqOoKd18HgOnAeSLSBGjk6vvBU3YyRZ8QcoH+qpqrqqOBfcCJACJSz/2+kzztv1TVGaqah3MT61RM10BV3aWq64CJxepD3+PwMN/HqAJUOoLAMMrJsUBLEdnlKauOYxTXAl4f9n6gjtvbLmu7EOuLHatlsbLi9R8Af8S5YdyMc2MJsRNoGEb/QzgGunjUTMhPvwb4yS37EbjdLVuvqms97be7RjvEfqCB+747MMW9gYQofl4aUJSy6kPfw3vujCqE9eiNeLMeWK2qh3teDVX1ch+2K56KdSPQ2vP56GL1XwGnun72X+H0lEPMB8JFvFwKPCYi1xQrn4zjhz+PQzefn4BulHTbROJyivrnK0t7YI2q7vFxn0YSYYbeiDczgL0i8og7WFpdRDqKyFkx2G4o8KiINBaRVsC93kpVzcHx6X8MzHDdHiFGA+eH2edCHDfPqyJypad8Ko5r5GZcQ6+qO4Gtbll5DH1PivrnK8v5wDc+7s9IMszQG3HF9b3/CseHvBrYBryDMzDp93b9gQy3/QQco36gWJsPcAZeBxcrHwmcJCItw2iZ52p5W0R6umVZwCygFrDA0/wHoAVRGnr36WJfsZtOZbkBeNPH/RlJhtjCI0ZVQUT+CFyvqud7yo7BiXk/srhrQ0R6Ayer6oNx1Pgw0ExVH/Zpf1cAt6jqdX7sz0hOzNAbKYuIHIUTWjkVaIfjDnlFVUNpDaoBLwCHqervEybUg4hcB/ysqosTrcVIHczQGymLiByLY9zb4EScfAI8qqoHRaQ+sBkn0qeHqhaPyDGMlMEMvWEYRopjg7GGYRgpTiAnTDVr1kzT0tISLcMwDCNpmDVr1jZVbR6uLpCGPi0tjfT09ETLMAzDSBpEZG1pdea6MQzDSHHM0BuGYaQ4EQ29iBwtIhNFZJGbH/yBMG3EzeW9QkTmi8gZnrrb3DzZy0XkNr+/gGEYhlE20fjo84C/qupsEWkIzBKR8aq6yNOmJ86ElHbA2cDrwNlueta+QGechFOzRGSEmwPEMAzDN3Jzc8nIyCAnJyfRUmJKnTp1aN26NTVr1ox6m4iG3s3bvdF9v1dEFuMsluA19FcBH6oTlD9NRA53ZyVeAIxX1R0AIjIeJyHUkKgVGoZhREFGRgYNGzYkLS2NogtspQ6qyvbt28nIyKBNmzZRb1cuH72IpAGn4yyy4KUVRXN9Z7hlpZUbhmH4Sk5ODk2bNk1ZIw8gIjRt2rTcTy1RG3oRaQB8DjwYi7zWItJbRNJFJH3r1q1+7x6A9DU7yC+wmcCGkaqkspEPUZHvGJWhF5GaOEb+I1X9IkyTDRRd1KG1W1ZaeQlU9S1V7ayqnZs3DxvzXymmrdrOb96Yyhvfr/R934ZhGEEmmqgbAd4FFqvqC6U0GwHc6kbfnAPsdn37Y4FL3YUfGuOszjPWJ+3lYuPubACWbw63DKhhGEbl2LVrF6+99lq5t7v88svZtSu2qzxG06PvBtwCXCQic93X5SJyt4jc7bYZDazCWZ3+beBPAO4g7ABgpvvqHxqYNQzDSCVKM/R5eXlhWh9i9OjRHH54bNdtjybq5kegTKeQG21zTyl1g4BBFVJnGIaRJPTp04eVK1fSqVMnatasSZ06dWjcuDFLlixh2bJl/PrXv2b9+vXk5OTwwAMP0Lt3b+BQypd9+/bRs2dPzj33XKZMmUKrVq0YPnw4devWrbS2QOa6MQzDqAxPjlzIokx/Y0ZObnkYfa/oUGr9wIEDWbBgAXPnzmXSpEn06tWLBQsWFIZBDho0iCZNmpCdnc1ZZ53FNddcQ9OmTYvsY/ny5QwZMoS3336b6667js8//5ybb7650toDaejzLDLGMIwkp0uXLkVi3V966SW+/PJLANavX8/y5ctLGPo2bdrQqVMnAM4880zWrFnji5ZAGvrFG32P3jQMowpRVs87XtSvX7/w/aRJk5gwYQJTp06lXr16XHDBBWFj4WvXrl34vnr16mRnZ/uixZKaGYZh+EDDhg3Zuzd8VN/u3btp3Lgx9erVY8mSJUybNi2u2gLZo48FtmKiYRixpGnTpnTr1o2OHTtSt25djjjiiMK6Hj168MYbb9C+fXtOPPFEzjnnnLhqqzKGPkRVmDlnGEZi+Pjjj8OW165dm2+++SZsXcgP36xZMxYsWFBY/re//c03XVXGdfOXofMSLcEwDCMhVBlDbxiGUVVJSUN/5oDxXPqf7xMtwzAMIxCkpI9+e9ZBtmcdTLQMwzCMQJCSPXrDMAzjEGboDcMwUpwqZ+gtuNIwjFhQ0TTFAC+++CL79+/3WdEhqpyhj5alm/bSqf84tuxJ7YWGDcPwhyAb+oiDsSIyCPgVsEVVO4apfwi4ybO/9kBzVd0hImuAvUA+kKeqnf0SXmGi7NK/99Nqdu3P5dslW7ihyzGx1WQYRtLjTVN8ySWX0KJFC4YOHcqBAwe4+uqrefLJJ8nKyuK6664jIyOD/Px8Hn/8cTZv3kxmZiYXXnghzZo1Y+LEib5riybq5n3gFeDDcJWq+jzwPICIXAH8udjiIheq6rZK6jQMw4ieb/rApp/93eeRp0DPgaVWe9MUjxs3jmHDhjFjxgxUlSuvvJLJkyezdetWWrZsyahRowAnB06jRo144YUXmDhxIs2aNfNXs0tE142qTgaiXRXqBmBIpRQZhmEkOePGjWPcuHGcfvrpnHHGGSxZsoTly5dzyimnMH78eB555BF++OEHGjVqFBc9vsXRi0g9oAdwr6dYgXEiosCbqvpWGdv3BnoD1DryeL9kGYZRFSmj5x0PVJVHH32Uu+66q0Td7NmzGT16NI899hjdu3fniSeeiLkePwdjrwB+Kua2OVdVzwB6AveIyHmlbayqb6lq50D48T1Y1kvDMKLBm6b4sssuY9CgQezbtw+ADRs2sGXLFjIzM6lXrx4333wzDz30ELNnzy6xbSzwc2bs9RRz26jqBvfvFhH5EugCTPbxmDHDklwahlEevGmKe/bsyY033kjXrl0BaNCgAf/73/9YsWIFDz30ENWqVaNmzZq8/vrrAPTu3ZsePXrQsmXLhA3GRkREGgHnAzd7yuoD1VR1r/v+UqC/H8czDMMIIsXTFD/wwANFPrdt25bLLrusxHb33Xcf9913X8x0RRNeOQS4AGgmIhlAX6AmgKq+4Ta7GhinqlmeTY8AvnTzv9cAPlbVMf5JNwzDMKIhoqFX1RuiaPM+Thimt2wVcFpFhcUKsbmxhmFUMWxmrGEYKYNWgeiJinxHM/QRUFL/wjGMVKBOnTps3749pY29qrJ9+3bq1KlTru0Cm4/+hrem0b19C+785XGV2s/Q9PX8sNwm5hpGqtO6dWsyMjLYunVroqXElDp16tC6detybRNYQz911XamrtpeKUM/d/0uHh42v1I6zKdvGMlBzZo1adOmTaJlBJKkdt20/fto0vqMYuveA2HrLfOkYRhGkhv6/ALHF/fh1DVRb2MToQzDqGoE1nWTKLbtO8DenLzCzzYYaxhGspPUPXq/2JuTS4H7dHDOM99y4b8mESlx/ZptWSk9um8YRupQ5Qx9cfO9M+sgp/Qbx4vfLgcgryCy8Z61dgcX/GsSH01fFwOFhmEY/pIShr4ybvftWc5A7qj5mVFvs3Krk+lh7vpdlTiyYRhGfEhKH/3a7Vks37wvYrvS+uZ7cnLZuCuHE49sWFi2abdF6BiGkZokpaE///lJldr++jensWjjHtYM7FVYlnUwP2xbc8MbhpHspITrprws2rgnYhsLwzQMI1VIDUNvVtkwDKNUIhp6ERkkIltEZEEp9ReIyG4Rmeu+nvDU9RCRpSKyQkT6+Ck81kTjsjG3jmEYyUA0Pfr3cRb9LosfVLWT++oLine truncated
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAXoAAAEGCAYAAABrQF4qAAAABHNCSVQICAgIfAhkiAAAAAlwSFlzAAALEgAACxIB0t1+/AAAADh0RVh0U29mdHdhcmUAbWF0cGxvdGxpYiB2ZXJzaW9uMy4xLjEsIGh0dHA6Ly9tYXRwbG90bGliLm9yZy8QZhcZAAAgAElEQVR4nO2dd3gVVfrHP2+QKr1JE4KAYEEREXBRsdFERdeG7ee6uqhrd3WFXbuuou66rKuouLJiQWSxAII0BUF6QKr0HmooCaEkkOT9/TGTMElukptkbsnN+3me+2TmnDNnvncy950z57znPaKqGIZhGLFLXKQFGIZhGKHFDL1hGEaMY4beMAwjxjFDbxiGEeOYoTcMw4hxzNAbhmHEOGboDaMAROQ+ERnibseLiIrIST7Ue6mIJBaj/Mci8koh+SoirQvIu0ZEviyJTiN2MENvGAEQkUrAM8CbQZSdLCJPe/abusY3UFqjYmhYIyKnF1e7F1UdD5wlIueUph6jbGOG3ohJfGh59wNWq+r2IMrOBC7x7F8CrA6Qtk5VdwVzchFpBVRQ1bVB6i2ML4ABPtRjlFHM0BthR0SaiMhXIpIkIptE5BE3/QURGS0in4hIqoisFJFORR3nOXaMiHwmIgeB34lIVREZISIHRGSViPw5u8tERJ4Ska/y6HpbRP7l7vYBfirkO9wgIptF5GwcQ99NRLJ/TxcDQ4BOedJm5qnjTyKyR0R2isjdeU7RF5jo2a8jIhPc6zLffRB4uVJE1olIsoi8KyLiyZvh1meUU8zQG2HFNXzjgaVAU+AK4DER6eUWuRYYBdQGxgHvBHkcOK3wMe6xnwPPA/HAaUAP4A5P2c+A3iJS263/JKA/8Imb3x5YU8B3uBt4HbhSVVcAC4DKwLlukUuAqcD6PGleQ98IqOV+l3uAd0Wkjif/KmCCZ78/8CJQx633b3lkXQ1cAJwD3Ax4r8sqIF5Eagb6PkbsY4beCDcXAA1U9SVVPaaqG4EPcQwZwM+qOlFVM4FPOWEoizoOYK6qfquqWap6FMfgvaqqB1Q1EXg7u6Cq7sQxvDe5Sb2Bvaq6yN2vDaQG0P8Y8BRwqaqud+tKB+YDl4hIXaCWq2+WJ+1Mcr8hHAdeUtXjqjoROAS0BRCRau73neEp/42qLlDVDJyHWIc8ugararKqbgWm58nP/h61A3wfoxxQag8CwygmLYAmIpLsSauAYxS3AN4+7CNAFbe1Xdhx2WzLc64medLy5o8AHsB5YNyB82DJ5gBQI4D+p3AMdF6vmex++s3AbDftZ+BuN22bqm7xlN/nGu1sjgDV3e0rgDnuAySbvNelOrkpLD/7e3ivnVGOsBa9EW62AZtUtbbnU0NVr/LhuLyhWHcCzTz7p+bJ/xY4x+1nvxqnpZzNMiCQx0tP4BkRuSFP+kycfvhLOPHwmQ10I3+3TVFcRe7++dJyBrBZVQ/6WKdRhjBDb4SbBUCqiDztDpZWEJGzReSCEBw3GhgkInVEpCnwkDdTVdNw+vRHAgvcbo9sJgLdA9S5Eqeb510RudaTPhena+QOXEOvqgeAJDetOIa+D7n750tLd+B7H+szyhhm6I2w4va9X43Th7wJ2Av8B2dg0u/jXgIS3fLTcIx6ep4yI3AGXj/Nkz4eaCciTQJoWepq+VBE+rhph4FFQCVghaf4LKAhQRp69+3iUJ6HTmm5FfjAx/qMMobYwiNGeUFEHgD6q2p3T1pzHJ/3Rnm7NkRkAHCmqj4WRo1/Buqr6p99qu8a4E5VvdmP+oyyiRl6I2YRkcY4rpVzgTY43SHvqGp2WIM44C2gpqr+PmJCPYjIzcByVV0VaS1G7GCG3ohZRKQFjnFvieNxMgoYpKrHRORkYDeOp09vVc3rkWMYMYMZesMwjBjHBmMNwzBinKicMFW/fn2Nj4+PtAzDMIwyw6JFi/aqaoNAeVFp6OPj40lISIi0DMMwjDKDiGwpKM+6bgzDMGIcM/SGYRgxTpGGXkROFZHpIvKrGx/80QBlxI3lvV5ElolIR0/eXW6c7HUicpffX8AwDMMonGD66DOAP6nqYhGpASwSkamq+qunTB+cCSltgC7Ae0AXNzzr80AnnIBTi0RknBsDxDAMwzeOHz9OYmIiaWlpkZYSUqpUqUKzZs2oWLFi0McUaejduN073e1UEVmFs1iC19D3Az5Rxyl/nojUdmclXgpMVdX9ACIyFScg1BdBKzQMwwiCxMREatSoQXx8PLkX2IodVJV9+/aRmJhIy5Ytgz6uWH30IhIPnIezyIKXpuSO9Z3ophWUbhiG4StpaWnUq1cvZo08gIhQr169Yr+1BG3oRaQ68BXwWCjiWovIABFJEJGEpKQkv6svOTt+gWOHI63CMIwgiGUjn01JvmNQhl5EKuIY+c9V9esARbaTe1GHZm5aQen5UNVhqtpJVTs1aBDQ5z/8HE2GYZfCV3+ItBLDMIwSE4zXjQAfAatU9a0Cio0D/s/1vukKpLh9+5OBnu7CD3VwVueZ7JP20JPhvh5tt8lbhmEUTnJyMkOHDi32cVdddRXJyaFd5TGYFn034E7gchFZ4n6uEpH7ReR+t8xEYCPO6vQfAn8EcAdhXwYWup+XsgdmDcMwYomCDH1GRkaA0ieYOHEitWuHdt32YLxufgYK7RRyvW0eLCBvODC8ROoMwzDKCAMHDmTDhg106NCBihUrUqVKFerUqcPq1atZu3Yt1113Hdu2bSMtLY1HH32UAQMGACdCvhw6dIg+ffpw0UUXMWfOHJo2bcrYsWOpWrVqqbVFZawbwzCM0vDi+JX8usNfn5Ezm9Tk+WvOKjB/8ODBrFixgiVLljBjxgz69u3LihUrctwghw8fTt26dTl69CgXXHABN9xwA/Xq1ctVx7p16/jiiy/48MMPufnmm/nqq6+44447Sq09tkIgHDsCKQHHeg3DMMJK586dc/m6v/3225x77rl07dqVbdu2sW7dunzHtGzZkg4dOgBw/vnns3nzZl+0xFaL/vMbYctseCEl0koMw4gghbW8w8XJJ5+csz1jxgymTZvG3LlzqVatGpdeemlAX/jKlSvnbFeoUIGjR4/6oiW2WvRbZkdagWEY5ZQaNWqQmpoaMC8lJYU6depQrVo1Vq9ezbx588KqLbZa9IZhGBGiXr16dOvWjbPPPpuqVatyyimn5OT17t2b999/nzPOOIO2bdvStWvXsGozQ18Ytp6uYRjFYOTIkQHTK1euzPfffx8wL7sfvn79+qxYsSIn/cknn/RNV2x13fjN19kzYmN/WrVhGLGLGfrC2Dwr0goMwzBKjRn6bJK3wQu1YNn/Iq3EMAzDV8zQZ7PHDa+/fHRkdRiGYfiMGXrDMIwYxwy9YRhGjGOG3jAMIwJUr149bOcyQ+8X+zfC4BZwYHOklRiGYeQimIVHhovIHhFZUUD+U5449StEJFNE6rp5m0VkuZsX26t3LBkJacmwzAZzDaM8MnDgQN59992c/RdeeIFXXnmFK664go4dO9K+fXvGjh0bEW3BzIz9GHgH+CRQpqq+CbwJICLXAI/nWVzkMlXdW0qd4cNmwxpG2ef7gbBrub91NmoPfQYXmH3LLbfw2GOP8eCDztIco0ePZvLkyTzyyCPUrFmTvXv30rVrV6699tqwr20bzMIjM0UkPsj6bgW+KI2gyFHIhS8HCw4bhlE6zjvvPPbs2cOOHTtISkqiTp06NGrUiMcff5yZM2cSFxfH9u3b2b17N40aNQqrNt9i3YhINaA38JAnWYEpIqLAB6o6rJDjBwADAJo3b+6XLMMwyiOFtLxDyU033cSYMWPYtWsXt9xyC59//jlJSUksWrSIihUrEh8fHzA8cajxczD2GmB2nm6bi1S1I9AHeFBELinoYFUdpqqdVLVTgwYNfJTlA9adYxhGENxyyy2MGjWKMWPGcNNNN5GSkkLDhg2pWLEi06dPZ8uWLRHR5aeh70+ebhtV3e7+3QN8A3T28XyGYRhRxVlnnUVqaipNmzalcePG3H777SQkJNC+fXs++eQT2rVrFxFdvnTdiEgtoDtwhyftZCBOVVPd7Z7AS36czzAMI1pZvvzEIHD9+vWZO3duwHKHDh0Kl6SiDb2IfAFcCtQXkUTgeaAigKq+7xa7Hpiiqoc9h54CfOOOLp8EjFTVSf5JDxXWTWMYRmwRjNfNrUGU+RjHDdObthE4t6TCwo551hiGEaPYzFi/sYFbwzCiDDP0vmFvBIZhRCexaeg/6Qez3y7ZsSmJzuf9i06kWbeOYRhlmNg09BtnwNRnS3Zs0mqY+WYJpk9bl41hGNFJbBr6ovj8JmfZQL9jYYC1/g2jnJKcnMzQoUNLdOyQIUM4cuSIz4pOUD4N/bopzt/xj0ZWh2EYMUM0G3rfYt2UfUrYEs84Fpo3A8MwyhQDBw5kw4YNdOjQgR49etCwYUNGjx5Neno6119/PS+++CKHDx/m5ptvJjExkczMTJ599ll2797Njh07uOyyy6hfvz7Tp0/3XZsZ+pKQlQnHj0Ll6s5YwPz34YxrCy5/eB/ExUHVOuHTaBjlmNcXvM7q/at9rbNd3XY83fnpAvMHDx7MihUrWLJkCVOmTGHMmDEsWLAAVeXaa69l5syZJCUl0aRJEyZMmABASkoKtWrV4q233mL69OnUr1/fV83ZlM+um9Iy9iF4ramzvXOp8/fI/oLLv3kavB4fclmGYUQHU6ZMYcqUKZx33nl07NiR1atXs27dOtq3b8/UqVN5+umnmTVrFrVq1QqLHmvRl4SlIyOtwDCMQiis5R0OVJVBgwZx33335ctbvHgxEydO5JlnnuGKK67gueeeC7me8tWiz8qChR+V7NhDe2DfhqLL2cxYwyiX1KhRg9TUVAB69erF8OHDcwKXbd++PWdRkmrVqnHHHXfw1FNPsXjx4nzHhoLy1aJf+gVMeKIEBwr8/XRA4YWUE8nbFnqKmFuLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
@@ -181,140 +194,12 @@
|
||||
"source": [
|
||||
"# Show split\n",
|
||||
"df_train['energy(kWh/hh)'].plot(label='train')\n",
|
||||
"df_val['energy(kWh/hh)'].plot(label='val')\n",
|
||||
"df_test['energy(kWh/hh)'].plot(label='test')\n",
|
||||
"plt.title('energy(kWh/hh)')\n",
|
||||
"plt.legend()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-02-16T00:24:15.905017Z",
|
||||
"start_time": "2020-02-16T00:24:15.747859Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"# Train helpers"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 9,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:41:41.584788Z",
|
||||
"start_time": "2020-03-01T03:41:41.503778Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def main(trial, train=True):\n",
|
||||
" # PyTorch Lightning will try to restore model parameters from previous trials if checkpoint\n",
|
||||
" # filenames match. Therefore, the filenames for each trial must be made unique.\n",
|
||||
" \n",
|
||||
" checkpoint_callback = pl.callbacks.ModelCheckpoint(\n",
|
||||
" os.path.join(MODEL_DIR, name, 'version_{}'.format(trial.number), \"chk\"), monitor='val_loss', mode='min')\n",
|
||||
"\n",
|
||||
" # The default logger in PyTorch Lightning writes to event files to be consumed by\n",
|
||||
" # TensorBoard. We create a simple logger instead that holds the log in memory so that the\n",
|
||||
" # final accuracy can be obtained after optimization. When using the default logger, the\n",
|
||||
" # final accuracy could be stored in an attribute of the `Trainer` instead.\n",
|
||||
" logger = DictLogger(MODEL_DIR, name=name, version=trial.number)\n",
|
||||
"# print(\"log_dir\", logger.experiment.log_dir)\n",
|
||||
" hparams = dict(**trial.params, **trial.user_attrs)\n",
|
||||
"\n",
|
||||
" trainer = pl.Trainer(\n",
|
||||
" logger=logger,\n",
|
||||
" val_percent_check=PERCENT_TEST_EXAMPLES,\n",
|
||||
" checkpoint_callback=checkpoint_callback,\n",
|
||||
" max_epochs=hparams['max_nb_epochs'],\n",
|
||||
" gpus=-1 if torch.cuda.is_available() else None,\n",
|
||||
" early_stop_callback=PyTorchLightningPruningCallback(trial, monitor='val_loss')\n",
|
||||
" )\n",
|
||||
" \n",
|
||||
" model = LSTMSeq2Seq_PL(hparams)\n",
|
||||
" if train:\n",
|
||||
" trainer.fit(model)\n",
|
||||
" return model, trainer"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 10,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:41:41.652356Z",
|
||||
"start_time": "2020-03-01T03:41:41.588888Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def add_suggest(trial):\n",
|
||||
" trial.suggest_loguniform(\"learning_rate\", 1e-5, 1e-2)\n",
|
||||
" trial.suggest_uniform(\"lstm_dropout\", 0, 0.75)\n",
|
||||
" trial.suggest_categorical(\"hidden_size\", [1, 2, 4, 8, 16, 32, 64, 128]) \n",
|
||||
" trial.suggest_categorical(\"lstm_layers\", [1, 2, 4, 8]) \n",
|
||||
" trial.suggest_categorical(\"bidirectional\", [False, True]) \n",
|
||||
" \n",
|
||||
"\n",
|
||||
" trial._user_attrs = {\n",
|
||||
" 'batch_size': 16,\n",
|
||||
" 'grad_clip': 40,\n",
|
||||
" 'max_nb_epochs': 200,\n",
|
||||
" 'num_workers': 4,\n",
|
||||
" 'num_extra_target': 24*4,\n",
|
||||
" 'vis_i': '670',\n",
|
||||
" 'num_context': 24*4,\n",
|
||||
" 'input_size': 18,\n",
|
||||
" 'input_size_decoder': 17,\n",
|
||||
" 'context_in_target': True,\n",
|
||||
" 'output_size': 1\n",
|
||||
" }\n",
|
||||
" \n",
|
||||
" # For manual experiment we will start at -1 and deincr by 1\n",
|
||||
" versions = [int(s.stem.split('_')[-1]) for s in (MODEL_DIR / name).glob('version_*')] + [-1]\n",
|
||||
" trial.number = min(versions)-1\n",
|
||||
" print('trial.number', trial.number)\n",
|
||||
" return trial"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 11,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:41:41.723534Z",
|
||||
"start_time": "2020-03-01T03:41:41.655316Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"\n",
|
||||
"def objective(trial):\n",
|
||||
" # see https://github.com/optuna/optuna/blob/cf6f02d/examples/pytorch_lightning_simple.py\n",
|
||||
" trial = add_suggest(trial)\n",
|
||||
"\n",
|
||||
" \n",
|
||||
" print('trial', trial.number, 'params', trial.params)\n",
|
||||
" \n",
|
||||
" model, trainer = main(trial)\n",
|
||||
" \n",
|
||||
" # also report to tensorboard & print\n",
|
||||
" print('logger.metrics', model.logger.metrics[-1:])\n",
|
||||
" model.logger.experiment.add_hparams(trial.params, logger.metrics[-1])\n",
|
||||
" model.logger.save()\n",
|
||||
" \n",
|
||||
" return model.logger.metrics[-1]['val_loss']\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
@@ -322,47 +207,6 @@
|
||||
"# Train"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 12,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:41:41.777701Z",
|
||||
"start_time": "2020-03-01T03:41:41.726164Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"PERCENT_TEST_EXAMPLES = 0.3\n",
|
||||
"EPOCHS = 2\n",
|
||||
"DIR = Path(os.getcwd())\n",
|
||||
"MODEL_DIR = DIR/ 'optuna_result'/ 'lstm'\n",
|
||||
"name = \"lstm1\"\n",
|
||||
"MODEL_DIR.mkdir(parents=True, exist_ok=True)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 13,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:41:41.829950Z",
|
||||
"start_time": "2020-03-01T03:41:41.779861Z"
|
||||
}
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"now run `tensorboard --logdir /media/wassname/Storage5/projects2/3ST/attentive-neural-processes/optuna_result/lstm\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"print(f\"now run `tensorboard --logdir {MODEL_DIR}\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
@@ -372,64 +216,317 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"execution_count": 9,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-02-01T12:52:22.492191Z",
|
||||
"start_time": "2020-02-01T12:52:22.416653Z"
|
||||
"end_time": "2020-03-15T01:51:46.861012Z",
|
||||
"start_time": "2020-03-15T01:04:31.130838Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-02-16T01:47:59.598921Z",
|
||||
"start_time": "2020-02-16T01:47:02.100Z"
|
||||
}
|
||||
},
|
||||
"source": [
|
||||
"Note that the LSTM has access to the y values for the first half of the plot (the context) to match the NP setup."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"start_time": "2020-03-01T03:47:34.500Z"
|
||||
},
|
||||
"scrolled": true
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"trial.number -21\n"
|
||||
"now run `tensorboard --logdir lightning_logs`\n",
|
||||
"trial.number -8\n",
|
||||
"INFO:root:GPU available: True, used: True\n",
|
||||
"INFO:root:VISIBLE GPUS: 0\n",
|
||||
"INFO:root:\n",
|
||||
" | Name | Type | Params\n",
|
||||
"-----------------------------------------------------------------\n",
|
||||
"0 | model | Seq2SeqNet | 8 M \n",
|
||||
"1 | model.norm_input | BatchNormSequence | 36 \n",
|
||||
"2 | model.norm_input.norm | BatchNorm1d | 36 \n",
|
||||
"3 | model.encoder | LSTM | 3 M \n",
|
||||
"4 | model.multihead_attn | MultiheadAttention | 263 K \n",
|
||||
"5 | model.multihead_attn.out_proj | Linear | 65 K \n",
|
||||
"6 | model.norm_target | BatchNormSequence | 34 \n",
|
||||
"7 | model.norm_target.norm | BatchNorm1d | 34 \n",
|
||||
"8 | model.decoder | LSTM | 3 M \n",
|
||||
"9 | model.mean | Linear | 257 \n",
|
||||
"10 | model.std | Linear | 257 \n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"data": {
|
||||
"application/vnd.jupyter.widget-view+json": {
|
||||
"model_id": "",
|
||||
"version_major": 2,
|
||||
"version_minor": 0
|
||||
},
|
||||
"text/plain": [
|
||||
"HBox(children=(FloatProgress(value=0.0, description='Validation sanity check', layout=Layout(flex='2'), max=5.…"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYIAAAEWCAYAAABrDZDcAAAABHNCSVQICAgIfAhkiAAAAAlwSFlzAAALEgAACxIB0t1+/AAAADh0RVh0U29mdHdhcmUAbWF0cGxvdGxpYiB2ZXJzaW9uMy4xLjEsIGh0dHA6Ly9tYXRwbG90bGliLm9yZy8QZhcZAAAgAElEQVR4nOydd3hUxdeA30k2vRcIIZ3ee+8gImIBBVQEREAQBSyIioBYUQQUUVApSpEI+ENB5AMLKGooQkB6JyGBEEgnvWx2vj/uZknIJoQku5uQ+z7PPuzO3Dtzcne5586ZU4SUEhUVFRWVmouVpQVQUVFRUbEsqiJQUVFRqeGoikBFRUWlhqMqAhUVFZUajqoIVFRUVGo4qiJQUVFRqeGoiqAGI4SQQogGlpZDRUXFsqiKQKVKIYSwE0J8I4RIFUJcE0JMK+XYB4QQYUKIFP2xK4UQLmUZSwhhK4TYJIS4pFeIfW4Zu68Q4k8hxA0hxKUyyF3i8UKI2kKI9UKIq/r+PUKIzrcZ7z0hxHEhhFYI8fYtfUIIMUsIEa3/2zYIIVxLGctTCLFZCJEhhIgSQjx5S/9UIUSkfqxwIUSPUsYSQoiPhBCJ+tdHQghRqL+NEOKQECJT/28bc4ylUjFURaBS1XgbaAgEAX2B14QQA0s41g14H6gLNAX8gAV3MFYYMAq4ZmTsDOAb4NUyyl3a8c7AQaA94AmsAf5PCOFcyngXgNeA/zPS9xQwGuiO8rc7AJ+XMtZSIBfwAUYCXwohmgPoFdI8YBjK9fwa2CyEsC5hrInAEKA10Ap4CHhWP5Yt8BOwDvDQ/50/6dtNPZZKRZBSqq8a+gIk0ED/3g1YC8QDUcBswErf1wD4C7gBJAAb9e0CWATEAanAcaBFBWW6Cgwo9Pk9YEMZz30UOH6nYwFXgD4ljNkfuHQH8pfpeP31al+G49YBb9/Stgl4tdDnbkA24GjkfCcUJdCoUNu3wDz9+8eBA7ccLwHfEuTZC0ws9Hk8sF//fgAQA4hC/dHAQFOPpb4q9lJXBCoFfI6iDOoBvVGeOsfq+94DfkN5MvPn5tPnAKAX0Eh/7mNAIoAQYobeZGP0ZUwAIYQH4AscLdR8FGhexr+hF3CyksYyGXoThy3KU3+5h7nlvR3K6udWGgFaKeW5Qm2Fr8MOwFoI0Vm/ChgHHEG/ShJCPCmEOFbo3OaUfE2bA8ek/q6t51hBf2WOpVK5aCwtgIrl0d8AngDaSCnTgDQhxMco5oevgTwU80pdKeUVFJMK+nYXoAnKU+XpgjGllPNQTA53QoGp5Eahthv6OW73N9wLjAEKbO/lHsuU6G353wLvSClv3O74EvgFxcz1PZAMvK5vdzRyrDPK6qMwha9DGvADyncqgBTg/oIbsJTyO+C7W8a79Zo66237t/YVmasyx1KpXNQVgQqAN2CDYhIqIArF5g6KrVoAB4QQJ4UQ4wCklH8AS1Bs0HFCiOWlbVqWgXT9v4XHcEW5WZWIEKILyg1mWKEn33KNVRaEEDOFEOn611d3cJ4D8DOK+ePDQu0nC43XswxDfQOsB3ajrID+1LdfMXJsOkWvARS9DuNRVn7NUVYpo4BtQoi6Jcx963iuQLpecdxuLlOOpVIBVEWgAordv+Cpv4BAFBstUsprUsoJUsq6KJt5Xwi926mU8jMpZXugGYoZ4lUodrMs9jImhJQyGYhF2TwsoDV6c48xhBBtga3AOCnlroqMVVaklB9IKZ31r0llOUcIYQdsQblZP3vLeM0LjfdPGebXSSnfklIGSyn9Uf6mGP3rVs4BGiFEYbNR4evQBtgmpTynH/cXlOvWrYTpT1LyNT0JtCrs+YOyCVzSNa/MsVQqgKoIVJBS5gPfA3OFEC5CiCBgGspGJUKI4UIIf/3hySibiTohREe9bdkGxWsmG9Dpxyx8syz2KkWctcBsIYSHEKIJMAFYbexAIUQLFDPJVCnlz3c6llDcS+31H22FEPYFNx4hhJW+z0b5KOxL81gp7Xj99dkEZAFjpJS6Uv7+gvFs9ONZodzI7Qs8eYTiDlpf737ZDPgEeNfYuFLKDOBH4F0hhJMQojswGMU8BYo30wNCiHr68e5FUegnShBtLTBNCOGnXzW8ws1ruhvIB17QX9sp+vY/zDCWSkWw9G61+rLci6JeQx4oN/544DIwh5teQ/NRnjbTgYvoPT2Ae1A28NJRVhWhgHMFZbJDMX2kAteBaaUcuwpF8aQXep0s61jAJf01KPwK1vf1MdK3uxRZSjweZfNdApm3yNqzlPFWGxnvaX1fI+Csfryo0q6R/nhPlNVIBornzZOF+gTwrr49DTgNjC7UP/KWayr0v4ck/Ws+RT172gKHUJTeYaCtKcZSX5X7EvoLrqKioqJSQ1FNQyoqKio1HFURqKioqNRwVEWgoqKiUsNRFYGKiopKDUdVBCoqKio1nGqXYsLb21sGBweX61ytFnS39eBWUVFRqXoIATY25T//0KFDCVLKWsb6qp0iCA4OJjw8vFznJiZCmhqgrqKiUg2xswNf3/KfL4SIKqlPNQ2pqKio1HBURaCioqJSw1EVgYqKikoNR1UEKioqKjUckykCoRQNjxNClJTFsOC4jkIp0D3MVLIYY8uWUHr0CKZePSt69Ahmy5ZQc06voqKiUmUw5YpgNVBS0XHAUBnrI5QyiGZjy5ZQZs6cSExMFFJKYmKimDlzoqoMVFRUaiQmUwRSyr9RUsuWxlSUMnlxppLDGAsXziIrK7NIW1ZWJgsXzjKnGCoqKipVAovtEQgh/IBHgC/LcOxEIUS4ECI8Pj6+wnNfvRp9R+0qKioqdzOW3Cz+FHhdlqFak5RyuZSyg5SyQ61aRgPj7oi6dQPvqF1FRUXlbsaSiqADsEEIcQkYhlIHd4g5Jp4+fS4ODo5F2hwcHJk+fa45pldRUVGpUlgsxYSUMqTgvRBiNUoB7S3mmHvIkJEAzJv3Otevx6DR2PDBB8sN7SoqKio1CVO6j64H9gGNhRBXhBDjhRCThBCTTDXnnTBkyEgCA+sBoNXm0b//YAtLpFJeVFdgFZWKYbIVgZRyxB0c+7Sp5CgJnU7HwYP/ADBmzFTKsFWhUgUpcAUu8AIrcAUG1BWeikoZqbGRxTqdjg4dehAQEMLOnVtp3dpdfZqshqiuwCoqFafGKgKNRsPIkZNISLiuBpZVY1RXYBWVilNjFQGoT5N3A6orsIpKxamxiiAtLZWYGPVpsrozffpc7O1VV2AVlYpQYxVBeHgYII32qU+T1YchQ0bSsWNPw2c/vyDVFVhF5Q6pdqUqKwsrK2usrTXk52uLtKtPk9UPT08vABo0aMrvv5+ysDQqKtWPGrsi6N37Pi5cyGPRonX4+QUhhFCfJqspn34aiqdnLS5ePMPOnVstLY6KSrWjxq4ICnj44RF07dqXtLQbNGjQ1NLiqJST7OxMpJTEx1+ztCgqKtWOGq8IsrOz6dLFD4DISON7BipVnxUrtiKlpFWrDpYWRUWl2lFjTUOLF79DgwYaHnmkEwBCCDIzM29zlkpVZPjwnrzyyhiSkxNxcXGztDgqKtWOGrsiSEpKID8/n/T0NHUlUM05ffoIGRnpJCcnWFoUFZVqSY1dEbz44lts2rSXlSt/trQoKhXkjTcWULduICtXfszq1Z9bWhwVlWpHjVUEnp7etG/flaZNW1laFJUKMnLkJNzcPIiOjmDv3l2WFkdFpdpRYxVBYfr0aUjz5s7s2aPeRKorjz76FP37P8ygQcMtLYqKSrWjxiqCRYveZuDAVixe/A5xcVfJzMwgNvaKpcW6LWru/aIkJSXw/PPDuXYthhUrflJjQFRUykGN3Szet+8Pzp49jqenN3PnLiM7O5Peve+ztFiloubeL86ZM8fYsWMT1tbWzJ79saXFUVGpltRYRTBs2Fi8vX3o02cQjzwyytLilInSsqXWVEXg7OxKo0bNycvLZcWKj/H1DeDBBx+ztFgqKtUKIaVpXCeFEN8ADwJxUsoWRvpHAq8DAkgDnpNSHr3duB06dJDh4eHlkikxEdLSynVqlaBePSuMfV9CCCIianaFtbfemsratUuoWzeAPXvU7LEqdx92duDrW/7zhRCHpJRGIy5NuUewGhhYSn8k0FtK2RJ4D1huQllK5bPP3mXUqHvZtu17S4lQJtTc+yUTEtIIFxc3AgPrW1oUFZVqh8kUgZTybyCplP69Uspk/cf9gL+pZDHG5s3rWLJkLmfOnGDLllD27NnJn3/+nzlFuGOmT5+LRmNTpK2mZ0uNjDzP0aMHefTR0Rw7lsL69X9aWiQVlWpHVfEaGg/sKKlTCDFRCBEuhAiPj4+vlAk/+mgGH388m02bVtG//8O0a9eVzp37VMrYpkSnyze89/DwqhLZUi3lybRlSygDB7ZkyJBOdOxYp8Z7UKmolBeLbxYLIfqiKIIeJR0jpVyO3nTUoUOHStnUqFevMTk52YSENGLkyGcrY0iTUuAxpNPd3AvIzs6yoEQKlvJkKpg3NzcHgNzcHGbOnIiUstps/quoVBVMtlkMIIQIBrYZ2yzW97cCNgP3SynPlWXMmrpZ3KNHMDExUcXa/fyCCAu7ZH6B9FhKrpLmBTWLrMrdiSk3iy22IhBCBAI/AqPLqgRMxeXLkZw4cZg6dfxp27azJUUpkZLqKFu6vrKl5LL0362icjdhsj0CIcR6YB/QWAhxRQgxXggxSQgxSX/IHMAL+EIIcUQIUb7H/HJS2MTywQev8vzzw5gx4xlzinBHVFWPIUvJVfK8ASadV0XlbsSUXkMjpJS+UkobKaW/lPJrKeVXUsqv9P3PSCk9pJRt9C+zVhRp0sSBkBDBrl3bqFs3EFtbO7y8aptThDti+vS52NjYFmmrCh5D06fPxcHBsUibOeQqad5XX/3QpPOqqNyNVBWvIbMjpbIicHJy5s03P+Hs2Wy++67qJp0bMmQkzz77msF9tKrLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {
|
||||
"needs_background": "light"
|
||||
},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"INFO:root:gpu available: True, used: True\n",
|
||||
"INFO:root:VISIBLE GPUS: 0\n"
|
||||
"step 0, {'val_loss': '0.010536636225879192', 'val/loss_mse': '0.0017317170277237892', 'val/loss_p': '0.010536636225879192', 'val/sigma': '1.0592342615127563'}\n",
|
||||
"\r"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"application/vnd.jupyter.widget-view+json": {
|
||||
"model_id": "cf60c316bc5640b6b8f4254625c9322c",
|
||||
"version_major": 2,
|
||||
"version_minor": 0
|
||||
},
|
||||
"text/plain": [
|
||||
"HBox(children=(FloatProgress(value=1.0, bar_style='info', layout=Layout(flex='2'), max=1.0), HTML(value='')), …"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"application/vnd.jupyter.widget-view+json": {
|
||||
"model_id": "",
|
||||
"version_major": 2,
|
||||
"version_minor": 0
|
||||
},
|
||||
"text/plain": [
|
||||
"HBox(children=(FloatProgress(value=0.0, description='Validating', layout=Layout(flex='2'), max=178.0, style=Pr…"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYIAAAEWCAYAAABrDZDcAAAABHNCSVQICAgIfAhkiAAAAAlwSFlzAAALEgAACxIB0t1+/AAAADh0RVh0U29mdHdhcmUAbWF0cGxvdGxpYiB2ZXJzaW9uMy4xLjEsIGh0dHA6Ly9tYXRwbG90bGliLm9yZy8QZhcZAAAgAElEQVR4nOydd3hUxdeA30lPSC+QntCrSJPeFQELoICCSJMiCvqhYqFINaICoj+UqhRNsIIIKogoRUBK6B0RCIQEQnrv8/1xdzcbshtCks0m5L7PM092Z+bOnL3ZvWfKmXOElBIVFRUVleqLhbkFUFFRUVExL6oiUFFRUanmqIpARUVFpZqjKgIVFRWVao6qCFRUVFSqOaoiUFFRUanmqIqgGiOEkEKIeuaWQ0VFxbyoikClUiGEsBVCrBZCJAshbgohXi+m7uNCiL1CiERN3S+EEE4laUsIYSOE+FEIcVWjELvf0XYPIcROIUSSEOJqCeQ2Wl8IUVMI8Y0QIkpTvk8I0e4u7c0TQpwSQuQKIWbfUSaEENOFENc0n+1bIYRzMW25CyF+EkKkCSEihBDP3VH+ihDiiqatcCFE52LaEkKID4UQcZr0oRBC6JW3EEIcEUKka/62qIi2VMqGqghUKhuzgfpAENADeEsI0cdIXRfgPcAXaAz4AQvuoa29wPPATQNtpwGrgTdLKHdx9R2Bw0BrwB1YB/wqhHAspr1LwFvArwbKRgDDgU4on90eWFJMW58D2UAtYBiwTAjRFECjkD4ABqHczy+Bn4QQlkbaGg8MAB4EmgNPAi9q2rIBfgZCATfN5/xZk2/qtlTKgpRSTdU0ARKop3ntAnwF3AYigBmAhaasHrAbSAJige80+QJYDMQAycApoFkZZYoCHtV7Pw/4toTXPg2cute2gEigu5E2HwGu3oP8JaqvuV+tS1AvFJh9R96PwJt67zsCmYCDgetroCiBBnp5XwMfaF4/Cxy6o74EfIzIsx8Yr/d+DHBA8/pR4AYg9MqvAX1M3ZaaypbUGYGKliUoyqAO0A1l1DlaUzYP2I4yMvOnYPT5KNAVaKC59hkgDkAI8Y5mycZgMiSAEMIN8AFO6GWfAJqW8DN0Bc6UU1smQ7PEYYMy6i91M3e8tkWZ/dxJAyBXSnlRL0//PmwFLIUQ7TSzgBeA42hmSUKI54QQJ/WubYrxe9oUOCk1T20NJ7Xl5dmWSvliZW4BVMyP5gEwBGghpUwBUoQQi1CWH74EclCWV3yllJEoSypo8p2ARiijynPaNqWUH6AsOdwL2qWSJL28JE0fd/sMvYCRgHbtvdRtmRLNWv7XwBwpZdLd6hthG8oy1/dAAvC2Jt/BQF1HlNmHPvr3IQXYgPI/FUAi0Ff7AJZSrgfW39HenffUUbO2f2dZob7Ksy2V8kWdEagAeALWKEtCWiJQ1txBWasWwCEhxBkhxAsAUsq/gM9Q1qBjhBAri9u0LAGpmr/6bTijPKyMIoRoj/KAGaQ38i1VWyVBCDFNCJGqScvv4Tp7YAvK8sd8vfwzeu11KUFTq4FvgF0oM6CdmvxIA3VTKXwPoPB9GIMy82uKMkt5HvhFCOFrpO8723MGUjWK4259mbItlTKgKgIVUNb9taN+LYEoa7RIKW9KKcdJKX1RNvOWCo3ZqZTyf1LK1kATlGWIN6HIw7JIMiSElDIBiEbZPNTyIJrlHkMIIVoCm4EXpJR/lqWtkiKlfF9K6ahJE0pyjRDCFtiE8rB+8Y72muq193cJ+s+XUs6SUgZLKf1RPtMNTbqTi4CVEEJ/2Uj/PrQAfpFSXtS0uw3lvnU00v0ZjN/TM0BzfcsflE1gY/e8PNtSKQOqIlBBSpkHfA+ECCGchBBBwOsoG5UIIQYLIfw11RNQNhPzhRAPadaWrVGsZjKBfE2b+g/LIqkYcb4CZggh3IQQjYBxwFpDFYUQzVCWSV6RUm6517aEYl5qp3lrI4Sw0z54hBAWmjJr5a2wK85ipbj6mvvzI5ABjJRS5hfz+bXtWWvas0B5kNtpLXmEYg5aV2N+2QT4GJhrqF0pZRqwEZgrhKghhOgE9EdZngLFmulxIUQdTXu9UBT6aSOifQW8LoTw08wa3qDgnu4C8oBXNfd2kib/rwpoS6UsmHu3Wk3mSxS2GnJDefDfBq4DMymwGvoIZbSZCvyHxtIDeBhlAy8VZVYRBjiWUSZblKWPZOAW8HoxddegKJ5UvXSmpG0BVzX3QD8Fa8q6GyjbVYwsRuujbL5LIP0OWbsU095aA+2N0pQ1AC5o2oso7h5p6rujzEbSUCxvntMrE8BcTX4KcA4Yrlc+7I57KjTfh3hN+ojClj0tgSMoSu8o0NIUbampfJPQ3HAVFRUVlWqKujSkoqKiUs1RFYGKiopKNUdVBCoqKirVHFURqKioqFRzVEWgoqKiUs2pci4mPD09ZXBwsLnFUFFRUalSHDlyJFZK6WWorMopguDgYMLDw80thoqKikqVQggRYaxMXRpSUVFRqeaoikBFRUWlmqMqAhUVFZVqjqoIVFRUVKo5JlMEQgkaHiOEMObFUFvvIaEE6B5kKlkMERYWRnBwMBYWFgQHBxMWFlaR3auoqKhUGkw5I1gLGAs6DugiY32IEgaxwggLC2P8+PFEREQgpSQiIoLx48erykBFRaVaYjJFIKXcg+JatjheQQmTF2MqOQwxffp00tPTC+Wlp6czffr0ihRDRUVFpVJgtj0CIYQf8BSwrAR1xwshwoUQ4bdv3y5z39euXbunfBUVFZX7GXNuFn8CvC1LEK1JSrlSStlGStnGy8vgwbh7IjAw8J7yVVRUVO5nzKkI2gDfCiGuAoNQ4uAOqIiOQ0JCcHBwKJTn4OBASEhIRXSvoqKiUqkwm4sJKWVt7WshxFqUANqbKqLvYcOGAfD2229z48YNrK2tWblypS5fRUVFpTphSvPRb4B/gIZCiEghxBghxAQhxART9XkvDBs2jDp16gCQk5ND//79zSyRSmlRTYFVVMqGyWYEUsqh91B3lKnkMEZ+fj5///03AK+88gr5+XfdqlCphGhNgbVWYFpTYECd4amolJBqe7I4Pz+fzp07U7t2bTZv3oyrq6s6mqyCqKbAKiplp9oqAisrKyZMmMCtW7fUg2VVGNUUWEWl7FRbRQDqaPJ+QDUFVlEpO9VWESQnJ6ujyfsA1RRYRaXsVFtFsHfvXqSUBsvU0WTVYdiwYXTp0kX3PigoSDUFVlG5R6pcqMrywtLSEisrK3Jzcwvlq6PJqoeHhwcAjRs35uzZs2aWRkWl6lFtZwS9e/cmJyeH0NBQgoKCEEKoo8kqSlhYGF5eXpw/f57NmzebWxwVlSpHtZ0RaBk6dCg9evQgKSmJxo0bm1sclVKSnp6OlJKbN2+aWxQVlSpHtVcEmZmZ+Pn5ARjdM1Cp/GzevBkpJW3atDG3KCoqVY5quzQ0Z84crKysaNu2LQBCiCKmpCpVgy5dujBy5Eji4uJwcXExtzgqKlWOajsjiI2NJS8vj5SUFHUmUMU5fvw4qampxMbGmlsUFZUqSbWdEcyaNYv9+/ezZcsWc4uiUkYWLFhAYGAgixYtYsmSJeYWR0WlylFtFYGnpycdOnSgefPm5hZFpYxMmDABNzc3Ll++zJ9//mlucVRUqhzVVhHoU79+fRwdHdWHSBVmxIgR9OvXj8GDB5tbFBWVKke1VQSzZ8+mefPmzJkzh6ioKNLS0oiMjDS3WHdF9b1fmNjYWAYPHsyNGzf4+eef1TMgKiqloNpuFv/111+cOnUKT09PVqxYQXp6Or179za3WMWi+t4vysmTJ/nxxx+xtLRk0aJF5hZHRaVKUm0VwejRo6lVqxaPPfYYzz//vLnFKRHFeUutrorA2dmZpk2bkp2dzaJFiwgICOCZZ54xt1gqKlUKYSrTSSHEauAJIEZK2cxA+TDgbUAAKcBLUsoTd2u3TZs2Mjw8vLzFrRJYWFgYNHUVQlT7CGuvvPIKn332GQEBAar3WBUVAwghjkgpDZ64NOUewVqgTzHlV4BuUsoHgHnAShPKUixz586lV69efP/99+YSoUSovveN06BBA1xcXKhbt665RVFRqXKYTBFIKfcA8cWU75dSJmjeHgD8TSWLIUJDQwkJCeH06dOEhYWxY8cOfv3114oU4Z4JCQnB2tq6UF5195b677//cvjwYYYPH05iYiI7d+40t0gqKlWOymI1NAbYaqxQCDFeCBEuhAi/fft2uXT4zjvvMGPGDNasWUO/fv3o0KED3bt3L5e2TUleXp7utYeHR6XwlmouS6awsDAeeOAB2rZti7e3d7W3oFJRKTVSSpMlIBg4fZc6PYBzgEdJ2mzdurUsD3r27Ck9PDzk8uXLy6U9UxMaGiodHBwkoEsODg4yNDS0WsplrN+vv/7apP2qqFRVgHBp5Llqss1iACFEMPCLNLBZrClvDvwE9JVSXixJm9V1szg4OJiIiIgi+UFBQVy9erXiBdJgLrmM9QuqF1kVFUMUt1lsNvNRIUQgsBEYXlIlYCquXLnC0aNH8ff3p127duYUxSiVNb6yueQy9+dWUbmfMNkegRDiG+AfoKEQIlIIMUYIMUEIMUFTZSbgASwVQhwXQlToMF/f3PLNN99k0KBBjB07tiJFuCcqq8WQueQy1n5AQIBJ+1VRuR8xpdXQUCmlj5TSWkrpL6X8Ukq5XEq5XFM+VkrpJqVsoUkVGlHE3t4eIQS//PILgYGB2NraUrNmzYoU4Z4ICQnBxsamUF5lsBgKCQnBwcGhUF5FyGWs3/nz55u0XxWV+xJjmweVNZXXZrGVlZUE5M6dO8ulvYpgxowZ0traWgIyKCjI7BvFWkJDQ6WdnZ0EZGBgYIXJFRoaKi0sLCQgXVxcKs39UFGpjGCuzWJTUF6bxfn5+aSmpuLg4ICVVbX1tFEu5OfnY2lpCcC3337Ls88+W2F979q1i+joaDp16kTLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {
|
||||
"needs_background": "light"
|
||||
},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"step 1826, {'val_loss': '0.018330058082938194', 'val/loss_mse': '0.001598153030499816', 'val/loss_p': '0.018330058082938194', 'val/sigma': '0.17532871663570404'}\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"application/vnd.jupyter.widget-view+json": {
|
||||
"model_id": "",
|
||||
"version_major": 2,
|
||||
"version_minor": 0
|
||||
},
|
||||
"text/plain": [
|
||||
"HBox(children=(FloatProgress(value=0.0, description='Validating', layout=Layout(flex='2'), max=178.0, style=Pr…"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYIAAAEWCAYAAABrDZDcAAAABHNCSVQICAgIfAhkiAAAAAlwSFlzAAALEgAACxIB0t1+/AAAADh0RVh0U29mdHdhcmUAbWF0cGxvdGxpYiB2ZXJzaW9uMy4xLjEsIGh0dHA6Ly9tYXRwbG90bGliLm9yZy8QZhcZAAAgAElEQVR4nOydd3hURdfAf5OekBBIAUIqSG/SRKSjKCgK8gIqRkpoIogdlfIhoBEVUHkREVApEhReC4KClaKIAqH3TiCNhITUTdvsfH/c3ZBNdjd1s4m5v+e5T7J37s6cvXt3zsycM+cIKSUqKioqKrUXO1sLoKKioqJiW1RFoKKiolLLURWBioqKSi1HVQQqKioqtRxVEaioqKjUclRFoKKiolLLURVBLUYIIYUQzWwth4qKim1RFYFKtUII4SyE+FwIkSaEiBdCvGTh2sFCiL1CiBT9tZ8KITxKU5cQwkkI8bUQ4qpeIfYrUnd/IcQuIUSqEOJqKeQ2e70QooEQ4kshRKy+/C8hxN0l1PemEOKEEEIrhJhXpEwIIWYLIa7pP9tXQoi6FuryEkJ8J4TIFEJECSGeLFI+XQhxRV9XpBCil4W6hBDiXSFEkv54VwghCpV3FEIcEkJo9H87VkVdKhVDVQQq1Y15QHMgGOgPvCqEGGTmWk/gLaAx0BrwBxaVoa69wFNAvIm6M4HPgRmllNvS9e7AQaAL4AWsA34UQrhbqO8i8Crwo4myMcBooCfKZ3cFllmoazmQCzQEQoEVQoi2AHqF9A4wAuV+fgZ8J4SwN1PXZOBR4E6gA/AI8LS+Lifge2ADUF//Ob/Xn7d2XSoVQUqpHrX0ACTQTP+/J7AeSASigDmAnb6sGbAHSAVuApv05wXwAZAApAEngHYVlCkWeKDQ6zeBr0r53v8AJ8paFxAN9DNT5wDgahnkL9X1+vvVpRTXbQDmFTn3NTCj0OseQDbgZuL9dVCUQItC574A3tH//zhwoMj1EvAzI88+YHKh1xOAf/T/PwDEAKJQ+TVgkLXrUo+KHeqMQMXAMhRl0BToizLqDNOXvQn8gjIyC+D26PMBoA/QQv/ex4AkACHE6/olG5OHKQGEEPUBP+BYodPHgLal/Ax9gFOVVJfV0C9xOKGM+stdTZH/nVFmP0VpAWillOcLnSt8H3YA9kKIu/WzgPHAUfSzJCHEk0KI44Xe2xbz97QtcFzqe209xw3llVmXSuXiYGsBVGyPvgN4AugopUwH0oUQS1CWHz4D8lCWVxpLKaNRllTQn/cAWqGMKs8Y6pRSvoOy5FAWDEslqYXOperbKOkz3A+MBQxr7+Wuy5ro1/K/AOZLKVNLut4MP6Esc20GbgGv6c+7mbjWHWX2UZjC9yEd+AblOxVACvCgoQOWUm4ENhapr+g9ddev7RctM2qrMutSqVzUGYEKgA/giLIkZCAKZc0dlLVqARwQQpwSQowHkFLuBD5CWYNOEEKssmS0LAUZ+r+F66iL0lmZRQjRHaWDGVFo5FuuukqDEGKWECJDf3xShve5AttQlj8WFjp/qlB9vUtR1efAl8BulBnQLv35aBPXZmB8D8D4PkxAmfm1RZmlPAX8IIRobKbtovXVBTL0iqOktqxZl0oFUBWBCijr/oZRv4EglDVapJTxUspJUsrGKMa8j4Xe7VRK+V8pZRegDcoyxAwo1lkWO0wJIaW8BcShGA8N3Il+uccUQohOwFZgvJTy94rUVVqklG9LKd31x5TSvEcI4QxsQemsny5SX9tC9f1ZivZ1Uso3pJQhUsoAlM8Uoz+Kch5wEEIUXjYqfB86Aj9IKc/r6/0J5b71MNP8Kczf01NAh8KePyhGYHP3vDLrUqkAqiJQQUqZD2wGwoUQHkKIYOAlFEMlQoiRQogA/eW3UIyJOiHEXfq1ZUcUr5lsQKevs3BnWeywIM56YI4Qor4QohUwCVhr6kIhRDuUZZLpUsptZa1LKO6lLvqXTkIIF0PHI4Sw05c5Ki+FiyWPFUvX6+/P10AWMFZKqbPw+Q31Oerrs0PpyF0MnjxCcQe9Q+9+2QZ4H1hgql4pZSbwLbBACFFHCNETGIqyPAWKN9NgIURTfX33oyj0k2ZEWw+8JITw188aXub2Pd0N5APP6e/ts/rzO6ugLpWKYGtrtXrY7sDYa6g+SsefCFwH5nLba+g9lNFmBnAJvacHcB+KAS8DZVYRAbhXUCZnlKWPNOAG8JKFa9egKJ6MQsep0tYFXNXfg8JHiL6sn4my3RZkMXs9ivFdApoisva2UN9aE/WN05e1AM7p64uydI/013uhzEYyUTxvnixUJoAF+vPpwBlgdKHy0CL3VOifh2T98R7Gnj2dgEMoSu8w0MkadalH5R5Cf8NVVFRUVGop6tKQioqKSi1HVQQqKioqtRxVEaioqKjUclRFoKKiolLLURWBioqKSi2nxoWY8PHxkSEhIbYWQ0VFRaVGcejQoZtSSl9TZTVOEYSEhBAZGWlrMVRUVFRqFEKIKHNl6tKQioqKSi1HVQQqKioqtRxVEaioqKjUclRFoKKiolLLsZoiEErS8AQhhLkohobr7hJKgu4R1pLFFBEREYSEhGBnZ0dISAgRERFV2byKiopKtcGaM4K1gLmk40BBZqx3UdIgVhkRERFMnjyZqKgopJRERUUxefJkVRmoqKjUSqymCKSUf6CElrXEdJQ0eQnWksMUs2fPRqPRGJ3TaDTMnj27KsVQUVFRqRbYzEYghPAHhgErSnHtZCFEpBAiMjExscJtX7t2rUznVVRUVP7N2NJY/CHwmixFtiYp5SopZVcpZVdfX5Mb48pEUFBQmc6rqKio/JuxpSLoCnwlhLgKjEDJg/toVTQcHh6Om5ub0Tk3NzfCw8OronkVFRWVaoXNQkxIKZsY/hdCrEVJoL2lKtoODQ0F4LXXXiMmJgZHR0dWrVpVcF5FRUWlNmFN99Evgb+BlkKIaCHEBCHEFCHEFGu1WRZCQ0Np2rQpAHl5eQwdOtTGEqmUF9UVWEWlYlhtRiClHFWGa8dZSw5z6HQ6/vzzTwCmT5+OTleiqUKlGmJwBTZ4gRlcgQF1hqeiUkpq7c5inU5Hr169aNKkCVu3bqVevXrqaLIGoroCq6hUnFqrCBwcHJgyZQo3btxQN5bVYFRXYBWVilNrFQGoo8l/A6orsIpKxam1iiAtLU0dTf4LUF2BVVQqTq1VBHv37kVKabJMHU3WHEJDQ+ndu3fB6+DgYNUVWEWljNS4VJWVhb29PQ4ODmi1WqPz6miy5uHt7Q1A69atOX36tI2lUVGpedTaGcHAgQPJy8tjw4YNBAcHI4RQR5M1lIiICHx9fTl79ixbt261tTgqKjWOWjsjMDBq1Cj69+9PamoqrVu3trU4KuVEo9EgpSQ+Pt7Woqio1DhqvSLIzs7G398fwKzNQKX6s3XrVqSUdO3a1daiqKjUOGrt0tD8+fNxcHCgW7duAAghirmSqtQMevfuzdixY0lKSsLT09PW4qio1Dhq7Yzg5s2b5Ofnk56ers4EajhHjx4lIyODmzdv2loUFZUaSa2dEbzxxhvs27ePbdu22VoUlQqyaNEigoKCWLJkCcuWLbO1OCoqNY5aqwh8fHy455576NChg61FUakgU6ZMoX79+ly+fJnff//d1uKoqNQ4aq0iKEzz5s1xd3dXO5EazJgxYxgyZAgjR460tSgqKjWOWqsI5s2bR4cOHZg/fz6xsbFkZmYSHR1ta7FKRI29b8zNmzcZOXIkMTExfP/99+oeEBWVclBrjcU7d+7kxIkT+Pj4sHLlSjQaDQMHDrS1WBZRY+8X5/jx43z99dfY29uzZMkSW4ujolIjqbWKICwsjIYNG/LQQw/x1FNP2VqcUmEpWmptVQR169albdu25ObmsmTJEgIDA3nsscdsLZaKSo1CWMt1UgjxOfAwkCClbGeiPBR4DRBAOvCMlPJYSfV27dpVRkZGVra4NQI7OzuTrq5CiFqfYW369Ol89NFHBAYGqtFjVVRMIIQ4JKU0uePSmjaCtcAgC+VXgL5SyvbAm8AqK8pikQULFnD//fezefNmW4lQKtTY++Zp0aIFnp6e3HHHHbYWRUWlxmE1RSCl/ANItlC+T0p5S//yHyDAWrKYYsOGDYSHh3Py5EkiIiL47bff+PHHH6tShDITHh6Oo6Oj0bnaHi31woULHDx4kNGjR5OSksKuXbtsLZKKSo2jungNTQB2mCsUQkwWQkQKISITExMrpcHXX3+dOXPmsGbNGoYMGcI999xDv379KqVua5Kfn1/wv7e3d7WIlmorT6aIiAjat29Pt27daNSoUa33oFJRKTdSSqsdQAhwsoRr+gNnAO/S1NmlSxdZGdx7773S29tbfvLJJ5VSn7XZsGGDdHNzk0DB4ebmJjds2FAr5TLX7hdffGHVdlVUaipApDTTr1rNWAwghAgBfpAmjMX68g7Ad8CDUsrzpamzthqLQ0JCiIqKKnY+ODiYq1evVr1Aemwll7l2QY0iq6JiCkvGYpu5jwohgoBvgdGlVQLW4sqVKxw+fJiAgADuvvtuW4piluqaX9lWctn6c6uo/Juwmo1ACPEl8DfQUggRLYSYIISYIoSYor9kLuANfCyEOCqEqNJhfmF3yxkzZjBixAgmTpxYlSKUierqMWQruczVHxgYaNV2VVT+jVjTa2iUlNJPSukopQyQUn4mpfxESvmJvnyilLK+lLKj/qjSjCKurq4IIfjhhx8ICgrC2dmZBg0aVKUIZSI8PBwnJyejc9XBYyg8PBw3Nzejc1Uhl7l2Fy5caNV2VVT+lZgzHlTXo7KMxQ4ODhKQu3btqpT6qoI5c+ZIR0dHCcjg4GCbG4oNbNiwQbq4uEhABgUFVZlcGzZskHZ2dhKQnp6e1eZ+qKhUR7CVsdgaVJaxWKfTkZGRgZubGw4OtTbSRqWg0+mwt7cH4KuvvuLxxx+vsrZ3795Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {
|
||||
"needs_background": "light"
|
||||
},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"step 3653, {'val_loss': '0.020962275564670563', 'val/loss_mse': '0.0016868686070665717', 'val/loss_p': '0.020962275564670563', 'val/sigma': '0.16759823262691498'}\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"application/vnd.jupyter.widget-view+json": {
|
||||
"model_id": "",
|
||||
"version_major": 2,
|
||||
"version_minor": 0
|
||||
},
|
||||
"text/plain": [
|
||||
"HBox(children=(FloatProgress(value=0.0, description='Validating', layout=Layout(flex='2'), max=178.0, style=Pr…"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYIAAAEWCAYAAABrDZDcAAAABHNCSVQICAgIfAhkiAAAAAlwSFlzAAALEgAACxIB0t1+/AAAADh0RVh0U29mdHdhcmUAbWF0cGxvdGxpYiB2ZXJzaW9uMy4xLjEsIGh0dHA6Ly9tYXRwbG90bGliLm9yZy8QZhcZAAAgAElEQVR4nOydd3hUxdrAf5NOCAmQUFPpHREVQUBQUUAuYAGVGwGpcmliQUVQQc3FAggXAQUvICagXlRABfwUEEUUCL0XgUBoCenJJtkkO98fZ3fZTXY3ZVtizu95zkP2lJk3k+W8M/M2IaVERUVFRaX64uFuAVRUVFRU3IuqCFRUVFSqOaoiUFFRUanmqIpARUVFpZqjKgIVFRWVao6qCFRUVFSqOaoiqMYIIaQQorm75VBRUXEvqiJQqVQIIXyFECuFEJlCiOtCiBds3DtACLFLCJGuv/dTIUStsrQlhPARQqwXQlzUK8Texdq+TwixQwiRIYS4WAa5rd4vhKgvhFgnhLiqv/67EOLuUtp7WwhxVAhRKISYXeyaEELMFEJc0v9uXwghAm20VVcI8a0QIkcIkSCE+Gex61OEEBf0bcULIXrYaEsIId4TQqToj/eEEMLkeichxH4hhEb/bydXtKViH6oiUKlszAZaAJHAfcDLQoh+Vu4NAt4BGgNtgFDgg3K0tQt4Grhuoe0cYCUwvYxy27o/ANgH3AHUBT4DfhBCBNho7xzwMvCDhWsjgOFAd5TfvQaw2EZbSwAt0ACIBpYJIdoB6BXSu8AQlPH8L/CtEMLTSlvjgUeA24COwEDgWX1bPsBGIBaoo/89N+rPO7stFXuQUqpHNT0ACTTX/xwErAGSgQRgFuChv9Yc2AlkADeBL/XnBfAhkARkAkeB9nbKdBV4yOTz28AXZXz2MeBoedsCEoHeVtrsA1wsh/xlul8/XneU4b5YYHaxc+uB6Saf7wHyAH8Lz9dEUQItTc59Dryr//lJYG+x+yXQyIo8u4HxJp/HAH/qf34IuAIIk+uXgH7Obks97DvUFYGKgcUoyqAp0Atl1jlKf+1t4P9QZmZh3Jp9PgTcC7TUP/sEkAIghHhVv2Vj8bAkgBCiDtAIOGxy+jDQroy/w73AcQe15TT0Wxw+KLP+CjdT7GdflNVPcVoChVLKMybnTMdhC+AphLhbvwoYDRxCv0oSQvxTCHHE5Nl2WB/TdsARqX9r6zliuO7ItlQci5e7BVBxP/oXwFNAJyllFpAlhJiPsv3wX6AAZXulsZQyEWVLBf35WkBrlFnlSUObUsp3UbYcyoNhqyTD5FyGvo/SfocHgZGAYe+9wm05E/1e/ufAHCllRmn3W2EryjbXV0Aa8Ir+vL+FewNQVh+mmI5DFvA1yt9UAOlAf8MLWEq5FlhbrL3iYxqg39svfs2sL0e2peJY1BWBCkAI4I2yJWQgAWXPHZS9agHsFUIcF0KMBpBSbgc+QtmDThJCLLdltCwD2fp/TdsIRHlZWUUI0RXlBTPEZOZbobbKghDiNSFEtv74uBzP1QC+Q9n+mGty/rhJez3L0NRKYB3wC8oKaIf+fKKFe7MxHwMwH4cxKCu/diirlKeB74UQja30Xby9QCBbrzhK68uZbanYgaoIVEDZ9zfM+g1EoOzRIqW8LqUcJ6VsjGLMWyr0bqdSyv9IKe8A2qJsQ0yHEi/LEoclIaSUacA1FOOhgdvQb/dYQghxO7AJGC2l3GZPW2VFSvlvKWWA/phQlmeEEL7ABpSX9bPF2mtn0t5vZehfJ6V8U0oZJaUMQ/mdruiP4pwBvIQQpttGpuPQCfheSnlG3+5WlHG7x0r3x7E+pseBjqaePyhGYGtj7si2VOxAVQQqSCmLgK+AGCFELSFEJPACiqESIcRQIUSY/vY0FGOiTghxl35v2RvFayYP0OnbNH1ZljhsiLMGmCWEqCOEaA2MA1ZbulEI0R5lm2SKlPK78rYlFPdSP/1HHyGEn+HFI4Tw0F/zVj4KP1seK7bu14/PeiAXGCml1Nn4/Q3teevb80B5kfsZPHmE4g7aTO9+2RZYALxlqV0pZQ7wDfCWEKKmEKI7MBhlewoUb6YBQoim+vYeRFHox6yItgZ4QQgRql81vMitMf0FKAKm6sd2sv78dhe0pWIP7rZWq4f7Dsy9huqgvPiTgcvAG9zyGnofZbaZDfyF3tMDeADFgJeNsqqIAwLslMkXZesjE7gBvGDj3lUoiifb5Dhe1raAi/oxMD2i9Nd6W7j2iw1ZrN6PYnyXgKaYrD1ttLfaQnvP6K+1BE7r20uwNUb6++uirEZyUDxv/mlyTQBv6c9nASeB4SbXo4uNqdB/H1L1x/uYe/bcDuxHUXoHgNud0ZZ6OPYQ+gFXUVFRUammqFtDKioqKtUcVRGoqKioVHNURaCioqJSzVEVgYqKiko1R1UEKioqKtWcKpdiIiQkREZFRblbDBUVFZUqxf79+29KKetZulblFEFUVBTx8fHuFkNFRUWlSiGESLB2Td0aUlFRUanmqIpARUVFpZqjKgIVFRWVao6qCFRUVFSqOU5TBEIpGp4khLCWxdBw311CKdA9xFmyWCIuLo6oqCg8PDyIiooiLi7Old2rqKioVBqcuSJYDVgrOg4YK2O9h1IG0WXExcUxfvx4EhISkFKSkJDA+PHjVWWgoqJSLXGaIpBS/oqSWtYWU1DK5CU5Sw5LzJw5E41GY3ZOo9Ewc+ZMV4qhoqKiUilwm41ACBEKPAosK8O944UQ8UKI+OTkZLv7vnTpUrnOq6ioqPydcaexeCHwiixDtSYp5XIp5Z1Syjvr1bMYGFcuIiIiynVeRUVF5e+MOxXBncAXQoiLwBCUOriPuKLjmJgY/P39zc75+/sTExPjiu5VVFRUKhVuSzEhpWxi+FkIsRqlgPYGV/QdHR0NwCuvvMKVK1fw9vZm+fLlxvMqKioq1Qlnuo+uA/4AWgkhEoUQY4QQE4QQE5zVZ3mIjo6madOmABQUFDB48GA3S6RSUVRXYBUV+3DaikBKOawc9z7jLDmsodPp+O233wCYMmUKOl2ppgqVSojBFdjgBWZwBQbUFZ6KShmptpHFOp2OHj160KRJEzZt2kTt2rXV2WQVRHUFVlGxn2qrCLy8vJgwYQI3btxQA8uqMKorsIqK/VRbRQDqbPLvgOoKrKJiP9VWEWRmZqqzyb8Bqiuwior9VFtFsGvXLqSUFq+ps8mqQ3R0ND179jR+joyMVF2BVVTKSZUrVekoPD098fLyorCw0Oy8OpusegQHBwPQpk0bTpw44WZpVFSqHtV2RdC3b18KCgqIjY0lMjISIYQ6m6yixMXFUa9ePU6dOsWmTZvcLY6KSpWj2q4IDAwbNoz77ruPjIwM2rRp425xVCqIRqNBSsn169fdLYqKSpWj2iuCvLw8QkNDAazaDFQqP5s2bUJKyZ133uluUVRUqhzVdmtozpw5eHl50aVLFwCEECVcSVWqBj179mTkyJGkpKQQFBTkbnFUVKoc1XZFcPPmTYqKisjKylJXAlWcQ4cOkZ2dzc2bN90tiopKlaTargjefPNNdu/ezXfffeduUVTs5IMPPiAiIoL58+ezePFid4ujolLlqLaKICQkhG7dutGxY0d3i6JiJxMmTKBOnTqcP3+ebdu2uVscFZUqR7VVBKa0aNGCgIAA9SVShRkxYgSDBg1i6NCh7hZFRaXKUW0VwezZs+nYsSNz5szh6tWr5OTkkJiY6G6xSkXNvW/OzZs3GTp0KFeuXGHjxo1qDIiKSgWotsbi7du3c/ToUUJCQvjkk0/QaDT07dvX3WLZRM29X5IjR46wfv16PD09mT9/vrvFUVGpklRbRTBq1CgaNGjAww8/zNNPP+1uccqErWyp1VURBAYG0q5dO7RaLfPnzyc8PJwnnnjC3WKpqFQphLNcJ4UQK4F/AElSyvYWrkcDrwACyAL+JaU8XFq7d955p4yPj3e0uFUCDw8Pi66uQohqX2FtypQpfPTRR4SHh6vZY1VULCCE2C+ltBhx6UwbwWqgn43rF4BeUsoOwNvAcifKYpO33nqLBx98kK+++spdIpQJNfe+dVq2bElQUBDNmjVztygqKlUOpykCKeWvQKqN67ullGn6j38CYc6SxRKxsbHExMRw7Ngx4uLi+Pnnn/nhhx9cKUK5iYmJwdvb2+xcdc+WevbsWfbt28fw4cNJT09nx44d7hZJRaXKUVm8hsYAW6xdFEKMF0LECyHik5OTHdLhq6++yqxZs1i1ahWDBg2iW7du9O7d2yFtO5OioiLjz8HBwZUiW6q7PJni4uLo0KEDXbp0oWHDhtXeg0pFpcJIKZ12AFHAsVLuuQ84CQSXpc077rhDOoL7779fBgcHy48//tgh7Tmb2NhY6e/vLwHj4e/vL2NjY6ulXNb6/fzzz53ar4pKVQWIl1beq04zFgMIIaKA76UFY7H+ekfgW6C/lPJMWdqsrsbiqKgoEhISSpyPjIzk4sWLrhdIj7vkstYvqFlkVVQsYctY7Db3USFEBPANMLysSsBZXLhwgQMHDhAWFsbdd9/tTlGsUlnrK7tLLnf/3ioqfyecZiMQQqwD/gBaCSEShRBjhBAThBAT9Le8AQQDS4UQh4QQLp3mm7pbTp8+nSFDhjB27FhXilAuKqvHkLvkstZ+eHi4U/tVUfk74kyvoWFSykZSSm8pZZiU8r9Syo+llB/rr4+VUtaRUnbSHy6tKFKjRg2EEHz//fdERETg6+tL/fr1XSlCuYiJicHHx8fsXGXwGIqJicHf39/snCvkstbv3LlzndqvisrfEmvGg8p6OMpY7OXlJQG5Y8cOh7TnCmbNmiW9vb0lICMjI91uKDYQGxsr/fz8JCAjIiJcJldsbKz08PCQgAwKCqo046GiUhnBXcZiZ+AoY7FOpyM7Oxt/f3+8vKptpg2HoNPp8PT0BOCLL77gySefdFnfv/zyC9eLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {
|
||||
"needs_background": "light"
|
||||
},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"step 5480, {'val_loss': '0.027184773236513138', 'val/loss_mse': '0.0013622333062812686', 'val/loss_p': '0.027184773236513138', 'val/sigma': '0.13322952389717102'}\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"application/vnd.jupyter.widget-view+json": {
|
||||
"model_id": "",
|
||||
"version_major": 2,
|
||||
"version_minor": 0
|
||||
},
|
||||
"text/plain": [
|
||||
"HBox(children=(FloatProgress(value=0.0, description='Validating', layout=Layout(flex='2'), max=178.0, style=Pr…"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYIAAAEWCAYAAABrDZDcAAAABHNCSVQICAgIfAhkiAAAAAlwSFlzAAALEgAACxIB0t1+/AAAADh0RVh0U29mdHdhcmUAbWF0cGxvdGxpYiB2ZXJzaW9uMy4xLjEsIGh0dHA6Ly9tYXRwbG90bGliLm9yZy8QZhcZAAAgAElEQVR4nOydd3hURdfAf5OekAIpQEilVxEBEemiCIiCCopIExHEAi/yiQooRYwogooIKCDNIFhfioivAtJEpfciLYFAIEBCet/z/XE3S0J2Q0iy2YTc3/PMQ3Zm7szJZXPPnTlnzlEigo6Ojo5OxcXO1gLo6Ojo6NgWXRHo6OjoVHB0RaCjo6NTwdEVgY6Ojk4FR1cEOjo6OhUcXRHo6OjoVHB0RVCBUUqJUqqOreXQ0dGxLboi0ClTKKWclVKLlFIJSqlLSqkxBfTtoZTarpS6buy7UCnlUZixlFJOSqkflFIRRoXY6aaxH1BK/aGUildKRRRCbov9lVJVlVIrlFIXje1/KqXuu8V4U5VSh5RSWUqpyTe1KaXUBKXUOePvtlIp5VnAWN5Kqf8qpZKVUpFKqWdvah+plDprHGu3UqpdAWMppdSHSqlrxvKhUkrlam+mlNqjlEox/tusNMbSKR66ItApa0wG6gIhwAPAG0qpbhb6egHvATWAhkAA8NFtjLUdGABcMjN2MrAIGFtIuQvq7w7sAloA3sBSYJ1Syr2A8U4BbwDrzLQNAgYCbdF+d1dgdgFjzQEygGpAf2CeUqoxgFEhfQD0QbufXwH/VUrZWxhrOPA4cDfQFHgMeNE4lhOwGggHqhh/z9XGemuPpVMcREQvFbQAAtQx/uwFLAOuAJHA24Cdsa0OsAWIB64C3xrrFfAJEAMkAIeAJsWU6SLwcK7PU4GVhbz2SeDQ7Y4FRAGdLIz5EBBxG/IXqr/xfrUoRL9wYPJNdT8AY3N9bgOkAW5mrq+EpgTq5ar7GvjA+HNfYOdN/QXwtyDPDmB4rs9Dgb+NPz8MXABUrvZzQDdrj6WX4hV9RaCTw2w0ZVAL6Ij21jnE2DYV+A3tzSyQG2+fDwMdgHrGa58GrgEopd4ybtmYLeYEUEpVAfyBA7mqDwCNC/k7dACOlNBYVsO4xeGE9tZf5GFu+tkZbfVzM/WALBH5N1dd7vuwHrBXSt1nXAU8D+zHuEpSSj2rlDqY69rGWL6njYGDYnxqGzmY016SY+mULA62FkDH9hgfAM8AzUQkEUhUSs1E2374CshE216pISJRaFsqGOs9gAZob5XHcsYUkQ/Qthxuh5ytkvhcdfHGOW71O3QBBgM5e+9FHsuaGPfyvwamiEj8rfpb4Fe0ba7vgDjgTWO9m5m+7mirj9zkvg+JwI9o/6cKuA50z3kAi8g3wDc3jXfzPXU37u3f3JZnrpIcS6dk0VcEOgC+gCPallAOkWh77qDtVStgp1LqiFLqeQAR2QR8jrYHHaOUml+Q0bIQJBn/zT2GJ9rDyiJKqdZoD5g+ud58izRWYVBKjVdKJRnLF7dxnSuwFm37Y1qu+iO5xmtfiKEWASuAzWgroD+M9VFm+iaR9x5A3vswFG3l1xhtlTIA+FkpVcPC3DeP5wkkGRXHreay5lg6xUBXBDqg7fvnvPXnEIy2R4uIXBKRYSJSA82YN1cZ3U5F5DMRaQE0QtuGGAv5Hpb5ijkhRCQOiEYzHuZwN8btHnMope4B1gDPi8jG4oxVWETkfRFxN5YRhblGKeUMrEJ7WL9403iNc423rRDzG0RkkoiEikgg2u90wVhu5l/AQSmVe9so931oBvwsIv8ax/0V7b61sTD9ESzf0yNA09yeP2hGYEv3vCTH0ikGuiLQQUSyge+AMKWUh1IqBBiDZqhEKfWUUirQ2D0OzZhoUErda9xbdkTzmkkDDMYxcz8s85UCxFkGvK2UqqKUagAMA5aY66iUaoK2TTJSRNbe7lhKcy91MX50Ukq55Dx4lFJ2xjZH7aNyKchjpaD+xvvzA5AKDBYRQwG/f854jsbx7NAe5C45njxKcwetbXS/bAR8DLxrblwRSQZ+At5VSlVSSrUFeqFtT4HmzdRDKVXLOF4XNIV+2IJoy4AxSqkA46rh/7hxTzcD2cAo47191Vi/qRTG0ikOtrZW68V2hbxeQ1XQHvxXgPPARG54DU1He9tMAk5j9PQAHkQz4CWhrSqWA+7FlMkZbesjAbgMjCmg72I0xZOUqxwp7FhAhPEe5C6hxrZOZto2FyCLxf5oxncBUm6StX0B4y0xM95zxrZ6wAnjeJEF3SNjf2+01UgymufNs7naFPCusT4ROAYMzNXe/6Z7qozfh1hjmU5ez557gD1oSm8vcI81xtJLyRZlvOE6Ojo6OhUUfWtIR0dHp4KjKwIdHR2dCo6uCHR0dHQqOLoi0NHR0ang6IpAR0dHp4JT7kJM+Pr6SmhoqK3F0NHR0SlX7Nmz56qI+JlrK3eKIDQ0lN27d9taDB0dHZ1yhVIq0lKbvjWko6OjU8HRFYGOjo5OBUdXBDo6OjoVHF0R6Ojo6FRwrKYIlJY0PEYpZSmKYU6/e5WWoLuPtWQxx/LlywkNDcXOzo7Q0FCWL19emtPr6OjolBmsuSJYAlhKOg6YMmN9iJYGsdRYvnw5w4cPJzIyEhEhMjKS4cOH68pAR0enQmI1RSAiW9FCyxbESLQ0eTHWksMcEyZMICUlJU9dSkoKEyZMKE0xdHR0dMoENrMRKKUCgCeAeYXoO1wptVsptfvKlSvFnvvcuXO3Va+jo6NzJ2NLY/GnwJtSiGxNIjJfRFqKSEs/P7MH426L4ODg26rX0dHRuZOxpSJoCaxUSkUAfdDy4D5eGhOHhYXh5uaWp87NzY2wsLDSmF5HR0enTGGzEBMiUjPnZ6XUErQE2qtKY+7+/fsD8Oabb3LhwgUcHR2ZP3++qV5HR0enImFN99EVwF9AfaVUlFJqqFJqhFJqhLXmvB369+9PrVq1AMjMzKRXr142lkinqOiuwDo6xcNqKwIR6XcbfZ+zlhyWMBgMbNu2DYCRI0diMNzSVKFTBslxBc7xAstxBQb0FZ6OTiGpsCeLDQYD7dq1o2bNmqxZs4bKlSvrb5PlEN0VWEen+FRYReDg4MCIESO4fPmyfrCsHKO7AuvoFJ8KqwhAf5u8E9BdgXV0ik+FVQQJCQn62+QdgO4KrKNTfCqsIti+fTsiYrZNf5ssP/Tv35/27dubPoeEhOiuwDo6t0m5S1VZUtjb2+Pg4EBWVlaeev1tsvzh4+MDQMOGDTl69KiNpdHRKX9U2BVB165dyczMJDw8nJCQEJRS+ttkOWX58uX4+flx/Phx1qxZY2txdHTKHRV2RZBDv379eOCBB4iPj6dhw4a2FkeniKSkpCAiXLp0ydai6OiUOyq8IkhLSyMgIADAos1Ap+yzZs0aRISWLVvaWhQdnXJHhd0amjJlCg4ODrRq1QoApVQ+V1Kd8kH79u0ZPHgw165dw8vLy9bi6OiUOyrsiuDq1atkZ2eTmJiorwTKOfv37ycpKYmrV6/aWhQdnXJJhV0RTJo0iR07drB27Vpbi6JTTD766COCg4OZOXMms2fPtrU4OjrljgqrCHx9fbn//vtp2rSprUXRKSYjRoygSpUqnDlzho0bN9paHB2dckeFVQS5qVu3Lu7u7vpDpBwzaNAgevbsyVNPPWVrUXR0yh0VVhFMnjyZpk2bMmXKFC5evEhycjJRUVG2FuuW6LH383L16lWeeuopLly4wOrVq/UzIDo6RaDCGos3bdrEoUOH8PX15csvvyQlJYWuXbvaWqwC0WPv5+fgwYP88MMP2NvbM3PmTFuLo6NTLqmwimDIkCFUq1aNRx55hAEDBthanEJRULTUiqoIPD09ady4MRkZGcycOZOgoCCefvppW4ulo1OuUNZynVRKLQIeBWJEpImZ9v7Am4ACEoGXROTArcZt2bKl7N69u6TFLRfY2dmZdXVVSlX4DGsjR47k888/JygoSI8eq6NjBqXUHhExe+LSmjaCJUC3AtrPAh1F5C5gKjDfirIUyLvvvkuXLl347rvvbCVCodBj71umXr16eHl5Ubt2bVuLoqNT7rCaIhCRrUBsAe07RCTO+PFvINBaspgjPDycsLAwDh8+zPLly9mwYQPr1q0rTRFum7CwMBwdHfPUVfRoqSdPnmTXrl0MHDiQ69ev88cff9haJB2dckdZ8RoaCqy31KiUGq6U2q2U2n3lypUSmfCtt97i7bffZvHixfTs2ZP777+fTp06lcjY1iQ7O9v0s4+PT5mIlmorT6bly5dz11130apVK6pXr17hPah0dIqMiFitAKHA4Vv0eQA4BvgUZswWLVpISdC5c2fx8fGRL774okTGszbh4eHi5uYmgKm4ublJeHh4hZTL0rxff/21VefV0SmvALvFwnPVasZiAKVUKPCzmDEWG9ubAv8FuovIv4UZs6Iai0NDQ4mMjMxXHxISQkREROkLZMRWclmaF/Qosjo65ijIWGwz91GlVDDwEzCwsErAWpw9e5a9e/cSGBjIfffdZ0tRLFJW8yvbSi5b/946OncSVrMRKKVWAH8B9ZVSUUqpoUqpEUqpEcYuEwEfYK5Sar9SqlRf83O7W44dO5Y+ffrwwgsvlKYIt0VZ9RiylVyWxg8KCrLqvDo6dyLW9BrqJyL+IuIoIoEi8pWIfCEiXxjbXxCRKiLSzFhKNaOIq6srSil+/vlngoODcXZ2pmrVqqUpwm0RFhaGk5NTnrqy4DEUFhaGm5tbnrrSkMvSvNOmTbPqvDo6dySWjAdltZSUsdjBwUEA+eOPP0pkvNLg7bffFkdHRwEkJCTE5obiHMLDw8XFxUUACQ4OLjW5wsPDxc7OTgDx8vIqM/dDR6csgq2MxdagpIzFBoOBpKQk3NzccHCosJE2SgSDwYC9vT0AK1eupG/fvqU29+bNm4mOjqZt27Z0796Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {
|
||||
"needs_background": "light"
|
||||
},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"step 7307, {'val_loss': '0.018462706357240677', 'val/loss_mse': '0.001305599114857614', 'val/loss_p': '0.018462706357240677', 'val/sigma': '0.15088564157485962'}\n",
|
||||
"Epoch 3: reducing learning rate of group 0 to 2.0000e-05.\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"application/vnd.jupyter.widget-view+json": {
|
||||
"model_id": "",
|
||||
"version_major": 2,
|
||||
"version_minor": 0
|
||||
},
|
||||
"text/plain": [
|
||||
"HBox(children=(FloatProgress(value=0.0, description='Validating', layout=Layout(flex='2'), max=178.0, style=Pr…"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYIAAAEWCAYAAABrDZDcAAAABHNCSVQICAgIfAhkiAAAAAlwSFlzAAALEgAACxIB0t1+/AAAADh0RVh0U29mdHdhcmUAbWF0cGxvdGxpYiB2ZXJzaW9uMy4xLjEsIGh0dHA6Ly9tYXRwbG90bGliLm9yZy8QZhcZAAAgAElEQVR4nOydd3hURdfAf5NOSIEUIKQCoRdpIggCojRRLICItBcpoiiv8oGvgAqKEQuoiBTBFxATQEWkir4KgiBK71WEBAKBACmQXvZ8f9zNkpDdEJJsNiH39zzzkJ2ZO3Nys9xzZ86Zc5SIoKOjo6NTcbGztQA6Ojo6OrZFVwQ6Ojo6FRxdEejo6OhUcHRFoKOjo1PB0RWBjo6OTgVHVwQ6Ojo6FRxdEVRglFKilAq1tRw6Ojq2RVcEOmUKpZSzUmqRUuq6UuqSUmpcAX17KaW2K6USjH2/VEq5F2YspZSTUmqlUirSqBA73zL2g0qp35RSiUqpyELIbbG/UqqaUmq5Uuqisf0PpdR9txlvmlLqsFIqSyk19ZY2pZSarJQ6Z/zdViilPAoYy0sp9YNSKlkpFaWUevaW9peVUmeNY+1RSnUoYCyllPpAKXXNWD5QSqlc7c2VUnuVUinGf5uXxlg6xUNXBDpljalAXSAYeBB4TSnVw0JfT+BdoCbQEPAHPrqDsbYDg4BLZsZOBhYBEwopd0H93YDdQCvAC/gK2KCUcitgvNPAa8AGM21DgMFAe7TfvRIwu4Cx5gAZQHVgIDBPKdUYwKiQ3gf6ot3P/wI/KKXsLYw1CngCuAdoBjwGPG8cywlYA4QDVY2/5xpjvbXH0ikOIqKXCloAAUKNP3sCS4ErQBTwBmBnbAsFtgKJwFXgG2O9Aj4BYoHrwGGgSTFlugh0y/V5GrCikNc+BRy+07GAaKCzhTEfBiLvQP5C9Tfer1aF6BcOTL2lbiUwIdfn+4E0wNXM9ZXRlEC9XHVfA+8bf+4P7LqlvwB+FuTZAYzK9Xk48Jfx527ABUDlaj8H9LD2WHopXtFXBDo5zEZTBrWBTmhvncOMbdOA/6G9mQVw8+2zG9ARqGe89mngGoBS6nXjlo3ZYk4ApVRVwA84mKv6INC4kL9DR+BoCY1lNYxbHE5ob/1FHuaWn53RVj+3Ug/IEpFTuepy34eNgL1S6j7jKuA54ADGVZJS6lml1KFc1zbG8j1tDBwS41PbyKGc9pIcS6dkcbC1ADq2x/gAeAZoLiI3gBtKqZlo2w//BTLRtldqikg02pYKxnp3oAHaW+XxnDFF5H20LYc7IWerJDFXXaJxjtv9Dl2BoUDO3nuRx7Imxr38r4G3RSTxdv0t8BPaNte3QDzwH2O9q5m+bmirj9zkvg83gO/R/qYKSAB65jyARWQZsOyW8W69p27Gvf1b2/LMVZJj6ZQs+opAB8AHcETbEsohCm3PHbS9agXsUkodVUo9ByAim4HP0fagY5VSCwoyWhaCJOO/ucfwQHtYWUQp1RbtAdM315tvkcYqDEqpSUqpJGOZfwfXVQLWoW1/TM9VfzTXeA8UYqhFwHJgC9oK6DdjfbSZvknkvQeQ9z4MR1v5NUZbpQwC1iulalqY+9bxPIAko+K43VzWHEunGOiKQAe0ff+ct/4cgtD2aBGRSyIyUkRqohnz5iqj26mIfCYirYBGaNsQEyDfwzJfMSeEiMQDMWjGwxzuwbjdYw6lVAtgLfCciGwqzliFRUTeExE3YxldmGuUUs7AarSH9fO3jNc413jbCjG/QUSmiEiIiASg/U4XjOVWTgEOSqnc20a570NzYL2InDKO+xPafbvfwvRHsXxPjwLNcnv+oBmBLd3zkhxLpxjoikAHEckGvgXClFLuSqlgYByaoRKlVD+lVICxezyaMdGglLrXuLfsiOY1kwYYjGPmfljmKwWIsxR4QylVVSnVABgJLDHXUSnVBG2b5GURWXenYynNvdTF+NFJKeWS8+BRStkZ2xy1j8qlII+Vgvob789KIBUYKiKGAn7/nPEcjePZoT3IXXI8eZTmDlrH6H7ZCPgYeMfcuCKSDKwC3lFKVVZKtQceR9ueAs2bqZdSqrZxvK5oCv2IBdGWAuOUUv7GVcP/cfOebgGygbHGe/uSsX5zKYylUxxsba3Wi+0Keb2GqqI9+K8A54G3uOk19CHa22YS8A9GTw/gITQDXhLaqiICcCumTM5oWx/XgcvAuAL6LkZTPEm5ytHCjgVEGu9B7hJibOtspm1LAbJY7I9mfBcg5RZZHyhgvCVmxvuXsa0ecNI4XlRB98jY3wttNZKM5nnzbK42BbxjrL8BHAcG52ofeMs9VcbvQ5yxfEhez54WwF40pbcPaGGNsfRSskUZb7iOjo6OTgVF3xrS0dHRqeDoikBHR0engqMrAh0dHZ0Kjq4IdHR0dCo4uiLQ0dHRqeCUuxATPj4+EhISYmsxdHR0dMoVe/fuvSoivubayp0iCAkJYc+ePbYWQ0dHR6dcoZSKstSmbw3p6OjoVHB0RaCjo6NTwdEVgY6Ojk4FR1cEOjo6OhUcqykCpSUNj1VKWYpimNPvXqUl6O5rLVnMERERQUhICHZ2doSEhBAREVGa0+vo6OiUGay5IlgCWEo6DpgyY32Algax1IiIiGDUqFFERUUhIkRFRTFq1ChdGejo6FRIrKYIROR3tNCyBfEyWpq8WGvJYY7JkyeTkpKSpy4lJYXJkyeXphg6Ojo6ZQKb2QiUUv7Ak8C8QvQdpZTao5Tac+XKlWLPfe7cuTuq19HR0bmbsaWx+FPgP1KIbE0iskBEWotIa19fswfj7oigoKA7qtfR0dG5m7GlImgNrFBKRQJ90fLgPlEaE4eFheHq6pqnztXVlbCwsNKYXkdHR6dMYbMQEyJSK+dnpdQStATaq0tj7oEDBwLwn//8hwsXLuDo6MiCBQtM9To6OjoVCWu6jy4H/gTqK6WilVLDlVKjlVKjrTXnnTBw4EBq164NQGZmJo8//riNJdIpKrorsI5O8bDaikBEBtxB339ZSw5LGAwGtm3bBsDLL7+MwXBbU4VOGSTHFTjHCyzHFRjQV3g6OoWkwp4sNhgMdOjQgVq1arF27VqqVKmiv02WQ3RXYB2d4lNhFYGDgwOjR4/m8uXL+sGycozuCqyjU3wqrCIA/W3ybkB3BdbRKT4VVhFcv35df5u8C9BdgXV0ik+FVQTbt29HRMy26W+T5YeBAwfywAMPmD4HBwfrrsA6OndIuUtVWVLY29vj4OBAVlZWnnr9bbL84e3tDUDDhg05duyYjaXR0Sl/VNgVQffu3cnMzCQ8PJzg4GCUUvrbZDklIiICX19fTpw4wdq1a20tjo5OuaPCrghyGDBgAA8++CCJiYk0bNjQ1uLoFJGUlBREhEuXLtlaFB2dckeFVwRpaWn4+/sDWLQZ6JR91q5di4jQunVrW4uio1PuqLBbQ2+//TYODg60adMGAKVUPldSnfLBAw88wNChQ7l27Rqenp62FkdHp9xRYVcEV69eJTs7mxs3bugrgXLOgQMHSEpK4urVq7YWRUenXFJhVwRTpkxhx44drFu3ztai6BSTjz76iKCgIGbOnMns2bNtLY6OTrmjwioCHx8f2rVrR7NmzWwtik4xGT16NFWrVuXMmTNs2rTJ1uLo6JQ7KqwiyE3dunVxc3PTHyLlmCFDhtC7d2/69etna1F0dModFVYRTJ06lWbNmvH2229z8eJFkpOTiY6OtrVYt0WPvZ+Xq1ev0q9fPy5cuMCaNWv0MyA6OkWgwhqLN2/ezOHDh/Hx8eGLL74gJSWF7t2721qsAtFj7+fn0KFDrFy5Ent7e2bOnGlrcXR0yiUVVhEMGzaM6tWr88gjjzBo0CBbi1MoCoqWWlEVgYeHB40bNyYjI4OZM2cSGBjI008/bWuxdHTKFcparpNKqUXAo0CsiDQx0z4Q+A+ggBvACyJy8Hbjtm7dWvbs2VPS4pYL7OzszLq6KqUqfIa1l19+mc8//5zAwEA9eqyOjhmUUntFxOyJS2vaCJYAPQpoPwt0EpGmwDRggRVlKZB33nmHrl278u2339pKhEKhx963TL169fD09KROnTq2FkVHp9xhNUUgIr8DcQW07xCReOPHv4AAa8lijvDwcMLCwjhy5AgRERH8+uuvbNiwoTRFuGPCwsJwdHTMU1fRo6X+/fff7N69m8GDB5OQkMBvv/1ma5F0dModZcVraDiw0VKjUmqUUmqPUmrPlStXSmTC119/nTfeeIPFixfTu3dv2rVrR+fOnUtkbGuSnZ1t+tnb27tMREu1lSdTREQETZs2pU2bNtSoUaPCe1Dp6BQZEbFaAUKAI7fp8yBwHPAuzJitWrWSkqBLly7i7e0t8+fPL5HxrE14eLi4uroKYCqurq4SHh5eIeWyNO/XX39t1Xl1dMorwB6x8Fy1mrEYQCkVAqwXM8ZiY3sz4Aegp4icKsyYFdVYHBISQlRUVL764OBgIiMjS18gI7aSy9K8oEeR1dExR0HGYpu5jyqlgoBVwODCKgFrcfbsWfbt20dAQAD33XefLUWxSFnNr2wruWz9e+vo3E1YzUaglFoO/AnUV0pFK6WGK6VGK6VGG7u8BXgDc5VSB5RSpfqan9vdcsKECfTt25cRI0aUpgh3RFn1GLKVXJbGDwwMtOq8Ojp3I9b0GhogIn4i4igiASLyXxGZLyLzje0jRKSqiDQ3llLNKFKpUiWUUqxfv56goCCcnZ2pVq1aaYpwR4SFheHk5JSnrix4DIWFheHq6pqnrjTksjTv9OnTrTqvjs5diSXjQVktJWUsdnBwEEB+++23EhmvNHjjjTfE0dFRAAkODra5oTiH8PBwcXFxEUCCgoJKTa7w8HCxs7MTQDw9PcvM/dDRKYtgK2OxNSgpY7HBYCApKQlXV1ccHCpspI0SwWAwYG9vD8CKFSvo379/qc29ZcsWYmJiaN++PT1Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {
|
||||
"needs_background": "light"
|
||||
},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"step 9134, {'val_loss': '0.024827884510159492', 'val/loss_mse': '0.0014371582074090838', 'val/loss_p': '0.024827884510159492', 'val/sigma': '0.14072538912296295'}\n",
|
||||
"INFO:root:Epoch 00005: early stopping\n",
|
||||
"\n",
|
||||
"Loading checkpoint lightning_logs/seq2seq/version_-8/_ckpt_epoch_0.ckpt\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"application/vnd.jupyter.widget-view+json": {
|
||||
"model_id": "0e4740ece90b4034b09cc5e414b3fec0",
|
||||
"version_major": 2,
|
||||
"version_minor": 0
|
||||
},
|
||||
"text/plain": [
|
||||
"HBox(children=(FloatProgress(value=0.0, description='Testing', layout=Layout(flex='2'), max=234.0, style=Progr…"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYIAAAEWCAYAAABrDZDcAAAABHNCSVQICAgIfAhkiAAAAAlwSFlzAAALEgAACxIB0t1+/AAAADh0RVh0U29mdHdhcmUAbWF0cGxvdGxpYiB2ZXJzaW9uMy4xLjEsIGh0dHA6Ly9tYXRwbG90bGliLm9yZy8QZhcZAAAgAElEQVR4nOydd3hUxdeA30lPSC+QntCrSJPeFQELoICCSJMiCvqhYqFINaICoj+UqhRNsIIIKogoRUBK6B0RCIQEQnrv8/1xdzcbshtCks0m5L7PM092Z+bOnL3ZvWfKmXOElBIVFRUVleqLhbkFUFFRUVExL6oiUFFRUanmqIpARUVFpZqjKgIVFRWVao6qCFRUVFSqOaoiUFFRUanmqIqgGiOEkEKIeuaWQ0VFxbyoikClUiGEsBVCrBZCJAshbgohXi+m7uNCiL1CiERN3S+EEE4laUsIYSOE+FEIcVWjELvf0XYPIcROIUSSEOJqCeQ2Wl8IUVMI8Y0QIkpTvk8I0e4u7c0TQpwSQuQKIWbfUSaEENOFENc0n+1bIYRzMW25CyF+EkKkCSEihBDP3VH+ihDiiqatcCFE52LaEkKID4UQcZr0oRBC6JW3EEIcEUKka/62qIi2VMqGqghUKhuzgfpAENADeEsI0cdIXRfgPcAXaAz4AQvuoa29wPPATQNtpwGrgTdLKHdx9R2Bw0BrwB1YB/wqhHAspr1LwFvArwbKRgDDgU4on90eWFJMW58D2UAtYBiwTAjRFECjkD4ABqHczy+Bn4QQlkbaGg8MAB4EmgNPAi9q2rIBfgZCATfN5/xZk2/qtlTKgpRSTdU0ARKop3ntAnwF3AYigBmAhaasHrAbSAJige80+QJYDMQAycApoFkZZYoCHtV7Pw/4toTXPg2cute2gEigu5E2HwGu3oP8JaqvuV+tS1AvFJh9R96PwJt67zsCmYCDgetroCiBBnp5XwMfaF4/Cxy6o74EfIzIsx8Yr/d+DHBA8/pR4AYg9MqvAX1M3ZaaypbUGYGKliUoyqAO0A1l1DlaUzYP2I4yMvOnYPT5KNAVaKC59hkgDkAI8Y5mycZgMiSAEMIN8AFO6GWfAJqW8DN0Bc6UU1smQ7PEYYMy6i91M3e8tkWZ/dxJAyBXSnlRL0//PmwFLIUQ7TSzgBeA42hmSUKI54QQJ/WubYrxe9oUOCk1T20NJ7Xl5dmWSvliZW4BVMyP5gEwBGghpUwBUoQQi1CWH74EclCWV3yllJEoSypo8p2ARiijynPaNqWUH6AsOdwL2qWSJL28JE0fd/sMvYCRgHbtvdRtmRLNWv7XwBwpZdLd6hthG8oy1/dAAvC2Jt/BQF1HlNmHPvr3IQXYgPI/FUAi0Ff7AJZSrgfW39HenffUUbO2f2dZob7Ksy2V8kWdEagAeALWKEtCWiJQ1txBWasWwCEhxBkhxAsAUsq/gM9Q1qBjhBAri9u0LAGpmr/6bTijPKyMIoRoj/KAGaQ38i1VWyVBCDFNCJGqScvv4Tp7YAvK8sd8vfwzeu11KUFTq4FvgF0oM6CdmvxIA3VTKXwPoPB9GIMy82uKMkt5HvhFCOFrpO8723MGUjWK4259mbItlTKgKgIVUNb9taN+LYEoa7RIKW9KKcdJKX1RNvOWCo3ZqZTyf1LK1kATlGWIN6HIw7JIMiSElDIBiEbZPNTyIJrlHkMIIVoCm4EXpJR/lqWtkiKlfF9K6ahJE0pyjRDCFtiE8rB+8Y72muq193cJ+s+XUs6SUgZLKf1RPtMNTbqTi4CVEEJ/2Uj/PrQAfpFSXtS0uw3lvnU00v0ZjN/TM0BzfcsflE1gY/e8PNtSKQOqIlBBSpkHfA+ECCGchBBBwOsoG5UIIQYLIfw11RNQNhPzhRAPadaWrVGsZjKBfE2b+g/LIqkYcb4CZggh3IQQjYBxwFpDFYUQzVCWSV6RUm6517aEYl5qp3lrI4Sw0z54hBAWmjJr5a2wK85ipbj6mvvzI5ABjJRS5hfz+bXtWWvas0B5kNtpLXmEYg5aV2N+2QT4GJhrqF0pZRqwEZgrhKghhOgE9EdZngLFmulxIUQdTXu9UBT6aSOifQW8LoTw08wa3qDgnu4C8oBXNfd2kib/rwpoS6UsmHu3Wk3mSxS2GnJDefDfBq4DMymwGvoIZbSZCvyHxtIDeBhlAy8VZVYRBjiWUSZblKWPZOAW8HoxddegKJ5UvXSmpG0BVzX3QD8Fa8q6GyjbVYwsRuujbL5LIP0OWbsU095aA+2N0pQ1AC5o2oso7h5p6rujzEbSUCxvntMrE8BcTX4KcA4Yrlc+7I57KjTfh3hN+ojClj0tgSMoSu8o0NIUbampfJPQ3HAVFRUVlWqKujSkoqKiUs1RFYGKiopKNUdVBCoqKirVHFURqKioqFRzVEWgoqKiUs2pci4mPD09ZXBwsLnFUFFRUalSHDlyJFZK6WWorMopguDgYMLDw80thoqKikqVQggRYaxMXRpSUVFRqeaoikBFRUWlmqMqAhUVFZVqjqoIVFRUVKo5JlMEQgkaHiOEMObFUFvvIaEE6B5kKlkMERYWRnBwMBYWFgQHBxMWFlaR3auoqKhUGkw5I1gLGAs6DugiY32IEgaxwggLC2P8+PFEREQgpSQiIoLx48erykBFRaVaYjJFIKXcg+JatjheQQmTF2MqOQwxffp00tPTC+Wlp6czffr0ihRDRUVFpVJgtj0CIYQf8BSwrAR1xwshwoUQ4bdv3y5z39euXbunfBUVFZX7GXNuFn8CvC1LEK1JSrlSStlGStnGy8vgwbh7IjAw8J7yVVRUVO5nzKkI2gDfCiGuAoNQ4uAOqIiOQ0JCcHBwKJTn4OBASEhIRXSvoqKiUqkwm4sJKWVt7WshxFqUANqbKqLvYcOGAfD2229z48YNrK2tWblypS5fRUVFpTphSvPRb4B/gIZCiEghxBghxAQhxART9XkvDBs2jDp16gCQk5ND//79zSyRSmlRTYFVVMqGyWYEUsqh91B3lKnkMEZ+fj5///03AK+88gr5+XfdqlCphGhNgbVWYFpTYECd4amolJBqe7I4Pz+fzp07U7t2bTZv3oyrq6s6mqyCqKbAKiplp9oqAisrKyZMmMCtW7fUg2VVGNUUWEWl7FRbRQDqaPJ+QDUFVlEpO9VWESQnJ6ujyfsA1RRYRaXsVFtFsHfvXqSUBsvU0WTVYdiwYXTp0kX3PigoSDUFVlG5R6pcqMrywtLSEisrK3Jzcwvlq6PJqoeHhwcAjRs35uzZs2aWRkWl6lFtZwS9e/cmJyeH0NBQgoKCEEKoo8kqSlhYGF5eXpw/f57NmzebWxwVlSpHtZ0RaBk6dCg9evQgKSmJxo0bm1sclVKSnp6OlJKbN2+aWxQVlSpHtVcEmZmZ+Pn5ARjdM1Cp/GzevBkpJW3atDG3KCoqVY5quzQ0Z84crKysaNu2LQBCiCKmpCpVgy5dujBy5Eji4uJwcXExtzgqKlWOajsjiI2NJS8vj5SUFHUmUMU5fvw4qampxMbGmlsUFZUqSbWdEcyaNYv9+/ezZcsWc4uiUkYWLFhAYGAgixYtYsmSJeYWR0WlylFtFYGnpycdOnSgefPm5hZFpYxMmDABNzc3Ll++zJ9//mlucVRUqhzVVhHoU79+fRwdHdWHSBVmxIgR9OvXj8GDB5tbFBWVKke1VQSzZ8+mefPmzJkzh6ioKNLS0oiMjDS3WHdF9b1fmNjYWAYPHsyNGzf4+eef1TMgKiqloNpuFv/111+cOnUKT09PVqxYQXp6Or179za3WMWi+t4vysmTJ/nxxx+xtLRk0aJF5hZHRaVKUm0VwejRo6lVqxaPPfYYzz//vLnFKRHFeUutrorA2dmZpk2bkp2dzaJFiwgICOCZZ54xt1gqKlUKYSrTSSHEauAJIEZK2cxA+TDgbUAAKcBLUsoTd2u3TZs2Mjw8vLzFrRJYWFgYNHUVQlT7CGuvvPIKn332GQEBAar3WBUVAwghjkgpDZ64NOUewVqgTzHlV4BuUsoHgHnAShPKUixz586lV69efP/99+YSoUSovveN06BBA1xcXKhbt665RVFRqXKYTBFIKfcA8cWU75dSJmjeHgD8TSWLIUJDQwkJCeH06dOEhYWxY8cOfv3114oU4Z4JCQnB2tq6UF5195b677//cvjwYYYPH05iYiI7d+40t0gqKlWOymI1NAbYaqxQCDFeCBEuhAi/fft2uXT4zjvvMGPGDNasWUO/fv3o0KED3bt3L5e2TUleXp7utYeHR6XwlmouS6awsDAeeOAB2rZti7e3d7W3oFJRKTVSSpMlIBg4fZc6PYBzgEdJ2mzdurUsD3r27Ck9PDzk8uXLy6U9UxMaGiodHBwkoEsODg4yNDS0WsplrN+vv/7apP2qqFRVgHBp5Llqss1iACFEMPCLNLBZrClvDvwE9JVSXixJm9V1szg4OJiIiIgi+UFBQVy9erXiBdJgLrmM9QuqF1kVFUMUt1lsNvNRIUQgsBEYXlIlYCquXLnC0aNH8ff3p127duYUxSiVNb6yueQy9+dWUbmfMNkegRDiG+AfoKEQIlIIMUYIMUEIMUFTZSbgASwVQhwXQlToMF/f3PLNN99k0KBBjB07tiJFuCcqq8WQueQy1n5AQIBJ+1VRuR8xpdXQUCmlj5TSWkrpL6X8Ukq5XEq5XFM+VkrpJqVsoUkVGlHE3t4eIQS//PILgYGB2NraUrNmzYoU4Z4ICQnBxsamUF5lsBgKCQnBwcGhUF5FyGWs3/nz55u0XxWV+xJjmweVNZXXZrGVlZUE5M6dO8ulvYpgxowZ0traWgIyKCjI7BvFWkJDQ6WdnZ0EZGBgYIXJFRoaKi0sLCQgXVxcKs39UFGpjGCuzWJTUF6bxfn5+aSmpuLg4ICVVbX1tFEu5OfnY2lpCcC3337Ls88+W2F979q1i+joaDp16kTLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {
|
||||
"needs_background": "light"
|
||||
},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"step 9135, {'val_loss': '-0.003565264632925391', 'val/loss_mse': '0.0002942961873486638', 'val/loss_p': '-0.003565264632925391', 'val/sigma': '0.17361457645893097'}\n",
|
||||
"----------------------------------------------------------------------------------------------------\n",
|
||||
"TEST RESULTS\n",
|
||||
"{}\n",
|
||||
"----------------------------------------------------------------------------------------------------\n",
|
||||
"\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"ename": "NameError",
|
||||
"evalue": "name 'plot_from_loader' is not defined",
|
||||
"output_type": "error",
|
||||
"traceback": [
|
||||
"\u001b[0;31m---------------------------------------------------------------------------\u001b[0m",
|
||||
"\u001b[0;31mNameError\u001b[0m Traceback (most recent call last)",
|
||||
"\u001b[0;32m<ipython-input-9-aa3a84af01a3>\u001b[0m in \u001b[0;36m<module>\u001b[0;34m\u001b[0m\n\u001b[1;32m 22\u001b[0m \u001b[0;34m'patience'\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;36m2\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 23\u001b[0m },\n\u001b[0;32m---> 24\u001b[0;31m \u001b[0mPL_MODEL_CLS\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mLSTMSeq2Seq_PL\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m 25\u001b[0m )\n",
|
||||
"\u001b[0;32m/media/wassname/Storage5/projects2/3ST/attentive-neural-processes/src/train.py\u001b[0m in \u001b[0;36mrun_trial\u001b[0;34m(name, PL_MODEL_CLS, params, user_attrs, MODEL_DIR)\u001b[0m\n\u001b[1;32m 107\u001b[0m \u001b[0mdset_test\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mloader\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mdataset\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 108\u001b[0m \u001b[0mlabel_names\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mdset_test\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mlabel_names\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 109\u001b[0;31m \u001b[0mplot_from_loader\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mloader\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mmodel\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mi\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0;36m670\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m 110\u001b[0m \u001b[0;32mreturn\u001b[0m \u001b[0mtrial\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mtrainer\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mmodel\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;31mNameError\u001b[0m: name 'plot_from_loader' is not defined"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"trial = optuna.trial.FixedTrial(\n",
|
||||
"trial, trainer, model = run_trial(\n",
|
||||
" name=\"seq2seq\",\n",
|
||||
" params={\n",
|
||||
" 'bidirectional': False,\n",
|
||||
" 'hidden_size': 128,\n",
|
||||
" 'learning_rate': 2e-5,\n",
|
||||
" 'lstm_dropout': 0.0,\n",
|
||||
" 'lstm_layers': 2,\n",
|
||||
" })\n",
|
||||
"trial._user_attrs = {\n",
|
||||
" 'hidden_size': 256,\n",
|
||||
" 'learning_rate': 2e-4,\n",
|
||||
" 'lstm_dropout': 0.4,\n",
|
||||
" 'lstm_layers': 8,\n",
|
||||
" },\n",
|
||||
" user_attrs = {\n",
|
||||
" 'batch_size': 16,\n",
|
||||
" 'grad_clip': 40,\n",
|
||||
" 'max_nb_epochs': 200,\n",
|
||||
@@ -440,11 +537,11 @@
|
||||
" 'input_size': 18,\n",
|
||||
" 'input_size_decoder': 17,\n",
|
||||
" 'context_in_target': True,\n",
|
||||
" 'output_size': 1\n",
|
||||
"}\n",
|
||||
"trial = add_suggest(trial)\n",
|
||||
"model, trainer = main(trial, train=False)\n",
|
||||
"trainer.fit(model)"
|
||||
" 'output_size': 1,\n",
|
||||
" 'patience':2,\n",
|
||||
" },\n",
|
||||
" PL_MODEL_CLS=LSTMSeq2Seq_PL\n",
|
||||
")"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -452,35 +549,8 @@
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:31:43.522361Z",
|
||||
"start_time": "2020-03-01T03:31:15.800Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:47:33.307396Z",
|
||||
"start_time": "2020-03-01T03:41:28.300Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# test\n",
|
||||
"trainer.test(model)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:47:33.308984Z",
|
||||
"start_time": "2020-03-01T03:41:28.300Z"
|
||||
"end_time": "2020-03-15T01:51:46.863848Z",
|
||||
"start_time": "2020-03-15T01:04:18.100Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
@@ -496,8 +566,8 @@
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:47:33.310698Z",
|
||||
"start_time": "2020-03-01T03:41:28.300Z"
|
||||
"end_time": "2020-03-15T01:51:46.865188Z",
|
||||
"start_time": "2020-03-15T01:04:18.100Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
@@ -531,8 +601,8 @@
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:47:33.312313Z",
|
||||
"start_time": "2020-03-01T03:41:28.300Z"
|
||||
"end_time": "2020-03-15T01:51:46.866502Z",
|
||||
"start_time": "2020-03-15T01:04:18.100Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
@@ -569,8 +639,8 @@
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:47:33.313923Z",
|
||||
"start_time": "2020-03-01T03:41:28.300Z"
|
||||
"end_time": "2020-03-15T01:51:46.867929Z",
|
||||
"start_time": "2020-03-15T01:04:18.100Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
@@ -594,8 +664,8 @@
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:47:33.315413Z",
|
||||
"start_time": "2020-03-01T03:41:28.300Z"
|
||||
"end_time": "2020-03-15T01:51:46.869282Z",
|
||||
"start_time": "2020-03-15T01:04:18.100Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
@@ -608,8 +678,8 @@
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:47:33.317152Z",
|
||||
"start_time": "2020-03-01T03:41:28.300Z"
|
||||
"end_time": "2020-03-15T01:51:46.870600Z",
|
||||
"start_time": "2020-03-15T01:04:18.100Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
@@ -630,8 +700,8 @@
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:47:33.318533Z",
|
||||
"start_time": "2020-03-01T03:41:28.300Z"
|
||||
"end_time": "2020-03-15T01:51:46.872993Z",
|
||||
"start_time": "2020-03-15T01:04:18.100Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
@@ -644,8 +714,8 @@
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:47:33.319880Z",
|
||||
"start_time": "2020-03-01T03:41:28.300Z"
|
||||
"end_time": "2020-03-15T01:51:46.874727Z",
|
||||
"start_time": "2020-03-15T01:04:18.100Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
@@ -660,8 +730,8 @@
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:47:33.321408Z",
|
||||
"start_time": "2020-03-01T03:41:28.300Z"
|
||||
"end_time": "2020-03-15T01:51:46.876536Z",
|
||||
"start_time": "2020-03-15T01:04:18.100Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
@@ -676,8 +746,8 @@
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:47:33.322849Z",
|
||||
"start_time": "2020-03-01T03:41:28.300Z"
|
||||
"end_time": "2020-03-15T01:51:46.878187Z",
|
||||
"start_time": "2020-03-15T01:04:18.100Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
@@ -690,8 +760,8 @@
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:47:33.324307Z",
|
||||
"start_time": "2020-03-01T03:41:28.300Z"
|
||||
"end_time": "2020-03-15T01:51:46.879862Z",
|
||||
"start_time": "2020-03-15T01:04:18.200Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
@@ -709,8 +779,8 @@
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:47:33.325661Z",
|
||||
"start_time": "2020-03-01T03:41:28.300Z"
|
||||
"end_time": "2020-03-15T01:51:46.881522Z",
|
||||
"start_time": "2020-03-15T01:04:18.200Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
@@ -723,8 +793,8 @@
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2020-03-01T03:47:33.327059Z",
|
||||
"start_time": "2020-03-01T03:41:28.300Z"
|
||||
"end_time": "2020-03-15T01:51:46.883268Z",
|
||||
"start_time": "2020-03-15T01:04:18.200Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
@@ -747,6 +817,18 @@
|
||||
"language": "python",
|
||||
"name": "jup3.7.3"
|
||||
},
|
||||
"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.7.3"
|
||||
},
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"npconvert_exporter": "python",
|
||||
|
||||
File diff suppressed because it is too large.
Load diff
@@ -155,7 +155,7 @@ def get_smartmeter_df(indir=Path('./data/smart-meters-in-london'), use_logy=Fals
|
||||
df['dayofweek'] = time.dt.dayofweek / 7.0
|
||||
|
||||
# Drop nan and 0's
|
||||
df = df[df['energy(kWh/hh)']!=0]
|
||||
df = df[df['energy(kWh/hh)'] != 0]
|
||||
df = df.dropna()
|
||||
|
||||
if use_logy:
|
||||
@@ -163,7 +163,9 @@ def get_smartmeter_df(indir=Path('./data/smart-meters-in-london'), use_logy=Fals
|
||||
df = df.sort_values('tstp')
|
||||
|
||||
# split data
|
||||
n_split = -int(len(df)*0.1)
|
||||
df_train = df[:n_split]
|
||||
df_test = df[n_split:]
|
||||
return df_train, df_test
|
||||
test_split= -int(len(df) * 0.1)
|
||||
val_split= int(len(df) * 0.15)
|
||||
df_test = df[:val_split]
|
||||
df_train = df[val_split:test_split]
|
||||
df_val = df[test_split:]
|
||||
return df_train, df_val, df_test
|
||||
+45
-48
@@ -69,7 +69,7 @@ class LatentModelPL(pl.LightningModule):
|
||||
|
||||
# agg and print self.train_logs HACK https://github.com/PyTorchLightning/pytorch-lightning/issues/100
|
||||
train_logs = self.agg_logs(self.train_logs)
|
||||
train_logs_str = {k: f"{v.mean()}" for k, v in train_logs.items()}
|
||||
train_logs_str = {k: f"{v}" for k, v in train_logs.items()}
|
||||
self.train_logs = []
|
||||
print(f"step val {self.trainer.global_step}, {tensorboard_logs_str} {train_logs}")
|
||||
return logs
|
||||
@@ -95,10 +95,10 @@ class LatentModelPL(pl.LightningModule):
|
||||
if isinstance(outputs[0][j], dict):
|
||||
# Take mean of sub dicts
|
||||
keys = outputs[0][j].keys()
|
||||
aggs[j] = {k: torch.stack([x[j][k] for x in outputs if k in x[j]]).mean() for k in keys}
|
||||
aggs[j] = {k: torch.stack([x[j][k] for x in outputs if k in x[j]]).mean().item() for k in keys}
|
||||
else:
|
||||
# Take mean of numbers
|
||||
aggs[j] = torch.stack([x[j] for x in outputs if j in x]).mean()
|
||||
aggs[j] = torch.stack([x[j] for x in outputs if j in x]).mean().item()
|
||||
return aggs
|
||||
|
||||
# # Log hparams with metric, doesn't work
|
||||
@@ -117,15 +117,14 @@ class LatentModelPL(pl.LightningModule):
|
||||
return self.validation_end(*args, **kwargs)
|
||||
|
||||
def configure_optimizers(self):
|
||||
optim = torch.optim.Adam(self.parameters(), lr=self.hparams["learning_rate"], weight_decay=0)
|
||||
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optim, patience=1, verbose=True, min_lr=1e-7) # note early stopping has patience 3
|
||||
optim = torch.optim.AdamW(self.parameters(), lr=self.hparams["learning_rate"], weight_decay=0)
|
||||
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optim, patience=self.hparams["patience"], verbose=True, min_lr=1e-7) # note early stopping has patience 3
|
||||
return [optim], [scheduler]
|
||||
|
||||
def _get_cache_dfs(self):
|
||||
if self._dfs is None:
|
||||
df_train, df_test = get_smartmeter_df()
|
||||
# self._dfs = dict(df_train=df_train[:600], df_test=df_test[:600])
|
||||
self._dfs = dict(df_train=df_train, df_test=df_test)
|
||||
df_train, df_val, df_test = get_smartmeter_df()
|
||||
self._dfs = dict(df_train=df_train, df_val=df_val, df_test=df_test)
|
||||
return self._dfs
|
||||
|
||||
def train_dataloader(self):
|
||||
@@ -144,7 +143,7 @@ class LatentModelPL(pl.LightningModule):
|
||||
)
|
||||
|
||||
def val_dataloader(self):
|
||||
df_test = self._get_cache_dfs()['df_test']
|
||||
df_test = self._get_cache_dfs()['df_val']
|
||||
data_test = SmartMeterDataSet(
|
||||
df_test, self.hparams["num_context"], self.hparams["num_extra_target"]
|
||||
)
|
||||
@@ -172,49 +171,47 @@ class LatentModelPL(pl.LightningModule):
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def add_model_specific_args(parent_parser):
|
||||
"""
|
||||
Specify the hyperparams for this LightningModule
|
||||
"""
|
||||
# MODEL specific
|
||||
parser = HyperOptArgumentParser(strategy=parent_parser.strategy, parents=[parent_parser], add_help=False)
|
||||
parser.opt_range("--learning_rate", default=1e-3, type=float, tunable=True, high=1e-2, low=1e-5, log_base=10)
|
||||
def add_suggest(trial):
|
||||
trial.suggest_loguniform("learning_rate", 1e-5, 1e-2)
|
||||
|
||||
trial.suggest_categorical("hidden_dim", [8*2**i for i in range(6)])
|
||||
trial.suggest_categorical("latent_dim", [8*2**i for i in range(6)])
|
||||
|
||||
parser.opt_list("--hidden_dim", default=128, type=int, tunable=True, options=[8*2**i for i in range(8)])
|
||||
parser.opt_list("--latent_dim", default=128, type=int, tunable=True, options=[8*2**i for i in range(8)])
|
||||
parser.add_argument("--num_heads", default=8, type=int)
|
||||
parser.add_argument("--attention_layers", default=1, type=int)
|
||||
parser.opt_list("--n_latent_encoder_layers", default=4, type=int, tunable=True, options=[1, 2, 4, 8, 16])
|
||||
parser.opt_list("--n_det_encoder_layers", default=4, type=int, tunable=True, options=[1, 2, 4, 8, 16])
|
||||
parser.opt_list("--n_decoder_layers", default=2, type=int, tunable=True, options=[1, 2, 4, 8, 16])
|
||||
trial.suggest_int("attention_layers", 1, 4)
|
||||
trial.suggest_categorical("n_latent_encoder_layers", [1, 2, 4, 8])
|
||||
trial.suggest_categorical("n_det_encoder_layers", [1, 2, 4, 8])
|
||||
trial.suggest_categorical("n_decoder_layers", [1, 2, 4, 8])
|
||||
trial.suggest_int("num_heads", 8, 8)
|
||||
|
||||
parser.opt_range("--dropout", default=0, type=float, tunable=True, low=0, high=0.75)
|
||||
parser.opt_range("--attention_dropout", default=0, type=float, tunable=True, low=0, high=0.75)
|
||||
parser.add_argument("--min_std", default=0.005, type=float)
|
||||
trial.suggest_uniform("dropout", 0, 0.9)
|
||||
trial.suggest_uniform("attention_dropout", 0, 0.9)
|
||||
|
||||
parser.opt_list(
|
||||
"--latent_enc_self_attn_type", default="multihead", type=str, tunable=True, options=['uniform', 'dot', 'multihead', 'ptmultihead']
|
||||
trial.suggest_categorical(
|
||||
"latent_enc_self_attn_type", ['uniform', 'multihead', 'ptmultihead']
|
||||
)
|
||||
parser.opt_list("--det_enc_self_attn_type", default="multihead", type=str, tunable=True, options=['uniform', 'dot', 'multihead', 'ptmultihead'])
|
||||
parser.opt_list("--det_enc_cross_attn_type", default="multihead", type=str, tunable=True, options=['uniform', 'dot', 'multihead', 'ptmultihead'])
|
||||
trial.suggest_categorical("det_enc_self_attn_type", ['uniform', 'multihead', 'ptmultihead'])
|
||||
trial.suggest_categorical("det_enc_cross_attn_type", ['uniform', 'multihead', 'ptmultihead'])
|
||||
|
||||
parser.opt_list("--use_lvar", default=False, type=bool, tunable=True, options=[False, True])
|
||||
parser.opt_list("--use_rnn", default=False, type=bool, tunable=True, options=[False, True])
|
||||
parser.opt_list("--use_deterministic_path", default=True, tunable=True, type=bool, options=[False, True])
|
||||
parser.opt_list("--use_self_attn", default=True, tunable=True, type=bool, options=[False, True])
|
||||
parser.opt_list("--batchnorm", default=True, tunable=True, type=bool, options=[False, True])
|
||||
|
||||
# training specific (for this model)
|
||||
parser.add_argument("--context_in_target", default=True, type=bool)
|
||||
parser.add_argument("--grad_clip", default=0, type=float)
|
||||
parser.add_argument("--num_context", type=int, default=24 * 2)
|
||||
parser.add_argument("--num_extra_target", type=int, default=24)
|
||||
parser.add_argument("--max_nb_epochs", default=20, type=int)
|
||||
parser.add_argument("--num_workers", default=4, type=int)
|
||||
trial.suggest_categorical("batchnorm", [False, True])
|
||||
trial.suggest_categorical("use_self_attn", [False, True])
|
||||
trial.suggest_categorical("use_lvar", [False, True])
|
||||
trial.suggest_categorical("use_deterministic_path", [False, True])
|
||||
trial.suggest_categorical("use_rnn", [True, False])
|
||||
|
||||
trial._user_attrs = {
|
||||
'batch_size': 16,
|
||||
'grad_clip': 40,
|
||||
'max_nb_epochs': 200,
|
||||
'num_workers': 4,
|
||||
'num_context': 24* 4,
|
||||
'vis_i': '670',
|
||||
'num_extra_target': 24*4,
|
||||
'x_dim': 18,
|
||||
'context_in_target': True,
|
||||
'y_dim': 1,
|
||||
'patience': 3,
|
||||
'min_std': 0.005,
|
||||
}
|
||||
return trial
|
||||
|
||||
parser.add_argument("--batch_size", default=16, type=int)
|
||||
parser.add_argument("--x_dim", default=16, type=int)
|
||||
parser.add_argument("--y_dim", default=1, type=int)
|
||||
parser.add_argument("--vis_i", default=670, type=int)
|
||||
return parser
|
||||
|
||||
+37
-24
@@ -132,7 +132,7 @@ class LSTM_PL(pl.LightningModule):
|
||||
def validation_end(self, outputs):
|
||||
# TODO send an image to tensroboard, like in the lighting_anp.py file
|
||||
if int(self.hparams["vis_i"]) > 0:
|
||||
loader = self.val_dataloader()[0]
|
||||
loader = self.val_dataloader()
|
||||
vis_i = min(int(self.hparams["vis_i"]), len(loader.dataset))
|
||||
if isinstance(self.hparams["vis_i"], str):
|
||||
image = plot_from_loader(loader, self, vis_i=vis_i, window_len=self.hparams["window_length"])
|
||||
@@ -163,15 +163,14 @@ class LSTM_PL(pl.LightningModule):
|
||||
def configure_optimizers(self):
|
||||
optim = torch.optim.Adam(self.parameters(), lr=self.hparams["learning_rate"])
|
||||
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
|
||||
optim, patience=2, verbose=True, min_lr=1e-5
|
||||
optim, patience=self.hparams["patience"], verbose=True, min_lr=1e-5
|
||||
) # note early stopping has patient 3
|
||||
return [optim], [scheduler]
|
||||
|
||||
def _get_cache_dfs(self):
|
||||
if self._dfs is None:
|
||||
df_train, df_test = get_smartmeter_df()
|
||||
# self._dfs = dict(df_train=df_train[:600], df_test=df_test[:600])
|
||||
self._dfs = dict(df_train=df_train, df_test=df_test)
|
||||
df_train, df_val, df_test = get_smartmeter_df()
|
||||
self._dfs = dict(df_train=df_train, df_val=df_val, df_test=df_test)
|
||||
return self._dfs
|
||||
|
||||
@pl.data_loader
|
||||
@@ -193,7 +192,7 @@ class LSTM_PL(pl.LightningModule):
|
||||
|
||||
@pl.data_loader
|
||||
def val_dataloader(self):
|
||||
df_test = self._get_cache_dfs()["df_test"]
|
||||
df_test = self._get_cache_dfs()["df_val"]
|
||||
dset_test = SequenceDfDataSet(
|
||||
df_test,
|
||||
self.hparams,
|
||||
@@ -216,27 +215,41 @@ class LSTM_PL(pl.LightningModule):
|
||||
return DataLoader(dset_test, batch_size=self.hparams.batch_size, shuffle=False)
|
||||
|
||||
@staticmethod
|
||||
def add_model_specific_args(parent_parser):
|
||||
def add_suggest(trial: optuna.Trial):
|
||||
"""
|
||||
Specify the hyperparams for this LightningModule
|
||||
Add hyperparam ranges to an optuna trial and typical user attrs.
|
||||
|
||||
Usage:
|
||||
trial = optuna.trial.FixedTrial(
|
||||
params={
|
||||
'hidden_size': 128,
|
||||
}
|
||||
)
|
||||
trial = add_suggest(trial)
|
||||
trainer = pl.Trainer()
|
||||
model = LSTM_PL(dict(**trial.params, **trial.user_attrs), dataset_train,
|
||||
dataset_test, cache_base_path, norm)
|
||||
trainer.fit(model)
|
||||
"""
|
||||
# MODEL specific
|
||||
parser = HyperOptArgumentParser(parents=[parent_parser])
|
||||
parser.add_argument("--learning_rate", default=0.002, type=float)
|
||||
parser.add_argument("--batch_size", default=16, type=int)
|
||||
parser.add_argument("--lstm_dropout", default=0.5, type=float)
|
||||
parser.add_argument("--hidden_size", default=16, type=int)
|
||||
parser.add_argument("--input_size", default=8, type=int)
|
||||
parser.add_argument("--lstm_layers", default=8, type=int)
|
||||
parser.add_argument("--bidirectional", default=False, type=bool)
|
||||
trial.suggest_loguniform("learning_rate", 1e-6, 1e-2)
|
||||
trial.suggest_uniform("lstm_dropout", 0, 0.75)
|
||||
trial.suggest_categorical(
|
||||
"hidden_size", [1, 2, 4, 8, 16, 32, 64, 128, 256, 512, 1024]
|
||||
)
|
||||
trial.suggest_categorical("lstm_layers", [1, 2, 3, 4, 6, 8])
|
||||
trial.suggest_categorical("bidirectional", [False, True])
|
||||
|
||||
# training specific (for this model)
|
||||
parser.add_argument("--window_length", type=int, default=12)
|
||||
parser.add_argument("--target_length", type=int, default=2)
|
||||
parser.add_argument("--max_nb_epochs", default=10, type=int)
|
||||
parser.add_argument("--num_workers", default=4, type=int)
|
||||
|
||||
return parser
|
||||
trial._user_attrs = {
|
||||
"batch_size": 16,
|
||||
"grad_clip": 40,
|
||||
"max_nb_epochs": 200,
|
||||
"num_workers": 4,
|
||||
"vis_i": 670,
|
||||
"input_size": 6,
|
||||
"output_size": 1,
|
||||
"patience": 2,
|
||||
}
|
||||
return trial
|
||||
|
||||
|
||||
def plot_from_loader(loader, model, vis_i=670, n=1, window_len=0):
|
||||
|
||||
+47
-28
@@ -20,7 +20,7 @@ import torch
|
||||
import io
|
||||
import PIL
|
||||
from torchvision.transforms import ToTensor
|
||||
|
||||
from src.models.modules import BatchNormSequence
|
||||
from src.data.smart_meter import get_smartmeter_df
|
||||
|
||||
from src.utils import ObjectDict
|
||||
@@ -41,6 +41,9 @@ class Seq2SeqNet(nn.Module):
|
||||
self.hparams = hparams
|
||||
self._min_std = _min_std
|
||||
|
||||
|
||||
|
||||
self.norm_input = BatchNormSequence(self.hparams.input_size)
|
||||
self.encoder = nn.LSTM(
|
||||
input_size=self.hparams.input_size,
|
||||
hidden_size=self.hparams.hidden_size,
|
||||
@@ -49,6 +52,9 @@ class Seq2SeqNet(nn.Module):
|
||||
bidirectional=self.hparams.bidirectional,
|
||||
dropout=self.hparams.lstm_dropout,
|
||||
)
|
||||
self.multihead_attn = nn.MultiheadAttention(self.hparams.hidden_size, num_heads=8)
|
||||
|
||||
self.norm_target = BatchNormSequence(self.hparams.input_size_decoder)
|
||||
self.decoder = nn.LSTM(
|
||||
input_size=self.hparams.input_size_decoder,
|
||||
hidden_size=self.hparams.hidden_size,
|
||||
@@ -66,9 +72,23 @@ class Seq2SeqNet(nn.Module):
|
||||
|
||||
def forward(self, context_x, context_y, target_x, target_y=None):
|
||||
x = torch.cat([context_x, context_y], -1)
|
||||
|
||||
# Sometimes input normalisation can be important, an initial batch norm is a nice way to ensure this
|
||||
x = self.norm_input(x)
|
||||
target_x = self.norm_target(target_x)
|
||||
|
||||
_, (h_out, cell) = self.encoder(x)
|
||||
# hidden = [batch size, n layers * n directions, hid dim]
|
||||
# cell = [batch size, n layers * n directions, hid dim]
|
||||
|
||||
# context_x, d_encoded, target_x = k, v, q
|
||||
|
||||
# query, key, value = target_x, context_x, d_encoded
|
||||
attn_output, _ = self.multihead_attn(h_out.permute(1, 0, 2), h_out.permute(1, 0, 2), h_out.permute(1, 0, 2))
|
||||
h_out = attn_output.permute(1, 0, 2).contiguous()
|
||||
attn_output, _ = self.multihead_attn(cell.permute(1, 0, 2), cell.permute(1, 0, 2), cell.permute(1, 0, 2))
|
||||
cell = attn_output.permute(1, 0, 2).contiguous()
|
||||
|
||||
outputs, (_, _) = self.decoder(target_x, (h_out, cell))
|
||||
# output = [batch size, seq len, hid dim * n directions]
|
||||
|
||||
@@ -155,7 +175,7 @@ class LSTMSeq2Seq_PL(pl.LightningModule):
|
||||
|
||||
def show_image(self):
|
||||
# https://github.com/PytorchLightning/pytorch-lightning/blob/f8d9f8f/pytorch_lightning/core/lightning.py#L293
|
||||
loader = self.val_dataloader()[0]
|
||||
loader = self.val_dataloader()
|
||||
vis_i = min(int(self.hparams["vis_i"]), len(loader.dataset))
|
||||
# print('vis_i', vis_i)
|
||||
if isinstance(self.hparams["vis_i"], str):
|
||||
@@ -174,15 +194,14 @@ class LSTMSeq2Seq_PL(pl.LightningModule):
|
||||
def configure_optimizers(self):
|
||||
optim = torch.optim.Adam(self.parameters(), lr=self.hparams["learning_rate"])
|
||||
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
|
||||
optim, patience=2, verbose=True, min_lr=1e-5
|
||||
optim, patience=self.hparams["patience"], verbose=True, min_lr=1e-5
|
||||
) # note early stopping has patient 3
|
||||
return [optim], [scheduler]
|
||||
|
||||
def _get_cache_dfs(self):
|
||||
if self._dfs is None:
|
||||
df_train, df_test = get_smartmeter_df()
|
||||
# self._dfs = dict(df_train=df_train[:600], df_test=df_test[:600])
|
||||
self._dfs = dict(df_train=df_train, df_test=df_test)
|
||||
df_train, df_val, df_test = get_smartmeter_df()
|
||||
self._dfs = dict(df_train=df_train, df_val=df_val, df_test=df_test)
|
||||
return self._dfs
|
||||
|
||||
@pl.data_loader
|
||||
@@ -203,7 +222,7 @@ class LSTMSeq2Seq_PL(pl.LightningModule):
|
||||
|
||||
@pl.data_loader
|
||||
def val_dataloader(self):
|
||||
df_test = self._get_cache_dfs()['df_test']
|
||||
df_test = self._get_cache_dfs()['df_val']
|
||||
data_test = SmartMeterDataSet(
|
||||
df_test, self.hparams["num_context"], self.hparams["num_extra_target"]
|
||||
)
|
||||
@@ -232,25 +251,25 @@ class LSTMSeq2Seq_PL(pl.LightningModule):
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def add_model_specific_args(parent_parser):
|
||||
"""
|
||||
Specify the hyperparams for this LightningModule
|
||||
"""
|
||||
# MODEL specific
|
||||
parser = HyperOptArgumentParser(parents=[parent_parser])
|
||||
parser.add_argument("--learning_rate", default=0.002, type=float)
|
||||
parser.add_argument("--batch_size", default=16, type=int)
|
||||
parser.add_argument("--lstm_dropout", default=0.5, type=float)
|
||||
parser.add_argument("--hidden_size", default=16, type=int)
|
||||
parser.add_argument("--input_size", default=8, type=int)
|
||||
parser.add_argument("--input_size_decoder", default=8, type=int)
|
||||
parser.add_argument("--lstm_layers", default=8, type=int)
|
||||
parser.add_argument("--bidirectional", default=False, type=bool)
|
||||
def add_suggest(trial):
|
||||
trial.suggest_loguniform("learning_rate", 1e-5, 1e-2)
|
||||
trial.suggest_uniform("lstm_dropout", 0, 0.75)
|
||||
trial.suggest_categorical("hidden_size", [1, 2, 4, 8, 16, 32, 64, 128, 256, 512])
|
||||
trial.suggest_categorical("lstm_layers", [1, 2, 4, 8])
|
||||
trial.suggest_categorical("bidirectional", [False, True])
|
||||
|
||||
|
||||
# training specific (for this model)
|
||||
parser.add_argument("--num_context", type=int, default=12)
|
||||
parser.add_argument("--num_extra_target", type=int, default=2)
|
||||
parser.add_argument("--max_nb_epochs", default=10, type=int)
|
||||
parser.add_argument("--num_workers", default=4, type=int)
|
||||
|
||||
return parser
|
||||
trial._user_attrs = {
|
||||
'batch_size': 16,
|
||||
'grad_clip': 40,
|
||||
'max_nb_epochs': 200,
|
||||
'num_workers': 4,
|
||||
'num_extra_target': 24*4,
|
||||
'vis_i': '670',
|
||||
'num_context': 24*4,
|
||||
'input_size': 18,
|
||||
'input_size_decoder': 17,
|
||||
'context_in_target': True,
|
||||
'output_size': 1
|
||||
}
|
||||
return trial
|
||||
+28
-27
@@ -165,7 +165,7 @@ class LSTM_PL(pl.LightningModule):
|
||||
def validation_end(self, outputs):
|
||||
# TODO send an image to tensroboard, like in the lighting_anp.py file
|
||||
if int(self.hparams["vis_i"]) > 0:
|
||||
loader = self.val_dataloader()[0]
|
||||
loader = self.val_dataloader()
|
||||
vis_i = min(int(self.hparams["vis_i"]), len(loader.dataset))
|
||||
if isinstance(self.hparams["vis_i"], str):
|
||||
image = plot_from_loader(loader, self, vis_i=vis_i)
|
||||
@@ -196,15 +196,14 @@ class LSTM_PL(pl.LightningModule):
|
||||
def configure_optimizers(self):
|
||||
optim = torch.optim.Adam(self.parameters(), lr=self.hparams["learning_rate"])
|
||||
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
|
||||
optim, patience=2, verbose=True, min_lr=1e-5
|
||||
optim, patience=self.hparams["patience"], verbose=True, min_lr=1e-5
|
||||
) # note early stopping has patient 3
|
||||
return [optim], [scheduler]
|
||||
|
||||
def _get_cache_dfs(self):
|
||||
if self._dfs is None:
|
||||
df_train, df_test = get_smartmeter_df()
|
||||
# self._dfs = dict(df_train=df_train[:600], df_test=df_test[:600])
|
||||
self._dfs = dict(df_train=df_train, df_test=df_test)
|
||||
df_train, df_val, df_test = get_smartmeter_df()
|
||||
self._dfs = dict(df_train=df_train, df_val=df_val, df_test=df_test)
|
||||
return self._dfs
|
||||
|
||||
@pl.data_loader
|
||||
@@ -226,7 +225,7 @@ class LSTM_PL(pl.LightningModule):
|
||||
|
||||
@pl.data_loader
|
||||
def val_dataloader(self):
|
||||
df_test = self._get_cache_dfs()["df_test"]
|
||||
df_test = self._get_cache_dfs()["df_val"]
|
||||
dset_test = SequenceDfDataSet(
|
||||
df_test,
|
||||
self.hparams,
|
||||
@@ -249,27 +248,29 @@ class LSTM_PL(pl.LightningModule):
|
||||
return DataLoader(dset_test, batch_size=self.hparams.batch_size, shuffle=False)
|
||||
|
||||
@staticmethod
|
||||
def add_model_specific_args(parent_parser):
|
||||
"""
|
||||
Specify the hyperparams for this LightningModule
|
||||
"""
|
||||
# MODEL specific
|
||||
parser = HyperOptArgumentParser(parents=[parent_parser])
|
||||
parser.add_argument("--learning_rate", default=0.002, type=float)
|
||||
parser.add_argument("--batch_size", default=16, type=int)
|
||||
parser.add_argument("--lstm_dropout", default=0.5, type=float)
|
||||
parser.add_argument("--hidden_size", default=16, type=int)
|
||||
parser.add_argument("--input_size", default=8, type=int)
|
||||
parser.add_argument("--lstm_layers", default=8, type=int)
|
||||
parser.add_argument("--bidirectional", default=False, type=bool)
|
||||
|
||||
# training specific (for this model)
|
||||
parser.add_argument("--window_length", type=int, default=12)
|
||||
parser.add_argument("--target_length", type=int, default=2)
|
||||
parser.add_argument("--max_nb_epochs", default=10, type=int)
|
||||
parser.add_argument("--num_workers", default=4, type=int)
|
||||
|
||||
return parser
|
||||
def add_suggest(trial):
|
||||
trial.suggest_loguniform("learning_rate", 1e-5, 1e-2)
|
||||
trial.suggest_uniform("lstm_dropout", 0, 0.75)
|
||||
trial.suggest_categorical("hidden_size", [1, 2, 4, 8, 16, 32, 64, 128, 256, 512])
|
||||
trial.suggest_categorical("lstm_layers", [1, 2, 4, 8])
|
||||
trial.suggest_categorical("bidirectional", [False, True])
|
||||
|
||||
# constants
|
||||
trial._user_attrs = {
|
||||
'batch_size': 16,
|
||||
'grad_clip': 40,
|
||||
'max_nb_epochs': 200,
|
||||
'num_workers': 4,
|
||||
'num_extra_target': 24*4,
|
||||
'vis_i': '670',
|
||||
'num_context': 24*4,
|
||||
'input_size': 18,
|
||||
'input_size_decoder': 17,
|
||||
'context_in_target': True,
|
||||
'output_size': 1,
|
||||
'patience': 3,
|
||||
}
|
||||
return trial
|
||||
|
||||
|
||||
def plot_from_loader(loader, model, vis_i=670, n=1, window_len=0):
|
||||
|
||||
+31
-7
@@ -5,7 +5,7 @@ from torch.utils.data import TensorDataset, DataLoader
|
||||
import math
|
||||
|
||||
from src.models.modules import LatentEncoder, DeterministicEncoder, Decoder
|
||||
|
||||
from src.models.modules import BatchNormSequence
|
||||
|
||||
def log_prob_sigma(value, loc, log_scale):
|
||||
"""A slightly more stable (not confirmed yet) log prob taking in log_var instead of scale.
|
||||
@@ -66,18 +66,32 @@ class LatentModel(nn.Module):
|
||||
self._use_rnn = use_rnn
|
||||
self.context_in_target = context_in_target
|
||||
|
||||
# Sometimes input normalisation can be important, an initial batch norm is a nice way to ensure this
|
||||
self.norm_x = BatchNormSequence(x_dim)
|
||||
self.norm_y = BatchNormSequence(y_dim)
|
||||
|
||||
if self._use_rnn:
|
||||
self._lstm = nn.LSTM(
|
||||
self._lstm_x = nn.LSTM(
|
||||
input_size=x_dim,
|
||||
hidden_size=hidden_dim,
|
||||
num_layers=attention_layers,
|
||||
dropout=dropout,
|
||||
batch_first=True
|
||||
)
|
||||
self._lstm_y = nn.LSTM(
|
||||
input_size=y_dim,
|
||||
hidden_size=hidden_dim,
|
||||
num_layers=attention_layers,
|
||||
dropout=dropout,
|
||||
batch_first=True
|
||||
)
|
||||
x_dim = hidden_dim
|
||||
y_dim2 = hidden_dim
|
||||
else:
|
||||
y_dim2 = y_dim
|
||||
|
||||
self._latent_encoder = LatentEncoder(
|
||||
x_dim + y_dim,
|
||||
x_dim + y_dim2,
|
||||
hidden_dim=hidden_dim,
|
||||
latent_dim=latent_dim,
|
||||
self_attention_type=latent_enc_self_attn_type,
|
||||
@@ -93,7 +107,7 @@ class LatentModel(nn.Module):
|
||||
)
|
||||
|
||||
self._deterministic_encoder = DeterministicEncoder(
|
||||
input_dim=x_dim + y_dim,
|
||||
input_dim=x_dim + y_dim2,
|
||||
x_dim=x_dim,
|
||||
hidden_dim=hidden_dim,
|
||||
self_attention_type=det_enc_self_attn_type,
|
||||
@@ -126,16 +140,26 @@ class LatentModel(nn.Module):
|
||||
|
||||
def forward(self, context_x, context_y, target_x, target_y=None):
|
||||
|
||||
# https://stackoverflow.com/a/46772183/221742
|
||||
target_x = self.norm_x(target_x)
|
||||
context_x = self.norm_x(context_x)
|
||||
context_y = self.norm_y(context_y)
|
||||
|
||||
if self._use_rnn:
|
||||
# see https://arxiv.org/abs/1910.09323 where x is substituted with h = RNN(x)
|
||||
# x need to be provided as [B, T, H]
|
||||
target_x, _ = self._lstm(target_x)
|
||||
context_x, _ = self._lstm(context_x)
|
||||
target_x, _ = self._lstm_x(target_x)
|
||||
context_x, _ = self._lstm_x(context_x)
|
||||
context_y, _ = self._lstm_y(context_y)
|
||||
|
||||
|
||||
dist_prior, log_var_prior = self._latent_encoder(context_x, context_y)
|
||||
|
||||
if target_y is not None:
|
||||
dist_post, log_var_post = self._latent_encoder(target_x, target_y)
|
||||
target_y2 = self.norm_y(target_y)
|
||||
if self._use_rnn:
|
||||
target_y2, _ = self._lstm_y(target_y2)
|
||||
dist_post, log_var_post = self._latent_encoder(target_x, target_y2)
|
||||
z = dist_post.loc
|
||||
else:
|
||||
z = dist_prior.loc
|
||||
|
||||
@@ -24,6 +24,22 @@ class LSTMBlock(nn.Module):
|
||||
return self._lstm(x)[0]
|
||||
|
||||
|
||||
class BatchNormSequence(nn.Module):
|
||||
"""Applies batch norm on features of a batch first sequence."""
|
||||
def __init__(
|
||||
self, out_channels
|
||||
):
|
||||
super().__init__()
|
||||
self.norm = nn.BatchNorm1d(out_channels)
|
||||
|
||||
def forward(self, x):
|
||||
# x.shape is (Batch, Sequence, Channels)
|
||||
# Now we want to apply batchnorm and dropout to the channels. So we put it in shape
|
||||
# (Batch, Channels, Sequence) so we can use BatchNorm1d
|
||||
x = x.permute(0, 2, 1)
|
||||
x = self.norm(x)
|
||||
return x.permute(0, 2, 1)
|
||||
|
||||
class NPBlockRelu2d(nn.Module):
|
||||
"""Block for Neural Processes."""
|
||||
|
||||
|
||||
@@ -19,9 +19,11 @@ from matplotlib import pyplot as plt
|
||||
import torch
|
||||
import io
|
||||
import PIL
|
||||
import optuna
|
||||
from torchvision.transforms import ToTensor
|
||||
|
||||
from src.data.smart_meter import get_smartmeter_df
|
||||
from src.models.modules import BatchNormSequence
|
||||
|
||||
from src.utils import ObjectDict
|
||||
|
||||
@@ -41,7 +43,6 @@ class TransformerSeq2SeqNet(nn.Module):
|
||||
self.hparams = hparams
|
||||
self._min_std = _min_std
|
||||
|
||||
# TODO project to 8*nhead
|
||||
hidden_out_size = self.hparams.hidden_out_size
|
||||
self.enc_emb = nn.Linear(self.hparams.input_size, hidden_out_size)
|
||||
layer_enc = nn.TransformerEncoderLayer(
|
||||
@@ -92,7 +93,7 @@ class TransformerSeq2SeqNet(nn.Module):
|
||||
log_sigma = torch.clamp(log_sigma, math.log(self._min_std), -math.log(self._min_std))
|
||||
|
||||
sigma = torch.exp(log_sigma)
|
||||
y_dist=torch.distributions.Normal(mean, sigma)
|
||||
y_dist = torch.distributions.Normal(mean, sigma)
|
||||
|
||||
# Loss
|
||||
loss_mse = loss_p = None
|
||||
@@ -188,15 +189,14 @@ class TransformerSeq2Seq_PL(pl.LightningModule):
|
||||
def configure_optimizers(self):
|
||||
optim = torch.optim.Adam(self.parameters(), lr=self.hparams["learning_rate"])
|
||||
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
|
||||
optim, patience=2, verbose=True, min_lr=1e-5
|
||||
) # note early stopping has patient 3
|
||||
optim, patience=self.hparams["patience"], verbose=True, min_lr=1e-7
|
||||
) # note early stopping has patience 3
|
||||
return [optim], [scheduler]
|
||||
|
||||
def _get_cache_dfs(self):
|
||||
if self._dfs is None:
|
||||
df_train, df_test = get_smartmeter_df()
|
||||
# self._dfs = dict(df_train=df_train[:600], df_test=df_test[:600])
|
||||
self._dfs = dict(df_train=df_train, df_test=df_test)
|
||||
df_train, df_val, df_test = get_smartmeter_df()
|
||||
self._dfs = dict(df_train=df_train, df_val=df_val, df_test=df_test)
|
||||
return self._dfs
|
||||
|
||||
@pl.data_loader
|
||||
@@ -217,7 +217,7 @@ class TransformerSeq2Seq_PL(pl.LightningModule):
|
||||
|
||||
@pl.data_loader
|
||||
def val_dataloader(self):
|
||||
df_test = self._get_cache_dfs()['df_test']
|
||||
df_test = self._get_cache_dfs()['df_val']
|
||||
data_test = SmartMeterDataSet(
|
||||
df_test, self.hparams["num_context"], self.hparams["num_extra_target"]
|
||||
)
|
||||
@@ -246,27 +246,41 @@ class TransformerSeq2Seq_PL(pl.LightningModule):
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def add_model_specific_args(parent_parser):
|
||||
def add_suggest(trial: optuna.Trial):
|
||||
"""
|
||||
Specify the hyperparams for this LightningModule
|
||||
Add hyperparam ranges to an optuna trial and typical user attrs.
|
||||
|
||||
Usage:
|
||||
trial = optuna.trial.FixedTrial(
|
||||
params={
|
||||
'hidden_size': 128,
|
||||
}
|
||||
)
|
||||
trial = add_suggest(trial)
|
||||
trainer = pl.Trainer()
|
||||
model = LSTM_PL(dict(**trial.params, **trial.user_attrs), dataset_train,
|
||||
dataset_test, cache_base_path, norm)
|
||||
trainer.fit(model)
|
||||
"""
|
||||
# MODEL specific
|
||||
parser = HyperOptArgumentParser(parents=[parent_parser])
|
||||
parser.add_argument("--learning_rate", default=0.002, type=float)
|
||||
parser.add_argument("--batch_size", default=16, type=int)
|
||||
parser.add_argument("--attention_dropout", default=0.5, type=float)
|
||||
parser.add_argument("--hidden_size", default=16, type=int)
|
||||
parser.add_argument("--hidden_out_size", default=16, type=int)
|
||||
parser.add_argument("--input_size", default=8, type=int)
|
||||
parser.add_argument("--nhead", default=8, type=int)
|
||||
parser.add_argument("--input_size_decoder", default=8, type=int)
|
||||
parser.add_argument("--nlayers", default=8, type=int)
|
||||
# parser.add_argument("--bidirectional", default=False, type=bool)
|
||||
trial.suggest_loguniform("learning_rate", 1e-6, 1e-2)
|
||||
trial.suggest_uniform("attention_dropout", 0, 0.75)
|
||||
trial.suggest_categorical("hidden_size", [1, 2, 4, 8, 16, 32, 64, 128, 256, 512, 1024, 2048])
|
||||
trial.suggest_categorical("hidden_out_size", [1, 2, 4, 8, 16, 32, 64, 128, 256, 512, 1024, 2048])
|
||||
trial.suggest_categorical("nlayers", [1, 2, 4, 8])
|
||||
trial.suggest_categorical("nhead", [1, 2, 8, 16])
|
||||
|
||||
# training specific (for this model)
|
||||
parser.add_argument("--num_context", type=int, default=12)
|
||||
parser.add_argument("--num_extra_target", type=int, default=2)
|
||||
parser.add_argument("--max_nb_epochs", default=10, type=int)
|
||||
parser.add_argument("--num_workers", default=4, type=int)
|
||||
|
||||
return parser
|
||||
trial._user_attrs = {
|
||||
'batch_size': 16,
|
||||
'grad_clip': 40,
|
||||
'max_nb_epochs': 200,
|
||||
'num_workers': 4,
|
||||
'num_extra_target': 24*4,
|
||||
'vis_i': '670',
|
||||
'num_context': 24*4,
|
||||
'input_size': 18,
|
||||
'input_size_decoder': 17,
|
||||
'context_in_target': True,
|
||||
'output_size': 1,
|
||||
'patience': 3,
|
||||
}
|
||||
return trial
|
||||
+114
@@ -0,0 +1,114 @@
|
||||
from pytorch_lightning.callbacks import EarlyStopping
|
||||
from optuna.integration.pytorch_lightning import _check_pytorch_lightning_availability
|
||||
from pathlib import Path
|
||||
import optuna
|
||||
import pytorch_lightning as pl
|
||||
import torch
|
||||
from .dict_logger import DictLogger
|
||||
from .utils import PyTorchLightningPruningCallback
|
||||
from .plot import plot_from_loader
|
||||
|
||||
|
||||
def main(
|
||||
trial: optuna.Trial,
|
||||
PL_MODEL_CLS: pl.LightningModule,
|
||||
name: str,
|
||||
MODEL_DIR: Path = Path("./lightning_logs"),
|
||||
train=True,
|
||||
prune=True,
|
||||
PERCENT_TEST_EXAMPLES=0.5,
|
||||
):
|
||||
# PyTorch Lightning will try to restore model parameters from previous trials if checkpoint
|
||||
# filenames match. Therefore, the filenames for each trial must be made unique.
|
||||
|
||||
checkpoint_callback = pl.callbacks.ModelCheckpoint(
|
||||
MODEL_DIR / name / "version_{}".format(trial.number) / "chk",
|
||||
monitor="val_loss",
|
||||
mode="min",
|
||||
)
|
||||
|
||||
# The default logger in PyTorch Lightning writes to event files to be consumed by
|
||||
# TensorBoard. We create a simple logger instead that holds the log in memory so that the
|
||||
# final accuracy can be obtained after optimization. When using the default logger, the
|
||||
# final accuracy could be stored in an attribute of the `Trainer` instead.
|
||||
logger = DictLogger(MODEL_DIR, name=name, version=trial.number)
|
||||
# print("log_dir", logger.experiment.log_dir)
|
||||
hparams = dict(**trial.params, **trial.user_attrs)
|
||||
|
||||
trainer = pl.Trainer(
|
||||
logger=logger,
|
||||
val_percent_check=PERCENT_TEST_EXAMPLES,
|
||||
checkpoint_callback=checkpoint_callback,
|
||||
max_epochs=hparams["max_nb_epochs"],
|
||||
gpus=-1 if torch.cuda.is_available() else None,
|
||||
early_stop_callback=PyTorchLightningPruningCallback(trial, monitor="val_loss")
|
||||
if prune
|
||||
else EarlyStopping(
|
||||
patience=hparams["patience"] * 2, monitor="val_loss", verbose=True
|
||||
),
|
||||
)
|
||||
|
||||
model = PL_MODEL_CLS(hparams)
|
||||
if train:
|
||||
trainer.fit(model)
|
||||
return model, trainer
|
||||
|
||||
|
||||
def objective(trial, PL_MODEL_CLS):
|
||||
# see https://github.com/optuna/optuna/blob/cf6f02d/examples/pytorch_lightning_simple.py
|
||||
trial = PL_MODEL_CLS.add_suggest(trial)
|
||||
|
||||
print("trial", trial.number, "params", trial.params)
|
||||
|
||||
model, trainer = main(trial)
|
||||
|
||||
# also report to tensorboard & print
|
||||
print("logger.metrics", model.logger.metrics[-1:])
|
||||
model.logger.experiment.add_hparams(trial.params, logger.metrics[-1])
|
||||
model.logger.save()
|
||||
|
||||
return model.logger.metrics[-1]["val_loss"]
|
||||
|
||||
|
||||
def add_number(trial: optuna.Trial, model_dir: Path):
|
||||
# For manual experiment we will start at -1 and deincr by 1
|
||||
versions = [int(s.stem.split("_")[-1]) for s in model_dir.glob("version_*")] + [-1]
|
||||
trial.number = min(versions) - 1
|
||||
print("trial.number", trial.number)
|
||||
return trial
|
||||
|
||||
|
||||
def run_trial(
|
||||
name: str,
|
||||
PL_MODEL_CLS: pl.LightningModule,
|
||||
params: dict = {},
|
||||
user_attrs: dict = {},
|
||||
MODEL_DIR: Path = Path("./lightning_logs"),
|
||||
):
|
||||
print(f"now run `tensorboard --logdir {MODEL_DIR}`")
|
||||
(MODEL_DIR / name).mkdir(parents=True, exist_ok=True)
|
||||
trial = optuna.trial.FixedTrial(params=params)
|
||||
trial = PL_MODEL_CLS.add_suggest(trial)
|
||||
trial = add_number(trial, MODEL_DIR / name)
|
||||
trial._user_attrs.update(user_attrs)
|
||||
model, trainer = main(
|
||||
trial, PL_MODEL_CLS, name=name, MODEL_DIR=MODEL_DIR, train=False, prune=False
|
||||
)
|
||||
trainer.fit(model)
|
||||
|
||||
# Load checkpoint
|
||||
checkpoint = sorted(Path(trainer.checkpoint_callback.dirpath).glob("*.ckpt"))[-1]
|
||||
device = next(model.parameters()).device
|
||||
print(f"Loading checkpoint {checkpoint}")
|
||||
model = model.load_from_checkpoint(checkpoint).to(device)
|
||||
|
||||
trainer.test(model)
|
||||
|
||||
# Plot
|
||||
loader = model.val_dataloader()
|
||||
dset_test = loader.dataset
|
||||
label_names = dset_test.label_names
|
||||
plot_from_loader(model.val_dataloader(), model, i=670, title='val 670')
|
||||
plot_from_loader(model.train_dataloader(), model, i=670, title='train 670')
|
||||
plot_from_loader(model.test_dataloader(), model, i=670, title='test 670')
|
||||
return trial, trainer, model
|
||||
+21
-7
@@ -1,5 +1,18 @@
|
||||
from pytorch_lightning.callbacks import EarlyStopping
|
||||
from optuna.integration.pytorch_lightning import _check_pytorch_lightning_availability
|
||||
from pathlib import Path
|
||||
import numpy as np
|
||||
import torch
|
||||
import optuna
|
||||
|
||||
|
||||
def init_random_seed(seed):
|
||||
# https://pytorch.org/docs/stable/notes/randomness.html
|
||||
np.random.seed(seed)
|
||||
torch.random.manual_seed(seed)
|
||||
torch.backends.cudnn.deterministic = True
|
||||
torch.backends.cudnn.benchmark = False
|
||||
|
||||
|
||||
class PyTorchLightningPruningCallback(EarlyStopping):
|
||||
"""Optuna PyTorch Lightning callback to prune unpromising trials.
|
||||
@@ -20,10 +33,10 @@ class PyTorchLightningPruningCallback(EarlyStopping):
|
||||
how this dictionary is formatted.
|
||||
"""
|
||||
|
||||
def __init__(self, trial, monitor):
|
||||
def __init__(self, trial, monitor, **kwargs):
|
||||
# type: (optuna.trial.Trial, str) -> None
|
||||
|
||||
super(PyTorchLightningPruningCallback, self).__init__(monitor)
|
||||
super().__init__(monitor, **kwargs)
|
||||
|
||||
_check_pytorch_lightning_availability()
|
||||
|
||||
@@ -41,25 +54,26 @@ class PyTorchLightningPruningCallback(EarlyStopping):
|
||||
message = "Trial was pruned at epoch {}.".format(epoch)
|
||||
raise optuna.exceptions.TrialPruned(message)
|
||||
|
||||
|
||||
class ObjectDict(dict):
|
||||
"""
|
||||
Interface similar to an argparser
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
|
||||
def __setattr__(self, attr, value):
|
||||
self[attr] = value
|
||||
return self[attr]
|
||||
|
||||
|
||||
def __getattr__(self, attr):
|
||||
if attr.startswith('_'):
|
||||
if attr.startswith("_"):
|
||||
# https://stackoverflow.com/questions/10364332/how-to-pickle-python-object-derived-from-dict
|
||||
raise AttributeError
|
||||
return dict(self)[attr]
|
||||
|
||||
|
||||
@property
|
||||
def __dict__(self):
|
||||
return dict(self)
|
||||
|
||||
|
||||
Reference in new issue
Block a user