prophet/notebooks/diagnostics.ipynb

903 lines
278 KiB
Text
Raw Normal View History

2017-09-02 17:53:38 +00:00
{
"cells": [
{
"cell_type": "code",
"execution_count": 1,
"metadata": {
"block_hidden": true
2017-09-02 17:53:38 +00:00
},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
2020-06-06 00:46:41 +00:00
"Importing plotly failed. Interactive plots will not work.\n"
]
}
],
2017-09-02 17:53:38 +00:00
"source": [
"%load_ext rpy2.ipython\n",
"%matplotlib inline\n",
"from fbprophet import Prophet\n",
"import pandas as pd\n",
"from matplotlib import pyplot as plt\n",
"import logging\n",
"logging.getLogger('fbprophet').setLevel(logging.ERROR)\n",
"import warnings\n",
"warnings.filterwarnings(\"ignore\")\n",
"df = pd.read_csv('../examples/example_wp_log_peyton_manning.csv')\n",
2017-09-02 17:53:38 +00:00
"m = Prophet()\n",
"m.fit(df)\n",
"future = m.make_future_dataframe(periods=366)"
]
},
{
"cell_type": "code",
"execution_count": 2,
"metadata": {
"block_hidden": true
},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"WARNING:rpy2.rinterface_lib.callbacks:R[write to console]: Loading required package: Rcpp\n",
"\n",
"WARNING:rpy2.rinterface_lib.callbacks:R[write to console]: Loading required package: rlang\n",
"\n",
"WARNING:rpy2.rinterface_lib.callbacks:R[write to console]: Disabling daily seasonality. Run prophet with daily.seasonality=TRUE to override this.\n",
"\n"
]
2017-09-02 17:53:38 +00:00
}
],
"source": [
"%%R\n",
"library(prophet)\n",
"df <- read.csv('../examples/example_wp_log_peyton_manning.csv')\n",
2017-09-02 17:53:38 +00:00
"m <- prophet(df)\n",
"future <- make_future_dataframe(m, periods=366)"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"Prophet includes functionality for time series cross validation to measure forecast error using historical data. This is done by selecting cutoff points in the history, and for each of them fitting the model using data only up to that cutoff point. We can then compare the forecasted values to the actual values. This figure illustrates a simulated historical forecast on the Peyton Manning dataset, where the model was fit to a initial history of 5 years, and a forecast was made on a one year horizon."
]
},
{
"cell_type": "code",
2020-06-06 00:46:41 +00:00
"execution_count": 2,
2017-09-02 20:07:49 +00:00
"metadata": {
"input_hidden": true
},
2017-09-02 17:53:38 +00:00
"outputs": [
{
"data": {
2020-06-06 00:46:41 +00:00
"application/vnd.jupyter.widget-view+json": {
"model_id": "a1e4d1b9644e4792b4bcbb0e24c1f0a1",
"version_major": 2,
"version_minor": 0
},
"text/plain": [
"HBox(children=(FloatProgress(value=0.0, max=3.0), HTML(value='')))"
]
},
"metadata": {},
"output_type": "display_data"
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"\n"
]
},
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAl4AAAFzCAYAAADv+wfzAAAABHNCSVQICAgIfAhkiAAAAAlwSFlzAAALEgAACxIB0t1+/AAAADh0RVh0U29mdHdhcmUAbWF0cGxvdGxpYiB2ZXJzaW9uMy4yLjEsIGh0dHA6Ly9tYXRwbG90bGliLm9yZy+j8jraAAAgAElEQVR4nOy9e3wU1f3//5qd3QQEEY2kgAQo1YJCMIFwWa24ioKiKH7Sam0/RgWNqFipba3Ur59PKJb4w4+SFqkSBEq0XhuUi1yqwCqQ1RguAl7wUmNQCJcgcs1eZs7vj90zmZmdvSV7md19Pz8PPyV7mTmzc+ac13m/3+f9FhhjDARBEARBEETCsaS6AQRBEARBENkCCS+CIAiCIIgkQcKLIAiCIAgiSZDwIgiCIAiCSBIkvAiCIAiCIJIECS+CIAiCIIgkYU11A6Lh3HPPRf/+/VPdDIIgiIzi+PHjmr/PPPPMFLUk+WTztROJp7GxEYcPHzZ8Ly2EV//+/dHQ0JDqZhAEQWQUTqdT87fD4UhJO1JBNl87kXhKSkpCvkeuRoIgCIIgiCRBwosgCIIgCCJJkPAiCIIgCIJIEgkTXpMnT0Z+fj6GDBmivPbYY49h6NChKCoqwrhx47Bv375EnZ4gCIIgCMJ0JEx43XHHHVi7dq3mtT/84Q/YuXMnduzYgeuvvx5//vOfE3V6giAIgiAI05Ew4TVmzBicc845mte6deum/PvkyZMQBCFRpycIgiAIgjAdSU8n8eijj6KmpgZnnXUWNm7cGPJz1dXVqK6uBgAcOnQoWc0jCIIgsgBKH0GkiqQH1//lL3/B3r178etf/xrPPPNMyM+Vl5ejoaEBDQ0N6NGjRxJbSBAEQRAEkRhStqvx17/+NWpra1N1eoIgCIIgiKSTVOH1xRdfKP9evnw5Bg0alMzTEwRBEARBpJSExXjdeuutcDqdOHz4MPr06YOZM2di9erV2LNnDywWC/r164fnnnsuUacnCIIgiKhwuVxwOp1wOByw2+2pbg6R4SRMeL388stBr02ZMiVRpyMIgiCImHG5XBg7diw8Hg9ycnKwfv16El9EQqHM9QRBEETWsWfPHuzZswcbNmxAv379IEkSPB5PUPFsgog3JLwIgiCIrGP//v3Yv38/+vbtiz59+kAUReTk5FCaCSLhJD2PF0EQBEGYhYKCApSVleHKK6+kGC8iKZDwIgiCILKagoIC3HbbbaluBpElkKuRIAiCIAgiSZDwIgiCUOFyuVBZWQmXy5XqphAEkYGQq5EgCCIApRYgCCLRkMWLIAgigNPphMfjodQCBEEkDBJeBEEQARwOB3Jycii1AEEQCYNcjQRBEAHsdjvWr19P5WMIgkgYJLwIgiBU2O12ElwEQSQMcjUSBEEQBEEkCRJeBEEQBEEQSYKEF0EQBEEQRJIg4UUQBEEQBJEkKLieIAiCyDqGDx+e6iYQWQoJL4IgCCLrOPPMM1PdBCJLIVcjQRAEQRBEkiDhRRAEQRAEkSRIeBEEQRAEQSQJEl4EQRAEQRBJgoLrCYIgiKxj3759mr979+6dopYQ2QYJL4LQUV1djerqagBAeXk5ysvLY/r+vn37cMMNNwAArr/+elRUVAAAPv/8czidTgCAw+HAT3/6U833Jk6ciP3796NXr15YuXJlzO2uqKjAqlWrAAArVqygiSQDaGpqwiuvvIL6+nocOHAAgiCgR48eKC4uxo033ojCwsKYjrdv3z6lj8QjncJLL72Ef/3rX2hubobH40HXrl2VPh7uPTPw+eefa/6m54VIFiS8CCJJ7NmzRxF0vXr1ChJeBKFmxYoVeOKJJ+DxeDSvf/PNN/jmm2/w/fff46mnnorpmPv379csKjrSB+vq6vD000/H/B5BZDskvAgizvTu3RsNDQ0xf689Vi4iM/nwww/x+OOPQ5ZlCIKAyZMno7S0FGeffTb279+P9evXo6mpKaVt/Oyzz5R/V1RU4LrrroMgCBHfI4hsh4QXQUSB2o23ePFivP7669i8eTMEQUBJSQn++Mc/Ii8vD4Cxq7G8vBzbtm1Tjjdz5kzMnDkTAPC///u/mDhxoqGr8fPPP8fChQvxxRdf4MiRI3C73TjrrLNw8cUX484778RFF12UzJ+BSBLPPPMMZFkGAPzyl7/Evffeq7zXt29f3HnnnZAkCQBQUlICABg2bJhizTJ6Xd2HAb9L/fjx4wD8/XTixIkAgPfeew+vvPIKPv30U5w+fRp5eXkYNWoU7rrrLsUdx/sqp6KiAhUVFRg2bBj2798f8j11+wgiWyHhRRAx8uCDDyoTFgBs2LABJ06cwN///ve4n6uxsREbN27UvHbkyBFs3LgRLpcLL7zwAn784x/H/bxE6jhy5Ag+/vhj5e/bbrvN8HOiKMb93EuWLMH8+fM1rx04cAArVqyA0+nE888/jwEDBsT9vAQRD1wuF5xOJxwOB+x2e6qbExISXgQRI71798acOXMgSRLuuusuHDlyBPX19Th8+DDOPfdcw+9UV1dj5cqVQVauSAwaNAjPPPMMLrjgAnTr1g1erxdr1qxBZWUlWltbsWzZMvzud7+L6/URqUVtLerSpQvy8/PjctyKigpMnDgR99xzD4DgGK+WlhY899xzAPzldJ566ikMHDgQNTU1WLRoEY4dO4annnoK8+fPx8qVKzWbUBYsWKAJ1g/3HkEkApfLhbFjx8Lj8SAnJwfr1683rfiiPF4x4HK5UFlZCZfLleqmEClk6tSpOO+889C3b18UFRUpr6snzHiRl5eH+vp6TJ06FQ6HA2PGjEFlZaXy/jfffBP3cxLZyccff6y4L6+77joMGzYMXbp0wT333IPu3bsDABoaGoKC/QnCDDidTng8HkiSBI/HY6odtHrI4hUl6aSmicTSr18/5d+dO3dW/p2ICemRRx4JK/RbW1vjfk4itfTq1Uv598mTJ3Ho0CH06NEjpmNwARULJ06cUP7ds2dP5d8WiwX5+fk4evQoJEnCDz/8EHN7CCLROBwO5OTkKHO0w+FIdZNCkjCL1+TJk5Gfn48hQ4Yor/3hD3/AoEGDMHToUNx00004evRook4fd9JJTROJxWptW6/EslMr1l1dx44dU0TXOeecg9deew319fV45ZVXYjoOkV6cc845GDx4sPL3Cy+8YPg5Lq5sNhsArfD/7rvvDL8Trg+eeeaZyr+bm5uVf8uyjIMHDwLwx5WdddZZkS6BIJKO3W7H+vXrMWvWLNMbRhImvO644w6sXbtW89rVV1+N3bt3Y+fOnfjpT3+qcZmYHa6mRVE0vZomzIl6wvrqq68iWiWsVqsyUVqtVnTt2hVHjx7Fs88+m9B2Eqnn/vvvh8XiH55feeUVVFdX49ChQ/D5fGhqasLixYvx+OOPA2izkH355ZfYv38/fD5fyD6i7oNff/01vF6v8vfgwYOVgP3Vq1djx44dOHnyJBYuXKgskkeMGIGcnJz4XzBBxAG73Y4ZM2aYWnQBCXQ1jhkzBo2NjZrXxo0bp/x79OjR+Ne//pWo08cdrqbTYccEYU4GDhwIm80Gr9eLF198ES+++CKA0FnmzzjjDIwYMQL19fU4ePAgJkyYAMCfToDIbEaOHIk//elPeOKJJ+Dz+TTB6pzLL78cAHDNNdeguroara2tmDRpkkaw6ykoKED37t1x9OhRvP3221i2bBkA4KGHHsLAgQMxdepUzJ8/H8eOHcNdd92l+W63bt3w0EMPJeBqU8fevXvR2NiI/v37p7opRBaRsuD6xYsX49prr03V6dtFuqhpwpzk5+dj5syZGDBgQNRWg8cffxzjxo1Dt27d0LVrV0yYMCGtLMVE+5k0aRJeeeUV/OIXv0Dfvn2Rm5uLzp07o1+/frjxxhtxxx13APB7F371q1+hR48esNlsKC4uxuLFiw2PmZOTg8rKSlx44YXo1KlT0Pt33nknnn76aYwYMQJdu3aFKIrIz8/HDTfcgBdffDGjUkns3bsXNTU12LhxI2pqamjTFJE0BMYYS9TBGxsbcf3112P37t2a1//yl7+goaEBy5YtC7kyU6/
2017-09-02 17:53:38 +00:00
"text/plain": [
"<Figure size 720x432 with 1 Axes>"
2017-09-02 17:53:38 +00:00
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"from fbprophet.diagnostics import cross_validation\n",
"df_cv = cross_validation(\n",
" m, '365 days', initial='1825 days', period='365 days')\n",
"cutoff = df_cv['cutoff'].unique()[0]\n",
"df_cv = df_cv[df_cv['cutoff'].values == cutoff]\n",
2017-09-02 17:53:38 +00:00
"\n",
"fig = plt.figure(facecolor='w', figsize=(10, 6))\n",
"ax = fig.add_subplot(111)\n",
"ax.plot(m.history['ds'].values, m.history['y'], 'k.')\n",
"ax.plot(df_cv['ds'].values, df_cv['yhat'], ls='-', c='#0072B2')\n",
"ax.fill_between(df_cv['ds'].values, df_cv['yhat_lower'],\n",
" df_cv['yhat_upper'], color='#0072B2',\n",
" alpha=0.2)\n",
"ax.axvline(x=pd.to_datetime(cutoff), c='gray', lw=4, alpha=0.5)\n",
2017-09-02 17:53:38 +00:00
"ax.set_ylabel('y')\n",
"ax.set_xlabel('ds')\n",
"ax.text(x=pd.to_datetime('2010-01-01'),y=12, s='Initial', color='black',\n",
" fontsize=16, fontweight='bold', alpha=0.8)\n",
"ax.text(x=pd.to_datetime('2012-08-01'),y=12, s='Cutoff', color='black',\n",
" fontsize=16, fontweight='bold', alpha=0.8)\n",
"ax.axvline(x=pd.to_datetime(cutoff) + pd.Timedelta('365 days'), c='gray', lw=4,\n",
2017-09-02 17:53:38 +00:00
" alpha=0.5, ls='--')\n",
"ax.text(x=pd.to_datetime('2013-01-01'),y=6, s='Horizon', color='black',\n",
" fontsize=16, fontweight='bold', alpha=0.8);"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"[The Prophet paper](https://peerj.com/preprints/3190.pdf) gives further description of simulated historical forecasts.\n",
"\n",
2017-09-02 20:07:49 +00:00
"This cross validation procedure can be done automatically for a range of historical cutoffs using the `cross_validation` function. We specify the forecast horizon (`horizon`), and then optionally the size of the initial training period (`initial`) and the spacing between cutoff dates (`period`). By default, the initial training period is set to three times the horizon, and cutoffs are made every half a horizon.\n",
2017-09-02 17:53:38 +00:00
"\n",
"The output of `cross_validation` is a dataframe with the true values `y` and the out-of-sample forecast values `yhat`, at each simulated forecast date and for each cutoff date. In particular, a forecast is made for every observed point between `cutoff` and `cutoff + horizon`. This dataframe can then be used to compute error measures of `yhat` vs. `y`.\n",
"\n",
"Here we do cross-validation to assess prediction performance on a horizon of 365 days, starting with 730 days of training data in the first cutoff and then making predictions every 180 days. On this 8 year time series, this corresponds to 11 total forecasts."
2017-09-02 17:53:38 +00:00
]
},
{
"cell_type": "code",
"execution_count": 4,
2017-09-02 20:07:49 +00:00
"metadata": {
"output_hidden": true
},
2017-09-02 17:53:38 +00:00
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"WARNING:rpy2.rinterface_lib.callbacks:R[write to console]: Making 11 forecasts with cutoffs between 2010-02-15 and 2015-01-20\n",
"\n"
]
},
2017-09-02 17:53:38 +00:00
{
"data": {
"text/plain": [
" ds y yhat yhat_lower yhat_upper cutoff\n",
"1 2010-02-16 8.242493 8.954992 8.423614 9.496403 2010-02-15\n",
"2 2010-02-17 8.008033 8.721365 8.226481 9.219106 2010-02-15\n",
"3 2010-02-18 8.045268 8.605072 8.103985 9.104483 2010-02-15\n",
"4 2010-02-19 7.928766 8.526855 8.023088 9.042035 2010-02-15\n",
"5 2010-02-20 7.745003 8.268741 7.757920 8.779416 2010-02-15\n",
"6 2010-02-21 7.866339 8.599935 8.084956 9.060284 2010-02-15\n"
2017-09-02 17:53:38 +00:00
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"%%R\n",
"df.cv <- cross_validation(m, initial = 730, period = 180, horizon = 365, units = 'days')\n",
2017-09-02 17:53:38 +00:00
"head(df.cv)"
]
},
{
"cell_type": "code",
2020-06-06 00:46:41 +00:00
"execution_count": 3,
2017-09-02 17:53:38 +00:00
"metadata": {},
"outputs": [
{
"data": {
"application/vnd.jupyter.widget-view+json": {
"model_id": "a27653066bb54f36a540ec95c975b083",
"version_major": 2,
"version_minor": 0
},
"text/plain": [
"HBox(children=(FloatProgress(value=0.0, max=11.0), HTML(value='')))"
]
},
"metadata": {},
"output_type": "display_data"
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"\n"
]
},
2017-09-02 17:53:38 +00:00
{
"data": {
"text/html": [
"<div>\n",
"<style scoped>\n",
" .dataframe tbody tr th:only-of-type {\n",
" vertical-align: middle;\n",
" }\n",
"\n",
" .dataframe tbody tr th {\n",
" vertical-align: top;\n",
" }\n",
"\n",
" .dataframe thead th {\n",
" text-align: right;\n",
" }\n",
"</style>\n",
2017-09-02 17:53:38 +00:00
"<table border=\"1\" class=\"dataframe\">\n",
" <thead>\n",
" <tr style=\"text-align: right;\">\n",
" <th></th>\n",
" <th>ds</th>\n",
" <th>yhat</th>\n",
" <th>yhat_lower</th>\n",
" <th>yhat_upper</th>\n",
" <th>y</th>\n",
" <th>cutoff</th>\n",
" </tr>\n",
" </thead>\n",
" <tbody>\n",
" <tr>\n",
" <th>0</th>\n",
" <td>2010-02-16</td>\n",
" <td>8.957284</td>\n",
" <td>8.471820</td>\n",
" <td>9.464242</td>\n",
" <td>8.242493</td>\n",
" <td>2010-02-15</td>\n",
2017-09-02 17:53:38 +00:00
" </tr>\n",
" <tr>\n",
" <th>1</th>\n",
" <td>2010-02-17</td>\n",
" <td>8.723736</td>\n",
" <td>8.215131</td>\n",
" <td>9.236662</td>\n",
" <td>8.008033</td>\n",
" <td>2010-02-15</td>\n",
2017-09-02 17:53:38 +00:00
" </tr>\n",
" <tr>\n",
" <th>2</th>\n",
" <td>2010-02-18</td>\n",
" <td>8.607496</td>\n",
" <td>8.059417</td>\n",
" <td>9.089188</td>\n",
" <td>8.045268</td>\n",
" <td>2010-02-15</td>\n",
2017-09-02 17:53:38 +00:00
" </tr>\n",
" <tr>\n",
" <th>3</th>\n",
" <td>2010-02-19</td>\n",
" <td>8.529364</td>\n",
" <td>8.019238</td>\n",
" <td>9.034076</td>\n",
" <td>7.928766</td>\n",
" <td>2010-02-15</td>\n",
2017-09-02 17:53:38 +00:00
" </tr>\n",
" <tr>\n",
" <th>4</th>\n",
" <td>2010-02-20</td>\n",
" <td>8.271329</td>\n",
" <td>7.779745</td>\n",
" <td>8.760580</td>\n",
" <td>7.745003</td>\n",
" <td>2010-02-15</td>\n",
2017-09-02 17:53:38 +00:00
" </tr>\n",
" </tbody>\n",
"</table>\n",
"</div>"
],
"text/plain": [
" ds yhat yhat_lower yhat_upper y cutoff\n",
"0 2010-02-16 8.957284 8.471820 9.464242 8.242493 2010-02-15\n",
"1 2010-02-17 8.723736 8.215131 9.236662 8.008033 2010-02-15\n",
"2 2010-02-18 8.607496 8.059417 9.089188 8.045268 2010-02-15\n",
"3 2010-02-19 8.529364 8.019238 9.034076 7.928766 2010-02-15\n",
"4 2010-02-20 8.271329 7.779745 8.760580 7.745003 2010-02-15"
2017-09-02 17:53:38 +00:00
]
},
2020-06-06 00:46:41 +00:00
"execution_count": 3,
2017-09-02 17:53:38 +00:00
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"from fbprophet.diagnostics import cross_validation\n",
"df_cv = cross_validation(m, initial='730 days', period='180 days', horizon = '365 days')\n",
2017-09-02 17:53:38 +00:00
"df_cv.head()"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"In R, the argument `units` must be a type accepted by `as.difftime`, which is weeks or shorter. In Python, the string for `initial`, `period`, and `horizon` should be in the format used by Pandas Timedelta, which accepts units of days or shorter.\n",
"\n",
"Custom cutoffs can also be supplied as a list of dates to to the `cutoffs` keyword in the `cross_validation` function in Python and R. For example, three cutoffs six months apart, would need to be passed to the `cutoffs` argument in a date format like below:\n",
"```python\n",
"#Python\n",
"cutoffs = [pd.Timestamp('2013-02-15'), pd.Timestamp('2013-08-15'), pd.Timestamp('2014-02-15')]\n",
"```\n",
"```r\n",
"#R\n",
"cutoffs = c(as.Date('2013-02-15'), as.Date('2013-08-15'), as.Date('2014-02-15')\n",
"```"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"### Cross-Validation in Parallel "
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"Cross-validation can also be run in parallel mode in Python, by setting specifying the `parallel` keyword. Three modes are supported\n",
"\n",
"* `parallel=\"processes\"`\n",
"* `parallel=\"threads\"`\n",
"* `parallel=\"dask\"`\n",
"\n",
"For problems that aren't too big, we recommend using `parallel=\"processes\"`. It will achieve the highest performance when the parallel cross validation can be done on a single machine. For large problems, a [Dask](https://dask.org) cluster can be used to do the cross validation on many machines. You will need to [install Dask](https://docs.dask.org/en/latest/install.html) separately, as it will not be installed with `fbprophet`.\n",
"\n",
"\n",
"```python\n",
"from dask.distributed import Client\n",
"\n",
"client = Client() # connect to the cluster\n",
"df_cv = cross_validation(m, initial='730 days', period='180 days', horizon='365 days',\n",
" parallel=\"dask\")\n",
"```\n",
"\n",
"The `performance_metrics` utility can be used to compute some useful statistics of the prediction performance (`yhat`, `yhat_lower`, and `yhat_upper` compared to `y`), as a function of the distance from the cutoff (how far into the future the prediction was). The statistics computed are mean squared error (MSE), root mean squared error (RMSE), mean absolute error (MAE), mean absolute percent error (MAPE), median absolute percent error (MDAPE) and coverage of the `yhat_lower` and `yhat_upper` estimates. These are computed on a rolling window of the predictions in `df_cv` after sorting by horizon (`ds` minus `cutoff`). By default 10% of the predictions will be included in each window, but this can be changed with the `rolling_window` argument."
]
},
{
"cell_type": "code",
"execution_count": 6,
"metadata": {
"output_hidden": true
},
"outputs": [
{
"data": {
"text/plain": [
" horizon mse rmse mae mape coverage\n",
"1 37 days 0.4971086 0.7050593 0.5075009 0.05882459 0.6765646\n",
"2 38 days 0.5029463 0.7091870 0.5125229 0.05940706 0.6765646\n",
"3 39 days 0.5252677 0.7247535 0.5186555 0.06001158 0.6751942\n",
"4 40 days 0.5326181 0.7298069 0.5215775 0.06032500 0.6788488\n",
"5 41 days 0.5401377 0.7349406 0.5226521 0.06041353 0.6838739\n",
"6 42 days 0.5438937 0.7374915 0.5230473 0.06043453 0.6891275\n"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"%%R\n",
"df.p <- performance_metrics(df.cv)\n",
"head(df.p)"
]
},
{
"cell_type": "code",
"execution_count": 3,
"metadata": {},
"outputs": [
{
"data": {
"text/html": [
"<div>\n",
"<style scoped>\n",
" .dataframe tbody tr th:only-of-type {\n",
" vertical-align: middle;\n",
" }\n",
"\n",
" .dataframe tbody tr th {\n",
" vertical-align: top;\n",
" }\n",
"\n",
" .dataframe thead th {\n",
" text-align: right;\n",
" }\n",
"</style>\n",
"<table border=\"1\" class=\"dataframe\">\n",
" <thead>\n",
" <tr style=\"text-align: right;\">\n",
" <th></th>\n",
" <th>horizon</th>\n",
" <th>mse</th>\n",
" <th>rmse</th>\n",
" <th>mae</th>\n",
" <th>mape</th>\n",
" <th>mdape</th>\n",
" <th>coverage</th>\n",
" </tr>\n",
" </thead>\n",
" <tbody>\n",
" <tr>\n",
" <th>0</th>\n",
" <td>37 days</td>\n",
" <td>0.493403</td>\n",
" <td>0.702427</td>\n",
" <td>0.504660</td>\n",
" <td>0.058481</td>\n",
" <td>0.050009</td>\n",
" <td>0.684102</td>\n",
" </tr>\n",
" <tr>\n",
" <th>1</th>\n",
" <td>38 days</td>\n",
" <td>0.499306</td>\n",
" <td>0.706616</td>\n",
" <td>0.509672</td>\n",
" <td>0.059061</td>\n",
" <td>0.049334</td>\n",
" <td>0.684102</td>\n",
" </tr>\n",
" <tr>\n",
" <th>2</th>\n",
" <td>39 days</td>\n",
" <td>0.521533</td>\n",
" <td>0.722172</td>\n",
" <td>0.515789</td>\n",
" <td>0.059662</td>\n",
" <td>0.049502</td>\n",
" <td>0.682732</td>\n",
" </tr>\n",
" <tr>\n",
" <th>3</th>\n",
" <td>40 days</td>\n",
" <td>0.528744</td>\n",
" <td>0.727148</td>\n",
" <td>0.518585</td>\n",
" <td>0.059961</td>\n",
" <td>0.049254</td>\n",
" <td>0.683874</td>\n",
" </tr>\n",
" <tr>\n",
" <th>4</th>\n",
" <td>41 days</td>\n",
" <td>0.536127</td>\n",
" <td>0.732207</td>\n",
" <td>0.519498</td>\n",
" <td>0.060030</td>\n",
" <td>0.049334</td>\n",
" <td>0.691412</td>\n",
" </tr>\n",
" </tbody>\n",
"</table>\n",
"</div>"
],
"text/plain": [
" horizon mse rmse mae mape mdape coverage\n",
"0 37 days 0.493403 0.702427 0.504660 0.058481 0.050009 0.684102\n",
"1 38 days 0.499306 0.706616 0.509672 0.059061 0.049334 0.684102\n",
"2 39 days 0.521533 0.722172 0.515789 0.059662 0.049502 0.682732\n",
"3 40 days 0.528744 0.727148 0.518585 0.059961 0.049254 0.683874\n",
"4 41 days 0.536127 0.732207 0.519498 0.060030 0.049334 0.691412"
]
},
"execution_count": 3,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"from fbprophet.diagnostics import performance_metrics\n",
"df_p = performance_metrics(df_cv)\n",
"df_p.head()"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"Cross validation performance metrics can be visualized with `plot_cross_validation_metric`, here shown for MAPE. Dots show the absolute percent error for each prediction in `df_cv`. The blue line shows the MAPE, where the mean is taken over a rolling window of the dots. We see for this forecast that errors around 5% are typical for predictions one month into the future, and that errors increase up to around 11% for predictions that are a year out."
]
},
{
"cell_type": "code",
"execution_count": 8,
"metadata": {
"output_hidden": true
},
"outputs": [
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAtAAAAGwCAIAAAAPKcUMAAAACXBIWXMAAAsSAAALEgHS3X78AAAgAElEQVR4nOydWYwr2V3/Ty3ey/vS7W73cu+d7rvOhEACo0kUJARZxAt5gweQ5iEPiIkIQkIkEPIAUiTmYQRBJCJCQigCCQUhEBL5Jw8wQSiDFCFImGTmzs3ce3u323uV91r+Dz/N4aTs9nV3u2x3+/t5cru9nCpXnfM9v1VyHIcBAAAAAHiJPO8BAAAAAOD6A8EBAAAAAM+B4AAAAACA50BwAAAAAMBz1HkPwE2v1zNNc7qf6ff7+/3+dD9zwZEkSZIk27bnPZCZoqqqbdvLdtSKoliWNe9RzBRFUSRJmvpEseAs4Q8tSZLP58PsvYBEIpELvGvhBEe/3+/1elP8QFmWI5FIo9GY4mcuPoqi+Hy+brc774HMlHg83uv1lm16ikQinU5n3qOYKaFQSFXVZTvqcDjc7XaXKq9QluVQKLRss7eqqoqiTHcdnDoXExxwqQAAAADAcyA4AAAAAOA5EBwAAAAA8BwIDgAAAAB4DgQHAAAAADwHggMAAAAAngPBAQAAAADPgeAAAAAAgOdAcAAAAADAcyA4AAAAAOA5EBwAAAAA8BwIDgAAAAB4DgQHAAAAADwHggMAAAAAnrNw7ekBWECOj4/L5TJjbHt7OxqNzns4AABw9YCFA4Bn0Gg0SG0wxp48eeI4znzHAwAAVxEIDgCewWAwEP+0bXteIwEAgKsLBAcAz0DTNP44FospijLHwQAAwBUFMRwAPINgMHjr1q1Go6EoSjqdnvdwAADgSgLBAcCzCYfD4XB43qMAAIArDFwqAAAAAPAcCA4AAAAAeA4EBwAAAAA8B4IDAAAAAJ4DwQEAAAAAz4HgAAAAAIDnQHAAAAAAwHMgOAAAAADgORAcAAAAAPAcCA4AAAAAeA4EBwAAAAA8B4IDAAAAAJ4DwQEAAAAAz4HgAAAAAIDnQHAAAAAAwHMgOAAAAADgORAcAAAAAPAcCA4AAAAAeA4EBwAAAAA8B4IDAAAAAJ4DwQEAAAAAz4HgAAAAAIDnQHAAAAAAwHMgOAAAAADgOaqnn25Z1muvvdZqtTY3N19++WV60nGcv/zLvyyVStFo9JVXXpEkydMxAAAAAGDueGvheOONNwqFwhe+8IWTk5ODgwP+ZCgU+tznPveTP/mTxWLR0wEAAAAAYBHw1sLx6NGju3fvMsZu3rz56NGjQqHAGHvrrbdkWf7TP/3Te/fura6u8lf+x3/8B2Psp3/6pzc3N6c4BrKghEKhKX7m4iPLsqIoy2Y9kmXZ7/crijLvgcwUVVWX7fL2+XyyLC/hUTPGHMeZ90BmhyRJkiQt2w8tv8e8BzJ9vBUchmFkMhnGWDqdNgyDnmy1Wq1W6+WXX/7KV76SzWbf9773Mcba7fbh4SFjrNPpTHfBoEV32RYhulGX8KiX7ZDZUh710l7e13IRGsNyzt6yLF/Xy9tbwRGJRCqVys2bNyuVSi6XoyfD4fBLL72Uy+U+8pGPvPPOOyQ4XnjhhRdeeIExpus6lyZTQZblQCAw3c9cfBRF8fl83W533gOZKfF4vNPp9Pv9eQ9kpkQikVarNe9RzJRQKKSq6rLd1OFwuNPpLJWFg2yWy/ZDq6qqKEqv15v3QMYRDAYv8C5v9fLOzs7jx48ZY0+fPt3Z2eFPvvPOO4yxd999d2VlxdMBAAAAAGAR8FZwvPjii0dHR6+++urq6mqhUHj48OGf/dmfvfjii48fP/7c5z5XrVZfeuklTwcAAAAAgEVAWjQDna7r0zUlybKcSqXK5fIUP3PxgUtlijiOY9v2wrpUl9alouv6vAcyU5bTpZJMJiuVyrwHMlOuhEuFojPPi7cxHABcdWq1GmV0JxKJQqGwbIk/AAAwLZYr5hmAc+E4Dq8fU6/Xm83mfMcDAABXFwgOAM7EsizxT9M05zUSAAC46kBwAHAmqqrG43H+ZywWm+NgAADgSoMYDgDGsbGxEYvFTNOMx+NU6hEAAMAFgOAAYBySJCUSiXmPAgAArjxwqQAAAADAcyA4AAAAAOA5EBwAAAAA8BwIDgAAAAB4DgQHAAAAADwHggMAAAAAngPBAQAAAADPgeAAAAAAgOdAcAAAAADAcyA4AAAAAOA5EBwAAAAA8BwIDgAAAAB4DgQHAAAAADwHggMAAAAAngPBAQAAAADPgeAAAAAAgOdAcAAAAADAcyA4AAAAAOA5EBwAAAAA8BwIDgAAAAB4DgQHAAAAADwHggMAAAAAngPBAQAAAADPgeAAAAAAgOdAcAAAAADAcyA4AAAAAOA5EBwAAAAA8BwIDgAAAAB4DgQHAAAAADwHggMAAAAAngPBAQAAAADPgeAAAAAAgOdAcAAAAADAcyA4AAAAAOA5EBwAAAAA8BwIDgAAAAB4DgQHAAAAADwHggMAAAAAngPBAQAAAADPgeAAAAAAgOdAcAAAAADAcyA4AAAAAOA5EBwAAAAA8Bx13gNwo6qqqk5zVJIkMcYikcgUP3PxkSRJlmVFUeY9kJmiKEowGPT5fPMeyEzx+XzLdnmrqirL8rIdtc/no9lseZAkSZKkZfuhZVmWJGm66+CCsHCHZJpmr9eb4gfKshwMBlut1hQ/c/FRFMXn83W73XkPZKaoqtrtdvv9/rwHMlMikciyXd6hUEhV1WU76nA43Ol0HMeZ90BmhyzLgUBg2X5oVVUVRZnuOjh1QqHQBd4FlwoAAAAAPAeCAwAAAACeA8EBAAAAAM+B4AAAAACA50BwAAAAAMBzIDgAAAAA4DkLlxYLwOJQq9Xa7XYoFEomk8tWAgEAAKYLBAcAoymXy8fHx/TYNM1cLjff8QAAwJUGLhUARiOWG2q323McCQAAXAMgOAAYjd/v54+vZZlhAACYJRAcAIwmm83G43HGWDweX1lZmfdwAADgaoN9GwCjUVV1c3Nz3qMAAIBrAiwcAAAAAPAcCA4AAAAAeA4EBwAAAAA8B4IDAAAAAJ4DwQEAAAAAz4HgAAAAAIDnQHAAAAAAwHMgOAAAAADgORAcAAAAAPAcCA4AAAAAeA4EBwAAAAA8B4IDAAAAAJ4DwQEAAAAAz4HgAAAAAIDnQHAAAAAAwHMgOAAAAADgORAcAAAAAPAcCA4AAAAAeA4EBwAAAAA8B4IDAAAAAJ4DwQEAAAAAz4HgAAAAAIDnQHAAAAAAwHMgOAAAAADgORAcAAAAAPAcCA4AAAAAeA4EBwAAAAA8B4IDAAAAAJ4DwQEAAAAAz4HgAAAAAIDnQHAAAAAAwHMgOAAAAADgORAcAAAAAPAcCA4AAAAAeA4EBwAAAAA8B4IDAAAAAJ4DwQEAAAAAz1E9/XTLsl577bVWq7W5ufnyyy+7/qvrejQa9XQAAAAAAFgEvLVwvPHGG4VC4Qtf+MLJycnBwYH4r3/913/96le/6um3AwAAAGBB8FZwPHr06ObNm4yxmzdvPnr0iD9fLBZff/11T78aAAAAAIuDty4VwzAymQxjLJ1OG4ZBT9q2/Vd/9Ve/+qu/+o//+I/8ld/+9rfJ4PGpT33qpZdemvpIEonE1D9zkZEkSZKkYDA474HMFEVRIpFIOBye90BmiizLPp9v3qOYKbIsS5K0bDe1LMt+v3/eo5g1S/hDS5LEGAuFQvMeyPTxVnBEIpFKpXLz5s1KpZLL5ejJr3/96x/96Edd0Rt379799Kc/zRjL5/OtVmuKY5AkKRaLTfczFx9FURRF6ff78x7ITIlEIv1+fzAYzHsgMyUYDHa73XmPYqYEAgFFUdrt9rwHMlOCwWCv13McZ94DmR2yLGuatoSztyzLCz6PxePxC7zLW8Gxs7Pz+PHjD37wg0+fPv3Qhz5ET7bb7X/+53/u9XqHh4f/8i//8olPfIIxls1ms9ksY0zX9V6vN8UxyLLMGFvwH2/q2LbNlu+oHccxTXPZjtrv9y/bIauqKknSsh21z+cbDAbLJjjYUs5jiqJcy6P2VnC8+OKLX/rSl1599dXV1dVCofDw4cNvfvObr7zyCmOsVCp97WtfI7UBAAAAgOuNtGh
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"%%R -w 10 -h 6 -u in\n",
"plot_cross_validation_metric(df.cv, metric = 'mape')"
]
},
{
"cell_type": "code",
"execution_count": 9,
"metadata": {},
"outputs": [
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAmQAAAF3CAYAAAALu1cUAAAABHNCSVQICAgIfAhkiAAAAAlwSFlzAAALEgAACxIB0t1+/AAAADl0RVh0U29mdHdhcmUAbWF0cGxvdGxpYiB2ZXJzaW9uIDMuMC4yLCBodHRwOi8vbWF0cGxvdGxpYi5vcmcvOIA7rQAAIABJREFUeJzsnX1sHMd5/7/3QvLuSN4dJeqNpExZpiiLciOblmLFLiy5gSBERtW6RVM3QREkcNUXAQka9CV9MxohRQIEcJFAQVs1bdLWTdyiQWIDSWQotiUFhi3VpmxXoSWStkiLR1Hkibw78o5H3tvvD/1mPLe3e7d3vOUuqe8HCGJSy93ZmdmZ7zzPM8+4CoVCAYQQQgghxDbcdheAEEIIIeROh4KMEEIIIcRmKMgIIYQQQmyGgowQQgghxGYoyAghhBBCbIaCjBBCCCHEZijICCGEEEJshoKMEEIIIcRmKMgIIYQQQmyGgowQQgghxGa8dhegWtrb27Ft2zbL7p9MJtHc3GzZ/VcrrJdSWCf6sF70Yb2UwjrRh/Wiz2qtl9HRUUSj0YrXrTpBtm3bNrzxxhuW3f/s2bM4ePCgZfdfrbBeSmGd6MN60Yf1UgrrRB/Wiz6rtV727t1r6jq6LAkhhBBCbIaCjBBCCCHEZijICCGEEEJshoKMEEIIIcRmKMgIIYQQQmyGgowQQgghxGYoyAghhBBCbIaCjBBCCCHEZijICCGEEEJshoKMEEIIIcRmKMgIcTjxeBxjY2OIx+N2F4UQQohFrLqzLAm5k4jH4zh37hzy+TzcbjcOHDiAUChkd7EIIYTUGVrICHEwsVgM+Xwe4XAY+XwesVjM7iIRQgixAAoyQhxMOByG2+1GLBaD2+1GOBy2u0iEEEIsgC5LQhxMKBTCgQMHEIvFEA6H6a4khJA1CgUZIQ4nFApRiBFCyBqHLktCCCGEEJuhICOEEEIIsRkKMkIIIYQQm6EgI4QQQgixGQoyQgghhBCboSAjhBBCCLEZCjJCCCGEEJuhICOEEEIIsRkKMkIIIYQQm6EgI4QQQgixGQoyQgghhBCboSAjhBBCCLEZCjJCCCGEEJuhICOEEEIIsRkKMkIIIYQQm6EgI4QQQgixGQoyQgghhBCboSAjhBBCCLEZSwXZ6dOnsXPnTvT09OBrX/ua7jX//d//jb6+PuzevRuf+tSnrCwOIYQQQogj8Vp141wuh+PHj+PMmTPo6urCvn37cPToUfT19clrhoeH8dWvfhWvvvoq2traMDU1ZVVxCCGEEEIci2UWsosXL6Knpwfbt29HY2MjnnzySTz//PNF1/zzP/8zjh8/jra2NgDAxo0brSoOIYQQQohjsUyQRSIRbN26Vf7c1dWFSCRSdM3Q0BCGhobwyCOPYP/+/Th9+rRVxSGEEEIIcSyWuSzNkM1mMTw8jLNnz2J8fByPPvoo/u///g/hcLjoulOnTuHUqVMAgPHxcZw9e9ayMs3Pz1t6/9UK66UU1ok+rBd9WC+lsE70Yb3os9brxTJB1tnZievXr8ufx8fH0dnZWXRNV1cXHnroITQ0NODuu+9Gb28vhoeHsW/fvqLrjh07hmPHjgEA9u7di4MHD1pVbJw9e9bS+69WWC+lsE70Yb3ow3ophXWiD+tFn7VeL5a5LPft24fh4WFcu3YNS0tLeO6553D06NGia379139dqt1oNIqhoSFs377dqiIRQgghhDgSywSZ1+vFyZMncfjwYezatQuf/OQnsXv3bjz99NN44YUXAACHDx/G+vXr0dfXh8ceewxf//rXsX79equKRAghhBDiSCyNITty5AiOHDlS9LsTJ07I/3a5XHjmmWfwzDPPWFkMQlYl8XgcsVgM4XAYoVDI7uIQQgixEFuD+gkh+sTjcZw7dw75fB5utxsHDhygKCOEkDUMj04ixIHEYjHk83mEw2Hk83nEYjG7i0QIIcRCKMgIcSDhcBhutxuxWAxut7skFQwhhJC1BV2WhDiQUCiEAwcOMIaMEELuECjICHEooVCIQowQQu4Q6LIkhBBCCLEZCjJCCCGEEJuhICOEEEIIsRkKMkIIIYQQm6EgI4QQQgixGQoyQgghhBCboSAjhBBCCLEZCjJCCCGEEJuhICOEEEIIsRkKMkIIIYQQm6EgI4QQQgixGQoyQgghhBCboSAjhBBCCLEZCjJCCCGEEJuhICOEEEIIsRkKMkIIIYQQm6EgI4QQQgixGQoyQgghhBCboSAjhBBCCLEZCjJCCCGEEJuhICOEEEIIsRkKMkIIIYQQm6EgI4QQQgixGQoyQgghhBCboSAjhBBCCLEZCjJCCCGEEJuhICOEEEIIsRkKMkIIIYQQm6EgI4QQQgixGQoyQgghhBCboSAjhBBCCLEZCjJCCCGEEJuxVJCdPn0aO3fuRE9PD772ta+V/Pt3v/tdbNiwAffffz/uv/9+fPvb37ayOIQQQgghjsRr1Y1zuRyOHz+OM2fOoKurC/v27cPRo0fR19dXdN1v//Zv4+TJk1YVgxBCCCHE8VhmIbt48SJ6enqwfft2NDY24sknn8Tzzz9v1eMIIYQQQlYtllnIIpEItm7dKn/u6urChQsXSq77wQ9+gPPnz6O3txd///d/X/Q3glOnTuHUqVMAgPHxcZw9e9aqYmN+ft7S+69WWC+lsE70Yb3ow3ophXWiD+tFn7VeL5YJMjP86q/+Kn7nd34HTU1N+Kd/+id85jOfwcsvv1xy3bFjx3Ds2DEAwN69e3Hw4EHLynT27FlL779aYb2UwjrRh/WiD+ulFNaJPqwXfdZ6vVjmsuzs7MT169flz+Pj4+js7Cy6Zv369WhqagIAPPXUU3jzzTetKg4hhBBCiGOxTJDt27cPw8PDuHbtGpaWlvDcc8/h6NGjRdfcuHFD/vcLL7yAXbt2WVUcQgghhBDHYpnL0uv14uTJkzh8+DByuRw+97nPYffu3Xj66aexd+9eHD16FN/85jfxwgsvwOv1Yt26dfjud79rVXEIIYQQQhyLpTFkR44cwZEjR4p+d+LECfnfX/3qV/HVr37VyiIQQgghhDgeZuonhBBCCLEZCjJCHEg8HsfY2Bji8bjdRSGEELIC2Jr2ghBSSjwex7lz55DP5+F2u3HgwAGEQiG7i0UIIcRCaCEjxGHEYjHk83mEw2Hk83nEYjG7i0QIIcRiKMgIcRjhcBhutxuxWAxutxvhcNjuIhFCCLEYuiwJcRihUAgHDhxALBZDOBymu5IQQu4AKMgIcSChUIhCjBBC7iDosiSEEEIIsRkKMkIIIYQQm6EgI4QQQgixGQoyQgghhBCboSAjhBBCCLEZCjJCCCGEEJuhICOEEEIIsRkKMkIIIYQQm6EgI4QQQgixGQoyQgghhBCboSAjhBBCCLEZCjJCCCGEEJuhICOEEEIIsRkKMkIIIYQQm6EgI4QQQgixGQoyQgghhBCboSAjhBBCCLEZCjJCCCGEEJuhICOEEEIIsRkKMkIIIYQQm6EgI4QQQgixGQoyQgghhBCboSAjhBBCCLEZCjJCCCGEEJuhICOEEEIIsRkKMkIIIYQQm6EgI4QQQgixGQoyQgghhBCboSAjhBBCCLEZSwXZ6dOnsXPnTvT09OBrX/ua4XU/+MEP4HK58MYbb1hZHEIIIYQQR2KZIMvlcjh+/Dh++tOfYnBwEN///vcxODhYct3c3By+8Y1v4KGHHrKqKIQQQgghjsYyQXbx4kX09PRg+/btaGxsxJNPPonnn3++5Lq/+Zu/wZ//+Z/D5/NZVRRCCCGEEEdjmSCLRCLYunWr/LmrqwuRSKTomoGBAVy/fh2PP/64VcUghBBCCHE8XrsenM/n8cUvfhHf/e53K1576tQpnDp1CgAwPj6Os2fPWlau+fl5S++/WmG9lMI60Yf1og/rpRTWiT6sF33Wer1YJsg6Oztx/fp1+fP4+Dg6Ozvlz3Nzc7h8+TIOHjwIAJicnMTRo0fxwgsvYO/evUX3OnbsGI4dOwYA2Lt3r/wbKzh79qyl91+tsF5KYZ3ow3rRh/VSCutEH9aLPmu9XixzWe7
"text/plain": [
"<Figure size 720x432 with 1 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"from fbprophet.plot import plot_cross_validation_metric\n",
"fig = plot_cross_validation_metric(df_cv, metric='mape')"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"The size of the rolling window in the figure can be changed with the optional argument `rolling_window`, which specifies the proportion of forecasts to use in each rolling window. The default is 0.1, corresponding to 10% of rows from `df_cv` included in each window; increasing this will lead to a smoother average curve in the figure. The `initial` period should be long enough to capture all of the components of the model, in particular seasonalities and extra regressors: at least a year for yearly seasonality, at least a week for weekly seasonality, etc.\n"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"### Hyperparameter Optimisation\n",
"\n",
"Auto parameter tuning can also be carried out by evaluating the parameter combinations in serial and using the in-built parallelization over cutoffs. An example implementation with multi-processing in Python is shown below, with a grid of six combinations of `changepoint_prior_scale` and `changepoint_range` parameters. The function `create_param_combinaitons` creates a dataframe of parameter combinations, which can be evaluated serially to call `single_cv_run` with the `parallel` keyword to parallelze over cutoffs. The best parameter combination is selected based on the best `rmse` score but can be switched to another performance metric depending on the use case.\n",
"\n",
"As an alternative for creating parameter combinations, one could also use the `ParameterGrid` class in `sklearn.model_selection`. This would need to be installed and imported separately if required, as it is not included with Prophet."
]
},
{
"cell_type": "code",
"execution_count": 2,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"\n",
" The best param combination is {'changepoint_prior_scale': 0.05, 'changepoint_range': 0.8}\n"
]
},
{
"data": {
"text/html": [
"<div>\n",
"<style scoped>\n",
" .dataframe tbody tr th:only-of-type {\n",
" vertical-align: middle;\n",
" }\n",
"\n",
" .dataframe tbody tr th {\n",
" vertical-align: top;\n",
" }\n",
"\n",
" .dataframe thead th {\n",
" text-align: right;\n",
" }\n",
"</style>\n",
"<table border=\"1\" class=\"dataframe\">\n",
" <thead>\n",
" <tr style=\"text-align: right;\">\n",
" <th></th>\n",
" <th>horizon</th>\n",
" <th>rmse</th>\n",
" <th>mape</th>\n",
" <th>params</th>\n",
" </tr>\n",
" </thead>\n",
" <tbody>\n",
" <tr>\n",
" <th>0</th>\n",
" <td>200 days</td>\n",
" <td>0.450030</td>\n",
" <td>0.034958</td>\n",
" <td>{'changepoint_prior_scale': 0.05, 'changepoint_range': 0.8}</td>\n",
" </tr>\n",
" <tr>\n",
" <th>1</th>\n",
" <td>200 days</td>\n",
" <td>0.453755</td>\n",
" <td>0.035471</td>\n",
" <td>{'changepoint_prior_scale': 0.05, 'changepoint_range': 0.9}</td>\n",
" </tr>\n",
" <tr>\n",
" <th>2</th>\n",
" <td>200 days</td>\n",
" <td>0.456887</td>\n",
" <td>0.035469</td>\n",
" <td>{'changepoint_prior_scale': 0.5, 'changepoint_range': 0.8}</td>\n",
" </tr>\n",
" <tr>\n",
" <th>3</th>\n",
" <td>200 days</td>\n",
" <td>0.490453</td>\n",
" <td>0.039134</td>\n",
" <td>{'changepoint_prior_scale': 0.5, 'changepoint_range': 0.9}</td>\n",
" </tr>\n",
" <tr>\n",
" <th>4</th>\n",
" <td>200 days</td>\n",
" <td>0.463969</td>\n",
" <td>0.036428</td>\n",
" <td>{'changepoint_prior_scale': 5.0, 'changepoint_range': 0.8}</td>\n",
" </tr>\n",
" <tr>\n",
" <th>5</th>\n",
" <td>200 days</td>\n",
" <td>0.512077</td>\n",
" <td>0.040488</td>\n",
" <td>{'changepoint_prior_scale': 5.0, 'changepoint_range': 0.9}</td>\n",
" </tr>\n",
" </tbody>\n",
"</table>\n",
"</div>"
],
"text/plain": [
" horizon rmse mape \\\n",
"0 200 days 0.450030 0.034958 \n",
"1 200 days 0.453755 0.035471 \n",
"2 200 days 0.456887 0.035469 \n",
"3 200 days 0.490453 0.039134 \n",
"4 200 days 0.463969 0.036428 \n",
"5 200 days 0.512077 0.040488 \n",
"\n",
" params \n",
"0 {'changepoint_prior_scale': 0.05, 'changepoint_range': 0.8} \n",
"1 {'changepoint_prior_scale': 0.05, 'changepoint_range': 0.9} \n",
"2 {'changepoint_prior_scale': 0.5, 'changepoint_range': 0.8} \n",
"3 {'changepoint_prior_scale': 0.5, 'changepoint_range': 0.9} \n",
"4 {'changepoint_prior_scale': 5.0, 'changepoint_range': 0.8} \n",
"5 {'changepoint_prior_scale': 5.0, 'changepoint_range': 0.9} "
]
},
"execution_count": 2,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"from fbprophet.diagnostics import cross_validation, performance_metrics\n",
"import itertools\n",
"\n",
"def create_param_combinations(**param_dict):\n",
" param_iter = itertools.product(*param_dict.values())\n",
" params =[]\n",
" for param in param_iter:\n",
" params.append(param) \n",
" params_df = pd.DataFrame(params, columns=list(param_dict.keys()))\n",
" return params_df\n",
"\n",
"def single_cv_run(history_df, metrics, param_dict, parallel):\n",
" m = Prophet(**param_dict)\n",
" m.fit(history_df)\n",
" df_cv = cross_validation(m, initial='2600 days', period='100 days', horizon = '200 days', parallel=parallel)\n",
" df_p = performance_metrics(df_cv, rolling_window=1)\n",
" df_p['params'] = str(param_dict)\n",
" df_p = df_p.loc[:, metrics]\n",
" return df_p\n",
"\n",
"\n",
"pd.set_option('display.max_colwidth', None)\n",
"param_grid = { \n",
" 'changepoint_prior_scale': [0.05, 0.5, 5],\n",
" 'changepoint_range': [0.8, 0.9],\n",
" }\n",
"metrics = ['horizon', 'rmse', 'mape', 'params'] \n",
"results = []\n",
"\n",
"\n",
"params_df = create_param_combinations(**param_grid)\n",
"for param in params_df.values:\n",
" param_dict = dict(zip(params_df.keys(), param))\n",
" cv_df = single_cv_run(df, metrics, param_dict, parallel=\"processes\")\n",
" results.append(cv_df)\n",
"results_df = pd.concat(results).reset_index(drop=True)\n",
"best_param = results_df.loc[results_df['rmse'] == min(results_df['rmse']), ['params']]\n",
"print(f'\\n The best param combination is {best_param.values[0][0]}')\n",
"results_df"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"Alternatively, in some cases one could benefit from parallelizing over parameter values instead, when the number of parameter combinations are large and the user has access to a large number of cores or a cluster. In the example below, parameter combinations are evaluated in parallel using `dask.distributed.Client`. The helper function `parallelize_param_combinations` parallelizes the calls to `single_cv_run` for each parameter combination. The cutoffs in `cross_validation` are then evaluated serially. To switch to other parallel modes in this example, import the `concurrent.futures` module and set `pool=concurrent.futures.ThreadPoolExecutor()`for execution in threads or `pool=concurrent.futures.ProcessPoolExecutor()` for multi-processing."
]
},
{
"cell_type": "code",
"execution_count": 3,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"\n",
" The best param combination is {'changepoint_prior_scale': 0.05, 'changepoint_range': 0.8}\n"
]
},
{
"data": {
"text/html": [
"<div>\n",
"<style scoped>\n",
" .dataframe tbody tr th:only-of-type {\n",
" vertical-align: middle;\n",
" }\n",
"\n",
" .dataframe tbody tr th {\n",
" vertical-align: top;\n",
" }\n",
"\n",
" .dataframe thead th {\n",
" text-align: right;\n",
" }\n",
"</style>\n",
"<table border=\"1\" class=\"dataframe\">\n",
" <thead>\n",
" <tr style=\"text-align: right;\">\n",
" <th></th>\n",
" <th>horizon</th>\n",
" <th>rmse</th>\n",
" <th>mape</th>\n",
" <th>params</th>\n",
" </tr>\n",
" </thead>\n",
" <tbody>\n",
" <tr>\n",
" <th>0</th>\n",
" <td>200 days</td>\n",
" <td>0.450030</td>\n",
" <td>0.034958</td>\n",
" <td>{'changepoint_prior_scale': 0.05, 'changepoint_range': 0.8}</td>\n",
" </tr>\n",
" <tr>\n",
" <th>1</th>\n",
" <td>200 days</td>\n",
" <td>0.453755</td>\n",
" <td>0.035471</td>\n",
" <td>{'changepoint_prior_scale': 0.05, 'changepoint_range': 0.9}</td>\n",
" </tr>\n",
" <tr>\n",
" <th>2</th>\n",
" <td>200 days</td>\n",
" <td>0.456887</td>\n",
" <td>0.035469</td>\n",
" <td>{'changepoint_prior_scale': 0.5, 'changepoint_range': 0.8}</td>\n",
" </tr>\n",
" <tr>\n",
" <th>3</th>\n",
" <td>200 days</td>\n",
" <td>0.490453</td>\n",
" <td>0.039134</td>\n",
" <td>{'changepoint_prior_scale': 0.5, 'changepoint_range': 0.9}</td>\n",
" </tr>\n",
" <tr>\n",
" <th>4</th>\n",
" <td>200 days</td>\n",
" <td>0.463969</td>\n",
" <td>0.036428</td>\n",
" <td>{'changepoint_prior_scale': 5.0, 'changepoint_range': 0.8}</td>\n",
" </tr>\n",
" <tr>\n",
" <th>5</th>\n",
" <td>200 days</td>\n",
" <td>0.512077</td>\n",
" <td>0.040488</td>\n",
" <td>{'changepoint_prior_scale': 5.0, 'changepoint_range': 0.9}</td>\n",
" </tr>\n",
" </tbody>\n",
"</table>\n",
"</div>"
],
"text/plain": [
" horizon rmse mape \\\n",
"0 200 days 0.450030 0.034958 \n",
"1 200 days 0.453755 0.035471 \n",
"2 200 days 0.456887 0.035469 \n",
"3 200 days 0.490453 0.039134 \n",
"4 200 days 0.463969 0.036428 \n",
"5 200 days 0.512077 0.040488 \n",
"\n",
" params \n",
"0 {'changepoint_prior_scale': 0.05, 'changepoint_range': 0.8} \n",
"1 {'changepoint_prior_scale': 0.05, 'changepoint_range': 0.9} \n",
"2 {'changepoint_prior_scale': 0.5, 'changepoint_range': 0.8} \n",
"3 {'changepoint_prior_scale': 0.5, 'changepoint_range': 0.9} \n",
"4 {'changepoint_prior_scale': 5.0, 'changepoint_range': 0.8} \n",
"5 {'changepoint_prior_scale': 5.0, 'changepoint_range': 0.9} "
]
},
"execution_count": 3,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"from dask.distributed import Client\n",
"import functools\n",
"from concurrent.futures import ThreadPoolExecutor, ProcessPoolExecutor\n",
"\n",
"\n",
"def parallelize_param_combinations(history_df, params_df, single_cv_callable, pool):\n",
" results = []\n",
" for param in params_df.values:\n",
" param_dict = dict(zip(params_df.keys(), param))\n",
" if isinstance(pool,(ThreadPoolExecutor, ProcessPoolExecutor)):\n",
" future = pool.submit(single_cv_callable, history_df, param_dict=param_dict)\n",
" results.append(future.result())\n",
" elif isinstance(pool, Client):\n",
" remote_df = pool.scatter(history_df)\n",
" future = pool.submit(single_cv_callable, remote_df, param_dict=param_dict)\n",
" results.append(future)\n",
" if isinstance(pool, Client):\n",
" results = pool.gather(results)\n",
" results_df = pd.concat(results).reset_index(drop=True)\n",
" \n",
" return results_df\n",
"\n",
"\n",
"single_cv_callable = functools.partial(single_cv_run, metrics=metrics, parallel=None)\n",
"\n",
"pool = Client()\n",
"results_df = parallelize_param_combinations(df, params_df, single_cv_callable, pool=pool)\n",
"best_param = results_df.loc[results_df['rmse'] == min(results_df['rmse']), ['params']]\n",
"print(f'\\n The best param combination is {best_param.values[0][0]}')\n",
"results_df"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"### Recommended Hyperparameter Ranges"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"In the examples above, we have used recommended initial settings for `changepoint_prior_scale:[0.05, 0.5, 5]` and \n",
"`changepoint_range: [0.8, 0.9]`. We could alternatively also use a random search to carry out a sweep between a range of values e.g. `np.random.uniform(0.05, 5, 3)`. Other parameters like the `seasonality_prior_scale`, `holidays_prior_scale` and `seasonality_mode` could also be optimised for. For the seasonality and holiday prior scales, recommended values to start with are `[0.1,1,10]`, and it is better to set these values on a log scale in the grid e.g. `np.random.uniform(-1, 1, 5)`."
]
2017-09-02 17:53:38 +00:00
}
],
"metadata": {
"kernelspec": {
2020-06-06 00:46:41 +00:00
"display_name": "Python (prophet)",
2017-09-02 17:53:38 +00:00
"language": "python",
2020-06-06 00:46:41 +00:00
"name": "prophet"
2017-09-02 17:53:38 +00:00
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 3
2017-09-02 17:53:38 +00:00
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
2020-06-06 00:46:41 +00:00
"version": "3.8.3"
2017-09-02 17:53:38 +00:00
}
},
"nbformat": 4,
"nbformat_minor": 1
}