prophet/notebooks/diagnostics.ipynb

265 lines
80 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
},
"outputs": [],
"source": [
"%load_ext rpy2.ipython\n",
"%matplotlib inline\n",
"from fbprophet import Prophet\n",
"import pandas as pd\n",
"import numpy as np\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_peyton_manning.csv')\n",
"df['y'] = np.log(df['y'])\n",
"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": [
{
"data": {
"text/plain": [
"Initial log joint probability = -19.4685\n",
"Optimization terminated normally: \n",
" Convergence detected: relative gradient magnitude is below tolerance\n"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"%%R\n",
"library(prophet)\n",
"df <- read.csv('../examples/example_wp_peyton_manning.csv')\n",
"df$y <- log(df$y)\n",
"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",
"execution_count": 3,
2017-09-02 20:07:49 +00:00
"metadata": {
"input_hidden": true
},
2017-09-02 17:53:38 +00:00
"outputs": [
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAmEAAAF3CAYAAADtkpxQAAAABHNCSVQICAgIfAhkiAAAAAlwSFlz\nAAALEgAACxIB0t1+/AAAIABJREFUeJzsvXt8FNX9//+and0EEJASoYAEoxVRIEIgXNa2uBQFRbH4\nobXVz8+gIPGGLbX9WKkfPw3FkhY/aKrw0QSBL6ut16CACiiQVSSrIQgC4r3GRAGBYOSavcyc3x+7\nZzIzO3tLdrOzu+/n4+FDspeZOTtnznmd9+0IjDEGgiAIgiAIolOxpPoCCIIgCIIgshESYQRBEARB\nECmARBhBEARBEEQKIBFGEARBEASRAkiEEQRBEARBpAASYQRBEARBECmARBhBEARBEEQKIBFGEARB\nEASRAkiEEQRBEARBpAASYQRBEARBECnAmuoLiIVzzjkHBQUFqb4MgiCIjOLEiROav3v06JGiK+l8\nsrntRPJpaGjA0aNHo34uLURYQUEB6uvrU30ZBEEQGYXL5dL87XA4UnIdqSCb204kn+Li4pg+R+5I\ngiAIgiCIFEAijCAIgiAIIgWQCCMIgiAIgkgBSRNhs2bNQt++fTF8+HDltQcffBCXXnopRo4cicmT\nJ+PAgQPJOj1BEARBEISpERhjLBkHfvvtt9G9e3eUlJRg3759AIDjx4+jZ8+eAIDHHnsM+/fvx5NP\nPhn1WMXFxRSYTxAEkWAoOJ0gkkOsuiVplrAJEyagd+/emte4AAOAU6dOQRCEZJ2eIAiCIAjC1HR6\niYoHHngATqcTZ599NmpqasJ+rqqqClVVVQCAI0eOdNblEQRBEARBdAqdHpj/17/+FU1NTfjP//xP\nLF26NOznSktLUV9fj/r6evTp06cTr5AgCIIgCCL5pCw78qabbkJ1dXWqTk8QBEFkMSdOnND8RxCp\noFPdkZ999hkGDx4MAFi3bh0uvvjizjw9QRAEQQAAdu7cqfmbkhKIVJA0EXbjjTfC5XLh6NGjGDhw\nIBYsWIDXX38dn3zyCSwWC84777yYMiMJgiAIIpk0NTWhvLwcDocDdrs91ZdDZBFJE2HPPvtsyGuz\nZ89O1ukIgiAIIm6amprgdDpRU1ODnJwcbNmyhYQY0WlQxXyCIAgia2loaIAkSZAkCV6vN6R2GkEk\nExJhBEEQRNZSUFAAURQhiiJycnIoNozoVDq9ThhBEARBmIX8/HyUlJTgZz/7GcWEEZ0OiTCCIAgi\nq8nPz8fNN9+c6ssgshByRxIEQRAEQaQAEmEEQRAq3G43ysvL4Xa7U30pBEFkOOSOJAiCCOJ2uzFp\n0iR4vV4qV0AQRNIhSxhBEEQQl8sFr9dL5QoIgugUyBJGEAQRxOFwICcnR7GEUbmCzIXuLWEGSIQR\nBEEEsdvt2LJlC1wuF5UrIAgi6ZAIIwiCUGG320l8EQTRKVBMGEEQBEEQRAogSxhBEASRdZw4cULz\nd48ePVJ0JUQ2QyKMIAiCyDp27typ+ZsC9YlUQO5IgiAIgiCIFEAijCAIgiAIIgWQCCMIgiAIgkgB\nJMIIgiAIgiBSAIkwgiAIgiCIFEAijCAIgiAIIgWQCCMIgiAIgkgBJMIIgiAIgiBSABVrJQgdVVVV\nqKqqAgCUlpaitLQ0ru8fOHAA1113HQDg2muvRVlZGQDg008/hcvlAhAoDHnRRRdpvjdt2jQcPHgQ\n/fv3x/r16+O+7rKyMrz66qsAgHXr1mHAgAFxH4MwF42NjXjuuedQV1eHb7/9FoIgoE+fPigqKsLP\nf/5zFBYWxnW8AwcOKH1k9OjRHb6+f/3rX3jppZdw6NAheL1edO/eXenjkd4jCCIAiTCC6CQ++eQT\nRdz1798/RIQRhJp169bhb3/7G7xer+b1r776Cl999RW+++47LFmyJK5jHjx4ULPA6EgfrK2txSOP\nPBL3ewRBtEEijCASzIABA1BfXx/399pj/SIykx07duChhx6CLMsQBAGzZs3CjBkz8IMf/AAHDx7E\nli1b0NjYmNJr/Pjjj5V/l5WV4ZprroEgCFHfMwu0TRFhBkiEEUQMqF19K1euxIsvvoh33nkHgiCg\nuLgYf/zjH5GXlwfA2B1ZWlqK999/XzneggULsGDBAgDAn//8Z0ybNs3QHfnpp59i+fLl+Oyzz3Ds\n2DF4PB6cffbZGDFiBG699VYMHTq0M38GopNYunQpZFkGAPz617/GnXfeqbw3aNAg3HrrrZAkCQBQ\nXFwMABg1apRi5TJ6Xd2HgYDbnW9ife2112LatGkAgLfffhvPPfccPvroI5w5cwZ5eXkYN24cbrvt\nNsXFzfsqp6ysDGVlZRg1ahQOHjwY9j319REEQSKMIOLmt7/9rTJ5AcDWrVtx8uRJ/N///V/Cz9XQ\n0ICamhrNa8eOHUNNTQ3cbjeefvppnH/++Qk/L5E6jh07hg8//FD5++abbzb8nCiKCT/3qlWrsGzZ\nMs1r3377LdatWweXy4WnnnoKF1xwQcLPSxCJwO12w+VyweFwwG63p/pyYoJEGEHEyYABA7B48WJI\nkoTbbrsNx44dQ11dHY4ePYpzzjnH8DtVVVVYv359iPUrGhdffDGWLl2KwYMHo2fPnvD5fNiwYQPK\ny8vR2tqKNWvW4Pe//31C20ekFrUV6ayzzkLfvn0TctyysjJMmzYNt99+O4DQmLDm5mY8+eSTAIAe\nPXpgyZIlGDJkCJxOJ1asWIHjx49jyZIlWLZsGdavX69JYKmsrNQE+kd6jyCSgdvtxqRJk+D1epGT\nk4MtW7akhRAjERYH6aiyicRzxx134NxzzwUAjBw5Elu3bgUQmDzDibD2kpeXh1deeQVLlizBgQMH\n4PF4NO9/9dVXCT0fkb18+OGHiovzmmuuwahRowAAt99+O6qrq9HS0oL6+nplkkt31NZsICA8ifTF\n5XLB6/VCkiR4vV64XK60mKdJhMVIuqpsIvGcd955yr+7du2q/FufxZYI7r//frjd7rDvt7a2Jvyc\nRGrp37+/8u9Tp07hyJEj6NOnT1zH4GIqHk6ePKn8u1+/fsq/LRYL+vbti5aWFkiShO+//z7u6zEj\nO3fu1PxNgfrpjcPhQE5OjjJHp8v9TFqx1lmzZqFv374YPny48tp//dd/4eKLL8all16K66+/Hi0t\nLck6fcIxUtlEdmK1tq1d4sn4ijc77Pjx44oA6927N1544QXU1dXhueeei+s4RHrRu3dvDBs2TPn7\n6aefNvwcF1o2mw2AdhHwzTffGH4nUh/s3r278u9Dhw4p/5ZlGYcPHwYQiEM7++yzozWBIDodu92O\nLVu2YOHChWllJEmaCLvllluwceNGzWtXXnkl9u3bhz179uCiiy5CeXl5sk6fcLjKFkUxrVQ2YR7U\nk9cXX3wR1VphtVqVSdNqtaJ79+5oaWnBE088kdTrJFLP3XffDYslMDw/99xzqKqqwpEjR+D3+9HY\n2IiVK1fioYceAtBmOfv8889x8OBB+P3+sH1E3Qe//PJL+P1+5e/hw4crwf6vv/46du/ejVOnTmH5\n8uXKgnnMmDEZ4YokMhO73Y758+enjQADkuiOnDBhAhoaGjSvTZ48Wfn3+PHj8dJLLyXr9AmHq2yK\nCSPay5AhQ2Cz2eDz+fDMM8/gmWeeARC+un23bt0wZswY1NXV4fDhw5g6dSqAQIkCIrMZO3Ys/vSn\nP+Hvf/87fD6fJtCdc/nllwMArrrqKlRVVaG1tRXTp0/XiHc9+fn56NWrF1paWvDmm29izZo1AIB7\n770XQ4YMwR133IFly5bh+PHjuO222zTf7dmzJ+69994ktDZ1NDU1oaGhAQUFBam+FCJLSdnekStX\nrsTVV1+dqtO3i3RU2YR56Nu3LxYsWIALLrggZmvCQw89hMmTJ6Nnz57o3r07pk6dmlYWZKL9TJ8+\nHc8++yx++ctfYtCgQcjNzUXXrl1x3nnn4ec//zluueUWAAGvw0033YQ+ffrAZrOhqKgIK1euNDxm\nTk4OysvLcckll6BLly4
"text/plain": [
"<matplotlib.figure.Figure at 0x7fc9d27e0d10>"
]
},
"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'] == cutoff]\n",
"\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=cutoff, c='gray', lw=4, alpha=0.5)\n",
"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=cutoff + pd.Timedelta('365 days'), c='gray', lw=4,\n",
" 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. This dataframe can then be used to compute error measures of `yhat` vs. `y`."
]
},
{
"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": [
{
"data": {
"text/plain": [
"Initial log joint probability = -53.0208\n",
"Optimization terminated normally: \n",
" Convergence detected: relative gradient magnitude is below tolerance\n",
" ds y yhat yhat_lower yhat_upper cutoff\n",
"1 2014-01-21 10.542574 9.439116 8.827709 10.077209 2014-01-20\n",
"2 2014-01-22 10.004283 9.266698 8.635896 9.888154 2014-01-20\n",
"3 2014-01-23 9.732818 9.263058 8.645846 9.861981 2014-01-20\n",
"4 2014-01-24 9.866460 9.277099 8.666444 9.847888 2014-01-20\n",
"5 2014-01-25 9.370927 9.087197 8.440665 9.734763 2014-01-20\n",
"6 2014-01-26 9.490545 9.460199 8.826654 10.070429 2014-01-20\n"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"%%R\n",
"df.cv <- cross_validation(m, horizon = 730, units = 'days')\n",
"head(df.cv)"
]
},
{
"cell_type": "code",
"execution_count": 5,
"metadata": {},
"outputs": [
{
"data": {
"text/html": [
"<div>\n",
"<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>2014-01-21</td>\n",
" <td>9.439510</td>\n",
" <td>8.799215</td>\n",
" <td>10.080240</td>\n",
" <td>10.542574</td>\n",
" <td>2014-01-20</td>\n",
" </tr>\n",
" <tr>\n",
" <th>1</th>\n",
" <td>2014-01-22</td>\n",
" <td>9.267086</td>\n",
" <td>8.645900</td>\n",
" <td>9.882225</td>\n",
" <td>10.004283</td>\n",
" <td>2014-01-20</td>\n",
" </tr>\n",
" <tr>\n",
" <th>2</th>\n",
" <td>2014-01-23</td>\n",
" <td>9.263447</td>\n",
" <td>8.628803</td>\n",
" <td>9.852847</td>\n",
" <td>9.732818</td>\n",
" <td>2014-01-20</td>\n",
" </tr>\n",
" <tr>\n",
" <th>3</th>\n",
" <td>2014-01-24</td>\n",
" <td>9.277452</td>\n",
" <td>8.693226</td>\n",
" <td>9.897891</td>\n",
" <td>9.866460</td>\n",
" <td>2014-01-20</td>\n",
" </tr>\n",
" <tr>\n",
" <th>4</th>\n",
" <td>2014-01-25</td>\n",
" <td>9.087565</td>\n",
" <td>8.447306</td>\n",
" <td>9.728898</td>\n",
" <td>9.370927</td>\n",
" <td>2014-01-20</td>\n",
" </tr>\n",
" </tbody>\n",
"</table>\n",
"</div>"
],
"text/plain": [
" ds yhat yhat_lower yhat_upper y cutoff\n",
"0 2014-01-21 9.439510 8.799215 10.080240 10.542574 2014-01-20\n",
"1 2014-01-22 9.267086 8.645900 9.882225 10.004283 2014-01-20\n",
"2 2014-01-23 9.263447 8.628803 9.852847 9.732818 2014-01-20\n",
"3 2014-01-24 9.277452 8.693226 9.897891 9.866460 2014-01-20\n",
"4 2014-01-25 9.087565 8.447306 9.728898 9.370927 2014-01-20"
]
},
"execution_count": 5,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"from fbprophet.diagnostics import cross_validation\n",
"df_cv = cross_validation(m, horizon = '730 days')\n",
"df_cv.head()"
]
}
],
"metadata": {
"kernelspec": {
"display_name": "Python 2",
"language": "python",
"name": "python2"
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 2
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython2",
"version": "2.7.13"
}
},
"nbformat": 4,
"nbformat_minor": 1
}