mirror of
https://github.com/wassname/Volt.git
synced 2026-10-03 12:10:38 +08:00
first
This commit is contained in:
1 parent
1f6f05ed40
commit
f293a6c489
721 files changed
+973301
No files matched your search
Vendored
BIN
Binary file not shown.
@@ -0,0 +1,637 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "0d67e226",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import seaborn as sns\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"from matplotlib.lines import Line2D\n",
|
||||
"from matplotlib.patches import Patch\n",
|
||||
"import torch\n",
|
||||
"import pandas as pd\n",
|
||||
"import copy\n",
|
||||
"from voltron.option_utils import GetTradingDays, GetTrainingData, Pricer, FindLastTradingDays\n",
|
||||
"from scipy.optimize import minimize"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "8ad5b4c2",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/Users/gregorybenton/miniconda3/lib/python3.8/site-packages/IPython/core/interactiveshell.py:3169: DtypeWarning: Columns (19) have mixed types.Specify dtype option on import or set low_memory=False.\n",
|
||||
" has_raised = await self.run_ast_nodes(code_ast.body, cell_name,\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"SPY = pd.read_csv(\"./data/SPY_prices.csv\")\n",
|
||||
"SPY['Date'] = pd.to_datetime(SPY['Date'])\n",
|
||||
"\n",
|
||||
"years = np.arange(2009, 2018)\n",
|
||||
"opts = pd.DataFrame()\n",
|
||||
"for year in years:\n",
|
||||
" dat = pd.read_csv(\"./data/SPY_\" + str(year) + \".csv\")\n",
|
||||
" dat = dat[dat.type == \"call\"]\n",
|
||||
" quotedate = dat.quotedate.unique()[0]\n",
|
||||
" dat = dat[dat.quotedate == quotedate] \n",
|
||||
" opts = pd.concat((opts, dat), ignore_index=True)\n",
|
||||
"\n",
|
||||
"# exps = [] \n",
|
||||
"# for idx, row in opts.iterrows():\n",
|
||||
"# eday = row.expiration\n",
|
||||
"# exps.append(SPY[SPY.Date == FindLastTradingDays(SPY, [pd.Timestamp(eday)])[0]].Close.item())\n",
|
||||
"# opts['exp_price'] = exps"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "f361417a",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Torch Attempt"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"id": "06969ba7",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ivol = torch.tensor(opts.impliedvol.to_numpy())\n",
|
||||
"Fs = torch.tensor(opts.underlying_last.to_numpy())\n",
|
||||
"Ks = torch.tensor(opts.strike.to_numpy())\n",
|
||||
"qdays = pd.to_datetime(opts.quotedate).dt.date.to_numpy()\n",
|
||||
"edays = pd.to_datetime(opts.expiration).dt.date.to_numpy()\n",
|
||||
"Ts = torch.tensor(([np.busday_count(qd, ed)/252. for qd, ed in zip(qdays, edays)]))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"id": "465286fb",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def BlackVol(pars, K, f, T):\n",
|
||||
" alpha = torch.exp(pars[0][0])\n",
|
||||
" rho = 2 * torch.sigmoid(pars[0][1]) - 1.\n",
|
||||
" v = torch.exp(pars[0][2])\n",
|
||||
" beta = 1.\n",
|
||||
" num = 1 + (alpha**2 * (1-beta)**2/(24 * (f*K)**(1-beta)) + 0.25 * rho*beta*v*alpha/((f*K)**(0.5*(1-beta))) + v**2*(2-3*rho**2)/24)*T\n",
|
||||
" num*= alpha\n",
|
||||
" \n",
|
||||
" denom = (f*K)**(0.5*(1-beta)) * (1 + (1-beta)**2/24 * torch.log(f/K)**2 + (1-beta)**4/1920 * torch.log(f/K)**4)\n",
|
||||
" \n",
|
||||
" z = v/alpha * (f*K)**(0.5*(1-beta)) * np.log(f/K)\n",
|
||||
" xi_z = torch.log((torch.sqrt(1 - 2 * rho * z + z**2) + z - rho)/(1-rho))\n",
|
||||
" \n",
|
||||
" return num/denom * z/xi_z"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 5,
|
||||
"id": "1ebbc187",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def MinVol(pars):\n",
|
||||
" return torch.mean((ivol - BlackVol(pars, Ks, Fs, Ts)).pow(2))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 6,
|
||||
"id": "57bc7a65",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"pars = [torch.tensor([-1., -5., -3.], requires_grad=True)]\n",
|
||||
"opt = torch.optim.SGD(pars, lr=0.1)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 7,
|
||||
"id": "ab9c2a60",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"iters = 1000\n",
|
||||
"stored_pars = torch.zeros(iters, 3)\n",
|
||||
"losses = []\n",
|
||||
"for e in range(iters):\n",
|
||||
" stored_pars[e, :] = pars[0]\n",
|
||||
" loss = MinVol(pars)\n",
|
||||
" opt.zero_grad()\n",
|
||||
" loss.backward()\n",
|
||||
" losses.append(loss.item())\n",
|
||||
" opt.step() "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 8,
|
||||
"id": "6f6aeab5",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"[<matplotlib.lines.Line2D at 0x7fc4a7df89d0>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7fc4a7df8a00>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7fc4a7df8b80>]"
|
||||
]
|
||||
},
|
||||
"execution_count": 8,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(stored_pars.detach())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 9,
|
||||
"id": "271add47",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"[<matplotlib.lines.Line2D at 0x7fc4a81a7700>]"
|
||||
]
|
||||
},
|
||||
"execution_count": 9,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(losses)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "93a71a4d",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"# Running SABR"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 10,
|
||||
"id": "0bb6a328",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import torch\n",
|
||||
"import pandas as pd\n",
|
||||
"import copy\n",
|
||||
"from voltron.option_utils import GetTradingDays, GetTrainingData, Pricer, FindLastTradingDays\n",
|
||||
"from scipy.optimize import minimize\n",
|
||||
"\n",
|
||||
"def BlackVol(pars, K, f, T):\n",
|
||||
" alpha = torch.exp(pars[0][0]) ## v0\n",
|
||||
" rho = 2 * torch.sigmoid(pars[0][1]) - 1. ##rho\n",
|
||||
" v = torch.exp(pars[0][2]) ## \"sigma\" \n",
|
||||
" beta = 1.\n",
|
||||
" num = 1 + (alpha**2 * (1-beta)**2/(24 * (f*K)**(1-beta)) +\\\n",
|
||||
" 0.25 * rho*beta*v*alpha/((f*K)**(0.5*(1-beta))) +\\\n",
|
||||
" v**2*(2-3*rho**2)/24)*T\n",
|
||||
" num*= alpha\n",
|
||||
" \n",
|
||||
" denom = (f*K)**(0.5*(1-beta)) * (1 + (1-beta)**2/24 * torch.log(f/K)**2 +\\\n",
|
||||
" (1-beta)**4/1920 * torch.log(f/K)**4)\n",
|
||||
" \n",
|
||||
" z = v/alpha * (f*K)**(0.5*(1-beta)) * np.log(f/K)\n",
|
||||
" xi_z = torch.log((torch.sqrt(1 - 2 * rho * z + z**2) + z - rho)/(1-rho))\n",
|
||||
" \n",
|
||||
" return num/denom * z/xi_z\n",
|
||||
"\n",
|
||||
"def MinVol(pars, Ks, Fs, Ts, ivol):\n",
|
||||
" return torch.mean((ivol - BlackVol(pars, Ks, Fs, Ts)).pow(2))\n",
|
||||
"\n",
|
||||
"def Calibrate(Fs, Ks, Ts, ivol, iters=1000):\n",
|
||||
" pars = [torch.tensor([-1., -5., -3.], requires_grad=True)]\n",
|
||||
" opt = torch.optim.SGD(pars, lr=0.1)\n",
|
||||
" stored_pars = torch.zeros(iters, 3)\n",
|
||||
" losses = []\n",
|
||||
" for e in range(iters):\n",
|
||||
" stored_pars[e, :] = pars[0]\n",
|
||||
" loss = MinVol(pars, Ks, Fs, Ts, ivol)\n",
|
||||
" opt.zero_grad()\n",
|
||||
" loss.backward()\n",
|
||||
" losses.append(loss.item())\n",
|
||||
" opt.step() \n",
|
||||
" \n",
|
||||
" return pars[0].detach().numpy()\n",
|
||||
"\n",
|
||||
"def SABRSim(Np, Nt, S0, V0, sigma, rho, dt=1./252.):\n",
|
||||
" dW = np.random.randn(Nt+1, Np) * np.sqrt(dt)\n",
|
||||
" dZ = rho * dW + np.sqrt(1-rho**2) * np.random.randn(Nt+1, Np) * np.sqrt(dt)\n",
|
||||
" \n",
|
||||
" S = np.zeros((Nt+1, Np))\n",
|
||||
" S[0] = S0\n",
|
||||
" V = np.zeros((Nt+1, Np))\n",
|
||||
" V[0] = V0\n",
|
||||
" \n",
|
||||
" for t in range(Nt):\n",
|
||||
" S[t+1] = S[t] + V[t]*S[t]*dW[t]\n",
|
||||
" V[t+1] = V[t] + sigma*V[t]*dZ[t]\n",
|
||||
" \n",
|
||||
" return S[1:]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 11,
|
||||
"id": "67d9cf31",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"logger = []\n",
|
||||
"full_logger = []\n",
|
||||
"SPY = pd.read_csv(\"./data/SPY_prices.csv\")\n",
|
||||
"SPY['Date'] = pd.to_datetime(SPY['Date'])\n",
|
||||
"Np = 10000\n",
|
||||
"ntrain = 252\n",
|
||||
"year = 2012\n",
|
||||
"options = pd.read_csv(\"./data/SPY_\" + str(year) + \".csv\")\n",
|
||||
"options.expiration = pd.to_datetime(options.expiration)\n",
|
||||
"options.quotedate = pd.to_datetime(options.quotedate)\n",
|
||||
"qday = options.quotedate.unique()[0]\n",
|
||||
"quote_price = SPY[SPY['Date']==qday].Close.item()\n",
|
||||
"options = options[(options.quotedate == qday) & (options.type=='call')]\n",
|
||||
"edays = options.expiration.sort_values().unique()\n",
|
||||
"testdays = (edays - qday)/np.timedelta64(1, \"D\")\n",
|
||||
"edays = edays[(testdays > 100) & (testdays < 365)]\n",
|
||||
"lastdays = FindLastTradingDays(SPY, edays)\n",
|
||||
"ntests = np.array([GetTradingDays(SPY, qday, \n",
|
||||
" pd.Timestamp(ld)) for ld in lastdays])\n",
|
||||
"fulltest = ntests[-1]\n",
|
||||
"train_y = torch.FloatTensor(GetTrainingData(SPY, qday, ntrain).to_numpy())\n",
|
||||
"test_y = torch.FloatTensor(GetTrainingData(SPY, \n",
|
||||
" pd.Timestamp(lastdays[-1]),\n",
|
||||
" fulltest).to_numpy())\n",
|
||||
"full_x = torch.arange(ntrain+fulltest).type(torch.FloatTensor)\n",
|
||||
"full_x = full_x/252.\n",
|
||||
"train_x = full_x[:ntrain]\n",
|
||||
"test_x = full_x[ntrain:]\n",
|
||||
"\n",
|
||||
"## extract data for calibration ##\n",
|
||||
"ivol = torch.tensor(options.impliedvol.to_numpy())\n",
|
||||
"Fs = torch.tensor(options.underlying_last.to_numpy())\n",
|
||||
"Ks = torch.tensor(options.strike.to_numpy())\n",
|
||||
"starts = options.quotedate.dt.date.to_numpy()\n",
|
||||
"ends = options.expiration.dt.date.to_numpy()\n",
|
||||
"Ts = torch.tensor(([np.busday_count(qd, ed)/252. for qd, ed in zip(starts, ends)]))\n",
|
||||
"\n",
|
||||
"pars = Calibrate(Fs, Ks, Ts, ivol)\n",
|
||||
"v0 = np.exp(pars[0])\n",
|
||||
"1/(1 + np.exp(-pars[1]))\n",
|
||||
"rho = (2/(1 + np.exp(-pars[1])) - 1.)\n",
|
||||
"sigma = np.exp(pars[2])\n",
|
||||
"px_paths = SABRSim(Np, fulltest, quote_price, v0, sigma, rho)\n",
|
||||
"px_samples = torch.tensor(px_paths[ntests-1])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 12,
|
||||
"id": "abc47786",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"option_output = Pricer(torch.tensor(px_paths), options, edays, test_y[ntests-1],\n",
|
||||
" quote_price)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 13,
|
||||
"id": "9bb8972b",
|
||||
"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>Expiry</th>\n",
|
||||
" <th>Strike</th>\n",
|
||||
" <th>Bid</th>\n",
|
||||
" <th>Ask</th>\n",
|
||||
" <th>Voltron</th>\n",
|
||||
" <th>Return</th>\n",
|
||||
" <th>ExpClose</th>\n",
|
||||
" <th>QuoteClose</th>\n",
|
||||
" <th>Year</th>\n",
|
||||
" <th>Sample_Percentile</th>\n",
|
||||
" </tr>\n",
|
||||
" </thead>\n",
|
||||
" <tbody>\n",
|
||||
" <tr>\n",
|
||||
" <th>0</th>\n",
|
||||
" <td>2013-03-16</td>\n",
|
||||
" <td>20.0</td>\n",
|
||||
" <td>121.39</td>\n",
|
||||
" <td>121.61</td>\n",
|
||||
" <td>126.884545</td>\n",
|
||||
" <td>136.729996</td>\n",
|
||||
" <td>156.729996</td>\n",
|
||||
" <td>141.449997</td>\n",
|
||||
" <td>2013</td>\n",
|
||||
" <td>0.781553</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>1</th>\n",
|
||||
" <td>2013-03-16</td>\n",
|
||||
" <td>25.0</td>\n",
|
||||
" <td>116.39</td>\n",
|
||||
" <td>116.61</td>\n",
|
||||
" <td>121.884545</td>\n",
|
||||
" <td>131.729996</td>\n",
|
||||
" <td>156.729996</td>\n",
|
||||
" <td>141.449997</td>\n",
|
||||
" <td>2013</td>\n",
|
||||
" <td>0.781553</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>2</th>\n",
|
||||
" <td>2013-03-16</td>\n",
|
||||
" <td>30.0</td>\n",
|
||||
" <td>111.39</td>\n",
|
||||
" <td>111.61</td>\n",
|
||||
" <td>116.884545</td>\n",
|
||||
" <td>126.729996</td>\n",
|
||||
" <td>156.729996</td>\n",
|
||||
" <td>141.449997</td>\n",
|
||||
" <td>2013</td>\n",
|
||||
" <td>0.781553</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>3</th>\n",
|
||||
" <td>2013-03-16</td>\n",
|
||||
" <td>35.0</td>\n",
|
||||
" <td>106.39</td>\n",
|
||||
" <td>106.61</td>\n",
|
||||
" <td>111.884545</td>\n",
|
||||
" <td>121.729996</td>\n",
|
||||
" <td>156.729996</td>\n",
|
||||
" <td>141.449997</td>\n",
|
||||
" <td>2013</td>\n",
|
||||
" <td>0.781553</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>4</th>\n",
|
||||
" <td>2013-03-16</td>\n",
|
||||
" <td>40.0</td>\n",
|
||||
" <td>101.39</td>\n",
|
||||
" <td>101.61</td>\n",
|
||||
" <td>106.884545</td>\n",
|
||||
" <td>116.729996</td>\n",
|
||||
" <td>156.729996</td>\n",
|
||||
" <td>141.449997</td>\n",
|
||||
" <td>2013</td>\n",
|
||||
" <td>0.781553</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>...</th>\n",
|
||||
" <td>...</td>\n",
|
||||
" <td>...</td>\n",
|
||||
" <td>...</td>\n",
|
||||
" <td>...</td>\n",
|
||||
" <td>...</td>\n",
|
||||
" <td>...</td>\n",
|
||||
" <td>...</td>\n",
|
||||
" <td>...</td>\n",
|
||||
" <td>...</td>\n",
|
||||
" <td>...</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>651</th>\n",
|
||||
" <td>2013-09-30</td>\n",
|
||||
" <td>169.0</td>\n",
|
||||
" <td>0.33</td>\n",
|
||||
" <td>0.42</td>\n",
|
||||
" <td>0.000000</td>\n",
|
||||
" <td>0.690002</td>\n",
|
||||
" <td>169.690002</td>\n",
|
||||
" <td>141.449997</td>\n",
|
||||
" <td>2013</td>\n",
|
||||
" <td>1.000000</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>652</th>\n",
|
||||
" <td>2013-09-30</td>\n",
|
||||
" <td>170.0</td>\n",
|
||||
" <td>0.28</td>\n",
|
||||
" <td>0.37</td>\n",
|
||||
" <td>0.000000</td>\n",
|
||||
" <td>0.000000</td>\n",
|
||||
" <td>169.690002</td>\n",
|
||||
" <td>141.449997</td>\n",
|
||||
" <td>2013</td>\n",
|
||||
" <td>1.000000</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>653</th>\n",
|
||||
" <td>2013-09-30</td>\n",
|
||||
" <td>175.0</td>\n",
|
||||
" <td>0.14</td>\n",
|
||||
" <td>0.21</td>\n",
|
||||
" <td>0.000000</td>\n",
|
||||
" <td>0.000000</td>\n",
|
||||
" <td>169.690002</td>\n",
|
||||
" <td>141.449997</td>\n",
|
||||
" <td>2013</td>\n",
|
||||
" <td>1.000000</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>654</th>\n",
|
||||
" <td>2013-09-30</td>\n",
|
||||
" <td>180.0</td>\n",
|
||||
" <td>0.08</td>\n",
|
||||
" <td>0.12</td>\n",
|
||||
" <td>0.000000</td>\n",
|
||||
" <td>0.000000</td>\n",
|
||||
" <td>169.690002</td>\n",
|
||||
" <td>141.449997</td>\n",
|
||||
" <td>2013</td>\n",
|
||||
" <td>1.000000</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>655</th>\n",
|
||||
" <td>2013-09-30</td>\n",
|
||||
" <td>185.0</td>\n",
|
||||
" <td>0.04</td>\n",
|
||||
" <td>0.08</td>\n",
|
||||
" <td>0.000000</td>\n",
|
||||
" <td>0.000000</td>\n",
|
||||
" <td>169.690002</td>\n",
|
||||
" <td>141.449997</td>\n",
|
||||
" <td>2013</td>\n",
|
||||
" <td>1.000000</td>\n",
|
||||
" </tr>\n",
|
||||
" </tbody>\n",
|
||||
"</table>\n",
|
||||
"<p>656 rows × 10 columns</p>\n",
|
||||
"</div>"
|
||||
],
|
||||
"text/plain": [
|
||||
" Expiry Strike Bid Ask Voltron Return ExpClose \\\n",
|
||||
"0 2013-03-16 20.0 121.39 121.61 126.884545 136.729996 156.729996 \n",
|
||||
"1 2013-03-16 25.0 116.39 116.61 121.884545 131.729996 156.729996 \n",
|
||||
"2 2013-03-16 30.0 111.39 111.61 116.884545 126.729996 156.729996 \n",
|
||||
"3 2013-03-16 35.0 106.39 106.61 111.884545 121.729996 156.729996 \n",
|
||||
"4 2013-03-16 40.0 101.39 101.61 106.884545 116.729996 156.729996 \n",
|
||||
".. ... ... ... ... ... ... ... \n",
|
||||
"651 2013-09-30 169.0 0.33 0.42 0.000000 0.690002 169.690002 \n",
|
||||
"652 2013-09-30 170.0 0.28 0.37 0.000000 0.000000 169.690002 \n",
|
||||
"653 2013-09-30 175.0 0.14 0.21 0.000000 0.000000 169.690002 \n",
|
||||
"654 2013-09-30 180.0 0.08 0.12 0.000000 0.000000 169.690002 \n",
|
||||
"655 2013-09-30 185.0 0.04 0.08 0.000000 0.000000 169.690002 \n",
|
||||
"\n",
|
||||
" QuoteClose Year Sample_Percentile \n",
|
||||
"0 141.449997 2013 0.781553 \n",
|
||||
"1 141.449997 2013 0.781553 \n",
|
||||
"2 141.449997 2013 0.781553 \n",
|
||||
"3 141.449997 2013 0.781553 \n",
|
||||
"4 141.449997 2013 0.781553 \n",
|
||||
".. ... ... ... \n",
|
||||
"651 141.449997 2013 1.000000 \n",
|
||||
"652 141.449997 2013 1.000000 \n",
|
||||
"653 141.449997 2013 1.000000 \n",
|
||||
"654 141.449997 2013 1.000000 \n",
|
||||
"655 141.449997 2013 1.000000 \n",
|
||||
"\n",
|
||||
"[656 rows x 10 columns]"
|
||||
]
|
||||
},
|
||||
"execution_count": 13,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"option_output"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 14,
|
||||
"id": "51771e55",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"plt.plot(train_x, train_y)\n",
|
||||
"plt.plot(test_x, test_y)\n",
|
||||
"plt.plot(test_x, px_paths[:, :20], c='gray', alpha=0.5);"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 15,
|
||||
"id": "a9503b79",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"(206, 20)"
|
||||
]
|
||||
},
|
||||
"execution_count": 15,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"px_paths[:, :20].shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 16,
|
||||
"id": "13d5c888",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"(206, 10000)"
|
||||
]
|
||||
},
|
||||
"execution_count": 16,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"px_paths.shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "5867b14f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,438 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "3d8b2c23",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Warning no robinhood utils.\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import datetime as dt\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns\n",
|
||||
"import torch\n",
|
||||
"import gpytorch\n",
|
||||
"import os\n",
|
||||
"# import robin_stocks.robinhood as r\n",
|
||||
"import pickle5 as pickle\n",
|
||||
"import pandas as pd\n",
|
||||
"# sns.set_style(\"white\")\n",
|
||||
"# sns.set_palette(\"bright\")\n",
|
||||
"sns.set(font_scale=1.5)\n",
|
||||
"sns.set_style(\"white\")\n",
|
||||
"\n",
|
||||
"import sys\n",
|
||||
"sys.path.append(\"../\")\n",
|
||||
"from voltron.means import LogLinearMean\n",
|
||||
"from voltron.models import BMGP, VoltronGP, MaternGP, SMGP\n",
|
||||
"from gpytorch.kernels import ScaleKernel, RBFKernel, MaternKernel\n",
|
||||
"from voltron.train_utils import TrainBasicModel"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "ba337d0d",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Options Helpers"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "61523f0d",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def GetTrainingData(SPY, date, N):\n",
|
||||
" idx = SPY[SPY[\"Date\"] == date].index.item()\n",
|
||||
" return SPY['Close'].iloc[(idx-N):idx]\n",
|
||||
"\n",
|
||||
"def GetTrueValue(SPY, date, strike):\n",
|
||||
" close_px = SPY['Close'][SPY[\"Date\"] == date].item()\n",
|
||||
" return np.maximum(close_px-strike, 0)\n",
|
||||
"\n",
|
||||
"def GetTradingDays(SPY, start, stop):\n",
|
||||
" start_idx = SPY[SPY[\"Date\"] == start].index.item()\n",
|
||||
" stop_idx = SPY[SPY[\"Date\"] == stop].index.item()\n",
|
||||
" return stop_idx-start_idx\n",
|
||||
"\n",
|
||||
"def FindLastTradingDays(SPY, dates):\n",
|
||||
" last_days = []\n",
|
||||
" for date in dates:\n",
|
||||
" last_days.append(np.max(np.where(SPY.Date < date)[0]))\n",
|
||||
" \n",
|
||||
" return np.array(SPY.Date[last_days])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "85bc2e35",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Data"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"id": "7d014904",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"SPY = pd.read_csv(\"./data/SPY_prices.csv\")\n",
|
||||
"SPY['Date'] = pd.to_datetime(SPY['Date'])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "19d1fee1",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Data Setup"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 15,
|
||||
"id": "eb96fc66",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ntrain = 252\n",
|
||||
"options = pd.read_csv(\"./data/SPY_\" + str(2013) + \".csv\")\n",
|
||||
"options.expiration = pd.to_datetime(options.expiration)\n",
|
||||
"options.quotedate = pd.to_datetime(options.quotedate)\n",
|
||||
"qday = options.quotedate.unique()[0]\n",
|
||||
"options = options[(options.quotedate == qday) & (options.type=='call')]\n",
|
||||
"edays = options.expiration.sort_values().unique()\n",
|
||||
"testdays = (edays - qday)/np.timedelta64(1, \"D\")\n",
|
||||
"edays = edays[(testdays > 100) & (testdays < 500)]\n",
|
||||
"lastdays = FindLastTradingDays(SPY, edays)\n",
|
||||
"ntests = np.array([GetTradingDays(SPY, qday, pd.Timestamp(ld)) \n",
|
||||
" for ld in lastdays])\n",
|
||||
"fulltest = ntests[-1]\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"train_y = torch.FloatTensor(GetTrainingData(SPY, qday, ntrain).to_numpy())\n",
|
||||
"test_y = torch.FloatTensor(GetTrainingData(SPY, \n",
|
||||
" pd.Timestamp(lastdays[-1]),\n",
|
||||
" fulltest).to_numpy())\n",
|
||||
"full_x = torch.arange(ntrain+fulltest).type(torch.FloatTensor)\n",
|
||||
"full_x = full_x/252.\n",
|
||||
"dt = full_x[1] - full_x[0]\n",
|
||||
"train_x = full_x[:ntrain]\n",
|
||||
"test_x = full_x[ntrain:]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 62,
|
||||
"id": "1991167e",
|
||||
"metadata": {
|
||||
"tags": []
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"## train data gp ##\n",
|
||||
"dmod, dlh = TrainBasicModel(train_x, train_y, train_iters=1000, model_type=\"SM\", mean_func=\"constant\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 63,
|
||||
"id": "0c218e4f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"nvol = 100\n",
|
||||
"npx = 100\n",
|
||||
"px_samples = torch.zeros(npx*nvol, len(edays))\n",
|
||||
"px_paths = torch.zeros(npx*nvol, fulltest)\n",
|
||||
"dmod.eval();\n",
|
||||
"\n",
|
||||
"for vidx in range(nvol):\n",
|
||||
" px_pred = dlh(dmod(test_x)).sample(torch.Size((npx,))).exp()\n",
|
||||
" px_paths[vidx*npx:(vidx*npx + npx), :] = px_pred.detach()\n",
|
||||
" px_samples[vidx*npx:(vidx*npx+npx), :] = px_pred[:, ntests-1].detach()\n",
|
||||
" "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 64,
|
||||
"id": "6b410709",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAjsAAAGACAYAAABLM6NwAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAA9hAAAPYQGoP6dpAADm40lEQVR4nOydd5hU1f3/X3f6zM72RltA2gpSFLBgQcXexdh+9hg11iQajSVRE03UJN/EWJNoLFhjJfYOglQB6b0v2/tO7+f3x+fOLOsuCEgRPK/n4WHnlnPPvTO75z2faiilFBqNRqPRaDT7KJY9PQGNRqPRaDSaXYkWOxqNRqPRaPZptNjRaDQajUazT6PFjkaj0Wg0mn0aLXY0Go1Go9Hs02ixo9FoNBqNZp9Gix2NRqPRaDT7NFrsaDQajUaj2afRYkej0Wg0Gs0+jRY7Gs0+zB133EF5eTl33HHHnp7KbuGxxx6jvLycSy+9dE9PZbuZPXs25eXllJeX/yivr9HsSmx7egIazY+Vxx57jMcff7zTdofDQX5+PkOGDOHMM8/klFNOwTCMPTDDPcfChQt5/fXX+eabb6itrSUej1NYWEhhYSHl5eUcfPDBjBkzhu7du+/pqf7g8fl8TJgwAYDLL7+cnJycPTwjjWb3o8WORvMDoKioKPOz3++nrq6Ouro6Jk+ezMSJE3niiSdwOBzbPW5xcTH77bcfxcXFO3O6uwylFA888AAvvPBCZpthGOTk5NDc3ExtbS1Lly7l7bffZvz48Tz00EMdzs/Pz2e//fbTImgzfD5fRlSPHz9+i2LH7Xaz33777c6paTS7DS12NJofANOnT8/8nEqlWLt2LQ8++CDTp09n6tSpPPzww9x+++3bPe6vf/1rfv3rX+/Mqe5Snn/++YzQOe6447j66qs54IADMkJv06ZNzJ49m48//hiLpbMX/pJLLuGSSy7ZrXPeVxg+fDgff/zxnp6GRrNL0GJHo/mBYbFYGDhwIP/85z8544wz2LhxI6+99hq//vWvsdn23V9ZpRTPPfccAEcddRRPPvlkp2PKysooKyvj3HPPJRKJ7O4pajSavZR99y+nRrOX43Q6Ofnkk/n3v/9NMBhk3bp1DBo0iMrKSo477jgAvvjiC1KpFE8//TTTp0+nvr6ekpISJk2aBEiA8sSJE7t0+aSpqanhxRdfZPr06VRWVhKPxykpKWHgwIGcdNJJnHLKKTidzk7nLVu2jBdffJE5c+bQ0NCAxWKhrKyMY489lssvv5yCgoLtut+Wlhbq6uoAGDdu3Hce73K5Om1Lx0EdcsghvPjiix32fftZvP3227z22musWbMGq9XKkCFDuOGGGzj44IMBSCQSvPrqq0ycOJENGzZgGAYjR47kV7/6FQcccECna7/99tvceeed9OzZM/P8v82337tevXp9532CWPvmz5/P5MmT+frrr6mtraW5uZmsrCwGDhzIaaedxrnnnovdbu9w3qWXXsrXX3+deZ2+dprNn9Ps2bO57LLLAFi5cmWX82hoaODZZ59l6tSpVFdXo5SiZ8+eHH300Vx55ZUd3LFbumeXy8W//vUvJk2aRENDA9nZ2Rx66KHceOON9O/fv8vr1tbW8uyzzzJ9+nSqqqpIJBLk5eVRUlLC6NGjOf300xk+fPg2PUvNjxMtdjSaHzClpaWZnwOBQKf98+fP55577iEUCuF2uzstdt/F//73P+655x6i0SgAdrsdl8vFpk2b2LRpE5MmTaK8vJzBgwd3OO/RRx/lySefRCkFSLxHPB5n5cqVrFy5krfeeounnnqKIUOGbO8tA2REz64iLXxsNhtOp5O2tjZmzpzJnDlzePzxxzniiCO47rrrmDZtGna7HbvdTjAYZOrUqcyZM4eXXnqJoUOH7tI5bk51dTUXXXRR5rXNZsPlctHa2sqcOXOYM2cO77//Ps8880wHEZibm0t+fj4tLS2AxDRZrdYO+7eVr7/+mhtuuAGfzwfIe24YBmvWrGHNmjW8+eabPPnkk4wePXqLY6xZs4a77rqLpqYm3G43AE1NTXz44YdMnTqVl19+mf3337/DOStWrOCyyy6jra0NAKvVitfrpbGxkYaGBpYuXYrP59NiR7NVtNjRaH7AVFVVZX7uamG65557GDhwIHfffTfDhg0DYP369ds09pQpU7jjjjtQSjFy5Eh+/etfM3LkSCwWC4FAgBUrVvDOO+90ElDPP/88TzzxBFlZWVx77bWMHz+e4uJikskky5cv569//SuzZs3iuuuu48MPPyQrK2ub5lNQUECvXr2orKzkxRdf5JBDDuGII47YpnO3hy+++IJYLMZ9993HWWedhcvlYt26ddx6660sXbqU+++/n2OPPZYlS5bwj3/8g+OPPx6bzcbSpUu5+eabqaio4E9/+hOvvvrqTp/blrDZbBx33HGcfvrpjBo1iuLiYiwWC8FgkE8++YSHH36YuXPn8vDDD3PnnXdmznv88cc7WFbefPPNbbYmbU5NTU1G6AwYMID77ruPUaNGATB37lx+97vfsX79em644QbefffdDiJ9c37zm9/Qv39//v3vfzNs2DASiQRff/01v/nNb2hoaOD+++/n5Zdf7nDOQw89RFtbGwcccAD33HMPI0aMwDAMYrEY1dXVTJo0iVQqtd33pPlxoevsaDQ/UAKBAO+99x4AeXl5XWbK5Ofn89xzz2WEDrBNGTWJRIL77rsPpRSjRo1iwoQJjB49OhP06/V6GT16NPfffz8DBgzInNfc3Mw//vEPDMPgiSee4JprrslkelmtVoYOHcozzzzDAQccQG1tLW+88cZ23fOvfvUrAILBIFdeeSXjxo3jN7/5DRMmTOCbb74hFott13hd4fP5uP/++7ngggsyVpB+/frxyCOPYBgGVVVVvPTSSzzxxBOccsop2O12DMNg6NCh3HfffQCZlPjdRbdu3XjyySc59dRTKS0tzbxPWVlZnHPOOZn4ptdffz1jpduZ/Otf/8Ln85Gbm8vzzz+fEToAo0eP5vnnn8fr9dLa2sq///3vLY5TWFjY4fNqs9k4/PDDM8917ty5nZ7r/PnzAbj77rs58MADM2UYHA4Hffv25corr+Sqq67aqfer2ffQYkej+YHh8/mYOXMml112GfX19YDEXnSVfXTxxRdvs+Vkc2bPnk1lZSUAd9555zantb/33nuEw2GGDh3KmDFjujzGZrNx+umnAzBt2rTtmtcZZ5zBww8/TLdu3QCxbL3zzjs88MAD/L//9/84+OCDufnmm1mxYsV2jbs5PXr04Iwzzui0vaysjN69ewOygHfljjnkkEMyz2pLcS17gmHDhlFYWEgoFGL58uU7dWylVCZL68ILL+yyjEG3bt248MILAfjggw+2ONaVV17ZZazV2LFjMxbEbz/X7OxsQOKFNJodRbuxNJofAFurWnvmmWdy3XXXdblv5MiRO3S99Lfl4uLiDlah72LevHkArF69eqsupnSmVHV19XbP7dRTT+WEE05gxowZzJw5k0WLFrFixQqCwSCRSIQPP/yQTz/9lHvvvZfzzz9/u8cfOnToFos0FhYWsnHjxi0+E6vVSn5+PnV1dZkYkt1FLBbjrbfe4rPPPmPVqlW0tbV1aena2RanyspKWltbAbYocAGOOOII/vOf/9Da2sqmTZsoKyvrdMyW4mpsNhsFBQVdPtdjjz2W119/ndtvv51vvvmGcePGMWzYsEzMj0azLWixo9H8ANg8iyVdQXnw4MGcccYZHHbYYVs8r7CwcIeul/6W3KNHj+06L21pikQi25T6vaPp4Xa7naOPPpqjjz4akGykFStWMHHiRF555RUSiQS///3vGT58eKeA1u9ia5awdGr/thyTSCS267rfh6amJq644gpWrVqV2eZ0OjsEHDc3N5NKpQiHwzv92mm2FIvz7X3Nzc1dip0dea633XYbGzduZPbs2Tz33HM899xzWK1W9t9/f4455hguuOCCrc5LowEtdjSaHwSbFxXcHrpybW0P29uGIh0IeuGFF/KHP/zhe117e7BYLAwZMoQhQ4aw//77c9ddd5FMJnnrrbf47W9/u9vmsad44IEHWLVqFXl5efzmN79h7NixndxJRx99NLW1tZkMuV3Btn5edmZ7k5ycHF544QXmzp3L5MmT+eabb1iyZAlLly5l6dKlPPPMM/zpT3/KuE41mq7QMTsazY+Q9EKZjtvZVtIWqM0tDLubs88+OxP3sa2ZZ7uDtIVlawHCXZUP+C7i8TifffYZINl3P/nJTzoJnWQymUkv39lsbj3cmots83IB+fn5O30eo0eP5rbbbuPVV19l7ty5PPnkkwwaNIhIJMJdd91FY2PjTr+mZt9Bix2N5kdIOtansbGRxYsXb/d5Cxcu7JAWvzuxWq2ZIoc70i9sV5EuDdDU1LTFrLGFCxdu97jNzc0ZAfXtekdp5s2bt0WRtbn1b0esPr169SIvLw+AmTNnbvG4GTNmAJI52JULa2fidDo57rjjMj2/otFoJp5Mo+kKLXY0mh8hhx56aGZBevDBB7c5pTtdlyaZTHLfffeRTCa3eGwqlcoUoNsWYrEYs2bN+s7jJk2alAli3dGihbuCdOyQUipjidmcSCTC888/v93jer3ejFuoqyy0RCLBww8/vNXz0/j9/u2+vmEYnHLKKQC89tprXWZF1dXV8dprrwHsVHdSIpHYag2dzTO7Ni+WqNF8Gy12NJofIVarlbvvvhvDMJg3bx5XXHEFc+fOzSwsgUCA2bNnc+utt7JmzZrMecXFxZnGol9++SU//elPmTdvXkb0KKVYu3Ytzz33HKeffjqTJ0/e5jnF43Euv/xyxo8fz3PPPceKFSsy46ZSKaqqqnj88ce55ZZbAFnEzzvvvJ3yPHYG3bp1y9SfefDBB5kxY0Zm/kuWLOGKK66gubl5u8fNysrKWNQeeughZs6cmXmfVq1axTXXXMOSJUvweDxdnp+Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 600x360 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"colors = [\"#1b4079\",\"#3db2ff\",\"#ffedda\",\"#ffb830\",\"#ff2442\",\"#61210f\",\"#32373b\"]\n",
|
||||
"colors = [\"#0c4767\",\"#8A95A5\", \"#F4442E\", \"#566e3d\",\"#b9a44c\",\"#fa7921\",\"#fe9920\"]\n",
|
||||
"colors = [\"#01295F\",\"#437F97\", \"#F4442E\", \"#566e3d\",\"#b9a44c\",\"#fa7921\",\"#fe9920\"]\n",
|
||||
"fs = 16\n",
|
||||
"\n",
|
||||
"fig, ax = plt.subplots(1,1, figsize=(6, 3.6), dpi=100)\n",
|
||||
"\n",
|
||||
"dt = 1./252\n",
|
||||
"\n",
|
||||
"# ax.plot(train_x, vol/dt**0.5, c=colors[0], label='GPCV Vol.')\n",
|
||||
"# vmod.eval();\n",
|
||||
"# samples = vmod(test_x).sample(torch.Size((20,))).exp().detach()\n",
|
||||
"# ax.plot(test_x, samples[0]/dt**0.5, c=colors[-2], alpha=0.5,\n",
|
||||
"# lw=0.7, label=\"Simulations\")\n",
|
||||
"# ax.plot(test_x, samples[1:8].T/dt**0.5, c=colors[-2], alpha=0.5,\n",
|
||||
"# lw=0.7)\n",
|
||||
"\n",
|
||||
"# ax.set_xlabel(\"Days\")\n",
|
||||
"# ax.set_ylabel(\"Volatility\")\n",
|
||||
"# ax.set_title(\"Volatility Simulations\")\n",
|
||||
"# ax.legend(fontsize=fs-2, frameon=False)\n",
|
||||
"\n",
|
||||
"ax.plot(train_x, train_y, c=colors[0], label=\"Train\")\n",
|
||||
"\n",
|
||||
"ax.plot(test_x, px_paths[1:20, :].T, c=colors[-1], alpha=0.5,\n",
|
||||
" lw=0.5)\n",
|
||||
"# ax[1].plot(test_x, px_paths[0, :].T, c=colors[-1], alpha=0.5,\n",
|
||||
"# lw=0.2, label=\"Simulations\")\n",
|
||||
"\n",
|
||||
"ax.plot(test_x, test_y, c=colors[1], lw=2., label=\"Test\")\n",
|
||||
"ax.plot(test_x, px_paths[0, :].T, c=colors[-2], alpha=0.5,\n",
|
||||
" lw=0.5, label=\"Simulations\")\n",
|
||||
"sns.rugplot(x=test_x[ntests-1], ax=ax, color=colors[2], height=0.1, label=\"Expirations\")\n",
|
||||
"ax.legend(fontsize=fs-2, frameon=False)\n",
|
||||
"ax.set_xlabel(\"Days\")\n",
|
||||
"ax.set_ylabel(\"Price\")\n",
|
||||
"ax.set_title(\"Price Simulations\")\n",
|
||||
"# plt.savefig(\"./option_diffusions.pdf\", bbox_inches=\"tight\")\n",
|
||||
"sns.despine()\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 65,
|
||||
"id": "2bac7474",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"logger = []\n",
|
||||
"days = np.datetime_as_string(edays, 'D')\n",
|
||||
"for day in range(px_samples.shape[1]):\n",
|
||||
" for smpl in range(px_samples.shape[0]):\n",
|
||||
" logger.append([px_samples[smpl, day].item(), days[day][5:]])\n",
|
||||
" \n",
|
||||
"df = pd.DataFrame(logger)\n",
|
||||
"df.columns = ['Price', 'Date']"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 66,
|
||||
"id": "876976aa",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAgsAAAF+CAYAAAAMWFkhAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAACvdElEQVR4nOy9d5xU1f3//zrnTtmZne0VdikLy1JEEJEiIEiLJZagxmjsJcZfRInto4n54ueT+ElMjGKE8MEeE2sSBaNYEQMiWAKoSO/sAtt3dnZ26r3n/P64O8MuOzs7fe5cz/Px8CF758697zO3nPd5V8I55xAIBAKBQCDoA5puAQQCgUAgEGgboSwIBAKBQCAIi1AWBAKBQCAQhEUoCwKBQCAQCMIilAWBQCAQCARhEcqCQCAQCASCsAhlQSCIgLq6OowcORJLly4Nu01L3H///Rg5cmRKzzlnzhxcc801KT1nvKTjOmrpftL6fSzQBoZ0CyAQ9MXnn3+Oa6+9tsc2q9WKqqoqXHzxxbj66qshSVKapIuPuro6rFy5EvPmzcPo0aPTLQ7mzJmDo0ePBv82Go0oLS3FmWeeiYULF2LAgAFplC5yuitHhBBYLBYUFhZi1KhRmDNnDr7//e8jKysrYed744034HA4cP311yfsmMlAa/ebIPMQyoJA81xwwQWYOXMmOOdobGzEypUr8dvf/hb79u3Db37zm7TJVVFRgW+++SYmheXo0aNYtmwZKioqNPPyLi8vx1133QUA6OzsxBdffIHXX38d69atw7/+9S8UFhb2e4z33nsv2WL2y+jRo3HDDTcAADweD44dO4ZPP/0Uv/zlL7FixQosXboUo0aNCu4fz3VcuXIljh49GrWyEM85YyHc/ZZqWQSZiVAWBJpnzJgxuPjii4N///jHP8Z5552Hf/zjH1i0aBGKi4tDfs/pdMJmsyVNLkIIzGZz0o6fanJycnr9zr/5zW/w4osv4o033sDNN98c8nt+vx+MMZjNZphMplSJ2ydlZWU9xgEAd955J959913ce++9uPnmm7F69Wrk5eUBSO11DNyTWrp3tCSLQLuImAVBxmGz2TBhwgRwzlFbWwvghK98x44duOmmmzBx4kRcdNFFwe8cOnQI9957L2bMmIGxY8dizpw5+P3vfw+Xy9Xr+P/5z39wxRVXYNy4cZg2bRp+/etfh9wvnK/3/fffxzXXXIMzzjgD48ePxznnnIOHHnoIPp8Pb7zxRtC98otf/AIjR47EyJEje/j6Oed4+eWXcckll2D8+PGYMGECrrnmGnz22We9zuX1evH73/8eM2bMwLhx43DZZZdhw4YN0f+wIZgxYwYA4MiRIwCApUuXYuTIkdi7dy9+97vfYebMmRg3bhy++uorAH3HLOzYsQN33HEHpk2bhrFjx2LWrFm46667gscNsHHjRtx4440444wzcOqpp+LCCy/EK6+8kpCxnHfeebjpppvQ1NSEl156Kbi9r+u4atUqXHbZZTjjjDNw2mmnYe7cubj77rvR2toaHOsXX3yBo0ePBq/hyJEj8fnnnwMArrnmGsyZMwe1tbW44447MHnyZEycODHsOQO8/fbbuPDCC3Hqqafi7LPPxtKlSyHLco99Asc/mZOP3d/91pcssizjqaeewvnnn49TTz0VU6ZMwW233Ybdu3f3eb6PP/4Yl156KU499VTMmDEDv//973vJvXfvXtxxxx0466yzMHbsWEyfPh3XXHMN/v3vf4f8LQTaQFgWBBkH5xyHDx8GABQUFAS3Hzt2DNdddx3OPfdcfO973wtO8N9++y2uu+465Obm4kc/+hHKysqwa9cu/O1vf8PWrVvxt7/9DUajEQDw9ddf44YbbkB2djZ+8pOfICcnB++88w7uu+++iOVbsmQJVqxYgerqalx//fUoKSnBkSNH8MEHH+COO+7ApEmTcOutt2LFihX40Y9+FJxAultI7r33XqxevRrnnHMOLrnkEvh8Prz11lu48cYbsXTpUsydOze471133YU1a9Zg9uzZOOuss3DkyBHcfvvtqKysjP1H7iLU7wwA99xzD7KysnDjjTcCAEpKSvo8xscff4zbb78dVqsVl112GYYMGYKmpiZs2LABe/bsweDBgwEAr732Gh588EGcdtppuPXWW2GxWLBx40b893//N44cORLVNeiLH/7wh1ixYgXWrVuHn/3sZ33u9+abb+K+++7DGWecgTvuuANZWVk4duwY1q9fj5aWFhQWFuKXv/wlHn30UbS1teEXv/hF8LvDhw8P/ruzsxNXX301Tj/9dPz85z8PKhrh+Pjjj/HCCy/gqquuQnFxMdauXYtly5bh2LFj+N3vfhf1mCO530Jxzz334N1338X06dNx5ZVXorm5GS+99BKuuOIKvPTSSxgzZkyP/detW4eXX34ZV1xxBS699FJ89NFHeO6555CXl4dbb70VANDW1obrrrsOAHDFFVdg4MCBaGtrw7fffouvv/4aZ599dtTjE6QILhBolM8++4zX1NTwpUuX8paWFt7S0sJ37tzJH3jgAV5TU8Mvv/zy4L6zZ8/mNTU1/O9//3uv41x44YX8nHPO4R0dHT22f/DBB7ympoa//vrrwW0/+tGP+CmnnMIPHDgQ3Ob1evmll17Ka2pq+BNPPBHcXltb22vb119/zWtqavg111zDPR5Pj/MxxjhjrMfYup/7ZLleffXVHtv9fj9fsGABnz17dvA4n3zyCa+pqeH33Xdfj30//PBDXlNTw2tqanodPxSzZ8/m5557bvB3PnLkCP/nP//JJ06cyMeMGcN3797NOef8iSee4DU1Nfzqq6/mfr8/5HGuvvrq4N8ul4tPmTKFT506ldfX1/faX1EUzjnnDQ0NfOzYsfyuu+7qtc9vfvMbPmrUKH748OF+x1FTU8NvueWWsPtMmDCBT548Ofh3qOt422238QkTJoQcY3euvvpqPnv27D4/q6mp4Y899livz0KdM7Bt1KhR/Ntvvw1uZ4zxn/3sZ7ympoZv3bq133OHOna4+y3U/hs2bOA1NTV80aJFwXuNc8537tzJR48eza+88spe3x8/fjyvra3tIff3v/99Pn369OC2NWvW8JqaGr569eqQv5lAuwg3hEDzLF26FGeeeSbOPPNMXHzxxXj99dcxZ84c/PnPf+6xX35+Pi655JIe23bv3o3du3fjggsugM/nQ2tra/C/iRMnwmq14tNPPwUAtLS0YOvWrZgzZw6qqqqCxzCZTBEHsP3rX/8CANx99929/MCEEBBCIjpGdnY25s2b10Neh8MRzFo4dOgQAGDNmjUAgJtuuqnHMebNm9djDJFw4MCB4O88b948/PKXv0RBQQGWL1+OmpqaHvted911MBj6N0xu2LABbW1tuOGGG1BWVtbrc0rVV9D7778Pn8+Hyy67rMeYW1tbMWfOHDDGsGnTpqjG0xc2mw1OpzPsPjk5OfB4PPj3v/8NHmdj3pOvTX9MmzYNp5xySvBvQkgwXuTDDz+MS5ZICZzn1ltv7XHPjho1CmeffTY2b97cy0oyd+7cHtYsQgimTJmCpqYmdHZ2AlB/VwD45JNP+r0GAm0h3BACzfOjH/0I5557bjAVbujQocjPz++136BBg3pFdO/fvx+AqnD05R9ubm4GgGD8w7Bhw3rtU11dHZGshw8fBiGkR7R9tOzfvx+dnZ2YNm1an/u0tLSgqqoKtbW1oJRi6NChvfYZPnw4Dh48GPF5Kyoq8NBDDwE4kTo5ZMiQkPuGOl8oAkrNySbrkwlcp3BKWeA6xUskga8//elP8eWXX+K2225Dfn4+Jk+ejJkzZ+K8886LKmi2sLAQubm5UcnX3Y0RIHD/Be7RZFNXVwdKaUhZRowYgY8++gh1dXU9MmQGDRrUa9/Ac2q325GdnY3JkyfjBz/4Ad544w289dZbGDt2LKZNm4bzzz8/4mdMkB6EsiDQPEOGDAk7cQawWCx9fnbjjTfirLPOCvlZ4GUeWEGGWv1HurrknEdkPejvGIWFhXj00Uf73GfEiBERHScarFZrRL8zgIhrFYT7TUPt9/vf/x6lpaUh9wk1GUVLXV0dOjs7MWHChLD7DR06FO+88w42bdqETZs24YsvvsCvfvUrPPHEE3jppZeCcRb9Ee6e7It47x9FUeL6PhD9vQMgbOpl9+P9/ve/x0033YR169Zh8+bNeP7557FixQr88pe/xNVXXx2TvILkI5QFga4JrIwppf1OhIEJILDK7U6obaGoqqrCJ598gt27d2PcuHF97hduQhgyZAgOHTqE8ePHIzs7O+z5Bg0aBMYYDh061EuBOHDgQEQyJ5OAlWbHjh2YPn16n/sFLBUFBQURKyyx8I9//AMAMGvWrH73NZlMmDVrVnDfdevW4ZZbbsHzzz+PBx98MGky7tu3r89t3RWm/Px8bN++vde+oawP0SoggwcPxoYNG7B///5eVrLAsxBPAG1NTQ1qamrwk5/8BA6HAz/84Q/x6KOP4qqrropbWRIkBxGzINA1Y8aMQU1NDV599dWQL1FZlmG32wEARUVFOO2007B27doe5nufz4e//OUvEZ3vwgsvBAA89thj8Pl8vT4PrLCsVisAoL29vdc+P/jBD8AYw2OPPRbyHN3N8YGsiGeffbbHPmvWrInKBZEspk+fjoKCAjz//PNobGzs9Xng9zjvvPNgMpmwdOlSeDyeXvt1dHSE/D2j4d1338Wzzz6L0tJSXHXVVWH3DZW1EHCldL9m2dnZaG9vjzuuoTsbN27soQRwzvHMM88AUGNRAgwdOhSdnZ345ptvgtsYYyHv1XD3WygC53nqqad6jG3Pnj1Yu3YtJk6cGFGRrpOx2+1gjPXYlpubi8rKSrjdbni93qiPKUgNwrIg0DWEEPzhD3/Addddh4suugiXXnopqqur4fF4cPjwYXz44Ye46667goGR999/P6655hpceeWVuOqqq4Kpk5GadseNG4eLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 576x360 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(figsize=(8,5))\n",
|
||||
"violins = sns.violinplot(x='Date', y='Price', data=df, color=colors[-1])\n",
|
||||
"for violin in violins.collections[::2]:\n",
|
||||
" violin.set_alpha(0.5)\n",
|
||||
" violin.set_edgecolor(colors[-2])\n",
|
||||
"# violin.set_facecolor(colors[-2])\n",
|
||||
"\n",
|
||||
"# parts = ax.violinplot(dataset=px_samples.T,positions=range(px_samples.shape[1]))\n",
|
||||
"# for pc in parts['bodies']:\n",
|
||||
"# pc.set_facecolor(colors[2])\n",
|
||||
"# pc.set_edgecolor(colors[3])\n",
|
||||
"# pc.set_alpha(1)\n",
|
||||
"plt.scatter(np.arange(px_samples.shape[1]), test_y[ntests-1], color=colors[1], zorder=4, s=80,\n",
|
||||
" label=\"Observed Price\")\n",
|
||||
"plt.xlabel(\"Expirations\")\n",
|
||||
"plt.xticks(rotation=45)\n",
|
||||
"sns.despine()\n",
|
||||
"plt.legend(loc=\"upper left\",frameon=False, fontsize=fs-4)\n",
|
||||
"plt.title(\"Predicted Price Distributions\")\n",
|
||||
"# plt.savefig(\"./distribution_plot.pdf\", bbox_inches='tight')\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 67,
|
||||
"id": "ef9e8f53",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def Pricer(mc_pxs, options, edays, true_pxs):\n",
|
||||
" logger = []\n",
|
||||
" for eday_idx, eday in enumerate(edays):\n",
|
||||
" eday = pd.Timestamp(eday)\n",
|
||||
" opts = options[options.expiration==pd.Timestamp(eday)]\n",
|
||||
" for idx, row in opts.iterrows():\n",
|
||||
" K = row.strike\n",
|
||||
" bid = row.bid\n",
|
||||
" ask = row.ask\n",
|
||||
" valuation = np.mean(np.maximum(mc_pxs[:, eday_idx].numpy() - K, 0))\n",
|
||||
" rtn = np.maximum(true_pxs[eday_idx] - K, 0)\n",
|
||||
" logger.append([eday, K, bid, ask, valuation, rtn.item()])\n",
|
||||
" \n",
|
||||
" df = pd.DataFrame(logger)\n",
|
||||
" df.columns = ['Expiry', \"Strike\", \"Bid\", \"Ask\", \"Voltron\", \"Return\"]\n",
|
||||
" return df"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 68,
|
||||
"id": "016d7372",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"option_output = Pricer(px_samples, options, edays, test_y[ntests-1])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 69,
|
||||
"id": "d4f6fcb4",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"2014-06-30T00:00:00.000000000\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"idx = 3\n",
|
||||
"print(edays[idx])\n",
|
||||
"dat = option_output[option_output.Expiry == option_output.Expiry.unique()[idx]]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 70,
|
||||
"id": "f9f405ee",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAZIAAAEWCAYAAABMoxE0AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAABnfUlEQVR4nO3dd1RUx9vA8e/Sq9gAKzYECwjYEFTsvWDvYsFQLLHXRGOLUbFXEGONJhoVLFGxKwKiYhSxROwiCioCAlKEff/Yn7whoK7uLiDO55w98baZuTewD3OnSaRSqRRBEARB+EJqBV0AQRAE4esmAokgCIKgEBFIBEEQBIWIQCIIgiAoRAQSQRAEQSEaBV2A/JSamkpERATGxsaoq6sXdHEEQRC+CpmZmbx48QIrKyt0dHRyHf+mAklERAQDBw4s6GIIgiB8lXbs2EH9+vVz7f+mAomxsTEgexhlypQp4NIIgiB8HZ4/f87AgQOzv0P/65sKJO9fZ5UpU4YKFSoUcGkEQRC+Lh9qEhCN7YIgCIJCRCARBEEQFFKggeTmzZuMHDmSJk2aYGtrS8eOHdmwYQPp6ek5zgsKCqJPnz7UqVMHBwcHZs2aRWJiYgGVWhAEQfi3AmsjuXfvHv369aNKlSrMmDGDEiVKcOHCBZYvX87du3dZvHgxAKGhobi5udGqVSvGjRtHbGwsS5Ys4c6dO+zcuRM1NVGpEgRBKEgFFkgOHz5MWloaq1evxszMDAAHBweio6P566+/+Pnnn9HU1MTLy4vq1auzYsWK7KBhbGzM8OHDOXr0KB07diyoWxCEfJGYmEhsbCwZGRkFXRShiNLU1MTExIRixYp90fUFFkg0NGRZGxgY5NhvaGiIhoYG6urqxMTEcP36daZNm5aj5tG4cWNMTU0JCAgQgUQo0hITE4mJiaF8+fLo6uoikUgKukhCESOVSnn79i1Pnz4F+KJgUmDvhZydnSlevDizZ8/myZMnJCUlceLECfz8/Bg2bBhqamrcuXMHgOrVq+e63sLCgsjIyHwrr+t+V1pva8391/fzLU9BiI2NpXz58ujp6YkgIqiERCJBT0+P8uXLExsb+0VpFFggKVeuHLt27eLevXu0bt2aevXqMWrUKFxcXBg3bhwA8fHxABgZGeW63sjIKPt4fnCu4cyl6EvUWV+HDWEbEOuBCfkhIyMDXV3dgi6G8A3Q1dX94tenBfZq6+nTp3h4eGBsbMzatWsxNDTk0qVL+Pj4IJFIsoMJ8MG/xPLzL7Sull257nmdYfuH4X7IHf/b/vza9VfKGpbNtzII3yZRExHygyI/ZwUWSJYuXUpycjL+/v7Zk4DZ29sDsHbtWnr16kXx4sUB8qx5JCQk5FlTUSUzIzOODz7O2otrmXJiClbrrVjfaT19avfJ13IIgiAUJgX2auvmzZuYm5vnmknSysqKrKws7t+/n902kldbyJ07d/JsO1E1NYkaY+zH8Lf731QrUY2+e/rSf29/4t7G5XtZBKGomjZtGu7u7gVdDEFOBRZITExMiIyM5O3btzn2//333wCYmppSpkwZrKysOHjwIFlZWdnnhISEEBMTQ9u2bfO1zP9Wo3QNgl2DmddiHntu7sFqnRVH7x4tsPIIQmHh4eHB0KFD8zx27949LC0tCQoK+qw0Bw8ezNy5c5VQOkEVCiyQuLi48OLFC1xdXQkICCA4OJjly5fz66+/4ujoiKWlJQCTJk3i9u3bTJgwgZCQEPz9/Zk8eTI2Nja0b9++oIoPgIaaBj86/UjoiFBK6pakw44OeB7yJCk9qUDLJQgFqVevXly4cIGoqKhcx/bs2UP58uVxcHBQSd5irE3BKLBA0rp1azZv3oyWlhZz5sxh5MiRnDhxAk9PT9auXZt9noODA97e3jx9+hQ3NzcWLlxI8+bN8fX1LTSLU9UtW5fLbpeZ5DAJnzAfbL1tCXr8eX9xCUJR0bx5c0qXLs2+ffty7M/IyGD//v306NGDsLAwevfujbW1NY6OjixYsCDX1EjvTZs2jYsXL7Jjxw4sLS2xtLQkKiqK0NBQLC0tOXv2LL169cLKyorz58+Tnp7Ozz//jKOjI9bW1vTp04fLly9np/f+upCQEHr37o2NjQ09evTgxo0bKn0uRZr0G/LkyROphYWF9MmTJyrL4+zDs9IqK6pI1eaoSacenypNzUhVWV5C0Xfz5s2CLsIX8fLykjZv3lyamZmZvS8gIEBao0YNaVRUlNTGxkY6c+ZM6d27d6WnTp2SOjo6Sn/55Zfsc6dOnSp1c3OTSqVSaWJiorRv377SadOmSWNjY6WxsbHSd+/eSS9cuCC1sLCQdu7cWRoYGCh9/Pix9NWrV9J58+ZJGzduLD19+rT07t270h9++EFqa2srjYmJkUql0uzrevbsKQ0JCZHevXtXOnz4cGn79u2lWVlZ+fugCpkP/bx96rvzm1qPJD84VXLimsc1Jh6byKKgRRyOPMz27tuxKWNT0EUTiojj16I4du1JvubZ1qYibWzkX8OnV69e+Pr6EhwcTJMmTQDZa63GjRuze/dujI2NmT17NmpqalSrVo2JEycya9Ysxo4dm2vcjKGhIZqamujq6ua5sNLo0aOz80hJSeGPP/5g/vz5NG/eHIA5c+Zw4cIFduzYwfjx47OvGzt2LI0aNQJg5MiRDBgwgJiYGLHo3RcQMx6qgKG2IRu6bOBQ/0O8SHlBA98G/BL4C++y3hV00QQhX1SuXJkGDRqwd+9eAGJiYjh//jy9e/fm3r172Nra5pj2qF69emRkZPDo0aPPzsvKyir7348fPyYjI4O6detm71NXV8fW1pZ79+7luO59OyzIOv8AvHr16rPzF76xFRLzWyeLTkR4RuD5lyczTs3g4J2DbO22leql8r/bslB0tLGp8Fm1g4LSq1cvZs6cSXx8PH5+fhgZGdGyZUsOHDig1EHGeY38zyud/+57P9/fv4/9u3eoID9RI1GxUnql2NVrFzt77OT2y9vY+tiy7tI6McWKUOS1b98ebW1tDhw4wN69e+nWrRuampqYm5tz9erVHF/aYWFhaGpqZs8E/l+amppkZmZ+Mk8zMzM0NTUJCwvL3peZmcnVq1epVq2a4jcl5EkEknwgkUjob92f657XaWrWlFGHR9F+R3uiEnN3jxSEokJHR4fOnTuzZs0aHj9+TK9evQAYMGAAsbGxzJ49m3v37nHmzBmWLl3KoEGDPjivWPny5bl+/TpRUVHExcV9sOagp6dH//79WbJkCWfPnuXevXvMnj2bV69eMWDAAJXd67dOBJJ8VL5YeY4MPIJ3J2/OPz6P9XprdoTvELUTocjq3bs3CQkJ2NnZZdcITE1N8fX15datWzg7OzNjxgw6derEhAkTPpjO8OHD0dTUpFOnTtnrFn3I5MmT6dChA9OnT8fZ2Zl//vkHX1/f7HYQQfkk0m/oWywqKopWrVpx8uRJKlQo2HfMd+PuMsR/CMFPgulVqxfrO62ntF7pAi2TUPjcunWLmjVrFnQxhG/Eh37ePvXdKWokBcS8pDnnhp5jYauF7L+9H6t1Vhy6c6igiyUIgvDZRCApQOpq6kxtMpXLbpcxNTCly+9dGHFgBIlpiQVdNEEQBLmJQFII1DGtw8URF5neZDqbr27GxtuGsw/PFnSxBEEQ5CICSSGhraHNglYLCBwWiLpEnRZbWzAxYCKp71ILumiCIAgfJQJJIeNY0ZFrHtfwrO/JsgvLqLehHmHRYZ++UBAEoYCIQCKnpQeuMXX7BaLjklWel76WPms7reXowKPEp8bT6NdGzD07l4xMMUW2IAiFjwgkcmpcowyRzxLw3BDI4SuP82XsRzvzdkR4RtC3dl9+OvMTjpscuf3ytsrzFQRB+BwikMipkYUp3u5O1KxQgpV/XWfmH5d49Ub17RcldEvwW4/f2N1rNw9eP8DOx46VF1aSJRVzAgmCUDiIQPIZTIx0WTCwIaPa1yb84SvcvM9xJuLDI2yVqXft3kSMjKB11daMCxhH622teRT/+TOlCoIgKJsIJJ9JTSKha4PKrHNrSoVS+vzi9zcL9l4hMSXv1d2UqYxBGQ70O8DGLhu5FH0J6/XWbLm6RUyxInzTWrZsya+//lrQxfimiUDyhSqUMmDZUAeGNLfg/O3nuPuc49LdWJXnK5FIcK3rSrhHOHZl7Ri2fxjdd3UnNln1eQuCPN4vh/uhz7Rp074o3X379mFnZ6fk0grKINYjUYC6mhoDmlanobkJXvuv8ePvl+hY1wy3NjXR1VLto61Sogqnh5xmxYUVzDg5g9rrarOh8wa61+yu0nwF4VPOnz+f/e8zZ87w448/5tino6OT4/yMjAw0NTXzrXyC8okaiRKYlzVi9YjG9HaoypErj/HcEEjE4ziV56smUWOCwwTC3MIwMzKjx+4eDPEfQnxqvMrzFoQPMTY2zv4YGhrm2JeWlkb9+vU5dOgQLi4u1KlTh127duVZ2wgNDcXS0pK4uDhCQ0OZPn06KSkp2TWb1atXZ5+blpbGrFmzqFu3Lk5OTmzcuDFf7/lbJwKJkmhpqDOidU28hjgglUqZtDWEjSdukf7u04vxKKq2SW0uuF5gltMsdoTvwHq9NSfun1B5voLwpZYtW8aAAQP466+/aN269SfPt7OzY8aMGejq6nL+/HnOnz/P8OHDs49v3boVCwsL/Pz8+O677/Dy8uLvv/9W5S0I/yJebSmZtVlJ1rs54XviFn+G3OfS3RdM6WZDtTJGKs1XU12TOS3m0MmiEy5+LrTZ3obRDUazqM0i9DT1VJq3kL+2XdvGpr835Wuew+2G42LjorT0Bg0aRPv27eU+X0tLC0NDQyQSCcbGxrmON27cmEGDBgEwePBgtm/fTkhIiGhTySeiRqICetoajO1kzbx+DUh8m86YX4PYGRhJZj6sB92wfEOuuF9hrP1Y1lxag52PHaFRoSrPVxA+h5WVlVLTs7S0zLFtYmJCXJzLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(dat.Strike, dat.Voltron, c='steelblue', label=\"Voltron\")\n",
|
||||
"plt.plot(dat.Strike, dat.Return, c='green', label=\"Truth\")\n",
|
||||
"plt.fill_between(dat.Strike, dat.Bid, dat.Ask, color='OrangeRed',\n",
|
||||
" label=\"Market Bid/Ask\")\n",
|
||||
"plt.ylabel(\"price (return)\")\n",
|
||||
"plt.xlabel(\"Strike\")\n",
|
||||
"plt.legend(fontsize=14)\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 47,
|
||||
"id": "3d718747",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAZIAAAEnCAYAAACDhcU8AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAABckElEQVR4nO3dd1xV9RvA8c+9l71BERQHCjLcAxfubY7QXGmurCzNshxltmyX5c/cMytXpuDWnJgmoqLmFsWFuCcgArLu748TV5HhRe5lPu/X677gnvE9Dwflued8v+f5qrRarRYhhBDiOakLOgAhhBBFmyQSIYQQeSKJRAghRJ5IIhFCCJEnkkiEEELkiSQSIYQQeZJtIrl27RqJiYn5GYsQQogiKNtE0rZtW7Zt26Z7P2jQIEJDQ/MlKCGEEEVHtonExMSElJQU3fsDBw5w586dfAlKCCFE0ZFtIilfvjzBwcE8ePBAt0ylUuVLUEIIIYoOVXYlUpYuXcpXX32Vq+ShUqk4deqUwYITQghR+Jlkt+KVV17Bw8OD0NBQbt26xerVq/Hz86NChQr5GZ8QQohCLtsrkqf5+Pjw448/0q1bN2PHJIQQogjJcdTWjh07dO979OiBl5dXvgQlhBCi6Mg2kVy/fp2HDx/q3q9Zs4azZ8/mS1BCCCGKjmwTiYuLS4bEIdOWCCGEyEq2fSRff/01S5YswdvbG3t7ew4cOICHhwelSpXKvjGVit9//91owQohhCh8sh21NXbsWOzs7Ni7dy/Xrl1DpVJx7949EhIS8jM+IYQQhZyM2hJCCJEneieS1atX06BBA8qXL2/smIQQQhQheieSJ92/f58rV64ASikVR0dHgwcmhBCiaMi2jyQr4eHhfP311xw6dCjDcj8/Pz7++GN8fHwMGpwQQojCT+8rkrNnz9K3b1+SkpJo1aoVVatWBeDcuXPs3LkTCwsLli9frlteGCUmJnLixAmcnZ3RaDQFHY4QQhQJqamp3L59mxo1amBhYZFpvd5XJNOmTcPU1JTly5fj7e2dYd3Zs2cZMGAA06ZNY/r06XmP2khOnDjBK6+8UtBhCCFEkbR06VL8/PwyLdc7kYSFhdG/f/9MSQTAy8uLfv36sXz58rxFaWTOzs6AcjJcXV0LOBohhLGMGDECgFmzZhVwJMXDjRs3eOWVV3R/Q5+mdyJJSEjIthGAMmXKFPpnTNJvZ7m6usroMyGKsbFjxwLI/3MDy65LINsSKU+rUKECO3fuzHb9zp07pcS8EKJQaNGiBS1atCjoMEoMvRNJQEAAe/bsYcyYMURERJCamkpqaipnz55lzJgxhISE0KNHD2PGKoQQejlx4gQnTpwo6DBKDL1vbb322mucOnWKjRs3smnTJtRqJQelpaWh1Wp54YUXGDp0qNECFUIIfU2cOBGAwMDAgg2khNA7kWg0Gn7++WdCQkLYvn07V65cQavVUrFiRdq1a4e/v3+uD37jxg0WLFjAyZMnCQ8PJz4+nkWLFtGoUaMM2w0cOJADBw5k2r9z585MmTIl18cVQghhOLl6IBGgadOmNG3a1CAHj4yMZOPGjVSrVo3GjRsTHByc7bbu7u788MMPGZbJE/VCCFHw9E4ks2fP5qWXXsLFxcVgB2/QoAGhoaEAbN++PcdEYmFhQZ06dQx2bCGEEIahdyKZOnUqM2bMoHnz5vTq1YvWrVvn+enw9H6WIuP2FUhJNkxbKtXjF6rs3+e0TqUClVp5qbP6qjJMrEIIkQO9E8mKFSsIDAxk06ZN7Nq1i1KlStG9e3d69uxJ5cqVjRkjABcvXqRBgwY8fPiQ8uXL0717d9544w1MTU2NfmwAEuNhyRdgiJkitVrlj7wWSP9bn96s6olteGJZdofVap9Yr33cljZNSSZqzX9JRfM4weiW/7dOYwoaE9BonvjeBEzSvzf97/v/vqZ/b2qmrE9vy8QUTMzA1Dzzy8xc2UaSm8gHH374YUGHUKLonUhq1apFrVq1mDBhAps3byYwMJAFCxbwyy+/UK9ePXr37k2nTp2yrMOSV/Xr16dz585UqVKF+Ph4tm/fzrRp0zh58iQzZ840+PGypNVCagqUqZQ/xzMErfa/hKTN+vv0V/IjSEpUko9W+/hr2lPvs/v69NVShsz3xHE0JmBmCWYWYG4FFlb/fbUGCxuwtAZLGzC1eLzuye1M8ulDgyjyGjRoUNAhlCi57my3sLCge/fudO/encjISAIDA1mzZg0fffQRX3/9NV27dqVv3774+voaLMj33nsvw/vWrVtTunRp5syZw8GDB7Os/SJ44g98IZGWBmkpkJoKiQ/hYQykpSoJOv1r+pUUT8SdnrRMzcDKHuydwdEFnFzB2l5JThZWSjKyK6VsJ0q0sLAwQBJKfsl1InmSm5sb1atX5/jx49y+fZv4+HhWrlzJn3/+SfPmzfn6668pU6aMoWLNoHv37syZM4cjR47kTyJRq5U/ULejnlioynhr6sm/2VnettJmf4vquT1xTyz9dpiKx9+n3/LSknM/jDq9v0XPr+r0NtRPtZsDtRrUZs/3ry79CiklCe5ehRsXIOnR43Z126WBrROULg8OZcDGAazswMoWSrmBfenClVyFUaSP8JTnSPLHcyWSiIgIAgMDWbduHdHR0ZQpU4bhw4fTu3dvTE1NWbZsGQsXLmTChAksWLDA0DEDyoOQkI8d9uaWMOQb5Q/Zk7eFgEy3i3K6lcST++VR+if1tLSMt5vS3z/99clP/k9+TUlWXmkpyrLUFOX9k19TUyA1fdmjx8vSUiDlv3ZUTyWWp89NmlZJQOr/+mOe/KrWPPF9FgMFVKr/+nAsld9FTuck+RHcuAhR4Ur82jQlLlBukVWqBhV8lERjV0quYoTII70TycOHD9m4cSOBgYEcP34ctVpN8+bN6dOnD61atcrwB33UqFFYWVkZtf9i7dq1ANSuXdtox8jEvnT+Hauo0Wr/S0hJyh/y5EfKFUPyU69H8ZDwEBLjlO8THyoDGR7990pKgOTkzFc46YlSY6L0oaR34MN/t8xSle9NzZU+GLNs+uqSH8Glk3Am7IljaMHBBdyqKq/SbsrVS3r7Qogc6Z1ImjVrRmJiIq6urrz99tv06tUrx1Lsbm5uJCYmPrPdzZs3A3D8+HFAubd5//59LC0tadmyJQcPHmTevHl06NABNzc34uPj2bFjB6tWraJTp07Ur19f3x9BGJNKpXyqNzVTOszzIjUFHiU8TjRPfh97D+5fh+jbEHtXuY1nag4m5kqyibkNyUnKVU36lZCpuXIVY26lfO/wVBVrrVZp/0wYnPjn8e0+F3eoXBPKeUKZCkp/jBAiE70TSePGjenbty8tWrTQ63ZS586d6dy58zO3GzVqVIb36RNjubm5ERwcrCtdP23aNO7fv49araZy5cqMHz+egQMH6hu+KEo0JkqfhpVt7vfVapURaHHREHdfSSw3I+FWJNy5plyRqFRK0tGYKknPwua/kWPWj9tJS4MH92Df+scJyb4UVKkNHnWVKxe5YhECyMVUu8XBlStXaNu2LTt27JB5CkoirRYS4iDmDkTfhOvn4UrEfwMo/ksWKo2SwCxtn+rE1yq33eKilSsmtQaq1ALPeuBSCZzKKn04olBIr/xbo0aNAo6keHjW3848jdoSokhRqR5f6ZStDL6NleWpKRB9C+5eV65eLp9SOuvTk4u5NVjbPX6uJX2fy+EQcVgZQKAxgfI+UL0puFfP++09kSeSQPJXrhLJoUOHmDdvHkePHiU2NpanL2ZUKhWnTp0yaIBCGJ3GBEqVU15e9YGXlH6WO1fh5iW4dAIun4bkRCWxmFkqw4qf7GtJTVGSz8XjypVMldpQszlU9M2+418Yze7duwFkcqt8kqs521999VVsbGyoXbs2u3btonHjxsTHx3Ps2DG8vLyoXr26MWMVIv+YmilXLWUrQ53WSp9JzG24cQnO/wvnjyrDoLVasLRTOuLThxKnpSpDj8/9qyQp7wZQzR/Ke8kw43wybdo0QBJJftE7kcyZMwdnZ2eCgoIA8Pf3580336RJkybs2bOHd999l88//9xogQpRoNRq5Wl6RxfwbaQ8nX/3Klw5CxGHlK9paUqtMfvSyjMqoFypRByG06HKiDHfJsottXIeSpIRohjQ+1/ysWPHGDJkCE5OTkRHRwPobm01a9aMgIAApk6dyqJFi4wSqBCFikYDZSoqr3rtlJFi185D+H44vU95nsbMEuxKK6VcQFl2MgSO/a08be/XCao1UW6TCVGE6Z1IkpKSdHORmJkpl+cPHz7Urff19WXdunUGDk+IIsLMQulkd68ObV5Rbm2dDIFzh5QrFXNrpXRLqXLK9o/iYfcK+Gcl+PpD7VZQtkrGkWJCFBF6JxJnZ2du3LgBgJWVFXZ2dpw9e5b27dsDyrS5JiZyqS4EZubgUVt5JcYro8CO7lI67VUoVykW1kol6dRUOHsATu5RbofV7wBefnKVIooUvf/y16xZk3///Vf3vmnTpvz++++4ubmRlpbG0qVLqVWrllGCFKLIsrBSEoOXn/JU/tkwOLRVGQ1mYqYkD6f/rlIS4mDHUvh7uXLby6/j8z2UKfj+++8LOoQSRe9E0qtXL1atWkViYiIWFhaMHj2agwcPMn78eABKly7NuHHjjBaoEEWenZOSHOq1h2vn4PhupT8lNQVsHJV+E0sbpWbZgU1weBs07aHc9pIhxLni6elZ0CGUKHl6sj0+Pp7Q0FA0Gg3169fH1rZwf3qSJ9tFoZPwUBlOfGCTMgrM1FK5SlGrlQ78+zeU5NL0JajuLwlFT1u3bgWgQ4cOBRxJ8ZDnJ9ujo6NZtWoVkZGRODk50aVLF122t7Kyom3btoaPWoiSwtIaajRTnoi/dg4ObVNuf6k14FhWKRz5KB62L4KQVdCsp/JMitT5ytG8efMASST5JcdEcuPGDfr06cPt27d1Q33nz5/P7Nmzad68eb4Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(dat.Strike, np.abs(dat.Return-dat.Voltron), label=\"Voltron\")\n",
|
||||
"plt.fill_between(dat.Strike, np.abs(dat.Return-dat.Ask),\n",
|
||||
" np.abs(dat.Return-dat.Bid),\n",
|
||||
" color=\"OrangeRed\", alpha=0.5, label='Bid/Ask')\n",
|
||||
"plt.ylabel(\"Abs. Distance to True Payoff\", fontsize=18)\n",
|
||||
"plt.xlabel(\"Strike\")\n",
|
||||
"plt.axvline(train_y[-1], label=\"ATM\", c='k', ls='--')\n",
|
||||
"plt.legend(fontsize=14)\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "918deec6",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "d3d5be72",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,97 @@
|
||||
import numpy as np
|
||||
import datetime as dt
|
||||
import matplotlib.pyplot as plt
|
||||
import seaborn as sns
|
||||
import torch
|
||||
import gpytorch
|
||||
import os
|
||||
# import robin_stocks.robinhood as r
|
||||
import pickle5 as pickle
|
||||
import pandas as pd
|
||||
import argparse
|
||||
|
||||
import sys
|
||||
sys.path.append("../")
|
||||
from voltron.likelihoods import VolatilityGaussianLikelihood
|
||||
from voltron.models import SingleTaskVariationalGP as SingleTaskCopulaProcessModel
|
||||
from voltron.kernels import BMKernel, VolatilityKernel
|
||||
from voltron.models import BMGP, VoltronGP
|
||||
from gpytorch.kernels import ScaleKernel, RBFKernel, MaternKernel
|
||||
from voltron.option_utils import GetTradingDays, GetTrainingData, Pricer, FindLastTradingDays
|
||||
from voltron.train_utils import TrainBasicModel
|
||||
|
||||
def main(args):
|
||||
years = [yr for yr in range(2006, 2018)]
|
||||
logger = []
|
||||
full_logger = []
|
||||
SPY = pd.read_csv("./data/SPY_prices.csv")
|
||||
SPY['Date'] = pd.to_datetime(SPY['Date'])
|
||||
ntrain = 252
|
||||
|
||||
nvol = 100
|
||||
npx = 100
|
||||
|
||||
for year in years:
|
||||
options = pd.read_csv("./data/SPY_" + str(year) + ".csv")
|
||||
options.expiration = pd.to_datetime(options.expiration)
|
||||
options.quotedate = pd.to_datetime(options.quotedate)
|
||||
qday = options.quotedate.unique()[0]
|
||||
quote_price = SPY[SPY['Date']==qday].Close.item()
|
||||
options = options[(options.quotedate == qday) & (options.type=='call')]
|
||||
edays = options.expiration.sort_values().unique()
|
||||
testdays = (edays - qday)/np.timedelta64(1, "D")
|
||||
edays = edays[(testdays > 100) & (testdays < 365)]
|
||||
lastdays = FindLastTradingDays(SPY, edays)
|
||||
ntests = np.array([GetTradingDays(SPY, qday, pd.Timestamp(ld)) for ld in lastdays])
|
||||
fulltest = ntests[-1]
|
||||
|
||||
train_y = torch.FloatTensor(GetTrainingData(SPY, qday, ntrain).to_numpy())
|
||||
test_y = torch.FloatTensor(GetTrainingData(SPY,
|
||||
pd.Timestamp(lastdays[-1]),
|
||||
fulltest).to_numpy())
|
||||
full_x = torch.arange(ntrain+fulltest).type(torch.FloatTensor)
|
||||
full_x = full_x/252.
|
||||
train_x = full_x[:ntrain]
|
||||
test_x = full_x[ntrain:]
|
||||
|
||||
dmod, dlh = TrainBasicModel(train_x, train_y, train_iters=500, model_type=args.model,
|
||||
mean_func=args.mean_func)
|
||||
|
||||
## figure out how to price options sanely ##
|
||||
|
||||
nvol = 100
|
||||
npx = 100
|
||||
px_samples = torch.zeros(npx*nvol, len(edays))
|
||||
px_paths = torch.zeros(npx*nvol, fulltest)
|
||||
dmod.eval();
|
||||
|
||||
for vidx in range(nvol):
|
||||
px_pred = dlh(dmod(test_x)).sample(torch.Size((npx,))).exp()
|
||||
px_paths[vidx*npx:(vidx*npx + npx), :] = px_pred.detach()
|
||||
px_samples[vidx*npx:(vidx*npx+npx), :] = px_pred[:, ntests-1].detach()
|
||||
|
||||
|
||||
|
||||
option_output = Pricer(px_samples, options, edays, test_y[ntests-1],
|
||||
quote_price)
|
||||
option_output.to_pickle("./output/" + args.model + "_options" + str(year) + ".pkl")
|
||||
print(str(year), "Done")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
parser.add_argument(
|
||||
"--mean_func",
|
||||
type=str,
|
||||
default="loglinear",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--model",
|
||||
type=str,
|
||||
default="matern",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
main(args)
|
||||
@@ -0,0 +1,536 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 81,
|
||||
"id": "6d823c39",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import seaborn as sns\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"from matplotlib.lines import Line2D\n",
|
||||
"from matplotlib.patches import Patch\n",
|
||||
"import torch\n",
|
||||
"import pandas as pd\n",
|
||||
"import copy\n",
|
||||
"from voltron.option_utils import GetTradingDays, GetTrainingData, Pricer, FindLastTradingDays"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 82,
|
||||
"id": "14a038a2",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def ComputeStratReturns(df):\n",
|
||||
" cumrtn = copy.deepcopy(df.Return.to_numpy() - df.Ask.to_numpy())\n",
|
||||
" cumrtn[df.Model < df.Ask] = 0.\n",
|
||||
" df['rtn'] = cumrtn\n",
|
||||
" df[\"cumrtn\"] = np.cumsum(cumrtn)\n",
|
||||
" \n",
|
||||
"def FilteredReturns(df):\n",
|
||||
" rtn = (df.Return.to_numpy() - df.Ask.to_numpy())/df.Ask.to_numpy()\n",
|
||||
" rtn[df.Model < df.Ask] = 0.\n",
|
||||
" return rtn\n",
|
||||
" \n",
|
||||
"def AddData(logger, data, name):\n",
|
||||
" cols = data.columns.to_numpy()\n",
|
||||
" cols[4] = \"Model\"\n",
|
||||
" data.columns = cols\n",
|
||||
" data['Close'] = SPY[SPY['Date']==quotedate].Close.item()\n",
|
||||
" data[\"Type\"] = name\n",
|
||||
" data['diffs'] = np.abs(data.Model - data.Return)\n",
|
||||
" data[\"Year\"] = year\n",
|
||||
" ComputeStratReturns(data)\n",
|
||||
" \n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 83,
|
||||
"id": "c0fd9d39",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def LoadData(fname, quotedate, name, all_dat, type_dat):\n",
|
||||
" curr_dat = pd.read_pickle(fname)\n",
|
||||
" cols = curr_dat.columns.to_numpy()\n",
|
||||
" cols[4] = \"Model\"\n",
|
||||
" curr_dat.columns = cols\n",
|
||||
" curr_dat['Close'] = SPY[SPY['Date']==quotedate].Close.item()\n",
|
||||
" curr_dat[\"Type\"] = name\n",
|
||||
" curr_dat['diffs'] = np.abs(mat_dat.Model - mat_dat.Return)\n",
|
||||
" curr_dat[\"Year\"] = year\n",
|
||||
" ComputeStratReturns(curr_dat)\n",
|
||||
" \n",
|
||||
" if name != \"SABR\":\n",
|
||||
" ec = [d.item() for d in curr_dat.ExpClose.to_numpy()]\n",
|
||||
" curr_dat.ExpClose = ec\n",
|
||||
" \n",
|
||||
" type_dat = pd.concat((type_dat, curr_dat), ignore_index=True)\n",
|
||||
" all_dat = pd.concat((all_dat, curr_dat), ignore_index=True)\n",
|
||||
" \n",
|
||||
" return all_dat, type_dat"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 182,
|
||||
"id": "4f87870c",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/Users/gregorybenton/miniconda3/lib/python3.8/site-packages/IPython/core/interactiveshell.py:3169: DtypeWarning: Columns (19) have mixed types.Specify dtype option on import or set low_memory=False.\n",
|
||||
" has_raised = await self.run_ast_nodes(code_ast.body, cell_name,\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"# years = [2010,2011,2012,2013,2014,2015]\n",
|
||||
"years = np.arange(2006, 2018)\n",
|
||||
"logger = []\n",
|
||||
"all_dat = pd.DataFrame()\n",
|
||||
"matern = pd.DataFrame()\n",
|
||||
"sm = pd.DataFrame()\n",
|
||||
"voltron = pd.DataFrame()\n",
|
||||
"sabr = pd.DataFrame()\n",
|
||||
"\n",
|
||||
"SPY = pd.read_csv(\"./data/SPY_prices.csv\")\n",
|
||||
"SPY['Date'] = pd.to_datetime(SPY['Date'])\n",
|
||||
"\n",
|
||||
"for year in years:\n",
|
||||
" \n",
|
||||
" dat = pd.read_csv(\"./data/SPY_\" + str(year) + \".csv\")\n",
|
||||
" quotedate = pd.Timestamp(dat.quotedate.unique()[0])\n",
|
||||
"\n",
|
||||
" all_dat, sm = LoadData(\"./output/SM_options\" + str(year) + \".pkl\",\n",
|
||||
" quotedate, \"SM\", all_dat, sm)\n",
|
||||
" all_dat, matern = LoadData(\"./output/matern_options\" + str(year) + \".pkl\",\n",
|
||||
" quotedate, \"Matern\", all_dat, matern)\n",
|
||||
" all_dat, voltron = LoadData(\"./output/options\" + str(year) + \".pkl\",\n",
|
||||
" quotedate, \"Voltron\", all_dat, voltron)\n",
|
||||
" all_dat, sabr = LoadData(\"./output/sabr\" + str(year) + \".pkl\",\n",
|
||||
" quotedate, \"SABR\", all_dat, sabr)\n",
|
||||
" \n",
|
||||
" \n",
|
||||
"# all_dat = pd.concat((all_dat, mat_dat), ignore_index=True)\n",
|
||||
" \n",
|
||||
"# ec = [d.item() for d in all_dat.ExpClose.to_numpy()]\n",
|
||||
"# all_dat.ExpClose = ec"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "626fcc17",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Calibration Plots"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 44,
|
||||
"id": "72e69c0f",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAABSYAAALUCAYAAAAWr+zXAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAABcSAAAXEgFnn9JSAAD6GElEQVR4nOzdd3TUVf7/8dek94RQQwhditKkI6L0gCIEXBvrF2mia8GGDVFZUdFVFAusuIhYEKS30DuCUgXpvQYILSE9mWTm9wc/xnzSyCSTmQDPxzmc5d75fO59Dwjn7ItbTFar1SoAAAAAAAAAcCI3VxcAAAAAAAAA4NZDMAkAAAAAAADA6QgmAQAAAAAAADgdwSQAAAAAAAAApyOYBAAAAAAAAOB0BJMAAAAAAAAAnI5gEgAAAAAAAIDTEUwCAAAAAAAAcDqCSQAAAAAAAABORzAJAAAAAAAAwOkIJgEAAAAAAAA4HcEkAAAAAAAAAKfzcHUBuLG88sorOnbsmGrUqKExY8a4uhwAAAAAAADcoAgmYZdjx45pz549ri4DAAAAAAAANzi2cgMAAAAAAABwOoJJAAAAAAAAAE5HMAkAAAAAAADA6QgmAQAAAAAAADgdwSQAAAAAAAAApyOYBAAAAAAAAOB0BJMAAAAAAAAAnI5gEgAAAAAAAIDTEUwCAAAAAAAAcDqCSQAAAAAAAABORzAJAAAAAAAAwOk8XF3AzSYzM1M9e/bUkSNHJEkHDhwo8ljJycmaO3eu1q9fr3379ikuLk5Wq1XlypVT/fr11blzZ/Xo0UNeXl6OKh8AAAAAAABwCoJJB5s0aZItlCyO+fPn6/3339eVK1dyfXbmzBmdOXNGK1eu1JdffqlPP/1UzZs3L/acAAAAAAAAgLOwlduBtmzZoi+//LLY43z++ed69dVX8wwlczp79qz69eunKVOmFHteAAAAAAAAwFlYMekg+/fv13PPPSez2VyscaZMmaJvvvnG0BcYGKiuXbuqVq1aSk9P1+7du7VmzRplZWVJkrKysvTBBx+oatWqateuXbHmBwAAAAAAAJyBYNIB/vrrLz355JOKj48v1jjHjx/XRx99ZOjr0qWLPvzwQwUFBRn6jxw5ohdeeEGHDh2SdDWcHDZsmJYvX57rWQAAAAAAAKC0YSt3MU2fPl3//Oc/ix1KStJnn32mjIwMW/vuu+/WF198kWfQWKtWLf3000+KiIiw9cXHx2vy5MnFrgMAAAAAAAAoaQSTRZSUlKThw4fr7bffNoSJRXXq1CktW7bM1vby8tL7778vd3f3fN8pU6aMPv/8c5lMJlvfDz/8oOTk5GLXAwAAAAAAAJQkgkk7WSwWzZgxQ5GRkZo1a5bhs+rVqxd53Dlz5shqtdrakZGRCgsLu+57DRs2VNu2bW3tpKQkrV27tsh1AAAAAAAAAM5AMGmH9PR09enTRyNGjNDFixcNn7Vs2VK//vprkcfOGSZ269at0O/mfHb58uVFrgMAAAAAAABwBi6/sUN6err27dtn6PPy8tLzzz+vQYMGFbjtuiCJiYnau3evrW0ymdSiRYtCv9+qVStDe/369bJarYYt3gAAAAAAx7NarYbdbwDgbCaT6YbNgAgmi6FNmzYaMWKEateuXaxxDh48KIvFYmuHh4crODi40O9XrVpV/v7+trMlExMTdeLEiWJtLQcAAAAAGFksFiUnJyspKUnJycnKzMwklARQKri5ucnT01OBgYEKCgqSt7e3q0sqFILJIqhfv76GDh2qjh07OmS8o0ePGto1atSwe4yqVasaVnMeO3aMYBIAAAAAHCArK0vnzp1TYmIiQSSAUslisSg9PV3p6em6ePGivL29ValSJfn5+bm6tAIRTNrB29tbkydPVps2bRw6bkxMjKFdqVIlu8eoWLGiIZg8c+ZMsesCAAAAgFud2WzWqVOnlJ6eLkny9PRUQECAAgIC5OXlJXd39xt2CyVQmlktFinL7OoyHM/dUyY3x175YrFYZLFYlJKSosTERCUnJys9PV2nTp1SREREqQ4nCSbt4O3t7fBQUpIuX75saJctW9buMUJDQwscEwAAAABgn4yMDJ08eVJms1keHh4KDw+Xr68vQSRQgqyWLGXFnpIl+YqU7di7m4abm9z8g+VeMUImt6LdVZJ7yKtBp5eXl0JCQpSVlaWYmBglJyeX+nCSW7lLgfj4eEM7ICDA7jH8/f0N7StXrhSnJAAAAAC45cXGxspsNsvLy0vVqlWTn58foSRQwrJiT8mSGHdzhpKSZLHIkhinrNhTJTaFu7u7qlSpIn9/f1ksFp07d67E5iougslS4NqWgGuKkmL7+voWOCYAAAAAoPAyMzOVlJQkSapSpYq8vLxcXBFw87NaLFdXSt4CLMlXrm5XLyFubm4KDw+XyWSynT1ZGhFMlgIZGRmGtru7/Ut5c76TmZlZrJoAAAAA4FaWkJAg6eoikBvldlvghpdlvnlXSubkhDM03d3dbTtsr/2dVtoQTJYCWVlZhnZRgkm3HAenWm6VP8gAAAAAUAKuHY8VHBzs4koAoOgCAwMlSYmJiS6uJG8Ek6WAp6enoV2U1Y4538k5JgAAAACg8K7tbCutF0YAQGFc+zvMbC6dN5xzK3cpkDNELMp/LDnf4fwTAAAAACgaq9Vq24VWlB1tABzPvUK45H4DxlhZmco6H+Oy6a/tsLVYLLJaraXuAq8b8Hf05hMUFGRop6Sk2D1GcnKyoZ3zlm4AAAAAQOFYrVbbz3MemwXARdw9ZHK/8XaHWq//SInK/ndYaQwm+Ru2FChTpoyhfe0sE3vkPMS0bNmyxaoJAAAAAAAAKEkEk6VAWFiYoX3p0iW7x7h48aKhXa5cuWLVBAAAAAAAAJQkgslSoEqVKob26dOn7R4j5zvVq1cvTkkAAAAAAABAiSKYLAXq1q1raB85csSu95OSkhQbG2trm0wm1apVyyG1AQAAAAAAACWBYLIUiIiIUGhoqK0dHx+vkydPFvr9Xbt22W6Mk6R69erJ19fXoTUCAAAAAAAAjkQwWUrcddddhvbatWsL/e66desM7TZt2jikJgAAAAAAAKCkEEyWEl27djW0Z8yYIav1+pfKp6Wlaf78+Ya+++67z6G1AQAAAAAAAI5GMFlKdOjQwXCT9oEDBzRlypTrvvfFF18YbuS+/fbb1bBhwxKpEQAAAAAAAHAUgslSwsvLS4MHDzb0jR49WitXrsz3nalTp2rSpEmGvhdeeKFE6gMAAAAAAAAciWCyFHn88cdVr149WzszM1PPPvusRo0apePHj0uSrFar9u/fr2HDhmnkyJGG97t27ar27ds7r2AAAAAAAACgiDxcXQD+5unpqa+//lqPP/64zp07J+lqEPnzzz/r559/lo+PjywWizIyMnK9W7t2bb3//vvOLhkAAAAAAKfp2LGjYmJinDbfc889p+eff95p8wG3GlZMljIRERGaPHmyatWqleuztLS0PEPJJk2a6Mcff1RwcLAzSgQAAAAAAACKjWCyFKpRo4bmzp2rl156SZUrV873ufDwcA0fPly//PKLypYt68QKAQAAAAAAgOJhK7eDHThwwCHjeHl56emnn9bTTz+tXbt26ejRo7pw4YIsFotCQ0PVoEED1alTR25uZMsAAAAAAAC48RBM3gAaNmyohg0buroMAAAAAABcatWqVYV+9vTp0+rUqZOhr3fv3vroo48cXRaAImK5HQAAAAAAAACnI5gEAAAAAAAA4HQEkwAAAAAAAACcjmASAAAAAAAAgNNx+Q0AAAAAAEAplpSUpK1bt+r48eNKTU2Vv7+/KleurEaNGqlChQoOnevYsWM6ePCgLl26pISEBAUHB6tcuXK64447VLlyZYfOBRBMAgAAAAAAZLNp0yb169fP0Pfyyy/rqaeesnus33//Xf3797e1TSaTli9froiICFtf3bp1bT93d3fX3r17JUmXL1/W2LFjNWfOHGVkZOQa22QyqUWLFho8eLDuvfdeu2u7JiEhQd99950WL16sEydO5PtcnTp11Lt3bz3++OPy8vIq8nzANWzlBgAAAAAAyKZly5YKDw839C1cuLBIY82bN8/QbtGihSGUzM+ePXvUs2dP/frrr3mGkpJktVq1efNmDRkyRM8995ySkpLsrm/69Onq0qWLvvnmmwJDSUk6ePCgPv74Y3Xr1k0bNmywey4gJ4JJAAAAAACAbEwmk6Kiogx9Bw8e1P79++0aJzU1VUuXLjX09e7d+7rvHTlyRAMGDNCFCxcKPdfy5cv12GOPKT4+vlDPWywWffjhh3r77bcL/c41MTExGjJkiGbMmGHXe0BObOUGAAAAAADIoXfv3ho/frysVqutb8GCBapXr16hx1ixYoVSUlJsbT8/P0VGRhb4jtVq1XPPPacrV67Y+urVq6cHH3xQ9evXl5eXl44dO6ZFixZp7dq1hncPHjyo5557Tj/++KPc3Apei/af//xHP/zwQ67+Fi1aqEuXLqpbt64CAwMVHx+vffv2adGiRdqzZ4/tuczMTI0YMUJBQUHX/U5AfggmAQAAAAAAcoiIiFDz5s21ZcsWW190dLSGDRsmk8lUqDFybuPu2rWr/P39C3zHYrHo6NGjkiQ3Nze9/PLLGjRokCFobNy4saKiorR27Vq98sorSkxMtH22ZcsW/fjjj4ZzLXNatWqVvv/+e0Nf+fLl9Z///Ed33XVXrufbtm2rwYMHa968eXr33XeVmppq++ytt97S7bffXqjt6UBObOUGAAAAAADIQ85t12fPnjUElQW5cOGCNm7caOjLuT38ekaNGqUnn3wy39WP9957ryZNmiQfHx9D/4QJEwwrNbPLyMjQu+++a+irWLGiZs6cmWcomV2vXr1yzZeYmKixY8cW4tsAuRFMAgAAAAAA5KFbt27y8/Mz9C1YsKBQ70ZHRysrK8vWDg8PV+vWrQs9d1RUlP7xj39c97lGjRrphRdeMPRLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 1500x750 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"colors = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"\n",
|
||||
"sns.set(font_scale=2.0, style=\"white\", rc={\"lines.linewidth\": 4.0})\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"vidx = 0\n",
|
||||
"sidx = 6\n",
|
||||
"midx = 2\n",
|
||||
"\n",
|
||||
"fig, ax = plt.subplots(1,1,dpi=150, figsize=(10, 5))\n",
|
||||
"sns.histplot(x=\"Sample_Percentile\", data=all_dat, hue=\"Type\",\n",
|
||||
" stat=\"density\", bins=20, common_norm=False, element=\"step\",\n",
|
||||
" palette=[colors[vidx], colors[sidx], colors[midx], 'k'],\n",
|
||||
" lw=3., alpha=0.25)\n",
|
||||
"sns.despine()\n",
|
||||
"\n",
|
||||
"# legend_elements = [Patch(facecolor=colors[vidx+1], edgecolor=colors[vidx], alpha=0.75,\n",
|
||||
"# label='Voltron', lw=2.),\n",
|
||||
"# Patch(facecolor=colors[sidx+1], edgecolor=colors[sidx],\n",
|
||||
"# label='SM', lw=2.),\n",
|
||||
"# Patch(facecolor=colors[midx+1], edgecolor=colors[midx],\n",
|
||||
"# label='Matern', lw=2.)]\n",
|
||||
"plt.xlabel(\"Expiration Price Percentiles\")\n",
|
||||
"# plt.legend(handles=legend_elements, bbox_to_anchor=(0.3, 0.9),\n",
|
||||
"# fontsize=18, frameon=False)\n",
|
||||
"plt.savefig(\"./option-percentiles.pdf\", bbox_inches='tight')\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "33d276ac",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Calibration Retry ##"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 183,
|
||||
"id": "8f86b39c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def Calibration(name, percentile=0.95):\n",
|
||||
" pcts = all_dat[all_dat.Type == name].Sample_Percentile.to_numpy()\n",
|
||||
"# upper = 0.5 + percentile/2\n",
|
||||
"# lower = 0.5 - percentile/2\n",
|
||||
"# in_band = np.where((pcts < upper) & (pcts > lower))[0].shape[0]\n",
|
||||
" in_band = np.where((pcts < percentile))[0].shape[0]\n",
|
||||
" \n",
|
||||
" return in_band/pcts.shape[0]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 197,
|
||||
"id": "de4c8965",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"pcts = np.linspace(0.05, 0.95, 19)\n",
|
||||
"names = [\"Voltron\", \"SM\", \"Matern\", \"SABR\"]\n",
|
||||
"logger = []\n",
|
||||
"for pct in pcts:\n",
|
||||
" for name in names:\n",
|
||||
" clb = Calibration(name, pct)\n",
|
||||
" if name == \"Voltron\":\n",
|
||||
" name = \"Volt\"\n",
|
||||
" logger.append([clb, np.round(pct, 2), name])\n",
|
||||
" \n",
|
||||
"df = pd.DataFrame(logger)\n",
|
||||
"df.columns = [\"Calibration\", \"Percentile\", \"Type\"]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 238,
|
||||
"id": "89ce5632",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAABDgAAAJVCAYAAAA7qkUxAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAABcSAAAXEgFnn9JSAAEAAElEQVR4nOzdd3hU1fbw8e+ZyUwy6SEJJaGHltBbKAkgHaQKCl67qNeOXvV61as/e7++ithQwQKCCoIFVKQKCTWQUBJ6S0JIIaSXqef9I2ZMID2ThCTr8zw8TNln7zUhM8xZZ++1FVVVVYQQQgghhBBCCCEaMU1DByCEEEIIIYQQQghRW5LgEEIIIYQQQgghRKMnCQ4hhBBCCCGEEEI0epLgEEIIIYQQQgghRKMnCQ4hhBBCCCGEEEI0epLgEEIIIYQQQgghRKMnCQ4hhBBCCCGEEEI0epLgEEIIIYQQQgghRKMnCQ4hhBBCCCGEEEI0epLgEEIIIYQQQgghRKMnCQ4hhBBCCCGEEEI0epLgEEIIIYQQQgghRKMnCQ4hhBBCCCGEEEI0ek4NHYC4ej3++OOcOXOGTp068c477zR0OEIIIYQQQgghRLkkwSHKdebMGWJjYxs6DCGEEEIIIYQQolKyREUIIYQQQgghhBCNniQ4hBBCCCGEEEII0ehJgkMIIYQQQgghhBCNniQ4hBBCCCGEEEII0ehJgkMIIYQQQgghhBCNniQ4hBBCCCGEEEII0ehJgkMIIYQQQgghhBCNniQ4hBBCCCGEEEII0ehJgkMIIYQQQgghhBCNniQ4hBBCCCGEEEII0ehJgkMIIYQQQgghhBCNnlNDByCaD1VVUVW1ocMQos4oioKiKA0dhhBCCCGEEM2SJDhEnbHZbOTl5ZGdnU1eXh5Wq7WhQxKizmm1Wtzc3PD09MTNzQ2NRibKCSGEEEIIUR8kwSHqhM1mIyEhgfz8/IYORYh6ZbVayc7OJjs7G1dXV9q1aydJDiGEEEIIIeqBJDiEw5VMbmg0Gry9vfHw8ECv18uJnmjSbDYbJpOJnJwcMjMzyc/PJyEhQZIcQgghhBBC1ANJcAiHy8vLsyc32rdvj8FgaOiQhKgXGo0GJycnXF1d8fT0JD4+nvz8fPLy8vDw8Gjo8IQQQgghhGjS5JKicLjs7GwAvL29Jbkhmi2DwYC3tzfw93tCCCGEEEIIUXckwSEcLi8vD0CuWItmr/g9UPyeEEIIIYQQQtQdSXAIh1JV1b5bil6vb+BohGhYxe8Bq9UqWyQLIYQQQogG01y+i0qCQzhUyTeOFFUUzV3J90Bz+U9FCCGEEEJcPSwWC1FRUSxcuJCUlJSGDqfOSZFRIYQQQgghhBCiCSksLCQqKopdu3bZl0vv3r2b6dOnN3BkdUsSHEIIIYQQQgghRBNQUFBAZGQkUVFRGI1GALy8vBg2bBgDBgxo4OjqniQ4hBBCCCGEEEKIJsBms7F7924sFgv+/v6EhYXRq1cvtFptQ4dWLyTBIYQQQgghhBBCNEIpKSkcP36cESNGAODm5sbYsWPx9vame/fuKIrSwBHWL0lwCCGEEEIIIYQQjUh8fDwRERGcOHECgE6dOtG2bVsAhg4d2pChNSjZ5kKIJuT++++ne/fudO/ene3bt1f7+BtuuIHu3bvTr18/cnNzaxxHYmKiPY6FCxeW2aawsJDExMQajyGEEEIIIURzoqoqJ06c4IsvvuCLL76wJzd69uyJi4tLA0d3dZAEhxBNyMyZM+23f/vtt2ode+7cOQ4ePAjAuHHjcHd3d2RopURGRjJ16lT27NlTZ2MIIYQQQgjRVGRmZrJo0SKWL19OfHw8Wq2WAQMG8NBDD3H99dfj5+fX0CFeFWSJihBNyOjRo/Hy8iIrK4uNGzfy4osvotPpqnTsL7/8Yr993XXX1VWIXLhwgXnz5tVZ/0IIIYQQQjQFqqraa2h4eHhQWFiIXq9n4MCBDBs2DA8PjwaO8OojMziEaEL0ej2TJ08GICsri8jIyCofu3btWgBatWrFsGHD6iQ+AKvVWmd9CyGEEEII0dgVFhYSERHBkiVL7N+dtVotN9xwA48++igTJkyQ5EY5JMEhRBNTcpnKr7/+WqVjDh8+zJkzZwCYMWMGGo18NAghhBBCCFGfcnNz2bRpE++99x6bNm0iMTGRI0eO2J8PDAzEYDA0YIRXP1miIkQT079/fzp27MjZs2fZtGkTJpMJvV5f4THFszegdIJECCGEEEIIUbcyMjLYsWMHMTExWCwWAPz9/QkLCyM4OLiBo2tcJMEhRBM0ffp03n//fXJzc9m2bRvjxo0rt63NZmPdunUA9OnTh6CgoFLPR0dH8+233xIVFUVqairOzs60a9eOUaNGcdttt9GiRYsqx9W9e/dS959++mmefvppAI4dO1blfoQQQgghhGgKLl26xAcffICqqkDRLI3w8HC6d+9ur78hqk4SHKLBqcZMyDvf0GHUD7dAFGfvOh9mxowZLFy4EFVV+e233ypMcOzevZvU1FSg9OwNi8XCSy+9xHfffVeqvclkIi4ujri4OJYtW8b//vc/rrnmmrp4GUIIIYQQQjQ5mZmZeHt7A9CiRQs6dOiAVqslPDycDh06SGKjFiTBIRqM7UIk1m0PQPrBhg6lHing2xvtyI/RtBleZ6O0bduWQYMGsXfvXjZv3kxhYWG5e2MXL0/R6XRMmTLF/vhzzz3H6tWrAWjfvj133nknwcHBFBQUsHXrVpYvX05OTg4PPPAAX3zxBUOGDKk0rh9//JHU1FT++c9/AvDwww8zduzY2r5cIYQQQgghrmqqqnLy5EkiIiJISkri0Ucfxc3NDYCbbrqpyjsfiopJgkM0CFvCBqzrpoDN3NCh1DMV0g9i/ekamLIOTbvxdTbSjBkz2Lt3L/n5+WzdupVJkyZd0cZkMvHHH38AMGbMGHsmeefOnfbkxoABA/j888/tH8AAw4cPZ9KkSdx5550UFhby1FNP8ccff1T6wRwcHFyq4nNAQICsKxRCCCGEEE2WzWYjNjaWiIgI+6xpjUZDQkICPXr0AJDkhgPJVgmi3qmqim33f5thcqMEmxnb7mfta+3qwuTJk+2zNsrbTeXPP/8kOzsbKL085YsvvgCKPmzfeeedUsmNYgMGDOCBBx4AICkpyZ4oEUIIIYQQormzWCxERUXxwQcfsHr1alJTU9Hr9QwbNoxHHnnEntwQjiUJDlH/TFmoqXsbOooGp6buAVNWnfXv7u5uX/7x559/kpeXd0WbX375BQBfX19GjhwJgNlsZu/eon+f8PBwAgICyh3jhhtusK8RjIyMdGj8QgghhBBCNFYFBQX8/vvvZGRk4OrqyujRo3n00UeZMGECnp6eDR1ekyUJDiGasBkzZgBQWFjIli1bSj2Xm5vL1q1bAZg2bRpOTkUr1pKSksjPzwegd+/eFfbfokUL2rVrB8CpU6ccGboQQgghhBCNRm5uLvv377ff9/DwICwsjEmTJvHII48wcuRIDAZDA0bYPEgNDlH/9F4oLQc3+1kcSstQ0HvV6Rjh4eH4+/uTlpbGb7/9xtSpU+3P/fHHHxiNRgCuu+46++OZmZn2276+vpWO4evrS3x8fKnjhBBCCCGEaA4yMjLYsWMHMTExWCwW2rRpQ5s2bQAYPXp0A0fX/MgMDlHvFEVBM+RV0DTjYjoaHZohr9b5FlBardae1Ni2bRu5ubn254p3T+nRo0epNYDVrQtitVoBZDsrIYQQQgjRbKSkpLB69WoWLlxIVFQUFouFwMBA+3dj0TBkBodoEJp242HGVqzb7of0Q0DdFdu8utTPNrElzZw5ky+++AKTycTGjRuZOXMmFy9eZNeuXUDp2RsAXl5/zyq5dOlSpf2np6cD2HdgEUIIIYQQoqnKycnhl19+4cSJE/bHgoKCCA8Pp0OHDnLRr4FJgkM0GE2b4WjmHkA1ZkLe+YYOp364BaI4e9frkD169KB79+4cO3aM33//nZkzZ7J+/XqsVitOTk5MmzatVPt27dphMBgoKCjg4MGDFfadnp7O+fNF/3adOnWqs9cghBBCCCHE1cBgMJCcnAxASEgI4eHh9iUpouFJgkM0OMXZG+r5pL+5mTlzJm+++SaRkZHk5uayadMmAEaMGHFFnQ0nJydCQ0P5888/iYiIIDk5mdatW5fZ78qVK+23hwwZUqVYNBpZGSeEEEIIIa5+NpuN2NhYYmNjmTNnDhqNBicnJ2bMmIG3t3eV6tWJ+iVnGkI0A9OmTUOr1WIymVi/fj179uwBihIfZbn11luBoi1jn3jiCfuuKiXt37+fjz/+GIDWrVszceLEKsWi1+vtt8vqVwghhBBCiIZksViIiorigw8+YPXq1Rw7doy4uDj780FBQZLcuErJDA4hmgF/f3+GDx/O9u3beffddzGbzXh5eTFmzJgy248YMYJZs2axevVq9u7dy8yZM7nzzjsJDg6moKCArVu3snz5ckwmExqNhjfffLPK2155e3vj5OSExWLh+++/txc4HThwoKxZFEIIIYQQDcZoNLJ371527dpFXl4eULQkZciQIQQFBTVwdKIqJMEhRDMxc+ZMtm/fTlpaGgDXXnttqdkUl3vppZfQaDSsWrWKc+fO8cILL1zRxtvbm//9738MHTq0ynE4OTkxYsQItmzZwrFjx7j55puBom1rO3ToUL0XJYQQQgghhANkZWXx8ccfYzQaAfD09GT48OH079+/wu/M4uoiCQ4hmolx48bh7u5u3yr28t1TLqfT6Xj11Ve57rrr+O6779i3bx9paWm4ubnRvn17Jk6cyKxZs/Dx8al2LG+++SZvvPEG27ZtIzMzEw8PD5KTkyXBIYQQQggh6o3RaMTZ2RkoSmi0bNmSwsJCwsLC6NWrF1qttoEjrFuqqja5GdSKqqrNZX9OUU2zZs0iNjaWnj17snr16iodY7PZOHbsGADdu3eXgpKiWZP3gxBCCCHE1ScLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 1200x600 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"colors = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"fig, ax = plt.subplots(1,1,dpi=150, figsize=(8, 4))\n",
|
||||
"sns.lineplot(x='Percentile', y=\"Calibration\", hue='Type', data=df, ax=ax,\n",
|
||||
" palette=[colors[7], colors[3], colors[1], colors[5]], alpha=0.5)\n",
|
||||
"sns.scatterplot(x='Percentile', y=\"Calibration\", hue='Type', data=df, ax=ax,\n",
|
||||
" palette=[colors[6], colors[2], colors[0], colors[4]], s=120, legend=False, zorder=4)\n",
|
||||
"# ax.get_legend().set_visible(False)\n",
|
||||
"x = np.linspace(0.05,0.95)\n",
|
||||
"y = np.linspace(0, len(pcts))\n",
|
||||
"ax.plot(x, x, color=\"gray\", lw=1., ls=\"--\")\n",
|
||||
"# ax.set_xticklabels(Rotation=45)\n",
|
||||
"plt.tick_params(labelsize=16)\n",
|
||||
"sns.despine()\n",
|
||||
"plt.legend(fontsize=14, frameon=\"False\")\n",
|
||||
"# plt.label(\"Percentile\")\n",
|
||||
"plt.savefig(\"./calibration.pdf\", bbox_inches=\"tight\")\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "9f00ead1",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Compute Bias of Bid/Ask"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 45,
|
||||
"id": "7e7dd945",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/Users/gregorybenton/miniconda3/lib/python3.8/site-packages/IPython/core/interactiveshell.py:3169: DtypeWarning: Columns (19) have mixed types.Specify dtype option on import or set low_memory=False.\n",
|
||||
" has_raised = await self.run_ast_nodes(code_ast.body, cell_name,\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"opts = pd.DataFrame()\n",
|
||||
"for year in years:\n",
|
||||
" dat = pd.read_csv(\"./data/SPY_\" + str(year) + \".csv\")\n",
|
||||
" dat = dat[dat.type == \"call\"]\n",
|
||||
" quotedate = dat.quotedate.unique()[0]\n",
|
||||
" dat = dat[dat.quotedate == quotedate] \n",
|
||||
" opts = pd.concat((opts, dat), ignore_index=True)\n",
|
||||
"\n",
|
||||
"exps = [] \n",
|
||||
"for idx, row in opts.iterrows():\n",
|
||||
" eday = row.expiration\n",
|
||||
" exps.append(SPY[SPY.Date == FindLastTradingDays(SPY, [pd.Timestamp(eday)])[0]].Close.item())\n",
|
||||
"opts['exp_price'] = exps"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 46,
|
||||
"id": "a5df0836",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ask_under_exp = (opts.ask < (opts.exp_price - opts.strike)).to_numpy()\n",
|
||||
"bid_under_exp = (opts.bid < (opts.exp_price - opts.strike)).to_numpy()\n",
|
||||
"\n",
|
||||
"types = []\n",
|
||||
"for idx in range(ask_under_exp.shape[0]):\n",
|
||||
" if ask_under_exp[idx]:\n",
|
||||
" types.append(\"Above\")\n",
|
||||
" elif bid_under_exp[idx]:\n",
|
||||
" types.append(\"Between\")\n",
|
||||
" else:\n",
|
||||
" types.append(\"Below\")\n",
|
||||
" \n",
|
||||
"opts[\"Val\"] = types"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 50,
|
||||
"id": "503e693e",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAnkAAAHUCAYAAABGYz4MAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAABMmElEQVR4nO3dd3wUdf7H8fcmJLSQQgu9GAiRO+mRoCBqpIqEjlThvNOfSDtURDkUPZHjlNMTTrGcJwoK0iUovUeBBJAa6iFI6KRBetnfH3ns3C7ZhACL2Uxez8eDx2Oy853PfBeW2XfmO/Mdi9VqtQoAAACm4lHcHQAAAIDrEfIAAABMiJAHAABgQoQ8AAAAEyLkAQAAmBAhDwAAwIQIeQAAACZEyAMAADChMsXdARSfNm3aKDMzU9WqVSvurgAAgCK6fPmyvL29FRMTU2g7Ql4plpGRoZycnOLuBgAAuAXZ2dkqygPLCHmlWPXq1SVJGzZsKOaeAACAogoPDy9SO67JAwAAMCFCHgAAgAkR8gAAAEyIkAcAAGBChDwAAAATIuQBAACYECEPAADAhAh5AAAAJkTIAwAAMCFTPfEiMzNTixcv1g8//KCjR48qNTVVfn5+uu+++9SrVy916dJFFoulwO2tVqsiIyO1ZMkSxcbGKjU1VdWqVVNoaKiGDBmiZs2a3bQP7lIDAACUbhZrUR5+VgJcvHhRzz77rGJjYwts07FjR73//vuqUKFCvnXp6ekaN26cNm/e7HRbT09PjR8/Xs8880yB9d2lRlHZHovCY80AACg5ivr9bYozeVlZWQ4Br379+urbt69q1qypU6dOacGCBYqPj9eWLVv0wgsv6KOPPspXY/LkyUawCgoK0oABA1S1alUdOnRICxYsUGpqqmbOnKnAwEBFREQ47Ye71AAAADDFmbxvv/1WU6ZMkSQ98sgjev/991WuXDljfXx8vEaOHKkjR45Ikv7973+rffv2xvqoqCj94Q9/kCSFhYXpk08+UdmyZY31J0+e1ODBg5WYmCh/f39t2LBBPj4+Dn1wlxq3gjN5AACUPEX9/jbFjRdr166VJHl4eOjNN990CHiSVLlyZU2ePDlfe5vPP/9cklSmTBm99dZbDsFKyjujZguRiYmJWrRoUb4+uEsNAAAAySQh7+zZs5Lywlz16tWdtmnevLmxHBcXZywnJibqxx9/lCR16NBBdevWdbp99+7dVaVKFUnS6tWrHda5Sw0AAAAbU4S8SpUqSZKuXr2qlJQUp23sg13lypWN5ZiYGOXm5krKGyItiIeHh0JDQyVJ+/btU1JSktvVAEqD3NwSf4UJ3BCfK5iRKW68aNasmfbv3y+r1arPP/9cY8aMydfms88+M5btr8c7fvy4sRwcHFzofho1aiQpb4qTY8eOGWHLXWoApYGHh0UL1x3T5YTU4u4KTKJaQAUN7FT4cRcoiUwR8p566iktXbpUqamp+vDDD5WcnKwnn3xStWrV0pkzZ/Tvf/9bK1askCTdf//96tGjh7Gt/Rm+2rVrF7qfGjVqOGxnC1fuUgMoLS4npOrcFedn7QEAeUwR8urVq6dPP/1UEyZM0MWLF/Xll1/qyy+/dGjj5eWlJ598Ui+88II8PT2N1+Pj443lgICAQvfj7+9vLCcmJrpdDQAAABtTXJMnSW3atNE//vEPhwBkz9fXVw0bNnQIeFLe5MM2N97NeiNvb2+n27lLDQAAABtTnMnLysrSyy+/rFWrVknKu3Ghc+fOCggI0NmzZ7VixQqdOHFCb775plavXq1PP/3UmGYlOzvbqGMfnpyxX2+/nbvUAAAAsDHFmbwXXnjBCHhTpkzR3LlzNWTIEHXv3l3PPPOMvvvuOw0cOFCStGvXLr366qvGtvZz6mVlZRW6n8zMTGPZPmi5Sw0AAACbEh/ydu3apTVr1kiSevfuraFDh+Zr4+npqddff13NmjWTJK1atcq4m9X+ObYZGRmF7ss+XNkPqbpLDQAAAJsSH/Lsn14xePDgAtt5eno6BEDb82F9fX2N1252E4P9evu59tylBgAAgE2JD3mnT582lps0aVJo26ZNmxrLtqdkNGjQwHjt/PnzhW5/4cIFY7lWrVrGsrvUAAAAsCnxIc9q/d8s5Tcb5vTw+N/btd1lGxQUZLxmPyGxM7b1FotFjRs3Nl53lxoAAAA2JT7k2U8MfPDgwULbHjt2zFi2nQFr2bKlvLy8JEk7d+4scNucnBxFR0dLkkJCQhyGV92lBgAAgE2JD3nt2rUzlm+cANme1WrV119/bfxse7SZr6+v8azYjRs36ty5c063X7VqlTFhcbdu3RzWuUsNAAAAmxIf8jp16mQ8BmzTpk366KOP8rWxWq165513tGvXLkl5AS8kJMRYP2LECEl5U5dMmDBB169fd9j+xIkTmjZtmiSpYsWK6t+/f759uEsNAAAAyQSTIXt7e2vGjBkaOXKksrKy9P7772v9+vXq0aOHAgMDdenSJX333Xc6dOiQpLy7Uf/617861Gjfvr26dOmiNWvWaO/evYqIiNCgQYNUs2ZNHT58WN98841SUvKekzlx4kSnd7S6Sw0AAABJsljt71wowaKiojRhwoRCpx9p0KCBPvzwQ4ebHGzS0tI0evRobd++3em2FotFo0eP1ujRowus7y41iio8PFyStGHDhjuuBfyWZn/7s85dSSnubsAkalWtqNEDWhR3N4AiK+r3d4k/k2fz4IMPat26dfrmm2+0adMm/fe//1VKSop8fX0VEhKizp07q2/fvgU+IaJ8+fL67LPPtHLlSi1fvlyxsbG6du2a/P391bp1aw0fPlytW7cutA/uUgMAAMA0Z/Jw6ziTh5KKM3lwJc7koaQp6vd3ib/xAgAAAPkR8gAAAEyIkAcAAGBChDwAAAATIuQBAACYECEPAADAhAh5AAAAJkTIAwAAMCFCHgAAgAkR8gAAAEyIkAcAAGBChDwAAAATIuQBAACYECEPAADAhAh5AAAAJkTIAwAAMCFCHgAAgAkR8gAAAEyIkAcAAGBChDwAAAATIuQBAACYECEPAADAhAh5AAAAJkTIAwAAMCFCHgAAgAkR8gAAAEyIkAcAAGBChDwAAAATIuQBAACYECEPAADAhAh5AAAAJkTIAwAAMCFCHgAAgAkR8gAAAEyIkAcAAGBChDwAAAATIuQBAACYECEPAADAhAh5AAAAJkTIAwAAMCFCHgAAgAkR8gAAAEyIkAcAAGBChDwAAAATIuQBAACYECEPAADAhAh5AAAAJkTIAwAAMCFCHgAAgAkR8gAAAEyIkAcAAGBChDwAAAATIuQBAACYECEPAADAhAh5AAAAJkTIAwAAMCFCHgAAgAkR8gAAAEyIkIe7IjfXWtxdAACgVCtT3B2AOXl4WLRw3TFdTkgt7q7ARILrBahzWP3i7gYAlAiEPNw1lxNSde5KSnF3AyZSzb98cXcBAEoMhmsBAABMyHRn8k6ePKlvvvlG27dv18WLFyVJdevW1SOPPKKnnnpKlStXLnDbrKwsffvtt1q5cqWOHz+urKws1ahRQw8++KCGDRume+6556b7d5caAACgdDNVyPviiy/07rvvKisry+H1o0eP6ujRo1q0aJE+/PBDtWjRIt+2CQkJ+tOf/qQDBw44vH769GmdPn1aS5cu1RtvvKFevXoVuH93qQEAAGCakPfVV19p+vTpkqTy5curX79+uu+++5Senq5Vq1Zp586dunr1qp555hmtWrVK1apVM7bNycnR6NGjjWDVvHlz9erVSz4+Ptq9e7eWLFmi9PR0TZ48WTVr1lTbtm3z7d9dagAAAEgmCXlnz57Vu+++K0mqXLmy5s6dq+DgYGP9wIED9dZbb+mrr75SUlKS5syZoylTphjrlyxZopiYGElSz549NWPGDHl4eBg/9+jRQ08//bQyMjL0xhtvKDIy0ljvbjUAAAAkk9x48a9//Uvp6emSpPfff98h4NlMnDjRuB7vhx9+cFj3+eefS5L8/Pz0+uuv5wtOoaGhev755yXlXfO3fv36fPXdpQYAAIBkgpCXmZmptWvXSpIeffTRAocwvb29NXr0aA0fPlxDhw5VZmamJCk2NlanTp2SJD3xxBPy8fFxuv2gQYPk6ekpSVq9erXDOnepAQAAYFPih2t/+uknXb9+XZLUu3fvQtsOGTIk32u7du0yltu1a1fgtr6+vmratKkOHDig7du3u2UNAAAAmxJ/Ju/IkSPGcvPmzY3l+Ph4xcTEKCoqSmfPni1w++PHjxvLjRs3LnRfjRo1kiQlJSXp3LlzblcDAADApsSfybOFI29vbwUGBurMmTP629/+pi1btig7O9tod9999+nVV19Vq1atHLaPi4uTJHl4eKhmzZqF7qtGjRrG8rlz51SrVi23qgEAAGBT4s/k2SY89vPzU3R0tCIiIrRhwwaHgCdJBw4c0LBhw7Rq1SqH1+Pj4yVJFStWlLe3d6H78vf3N5YTExPdrgYAAIBNiQ95KSl5z0ZNS0vT6NGjlZqaqn79+ikyMlIHDhzQ+vXr9ac//UkeHh7Kzs7WpEmTFBsba2xvuyu3bNmyN92XffhKS0tzuxoAAAA2pgl5169fV2JiosaOHatp06apcePG8vb2Vt26dfXiiy/q9ddfl5R3N+4777xjbG8743ezs2c3tsnJyXG7GgAAADYlPuTZCw4O1qhRo5yue/LJJ40bM6KionT+/HlJUrly5SQp36PQnLFNuyJJXl5exrK71AAAALAp8SGvfPnyxvLjjz8ui8VSYNuuXbsay3v27JEkVahQQZKUkZFx033Zhyv7YVV3qQEAAGBT4kOe/aTB99xzT6FtGzZsaCzbbtjw9fWVlDfce+PNGjeyv8nB9vQMd6oBAABgU+JDXp06dYrc1v5attzcXElSgwYNjJ9twa8gFy5cMJZr165tLLtLDQAAAJsSH/Lsn1Nrm2uuIFeuXDGWAwMDJUlBQUHGaydLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 640x480 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"sns.histplot(x=\"Val\", data=opts)\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "69e903c8",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Buy"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 74,
|
||||
"id": "44107413",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"mat_returns = FilteredReturns(matern)\n",
|
||||
"sm_returns = FilteredReturns(sm)\n",
|
||||
"volt_returns = FilteredReturns(voltron)\n",
|
||||
"sabr_returns = FilteredReturns(sabr)\n",
|
||||
"all_returns = (voltron.Return.to_numpy() - voltron.Ask.to_numpy())/voltron.Ask.to_numpy()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 75,
|
||||
"id": "1e7e1cc8",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"(5134,)"
|
||||
]
|
||||
},
|
||||
"execution_count": 75,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"sm_returns.shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 76,
|
||||
"id": "be55c590",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"(5134,)"
|
||||
]
|
||||
},
|
||||
"execution_count": 76,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"volt_returns.shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 78,
|
||||
"id": "03eb271f",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAlYAAAGvCAYAAACZ0JtTAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAADvG0lEQVR4nOyddXwbV9aGH5GZmRI75DA4zNykgQZK26aU4rZb7m5ht1vGFLbd0nbLX7lN03bbpBCuw+w4ccBOYmaS2bJgvj8Uy1Y0kmVbpuQ+/eVXaebOzJUtS++cc+57FJIkSQgEAoFAIBAI2o2yqycgEAgEAoFAcL4ghJVAIBAIBAKBixDCSiAQCAQCgcBFCGElEAgEAoFA4CKEsBIIBAKBQCBwEUJYCQQCgUAgELgIIawEAoFAIBAIXIS6qydwITF27FgaGhoIDQ3t6qkIBAKBQCBwkuLiYtzc3Ni/f3+LY4Ww6kR0Oh1Go7GrpyEQCAQCgaAVGAwGnPVTF8KqEwkLCwNg06ZNXTwTgUAgEAgEzjJnzhynx4oaK4FAIBAIBAIXIYSVQCAQCAQCgYsQwkogEAgEAoHARQhhJRAIBAKBQOAihLASCAQCgUAgcBFCWAkEAoFAIBC4CCGsBAKBQCAQCFyEEFYCgUAgEAgELkIIK4FAIBAIBAIXIYSVQCAQCAQCgYsQwkogEAgEAoHARQhhJRAIBAKBQOAiRBPmHowkSRgMBoxGY1dPRSDoclQqFWq1GoVC0WVzaNAbWbM5jWPpZUSGevOnufF4e2pIPlVCTmEVbhoVU0dGE+Dr3mVzFAgEHYsQVj2QhoYGtFotFRUVGAyGrp6OQNBtUKvV+Pv7ExAQgJubW6df/6XP9rMnpQCApLRift2ZYTPmq/UnefDaMYyKD+vk2QkEgs5ACKsehk6nIyMjAwB/f398fHxQqVRdepcuEHQ1kiRhNBqprq6mvLyc8vJy4uLicHfvvMhQel6FRVQ5orKmgX99eZD3/j4XD3fxESwQnG+Iv+oehMFgIDs7G41GQ2xsLCqVqqunJBB0K3x8fAgNDSUzM5Ps7Gzi4uJQqzvnY25bUq7TY8urdOxJKWDG6JgOnJFAIOgKRPF6D6Ix9RcTEyNElUBgB5VKRUxMDAaDgYqKik65piRJ7Dic16pjdh5p3XiBQNAzEMKqB1FdXY23t3eX1I4IBD0JNzc3vL29qa6u7pTrZRZUkVdS06pjDpwoQqcXC08EgvMNIax6CCaTibq6Ory9vbt6KgJBj8Db25u6ujpMJlOHX2v7YftpwOhQH9ntugYjSSeLOmpKAoGgixA1Vj0Eg8GAJEmdWowrEPRk3N3dLZYkHRnlNZokNu3Llt33f0/MJ8jPg9Sscv7670Sb/buO5jNhWGSHzU0gEHQ+ImLVQ2i861Yqxa9MIHCGxr+VjopYlVbU8euuDF7+fD8l2jqb/ROGRhDk5wFA/5gAgv09bMZs2pdNaUUdeoNICQoE5wsiYtXDELYKAoFzdOTfysnMMh777y7qdPZ95OZPjLU8VioVTBwWybod6TbjVj69HqUCBvQKZPG0vsxIiBZ/5wJBD0aEPwQCgaCVfPhTikNRFRbkxehB4VbbJjlI+ZkkOJlVzqtfHOC9H48gSZLL5ioQCDoXIawEAoGgFVRU6zieUeZwzJJpfVEpraNOQ/sF4+ulafH8a7ens+9YYbvmKBAIug4hrAQCgaAVpGVrHe4PCfBkwaQ4m+1qlZKpo6Kdusb3W0+1ak6SJJFTmU9ela0g0xv1IgImEHQiosZKIBAIWsGpHK3D/SsXDcFNI2/ge/VFAzl6upTswiqH50g5U0pZZb2l+N0ekiSxNzeJb478TE5lvmX7gKA4fNx9OFWWQZWuGo1SjbebF4vi5zAsfCB9AnuhVIj7aoGgIxDCSiAQCFrBKQcRq4snxTE9wX5UKtDPg9fun8HR0yUUl9eh0xvZm1JA8qkSm7F7juazYHIf2fMYTEaSC46z+uhaTpdn2uxPK8uweq43GdDWV/JF8g8AhHoFcd2oy5jYa7TduQoEgrYhhJVA0Eruvvtu1q9fD8A//vEPbrjhhlYdX11dzeTJk9HpdISGhrJ169Z29bPbs2cP119/PQB33XUXd999t82YsrIy1q1bx3XXXdfm63Q3Gox66vR1gAIfN29UnWBFIkmS3YjV7cuHs3BKH5sVfXqjnpMlZ0gpSqWktgylQolKqULtpkLloSJutIljNbmYqgIwVQcC5uN3HMlj0GA1O7P2k1xwnPL6CvQmAwajgQZT+9J7xbVl/Gvn+9w8+irmD5jR5vMIBAJbhLASCFrJ0qVLLcLq119/bbWwWr9+PTqdDoBLLrmkw5sE//TTTzz33HPEx8efN8KqTl9PQXUxJsnsUVVSW0aQZ4BVeqtBp6NOX88f6bsxqszjJEmiSY5INGoT81bp7JjG51iOaRxTWdNAhddp1N5YxgPE9w6kNuAYq1OOWc5tNJlIL8/iePEpdMYGh69H06vpsVEbiqR344RHOQ+vr23FT6X1fHToG/oE9iI+pG+HXkcguJAQwkogaCUzZswgMDCQ8vJykpKSyM/PJzLSeffstWvXWh4vW7asA2Zozeuvv45Wq+3w63Qm2vpKi6hqpKxOa/Xc2GCguqGGH079Rpnedc2Ym4ugRtJNkJ7imvOrAopdcyInkCSJjw5+wwsXPSK8swQCFyGqFwWCVqLRaFi0aBFg/mL69ddfnT62pKSE3bt3AzBkyBAGDhzYIXM839EZdF09hW6NSqHE282LuIAYfN28HYqmM+VZpBSd7MTZCQTnNyJiJRC0gaVLl/L5558D8Msvv3DTTTc5ddy6deswGs3tSzojWnW+YpQ6vrFyT8Hfw4/lg+czt980NEo1RpMRtcr6o72moRZtfSW/pm1h/SnbnoVPb/03t4+7jll9JonIlUDQToSwEgjawIgRI+jXrx+nT5/myJEjZGdn06uXTI7oHBrTgBqNhksuuaSjp3lecm4KsCehUapRK9UYJCNGk9Hp1yIZlZhq/ZB0XmBSolaquXz6MIaE9WdQSD8rIXWuqALwdvPC282LG0Zdzu7sg1Tqqm3GvLvvM44Vp/KX8dcLKwZBj6OwrJY9R/Px8tAwY3Q0GrW85UlnIISVQNBGli5dyr/+9S/AXMR+2223ORyflZVFcnIyANOmTSMoKMhqf0FBAZ9//jnbt28nOzsbvV5PSEgIo0eP5rLLLmPSpEmtmt/s2bPJzc21PN+7d68l9Whv9WBPoCeZXaqVavoHxTIsfCAjwgfTP7gPamXTB75JMlGv17E6ZR0bT29HZzSnOCWjEsnghqkqEJM2DKM2FExNH9d6YOBF4xkWHmrZVlnTgEIBvl5uduejUWmYFjuBdambZPcnZuwhLiCGxQPntvOVCwSdx+G0Yp76YDd6g/lG5Yc/TvGv+2bgbsdPrqMRwkogaCNLlizh9ddfx2Qy8csvv7QorH7++WfL4+XLl1vt+/LLL3nxxRctqwUbyc3NJTc3l59//pn58+ezatUqPD09XfcieiD2ojz+7r74efjSmMiqr6+nwbOOx2bci8ZdAwqFZZ8CRaOrAYqz/519YnmsADibFjtwvIjXvz5kc83IYG9evXcGCoX8edxVbigd2EAoFUq83Dy5IeFyrh91GcmnC/n811RKtPUE+3sQ3zuQkCGefLzWtjJ+d0o+I+NDSTlType/nyD5VAkKBYwZFM6flw8nIthb9ppz+k7h17Qtdn+Oa479yty+U/HQODYnFQi6Cx/876hFVAFkFVTxU+JprpgT3yXzEcLqPCUptYiNe7MpKKvp6ql0OhFB3swd34tR8WEdep3IyEjGjx/P7t27OX78OOnp6fTpI2/oCE1pwICAAGbOnGnZ/uWXX/LUU08BoFAomD9/PlOnTsXLy4sTJ06wevVqysvL+f333ykvL+eTTz5BpWr5Tuzpp5+mvr6exx57jLKyMgYMGMB9990H4HCe3R17gsBd7YabqqkXn1FlRKVUEewdiIdH20SC3mDiX18eYPvhPMA2ErRwwkB83eUFTGtRKBSM7B/ByLsjrLYbjSa+25xKVa3eavva7enkl9Rw4ESRZZskwf7jheQVV/PvB2bi4W77ER/jH8mKEcv4/PD3svOoaahld84hZvZpXYRUIOgK8ktqyMivtNm+IzlPCCuB60hKLeLJ93djNPWclIkrOZlZzvbDuTx16yRGxoe2fEA7WLZsmWWV3y+//MKdd94pOy4lJYUzZ84AsHjxYtzczF/SOTk5PP/88wB4eXnx9ttvM3nyZMtxixYt4qabbuK2224jOTmZvXv38sEHH/DnP/+5xblNnToVwHL+wMBA5s7t+Skek51UYEfUBf2yM/2sqLLF013FvAm9XX7Nc1GplIwbEsHm/dk2+5qLqubkldSwaV8Wi6bK+1MtGXQRE2JG8dOJDWw4vc1m/9b0XUJYCXoEB0/K/w2czqmgtKKOYP/Oj/B3ibB68803eeutt1p93PLly3nxxRctz/fu3eu04eHUqVP58MMPZfdJksTatWtZs2YNx48fp7a2ltDQUMaNG8c111zDiBEjWj3XrmTj3uwLVlQ1YjRJbNib1eHCat68eTz99NPU1tby66+/2hVW9ryrPvzwQ/R6cyTioYceshJVjQQGBvLWW2+xcOFCqqur+fjjj7nhhhvaHIXp6diLWCk7YDXbpn1ZdvddMq0fXh4au/tdyeThkbLCyhE7kvPtCiuAcJ9Qbh27gtTSdDK1OVb7jhefQltXQYCnf5vmKxB0Fgft3FwAlFbUd4mw6lFLP85dBnzyZPu9V+rr67n99tv529/+xq5du9BqtTQ0NJCbm8uPP/7IVVddxXvvvdfu6wjOT7y9vbnooosASEtLIy0tzWaMyWRi3bp1APTv35/hw4db9v3xxx+AOT142WWX2b1OeHi4RZCVl5dz8OBBV72EHoc9YaVwccSquLyO9DzbFANAVIgLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 640x480 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(np.cumsum(volt_returns), label=\"Volt\")\n",
|
||||
"plt.plot(np.cumsum(sm_returns), label=\"SM\")\n",
|
||||
"plt.plot(np.cumsum(mat_returns), label=\"Matern\")\n",
|
||||
"plt.plot(np.cumsum(sabr_returns), label=\"SABR\")\n",
|
||||
"plt.plot(np.cumsum(all_returns), label=\"All\")\n",
|
||||
"plt.legend()\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 79,
|
||||
"id": "66b1e899",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"786.2182778724494"
|
||||
]
|
||||
},
|
||||
"execution_count": 79,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"np.cumsum(volt_returns)[-1]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 80,
|
||||
"id": "339b1a97",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"718.9145160871468"
|
||||
]
|
||||
},
|
||||
"execution_count": 80,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"np.cumsum(all_returns)[-1]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "8ba78ce9",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## ITM"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 104,
|
||||
"id": "be1a215d",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"itm_dat = full_dat[full_dat.Strike < full_dat.Close]\n",
|
||||
"rtn_pct = (itm_dat.Return - itm_dat.Ask)/itm_dat.Ask\n",
|
||||
"volt_return = copy.deepcopy(rtn_pct)\n",
|
||||
"volt_return[itm_dat.Model < itm_dat.Ask] = 0."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 105,
|
||||
"id": "60bdb8f9",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"cum_rtn = np.cumsum(rtn_pct.to_numpy())\n",
|
||||
"volt_cum_rtn = np.cumsum(volt_return.to_numpy())"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,835 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "753ddfb2",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import datetime as dt\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"from matplotlib.lines import Line2D\n",
|
||||
"from matplotlib.patches import Patch\n",
|
||||
"import seaborn as sns\n",
|
||||
"import torch\n",
|
||||
"import gpytorch\n",
|
||||
"import os\n",
|
||||
"# import robin_stocks.robinhood as r\n",
|
||||
"import pickle5 as pickle\n",
|
||||
"import pandas as pd\n",
|
||||
"# sns.set_style(\"white\")\n",
|
||||
"# sns.set_palette(\"bright\")\n",
|
||||
"\n",
|
||||
"sns.set_style(\"white\")\n",
|
||||
"\n",
|
||||
"import sys\n",
|
||||
"sys.path.append(\"../\")\n",
|
||||
"from voltron.likelihoods import VolatilityGaussianLikelihood\n",
|
||||
"from voltron.models import SingleTaskVariationalGP as SingleTaskCopulaProcessModel\n",
|
||||
"from voltron.kernels import BMKernel, VolatilityKernel\n",
|
||||
"from voltron.models import BMGP, VoltronGP\n",
|
||||
"from voltron.train_utils import LearnGPCV, TrainVolModel, TrainDataModel, TrainBasicModel\n",
|
||||
"from gpytorch.kernels import ScaleKernel, RBFKernel, MaternKernel"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "1d4e4dd6",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Options Helpers"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "e15487d9",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def GetTrainingData(SPY, date, N):\n",
|
||||
" idx = SPY[SPY[\"Date\"] == date].index.item()\n",
|
||||
" return SPY['Close'].iloc[(idx-N):idx]\n",
|
||||
"\n",
|
||||
"def GetTrueValue(SPY, date, strike):\n",
|
||||
" close_px = SPY['Close'][SPY[\"Date\"] == date].item()\n",
|
||||
" return np.maximum(close_px-strike, 0)\n",
|
||||
"\n",
|
||||
"def GetTradingDays(SPY, start, stop):\n",
|
||||
" start_idx = SPY[SPY[\"Date\"] == start].index.item()\n",
|
||||
" stop_idx = SPY[SPY[\"Date\"] == stop].index.item()\n",
|
||||
" return stop_idx-start_idx\n",
|
||||
"\n",
|
||||
"def FindLastTradingDays(SPY, dates):\n",
|
||||
" last_days = []\n",
|
||||
" for date in dates:\n",
|
||||
" last_days.append(np.max(np.where(SPY.Date < date)[0]))\n",
|
||||
" \n",
|
||||
" return np.array(SPY.Date[last_days])\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"def BlackVol(pars, K, f, T):\n",
|
||||
" alpha = torch.exp(pars[0][0]) ## v0\n",
|
||||
" rho = 2 * torch.sigmoid(pars[0][1]) - 1. ##rho\n",
|
||||
" v = torch.exp(pars[0][2]) ## \"sigma\" \n",
|
||||
" beta = 1.\n",
|
||||
" num = 1 + (alpha**2 * (1-beta)**2/(24 * (f*K)**(1-beta)) +\\\n",
|
||||
" 0.25 * rho*beta*v*alpha/((f*K)**(0.5*(1-beta))) +\\\n",
|
||||
" v**2*(2-3*rho**2)/24)*T\n",
|
||||
" num*= alpha\n",
|
||||
" \n",
|
||||
" denom = (f*K)**(0.5*(1-beta)) * (1 + (1-beta)**2/24 * torch.log(f/K)**2 +\\\n",
|
||||
" (1-beta)**4/1920 * torch.log(f/K)**4)\n",
|
||||
" \n",
|
||||
" z = v/alpha * (f*K)**(0.5*(1-beta)) * np.log(f/K)\n",
|
||||
" xi_z = torch.log((torch.sqrt(1 - 2 * rho * z + z**2) + z - rho)/(1-rho))\n",
|
||||
" \n",
|
||||
" return num/denom * z/xi_z\n",
|
||||
"\n",
|
||||
"def MinVol(pars, Ks, Fs, Ts, ivol):\n",
|
||||
" return torch.mean((ivol - BlackVol(pars, Ks, Fs, Ts)).pow(2))\n",
|
||||
"\n",
|
||||
"def Calibrate(Fs, Ks, Ts, ivol, iters=1000):\n",
|
||||
" pars = [torch.tensor([-1., -5., -3.], requires_grad=True)]\n",
|
||||
" opt = torch.optim.SGD(pars, lr=0.1)\n",
|
||||
" stored_pars = torch.zeros(iters, 3)\n",
|
||||
" losses = []\n",
|
||||
" for e in range(iters):\n",
|
||||
" stored_pars[e, :] = pars[0]\n",
|
||||
" loss = MinVol(pars, Ks, Fs, Ts, ivol)\n",
|
||||
" opt.zero_grad()\n",
|
||||
" loss.backward()\n",
|
||||
" losses.append(loss.item())\n",
|
||||
" opt.step() \n",
|
||||
" \n",
|
||||
" return pars[0].detach().numpy()\n",
|
||||
"\n",
|
||||
"def SABRSim(Np, Nt, S0, V0, sigma, rho, dt=1./252.):\n",
|
||||
" dW = np.random.randn(Nt+1, Np) * np.sqrt(dt)\n",
|
||||
" dZ = rho * dW + np.sqrt(1-rho**2) * np.random.randn(Nt+1, Np) * np.sqrt(dt)\n",
|
||||
" \n",
|
||||
" S = np.zeros((Nt+1, Np))\n",
|
||||
" S[0] = S0\n",
|
||||
" V = np.zeros((Nt+1, Np))\n",
|
||||
" V[0] = V0\n",
|
||||
" \n",
|
||||
" for t in range(Nt):\n",
|
||||
" S[t+1] = S[t] + V[t]*S[t]*dW[t]\n",
|
||||
" V[t+1] = V[t] + sigma*V[t]*dZ[t]\n",
|
||||
" \n",
|
||||
" return S[1:]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "daf6bcbc",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Data"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"id": "15abd7f3",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"SPY = pd.read_csv(\"./data/SPY_prices.csv\")\n",
|
||||
"SPY['Date'] = pd.to_datetime(SPY['Date'])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "37b2df7b",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Data Setup"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"id": "31609ba5",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# ntrain = 375\n",
|
||||
"ntrain = 252\n",
|
||||
"options = pd.read_csv(\"./data/SPY_\" + str(2012) + \".csv\")\n",
|
||||
"options.expiration = pd.to_datetime(options.expiration)\n",
|
||||
"options.quotedate = pd.to_datetime(options.quotedate)\n",
|
||||
"qday = options.quotedate.unique()[0]\n",
|
||||
"quote_price = SPY[SPY['Date']==qday].Close.item()\n",
|
||||
"options = options[(options.quotedate == qday) & (options.type=='call')]\n",
|
||||
"edays = options.expiration.sort_values().unique()\n",
|
||||
"testdays = (edays - qday)/np.timedelta64(1, \"D\")\n",
|
||||
"edays = edays[(testdays > 100) & (testdays < 1000)]\n",
|
||||
"lastdays = FindLastTradingDays(SPY, edays)\n",
|
||||
"ntests = np.array([GetTradingDays(SPY, qday, pd.Timestamp(ld)) \n",
|
||||
" for ld in lastdays])\n",
|
||||
"fulltest = ntests[-1]\n",
|
||||
"\n",
|
||||
"train_y = torch.FloatTensor(GetTrainingData(SPY, qday, ntrain).to_numpy())\n",
|
||||
"test_y = torch.FloatTensor(GetTrainingData(SPY, \n",
|
||||
" pd.Timestamp(lastdays[-1]),\n",
|
||||
" fulltest).to_numpy())\n",
|
||||
"full_x = torch.arange(ntrain+fulltest).type(torch.FloatTensor)\n",
|
||||
"full_x = full_x/252.\n",
|
||||
"dt = full_x[1] - full_x[0]\n",
|
||||
"train_x = full_x[:ntrain]\n",
|
||||
"test_x = full_x[ntrain:]\n",
|
||||
"\n",
|
||||
"## SABR STUFF ##\n",
|
||||
"train_x = full_x[:ntrain]\n",
|
||||
"test_x = full_x[ntrain:]\n",
|
||||
"\n",
|
||||
"## extract data for calibration ##\n",
|
||||
"ivol = torch.tensor(options.impliedvol.to_numpy())\n",
|
||||
"Fs = torch.tensor(options.underlying_last.to_numpy())\n",
|
||||
"Ks = torch.tensor(options.strike.to_numpy())\n",
|
||||
"starts = options.quotedate.dt.date.to_numpy()\n",
|
||||
"ends = options.expiration.dt.date.to_numpy()\n",
|
||||
"Ts = torch.tensor(([np.busday_count(qd, ed)/252. for qd, ed in zip(starts, ends)]))\n",
|
||||
"\n",
|
||||
"pars = Calibrate(Fs, Ks, Ts, ivol)\n",
|
||||
"v0 = np.exp(pars[0])\n",
|
||||
"1/(1 + np.exp(-pars[1]))\n",
|
||||
"rho = (2/(1 + np.exp(-pars[1])) - 1.)\n",
|
||||
"sigma = np.exp(pars[2])\n",
|
||||
"sabr_paths = SABRSim(1000, fulltest, quote_price, v0, sigma, rho)\n",
|
||||
"sabr_samples = torch.tensor(sabr_paths[ntests-1])\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 5,
|
||||
"id": "c196bbc1",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Using gp-exp parameterization.\n",
|
||||
"Iter 1/500 - Loss: 14.952\n",
|
||||
"Iter 51/500 - Loss: -0.533\n",
|
||||
"Iter 101/500 - Loss: -0.591\n",
|
||||
"Iter 151/500 - Loss: -0.593\n",
|
||||
"Iter 201/500 - Loss: -0.594\n",
|
||||
"Iter 251/500 - Loss: -0.594\n",
|
||||
"Iter 301/500 - Loss: -0.594\n",
|
||||
"Iter 351/500 - Loss: -0.594\n",
|
||||
"Iter 401/500 - Loss: -0.594\n",
|
||||
"Iter 451/500 - Loss: -0.594\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"## learn vol with GPCV ##\n",
|
||||
"vol = LearnGPCV(train_x[1:], train_y, train_iters=500)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 6,
|
||||
"id": "c504764d",
|
||||
"metadata": {
|
||||
"tags": []
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"## train vol GP ## \n",
|
||||
"vmod, vlh = TrainVolModel(train_x[1:], vol, train_iters=500, printing=False)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 7,
|
||||
"id": "3a1ee561",
|
||||
"metadata": {
|
||||
"tags": []
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"## train data gp ##\n",
|
||||
"dmod, dlh = TrainDataModel(train_x[1:], train_y[1:], vmod, vlh, vol, train_iters=500)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 8,
|
||||
"id": "6b26da10",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"mat_mod, mat_lh = TrainBasicModel(train_x, train_y, train_iters=500)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 43,
|
||||
"id": "166e0a09",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"sm_mod, sm_lh = TrainBasicModel(train_x, train_y, train_iters=750, model_type=\"SM\", mean_func=\"constant\",\n",
|
||||
" num_mixtures=5)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 78,
|
||||
"id": "30798414",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/Users/gregorybenton/miniconda3/lib/python3.8/site-packages/gpytorch/utils/cholesky.py:44: NumericalWarning: A not p.d., added jitter of 1.0e-06 to the diagonal\n",
|
||||
" warnings.warn(f\"A not p.d., added jitter of {jitter_new:.1e} to the diagonal\", NumericalWarning)\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"nvol = 100\n",
|
||||
"npx = 10\n",
|
||||
"px_samples = torch.zeros(npx*nvol, len(edays))\n",
|
||||
"px_paths = torch.zeros(npx*nvol, fulltest)\n",
|
||||
"vol_paths = torch.zeros(nvol, fulltest)\n",
|
||||
"dmod.vol_model.eval();\n",
|
||||
"dmod.eval();\n",
|
||||
"\n",
|
||||
"mat_samples = torch.zeros(npx*nvol, len(edays))\n",
|
||||
"mat_paths = torch.zeros(npx*nvol, fulltest)\n",
|
||||
"mat_mod.eval();\n",
|
||||
"mat_lh.eval();\n",
|
||||
"\n",
|
||||
"sm_samples = torch.zeros(npx*nvol, len(edays))\n",
|
||||
"sm_paths = torch.zeros(npx*nvol, fulltest)\n",
|
||||
"sm_mod.eval();\n",
|
||||
"sm_lh.eval();\n",
|
||||
"\n",
|
||||
"mat_pred = mat_mod(test_x).sample(torch.Size((nvol*npx,))).exp()\n",
|
||||
"mat_paths = mat_pred.detach()\n",
|
||||
"mat_samples = mat_pred[:, ntests-1].detach()\n",
|
||||
"\n",
|
||||
"sm_pred = sm_lh(sm_mod(test_x)).sample(torch.Size((nvol*npx,))).exp()\n",
|
||||
"sm_paths = sm_pred.detach()\n",
|
||||
"sm_samples = sm_pred[:, ntests-1].detach()\n",
|
||||
"\n",
|
||||
"for vidx in range(nvol):\n",
|
||||
"# print(vidx)\n",
|
||||
" vol_pred = dmod.vol_model(test_x).sample().exp()\n",
|
||||
" vol_paths[vidx, :] = vol_pred.detach()\n",
|
||||
" \n",
|
||||
" px_pred = dmod.GeneratePrediction(test_x, vol_pred, npx).exp()\n",
|
||||
" px_paths[vidx*npx:(vidx*npx + npx), :] = px_pred.detach().T\n",
|
||||
" px_samples[vidx*npx:(vidx*npx+npx), :] = px_pred[ntests-1].detach().T\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 50,
|
||||
"id": "7a295328",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"plt_x = train_x[:]\n",
|
||||
"plt_y = train_y[:]\n",
|
||||
"plt_vol = vol[:]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 107,
|
||||
"id": "11bcade9",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def SimPlotter(sims, clr_idx, label, ax):\n",
|
||||
"# ax.plot(plt_x, plt_y, c=colors[0], label=\"Train\")\n",
|
||||
"# ax.plot(test_x, sims[1:20, :].T, c=colors[clr_idx], alpha=0.5,\n",
|
||||
"# lw=0.5)\n",
|
||||
"# ax.plot(test_x, test_y, c=colors[0], lw=2., label=\"Test\", alpha=0.75)\n",
|
||||
"# ax.plot(test_x, sims[0, :].T, c=colors[clr_idx], alpha=0.5,\n",
|
||||
"# lw=1., label=label)\n",
|
||||
"\n",
|
||||
" ax.plot(plt_x, plt_y, c='k', label=\"Train\")\n",
|
||||
" ax.plot(test_x, sims[1:15, :].T, c=colors[clr_idx], alpha=0.5,\n",
|
||||
" lw=0.5)\n",
|
||||
" ax.plot(test_x, test_y, c='k', lw=2., label=\"Test\", alpha=0.75, ls=\"--\")\n",
|
||||
" ax.plot(test_x, sims[0, :].T, c=colors[clr_idx], alpha=0.5,\n",
|
||||
" lw=1., label=label)\n",
|
||||
"\n",
|
||||
" \n",
|
||||
" sns.rugplot(x=test_x[ntests-1], ax=ax, color=\"#D72638\", height=0.1, label=\"Expirations\", lw=1.5)\n",
|
||||
" # ax.legend(fontsize=fs-2, frameon=False)\n",
|
||||
" ax.set_xlabel(\"Years\")\n",
|
||||
" ax.set_ylabel(\"Price\")\n",
|
||||
" ax.set_title(label + \" Simulations\")\n",
|
||||
" ax.set_ylim(100, 275)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 108,
|
||||
"id": "ea88c55b",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAABMUAAAMbCAYAAABNPSKwAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdd5gc1ZXw/2+FjpOTRjlnkIiSQCCSTLYxlsEBg9cBvwYTdjH2GocX+3UOP9us8a7DLsk2a7ANmJyDEEESQhHlPNJIk6e7p2PF3x81U5rW5NEowfk8zz473V1163bP4L46de45iuu6LkIIIYQQQgghhBBCfICoR3sCQgghhBBCCCGEEEIcaRIUE0IIIYQQQgghhBAfOBIUE0IIIYQQQgghhBAfOBIUE0IIIYQQQgghhBAfOBIUE0IIIYQQQgghhBAfOBIUE0IIIYQQQgghhBAfOBIUE0IIIYQQQgghhBAfOBIUE0IIIYQQQgghhBAfOBIUE0IcdaZpHu0piMNMfsdCCCGEEEeerMGE6J1+tCcghBicW2+9leeffx6AM844gwceeGBQ4zz00EN897vfBWDmzJk89thjQzbHCy64gNraWgA2b97c7TGrVq3izjvv5Mknnxyy665YsYJXXnmFZcuW0dDQQGtrK9FolLKyMiZMmMDZZ5/NwoULGTFiRK/j3HHHHf7n8ac//Yl58+YN2RyPlr1797Jw4UIA5s6dy5///OfDer1YLMYvfvEL5syZw5VXXtnl9euuu47ly5cD8PLLLzN69OjDOh8hhBDvLxs2bOCll17irbfeoq6ujpaWFgKBABUVFYwaNYqzzjqLhQsXMmHChEFf46Mf/SibNm0CoKioiNdff51oNNrv86dNm9bnMYFAgIKCAsaMGcOpp57KokWLmD59eq/ndF6n9EXXdUKhEOXl5UycOJG5c+fyyU9+kqKion6d3187d+7kxRdf5I033mDv3r20tLSgKArl5eUMHz6cM888kwsuuICZM2f2Os6jjz7KN7/5TQBuvvlmbrnlliGd59HSn7XxULEsiwceeIDa2lruvPPOLq/ffffd/Pa3vwXgJz/5CYsWLTqs8xHiWCVBMSGOU1dddZUfFFu+fDn19fVUV1cPeJx//vOfeWMeST/96U+5//77cV13SMbbvn07P/jBD3j77be7vBaPx4nH4+zatYtXX32Vn/70p1x33XXcdNNNFBYWDsn1Rb63336br371q7S0tHDaaacd7ekIIYR4H6mrq+MnP/kJzz33XJfXcrkcyWSS3bt389Zbb/HLX/6SK6+8kttvv53KysoBXee9997zA2IAbW1tPP3001x99dWH/B46M02TWCxGLBZj3bp1/OUvf+Gaa67hm9/8JpqmHfL4lmVhWRapVIo9e/awePFi7rnnHn74wx/6N8sORTwe55e//CX/+Mc/sG27y+vpdJq9e/eyYsUK7r77bhYuXMgdd9zB2LFjD/naoqu6ujq+/OUvs2nTJj72sY8d7ekIcUyToJgQx6mzzz6bESNGsH//fhzH4emnn+YLX/jCgMbYuXMnq1atAiAcDnPFFVccjqn26IUXXhiygNjWrVu59tpricViABQUFHDmmWcyceJEioqKMAyD+vp6li9fzq5duzBNk3vvvZdVq1Zx7733DuiOr+ifFStW0NLScrSnIYQQ4n2mvr6ez3zmM+zduxeAUCjEvHnzmDx5MqWlpViWRVNTE6tWrWLjxo04jsOjjz7KypUrefDBBwcUGHvkkUcACAaDqKpKNpvlr3/966CDYjfccAPFxcV5z9m2TTabpampiZUrV7J161Zs2+bPf/4zjuN0m+VzsLPOOouzzjqrx9cty6K5uZnVq1ezZs0aAFpaWrjtttv405/+xMknnzyo9wOQTCb5whe+wHvvvQd4WWmnnXYa06dPp6ysDEVRaG5uZt26daxevRrXdXn55ZdZvXo1f/nLX5g4ceKgry26t3v37rxgrhCiZxIUE+I4paoqH/vYx/iv//ovAJ544okBB8U6Z4ldfPHFQ55Cf6TkcjluuOEGPyD26U9/mq9//esUFBR0e/zLL7/MN7/5TeLxOKtWreJb3/oWd911V5fjfvrTn/LTn/70MM5cHO7tm0IIId5fXNflX//1X/2A2IUXXsj3v/99ysvLuz1+xYoV/Pu//zu1tbXs2rWLm266iYceeghFUfq8Vi6X4+mnnwZg9uzZFBcX88orr7B+/XrWrl3L7NmzBzz/q6++us9SAZ1LWzz44INceeWVfV7rlFNO4Ytf/GK/5rBixQpuvfVWmpubyeVy/PjHP+Zvf/tb/95AN+68804/IHbaaafx85//vMf3uGXLFr72ta+xefNmmpubuf7663n66aeJRCJ5xy1atEi28x1mt9xyy/tmW6oQh0IK7QtxHFu0aJG/qNu4cSPbtm3r97mu6+bV8RrqbQBH0pNPPukvjs8//3y+973v9RgQA1i4cCH/+Z//6X92zz77rNxNE0IIIY4DS5cu9bPcZ8yYwV133dVjQAzg9NNP59577yUcDgOwevVqXn311X5d68UXXyQej/vjXHLJJf5rDz300GDfQp8+9alP5V3rr3/965COf/rpp/PLX/7Sf7xmzRrWr18/qLFqamp45plnABg2bBh//OMfew36TZ06lfvvv9/P1qutreUf//jHoK4thBBDQYJiQhzHxowZwxlnnOE/Hkix+mXLlvmFPsePH8+cOXOGfH5HSucaYt0Vc+/OnDlzOPvss/3HixcvHuppCSGEEGKIdf7O/8hHPoKu973xZfz48Xnrg9dee61f13r00Uf9ny+44AIWLlzoB9eeeeYZ2tra+jfpQfjQhz7k/3w4CrKfeeaZjBo16pCv8fbbb/ulMC688MJ+1WktLy/nX/7lX/zH/f19CCHE4SBBMSGOcx//+Mf9n5988sl+1+jqvHWy8xidJZNJHnjgAT7/+c9z1llnceKJJzJv3jw+/vGP86tf/coPqg3UtGnTmDZtWt75Hc/1p0PTwTq2TQJkMpl+n3fWWWcRDAYZNmxYt+2q77jjDn9Oy5Yty3tt2bJl/msdRX7XrFnDN77xDRYuXMisWbM4++yz+fznP9+lCHBLSwt33303H/nIRzjllFM49dRTufrqq/nTn/6EZVndzrW3uRzs7rvv9o/tvKAfqLfffpsf/OAHXHnllcyfP58TTzyR0047jYULF/Kv//qvPPXUU90W0+2Ya0dHI4BvfvOb3c7puuuu85/vyPbrzqH+LXZcp+POu2EYPPjgg1x77bXMnz+fWbNmcd5553H77bezdOnSPj+b2tpafv3rX3P11Vdz+umnc+KJJzJ//nyuvvpqfvWrX7F79+4+xxBCCDFwh/KdHwgEqKysxHGcPo/ft2+fH4AbMWIEs2fPprCw0A9WZTKZIe3YfbCD644dDp1rqzU2Ng5qjEP5fei6Tnl5Oara9Z+kjz76qL8+uPvuu7u83vHa97//fcD7ff3iF7/gsssu45RTTmHu3LksWrSI+++/n2w2659nWRYPP/wwn/70p5k3bx6zZ8/mkksu4ec//zmtra3dzrWvuXTWeX14xx139PvzONiOHTu46667+MxnPsOCBQuYPXs2J510EgsWLOALX/gC9913H8lksse5fvazn/Wfe+yxx7qdU3/Xi47j8Nxzz3Hrrbdy/vnnM3v2bE499VQuueQS7rzzTj9zsyedr9Pxd/byyy9z8803c9555/lrqC9+8Ys89thjff73mUwm+fOf/8znP/95zjzzTE444QTmzp3L5Zdfzne/+12/q7kQ/SU1xYQ4zl100UUUFxeTSCSora3l3Xff5fTTT+/1nHQ67Xeu1HW925oNL730Et/5zne6LBA6OiO999573HfffXzlK1/hxhtvHLo3NAidu27+7W9/6/ed48997nN8/vOfH5I5/PrXv+aPf/xj3hd5Y2MjjY2NvPXWW3z605/me9/7HqtXr+aWW26hoaEh7/y1a9eydu1aFi9ezB//+Mch6TQ1WI2Njfzbv/0bK1as6PKaaZokk0n27t3Lc889x7333ssf//jHAXfzGoih/lvcs2cPX/nKV9iyZUve8/v37+epp57iqaee4lOf+hTf+973uq0588gjj/C9730PwzDynm9ubqa5uZm1a9dyzz33cMMNN0itDiGEGGKdv/OfeOIJPve5z/UrO+nCCy/06171x6OPPup/p1922WX+98HHPvYxnnrqKQAefvjhvODDUNqxY4f/8/Tp04d8fMdx2LNnj/+4oqJiUON0/n288sor1NXVMXz48D7PO+GEE3jvvff6VdutLy+++CJ33HFHlyBRPB5n/fr1PPvss/zP//wPuVyOm2++uUsQZ+fOndxzzz0899xz/PWvfx1UN/ehYpomP/jBD/j73//ebXAom83S0NDAm2++yR//+Ef+67/+i1NOOeWwzWfbtm3cfvvtXcqM5HI5du7cyc6dO3n44Ye59NJL+dGPftRr+ZKO82699Vb/3yEdmpubeeONN3jjjTf461//yj333NNtreP33nuPG2+8scs6uqPL/LZt23jooYe44IIL+NWvftWlVp0Q3ZGgmBDHuVAoxEc+8hEefPBBwMsW6yso9sILL5BOpwE477zzugQ0nnrqKb72ta/5WWdVVVVccMEFjBw5klgsxpIlS9i2bRuGYXDXXXexf/9+/05df/z7v/87AH/4wx/8Wh0dzw3Gueee63eHWrlyJZ///Oe56aabmDdvXq+LraFYiAHcf//9/gJrzpw5nHrqqViWxdKlS/0aHX/961+ZMGEC//Vf/0UsFmPGjBksWLCAUCjE8uXL/eyvN954gwcffPCwLbL7kk6nueaaa6ipqQG8O9Xnnnsu48ePJxgM0tDQwFtvvcX27dsBWL9+Pd/+9rf5wx/+4I9x2WWXMWXKFN58803efPNN/7kTTzwRgFmzZvV7PkP9t5hOp7n++uvZtWsXxcXLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 1400x800 with 4 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"colors = [\"#1b4079\",\"#c6ddf0\",\"#50723C\",\"#b9e28c\",\"#8c2155\",\"#af7595\",\"#e6480f\",\"#fa9500\", \"#808080\"]\n",
|
||||
"colors = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"fs = 16\n",
|
||||
"\n",
|
||||
"sns.set(font_scale=2.)\n",
|
||||
"sns.set_style(\"white\")\n",
|
||||
"\n",
|
||||
"fig, ax = plt.subplots(2,2, figsize=(14, 8), dpi=100)\n",
|
||||
"plt.subplots_adjust(hspace=0.6)\n",
|
||||
"vol_scale = 1.\n",
|
||||
"\n",
|
||||
"## VOLTRON PLOT ##\n",
|
||||
"SimPlotter(px_paths, 6, \"Volt\", ax[0,0])\n",
|
||||
"SimPlotter(sabr_paths.T, 4, \"SABR\", ax[0,1])\n",
|
||||
"SimPlotter(sm_paths, 2, \"SM\", ax[1,0])\n",
|
||||
"SimPlotter(mat_paths, 0, \"Matern\", ax[1,1])\n",
|
||||
"lw = 3.\n",
|
||||
"\n",
|
||||
"custom_lines = [Line2D([0], [0], color='k', lw=lw),\n",
|
||||
" Line2D([0], [0], color=colors[6], lw=lw),\n",
|
||||
" Line2D([0], [0], color=colors[4], lw=lw),\n",
|
||||
" Line2D([0], [0], color=colors[2], lw=lw),\n",
|
||||
" Line2D([0], [0], color=colors[0], lw=lw),\n",
|
||||
" Line2D([0], [0], color=\"#D72638\", lw=lw),]\n",
|
||||
"\n",
|
||||
"# fig.legend(custom_lines, [\"Train\", \"Test\", \"Voltron\", \"SM\", \"Matern\", \"Expirations\"],\n",
|
||||
"# ncol=1, frameon=False, bbox_to_anchor=(1.1, 0.65), fontsize=fs+1)\n",
|
||||
"\n",
|
||||
"fig.legend(custom_lines, [\"Prices\", \"Volt\", \"SABR\", \"SM\", \"Matern\", \"Expirations\"],\n",
|
||||
" ncol=6, bbox_to_anchor=(0.92, 0.02), fontsize=fs+2)\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"sns.despine()\n",
|
||||
"plt.savefig(\"./full_diffusions.pdf\", bbox_inches=\"tight\")\n",
|
||||
"plt.show()\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 81,
|
||||
"id": "0c928d57",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"logger = []\n",
|
||||
"days = np.datetime_as_string(edays, 'D')\n",
|
||||
"for day in range(px_samples.shape[1]):\n",
|
||||
" for smpl in range(px_samples.shape[0]):\n",
|
||||
" logger.append([px_samples[smpl, day].item(), days[day], \"Volt\"])\n",
|
||||
" \n",
|
||||
" for smpl in range(mat_samples.shape[0]):\n",
|
||||
" logger.append([mat_samples[smpl, day].item(), days[day][5:], \"Matern\"])\n",
|
||||
" \n",
|
||||
" for smpl in range(sm_samples.shape[0]):\n",
|
||||
" logger.append([sm_samples[smpl, day].item(), days[day][5:], \"SM\"])\n",
|
||||
" \n",
|
||||
" for smpl in range(sabr_samples.shape[0]):\n",
|
||||
" logger.append([sabr_samples.t()[smpl, day].item(), days[day], \"SABR\"])\n",
|
||||
"df = pd.DataFrame(logger)\n",
|
||||
"df.columns = ['Price', 'Date', \"Type\"]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 100,
|
||||
"id": "f8754ffd",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAABBoAAAJqCAYAAACFLiOkAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdd3hUZfbA8e/U9EIHEQUBUbBgQ1xBLKgsFlDsKGtbV1f9gWVtu6i79oKua8G2FuyLShXBBgiINAvSu5AQSEhPJjNz2++PO3MzQ3oyyRTO53l8nMzcuffNzU2Y99zznmMzDMNACCGEEEIIIYQQIgLs0R6AEEIIIYQQQgghEocEGoQQQgghhBBCCBExEmgQQgghhBBCCCFExEigQQghhBBCCCGEEBEjgQYhhBBCCCGEEEJEjAQahBBCCCGEEEIIETESaBBCCCGEEEIIIUTESKBBCCGEEEIIIYQQESOBBiGEEKIFDMNA07RoD0OINiPXPCiKEu0hCCFETHNGewBCCCGa7/PPP+f+++9vcDuHw0FSUhLt27enT58+nHzyyYwaNYoOHTq0wShbxzXXXMPy5csB+Pbbbzn44IOt15YtW8a4ceMAuOiii3jyySdbZQxbtmzhoYce4qmnngo7fjSFXhO33XYbt99+e5P3EXr+6mO320lOTiYrK4tevXpx0kknccEFF9CjR48mH7O+MbTmz7At9evXr97XbTYbLpeLtLQ0OnfuzJFHHsmIESMYMmQILper3vfG4zUfPB/du3fnu+++C3utvt/vaPv8889ZsGAB//nPf2p9raW/f0IIkQgko0EIIQ4Amqbh8XjIyclhwYIFPPXUUwwfPpwPPvgg2kOLW2+//TajR49m5cqV0R5K1Oi6jsfjIS8vjx9++IEXXniBP/7xj/z73/+O9tDikmEY+P1+iouL2bhxI9OnT+fmm29mxIgRLF26NNrDO+Cv+crKSq655hruv/9+iouLoz0cIYSIaZLRIIQQCaJHjx5ceeWVtb6m6zpVVVXk5eXx1VdfUVFRgcfj4V//+hd2u73O94m6fffddwdE+nRWVhZ/+ctfan3NMAy8Xi/5+fl8++237Nu3D0VRmDx5Mqqqcvfdd7fxaOPHzTffTGZmZthzuq7j8/koLCxk8+bN/PLLLyiKQk5ODtdddx2PPfYYY8aMidKID5xrvi7FxcVWloUQQoj6SaBBCCESRLdu3bjhhhsa3O6+++7j1ltvZcWKFQA8/fTTDB8+nE6dOrX2ENvMySefzMaNG6M9jISQnp7eqOvq/vvv54EHHmDOnDkAvPXWW5x//vkcccQRzTpuov8ML7300gaXA+Tm5vLQQw+xaNEiDMPgwQcfpEOHDpx++uk1to3H8xVv422Miy++mIsvvjjawxBCiKiTpRNCCHGAycrK4uWXXyY9PR0Aj8fDjBkzojwqEe9SUlJ4+umnrfoMmqbxySefRHlU8a179+68/vrrnHPOOQCoqso//vEPqqqqojwyIYQQon4SaBBCiANQVlYW559/vvV1LKz/FvHP5XJx2WWXWV//8MMPURxNYrDb7TzxxBN069YNgIKCAt57770oj0oIIYSonyydEEKIA1Tv3r2tx3v27Al77cUXX+Sll14CYPHixfh8Pp588kl++OEHXC4XPXv25OKLL+byyy8Pe5+iKMyYMYNvvvmGdevWUVxcTGpqKt27d2fIkCFcddVVdO3atcGx6brOnDlzmDZtGmvXrqWiooJOnTpxyimncO2113L44YfX+/6mVOBfsWIFn3/+OT/99BN5eXkAdOrUieOPP57LLruMk046KWz7M888k9zc3LDnzjrrLOtxXRXyf/31Vz7//HOWLVtGfn4+mqbRsWNHTjjhBC688EKGDBlS/0kJ2LhxIx988AFLly4lLy+PtLQ0Dj/8cC677DIuuOCCRu2jNYVeV3v37g17LbQi/yeffMKhhx7K008/zbfffoumaRxyyCGcc8453HLLLa36M9yfYRjMmzePL7/8ktWrV1NYWIjb7aZbt24MHjyYK664Iuz7amvp6en85S9/4eGHHwbg/fff56abbgrbpjHny+/3M2vWLL7++mvWrFlDSUkJSUlJdOzYkeOOO45zzjmHM888s8b7GnvN5+TkWM/fdddd3HDDDbz22mt88sknlJSU0LVrV04++WTuvvtuMjMz6+06sT+Px8O7777Ll19+ya5du3A6nfTo0YMzzzyTq666ivbt29f6vqZcR6HjHzRokBXQqa0Ly/Lly63xh27blK4TS5YsYfr06fzyyy8UFBQA0LFjR44//njOP/98TjvttDrfu//v0sCBA1mxYgX/+9//WLVqFfv27SMlJYU+ffowYsQILr/8ctxud537a+61IYQQdZFAgxBCHKAcDketj/dXWlrKTTfdFDbR+OWXXzj++OPDtlu/fj0TJkxgx44dYc/7/X5KSkpYu3Yt77zzDnfeeSfXXnttvccLrSERtHv3bj777DNmzpzJxIkTG/Ed1q+4uJj777+f+fPn13gtJyeHnJwcZs6cycUXX8y//vWvBtsL1sXn8zFx4sRal6cEjzNjxgyGDh3KpEmTyMrKqnNf//3vf3n22WfRdd16rqSkhOXLl7N8+XJmzpzJGWec0axxRkpjrytFUbjxxhtZs2aN9dy6deua1BozEj/D3Nxcxo8fz2+//Rb2vM/no7y8nE2bNvHBBx9w3XXXcdddd2G3RycZ9LzzzuORRx5B0zT27t3L1q1bmxT82LVrFzfddBPbtm0Le15RFCoqKtixYwfTpk1j4MCBTJ48uc6Je1M88cQTYdkXO3bsoLy8nIceeqhJ+9m1axfXXnstu3btCnt+7dq1rF27lilTpvDoo49aS0xi3Z49e7jnnntYtmxZjdd27drFrl27mDFjBieffDLPPfccHTt2rHd/hmHwyCOP8P7774c97/P5WLlyJStXruT999/nnXfesTJj9j9mW18bQojEJ4EGIYQ4QG3YsMF6fMghh9S53ZNPPlnjbibAiBEjrMerV6/m2muvpbKyEoDOnTtzxhlncNBBB1FRUcHKlSv5+eef8fl8PPHEE5SWljJ+/Pga+/R4PIwdO5bNmzcD4Ha7Oeusszj88MMpLy9n/vz5bN++nYcffpiMjIxmf++VlZWMGzeOTZs2AWCz2TjllFM49thjMQyD3377jSVLlgDmnUMwJ01gdgsoLy/no48+siY+oR0EsrOzreP4/X6uu+46Vq1aBZhLC4YOHUr//v2x2Wxs3bqVBQsW4PF4WLRoEWPHjuXjjz+26meEeumll3jxxRetr4866ihOPfVU3G43v/32G99//z3ff/89q1evbvZ5iYTGXlevvfZaWJAhKPS6qk9LfoZBu3bt4sorr7TuJmdnZ3PGGWdw6KGH4vV6Wb16NUuXLkXTNN58800KCgp4+umnGzW+SMvMzOSII45g7dq1gLncqbGBBr/fz80332xNJLt168bpp59Ot27d8Hg8bNq0iQULFqDrOr/88gu33XYbH374ofX+plzzQcuXL2fRokU1nj/77LPrDUDVZvz48ZSWlpKSksLw4cM57LDDKCoq4quvvmLv3r2UlpZyxx138MorrzBs2LAm7bsxDjnkEO655x7Kysp49dVXgfAuP7VN3utSUFDAVVddZf1NdTqdYX8T1q5dy6JFi1BVlWXLlnHZZZfxv//9r95gw/PPP8+yZcuw2WwMHjyY4447DrvdzurVq61Cojt27OCOO+7g448/DntvS68NIYSoiwQahBDiALRv3z6rOwBQb4ruokWL6NSpE//85z8ZPHiw9QH/2GOPBaCiooIJEyZYQYYbbriBCRMm1EjTXbhwIXfddRfl5eVMnjyZQYMGccopp4Rt8+qrr1pBhu7du/PGG2+ETabuvvtuXnrpJV555RVKS0ub/f1PmjTJmqB26tSJF198keOOOy5sm++++47bb78dVVX5/PPPueiiixg0aJBVg2DBggXWpKuuDgKTJk2yggz9+/fnhRdeqDH5zs/P56677mL58uVs3ryZRx55hKeeeipsmy1btjB58mTAXLM/ceJErrrqqrBtfv75Z2655RaKi4ube1pazOv1MnXqVOvrhq6r1NRUHnroIYYPH05lZSVffvllrR0VatOSnyGYxSrvuOMOK8hwwQUX8PDDD9cI8qxevZrbbruNvXv3WneZo9VislevXlag4ffff2/0++bNm8eWLVsAM83/zTffJCkpKWyb1atX86c//QmPx8OqVatYsWKFteSkKdd8UDDIcNNNNzFu3Djcbjc//PBDvcGnupSWlnL44YczefLksGPeddddVqcTVVX5+9//zty5c2sN1LVEsKNPTk6OFWhobJef/d15551WkKFnz568/PLL9OnTJ2ybjRs3cuutt7Jr1y5yc3O56667ePfdd+vc57Jly+r8Hfj++++55ZZbUFWVn3/+mZ9++iksG62l14YQQtRFikEKIcQB5vfff+emm26ioqICMLMPRo0aVe97XnzxRc466yzS0tLo0aNH2Afsjz76yPrgfPHFF3PPPffUuhZ42LBhPProo4CZ6ht6dx7MJQBvv/02YN75nzx5co07tg6Hg/Hjx7eofVxpaSn/+9//rP29/PLLNT6cg7ku/a9//av1dVM7KOzdu5cPPvgAgPbt2/Pf//631klW586dmTx5stVedObMmTWWn/znP/9BVVXADOTsH2QAOO644/jPf/6DzWZr0jgjpaCggFtvvZWcnBzA7EJx9dVX1/ueRx55hNGjR5Oenk6XLl249tprSU1NbfBYkfgZzps3z1ouMXjwYJ5++ulaJ6jHHHMML730knVeX375ZTRNa3CMrSH0rnZTAkq//vqr9fhPf/pTjYkkmN/n9ddfD2DdDW+pyy+/nLvuuotOnTqRlZXFH//4RwYMGNDk/WRlZfHmm2/WCGykpKTw7LPPWvssKChg2rRpLR53a1m8eDHLly8Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 1200x500 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fs = 16\n",
|
||||
"fig, ax = plt.subplots(figsize=(12,5))\n",
|
||||
"# violins = sns.violinplot(x='Date', y='Price', data=df[df.Type==\"Voltron\"])\n",
|
||||
"violins = sns.violinplot(x='Date', y='Price', data=df[df.Type.isin([\"Volt\", \"SABR\"])],\n",
|
||||
" split=False, hue=\"Type\", c=[colors[6], colors[2]])\n",
|
||||
"# violins = sns.violinplot(x='Date', y='Price', data=df[df.Type == \"Matern\"], hue=\"Type\")\n",
|
||||
"# violins = sns.violinplot(x='Date', y='Price', data=df, hue=\"Type\")\n",
|
||||
"for violin in violins.collections[::4]:\n",
|
||||
" violin.set_alpha(0.5)\n",
|
||||
" violin.set_edgecolor(colors[6])\n",
|
||||
" violin.set_facecolor(colors[7])\n",
|
||||
" \n",
|
||||
"for violin in violins.collections[2::4]:\n",
|
||||
" violin.set_alpha(0.5)\n",
|
||||
" violin.set_edgecolor(colors[4])\n",
|
||||
" violin.set_facecolor(colors[5])\n",
|
||||
" \n",
|
||||
" \n",
|
||||
"plt.scatter(np.arange(px_samples.shape[1]), test_y[ntests-1], color=colors[1], zorder=4, \n",
|
||||
" s=160,\n",
|
||||
" edgecolor='k',lw=2.,\n",
|
||||
" label=\"Observed Price\")\n",
|
||||
"plt.plot(np.arange(px_samples.shape[1]), test_y[ntests-1], \n",
|
||||
" color=colors[0], zorder=3, lw=2.)\n",
|
||||
"plt.xlabel(\"Expirations\")\n",
|
||||
"plt.xticks(rotation=45)\n",
|
||||
"sns.despine()\n",
|
||||
"ax.get_legend().set_visible(False)\n",
|
||||
"\n",
|
||||
"legend_elements = [Patch(facecolor=colors[7], edgecolor=colors[6], alpha=0.75,\n",
|
||||
" label='Volt'),\n",
|
||||
" Patch(facecolor=colors[5], edgecolor=colors[4],\n",
|
||||
" label='SABR'),\n",
|
||||
" Line2D([0], [0], marker='o', color=colors[0], label='SPY Value',\n",
|
||||
" markerfacecolor=colors[1], markersize=15, lw=2.)]\n",
|
||||
"\n",
|
||||
"fig.legend(handles = legend_elements, fontsize=fs+2,\n",
|
||||
" bbox_to_anchor=(0.35, 0.95), frameon=False)\n",
|
||||
"\n",
|
||||
"plt.title(\"Predicted Price Distributions\")\n",
|
||||
"# plt.ylim(0, 500)\n",
|
||||
"plt.savefig(\"./voltron-sabr-distribution-plot.pdf\", bbox_inches='tight')\n",
|
||||
"\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 85,
|
||||
"id": "7c15a641",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAA5YAAAJqCAYAAABQLN+oAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdd3gU1f4G8He2pDcgoSMBQrdQpCggKKiIKEVREeRasYAXrNdyUfxZsVz1WrBwLVhAUboKihRp0hHpCQRI75tks9k2c35/bDLZTXY3m2TDpryf5/FxszszezKZhHn3fM85khBCgIiIiIiIiKiWNIFuABERERERETVuDJZERERERERUJwyWREREREREVCcMlkRERERERFQnDJZERERERERUJwyWREREREREVCcMlkRERERERFQnDJZERERERERUJwyWREQEABBCQJblQDeD6LzhNQ/YbLZAN4GImghdoBtARAQAy5cvx9NPP13tdlqtFsHBwWjZsiUSEhIwZMgQTJgwAa1atToPrawfd9xxB3bv3g0A+P3339GxY0f1tV27dmHGjBkAgEmTJuG1116rlzYkJSXh+eefx4IFC1zeP5Ccr4nZs2fj4YcfrvExnM+fNxqNBiEhIYiOjkaXLl0waNAg3HDDDejUqVON39NbG+rzZ3g+9ezZ0+vrkiRBr9cjPDwcrVu3Ru/evTF27FgMHz4cer3e676N8ZovPx8dOnTAxo0bXV7z9vsdaMuXL8fmzZvx3//+1+1rdf39I6LmhT2WRNSoyLIMk8mE1NRUbN68GQsWLMCYMWPwzTffBLppjdbnn3+OiRMnYu/evYFuSsAoigKTyYSMjAzs2LED7777Lq677jq88847gW5aoySEgNVqRUFBAU6cOIGVK1figQcewNixY7Fz585AN6/ZX/MlJSW444478PTTT6OgoCDQzSGiJoI9lkTU4HTq1AlTp051+5qiKCgtLUVGRgZ+/fVXGI1GmEwm/N///R80Go3H/cizjRs3NotyuOjoaNx///1uXxNCwGw2Izs7G7///jtyc3Nhs9mwcOFC2O12PP744+e5tY3HAw88gKioKJfnFEWBxWJBXl4eEhMTcfDgQdhsNqSmpuKuu+7Cyy+/jJtuuilALW4+17wnBQUFai8qEZG/MFgSUYPTrl073HPPPdVu99RTT2HWrFnYs2cPAOD111/HmDFjEBcXV99NPG+GDBmCEydOBLoZTUJERIRP19XTTz+NZ555Bj///DMA4LPPPsP48ePRq1evWr1vU/8ZTpkypdryzrS0NDz//PPYunUrhBB47rnn0KpVK4waNarKto3xfDW29vpi8uTJmDx5cqCbQUSNCEthiajRio6OxgcffICIiAgAgMlkwqpVqwLcKmrsQkND8frrr6vjK2VZxnfffRfgVjVuHTp0wCeffIJrrrkGAGC32/Hvf/8bpaWlAW4ZERH5C4MlETVq0dHRGD9+vPp1Qxi/RY2fXq/HLbfcon69Y8eOALamadBoNHj11VfRrl07AEBOTg6++uqrALeKiIj8haWwRNTodevWTX2cmZnp8tp7772H999/HwCwbds2WCwWvPbaa9ixYwf0ej3i4+MxefJk3HrrrS772Ww2rFq1Chs2bMDRo0dRUFCAsLAwdOjQAcOHD8ftt9+Otm3bVts2RVHw888/Y8WKFThy5AiMRiPi4uJw2WWX4c4770SPHj287l+TGTL37NmD5cuXY//+/cjIyAAAxMXFYcCAAbjlllswaNAgl+2vuuoqpKWluTw3evRo9bGnGSz/+usvLF++HLt27UJ2djZkWUZsbCwGDhyIG2+8EcOHD/d+UsqcOHEC33zzDXbu3ImMjAyEh4ejR48euOWWW3DDDTf4dIz65HxdZWVlubzmPGPmd999h86dO+P111/H77//DlmWccEFF+Caa67Bgw8+WK8/w8qEEFi/fj1++eUXHDp0CHl5eQgKCkK7du0wdOhQ3HbbbS7f1/kWERGB+++/H/PnzwcAfP3115g5c6bLNr6cL6vVijVr1uC3337D4cOHYTAYEBwcjNjYWPTv3x/XXHMNrrrqqir7+XrNp6amqs8/9thjuOeee/Dxxx/ju+++g8FgQNu2bTFkyBA8/vjjiIqK8jorbGUmkwlffvklfvnlF6SkpECn06FTp0646qqrcPvtt6Nly5Zu96vJdeTc/sGDB6sB3t0sybt371bb77xtTWaF3b59O1auXImDBw8iJycHABAbG4sBAwZg/PjxuOKKKzzuW/l3qV+/ftizZw++//577Nu3D7m5uQgNDUVCQgLGjh2LW2+9FUFBQR6PV9trg4jqjsGSiBo9rVbr9nFlhYWFmDlzpsuN5cGDBzFgwACX7Y4dO4a5c+fizJkzLs9brVYYDAYcOXIEX3zxBR599FHceeedXt/PeQxoufT0dPz4449YvXo15s2b58N36F1BQQGefvppbNq0qcprqampSE1NxerVqzF58mT83//9X7XLPXhisVgwb948t+XG5e+zatUqjBgxAm+99Raio6M9Hut///sf3nzzTSiKoj5nMBiwe/du7N69G6tXr8aVV15Zq3b6i6/Xlc1mw7333ovDhw+rzx09erRGS5X442eYlpaGOXPm4O+//3Z53mKxoLi4GCdPnsQ333yDu+66C4899hg0msAULV1//fV48cUXIcsysrKycOrUqRqF3ZSUFMycOROnT592ed5ms8FoNOLMmTNYsWIF+vXrh4ULF3oMajXx6quvuvSunjlzBsXFxXj++edrdJyUlBTceeedSElJcXn+yJEjOHLkCBYvXoyXXnpJLRlu6DIzM/Hkk09i165dVV5LSUlBSkoKVq1ahSFDhuA///kPYmNjvR5PCIEXX3wRX3/9tcvzFosFe/fuxd69e/H111/jiy++UHu+K7/n+b42iKgCgyURNXrHjx9XH19wwQUet3vttdeq9FYAwNixY9XHhw4dwp133omSkhIAQOvWrXHllVeiffv2MBqN2Lt3Lw4cOACLxYJXX30VhYWFmDNnTpVjmkwmTJs2DYmJiQCAoKAgjB49Gj169EBxcTE2bdqE5ORkzJ8/H5GRkbX+3ktKSjBjxgycPHkSgGP9wMsuuwyXXHIJhBD4+++/sX37dgCOngHAcZMMOGbzLC4uxpIlS9QbXecZPmNiYtT3sVqtuOuuu7Bv3z4AjlLRESNGoE+fPpAkCadOncLmzZthMpmwdetWTJs2DUuXLlXHvzp7//338d5776lfX3jhhRg2bBiCgoLw999/448//sAff/yBQ4cO1fq8+IOv19XHH3/sEirLOV9X3tTlZ1guJSUFU6dOVXuLYmJicOWVV6Jz584wm804dOgQdu7cCVmWsWjRIuTk5OD111/3qX3+FhUVhV69euHIkSMAHOXrvgZLq9WKBx54QA0O7dq1w6hRo9CuXTuYTCacPHkSmzdvhqIoOHjwIGbPno1vv/1W3b8m13y53bt3Y+vWrVWev/rqq71+4ODOnDlzUFhYiNDQUIwZMwZdu3ZFfn4+fv31V2RlZaGwsBCPPPIIPvzwQ4wcObJGx/bFBRdcgCeffBJFRUX46KOPALjOwu0urHmSk5OD22+/Xf2bqtPpXP4mHDlyBFu3boXdbseuXbtwyy234Pvvv/caLt9++23s2rULkiRh6NCh6N+/PzQaDQ4dOqRO/HTmzBk88sgjWLp0qcu+db02iKjuGCyJqFHLzc1VZ+8E4LXkauvWrYiLi8MLL7yAoUOHqjd0l1xyCQDAaDRi7ty5aqi85557MHfu3CplV1u2bMFjjz2G4uJiLFy4EIMHD8Zll13mss1HH32khsoOHTrg008/dbl5fvzxx/H+++/jww8/RGFhYa2//7feeksNJHFxcXjvvffQv39/l202btyIhx9+GHa7HcuXL8ekSZMwePBgdQzh5s2b1ZtsTzN8vvXWW2qo7NOnD959990qYSs7OxuPPfYYdu/ejcTERLz44otYsGCByzZJSUlYuHAhAMeYu3nz5uH222932ebAgQN48MEHA7q+ntlsxrJly9Svq7uuwsLC8Pzzz2PMmDEoKSnBL7/84nbGU3fq8jMEHJMLPfLII2qovOGGGzB//vwqof7QoUOYPXs2srKy1F6kQC350aVLFzVYnj171uf91q9fj6SkJACOss1FixYhODjYZZtDhw7hH//4B0wmE/bt24c9e/aoJcQ1uebLlYfKmTNnYsaMGQgKCsKOHTu8ftjgSWFhIXr06IGFCxe6vOdjjz2mzkRst9vx7LPPYt26dW4/mKmL8hm3U1NT1WDp6yzclT366KNqqIyPj8cHH3yAhIQEl21OnDiBWbNmISUlBWlpaXjsscfw5Zdfejzmrl27PP4O/PHHH3jwwQdht9tx4MAB7N+/36XapK7XBhHVHSfvIaJG6+zZs5g5cyaMRiMAR+/ihAkTvO7z3nvvYfTo0QgPD0enTp1cbqiWLFmi3ihNnjwZTz75pNuxPCNHjsRLL70EwFG65dz7BjhKOj///HMAjp69hQsXVumR0Wq1mDNnTp2m8y8sLMT333+vHu+DDz6ocjMGOMaVPfTQQ+rXNZ3hNCsrC9988w0AoGXLlvjf//7n9qa6devWWLhwobrcy+rVq6uUE//3v/+F3W4H4AjulUMlAPTv3x///e9/IUlSjdrpLzk5OZg1axZSU1MBOGaJnT59utd9XnzxRUycOBERERFo06YN7rzzToSFhVX7Xv74Ga5fv14tfx06dChef/11t4Hk4osvxvvvv6+e1w8++ACyLFfbxvrg3GtVkw8Q/vrrL/XxP/7xjyrBAXB8n3fffTcAqL1ddXXrrbfiscceQ1xcHKKjo3Hdddehb9++NT5OdHQ0Fi1aVCXIhoaG4s0331SPmZOTgxUrVtS53fVl27Zt6jqYUVFR+OKLL6qESgDo2bMnPv/8c/V6/PPPP9Xed0/eeOMNt78DV1xxBW6Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 1000x500 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fs = 16\n",
|
||||
"fig, ax = plt.subplots(figsize=(10,5))\n",
|
||||
"# violins = sns.violinplot(x='Date', y='Price', data=df[df.Type==\"Voltron\"])\n",
|
||||
"# violins = sns.violinplot(x='Date', y='Price', data=df[df.Type.isin([\"Voltron\", \"SM\"])],\n",
|
||||
" # split=True, hue=\"Type\", c=[colors[6], colors[2]])\n",
|
||||
"# violins = sns.violinplot(x='Date', y='Price', data=df[df.Type == \"Matern\"], hue=\"Type\")\n",
|
||||
"violins = sns.violinplot(x='Date', y='Price', data=df[df.Type == \"Volt\"])\n",
|
||||
"for violin in violins.collections[::2]:\n",
|
||||
" violin.set_alpha(0.5)\n",
|
||||
" violin.set_edgecolor(colors[6])\n",
|
||||
" violin.set_facecolor(colors[7])\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\n",
|
||||
" \n",
|
||||
" \n",
|
||||
"plt.scatter(np.arange(px_samples.shape[1]), test_y[ntests-1], color=colors[1], zorder=4, \n",
|
||||
" s=160,\n",
|
||||
" edgecolor='k',lw=2.,\n",
|
||||
" label=\"Observed Price\")\n",
|
||||
"plt.plot(np.arange(px_samples.shape[1]), test_y[ntests-1], \n",
|
||||
" color=colors[0], zorder=3, lw=2.)\n",
|
||||
"plt.xlabel(\"Expirations\")\n",
|
||||
"plt.xticks(rotation=45)\n",
|
||||
"sns.despine()\n",
|
||||
"\n",
|
||||
"# plt.legend(loc=\"upper left\",frameon=False, fontsize=fs)\n",
|
||||
"\n",
|
||||
"legend_elements = [Patch(facecolor=colors[7], edgecolor=colors[6], alpha=0.75,\n",
|
||||
" label='Voltron'),\n",
|
||||
" Line2D([0], [0], marker='o', color=colors[0], label='SPY Value',\n",
|
||||
" markerfacecolor=colors[1], markersize=15, lw=2.)]\n",
|
||||
"\n",
|
||||
"fig.legend(handles = legend_elements, fontsize=fs+2, bbox_to_anchor=(0.35, 0.85), frameon=False)\n",
|
||||
"\n",
|
||||
"plt.title(\"Predicted Price Distributions\")\n",
|
||||
"# plt.ylim(0, 500)\n",
|
||||
"plt.savefig(\"./volton-distribution-plot.pdf\", bbox_inches='tight')\n",
|
||||
"\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 56,
|
||||
"id": "b944b038",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"true_pxs = test_y[ntests-1]\n",
|
||||
"logger = []\n",
|
||||
"for eday_idx, eday in enumerate(edays):\n",
|
||||
" eday = pd.Timestamp(eday)\n",
|
||||
" opts = options[options.expiration==pd.Timestamp(eday)]\n",
|
||||
" for idx, row in opts.iterrows():\n",
|
||||
" K = row.strike\n",
|
||||
" bid = row.bid\n",
|
||||
" ask = row.ask\n",
|
||||
" valuation = np.mean(np.maximum(px_samples[:, eday_idx].numpy() - K, 0))\n",
|
||||
" rtn = np.maximum(true_pxs[eday_idx] - K, 0)\n",
|
||||
" logger.append([eday, K, bid, ask, valuation, rtn.item(), \"Voltron\"])\n",
|
||||
"\n",
|
||||
" valuation = np.mean(np.maximum(mat_samples[:, eday_idx].numpy() - K, 0))\n",
|
||||
" rtn = np.maximum(true_pxs[eday_idx] - K, 0)\n",
|
||||
" logger.append([eday, K, bid, ask, valuation, rtn.item(), \"Matern\"])\n",
|
||||
"\n",
|
||||
" valuation = np.mean(np.maximum(sm_samples[:, eday_idx].numpy() - K, 0))\n",
|
||||
" rtn = np.maximum(true_pxs[eday_idx] - K, 0)\n",
|
||||
" logger.append([eday, K, bid, ask, valuation, rtn.item(), \"SM\"])\n",
|
||||
"\n",
|
||||
" valuation = np.mean(np.maximum(sabr_samples.t()[:, eday_idx].numpy() - K, 0))\n",
|
||||
" rtn = np.maximum(true_pxs[eday_idx] - K, 0)\n",
|
||||
" logger.append([eday, K, bid, ask, valuation, rtn.item(), \"SABR\"])\n",
|
||||
"\n",
|
||||
"full_dat = pd.DataFrame(logger)\n",
|
||||
"full_dat.columns = ['Expiry', \"Strike\", \"Bid\", \"Ask\", \"Valuation\", \"Return\", \"Type\"]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 118,
|
||||
"id": "cd4603f4",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def PlotPayoff(exp_idx, title, ax, ylabel=False):\n",
|
||||
" dat = full_dat[full_dat.Expiry == full_dat.Expiry.unique()[exp_idx]]\n",
|
||||
" g = sns.lineplot(x=\"Strike\", y='Valuation', data=dat, hue=\"Type\", \n",
|
||||
" palette=[colors[6], colors[0], colors[2], colors[4]], lw=2.,\n",
|
||||
" ax=ax)\n",
|
||||
" ax.fill_between(dat.Strike, dat.Bid, dat.Ask, color='gray',\n",
|
||||
" label=\"Market Bid/Ask\", lw=3.)\n",
|
||||
" ax.plot(dat.Strike, dat.Return, c='k', label=\"Payoff\", ls=\"--\")\n",
|
||||
" \n",
|
||||
" if ylabel:\n",
|
||||
" ax.set_ylabel(\"Price (Payoff)\")\n",
|
||||
" else:\n",
|
||||
" ax.set_ylabel(\"\")\n",
|
||||
" ax.set_xlabel(\"Strike\")\n",
|
||||
" sns.despine()\n",
|
||||
" ax.set_title(title)\n",
|
||||
" ax.get_legend().set_visible(False)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 123,
|
||||
"id": "b29bb187",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAB8gAAAKuCAYAAADNUUZzAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAABcSAAAXEgFnn9JSAAEAAElEQVR4nOzddXQUydoG8KeTTNxDcHcItri7uyXBfVmcxWXRZXF2gSW4W4jhrruLu7t7SCDuycj3Rz7m0pkJTCaT6UCe3zk5Z+vtruoX7t57K/12VQkqlUoFIiIiIiIiIiIiIiIiIiKiH5yJ1AkQEREREREREREREREREREZAwvkRERERERERERERERERESUJbBATkREREREREREREREREREWQIL5ERERERERERERERERERElCWwQE5ERERERERERERERERERFkCC+RERERERERERERERERERJQlsEBORERERERERERERERERERZAgvkRERERERERERERERERESUJbBATkREREREREREREREREREWQIL5ERERERERERERERERERElCWwQE5ERERERERERERERERERFkCC+RERERERERERERERERERJQlsEBORERERERERERERERERERZAgvkRERERERERERERERERESUJbBATkREREREREREREREREREWYKZ1AkQEUlBLpfj33//xb///ou7d+/i3bt3iI2NhaWlJZycnODm5obq1aujTZs2sLW1zfB8li1bBi8vL63XKleujO3btxv0eQ0bNsS7d+804h06dMC8efMM+qzvWVxcHKysrL56T4kSJUTtuXPnomPHjhmZFhEREelp1apVWLx4MYCMnff8/PPPOH36tChmZ2eHgwcPIkeOHAZ7TmBgIFq1aoWYmBhRfPLkyejdu7fBnkNERET0WWJiIk6cOIHLly/j9u3b+PjxIyIiImBiYgJHR0fkyJEDFStWRM2aNVG7dm0IgmDwHDjXIiKi9OIKciLKcvbt24cmTZpg6NCh8Pf3x4MHDxAZGQm5XI7o6Gi8efMGR44cwYwZM1CvXj2sWbMGSqVSsnyvX7+O4OBgg4138+ZNrcVx+p/AwEAMGzYMR44ckToVIiIiMpDnz59j9erVRnnWrFmzYGdnJ4pFRUVh5syZBn3O1KlTNV7YVq9eHb169TLoc4iIiIjkcjnWrFmD+vXrY9SoUdixYwfu3buH4OBgJCQkIC4uDoGBgbh58yY2bNiAAQMGoHXr1jh69KjBc+Fci4iI0osFciLKMhISEjBq1CiMGzcO79+/16lPdHQ0/vzzT/Tt2xdxcXEZnKF2SqUSx44dM9h4Bw8eNNhYPxq5XI4NGzagZcuWOH78OFQqldQpERERkQFER0dj5MiRiI2NNcrzcubMiQkTJmjET548iUOHDhnkGbt27cKZM2dEMTs7O8ydOzdDVmoRERFR1hUUFITu3bvjzz//REhIiM79nj59ihEjRmDGjBlITEw0WD6caxERUXqxQE5EWUJiYiIGDRqkdZJsZ2eHqlWromHDhqhQoQJkMpnGPRcvXsSwYcMkW0luqJXMSqWSq6K/okOHDpg/f77RXp4TERFRxktISMCwYcPw+PFjoz7X3d0dtWvX1ojPmjULYWFh6Ro7ODhY6/bwv/32G3Lnzp2usYmIiIi+FB4ejr59++LmzZsa12QyGcqUKYMGDRqgdu3aKFiwoNYxduzYgcmTJxs0L861iIgoPVggJ6IsYe7cuTh//rwo5urqioULF+LChQvYunUrVq5cCV9fX1y8eBG//vorzM3NRfefPXsW3t7exkxb7dq1awbZZv3KlSsG3a79R2PsF+dERESUsWJiYjBw4EBcuHBBkuf/8ccfsLW1FcVCQ0Mxe/bsdI07c+ZMREREiGKNGzdGhw4d0jUuERERUUoTJkzAs2fPRLFs2bJhxowZuHz5Mnbu3IlVq1Zh/fr1OHr0KE6dOoVu3bpprLLev38/1q9fb9DcONciIiJ9sUBORD+8CxcuaBS23dzcsHv3brRt21ZjxbitrS0GDx6MTZs2wcrKSnRt2bJlBt0S6mscHBzU/2yobda5vToRERFlFY8fP0bnzp1x8eJFyXLIlSsXJk6cqBHfv38//vvvP73GPHjwIE6cOCGKubi4YNasWXqNR0RERJSa06dP499//xXFihcvjl27dqFr166wtrbW6JMnTx5Mnz4dy5cv13jntnz58jRt0f4tnGsREZG+WCAnoh/eggULRO28efNi3bp1cHV1/Wq/SpUqYeTIkaJYeHi4xiQ5ozRu3FjUTu/W6HK5XFRkz5UrF7Jnz56uMYmIiIgyoz179sDDwwPPnz+XOpVUt/+cNm0aoqOj0zRWaGio1pezs2bNgrOzs945EhEREWmzceNGUdvGxgYrVqxAjhw5vtm3UaNGGueEx8TEwMfHx6A5cq5FRET6YIGciH5oFy5cwP3790WxefPm6Typ9fT01NiqKeVW7RmlRYsWova1a9fw8eNHvcc7f/686AymFi1aaGx3RURERPQ9e/78OYYMGYIJEyYgLi5O6nTUZs2aBRsbG1Hsw4cPGh9y6jJOyjM1O3bsiEaNGqU7RyIiIqIvhYaG4tKlS6JYt27dkC9fPp3H0Hb/P//8Y5D8vsS5FhERpRUL5ET0QwsICBC1GzRogCpVqujc39raGrVr14aVlRVy586N0qVLa5xNnlGqVasGFxcXdVupVOLo0aN6j3fo0CFRu1WrVnqPRURERJSZhISEYObMmWjTpg1Onjypcb1nz56oVKmSBJkly507t8YKKgDw8/PD5cuXdRrjxIkTGvO5PHny4LfffjNIjkRERERfunLlChQKhSjWsmXLNI1hamqKhg0bimIPHz6ESqVKd35f4lyLiIjSykzqBIiIMopcLtc4b6hnz55pHmfJkiWSrLQ2MTFBs2bNROenHzlyBD169EjzWImJiaKt4QsWLIgyZcoYJM/P49+5cwevX79GWFgY5HI5nJ2d4erqip9++gn29vYGe9aXgoODcfv2bbx9+xZxcXFwdHSEq6srKlasmKm2vnr+/DkePnyIjx8/Ii4uDnZ2dihUqBDKly+v8YUzERERpd2qVatEc6bPbGxs8Ntvv6FTp056zQMNydPTE0ePHsW5c+fUMZVKhSlTpmDfvn2wtLRMtW9kZCRmzJghigmCgLlz52rsdvQtL168wOPHjxESEoLIyEg4ODggW7ZscHNzQ+7cudM01rdER0fjzp07CAkJQUREBCIjI2FqagobGxvkzJkTxYsXT9MqNF1cv34d9+7dQ1JSEgoVKoSqVatyvkVERKSHp0+fitoymQwlS5ZM8zgp/78+KSkJ4eHhcHJySld+KXGulbnnWm/evFHnGh0dDTs7Ozg5OSFnzpwoW7as0RYkERF9xgI5Ef2wbt++jaioKHXbwcEB1atXT/M4Um5D3qJFC9HL3s/brH/r/PSUTp8+Lfq7SOsXv6m5ceMGNmzYgDNnzqS6jampqSkqVKiALl26oE2bNjr/fU6cOBG7d+9Wt5cvX64+l/2///7DmjVrcPXq1VSfWalSJQwaNAi1atVK9RmXLl1Cr169Ur0+adIkTJo0Sd2uWrUqtm7dqlP+CoUCO3bsgLe3N549e6b1HnNzczRt2hSDBw9G0aJFdRqXiIiIdFOvXj3MnDkTuXLlkjoVtT/++AOtW7dGTEyMOvbq1SssXbpU66qnz+bOnatx1E6fPn1QrVo1nZ4bGRmJ9evX4/Dhw3j16lWq9xUvXhwdOnRAjx499H5JGRwcDD8/P/z777+4f/++xsqzlPLmzYv27dujZ8+ecHR0/Ob4qc0Rg4ODMW7cOFy8eFF0v7W1NTp16oQhQ4Zkqg8oiYiIMjs3Nzf07NkTwcHBCA4OhlKphIlJ2jekTUhIyIDstONcS5OUc63Q0FBs3rwZ+/fvx7t371J9prW1NapUqYKOHTuiadOmev17RkSUVvxfGiL6Yd2+fVvULlu2LExNTSXKRj+VK1cWFcOVSiWOHTuW5nFSbhHVunXrdOUVGhqKIUOGoEuXLjh27NhXz/hUKBS4du0axo0bh3bt2uHevXt6Pzc+Ph7jxo3DwIEDUy2Of37m5cuX0a9fP4waNQpJSUl6P1Mfz549Q/v27TFr1qxUi+NA8sr7AwcOoH379ti8ebMRMyQiIvpxFS9eHKtWrcKaNWsyVXEcSN7+c/z48RrxzZs34/79+1r7XLhwAbt27RLFihYtilGjRun0TD8/PzRp0gSrVq366gtbAHj8+DHmz5+P5s2bi1Zf6UKpVGLlypVo0qQJli1bhjt37nzzhS0AvH37Fl5eXmjWrBn+/fffND3zs9jYWPTr10/jhe3na1u3btVYBUdERERfV79+fUyZMgV///03fHx84Ofnp9c4Kf8/2Nzc3OCrxz/jXEuTVHOtEydOoFmzZli1atVXi+Ofx/jvv/8wcuRIuLu7f/PvkYjIEFggJ6If1oMHD0TtEiVKaNzz8eNH7N69G3PnzsXYsWMxefJkLF68GEePHkVsbKyxUk2ViYkJmjdvLoodOXIkTWPExcXhn3/+UbdLliyJIkWK6J3T48eP0bZtW63ne37Lo0eP0L17d72K/ImJiRg8eDD27duXpn6HDh3CuHHj0vw8fd25cwfdunXD48ePde6TlJSEOXPmwMfHJwMzIyIi+rGVLVsWCxcuxN69e9GgQQOp00lVly5dUKNGDVFMoVBg2rRpUCqVonhCQgKmT58uislkMixYsAAWFhZffY5SqcScOXMwdepUhIeHpynHd+/eYeDAgfD399fpfoVCgV9//RVLlixBfHx8mp71WXh4OIYNG6bxkasuFi1ahCdPnqR63dXVFZUrV9YrLyIiItJfYmKiRlG2WLFiGfpMzrW0M+Zc6/Dhwxg+fDgiIyPT/Ky7d++ia9euePv2bZr7EhGlBbdYJ6IfVsqvDb9cQfTo0SMsW7YMp06dSvVLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 2250x525 with 3 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(1, 3, dpi=150, figsize=(15, 3.5))\n",
|
||||
"\n",
|
||||
"PlotPayoff(3, \"6 Month\", ax[0], ylabel=True)\n",
|
||||
"PlotPayoff(6, \"1 Year\", ax[1])\n",
|
||||
"PlotPayoff(8, \"2 Years\", ax[2])\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"lw = 2\n",
|
||||
"custom_lines = [Line2D([0], [0], color=colors[0], lw=lw, ls=\"--\"),\n",
|
||||
" Line2D([0], [0], color='gray', lw=lw),\n",
|
||||
" Line2D([0], [0], color=colors[6], lw=lw),\n",
|
||||
" Line2D([0], [0], color=colors[4], lw=lw),\n",
|
||||
" Line2D([0], [0], color=colors[0], lw=lw),\n",
|
||||
" Line2D([0], [0], color=colors[2], lw=lw)]\n",
|
||||
"fig.legend(custom_lines, [\"Payoff\", \"Bid/Ask\", \"Volt\", \"SABR\", \"Matern\", \"SM\"],\n",
|
||||
" frameon=True, ncol=6, bbox_to_anchor=(0.95, -0.08))\n",
|
||||
"plt.savefig(\"./payoffs.pdf\", bbox_inches=\"tight\")\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 67,
|
||||
"id": "4ef1cb0a",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"numpy.datetime64('2012-12-03T00:00:00.000000000')"
|
||||
]
|
||||
},
|
||||
"execution_count": 67,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"qday"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 71,
|
||||
"id": "c1caf7e6",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"2013-06-28T00:00:00.000000000\n",
|
||||
"2013-12-21T00:00:00.000000000\n",
|
||||
"2014-12-20T00:00:00.000000000\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"print(edays[3])\n",
|
||||
"print(edays[6])\n",
|
||||
"print(edays[8])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "4779eac9",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"# Calibration"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 20,
|
||||
"id": "7682f20f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"temp = px_samples[:, 0]\n",
|
||||
"yy = test_y[0]\n",
|
||||
"smp = temp.log().sort(0)[0]\n",
|
||||
"log_px = yy.log()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 21,
|
||||
"id": "9643581d",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"tensor(0.3070)\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"print(torch.sum(smp < log_px)/smp.shape[0])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 22,
|
||||
"id": "dfb68818",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"tensor(0.3070)"
|
||||
]
|
||||
},
|
||||
"execution_count": 22,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"torch.sum(smp < log_px)/smp.shape[0]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 128,
|
||||
"id": "887c515b",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"ConstantMean()"
|
||||
]
|
||||
},
|
||||
"execution_count": 128,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"gpytorch.means.ConstantMean()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "3d7d33a1",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
BIN
Binary file not shown.
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
@@ -0,0 +1,551 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "3d8b2c23",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Warning no robinhood utils.\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import datetime as dt\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns\n",
|
||||
"import torch\n",
|
||||
"import gpytorch\n",
|
||||
"import os\n",
|
||||
"# import robin_stocks.robinhood as r\n",
|
||||
"import pickle5 as pickle\n",
|
||||
"import pandas as pd\n",
|
||||
"# sns.set_style(\"white\")\n",
|
||||
"# sns.set_palette(\"bright\")\n",
|
||||
"sns.set(font_scale=1.5)\n",
|
||||
"sns.set_style(\"white\")\n",
|
||||
"\n",
|
||||
"import sys\n",
|
||||
"sys.path.append(\"../\")\n",
|
||||
"from voltron.likelihoods import VolatilityGaussianLikelihood\n",
|
||||
"from voltron.models import SingleTaskVariationalGP as SingleTaskCopulaProcessModel\n",
|
||||
"from voltron.kernels import BMKernel, VolatilityKernel\n",
|
||||
"from voltron.models import BMGP, VoltronGP\n",
|
||||
"from voltron.train_utils import LearnGPCV, TrainVolModel, TrainDataModel\n",
|
||||
"from gpytorch.kernels import ScaleKernel, RBFKernel, MaternKernel"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "ba337d0d",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Options Helpers"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "61523f0d",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def GetTrainingData(SPY, date, N):\n",
|
||||
" idx = SPY[SPY[\"Date\"] == date].index.item()\n",
|
||||
" return SPY['Close'].iloc[(idx-N):idx]\n",
|
||||
"\n",
|
||||
"def GetTrueValue(SPY, date, strike):\n",
|
||||
" close_px = SPY['Close'][SPY[\"Date\"] == date].item()\n",
|
||||
" return np.maximum(close_px-strike, 0)\n",
|
||||
"\n",
|
||||
"def GetTradingDays(SPY, start, stop):\n",
|
||||
" start_idx = SPY[SPY[\"Date\"] == start].index.item()\n",
|
||||
" stop_idx = SPY[SPY[\"Date\"] == stop].index.item()\n",
|
||||
" return stop_idx-start_idx\n",
|
||||
"\n",
|
||||
"def FindLastTradingDays(SPY, dates):\n",
|
||||
" last_days = []\n",
|
||||
" for date in dates:\n",
|
||||
" last_days.append(np.max(np.where(SPY.Date < date)[0]))\n",
|
||||
" \n",
|
||||
" return np.array(SPY.Date[last_days])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "85bc2e35",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Data"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"id": "7d014904",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"SPY = pd.read_csv(\"./data/SPY_prices.csv\")\n",
|
||||
"SPY['Date'] = pd.to_datetime(SPY['Date'])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "19d1fee1",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Data Setup"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"id": "eb96fc66",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ntrain = 252\n",
|
||||
"options = pd.read_csv(\"./data/SPY_\" + str(2014) + \".csv\")\n",
|
||||
"options.expiration = pd.to_datetime(options.expiration)\n",
|
||||
"options.quotedate = pd.to_datetime(options.quotedate)\n",
|
||||
"qday = options.quotedate.unique()[0]\n",
|
||||
"options = options[(options.quotedate == qday) & (options.type=='call')]\n",
|
||||
"edays = options.expiration.sort_values().unique()\n",
|
||||
"testdays = (edays - qday)/np.timedelta64(1, \"D\")\n",
|
||||
"edays = edays[(testdays > 100) & (testdays < 500)]\n",
|
||||
"lastdays = FindLastTradingDays(SPY, edays)\n",
|
||||
"ntests = np.array([GetTradingDays(SPY, qday, pd.Timestamp(ld)) \n",
|
||||
" for ld in lastdays])\n",
|
||||
"fulltest = ntests[-1]\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"train_y = torch.FloatTensor(GetTrainingData(SPY, qday, ntrain).to_numpy())\n",
|
||||
"test_y = torch.FloatTensor(GetTrainingData(SPY, \n",
|
||||
" pd.Timestamp(lastdays[-1]),\n",
|
||||
" fulltest).to_numpy())\n",
|
||||
"full_x = torch.arange(ntrain+fulltest).type(torch.FloatTensor)\n",
|
||||
"full_x = full_x/252.\n",
|
||||
"dt = full_x[1] - full_x[0]\n",
|
||||
"train_x = full_x[:ntrain]\n",
|
||||
"test_x = full_x[ntrain:]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 5,
|
||||
"id": "e4fa4d35",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"## learn vol with GPCV ##\n",
|
||||
"vol = LearnGPCV(train_x, train_y, train_iters=400)/(dt**0.5)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 6,
|
||||
"id": "a01cdb88",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAZkAAAEACAYAAABhzAtFAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAAB8BElEQVR4nO29eZwdVZk+/tR2t16zhyyQpZNgCIsGCNsQIQSRYVQQf+MAQb6CKBLFUVxmGBj5qCPbN3HBBBQcQaPfGQWjCAIJSETCJkogBLICCYHsvaS771pVvz+q3lPn1K26t273vd19u8/z+fChc29V3XOqTp33PO/zvu9RbNu2ISEhISEhUQOog90ACQkJCYnhC2lkJCQkJCRqBmlkJCQkJCRqBmlkJCQkJCRqBmlkJCQkJCRqBn2wGzBUkMlksHHjRowbNw6apg12cyQkJCTqAqZpYv/+/Zg3bx4SiUTR99LIuNi4cSMuvfTSwW6GhISERF1i1apVOPHEE4s+l0bGxbhx4wA4N2rixImD3BoJCQmJ+sCePXtw6aWXsjnUj5oamZ6eHixfvhyPPvoourq60NbWhmuvvRaLFi0qe+7OnTtxyy234Pnnn4dlWTjxxBPx9a9/HW1tbeyYN998E//v//0/PP/889i1axd0XcfMmTNx5ZVXRvoNHuQimzhxIqZMmVJZRyUkJCRGOMJkhpoK/0uXLsVDDz2E6667DnfffTfa2tqwdOlSrFu3ruR5Bw8exCWXXILdu3fj1ltvxbJly9DZ2YnLLrsMe/bsYcc988wz+POf/4zzzjsPP/jBD3Dbbbdh4sSJ+PznP4+f/exnteyahISEhEQE1IzJrFu3DuvXr8edd96JxYsXAwBOOeUU7Nq1C7fccgsWLlwYeu69996Lrq4uPPDAA5gwYQIA4IQTTsCiRYuwcuVK3HzzzQCA888/H5deeikURWHnLly4EPv378fKlStxxRVX1Kp7EhISEhIRUDMms2bNGjQ1NQluK0VRcOGFF2LHjh3Ytm1b6Llr167FaaedxgwMAIwaNQpnnXUW1qxZwz4bPXq0YGAIxx57LDo6OpDJZKrUGwkJCQmJvqBmRmbr1q1oa2uDqoo/MWfOHADAli1bAs/LZDLYuXMnZs+eXfTdnDlzcPDgQRw8eDD0d23bxvPPP4+pU6cGhtNJSEhISAwcamZkOjo60NLSUvQ5fdbR0RF4XmdnJ2zbDjy3tbW15LkAcN9992Hjxo245pprKm6zhISEhER1UVPhP8iVFeW7KN8HYe3atbjttttw0UUX4eMf/3jF50tISEhIVBc1MzKtra2BjKOzsxMAApkKfa4oSuC59BkxGh5PPfUUvvSlL2Hx4sX49re/3ddmS0hISAwofvaH1/Cj32wY7GbUDDUzMm1tbdi+fTssyxI+Jy0mSHMBgEQigalTpwZqNlu2bMHo0aMxZswY4fN169Zh6dKlOPPMM3HHHXfIsjASEhJ1g627OrDtnY7BbkbNUDMjs3jxYnR1deHJJ58UPl+9ejWmT58uJFX6cc4552D9+vXYv38/+6yjowN/+tOfWDg04emnn8bSpUtx2mmn4Xvf+x4Mw6huRyQkJCRqiHzBwnDeoLhmeTILFy7EggULcMMNN6CjowNTpkzB6tWr8dJLL2HFihXsuCVLluCFF17A5s2b2WdXXnklfv/73+Pqq6/GtddeC13XsXLlSui6js997nPsuL/+9a9YunQpJkyYgKuuugqbNm0S2jB37lzEYrFadVFCQkKi38ibFixLGpmKoSgKVqxYgWXLlmH58uWsrMydd96Js88+u+S5Y8eOxapVq3Drrbfia1/7Gmzbxvz58/GLX/wCkyZNYsc9++yzyGQy2LVrF5YsWVJ0nSeeeEKWiJGQkBjSKBSs8gfVMRR7OPO0CvDOO+9g0aJF0jBJSEgMKD53yxNQVQUrvlZ68T1UUW7ulJuWSUhISAwi8ubw1mSkkZGQkJAYRBQKw1uTkUZGQkJCYhDhRJcNditqB2lkJCQkJAYRBdOEOYytjDQyEhISElXGwc40du/vjnTscM+TkUZGQkJCosr42cObcMeql8oeZ9s2CqYtNRmJ2mHfoV68tiN86wIJCYn6QyZbQDZXKHtcwXRyZCSTkagZHvjT1kgrHgkJifqBbSOSmJ93EzGtYZyPKY3MICOXt5DLm4PdDAkJiSrCsu1I7IQZGclkJGoFy7ZhDmN/rITESIRtA1Fea3KXSU1GomYwTbtoOwQJCYn6RqVMRmoyEjWDw2QGuxUSEhLVhG3ZkZiMdJdJ1ByWJZmMhMRwg20jkvLP3GXD18ZIIzPYkJqMhMTwg2VXyGSG8RwgjcwgwzRtRyQcxoNMQmKkwQlhlpoMII3MoIN8scPZJyshMdIQVfgvSCYjUWvQ4JIuMwmJ4QM7qrtMajIStYZpDf+VjITESIMj/Jc/Ll/wErGH6xwgjcwggwLLJJORkBg+cIT/6JoMMHx1GWlkBhk0EE2ZLCMhMWxgR9VkuPd+uOqy0sgMMswRUFbCj1e27cdLb+wd7GZISNQMVsSyMjyTGa5TgD7YDRjpGInRZb9+YivS2QLmHz1hsJsiIVET2BHLMAtGZphaGclkBhksuswcngMsCJZlC24CCYnhBtuqnMlITUaiJjBHYAizadkjyqhKjDxEzpMxJZORqDFoYA1Vd9n+9jQe/NPWqr4AkslIDHdEzpMZAZqMNDKDjKEeXfarx9/Af/9hE/76evWEeksyGYlhDqvCsjKAZDISNQJNtkPVXTauNQkAeOaVd6t2TdO2UZCVpyWGMZwQ5vLH8YxeajISNQGLLhuiRkbXnSHy/Gt7qmYILctmNZskJIYjLKuyTcuAoesy7y+kkRlkDPXaZdSunnQeW3e2V+WajiYzNPsrIVENVFqFGfCqfww3SCMzyKBJfKgyGV47aT+crc41LZvVbJOQGI6Ivp8MV7tMMhmJWmDoMxnPGFTLxSWZjMRwB7GYUmxm9/7uEZEnIzP+BxlDXZPhjUG+ShFwlm0P2Wg6icGDs6pXYOj1v/al19m2AUUp/n7X3sP4/G1PCp8VTAt/fX0v5h89HkrQSXWK+n+adQ6PyQzNSZdvV76KTMaMKIxKjBx8579fwN2/fWWwm1EVlGMy6Wyh6LNH1r+Fm+95Dn/++24AwIYt+/HI+jfZ989seBd7D/XWoLW1hTQyNYJp2dj89qFIx/H/H2rgNZlqJVCy3KAh2meJwcHBzgz21eEkGgSb1SQM/l5Vi5nKoc4MAMeNBgBrX9yJXz+xlV3v9l/8FY8991b1G1tj1NTI9PT04Nvf/jbOOOMMHHfccbjooovwxBNPRDp3586d+PznP4/58+fj/e9/Pz7zmc9g27ZtRcctX74cV111FU499VTMmTMHP/zhD6vdjT7hpTf24vofPI09B3tKHmcNceG/YFrMfVEwLXR2Z9HVk+vXNclwyax/CR6mZSFXp6Ht6WwBq9dt5yp4OJ+HMZkgd3E8pgEAejMOy8mbFjsukzNhWjZy+fq7PzU1MkuXLsVDDz2E6667DnfffTfa2tqwdOlSrFu3ruR5Bw8exCWXXILdu3fj1ltvxbJly9DZ2YnLLrsMe/bsEY69//770d3djXPOOaeWXakYvek8AKDb/X8YaiH8/9NXfof/u+qlqlzLsmwk3MGfL1hY9qu/4Ue/ebl/17TJyAyOYX39zUPIBLgrJAYXplm/+VN/e2Mf7v39RuzcexgAz2SCx3jQ2Kf3jFxppmmx4+izelyY1Uz4X7duHdavX48777wTixcvBgCccsop2LVrF2655RYsXLgw9Nx7770XXV1deOCBBzBhglMO/oQTTsCiRYuwcuVK3HzzzezYl156CaqqoqurC//7v/9bq+5UDBoc+TIrj1q5y5762zv4yqXzIx//5rudeOmNfbj47FnC5wXTQjym43BvHgXTQsfhLJLx/g0br/L0wL8wBzvT+NqdT+PM90/GVy87UfiuYFrQVGVYia71BNOykeNCeusJ1G4KSWZSZshrHWQsYoZoZAqmF+rfm8mHnjfUUTMms2bNGjQ1NWHRokXsM0VRcOGFF2LHjh2Bri/C2rVrcdpppzEDAwCjRo3CWWedhTVr1gjHqurQlJXIaJR7aYZKdNlfNryL+x7eVETvTdNGTFehKk4Ic75g9Xu1SX0djBeGXH079xwWPrdtG5/5r7X447NvDXibJBzUqzsI8ML7CwVR8A9jMkG1++hYMihBTKZawTcDiZrN0Fu3bkVbW1uREZgzZw4AYMuWLYHnZTIZ7Ny5E7Nnzy76bs6cOTh48CAOHjxY/QZXGZa7Aik1KJzSE87fgy2Csx06fc0wLRuapkDXVGZg+hvK7BUFHfg+5/KO0feHyaazBRzoSGPvweEhPA8UCqaFf1vxF2zcfqDf17IsS0hOrCfQgolq8nnRZSHHB0STkqEiTcbZEsP5rJ7dZTUzMh0dHWhpaSn6nD7r6OgIPK+zsxO2bQee29raWvLcoQTmLivx0vCrHGuQQ5jJuPjb4biQVBi6ioJpIVcw+72aYkxmEPpMwjK5JgiHe/Pu9/U5yQ0WetJ5bNx+EFt3dfT7WqZl1+VKHfAWk2Qoygn/Qd6AINZCof4sGKAO709NfU2lfNvl/N717hdn7rIS9J93kQ126fuwHTpNy4auKdB1h8nkC1a/V1NmyG8NBEgj8zOZw64brZovcXc6j7sffIWxp+EIL4ij//fNNO26jS6j/tPYZkwm5PigsU/X6CXh312EFUxbMpkgtLa2BjKOzs5OAAhkKvS5oiiB59JnxGiGMjx3WQkmwxmZwa5bRAPa77YzicloDpMhQ9MfDKYmQ0wlpotMpqu3+kZm046D+MMzb2L7O51Vu+ZQg/cs+z9+TctGvk4NMrmQ8z4mE6a1EotvShk4/fhJzmfuuWkm8nsBMsRk6jH6rmZGpq2tDdu3by9yv5AWE6S5AEAikcDUqVMDNZstW7ZLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"log_returns = torch.log(train_y[1:]/train_y[:-1])\n",
|
||||
"plt.plot(log_returns)\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 7,
|
||||
"id": "c0436299",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYsAAAEACAYAAABCl1qQAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAABMOklEQVR4nO29eWBU9b33/z6z71tWkpAQMoQdVJAgKqksyqV1Qb239aLWVm97q/Ghj7W//vp4y9Pe6k/q9YqtFJdqbytdru21olaLolRUtggKEhBIAoSQPZnMvs+c3x9nzslMZpKZJDOZmfB5/QM5c86Z78lMzvt8doZlWRYEQRAEMQqibC+AIAiCyH1ILAiCIIikkFgQBEEQSSGxIAiCIJJCYkEQBEEkRZLtBaQbr9eLpqYmFBUVQSwWZ3s5BEEQeUEoFEJfXx8WLFgAhUIR9/qUE4umpiZs3Lgx28sgCILIS37/+99j6dKlcdunnFgUFRUB4C64tLQ0y6shCILID7q7u7Fx40bhHjqcKScWvOuptLQUFRUVWV4NQRBEfjGS+54C3ARBEERSSCwIgiCIpJBYEARBEEkhsSAIgiCSQmJBEARBJIXEgiAIgkgKiQVxyfEfOw7jt2+dzPYyCCKvILEgLilYlsXhUz04dKIr20shiLxiyhXlEcRo2Jx+uL1BeHxOeH1BKOT0J0AQqUCWBXFJ0dnvBACwLHCu057l1RBE/kBiQVxSdPa5hP+3dliztxCCyDPIBicuKTr7nRCJGGiUUrRetGV7OQSRN5BYEJcUnf0ulJhUmFaoRstFa7aXQxB5A7mhiEuKrj4XygrVmF6sRdeACyzLZntJRJ7wHzsO4+NjHdleRtYgsSAuGViWRWe/E+VFGhQZlfD5Q3C4A9leFpEH+AIhfHi0A4e/6Mn2UrIGiQVxydBv9cLrD6G8WINCgzKyzZPlVRH5wEDkezJo92V5JdmDxIK4ZDjfxQW0q0p1KIqIRd+gO5tLIvKEvsGIWDi8WV5J9qAAN3HJcL6Lq6uomqZDIBACAPSRZUGkQJ+Ve6i4lC0LEgvikqGty4EioxIapRRhuQQSsYjcUERK8JaFzeVDKBSGWHzpOWUuvSsmLlnOd9lQVaoDAIhEDIoMSuEmQBCJGHR48cru0+i2cJYFywJW56VpXZBlQVwSBIJhXOx14sp5pcK2QoOS3FDEqOw/1onf7ToFpVwsbLPYvSjQK7O4quxAlgVxSdDR50QozKJqmk7YVmQksSBGZ8DOBbQ9vhCKjZxADDrIshgRl8uFrVu3YteuXbDb7TCbzXjggQewevXqUY87fPgwXn31VZw8eRItLS0IBoM4ffp03H4WiwU//vGP8cUXX6C/vx8Mw6CyshK333477rjjDojF4gRnJ4hYPj3di4+PduDBf7oMDMPEvNYzwPWEKitUC9sKDUpY7F4EgmFIJfTcRMQzYBvKfqqtNKJ30INB+6WZEZXSX0hDQwPefPNNbNq0Cc8//zzMZjMaGhqwd+/eUY87ePAgGhsbUVVVhTlz5oy4n9/vh0wmw7e+9S388pe/xNNPP42lS5fipz/9KX7yk5+M7YqIS5YPjrRjd+MFdA244l7jnwZNOoWwrbxIg3CYxX2P7cbFXse435dlWXx6uhehMFWD5zO+QAgvvdEEl2eoUDM6AWLWdCMAwHKJZkQltSz27t2L/fv3Y9u2bVi7di0AYPny5Whvb8eWLVtQX18/4rH3338/GhoaAACPPfYYmpqaEu5XWlqKJ598MmbbypUrMTAwgL/85S/YvHkzJBIKrxCj09bF3fCbWgdQVqiJeY1/GjRo5cK2L11RgXA4jJ+/chTN7VZUFGvH9b5HTvXiJy8exI/urcOyqJgIkV+caRvEzr2tqCnX40tLpgPgLIu5M0ywOnxYaC6AViUjy2Ikdu/eDa1WG+NyYhgGGzZswNmzZ9HS0jLyyUUTM+2NRiNEItGEz0NMfUKhMNp7ebHoj3vd4vBBp5ZBEpXyKBIxWL5gGgDA7vKP+70/Osr1C+rsc477HET2cXs5i6K9d+hztNg9mDXdgBf+zxrMmm6ESSeHhcQiMc3NzTCbzXE37NmzZwMAzpw5k7bFsCyLYDAIm82Gt99+G6+99hruvfdeEgsiKZ39LgSCYcgkIhxvHYhrEDho98a4oHhUCilEDOAYp1gEgmEcOtENAOgZoGrwfMbjCwKA4JJ0ewPw+EIo0A99b0pMapzrsiN8Cbock96FrVYr9Hp93HZ+m9VqTdtifv/732P+/PlYtmwZHnroIdx9993YtGlT2s5PTF0udHN/4PVXVKDf6sG2Px+DNSprZdDhjXFB8YhEDLRq2bgsi4+PdeBHz++HyxOASMQIufhEfuKOiEV7D2dZ8MHt6DTZay8rQ6/FjRNnByZ/gVkmpUDA8MySVF8bK+vXr8fixYtht9tx6NAh/PrXv4bT6cSPfvSjtL0HMTU532WHiAHuXj8P/kAY7x5qQ9U0LW66tgYAF+AeKSahVclgd49NLI419+E/fncEhXoF6uaXIhRm0UNikde4vZxYdPU7EQqFMWDjgtvRlsVVi8qg/svneLexDQvNhVlZZ7ZIKhYGgyGh9WCzcU3ZElkd48VkMsFkMgEArr76ahgMBvzsZz/Dbbfdhnnz5qXtfYipR1u3HdMK1TBo5fjexiuw/3gnBqzckyHLshi0+2BMYFkAnFiM1Q317Kufo6xQjf/ctBIqhRQvvdGE4639YFk2rQ9QxOTBu6GCIRZdA66EloVcKkb9FRV4e/95tPc4wLLAA7cvRm2lMStrnkySuqHMZjNaW1sRDodjtvOxitra2sysDMCiRYsAAOfPn8/YexBTgwGbB8VGFQDO2jXpFMIfu8MdQDAUhjFBzAIAdGN0Q7k8AXT0OXHdkulQKaQAgBKTCj5/6JJtBTEV4APcAOeK6o9YFiZ97PfmG1+Zj69/eR5kEjHOdthw8pxlUteZLZKKxdq1a2G327Fnz56Y7Tt37kR1dTXMZnPGFnfw4EEAQGVlZcbeg5gaOFwBaFUy4ecCvQID9ti20iZtesTibCdnVc8sH7KqS0ycUJErKn9xe4PQKDnxv9BjR6/FA61KBrk0tihYIZfg9lWz8PgD10SOuzQGaCV1Q9XX16Ourg6PPPIIrFYrKioqsHPnThw5cgTbt28X9rvrrrvQ2NgYU6FtsVjQ2NgIALhw4QIAYNeuXQCA8vJyLFy4EADw0ksvobW1FcuXL0dJSQkcDgf27duHV155BTfccAMWLFiQvismpiQOtx9adbRYKIUZ20KNhS6xG0qnlsHh9qfsQjrbwYlFTSKxGHBjTpVpXNdAZBePL4gCvQIGrRyn2wbRN+iBuWJkN7tYxEClkMQU8U1lkooFwzDYvn07nnrqKWzdulVo97Ft2zasWrVq1GObm5vjspn4nzds2IAtW7YAAObOnYv9+/fjiSeegNVqhVQqxcyZM/GDH/wAGzduHO+1EZcIoTALlzfesjh0gusY+rcD5wEgYeoswMUsAsEwfP4QFPLEfxKn2yx452AbvnHjfJztsMGolce4tYojYtFt4arHvzhnwQuvH8dPv71CeFolchuPNwilXILpJVrs+7wTHl8QKxbOHvUYlUIKF1kWQ2g0GmzevBmbN28ecZ8dO3bEbaurq0vYC2o4K1aswIoVK1JZCkHE4XT7wbKAVjV0Uy7QK+APhPD2/vOw2L2QSkQjioUuYpHYXf4RxeIP75zGp6d78cV5C/yBUIwLCgAUMgmKTSpciFSR/27XF2hpt+LkuQGq6s4T3L4ANEoZ5s4wYXcj5wmZM2N0K1GjlF4ylgVVuxF5jzPyxxrjhtJxGSwWuxc3r6zB0/+7HsoRhIA/bqT0WYvdi6NnerFsXimc7gB6Bz1xYgEA1dN0ONdlQ0u7FZ+3cFXkzRes474uYnLx+IJQKiSYW80JBMMAs6tGz3Li3FDByVhe1qGGS0Tew6e9RruhojNYFs8qRGWpLu44Hv64kYLcez+9iDALfOPGeTBoFdh9qA0rLy+P229GmQ6fnOzGzr2tUMrF0GvkQtyEyH3c3iBUcgnKizTQqWUw6RRCtttIqJXSmM60UxkSCyLv4S2C4W4oHnOFYdTjeTdUdK0Fy7L424HzWHlZOT5v6UdVqVYo6tvwpcQZgNVleoRZ4KOjF3HNZeWQSkQ4/EUP1V7kCW4vZ1kwDINvfGVeUqEAALVCKnQPmOqQWBB5j5MXC3VsgBsACvWKEesreKJjFjwXehx49tXPwTAMbE4fCg3JJ6NVRwYrhVlg+fxpsLv9eP+TdvQNeoQAOJGbhMMsPL4gVHJOINYsq0rpOLVSesmkzlLMgsh77C7uj1UX5YaSSsQwaOQwTzckPV6jkoFhYsWiq5/LarI5fbC7/IKgjEZpgRoKmRgSMYMlc4sxO1LV++8vHcSZC4NjuSRikvH6ubiDSjG252eVQgKXNxjXuHIqQmJB5D0Otx8iBnFug/99xxW4e33yNjFiEQOjVh4z6KZ7YLhYJK7RiEYkYjCvugBXziuFSiGFeboBD9y+GBa7F/+zp3mMV0VMJnyrj5GSIEZCo5QiHGbh9YcysaycgtxQRN7jcPmhVsogEsXGBa6YU5zyOUpM6pjqa96yGLB54fEFU7IsAODfvrkMwNA61l01AyfODeDomb4JxS7eOXgetZVGVJelrxcbMQTfRHDslgX3gOLyBMYsNPkGWRZE3uNw+6FTT6zwraRAhR7L0DjW7shsivYeLnipTVEspBJx3DzvOVXcpLXxtgJxewP45f8cw+/+dmpcxxPJ4S2LVILa0agjBZe/fesktv35aLqXlVOQWBB5j8Ptj0mbHQ8lJhX6rR4EQ1zDTH6ON29hpGpZJGJupLCr8UQ3jieY4peMc512sCzw2ZneSyaYOtnwv9exWgfqiLh8eLQDn5zsSfu6cgkSCyLvcbgC0ExQLEpNKoRZoG/Qg1AojN6IFRCKTESbiFhLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"log_returns = torch.log(train_y[1:]/train_y[:-1])\n",
|
||||
"plt.plot(vol)\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 8,
|
||||
"id": "adb31466",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYYAAAEACAYAAAC3adEgAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAABC20lEQVR4nO3daWBb1Zk38L/2ffO+xatiJ46zkT2EpGQhKTBlndKWQjOFoe00EMqEt6VMQ6dDCwWaQBNSSCgta9vptKSkKQlZIKTZV7IYx1viXV61b7ak+36QdG1Zsi3bsrY8vy/EV1fSuZbRc89zznkOh2EYBoQQQogPN9YNIIQQEl8oMBBCCAlAgYEQQkgACgyEEEICUGAghBASgB/rBoyHw+HApUuXkJ6eDh6PF+vmEEJIQnC73ejs7ERFRQXEYnHQ4wkdGC5duoT7778/1s0ghJCE9N5772Hu3LlBxxM6MKSnpwPwXlxWVlaMW0MIIYlBp9Ph/vvvZ79DB0vowOBPH2VlZSEvLy/GrSGEkMQyVAqeBp8JIYQEoMBACCEkAAUGQgghASgwEEIICUCBgRBCSAAKDIQQQgJQYCCEkASi67biwZ/uwcXargl7DwoMhBCSQPYcuwa92Yn3P66asPdI6AVuhBByvTh2sQ3nqjtw/GIbJCIeLtV1o6qhB1MKUiL+XtRjIISQOMcwDN75qBIfHfX2Fv7j3lnISJGitdMyIe9HPQZCCIlz9S1GNLVbcPuNRVArRbhpZg5umpkDHm9i7u0pMBBCSJz79GwzeFwOvr56CpQy4YS/H6WSCCEkzp2qbMes0vSoBAWAAgMhhMQ1u9OF1i4LphRGfpB5KBQYCCEkjjW0mcAwQHGOKmrvSYGBEELiWH2rEQBQmKOM2ntSYCCEkDhW32KEXCJAuloStfekwEAIIXHsaqsRxbkqcDicqL0nBQZCCIkzjToT3v5HJexOF661maOaRgJoHQMhhMSd/aea8MGntTjzRQd6+9xYPD0nqu9PPQZCCIkzDW0mAN6B58UzsjGtODWq7089BkIIiTMNOhOml6QhRSnGg7dOjfr7U2AghJA4YrH1otvowL8sKcY9yyfHpA2USiKEkDjSoDMDAAqyozvgPBAFBkIIiSPXfOMLBVkUGAgh5LrncntwvroDMjEfaWpxzNpBYwyEEBIjdc0GVF3rwYr5+Xjr75U4UalDp96OrywtjuqCtsEoMBBCSAx8cbUHz+w4BrvThb98Wosugx3zy7PwnTunY0FFdkzbRoGBEEJiYNtfPodKLsRNs3Lx8YkG/PudFfjKTSWxbhYACgyEEBJ1rV0WXGsz4eE7KvCVm4px36pSZGiksW4WiwafCSEkyo5fbAMALKrIBofDiaugAFBgIISQqHJ7GBw+3wJtngoZKfEVEPwoMBBCSJR4PAxefPc0apuNWLOoMNbNGRIFBkIIiZLKq9048nkr7l8zBasXFsa6OUOiwEAIIVFysbYLHA5w+5LiWDdlWBQYCCEkSi7UdaE4VwW5RBDrpgyLAgMhhESBs8+Nqmt6TC9Ji3VTRkSBgRBCoqDqWg9cbg9maCkwEEIIgbcuEgBMKUyJbUPCQIGBEEKioKndArVCBIVUGOumjIgCAyEk4VU19ODjEw2xbkZIjToTuo12NHWYkZ+piHVzwkK1kgghCe//DtTg7JUOrJg7CTxefNzvMgyD3/zlAj46dg3TilPR1G7Gl27Ii3WzwhIfv0FCCBmHuhYj+lweNHdaYt0Ult3pwkfHrkEm5uNyfTdsDlfC9BgoMBBCEprJ2osugx0AcLXVFOPW9DOYnQCAZQN6CXkUGAghZOz0ZgfOXekY8bz6FgP772utxgls0ejofYFhXnkWVHLvgDP1GAghZBz+d381Nm4/Bl23ddjz6lu8wSBdI4mvHoPFGxhSVWLML8+CRiGCWiGKcavCM2JgOHbsGH70ox9h9erVmDlzJpYuXYp169bhypUrQeceOXIEX/3qVzFjxgwsWrQIGzduhMkU/EFZrVY8++yzWLJkCWbMmIG7774bBw4ciMwVEUISksfDBPxcebUHALDvZOOwz6trMSJNLcH0kjRcjaMegz+VpJaL8PAdFXjh0Ztiuo/zaIwYGP7whz+gtbUVa9euxY4dO/CjH/0Ira2tuPfee3H+/Hn2vBMnTuCRRx5BVlYWXnvtNfzwhz/EwYMH8cgjj8Dj8QS85rp167Br1y6sX78er7/+OrRaLdatW4dDhw5F/AIJIfHv5GUd7v7hLmz53/Po0Ntgc/ThWqsRHA6w/2Qj3G5PyOe199hw5ot2lOVrUJSjgt7shNF3px5rBrMTHA6glAkhFQuQlSqLdZPCNuJ01WeeeQapqakBx5YsWYIVK1bgt7/9LbZs2QIAePHFFzF58mS8/PLL4HK98SY9PR3f/va3sWfPHtx6660AgEOHDuHo0aPYunUrVq1aBQBYuHAhmpqa8Pzzz2PZsmURvUBCSPz76Ng1CAU8HDzdhIOnG7FkZi48DLBqfj72nWxEdaMBU4sCVwwzDIMX3zkNAFh7ezmaO7wzktq6rFDJY5+yMVicUMqEcTN9djRGbPHgoAAASqUSBQUF0Ol0AID29nZcvHgRd9xxBxsUAODGG29EZmYm9u7dyx7bt28fFAoFVqxYwR7jcDi46667UF9fj9ra2nFdECEk/tU2GfDfbxzH1j+fx+fVnTh7pQO3LynC9qdWoqI4DZ+ebQaHA9x2YxEAoLHdDAA4cqEVe49fA+Ad3L3SqMdXV5YhK1WGTN9uaLoeW0yuaTCD2QF1HASosRhTKOvp6UFNTQ0mT54MAKiurgYA9ueBSktLUVNTw/5cU1MDrVYbEEAAoKysLOC1CCHJqbXTgg2//gzVjXp8cqYZ//X6UXg8DL50Qx7SNRJs+OYcaBQiFGQpUZijglDAQ3OHGScrdXjh7VN4a/cXYBgGbV3eQemCbO9MH/82me0jDFZHi8HsTJjB5sFGvfKZYRj85Cc/gcfjwUMPPQQAMBgMAACVShV0vkqlQmVlJfuzwWBAYWFhyPMGvhYhJDldqO2C28Pgl+uWQC0X4a+f1sLZ60Z+lhIAoJKL8Pz3l8DDMOBxOchLl6NRZ8ahs83gcjkw23qhNzvZ2UrZvty9SMBDilIMXXec9BgsTpTlx3/BvFBGHRheeOEF7N+/H8899xxKSkoCHhtqxH3w8eFG5hNl1J4QMja1zQbIJQLkpsvB4XDw4K3lQefkpMvZf+dlynH0QitcbgYr5k3CgVNNuNZmQlu3FVwOkK6RsudmpUqh66Eew3iNKpW0efNmvPnmm3j66adx9913s8fVajWA0Hf7RqMxoCehVquHPA8I3esghCSPmiYDtJPUYd8E5mcq4HJ7p7LetUwLAGhoM0HXZUOaRgoBv/9rLCtVxvYYDGYnO2U02hxOFxy97uQPDK+88gpee+01PPnkk3jwwQcDHvOPLQwcS/Crrq4OGHvQarWoq6sLmsLqH1soLS0Nv/WEkITx5wPV+OmOY2hoM2HyJHXYz/OXkSjMVqIgWwmNQoRrbSbouq3ITpUGnJuVIkW30Y4+lxub3j+DV/50LpKXEDb/4rakHnzeunUrtm3bhvXr1+Phhx8OejwrKwsVFRXYtWtXwBf+sWPH0N7ejltuuYU9tmrVKphMJhw8eDDgNXbu3ImioiJotdqxXgshJE71uTz44NM6nKnqgNvDQJunDvu5/jISs8syAAAF2Uo06LyppMFrAzJTZWAYoENvR4POzNZQirZuowMAoFEmZmAYcYzhzTffxJYtW3DzzTdj8eLFAYvahEIhysu9+cENGzbgoYcewhNPPIH77rsP7e3teOmllzBz5kysWbOGfc6yZcuwYMECPP300zAYDMjLy8POnTtx5swZbNu2LfJXSAiJudNftMNs60V2mgxtXVZMnqQJ+7m56XLcv2YKls+dBMDbc9h1uB5uD8MOPPtl+XoQjToTekwOcLmxGbP0r6nIHTBWkkhGDAyffPIJ+1//v/1yc3PZO/9Fixbhtddew5YtW/DII49AJpNh5cqVePLJJ8Hj8djncDgcbNu2DZs2bcLmzZthMpmg1WqxdetWLF++PJLXRgiJEwdPN0KtEGHT+qWoaTIgXSMJ+7lcLgdfW1XG/nzTrFzsPd4Au9OFrLTAwOD/Ij79hbf4ntXeF4HWj15zhxlCAS9gYDyRjBgY3nnnnbBfbOnSpVi6dOmI58nlcmzcuBEbN24M+7UJIYnJaHHi9BftuH1JMeRSIZsSGqvSfA1eeeJL2HeyAXMGvZZKLkKKUoyTl72Lb+1OF9xuT9RXHze2m5GXLgcvRj2W8Uq8tdqEkIiy2vvw5K8/w9qf7cXuf9ZH/PUPn2+By82wqaBIyE6T4cFbyyEWBd/bFuYo2cFfALA6XBF733A1t5uRl5mYaSSAAgMh172Pjl1DVYMebg+DT882R/z1D5xuQnGOCkU50ZmKXpStDPg5mumkBp0JRy+0okNvT5i9F0KhPZ8JSQI9JgfEQh6kYsGontfb58aHn9VhVmk6inJU2HW4Hn0uNwR83shPDkO30Y7aJgP+7fZpEXm9cBTGMDD88u3TaPLVdUqU3dpCoR4DIUngR6/+E9v+78Kon3eyUge92Ym7v6TFlAINXG4Pu/FNJNQ2GQAAUwujVxqi0NczkUu8QTKagWHgmr1JGZRKIoTESLfRjrYuK05W6tDnCr1vAQBY7H346Y5jaO20sMcu1XVDLORhhjYNZQXeKaRVDfqIta222QguByjKVY58coTkpsshEvLY6+k22XHg1NB7OkRSn8uDNJUYX7ohL2GnqgKUSiIk4dX47srtThcq67sxszQ95HmnKnU4U9WBi3Xd4HI5YBjgcn03phSkgMfjIlUlQZpagr9+UoNTlTr89N8XgT/O2Ty1zQbkZSogFkbvq0bA5+KFdTeBy+XgTFUH9p9swsW6LtidLty+pHhC39tgdmLV/Hz8+53Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(train_y)\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 9,
|
||||
"id": "1991167e",
|
||||
"metadata": {
|
||||
"tags": []
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"## train vol GP ## \n",
|
||||
"vmod, vlh = TrainVolModel(train_x, vol, train_iters=500, printing=False)\n",
|
||||
"\n",
|
||||
"## train data gp ##\n",
|
||||
"dmod, dlh = TrainDataModel(train_x, train_y, vmod, vlh, vol, train_iters=500)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 10,
|
||||
"id": "a1b37e3b",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYsAAAEACAYAAABCl1qQAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAACEXUlEQVR4nO29eZwkV3Xn+8t932vJ2npBJbWEVhC4BQLaSMhgxh5oYWbGBvH0DF6QWiM/lvHwZITHA2PB2MiA3EhjeUPGnmcbqzEwlhESiEVLIwEG7S31UmtW5b5HZERGvD+yz60bkZFLVWV2VVff7+ejj7pyiYyIyjrn3rP8jk3XdR0CgUAgEHTBvtUnIBAIBILtj3AWAoFAIOiJcBYCgUAg6IlwFgKBQCDoiXAWAoFAIOiJc6tPYNBIkoSnn34ao6OjcDgcW306AoFAcFbQbDaRTqdxySWXwOv1tj2/45zF008/jfe85z1bfRoCgUBwVvLlL38Zr3nNa9oe33HOYnR0FEDrgpPJ5BafjUAgEJwdpFIpvOc972E21MyOcxYUekomk5ient7isxEIBIKzi07he5HgFggEAkFPhLMQCAQCQU+EsxAIBAJBT/rKWVSrVdx555144IEHUCqVMDs7i5tvvhnXXntt1/c9+eST+MpXvoJnn30WL730ElRVxQsvvNDx9fPz8/j85z+PRx99FMViEaOjozhw4AB+//d/f10XJRAIBILB0pezOHToEJ599ll85CMfwfT0NO6//34cOnQId999Nw4cONDxfY8//jiOHj2Kiy++GE6nE08//XTH1z7//PN43/veh0suuQQf//jHEY/HsbS0hOeee279VyUQCASCgdLTWTzyyCN49NFHcdddd+G6664DAFx11VWYn5/HHXfc0dVZ3HTTTTh06BAA4FOf+lRHZ6HrOj760Y/iVa96Fe6++27YbDb23Dvf+c71XI9AIBAIhkDPnMWDDz6IUChkCDnZbDYcPHgQx48fx0svvdT54Pb+UiJHjx7Fiy++iPe///0GRyEQCM4cxWIR8/PzW30agm1KT2t+7NgxzM7Othn+ffv2AQBefPHFTZ/ED3/4QwCApmn41V/9VVxyySV47Wtfiw996ENYWVnZ9PEFAkFv8vk8ms0mZFne6lMRbEN6OotCoYBIJNL2OD1WKBQ2fRKrq6sAgFtuuQWvetWrcO+99+KjH/0oHn30Udxwww2o1+ub/gyBQNAdj8cDAOLvTWBJXwnubqGhQYSNaLLrL/7iL+K//Jf/AqCVFxkbG8Nv/dZv4etf/zre/e53b/pzBAJBb5rN5lafgmAb0nNnEY1GLXcPxWIRACx3HeslGo0CAN74xjcaHr/66qvhcDjwzDPPbPozBAJBd8hJlMtllMvlLT4bwXajp7OYnZ3Fyy+/DE3TDI9TruKCCy7Y9En0Oka/iXKBQLBx+L/xbDYrchcCAz2t8HXXXYdSqYSHH37Y8PiRI0ewd+9ezM7Obvok3vSmN8Hr9eKRRx4xPP69730PzWYTl1122aY/QyAQdEbX9bYF4fLy8hadjWA70jNnceDAAezfvx+33XYbCoUCpqenceTIETz11FM4fPgwe90NN9yAo0ePGjq0c7kcjh49CgCYm5sDADzwwAMAgKmpKVx66aUAWqGsm2++GXfeeSeCwSDe9KY34eTJk/jc5z6HCy+8EG9/+9sHd8UCgaANs6MQCMz0dBY2mw2HDx/GZz/7Wdx5551M7uOuu+7CNddc0/W9x44dw6233mp4jH4+ePAg7rjjDvb4b/7mbyIUCuG+++7D3/zN3yAcDuMXfuEX8OEPfxhut3sj1yYQCPqEnIXL5YKiKOxxXddF75MAAGDTqRRph7CwsIBrr70WDz30kJhnIRD0iSzLWF5extjYGGq1GiqVCgBgenoaTueOG3sjsKCX7RSZY4HgHEBRlK4Ja1oz2u12Q0GJKKMVEMJZCATnAIuLi10T1hSGstls4IMNqqoO/dwEZwfCWQgEO5x+Is38zoLPEYqdhYAQzkIg2OE0Go2er+F3FqFQCBMTEwCEsxCsIZyFQLDD6cfgk7OgfIXH44HD4RBhKAFDOAuBYIfTTw8FhaH4Mlm73Y5qtcqkfQTnNsJZCAQ7HD5n0Sl/oWkabDabwVlQv0UvFdpmsyma+s4BhLMQCHY4vCHv5Cx0XW/TYAuHwwDQsyl2fn6ejRkgRPhq5yGchUCww+EdRKcdAO0seOLxOOx2e9dqKnIKkiSxx0qlEhYWFvpKrAvOHkRrpkCww+lnZ6FpmqW6s7nvwgw1+jkcDmSzWYO0uaqqQqpnByF2FgLBDqebs9B1HfV6HZIkWRr2Xs6Cdg9Op7NtBobIY+wsxM5CINjhdAtD5XI5ZuRDoVDbe3s5C0qCW71G9GjsLMTOQiDY4fAOolqtIpvNMuPOVzo5HI629/bKWZCzsApjCWexsxDOQiDY4fCGnEamkiE391WY6baz0HXdsLOg909NTcHpdApnscMQzkIg2OHout62a7DKJ6zXWZAzsNvt0DQNuq4jFArB5XLB4XAIZ7HDEM5CINjBKIqCRqMBr9dreLzf5LPNZjO8VlEU5jzIGTidTui6jmazyRyOw+EQCe4dhnAWAsEOplQqwWazIRKJGB5fj7PgncPi4iJyuRz7GWhN1+NfD6ztNgQ7B+EsBIIdiq7rqFQqCAQCcDqdGB8fh9/vZ8/1AzkLTdNQrVYBrCXFrZwF7SzsdrsIQ+0wROmsQLAN0HUdy8vL8Pl8CIVCAxll2mw2oes6PB4PAMDn88HtdqNWq7FVf6/Vv81mQ7PZxNzcHHMK9B7q3ubPlZyFpmlQVVXM8N5BiJ2FQLANaDQaaDQaKBaLSKfTbc/ruo58Ps+qj/qBXmsVJtI0DbIs91z9WwkLapqGTCaDYrEIh8NhcBaUSF9ZWUE6nRahqB2E2FkIBNsAXkfJysCSI5FlGclksq9jWjkLWvlXq1Xk83kAQCQSQTAYtDxGp11BpVIBAFb9RNC/+R2GVf+G4OxDOAuBYBvAOwuqLlpeXkYgEECj0WC5hvWEdOr1Oux2e5uxttvt7PPC4TBisVjHY5g/z+l0svDT2NgYfD6f4TW0y6DHFEUxOJPtRrVaZaE5cxGAwIhwFgLBNqDZbMLtdjNDfurUKQBrToSSy1a9EFaoqop6vY5oNNrxNS6Xq6eBNH+ey+VizsLsKAi+QU+WZebothOapmFlZYUJIQJAMBgUu6AuCGchEGwDqEehVxVRvzsLOoaVOGA8Hocsy4jFYj2dDyXHCd6Y8ucyOTnJ/s13jMuyjGKxCI/H09brsZWUSiWDowBEyKwXwlkIBNuAZrPZZpit6JQwpml1zWYTmUyG7SisjF8wGOyYozBjdjZer5flKzq9rtlsss+VZZnNutizZ09fn9kLKuVdr2EvlUrQdR2RSMTS6YpkfHeEsxAItgFkYDv1P8zMzGB1dbWjQVtdXYUsy3C73VBVlYWtNrtSttlsmJmZYf0WDocDmUym63tokJLD4UCj0RhozqJcLiObzcJms2FycnJdx6ZmwkgkYrmjWl5exvT09EDKlncionRWINhiSFfJ4XAY5DKmpqbYa+g5SZJQq9UArOUlgLXcBv2fVvODCKvQZ/d7LF4GpNf87vVCoSNqONwonZwyP/FPYEQ4C4Fgi6HdAuUsgJahpVUzxfrJwNG865WVFaysrFjOz6bHhtEQNzExgenp6Y7Pk7PweDzMsQ0KTdPgdrvhdrvXNbbV7BzonptnePRbQHAuIvZbAsEWQ8aVV2qlHMD09DQzYFSFZP650WhYJsWHlaztllup1Wos3EN0Gtm6EXixwvXkGOheEeQ8EokEHA4HCoXCQM5vJyPcqECwxZDhstlsbCfg8/kAtHYYZBzpdRRTJ2dAoRNztdFWzL/OZrPs35RkH2TimHI761W15cNhuq4bdmN8+fDZkuSuVqs4depU3xpfg0A4C4Fgi6E/eLvdjnA4jLGxMcvehPHxcQBGZVdgzRCancVWlIHyyWFq9huGs7Db7ahUKigWi329j3+dpmksCQ8YS4DPpPHdDOl0Grquo1QqbSp3sx6EsxAIthgyprSz6NTE5na7EQgEoGkaVldXWcyeZD348JDL5bKcqT1seAdF/x6U+iy/I6hUKsjn8yzkJcsyTp48aamdRSXFlAOi42zH8tlarcZmovdDPp/vWZ02KETOQiDYYvgwVC9oGBGfOFZVlanDejweRCKRLeuaJmM7OTkJh8MBj8fDHtusAi2f2yGDSveOdK5kWW4rp+WdqaIorPqMz6OMjo6y1fpWQsUL3Ry91TmeCXVfsbMQCLaY9TgLqw5vSnBTr8V6qoQGjaZpTArdbrcjGo0iEAiw5zZ7bACGMl4KKXW7ZnqOdl7mMBQABAKBriNkzzTd1IXNyfperx8UwlkIBFtMv87CqkTWZrOh0Wiw3QVgbUzOFHzlk81mazPsAJBKpdYVauGPDcBQEqxpGvL5PHtOluWOZbJ8GMrsLOh8tzoMRXRzfla/X2rCHCbCWQgEWwyf4O6EpmmYn5836BnNzMwgHA5D13UUi0WW89jKCXXmMlm+d0RRFCiKAkmSDFVT6zk2sOaE6DHe8ZTL5bakNzkGek+pVEKj0WhzDNthZ0HnSL9DmqHOY/79ut1uJmUy1HMb6tEFAkFPeCNo9VytVmPGjX+Nw+GA2+02PO52u7fMWdCK3ewsqEJqdXUVpVJpw8cn3atGo8GaFq1W2eZGQOrNoB2OeSwsf67bZWdB57G4uIilpSXDc+ZrDoVC0HUd1Wp1qOEo4SwEgi2mW3KyUqlgdXWVJXBdLhfi8ThGR0cBtGLtiUSChWZcLteWGTz6XL4iym63Q9d11jfCG/L1OrVarYZMJoPjx4/DbrfD7/dbHsN8/eTAzDu3sbExw8/bYWdBn2++Bv46m82m4ftC4bVMJrOhHVu/CGchEGwx3ZwFH4unn8PhMEsaA2BVR06Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(train_x, vol)\n",
|
||||
"vmod.eval();\n",
|
||||
"samples = vmod(test_x).sample(torch.Size((10,))).exp().detach()\n",
|
||||
"plt.plot(test_x, samples.T, c='gray', alpha=0.25)\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 11,
|
||||
"id": "0c218e4f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"nvol = 100\n",
|
||||
"npx = 100\n",
|
||||
"px_samples = torch.zeros(npx*nvol, len(edays))\n",
|
||||
"px_paths = torch.zeros(npx*nvol, fulltest)\n",
|
||||
"vol_paths = torch.zeros(nvol, fulltest)\n",
|
||||
"dmod.vol_model.eval();\n",
|
||||
"dmod.eval();\n",
|
||||
"\n",
|
||||
"for vidx in range(nvol):\n",
|
||||
"# print(vidx)\n",
|
||||
" vol_pred = dmod.vol_model(test_x).sample().exp()\n",
|
||||
" vol_paths[vidx, :] = vol_pred.detach()\n",
|
||||
" \n",
|
||||
" px_pred = dmod.GeneratePrediction(test_x, vol_pred, npx).exp()\n",
|
||||
" px_paths[vidx*npx:(vidx*npx + npx), :] = px_pred.detach().T\n",
|
||||
" px_samples[vidx*npx:(vidx*npx+npx), :] = px_pred[ntests-1].detach().T\n",
|
||||
" "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 12,
|
||||
"id": "6b410709",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAABCIAAAGACAYAAAB4PcMRAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdd3hUZfrw8e+ZzEzapCcESOgllBg6ghSl2EUBabLqoqv+VsVddV3Fiq/urm5RVwHrKqCiggJWmtIFQpfeIYRU0vv05/3jyUwyZFJJCOX5XJeXZOacM88kkDnnPnfRhBACRVEURVEURVEURVGUC0DX3AtQFEVRFEVRFEVRFOXKoQIRiqIoiqIoiqIoiqJcMCoQoSiKoiiKoiiKoijKBaMCEYqiKIqiKIqiKIqiXDAqEKEoiqIoiqIoiqIoygWjAhGKoiiKoiiKoiiKolwwKhChKIqiKIqiKIqiKMoFowIRiqIoiqIoiqIoiqJcMCoQoSiKoiiKoiiKoijKBaMCEYpykdq6dStxcXHExcVddK/f0OcuByNHjiQuLo4lS5Y091LqbcaMGcTFxTFjxowr8vUVRVEuV1fa79dZs2YRFxfHPffc09xLqbfmPk9q7tdXFBd9cy9AUS5WL7zwAl9//TWhoaFs3LgRo9FYp/2uv/56kpOTGTFiBO+//34Tr7L+CgsLmT9/PgC///3vCQ4OvmCvfejQIX755ReCgoKYNm3aBXvdc9lsNr777jtWrlzJ4cOHycvLw9fXl8jISFq0aEHv3r3p378/gwYNwtfXt9nWeanYunUr27ZtIyYmhvHjxzf3chRFUS4qs2bNYvbs2VUeNxqNhIWF0aNHD26//XZuvvlmNE1rhhU2nz179rBo0SJ27dpFRkYGNpuNiIgIIiIiiIuLY8CAAQwePJhWrVo191Ives15fqcoDaECEYpSjQkTJvD111+Tn5/PL7/8wi233FLrPtu2bSM5Odm9/8WosLDQfUI0bty4aj+o/P396dChQ72PX9N+hw4dYvbs2cTExDRbICI9PZ2HHnqIo0ePuh8zGAz4+PiQnJxMUlIS27Zt48MPP+TTTz/l6quv9ti/TZs2GI1GgoKCLvTSL1rbtm1j9uzZDBw4sMZARFRUFB06dCAqKuoCrk5RFOXiERkZ6f5zUVERmZmZZGZmsnbtWpYuXcqcOXPqfOOjskvt96sQgn/84x98+umn7sc0TSM4OJjc3FwyMjI4cOAAS5YsYdy4cbz++use+4eFhdGhQwcVoKikqc/vFKWxqUCEolSjd+/edO7cmePHj7NkyZI6BSJc6fqRkZFcd911TbzCppWQkMCKFSsu2H4XgsPh4JFHHuHo0aP4+/vz0EMPMW7cOFq2bImmaVitVg4fPsyGDRv47rvvvB7DdbdBqb+//OUv/OUvf2nuZSiKojSbTZs2uf/sdDo5ceIEr732Gps2bWLDhg289dZbPPPMM/U+7qX2+3XevHnuIMSoUaN48MEH6dmzpzsIc+bMGbZu3cqKFSvQ6apWkt99993cfffdF3TNl4uL+TxNubKoQISi1GDChAm8/vrrbNq0iYyMDFq2bFnttsXFxaxcuRKAO+64A71e/fO62CQmJnLw4EEA/v73v3Prrbd6PG80GklISCAhIYFHH30Um83WHMtUFEVRrgA6nY4uXbrw3nvvMWbMGE6fPs3ChQv5y1/+clmfQwghmDt3LgDDhg3j3XffrbJNmzZtaNOmDRMmTMBsNl/oJSqKcgFcvr/lFKUR3HHHHbzxxhvYbDaWLl3Kww8/XO22y5cvp7S0FIA777zT47mDBw8yb948tm/fTnZ2Nn5+fnTq1ImbbrqJqVOn1jsN0+l0snv3btauXcu2bdvIyMggNzeXwMBAunTpwq233sqECRMwGAwe+91zzz1s27bN/fWoUaM8nh84cCCfffYZIOv+7733XgCOHDlS57VVt1/lpkipqalVmiRNnz6dRx55hBEjRpCZmclTTz3Fgw8+WO3rfP3117zwwgsEBATw66+/EhgYWOvaDh065P7zue/9XJqmef25jBw5ktTUVF577bUqZQiu9/Tpp5/StWtX3n//fVavXs3Zs2eJiIhgxIgRTJ8+nfDwcEB+Hz788EM2btxIVlYWERER3HTTTUyfPh2TyVTltV0/v+nTp/PYY495XberFrnyz7IuioqK2LBhA2vWrOHo0aNkZmZSVlZGZGQkffv25Z577qF3794e+6SkpHh8H7dt21bl51r5+zRjxgyWLl3qNc3WZevWrSxYsIDdu3eTl5dHYGAg3bp14/bbb2fs2LH4+PjU+p63bNnC3Llz2bt3LyUlJcTGxnLrrbfy4IMPVtvzY+PGjSxcuJC9e/eSm5vrrt1u164dQ4YM4c477yQ0NLTO309FUZS68vX15aabbuKDDz6gpKSEkydP0rVrV4/fsatXr8bpdPLRRx+xadMmzp49S4sWLVizZg1Qt9+v6enpfPbZZ2zatImUlBRsNhstWrSgS5cu3Hjjjdx8881ef0cePHiQzz77jO3bt5OVlYVOp6NNmzaMGDGC3//+9+7PtLrKy8sjMzMTkJ+ptfHz86vyWE2fded+L5YsWcLChQs5fvw4Pj4+9OjRg0cffZQBAwYAYLfb+fLLL1m6dClJSUlomkbfvn15/PHH6dmzZ5XXXrJkCc8++ywxMTHu7/+5zv3ZxcbG1vo+4eI5v8vKyuKTTz5hw4YNpKWlIYQgJiaGa6+9lvvvv9+jxKi69+zn58f777/PmjVryMrKIigoiKuvvprp06fTqVMnr6+bkZHBJ598wqZNm0hNTcVutxMaGkqLFi3o378/t912GwkJCXX6XioXPxWIUJQahIeHM3LkSFauXFlrIGLx4sUA9O3b1+MX7Lx583j99dcRQgAQFBREWVkZu3fvZvfu3SxZsoT//e9/tGjRos7rSktLY+rUqe6v9Xo9fn5+5Ofns337drZv386PP/7Ixx9/7PEBHhISQlhYGHl5eYCssax8YRcSElLnNdRXZGQkZrOZ4uJidDpdlROXgIAAfHx8mDhxIrNnz+abb77hgQceqLZx19dffw3AbbfdVqcgxLkyMjJo3759vferi/T0dJ5++mkyMjIICAjA6XSSlpbGggULSExM5KuvviIpKYmHHnqIvLw8TCYTTqeT9PR05s6dy549e/j888+9XnQ3lXnz5nk0UwsICADk37W0tDR++uknnnvuOffJC4CPjw+RkZGUlpZSWlqKwWCo8nfI2wlkdV577TXmzZsHyEBQUFAQRUVFJCYmkpiYyPfff8+cOXO8Bmlc/ve///Gf//wHkP/WbDYbJ0+eZNasWWzbto25c+dW+b7Onj2bWbNmub/29/dHCEFKSgopKSls2rSJ+Pj4Kv1CFEVRGkt0dLT7z8XFxVWe3717Ny+99BKlpaX4+/tXuRCtzbfffstLL72ExWIBZG8kPz8/zpw5w5kzZ1izZg1xcXF0797dY7933nmHd999130O4+/vj81m48iRIxw5coTFixfz4Ycf0qNHj/q+ZQB3QKKpuIISer0eX19fCgoK2LJlC9u3b2f27NkMGTKEhx9+mF9//RWDwYDBYKCkpIQNGzawfft2Pv/8c+Lj45t0jZVdDOd327Zt49FHH6WwsBCQP3NN0zh+/DjHjx/nm2++4d1336V///7VHuP48eM899xz5OTk4O/vD0BOTg7Lli1jw4YNLFiwgG7dunnsc/jwYe69914KCgoAeY5hMpnIzs4mKyuLAwcOUFhYqAIRlxE1vlNRauHKbjh9+jTbt2/3us3JkyfZvXu3x/YAa9eu5bXXXkMIwahRo/jll1/YsWMHu3bt4p///CeBgYEcOXKEP/3pTzgcjjqvSa/XM2rUKN566y02bNjAvn372LlzJ7t27eK1116jRYsW7Nixg7feestjP9cFvss333zDpk2b3P956+rdWDZt2sTzzz8PQKtWrTxed9OmTfzhD38AYNKkSej1epKSkti6davXYx05coQ9e/YAMHny5DqvofKH18yZM5vsBOjvf/87YWFhLFq0yB1wevPNN/H39+fEiRO8/fbbPP7448TFxfHjjz+6f3YvvvgiPj4+7Nq164KPB42MjGTatGksWrSI7du3s3v3bvbu3csvv/ziDj68/vrr7tIWqPg53n///QD06dOnys+1Lr1VAD7//HN3EGLy5Mls3LiR7du3s2PHDp599ln0ej2JiYm8+OKL1R7j8OHDvPHGGzz00ENs3rzZvf+jjz4KyLtAS5cu9dgnNTWVOXPmAHDfffexYcMGfvvtN3bv3s2OHTtYsGABU6dObVCwS1EUpa5SU1Pdf/Z20fjSSy/RpUsXvvnmG/fvqI8//rhOx16/fj0zZszAYrHQt29fFixYwN69e9mxYwc7d+5kwYIFTJo0qUpwY968ecyZM4eAgAD+8pe/8Ouvv/Lbb7+xZ88eFi9ezKBBg8jKyuLhhx+mpKSkzu81PDzcnSHgytBoCqtXr2b58uW88sor7s/Z5cuX07NnT+x2O6+++ir//Oc/2b9/P//973/ZvXs3u3btYvHixbRt25aysjL+/ve/N8naqtPc53fp6enuIETnzp354osv3H/fFixYQIcOHSgoKODRRx+t8Rzq6aefpl27dh5/X+fOnUtUVBTFxcW8+uqrVfZ5/fXXKSgooGfPnixcuJADBw6wbds29u7dy8qVK3nmmWfo3LlzHb+TyqVABSIUpRbDhg1z94ZwZT2cy/V4QEAAN998s/tx153Zfv36MWvWLNq0aQPIXgRjx451P797925+/vnnOq+pZcuWvPvuu9xyyy1ER0e7GzkFBgYyfvx4d73lokWL3Hc/LhXR0dHuRp+LFi3yuo3r8Z49e9brTsXAgQMZMmQIIPtFjBgxgilTpvCPf/yD7777jqSkpPNau4vRaGTu3Ln06tULkHeebr31Vu677z5AXnQHBATw0Ucf0aVLF0CLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 1200x360 with 2 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"colors = [\"#1b4079\",\"#3db2ff\",\"#ffedda\",\"#ffb830\",\"#ff2442\",\"#61210f\",\"#32373b\"]\n",
|
||||
"colors = [\"#0c4767\",\"#8A95A5\", \"#F4442E\", \"#566e3d\",\"#b9a44c\",\"#fa7921\",\"#fe9920\"]\n",
|
||||
"colors = [\"#01295F\",\"#437F97\", \"#F4442E\", \"#566e3d\",\"#b9a44c\",\"#fa7921\",\"#fe9920\"]\n",
|
||||
"fs = 16\n",
|
||||
"\n",
|
||||
"fig, ax = plt.subplots(1,2, figsize=(12, 3.6), dpi=100)\n",
|
||||
"\n",
|
||||
"ax[0].plot(train_x, vol, c=colors[0], label='GPCV Vol.')\n",
|
||||
"vmod.eval();\n",
|
||||
"samples = vmod(test_x).sample(torch.Size((20,))).exp().detach()\n",
|
||||
"ax[0].plot(test_x, samples[0], c=colors[-2], alpha=0.5,\n",
|
||||
" lw=0.7, label=\"Simulations\")\n",
|
||||
"ax[0].plot(test_x, samples[1:8].T, c=colors[-2], alpha=0.5,\n",
|
||||
" lw=0.7)\n",
|
||||
"\n",
|
||||
"ax[0].set_xlabel(\"Days\")\n",
|
||||
"ax[0].set_ylabel(\"Volatility\")\n",
|
||||
"ax[0].set_title(\"Volatility Simulations\")\n",
|
||||
"ax[0].legend(fontsize=fs-2, frameon=False)\n",
|
||||
"\n",
|
||||
"ax[1].plot(train_x, train_y, c=colors[0], label=\"Train\")\n",
|
||||
"\n",
|
||||
"ax[1].plot(test_x, px_paths[1:20, :].T, c=colors[-1], alpha=0.5,\n",
|
||||
" lw=0.5)\n",
|
||||
"# ax[1].plot(test_x, px_paths[0, :].T, c=colors[-1], alpha=0.5,\n",
|
||||
"# lw=0.2, label=\"Simulations\")\n",
|
||||
"\n",
|
||||
"ax[1].plot(test_x, test_y, c=colors[1], lw=2., label=\"Test\")\n",
|
||||
"ax[1].plot(test_x, px_paths[0, :].T, c=colors[-2], alpha=0.5,\n",
|
||||
" lw=0.5, label=\"Simulations\")\n",
|
||||
"sns.rugplot(x=test_x[ntests-1], ax=ax[1], color=colors[2], height=0.1, label=\"Expirations\")\n",
|
||||
"ax[1].legend(fontsize=fs-2, frameon=False)\n",
|
||||
"ax[1].set_xlabel(\"Days\")\n",
|
||||
"ax[1].set_ylabel(\"Price\")\n",
|
||||
"ax[1].set_title(\"Price Simulations\")\n",
|
||||
"plt.savefig(\"./option_diffusions.pdf\", bbox_inches=\"tight\")\n",
|
||||
"sns.despine()\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 13,
|
||||
"id": "2bac7474",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"logger = []\n",
|
||||
"days = np.datetime_as_string(edays, 'D')\n",
|
||||
"for day in range(px_samples.shape[1]):\n",
|
||||
" for smpl in range(px_samples.shape[0]):\n",
|
||||
" logger.append([px_samples[smpl, day].item(), days[day][5:]])\n",
|
||||
" \n",
|
||||
"df = pd.DataFrame(logger)\n",
|
||||
"df.columns = ['Price', 'Date']"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 14,
|
||||
"id": "876976aa",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAgsAAAF+CAYAAAAMWFkhAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAACphklEQVR4nOydd3xUVfr/P+feKem9QUIJhAQQCEiX3hRYFEFsa0HFZX0plkX9qquy35+6dsEFZJG1rGtZ+VrAxYIKuCiCBSwgvQXSSJ9MMpNMuef8/rgzQyaZmcwkk7n3wnm/Xi47d87c+8zNnXOe81TCGGPgcDgcDofD8YOgtAAcDofD4XDUDVcWOBwOh8PhBIQrCxwOh8PhcALClQUOh8PhcDgB4coCh8PhcDicgHBlgcPhcDgcTkC4ssDhBEFJSQkKCgqwatWqgMfUxIMPPoiCgoKIXnPq1Km44YYbInrNzqLE31FNz5Pan2OOOtApLQCH44/vv/8eN954o9exmJgY5ObmYu7cubj++ushiqJC0nWOkpISbNiwAdOnT8eAAQOUFgdTp05FaWmp57Ver0dGRgbGjh2LJUuWoFu3bgpKFzwtlSNCCKKjo5GSkoL+/ftj6tSp+N3vfoeoqKiwXe/DDz+E2WzGTTfdFLZzdgVqe9442oMrCxzVM2fOHEycOBGMMVRWVmLDhg148skncezYMTz++OOKyZWdnY29e/d2SGEpLS3F6tWrkZ2drZrJOysrC0uXLgUAWCwW/PDDD/jggw+wfft2/Oc//0FKSkq759i8eXNXi9kuAwYMwM033wwAaG5uRllZGb799lv8+c9/xtq1a7Fq1Sr079/fM74zf8cNGzagtLQ0ZGWhM9fsCIGet0jLwtEmXFngqJ6BAwdi7ty5nte///3vMWvWLLz33nu4++67kZaW5vNzjY2NiIuL6zK5CCEwGo1ddv5IEx8f3+Y+P/7443jrrbfw4Ycf4tZbb/X5OYfDAUopjEYjDAZDpMT1S2Zmptf3AIA//elP+Oyzz3D//ffj1ltvxSeffILExEQAkf07up9JNT07apKFo154zAJHc8TFxWHYsGFgjKG4uBjAWV/5gQMHsGjRIgwfPhyXXXaZ5zNFRUW4//77MX78eAwaNAhTp07FM888A6vV2ub8u3fvxjXXXIMhQ4bgoosuwmOPPeZzXCBf7+eff44bbrgBI0aMQGFhIS655BI88cQTsNvt+PDDDz3ulYceeggFBQUoKCjw8vUzxvDOO+9g/vz5KCwsxLBhw3DDDTfgu+++a3Mtm82GZ555BuPHj8eQIUOwYMEC7NixI/Qb64Px48cDAE6fPg0AWLVqFQoKCnD06FE89dRTmDhxIoYMGYJffvkFgP+YhQMHDuCuu+7CRRddhEGDBmHSpElYunSp57xudu7ciVtuuQUjRozA4MGDcemll+Lf//53WL7LrFmzsGjRIlRVVeHtt9/2HPf3d9y4cSMWLFiAESNGYOjQoZg2bRruvfde1NbWer7rDz/8gNLSUs/fsKCgAN9//z0A4IYbbsDUqVNRXFyMu+66C6NGjcLw4cMDXtPNxx9/jEsvvRSDBw/G5MmTsWrVKjidTq8x7vO3pvW523ve/MnidDqxbt06zJ49G4MHD8bo0aNxxx134PDhw36v99VXX+GKK67A4MGDMX78eDzzzDNt5D569CjuuusuTJgwAYMGDcK4ceNwww034L///a/Pe8FRB9yywNEcjDGcOnUKAJCcnOw5XlZWhoULF2LmzJm4+OKLPQv8b7/9hoULFyIhIQFXX301MjMzcejQIbz55pv4+eef8eabb0Kv1wMAfv31V9x8882IjY3FH/7wB8THx+PTTz/FAw88ELR8K1aswNq1a5GXl4ebbroJ6enpOH36NL744gvcddddGDlyJG677TasXbsWV199tWcBaWkhuf/++/HJJ5/gkksuwfz582G327Fp0ybccsstWLVqFaZNm+YZu3TpUmzZsgVTpkzBhAkTcPr0adx5553Iycnp+E124es+A8B9992HqKgo3HLLLQCA9PR0v+f46quvcOeddyImJgYLFixAr169UFVVhR07duDIkSPo2bMnAGD9+vX4y1/+gqFDh+K2225DdHQ0du7cif/93//F6dOnQ/ob+OPKK6/E2rVrsX37dtx+++1+x3300Ud44IEHMGLECNx1112IiopCWVkZvv76a9TU1CAlJQV//vOf8cILL6Curg4PPfSQ57N9+/b1/H+LxYLrr78eF154Ie655x6PohGIr776Cm+88Qauu+46pKWlYdu2bVi9ejXKysrw1FNPhfydg3nefHHffffhs88+w7hx43Dttdeiuroab7/9Nq655hq8/fbbGDhwoNf47du345133sE111yDK664Alu3bsVrr72GxMRE3HbbbQCAuro6LFy4EABwzTXXoHv37qirq8Nvv/2GX3/9FZMnTw75+3EiBONwVMp3333H8vPz2apVq1hNTQ2rqalhBw8eZA8//DDLz89nV111lWfslClTWH5+Pvu///u/Nue59NJL2SWXXMIaGhq8jn/xxRcsPz+fffDBB55jV199NbvgggvYiRMnPMdsNhu74oorWH5+Plu5cqXneHFxcZtjv/76K8vPz2c33HADa25u9roepZRRSr2+W8trt5br3Xff9TrucDjYvHnz2JQpUzzn+eabb1h+fj574IEHvMZ++eWXLD8/n+Xn57c5vy+mTJnCZs6c6bnPp0+fZu+//z4bPnw4GzhwIDt8+DBjjLGVK1ey/Px8dv311zOHw+HzPNdff73ntdVqZaNHj2ZjxoxhZ86caTNekiTGGGMVFRVs0KBBbOnSpW3GPP7446x///7s1KlT7X6P/Px8tnjx4oBjhg0bxkaNGuV57evveMcdd7Bhw4b5/I4tuf7669mUKVP8vpefn8+WL1/e5j1f13Qf69+/P/vtt988xyml7Pbbb2f5+fns559/bvfavs4d6HnzNX7Hjh0sPz+f3X333Z5njTHGDh48yAYMGMCuvfbaNp8vLCxkxcXFXnL/7ne/Y+PGjfMc27JlC8vPz2effPKJz3vGUS/cDcFRPatWrcLYsWMxduxYzJ07Fx988AGmTp2Kl156yWtcUlIS5s+f73Xs8OHDOHz4MObMmQO73Y7a2lrPf8OHD0dMTAy+/fZbAEBNTQ1+/vlnTJ06Fbm5uZ5zGAyGoAPY/vOf/wAA7r333jZ+YEIICCFBnSM2NhbTp0/3ktdsNnuyFoqKigAAW7ZsAQAsWrTI6xzTp0/3+g7BcOLECc99nj59Ov785z8jOTkZa9asQX5+vtfYhQsXQqdr3zC5Y8cO1NXV4eabb0ZmZmab9wVBnoI+//xz2O12LFiwwOs719bWYurUqaCUYteuXSF9H3/ExcWhsbEx4Jj4+Hg0Nzfjv//9L1gnG/O2/tu0x0UXXYQLLrjA85oQ4okX+fLLLzslS7C4r3Pbbbd5PbP9+/fH5MmTsWfPnjZWkmnTpnlZswghGD16NKqqqmCxWADI9xUAvvnmm3b/Bhx1wd0QHNVz9dVXY+bMmZ5UuN69eyMpKanNuB49erSJ6D5+/DgAWeHw5x+urq4GAE/8Q58+fdqMycvLC0rWU6dOgRDiFW0fKsePH4fFYsFFF13kd0xNTQ1yc3NRXFwMQRDQu3fvNmP69u2LkydPBn3d7OxsPPHEEwDOpk726tXL51hf1/OFW6lpbbJujfvvFEgpc/+dOkswga9//OMf8eOPP+KOO+5AUlISRo0ahYkTJ2LWrFkhBc2mpKQgISEhJPlaujHcuJ8/9zPa1ZSUlEAQBJ+y9OvXD1u3bkVJSYlXhkyPHj3ajHX/Tk0mE2JjYzFq1Chcfvnl+PDDD7Fp0yYMGjQIF110EWbPnh30b4yjDFxZ4KieXr16BVw43URHR/t975ZbbsGECRN8vueezN07SF+7/2B3l4yxoKwH7Z0jJSUFL7zwgt8x/fr1C+o8oRATExPUfQYQdK2CQPfU17hnnnkGGRkZPsf4WoxCpaSkBBaLBcOGDQs4rnfv3vj000+xa9cu7Nq1Cz/88AMeeeQRrFy5Em+//bYnzqI9Aj2T/ujs8yNJUqc+D4T+7AAImHrZ8nzPPPMMFi1ahO3bt2PPnj14/fXXsXbtWvz5z3/G9ddf3yF5OV0PVxY45zTunbEgCO0uhO4FwL3LbYmvY77Izc3FN998g8OHD2PIkCF+xwVaEHr16oWioiIUFhYiNjY24PV69OgBSimKioraKBAnTpwISuauxG2lOXDgAMaNG+d3nNtSkZycHLTC0hHee+89AMCkSZPaHWswGDBp0iTP2O3bt2Px4sV4/fXX8Ze//KXLZDx27JjfYy0VpqSkJOzfv7/NWF/Wh1AVkJ49e2LHjh04fvx4GyuZ+7fQmQDa/Px85Ofn4w9/+APMZjOuvPJKvPDCC7juuus6rSxxugYes8A5pxk4cCDy8/Px7rvv+pxEnU4nTCYTACA1NRVDhw7Ftm3bvMz3drsd//znP4O63qWXXgoAWL58Oex2e5v33TusmJgYAEB9fX2bMZdffjkopVi+fLnPa7Q0x7uzIl599VWvMVu2bAnJBdFVjBs3DsnJyXj99ddRWVnZ5n33/Zg1axYMBgNWrVqF5ubmNuMaGhp83s9Q+Oyzz/Dqq68iIyMD1113XcCxvrIW3K6Uln+z2NhY1NfXdzquoSU7d+70UgIYY3jllVcAyLEobnr37g2LxYK9e/d6jlFKfT6rgZ43X7ivs27dOq/vduTIEWzbtg3Dhw8PqkhXa0wmEyilXscSEhKQk5ODpqYm2Gy2kM/JiQzcssA5pyGE4Nlnn8XChQtx2WWX4YorrkBeXh6am5tx6tQpfPnll1i6dKknMPLBBx/EDTfcgGuvvRbXXXedJ3UyWNPukCFD8Ic//AH/+Mc/MH/+fMyaNQvp6ekoKSnB559/jvfeew8JCQnIy8tDbGwLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 576x360 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(figsize=(8,5))\n",
|
||||
"violins = sns.violinplot(x='Date', y='Price', data=df, color=colors[-1])\n",
|
||||
"for violin in violins.collections[::2]:\n",
|
||||
" violin.set_alpha(0.5)\n",
|
||||
" violin.set_edgecolor(colors[-2])\n",
|
||||
"# violin.set_facecolor(colors[-2])\n",
|
||||
"\n",
|
||||
"# parts = ax.violinplot(dataset=px_samples.T,positions=range(px_samples.shape[1]))\n",
|
||||
"# for pc in parts['bodies']:\n",
|
||||
"# pc.set_facecolor(colors[2])\n",
|
||||
"# pc.set_edgecolor(colors[3])\n",
|
||||
"# pc.set_alpha(1)\n",
|
||||
"plt.scatter(np.arange(px_samples.shape[1]), test_y[ntests-1], color=colors[1], zorder=4, s=80,\n",
|
||||
" label=\"Observed Price\")\n",
|
||||
"plt.xlabel(\"Expirations\")\n",
|
||||
"plt.xticks(rotation=45)\n",
|
||||
"sns.despine()\n",
|
||||
"plt.legend(loc=\"upper left\",frameon=False, fontsize=fs-4)\n",
|
||||
"plt.title(\"Predicted Price Distributions\")\n",
|
||||
"plt.savefig(\"./distribution_plot.pdf\", bbox_inches='tight')\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 15,
|
||||
"id": "ef9e8f53",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def Pricer(mc_pxs, options, edays, true_pxs):\n",
|
||||
" logger = []\n",
|
||||
" for eday_idx, eday in enumerate(edays):\n",
|
||||
" eday = pd.Timestamp(eday)\n",
|
||||
" opts = options[options.expiration==pd.Timestamp(eday)]\n",
|
||||
" for idx, row in opts.iterrows():\n",
|
||||
" K = row.strike\n",
|
||||
" bid = row.bid\n",
|
||||
" ask = row.ask\n",
|
||||
" valuation = np.mean(np.maximum(mc_pxs[:, eday_idx].numpy() - K, 0))\n",
|
||||
" rtn = np.maximum(true_pxs[eday_idx] - K, 0)\n",
|
||||
" logger.append([eday, K, bid, ask, valuation, rtn.item()])\n",
|
||||
" \n",
|
||||
" df = pd.DataFrame(logger)\n",
|
||||
" df.columns = ['Expiry', \"Strike\", \"Bid\", \"Ask\", \"Voltron\", \"Return\"]\n",
|
||||
" return df"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 16,
|
||||
"id": "016d7372",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"option_output = Pricer(px_samples, options, edays, test_y[ntests-1])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 17,
|
||||
"id": "d4f6fcb4",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"2015-06-20T00:00:00.000000000\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"idx = 3\n",
|
||||
"print(edays[idx])\n",
|
||||
"dat = option_output[option_output.Expiry == option_output.Expiry.unique()[idx]]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 18,
|
||||
"id": "f9f405ee",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAZwAAAEWCAYAAABSaiGHAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAABaVklEQVR4nO3ddVgV6fvH8fehQ8JAURRFEWxEwe5eDEKwY1dXXXd1LezubnRVsBMVBBMDEwkFuxW7MBEVg/r9wcpv+YIKu8ABvF/Xda7lzMyZ85nZ47nPzDzzPIqEhIQEhBBCiEymouwAQgghfgxScIQQQmQJKThCCCGyhBQcIYQQWUIKjhBCiCyhpuwA2dHHjx+5dOkSRkZGqKqqKjuOEELkCHFxcTx//pwKFSqgpaWVYr4UnFRcunSJzp07KzuGEELkSBs3bsTGxibFdCk4qTAyMgISd5qxsbGS0wghRM7w9OlTOnfunPQd+r+k4KTiy2k0Y2NjihYtquQ0QgiRs3ztUoQ0GhBCCJElpOAIIYTIElJwhBBCZAkpOEIIIbKENBoQIpeIiori2bNnxMTEKDuKyKXU1dUpWLAg+vr6/+r1UnCEyAWioqKIiIjAxMQEbW1tFAqFsiOJXCYhIYEPHz7w6NEjgH9VdOSUWga7+vwqpvNN2XRxk7KjiB/Is2fPMDExQUdHR4qNyBQKhQIdHR1MTEx49uzZv1qHFJwMZmpgSgnDEnT27sywg8OIi49TdiTxA4iJiUFbW1vZMcQPQFtb+1+ftpWCk8F0NXQ51O0QfW36MjtwNnab7Hj94bWyY4kfgBzZiKzwXz5nUnAygYaqBktbLmVFqxUcuXMEW3dbLj+7rOxYQgihVFJwMlGvqr040v0I7z6/o8bKGvhc81F2JCFylREjRtCnTx9lxxBpJAUng334HIt38G3efUw8x1nbtDZhvcMoW6Asjp6OTDw6kfiEeCWnFEL5fvvtN37++edU54WHh2NpacnJkyfTtc6uXbsyadKkDEgnMoMUnAz2POojHv7XGLDyJPefvwXARN+E478cp7tVdyYcm0DbrW15++mtkpMKoVzOzs4EBwfz8OHDFPO2b9+OiYkJNWvWzJT3lnuVlEMKTgYzLZCHmV2q8/ZjDANWBRJ8IwIALTUtVtuvZkHzBey6vosaK2tw69UtJacVQnkaNGhAgQIF8Pb2TjY9JiYGX19fnJycCAsLw8XFhYoVK1KrVi2mTZvG58+fU13fiBEjOHXqFBs3bsTS0hJLS0sePnxISEgIlpaWHDt2DGdnZypUqEBAQACfP39m6tSp1KpVi4oVK9KuXTtCQ0OT1vfldUFBQbi4uGBlZYWTkxOXL8v12H9LCk4mqBj/BLefjCmST4cJnqFsOnGThIQEFAoFA2oM4EDXAzx99xRbd1v239qv7LhCKIWamhoODg7s2LGD+Pj/P8185MgRXr9+jaOjI7169aJs2bL4+PgwdepU9uzZw7x581Jd3+jRo7G2tsbJyYmAgAACAgIoXLhw0vw5c+YwcOBA9u3bh5WVFbNmzWLfvn1MmzYNHx8fLCws6NWrV4p7TObOncuQIUPw9vYmb968uLq6kpCQkDk7JZeTngYyw9aZFPTfwNxBq1lQwIq1R28Q/jQKV3srtDXUaGTWiNBeoTh4OmC3yY4ZjWfgWstVmrWKDHXw/EMOnH+Qpe/ZzKoYTa3SPoaUs7Mz7u7uBAYGUqdOHSDxdFrt2rXZunUrRkZGTJgwARUVFUqVKsWQIUMYN24cAwYMSHHfkZ6eHurq6mhra6c6AFi/fv2S3iM6OpotW7YwZcoUGjRoAMDEiRMJDg5m48aNDBo0KOl1AwYMoEaNGgD8/vvvdOrUiYiICBmc8V+QI5zMMMgdipdHa/4vDL80g16NyxB4/SmDVgfy5HU0AGZ5zQjsEUjbsm0ZdmgYnb07Ex0TreTgQmStEiVKYGtri5eXFwAREREEBATg4uJCeHg4lStXRkXl/7+mqlatSkxMDPfu3Uv3e1WoUCHp7/v37xMTE0OVKlWSpqmqqlK5cmXCw8OTvc7S0jLp74IFCwLw8uXLdL+/kCOczKGhBSsuwurRKLZMxzlkF2YzrzJt73X6eQQwum0VqpQsgK6GLp7OnlgHWDP68GiuvbjGjvY7KG5YXNlbIHKBplZF03W0oSzOzs6MHTuWyMhIduzYgYGBAY0aNWLnzp1fPer/N2cDUuuJIbX1/O80NTW1FPP+eQpQpJ0c4WQWhQJ6TINp++HzB6oOKsHilkXIr6fJ6E0heAffTrquM7LuSHZ13EX463Bs3G04dveYstMLkWVatGiBpqYmO3fuxMvLCwcHB9TV1TE3N+fcuXPJvtzDwsJQV1fH1NQ01XWpq6sTF/f97qRMTU1RV1cnLCwsaVpcXBznzp2jVKlS/32jRKqk4GQ2m2aw7g4ARQZWZEGZV9S0KMTyg1eZ7XueTzGJ/zhaWrTk1K+nyK+dnybrm7Dk1BK5MCl+CFpaWrRq1Qo3Nzfu37+Ps7MzAJ06deLZs2dMmDCB8PBwjh49yty5c+nSpctX+40zMTHh4sWLPHz4kFevXn31SERHR4eOHTsyZ84cjh07Rnh4OBMmTODly5d06tQp07b1RycFJysYlwDfd5CvMDrTnBjzZCVd65XG/+IjXNcG8ezNBwAsC1gS8msIzUs1p9++fvTa1YtPsZ+Um12ILODi4sKbN2+wtrZOOsIoVKgQ7u7uXL16FXt7e0aNGkXLli0ZPHjwV9fTo0cP1NXVadmyJTVr1uTx48dfXXbo0KH89NNPjBw5Ent7e65fv467u3vSdRqR8RQJ8jM6hYcPH9K4cWP8/f0pWjQDz4EnJIDbH7DrL8hXmMDRQczaexVNdVXGOlelgmk+AOLi4xh/dDxTT0ylZtGaeLXzorBe4e+sXPzIrl69StmyZZUdQ/wgvvZ5+953pxzhZCWFAvovhXHe8OoJtYaUYGHrEuhoqjF8fTB7whJb3qiqqDKl0RS2uWzjfMR5qq6oSsjDECWHF0KI/0YKjjLUcYTVNwAoPrAsiyp/oLJZARbtvcTCPReJiUs87+xczpmgnkFoqWlRb009Vp9drczUQgjxn0jBURaT0rDjDWjroTexJZPebsOlZkn2nrnP8PXBvH6XeO2mUqFKnO51mrqmdemxswd/7vuTmDjpB0oIkfNIwVEmXX3YEQmNu6DqNYdfN7RhRKty3Hryhn4rA7jxOBKA/Dr58evix6Aag1h8ajHNNzTnRfQLpUYXQoj0koKjbCoqMHw9jNgIj2/RcHhJ5tmXQkWhYMjaIPwvJPakq6aixrzm81jrsJbAB4HYrLDh3NNzys0uhBDpoNSC8/TpU6ZMmULHjh2xtrbG0tKSkJCUF8e7du2a1PvrPx//7O/oi/fv3zNlyhTq1KlDpUqVcHJywt/fPys2579p1AlWXALAfKAli2soKGNiyCzf8yw/eIW4v+8n6GbVjRO/nCA2PpZaK2vheclTmamFECLNlFpw7t27x549e9DR0UnqHO9rSpQogaenZ7LHwIEDUyzXr18/du3axYABA1i+fDnm5ub069ePY8dywN37JcqD1ysADMc0YrqKP21si+MdfIfRm04TFZ3YLbutiS2hvUOpUrgKHbw6MPLQSOLiv393tRBCKJNS+1KztbUlKCgIgEOHDnH48OGvLqulpUXlypW/ub5jx44RGBiIm5sbTZs2BaBGjRo8ePCAGTNmUL9+/QzLnmn08sK+WJjggNq6sfxhsRPznz1ZfOAa/VcGMKGdDWaF9DHOY8zh7ofpv7c/M07O4HzEeTa13YShlqGyt0AIIVKl1COcf/YCmxEOHjyInp4ejRs3TpqmUChwdHTk9u3b3LqVQwY8U1WFybtgwDK4cZrmo0oy29GCz7HxDFwdyImrTwDQUNVgeevlLGu5jIO3D1LNvRpXn19VcnghhEhdjmk0cOfOHWxtbSlXrhzNmjVj6dKlKYaJvXnzJubm5ikK2ZfuxW/cuJFleTNEyz7gljgCYdkB5rg10KVEQT2mbD/D2iPXif+7k4g+Nn040v0Ibz69obpHdXZd36XM1EJkS40aNWLlypXKjvFDyxEFp2rVqowcORI3NzeWLFmCra0tixYtSnENJzIyEgMDgxSv/zItMjIyC9JmMIuq4Jk4THX+4bWYbXiOZlZF2RRwi4meobz/lFh065jWIbRXKBb5LbDfYs+U41OIT5Au1EX2lVpDoH8+RowY8a/W6+3tjbW1dQanFRkhR4yH87+FpWHDhhQoUIBly5YRGhqKjY1N0rxvjZORY0fUzFsQ9n6GoY3QWNaPwVWbY+6whGUHrzNg5UkmtLehaP48FDMoxolfTtB7d2/GHhnL2adnWeuwljwaeZS9BUKkEBAQkPT30aNHGTNmTLJpWlpayZaPiYlBXV09y/KJjJcjjnBS4+DgAMC5c+eSphkaGqZ6FPPmzRuAVI9+cgw1dZh/AnrORBG2H/ux5sxwKkPUhxj+XHmSUzcTx2HXVtdmncM65jSdg881H2qurEn4q/DvrFyIrGdkZJT00NPTSzbt06dP2NjYsHv3brp160alSpXw9PRM9eglJCQES0tLXr16RUhICCNHjiQ6OjrpSGnx4sVJy3769Ilx48ZRpUoV6tWrh4eHR5Zu848uxxacL+Nc/PN6jbm5OeHh4SnGwPhy7cbCwiLrAmaW9sNg3gkArAaWYnHT/Bgb6jBuy2k8T95KGtRtSK0h+HX241HUI2zdbTkYflDJwYVIv3nz5tGpUyf27NlDkyZNvru8tbU1o0aNQltbm4CAAAICAujRo0fS/LVr12JhYcGOHTvo1asXs2fP5uzZs5m5CeIfcsQptdT4+voCYGVllTStadOmbN++ncOHDyf7cPr4+GBmZoa5uXmW58wUFerApofQqSiFXKsyb/Ba5uWvwKrD17n1NIohrSuhpaFG01JNOd3rNA6eDrTY2ILZTWczqMaLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(dat.Strike, dat.Voltron, c='steelblue', label=\"Voltron\")\n",
|
||||
"plt.plot(dat.Strike, dat.Return, c='green', label=\"Truth\")\n",
|
||||
"plt.fill_between(dat.Strike, dat.Bid, dat.Ask, color='OrangeRed',\n",
|
||||
" label=\"Market Bid/Ask\")\n",
|
||||
"plt.ylabel(\"price (return)\")\n",
|
||||
"plt.xlabel(\"Strike\")\n",
|
||||
"plt.legend(fontsize=14)\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 19,
|
||||
"id": "3d718747",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYcAAAEnCAYAAABCAo+QAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAABhNElEQVR4nO3dd3gUZdfA4d+m915IJSEhgdCl916kCCiIDUR97Sh2UWwvr58FCyKICIgoYiMISJGuID2gdIEQSEghjfRedr8/JlkIaZuy2ZRzX1euZGdmZ0+GZU/mKedRaTQaDUIIIcRNjAwdgBBCiMZHkoMQQohyJDkIIYQoR5KDEEKIciQ5CCGEKEeSgxBCiHIqTQ5xcXHk5eU1ZCxCCCEaiUqTw/Dhw9m5c6f28YwZMzh06FCDBCWEEMKwKk0OJiYmFBUVaR8fPXqU5OTkBglKCCGEYVWaHLy9vdmzZw+ZmZnabSqVqkGCEkIIYViqyspnrFmzhv/97381SggqlYpz587VW3BCCCEMw6SyHffffz8BAQEcOnSIxMRE1q9fT48ePfDx8WnI+IQQQhhApXcOt2rXrh0fffQREyZM0HdMQgghDKzK0Uq7d+/WPp48eTJBQUENEpQQQgjDqjQ5XLt2jezsbO3jDRs2cPHixQYJSgghhGFVmhzc3d3LJANZ9kEIIVqOSvsc3n33Xb7//nuCg4Oxt7fn6NGjBAQE4OzsXPnJVCq+/fZbvQUrhBCiYVQ6Wumll17Czs6OgwcPEhcXh0qlIiUlhdzc3IaMTwghhAHIaCUhhBDl6Jwc1q9fT8+ePfH29tZ3TEIIIQxM5+Rws9TUVGJiYgClzIajo2O9ByaEEMJwKu1zqMj58+d59913OX78eJntPXr0YO7cubRr165egxNCCGEYOt85XLx4kWnTplFQUMCQIUNo27YtAJcuXeKPP/7AwsKCn376Sbu9oeTl5XHmzBlcXV0xNjZu0NcWQoimqri4mKSkJDp27IiFhUW5/TrfOXz++eeYmpry008/ERwcXGbfxYsXeeCBB/j8889ZtGhR3aOugTNnznD//fc36GsKIURzsWbNGnr06FFuu87JISwsjPvuu69cYgAICgri3nvv5aeffqpblLXg6uoKKL9gq1atGvz1hRAN56mnngJgyZIlBo6k6YuPj+f+++/XfobeSufkkJubW+lJANzc3AwyB6K0KalVq1YykkqIZu6ll14CkP/r9aiy5vhKy2fcysfHhz/++KPS/X/88YeU8xZC6NWgQYMYNGiQocNoEXRODhMnTmT//v28+OKLhIeHU1xcTHFxMRcvXuTFF1/kwIEDTJ48WZ+xCiFauDNnznDmzBlDh9Ei6Nys9Mgjj3Du3Dm2bNnC1q1bMTJS8oparUaj0XD77bfz8MMP6y1QIYR45513AAgNDTVsIC2AzsnB2NiYzz77jAMHDrBr1y5iYmLQaDT4+voyYsQI+vXrp884hRBCNKAaTYID6N+/P/3799dHLHqTkZFBYmIihYWFhg5FAKampri5uWFnZ2foUIQQldA5OXz55ZfceeeduLu76zOeepeRkUFCQgJeXl5YWlqiUqkMHVKLptFoyM3NJTY2FkAShBCNlM4d0gsXLmTYsGE88cQT7Nq1i+LiYn3GVW8SExPx8vLCysqq4sSgLoaC/Mq/CvOhqACKCqG4SDlerQZZ/KhWVCoVVlZWeHl5kZiYaOhwhBCV0PnO4ZdffiE0NJStW7eyd+9enJ2dmTRpEnfddRf+/v76jLFOCgsLsbS0rPyAnEzITAGVrnny5qTQTO5CVKqbvozKftcTS42GwvTrsOH3+jmhkTGYWdz0Zal8NzVT9lX1b2ViqnwZm1b9s7GJXq+JqN6rr75q6BBaDJ2TQ+fOnencuTOvv/4627ZtIzQ0lBUrVvD1119z2223MXXqVMaMGVNhjQ5Dq7IpSaNRPghNzGp41mZ251B6J6RRgxqgSK8vpwLlTiymntYl12hKYi+5s1MXK49VKqpO4hplv4qKj9NolGM0JceZmoKJOVhYQZsu4N8JWrUBS+v6+T1ElXr27GnoEFqMGndIW1hYMGnSJCZNmkRUVBShoaFs2LCB1157jXfffZfx48czbdo02rdvr494G5Fm9hfkzQm0oX41lRHYVb7sbKOj0ZQ0KRYrzYwn/oC/dwIqcG8NwT3Bpx24+Sp3GaLehYWFAZIkGkKd3sFeXl506NCB06dPk5SURE5ODmvXruXnn39m4MCBvPvuu7i5udVXrKIG5rz7Aanp6Xz10fuGDqX5UKnA2BgwVu40LUruFtRqyEqDfaFKYjUyBnc/8A6GVv7g7AEObpIw6sGHH34IyDyHhlCrd2t4eDihoaH89ttvpKWl4ebmxpNPPsnUqVMxNTXlhx9+YOXKlbz++uusWLGivmNu9p545XXy8vNZtfCTcvsiIqMYe/9MVi74iP69yldSrMz0Wc/R1t+ft16cXZ+hCgAjI7C2V75AaS5LT4b4K0rzVmnTlquPcmfh1xG82oKZuUHDFqIqOieH7OxstmzZQmhoKKdPn8bIyIiBAwdy9913M2TIEO2MaYDZs2djZWXFF198oZegm7spE8Yy67W3iLkWj7dH2UqzoZu34tXKnb49btPLaxcWFWFqIn/h1omxSdlkAcrdRU4m/LMbjm8HIxOlzyKkj3KHYWVruHiFqIDOQ1kHDBjA22+/TVJSEk8//TS7d+9m6dKlDBs2rExiKOXl5UVeXl69BttSDOnbFxcnR37dUnYkT2FRERu37eTOcbdz/NRppj76JJ2GjqLf+Dt5b+EXFFQyyW/Oux9w9J+TrPl1A8H9hxLcfygx1+I58vcJgvsPZe/Bw0z5z5N0HDyS/UfCKCgo4P8+W0y/8XfSaego7n70KY6dPK09X+nzDh07ztRHn6TLsDHc+fDjnL1QT53LzZGRkZIAXLzArTU4usPVf2HjF7DkWfjxPTi1F3KzDB2pEEAN7hz69OnDtGnTGDRoUIXJ4FZjx45l7NixdQqupTIxMWbS7aNZv3U7sx5+UHu9/9h/kNT0dCaPHc24+x/ijtEj+WDuHK7GxvHGBx9hZKRizjNPlTvf3OdmERkdg39rH154/FEAnBzsib0WD8DHXy7j1VlP0trbC2srK+Yv+Ypte/7kvddfxsfTk29+WsujL77C9p++x83lRgfyJ0tX8NKTj+Hq4sx7ny3ipf/+H1vXrJKJhrowNgGHkhL4ajWkJMD2b2DPD9BnAnQZKiOghEHVaIZ0c7Hn2FV2Hr2qPCgqAnVRDeY51M7Izi4M6+Si8/FTxo9l+fc/cjDsOAN6KyMzQjdvpX+vHvyycTOuzk6889JzGBkZEeDXmhefeIy3PvqU2Y8+jOUtw4ltbWwwNTXB0twCV2encq816+EHta+Rk5vLT+t/4905LzGkX18A/vvy8xw+/g9rft3A8489on3e7Ecfok/3bgA89dAM7nvyWRKSkmnlVvm6H6ICRkZg66h8FeTD/nVwZDP0vQM6D1GGzQrgRuE9oX/6/UQUtebn403Prp1ZV9K0lJCUzP6jYUwdP5aIyKt07dihzB1c9y6dKCwsJComtsav1bH9jdX9rsbGUVhUxG2dOmq3GRsb07VjCBFXoso8LzggQPuzm4uS+K6nptb49cVNzMyVZidrB9i3Fr56EcK2QV6OoSNrFDp27EjHjh2rP1DUWY16Ho8fP86yZcs4efIkGRkZaG4pIaFSqTh37ly9BqgPw3r4MqyHr/IgMxWyUqsZZnjTJKkyk6X023wyZcI43vzwY9IyMli/dRv2tnYMG9if33bsqnSibm2adCwtbppBXvJvWtF5bt1kclPHdenxanUzmxxoKGYWSpIoyIM/f4ZDv8HIGdCud4uepb1v3z4AWfCnAdRoDemHHnoIGxsbunTpwt69e+nTpw85OTmcOnWKoKAgOnTooM9Y9cPUHMyrum0vmR2rnYF70896NmZgH95dYMZvW39n3ZatTBo9HFOKCfT15vc/9qEuyNXePRz/529MTU3xdXeGonxlopZGrfwMmBobU1xUqH0MQHGB8r2opH4U4NvKDVNTU47/8w8+7srdQHFxMSfOnGX8iKHKccUlHd9FBdrnab8XF974uTrqYkiOqf0Fupl2hvfNM5opKQNSxQ2yClAZK007RiXfVTf9rN1W+r2BP5jNLJQJdvk58NsXEHESht8PljYNG0cj8fnnnwOSHBqCzslh6dKluLq6sm7dOgD69evH448/Tt++fdm/fz/PPvssb7/9tt4C1RsLq9q16TZA4T0LYPyEO1j83Y+kp6czZcbD4O7HfY8+ybfrNvLO0m95cMYMoqOj+WTFdzxw//1Ytm5X8mQbyC1UJmMBXv4BnD57lphCE6ysrXGwtwfHBOVYVx9wdACNBiuNhnunTePjFd/i6O2Ht7cXq777nuupadw382FwcgW7aOV5jm7g6Kj8nFtSiNHeGZzKDr+tVFIGTKunWjlqdUlxxMKy3/NzlWRV2b+XRg2FBVCYp/yVXpCnFFss/bmoAPIKlARaWAholCRR1V2jRqPs1miUpGJhpfwBYmpe++RibgXu/nDhKMRcgAlPKnMlhNATnZPDqVOnmDlzJk5OTqSlpQFom5UGDBjAxIkTWbhwId99951eAm10GugvyKlTp/Ljjz/SrVs3AgIDAXBv1Yrly5czf/58Jk6ahJ2dHePHj+eFF1+8EdfNxfSAhx95hDlz5jBu/Hjy8vLYvXt32WONbiwy/vKrr4KREa+99TYZGRmEhISwfMUK3LxK1ggvrUNlaq78ZVv6c+k+Mx3ra5mYNq0POI1GmeBWmnwqU1wM2WnKRLjrsRAfqdwhpSXcuOa2zjdmWOvKyEhJ5FlpsOZdGHgX9BpXMmtbiPqlc3IoKCjQruVgZqZ8OGRnZ2v3t2/fnt9++62ewxMdOnTgwoUL5bb37NmTtWvXVvq8Dz74oMxjf39/fv755zLbvL2Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(dat.Strike, np.abs(dat.Return-dat.Voltron), label=\"Voltron\")\n",
|
||||
"plt.fill_between(dat.Strike, np.abs(dat.Return-dat.Ask),\n",
|
||||
" np.abs(dat.Return-dat.Bid),\n",
|
||||
" color=\"OrangeRed\", alpha=0.5, label='Bid/Ask')\n",
|
||||
"plt.ylabel(\"Abs. Distance to True Payoff\", fontsize=18)\n",
|
||||
"plt.xlabel(\"Strike\")\n",
|
||||
"plt.axvline(train_y[-1], label=\"ATM\", c='k', ls='--')\n",
|
||||
"plt.legend(fontsize=14)\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "918deec6",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "d3d5be72",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,97 @@
|
||||
import numpy as np
|
||||
import datetime as dt
|
||||
import matplotlib.pyplot as plt
|
||||
import seaborn as sns
|
||||
import torch
|
||||
import gpytorch
|
||||
import os
|
||||
# import robin_stocks.robinhood as r
|
||||
import pickle5 as pickle
|
||||
import pandas as pd
|
||||
sns.set_style("whitegrid")
|
||||
sns.set_palette("bright")
|
||||
sns.set(font_scale=2.0)
|
||||
|
||||
|
||||
import sys
|
||||
sys.path.append("../")
|
||||
from voltron.likelihoods import VolatilityGaussianLikelihood
|
||||
from voltron.models import SingleTaskVariationalGP as SingleTaskCopulaProcessModel
|
||||
from voltron.kernels import BMKernel, VolatilityKernel
|
||||
from voltron.models import BMGP, VoltronGP
|
||||
from gpytorch.kernels import ScaleKernel, RBFKernel, MaternKernel
|
||||
from voltron.option_utils import GetTradingDays, GetTrainingData, Pricer, FindLastTradingDays
|
||||
from voltron.train_utils import LearnGPCV, TrainVolModel, TrainDataModel
|
||||
|
||||
def main():
|
||||
years = [yr for yr in range(2006, 2018)]
|
||||
logger = []
|
||||
full_logger = []
|
||||
SPY = pd.read_csv("./data/SPY_prices.csv")
|
||||
SPY['Date'] = pd.to_datetime(SPY['Date'])
|
||||
ntrain = 375
|
||||
|
||||
nvol = 100
|
||||
npx = 100
|
||||
|
||||
for year in years:
|
||||
options = pd.read_csv("./data/SPY_" + str(year) + ".csv")
|
||||
options.expiration = pd.to_datetime(options.expiration)
|
||||
options.quotedate = pd.to_datetime(options.quotedate)
|
||||
qday = options.quotedate.unique()[0]
|
||||
quote_price = SPY[SPY['Date']==qday].Close.item()
|
||||
options = options[(options.quotedate == qday) & (options.type=='call')]
|
||||
edays = options.expiration.sort_values().unique()
|
||||
testdays = (edays - qday)/np.timedelta64(1, "D")
|
||||
edays = edays[(testdays > 100) & (testdays < 365)]
|
||||
lastdays = FindLastTradingDays(SPY, edays)
|
||||
ntests = np.array([GetTradingDays(SPY, qday, pd.Timestamp(ld)) for ld in lastdays])
|
||||
fulltest = ntests[-1]
|
||||
|
||||
train_y = torch.FloatTensor(GetTrainingData(SPY, qday, ntrain).to_numpy())
|
||||
test_y = torch.FloatTensor(GetTrainingData(SPY,
|
||||
pd.Timestamp(lastdays[-1]),
|
||||
fulltest).to_numpy())
|
||||
full_x = torch.arange(ntrain+fulltest).type(torch.FloatTensor)
|
||||
full_x = full_x/252.
|
||||
train_x = full_x[:ntrain]
|
||||
dt = train_x[1]-train_x[0]
|
||||
test_x = full_x[ntrain:]
|
||||
|
||||
## learn vol with GPCV ##
|
||||
vol = LearnGPCV(train_x, train_y, train_iters=750)/(dt**0.5)
|
||||
|
||||
## train vol GP ##
|
||||
vmod, vlh = TrainVolModel(train_x, vol, train_iters=750)
|
||||
|
||||
## train data gp ##
|
||||
dmod, dlh = TrainDataModel(train_x, train_y, vmod, vlh, vol,
|
||||
printing=False, train_iters=750)
|
||||
|
||||
## figure out how to price options sanely ##
|
||||
|
||||
px_samples = torch.zeros(npx*nvol, len(edays))
|
||||
px_paths = torch.zeros(npx*nvol, fulltest)
|
||||
vol_paths = torch.zeros(nvol, fulltest)
|
||||
dmod.vol_model.eval();
|
||||
dmod.eval();
|
||||
|
||||
for vidx in range(nvol):
|
||||
# print(vidx)
|
||||
vol_pred = dmod.vol_model(test_x).sample().exp()
|
||||
vol_paths[vidx, :] = vol_pred.detach()
|
||||
|
||||
px_pred = dmod.GeneratePrediction(test_x, vol_pred, npx).exp()
|
||||
px_paths[vidx*npx:(vidx*npx + npx), :] = px_pred.detach().T
|
||||
px_samples[vidx*npx:(vidx*npx+npx), :] = px_pred[ntests-1].detach().T
|
||||
|
||||
|
||||
option_output = Pricer(px_samples, options, edays, test_y[ntests-1],
|
||||
quote_price)
|
||||
option_output.to_pickle("./output/options" + str(year) + ".pkl")
|
||||
print(str(year), "Done")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
@@ -0,0 +1,110 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
import pandas as pd
|
||||
import copy
|
||||
from voltron.option_utils import GetTradingDays, GetTrainingData, Pricer, FindLastTradingDays
|
||||
from scipy.optimize import minimize
|
||||
|
||||
def BlackVol(pars, K, f, T):
|
||||
alpha = torch.exp(pars[0][0]) ## v0
|
||||
rho = 2 * torch.sigmoid(pars[0][1]) - 1. ##rho
|
||||
v = torch.exp(pars[0][2]) ## "sigma"
|
||||
beta = 1.
|
||||
num = 1 + (alpha**2 * (1-beta)**2/(24 * (f*K)**(1-beta)) +\
|
||||
0.25 * rho*beta*v*alpha/((f*K)**(0.5*(1-beta))) +\
|
||||
v**2*(2-3*rho**2)/24)*T
|
||||
num*= alpha
|
||||
|
||||
denom = (f*K)**(0.5*(1-beta)) * (1 + (1-beta)**2/24 * torch.log(f/K)**2 +\
|
||||
(1-beta)**4/1920 * torch.log(f/K)**4)
|
||||
|
||||
z = v/alpha * (f*K)**(0.5*(1-beta)) * np.log(f/K)
|
||||
xi_z = torch.log((torch.sqrt(1 - 2 * rho * z + z**2) + z - rho)/(1-rho))
|
||||
|
||||
return num/denom * z/xi_z
|
||||
|
||||
def MinVol(pars, Ks, Fs, Ts, ivol):
|
||||
return torch.mean((ivol - BlackVol(pars, Ks, Fs, Ts)).pow(2))
|
||||
|
||||
def Calibrate(Fs, Ks, Ts, ivol, iters=1000):
|
||||
pars = [torch.tensor([-1., -5., -3.], requires_grad=True)]
|
||||
opt = torch.optim.SGD(pars, lr=0.1)
|
||||
stored_pars = torch.zeros(iters, 3)
|
||||
losses = []
|
||||
for e in range(iters):
|
||||
stored_pars[e, :] = pars[0]
|
||||
loss = MinVol(pars, Ks, Fs, Ts, ivol)
|
||||
opt.zero_grad()
|
||||
loss.backward()
|
||||
losses.append(loss.item())
|
||||
opt.step()
|
||||
|
||||
return pars[0].detach().numpy()
|
||||
|
||||
def SABRSim(Np, Nt, S0, V0, sigma, rho, dt=1./252.):
|
||||
dW = np.random.randn(Nt+1, Np) * np.sqrt(dt)
|
||||
dZ = rho * dW + np.sqrt(1-rho**2) * np.random.randn(Nt+1, Np) * np.sqrt(dt)
|
||||
|
||||
S = np.zeros((Nt+1, Np))
|
||||
S[0] = S0
|
||||
V = np.zeros((Nt+1, Np))
|
||||
V[0] = V0
|
||||
|
||||
for t in range(Nt):
|
||||
S[t+1] = S[t] + V[t]*S[t]*dW[t]
|
||||
V[t+1] = V[t] + sigma*V[t]*dZ[t]
|
||||
|
||||
return S[1:]
|
||||
|
||||
def main():
|
||||
years = [yr for yr in range(2006, 2018)]
|
||||
logger = []
|
||||
full_logger = []
|
||||
SPY = pd.read_csv("./data/SPY_prices.csv")
|
||||
SPY['Date'] = pd.to_datetime(SPY['Date'])
|
||||
Np = 10000
|
||||
for year in years:
|
||||
options = pd.read_csv("./data/SPY_" + str(year) + ".csv")
|
||||
options.expiration = pd.to_datetime(options.expiration)
|
||||
options.quotedate = pd.to_datetime(options.quotedate)
|
||||
qday = options.quotedate.unique()[0]
|
||||
quote_price = SPY[SPY['Date']==qday].Close.item()
|
||||
options = options[(options.quotedate == qday) & (options.type=='call')]
|
||||
edays = options.expiration.sort_values().unique()
|
||||
testdays = (edays - qday)/np.timedelta64(1, "D")
|
||||
edays = edays[(testdays > 100) & (testdays < 365)]
|
||||
lastdays = FindLastTradingDays(SPY, edays)
|
||||
ntests = np.array([GetTradingDays(SPY, qday,
|
||||
pd.Timestamp(ld)) for ld in lastdays])
|
||||
fulltest = ntests[-1]
|
||||
|
||||
test_y = torch.FloatTensor(GetTrainingData(SPY,
|
||||
pd.Timestamp(lastdays[-1]),
|
||||
fulltest).to_numpy())
|
||||
|
||||
## extract data for calibration ##
|
||||
ivol = torch.tensor(options.impliedvol.to_numpy())
|
||||
Fs = torch.tensor(options.underlying_last.to_numpy())
|
||||
Ks = torch.tensor(options.strike.to_numpy())
|
||||
starts = options.quotedate.dt.date.to_numpy()
|
||||
ends = options.expiration.dt.date.to_numpy()
|
||||
Ts = torch.tensor(([np.busday_count(qd, ed)/252. for qd, ed in zip(starts, ends)]))
|
||||
|
||||
pars = Calibrate(Fs, Ks, Ts, ivol)
|
||||
v0 = np.exp(pars[0])
|
||||
1/(1 + np.exp(-pars[1]))
|
||||
rho = (2/(1 + np.exp(-pars[1])) - 1.)
|
||||
sigma = np.exp(pars[2])
|
||||
px_paths = SABRSim(Np, fulltest, quote_price, v0, sigma, rho)
|
||||
# px_samples = torch.tensor(px_paths[ntests-1])
|
||||
|
||||
option_output = Pricer(torch.tensor(px_paths), options, edays, test_y[ntests-1],
|
||||
quote_price)
|
||||
|
||||
option_output.to_pickle("./output/sabr" + str(year) + ".pkl")
|
||||
print(str(year), "Done")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
@@ -0,0 +1,115 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
import pandas as pd
|
||||
import gpytorch
|
||||
import argparse
|
||||
import datetime
|
||||
import warnings
|
||||
|
||||
from voltron.data import make_ticker_list, GetStockHistory
|
||||
import sys
|
||||
sys.path.append("../trading/")
|
||||
from GenerateMultiMeanPreds import GenerateStockPredictions, GenerateBasicPredictions
|
||||
from gpytorch.utils.warnings import NumericalWarning
|
||||
warnings.simplefilter("ignore", NumericalWarning)
|
||||
|
||||
def main(args):
|
||||
|
||||
|
||||
data_path = "../../voltron/data/"
|
||||
ticker_file = args.ticker_fname + ".txt"
|
||||
tckr_list = make_ticker_list(data_path + ticker_file)
|
||||
|
||||
if args.end_date.lower() == "none":
|
||||
end_date = datetime.date.today()
|
||||
else:
|
||||
end_date = datetime.datetime.strptime(args.end_date, "%Y-%m-%d")
|
||||
|
||||
|
||||
for tckr in tckr_list:
|
||||
try:
|
||||
data = GetStockHistory(tckr, history=args.ntrain + args.lookback)
|
||||
if args.kernel.lower() == 'volt':
|
||||
GenerateStockPredictions(tckr, data, forecast_horizon=args.forecast_horizon,
|
||||
train_iters=args.train_iters,
|
||||
nsample=args.nsample, mean_name=args.mean,
|
||||
ntrain=args.ntrain, save=args.save)
|
||||
else:
|
||||
GenerateBasicPredictions(tckr, data, forecast_horizon=args.forecast_horizon,
|
||||
kernel_name=args.kernel, mean_name=args.mean,
|
||||
k=args.k, train_iters=args.train_iters,
|
||||
nsample=args.nsample, ntimes=args.ntimes,
|
||||
ntrain=args.ntrain, save=args.save)
|
||||
|
||||
except:
|
||||
print("FAILED ", tckr)
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--ticker_fname",
|
||||
type=str,
|
||||
default='test_tickers',
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ntrain",
|
||||
type=int,
|
||||
default=400,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ntimes",
|
||||
type=int,
|
||||
default=25,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--forecast_horizon",
|
||||
type=int,
|
||||
default=100,
|
||||
)
|
||||
parser.add_argument(
|
||||
'--kernel',
|
||||
type=str,
|
||||
default="volt",
|
||||
)
|
||||
parser.add_argument(
|
||||
'--mean',
|
||||
type=str,
|
||||
default="ewma",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--nsample",
|
||||
type=int,
|
||||
default=1000,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--printing",
|
||||
type=bool,
|
||||
default=False
|
||||
)
|
||||
parser.add_argument(
|
||||
"--train_iters",
|
||||
type=int,
|
||||
default=300,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--end_date",
|
||||
default="none",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--lookback",
|
||||
type=int,
|
||||
default=500,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--save",
|
||||
type=bool,
|
||||
default=True,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--k",
|
||||
type=int,
|
||||
default=100,
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
main(args)
|
||||
@@ -0,0 +1,133 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
import pandas as pd
|
||||
import gpytorch
|
||||
import argparse
|
||||
import datetime
|
||||
import warnings
|
||||
import os
|
||||
from voltron.data import make_ticker_list, GetStockHistory
|
||||
from LSTMUtils import SequenceDataset, LSTM, TrainLSTM, LSTMRollouts, NLL
|
||||
from torch.utils.data import DataLoader
|
||||
|
||||
def main(args):
|
||||
|
||||
data_path = "../../voltron/data/"
|
||||
ticker_file = args.ticker_fname + ".txt"
|
||||
tckr_list = make_ticker_list(data_path + ticker_file)
|
||||
# tckr_list = ['TSLA']
|
||||
|
||||
use_cuda = False
|
||||
if torch.cuda.is_available():
|
||||
use_cuda = True
|
||||
|
||||
ntest = args.forecast_horizon
|
||||
ntrain = args.ntrain
|
||||
seq_len = args.seq_length
|
||||
|
||||
if args.end_date.lower() == "none":
|
||||
end_date = datetime.date.today()
|
||||
else:
|
||||
end_date = datetime.datetime.strptime(args.end_date, "%Y-%m-%d")
|
||||
|
||||
|
||||
for tckr in tckr_list:
|
||||
try:
|
||||
data = GetStockHistory(tckr, history= ntrain + args.lookback)
|
||||
end_idxs = torch.arange(args.ntrain, data.shape[0],
|
||||
int((data.shape[0]-args.ntrain)/args.ntimes))
|
||||
|
||||
savepath = "./saved-outputs/" + tckr + "/"
|
||||
if not os.path.exists(savepath):
|
||||
os.mkdir(savepath)
|
||||
|
||||
for last_day in end_idxs:
|
||||
date = str(data.index[last_day.item()].date())
|
||||
raw_y = data.Close[last_day.item()-ntrain:last_day.item()].to_numpy()
|
||||
raw_y = torch.FloatTensor(raw_y).log()
|
||||
train_y = (raw_y - raw_y.mean())/raw_y.std()
|
||||
|
||||
## make trainloader ##
|
||||
dset = SequenceDataset(train_y, seq_len)
|
||||
trainloader = DataLoader(dset, batch_size=args.batch_size, shuffle=True)
|
||||
|
||||
model = LSTM(2, seq_len, 128, 1)
|
||||
if use_cuda:
|
||||
model = model.cuda()
|
||||
optimizer = torch.optim.Adam(model.parameters(), lr=0.01)
|
||||
TrainLSTM(trainloader, model, NLL, optimizer, epochs=args.train_epochs,
|
||||
printing=True, use_cuda=use_cuda)
|
||||
|
||||
rollouts = LSTMRollouts(model, args.nsample, ntest,
|
||||
dset, use_cuda).cpu()
|
||||
rollouts = rollouts * raw_y.std() + raw_y.mean()
|
||||
torch.save(rollouts, savepath + "lstm_" + date + ".pt")
|
||||
|
||||
del model
|
||||
except:
|
||||
print("FAILED ", tckr)
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--ntimes",
|
||||
type=int,
|
||||
default=25,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--forecast_horizon",
|
||||
type=int,
|
||||
default=20,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--seq_length",
|
||||
type=int,
|
||||
default=25,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ticker_fname",
|
||||
type=str,
|
||||
default='test_tickers',
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ntrain",
|
||||
type=int,
|
||||
default=400,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch_size",
|
||||
type=int,
|
||||
default=128,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--nsample",
|
||||
type=int,
|
||||
default=1000,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--printing",
|
||||
type=bool,
|
||||
default=False
|
||||
)
|
||||
parser.add_argument(
|
||||
"--train_epochs",
|
||||
type=int,
|
||||
default=200,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--end_date",
|
||||
default="none",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--lookback",
|
||||
type=int,
|
||||
default=500,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--save",
|
||||
type=bool,
|
||||
default=False,
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
main(args)
|
||||
@@ -0,0 +1,107 @@
|
||||
import numpy as np
|
||||
import matplotlib.pyplot as plt
|
||||
import pandas as pd
|
||||
import torch
|
||||
from torch import nn
|
||||
import seaborn as sns
|
||||
import time
|
||||
import copy
|
||||
import sys
|
||||
from torch.utils.data import DataLoader
|
||||
from torch.utils.data import Dataset
|
||||
from torch.autograd import Variable
|
||||
|
||||
def NLL(targets, outputs):
|
||||
dist = torch.distributions.Normal(outputs[:, 0], outputs[:, 1])
|
||||
return -dist.log_prob(targets).sum()
|
||||
|
||||
class SequenceDataset(Dataset):
|
||||
def __init__(self, data, sequence_length=5):
|
||||
self.sequence_length = sequence_length
|
||||
self.X = data.float()
|
||||
|
||||
def __len__(self):
|
||||
return self.X.shape[0]-1
|
||||
|
||||
def __getitem__(self, i):
|
||||
if i >= self.sequence_length - 1:
|
||||
i_start = i - self.sequence_length + 1
|
||||
x = self.X[i_start:(i + 1)]
|
||||
else:
|
||||
padding = self.X[0].repeat(self.sequence_length - i - 1, 1).squeeze(-1)
|
||||
x = self.X[0:(i + 1)]
|
||||
x = torch.cat((padding, x), 0)
|
||||
|
||||
return x.unsqueeze(0), self.X[i+1]
|
||||
|
||||
class LSTM(nn.Module):
|
||||
def __init__(self, num_classes, seq_len, hidden_size, num_layers):
|
||||
super(LSTM, self).__init__()
|
||||
self.num_classes = num_classes #number of classes
|
||||
self.num_layers = num_layers #number of layers
|
||||
self.input_size = seq_len #input size
|
||||
self.hidden_size = hidden_size #hidden state
|
||||
|
||||
self.lstm = nn.LSTM(input_size=seq_len, hidden_size=hidden_size,
|
||||
num_layers=num_layers, batch_first=True) #lstm
|
||||
self.fc_1 = nn.Linear(hidden_size, 128) #fully connected 1
|
||||
self.fc = nn.Linear(128, num_classes) #fully connected last layer
|
||||
|
||||
self.relu = nn.ReLU()
|
||||
self.softplus = nn.Softplus()
|
||||
|
||||
def forward(self,x):
|
||||
h_0 = Variable(torch.zeros(self.num_layers, x.size(0), self.hidden_size)).to(x.device) #hidden state
|
||||
c_0 = Variable(torch.zeros(self.num_layers, x.size(0), self.hidden_size)).to(x.device) #internal state
|
||||
# Propagate input through LSTM
|
||||
output, (hn, cn) = self.lstm(x, (h_0, c_0)) #lstm with input, hidden, and internal state
|
||||
|
||||
hn = hn[self.num_layers-1]
|
||||
hn = hn.view(-1, self.hidden_size) #reshaping the data for Dense layer next
|
||||
out = self.relu(hn)
|
||||
out = self.fc_1(out) #first Dense
|
||||
out = self.relu(out) #relu
|
||||
out = self.fc(out) #Final Output
|
||||
|
||||
output = torch.zeros_like(out)
|
||||
output[:, 0] = out[:, 0]
|
||||
output[:, 1] = self.softplus(out[:, 1])
|
||||
return output
|
||||
|
||||
def TrainLSTM(data_loader, model, loss_function, optimizer, epochs=200,
|
||||
printing=False, use_cuda=False):
|
||||
num_batches = len(data_loader)
|
||||
total_loss = 0
|
||||
model.train()
|
||||
for epoch in range(epochs):
|
||||
for X, y in data_loader:
|
||||
if use_cuda:
|
||||
X = X.cuda()
|
||||
y = y.cuda()
|
||||
output = model(X)
|
||||
loss = loss_function(y, output)
|
||||
|
||||
optimizer.zero_grad()
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
|
||||
total_loss += loss.item()
|
||||
if printing:
|
||||
if epoch%10 == 0:
|
||||
avg_loss = total_loss / num_batches
|
||||
print(f"Train loss: {avg_loss}, Epoch: {epoch}")
|
||||
|
||||
def LSTMRollouts(model, nrollout, rollout_len, dset, use_cuda=False):
|
||||
xin, xout = dset[len(dset)-1]
|
||||
xx = torch.cat((xin[0, 1:], xout.unsqueeze(0)))
|
||||
xx = xx.repeat(nrollout, 1).unsqueeze(1)
|
||||
if use_cuda:
|
||||
xx = xx.cuda()
|
||||
roll_pxs = torch.zeros(nrollout, rollout_len)
|
||||
with torch.no_grad():
|
||||
for idx in range(rollout_len):
|
||||
out = model(xx)
|
||||
smpl = torch.normal(out[:, 0], out[:, 1])
|
||||
roll_pxs[:, idx] = smpl
|
||||
xx = torch.cat((xx[..., 1:], smpl.unsqueeze(-1).unsqueeze(-1)), -1)
|
||||
return roll_pxs
|
||||
@@ -0,0 +1,142 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "0f0afec5-bfbe-437e-8c42-5702d8c6be91",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Warning no robinhood utils.\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import torch\n",
|
||||
"import pandas as pd\n",
|
||||
"import gpytorch\n",
|
||||
"import argparse\n",
|
||||
"import datetime\n",
|
||||
"import warnings\n",
|
||||
"import os\n",
|
||||
"from voltron.data import make_ticker_list, GetStockHistory\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 29,
|
||||
"id": "4fb8850d-b17f-41a8-8285-d4e74e6d7b16",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ntrain = 400\n",
|
||||
"lookback = 1000\n",
|
||||
"tckr = \"MSFT\"\n",
|
||||
"data = GetStockHistory(tckr, history= ntrain + lookback)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 30,
|
||||
"id": "ef536909-31a6-46ef-99b2-c9c956e3575a",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"pxs = data.Close[:-1].to_numpy()\n",
|
||||
"date = \"2021-12-29\"\n",
|
||||
"preds = torch.load(\"./saved-outputs/\" + tckr + \"/sm_constant100_\" + date + \".pt\")\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 31,
|
||||
"id": "28d2291f-677e-441d-a387-9ddd644828ec",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"cutoff = np.where(data.index == np.datetime64(datetime.datetime.strptime(date, \"%Y-%m-%d\").date()))[0][0]\n",
|
||||
"pxs = data.Close[cutoff-100:cutoff]\n",
|
||||
"trx = np.arange(pxs.shape[0])\n",
|
||||
"tex = np.arange(pxs.shape[0], pxs.shape[0] + preds.shape[-1])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 32,
|
||||
"id": "c779bb74-3b07-4c78-9924-911efae14177",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"[<matplotlib.lines.Line2D at 0x7f9d88a95f40>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7f9d88ab70a0>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7f9d88ab71c0>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7f9d88ab72e0>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7f9d88ab7400>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7f9d88ab7520>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7f9d88ab7640>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7f9d88ab7760>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7f9d88ab7880>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7f9d88ab79a0>]"
|
||||
]
|
||||
},
|
||||
"execution_count": 32,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAXcAAAD4CAYAAAAXUaZHAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMSwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Z1A+gAAAACXBIWXMAAAsTAAALEwEAmpwYAAA+6klEQVR4nO3deXxb13Xo+9/GDBATQYAzRUqUKMm2IlmSHceOnXly26Rzknebpk17/dqm8/Ta5r7eJK/t+6Tp8G46J01e0960mZo0jp3RzWTfVrZlW5I1WKJEieIIEiTmkQD2/QPAEUeRlDiC6/v58CPw4ADYOKQWNtfee22ltUYIIUR9MW12A4QQQqw9Ce5CCFGHJLgLIUQdkuAuhBB1SIK7EELUIctmNwAgGAzqnp6ezW6GEEJsK88991xEax1a7L4tEdx7eno4efLkZjdDCCG2FaXU4FL3SVpGCCHqkAR3IYSoQxLchRCiDklwF0KIOiTBXQgh6pAEdyGEqEMS3IUQog5JcBdCiDokwV0IIdZJIpEgFosZ32utmZmZ2ZDX3hIrVIUQoh4VCgVqGyIlEgny+Tz5fJ7Ozs51f23puQshxG1Yajc7rTXFYhGAmZkZcrkcwWAQk8lEKpVa93ZJcBdCiFuUSCQ4derUkvdprZmenub06dOUSiUmJycJBoOMjIwQiUQ4f/78urVNgrsQQtyCqakplFJYrVby+Txaa65fv04qlaJcLpPJZEgkEvj9fg4fPozb7UZrjcVioampCZ/Ph91uX7f2SXAXQohbUC6XKRQK9Pb2MjY2Rjwep62tjVQqxdWrV3G73TgcDnw+H9lslmw2S3NzM+Pj42itGRkZoa2tbd3aJ8FdCCGWobUmHo8b36fTaRoaGtBa43Q6KRQK5PN5ACwWC93d3czMzNDQ0IDD4WB6eppgMIhSiubmZrTWFAoFXC7XurVZZssIIQQQjUbx+XyUy2XGxsYIBAK4XC7Gx8dJJpMAeDweY0DU5XKRy+UA8Hq9pNNpAGw2G6VSCYfDYTx3T08P+XyeYrFofCj4fL51fT8S3IUQO146ncZsNjM5OUm5XKazs5OpqSlisRherxe3202xWGRsbIxyuUw8HicUClEsFpmYmCAYDFIsFonH47hcLmKxGGazGbvdbjx3IpHAZDIZwb32wbBeJLgLIepKuVzGZFo+41zLe1utVkqlEu3t7Xi9XuP+YDDIzMwM0WiUUChENpslEonQ0NBAMBhkYGCAgwcP4vV6iUQiWK1WWlpaAGhsbGR4eBiHw0EikcDj8TAzM4Pb7SaRSOB2u7Hb7czMzGC1WtflOkjOXQhRV8bHxxedR57NZo3VoblcjsHBQVpaWigWi5jNZuO8aDTK5OQkkUiESCRCIBAwnrOrqwun04nT6WTv3r0Ui0USiQTNzc00NjYaz2Gz2QgGg/T39zM4OEipVMJisTA4OIjJZMLtduP1ejl37ty6XQfpuQsh6srMzAzpdBq3220cy+VyZDIZisUiLS0tRKNRLBYLp0+f5siRI2QyGbLZLMlkksbGRqxWq5E6iUQitLa2AjA8PIzb7SaXy9HS0sLly5cJBAJzBljHx8cpFotEIhFsNpsxuBoMBgG4dOkSVquVTCZDc3Pzul0HCe5CiLpRLpfJ5XLkcjmam5tRShkzXQKBgNFjN5lMtLe343Q6jbTLhQsX8Pv9c2bFaK1pbW2lXC4TiUTo6Oggk8kwPDyM3+/H4XBQLBYpFovkcjmmpqbwer0UCgXMZjMul4tisUhXV5fxnOl0ms7OTpRSZDKZdbsWy6ZllFIOpdQzSqnTSqlzSqkPVI8rpdQfKqUuKaUuKKV+edbxjyilLiulziiljq5b64UQO0KhUDBmrBSLRS5fvrzosv/BwUEaGxtpbGxkcnISgEgkQigUIhaLUSgUAGhpaSEWi9HY2IjZbGZ0dJS9e/fS1tZGMBg0vkKhEADT09MAmEwmRkdHaW1tZXBwkIaGBiYmJmhubsblcuH1ehkdHSWXy6GUQimF2Wwmk8kQjUZJp9M0NjaSSqUYGxubM6Nmra2k554HXqu1TimlrMBTSqmvAgeBLuCA1rqslKr9ffEWYF/16+XA31T/FUKIWxKNRjGbzXg8HuLxOLt37yYcDhvpErhRcbFcLmOz2ZiammJ6etoI0EopmpqaaGpqMs43mUwEAgG01iilFn1trTVaa/x+P5OTk6TTaXK5HCaTiYmJCUKhEJFIBK01drvdWKHq9Xppa2sjHo8zPT2N1ppgMEi5XGZqaorGxsYlX3MtLNtz1xW10Qlr9UsDPw98UGtdrp43UT3nbcA/Vh93AvArpdZvGZYQoq4lk8k5+XOtNWazGa/Xy9jYGNlslsnJSQYHB3E4HFgsFsrlMvv27aOvr49CocDFixdxOp2USiUuXrxIqVQiGo0aPXKllLEICWB0dNR4rXA4jM/n4/Lly1y+fJmZmRlCoRDRaJREIsH09DRXrlwhkUgwPDzM4cOH6enpYWZmBo/Hg8PhIJfLobXG4XBQLpfp7u7G5XLR0NCwbtdtRTl3pZQZeA7YC/yV1vpppVQv8Hal1A8Bk8Ava637gQ5gaNbDh6vHxuY95yPAIwC7du263fchhKhDmUyGQqGAw+EgFosRCASAStCtDYoWi0VCoZBR58XhcBgpnNq8cpPJZAR0s9nM6dOn6ezs5PLly3R1dWGxWLh06RJ9fX34fD6ee+45tNbYbDZMJhP9/f3s3buXkZERGhoaKJVK9PX10djYSDQapVgsUigUaG1txWq1Eo1GjbhW+2Cy2+3E43FjuqTNZiOXy61bamZFUyG11iWt9RGgE7hXKXUXYAdyWuvjwMeAT6zmhbXWH9VaH9daH6/92SSEEFDprUciEfL5PG63m2w2S1dXF5cuXaJQKDA5OUlbWxuNjY1ks1mi0SilUomZmRnsdjsNDQ1GeqQ2s6VUKpFKpdi9ezfFYtEow/vCCy8wODhoDHKOj4/jdDoZGBgw0jXpdJpz586RzWax2+1MT0/T3NzM1NQUPp+PlpYWPB6PsfApFArh8XgAjPIDHo8Hv99PNBoFMNqzXlY1W0ZrHVNKfRt4M5Ue+Reqd30R+P+rt0eo5OJrOqvHhBA7UK1aYnd3t3EskUhgsVgWra1SKpUoFAoEg0FKpRKRSMTo7ZrNZiPY1iosFgoFstksJpOJbDaL2+3G6XQyNjbGqVOncDqdxmrSQqFANBrF6/USi8XYtWsXo6OjTE5OYrPZjA+E3bt3Mzo6avTKm5qacLvduFwuIpEIBw8eRClFS0sLkUhkzpTG2kKn2fbu3QtUeu/ZbJZEIkEul2M9O7YrmS0TUkr5q7edwBuAl4B/A15TPe1VwKXq7UeBn6zOmrkPiGut56RkhBA7RzKZNOaNA8b88Ww2a8w/j0QixnZ0U1NTBAIByuWyMRNFa00mkzFmt3R3d5NKpXC73XR2dho9ZYvFwokTJ7h06RLBYBCLxYLVajVmtvj9foaGhti1axdTU1OYTCaOHj3K0aNHaWlpIZ/Pk06nsVgs+P1+o3RAS0uL0Tv3eDzGoieTyWRMt6y9t8XMXjHr9/sxm804nc51HVBdSc+9DfhkNe9uAj6rtX5MKfUU8Cml1K8BKeBnq+d/BXgYuAxkgJ9e+2YLIbaLQqFAc3OzMfVwamqKpqYmzGaz0TP2er2YTCbC4TBWq5VLly7hcrmMVMno6CixWIyOjg601kxOTtLS0kIikSAUCuFyuZiamuLatWu0trYave5cLoff76e/vx+ozIPv6+vj+vXr7N27l/b2diYmJmhoaGB8fBybzcb09DR+v5+GhgYCgQAmk8mo4mg2mxeUNmhsbDTeWywWw+/3L3tN1nMgtWbZ4K61PgPcvcjxGPB9ixzXwHvXonFCiO2pXC4TDoeNeuUWi4VSqUQikcBqtRo9X4vFglKKeDxuzBWvTXGcmZkxSgA4HA4OHDjAhQsXCAQCpFIpkskkNpsNgHg8zuTkJFarlVQqZcwlv/POOykUCjzwwAOcP3+e3t5e0uk0TU1NKKXI5XLY7Xbcbjfd3d1cuHABn8+Hz+cjGo0aOftkMmkUAWtvb5/zXmuDulpr8vn8nDIEm0lqywgh1tz09DROp5N0Om2kHmw2G2az2Sh1WygUGBsbw2w2Y7VauXr1KvF4HJ/PRyaTwWKxMDMzY9RFN5lMRpGt2hzydDrN9evXUUoZA5kHDhxgZmYGl8tlBP1UKmVsc1crQ+DxeEgkEkZ7rFYrhw4d4u67K31Zn8/H8PAwpVKJlpYW/H4/HR0di6ZS7Hb7gtz7ZpPgLoRYc7VFP+FwGI/Hw9mzZ43BSqjUf+nv72ffvn1GT97n8xEKhbBYLAQCAeLxOIlEgpmZGYaHh0mn05TLZa5du0apVCKTyVAqlYzphFNTU+zfvx+v14vP5+Ouu+4CMMoE7N27l87OThoaGmhra6NcLuN0Oue0WymFxVJJaJhMJnp6eoxZLzfj9XoJhUIrqka5UbZOS4QQW1ptUHM5tXPOnz/Pnj17iEajxsYX2WwWrTVXrlxh9+7dKKWw2Wy43W4CgQDFYpFYLIbdbieXy+H1enE4HBQKBfr7+43phIBRrndwcJDHH3+cAwcOGL3qQCBgzGVvb283tr1zuVzGDJVkMrmiwL1dSeEwIcRNZTIZI2AXi8VFpy+Wy2Wi0Shaa2Mzi9q0v3A4THNzMxMTE9jtds6cOcPevXtxuVwMDQ1hs9mwWCzGwqRa7fPaIKbdbqdYLOJwOOjo6GB0dJSmpiY6Ojrw+Xzs37/fCPw1tTK9tQVCbrebaDSK3++nUCis+xZ3W4FaaurORjp+/Lg+efLkZjdDCDFPbWZKLZecyWTQWs+Z7VEsFpmcnCQYDBKLxYyFP16vlwsXLmC32zl48CDJZNIoqVsbRM3lcgwPD9PT00NHRwff+973sFgsNDU1YbFYiEQi3HPPPcYepOPj45hMJjKZDIFAgGw2awyLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {
|
||||
"needs_background": "light"
|
||||
},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(trx, pxs)\n",
|
||||
"plt.plot(tex, preds[:10, :].T.exp(), color='gray', alpha=0.75, lw=0.2)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "eb8f6699-b9ff-45cc-8901-e2becfcc01b0",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.8"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,588 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "8d4521d7-e61a-4820-b704-4908e9210ee7",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"from matplotlib.lines import Line2D\n",
|
||||
"import pandas as pd\n",
|
||||
"import torch\n",
|
||||
"import gpytorch\n",
|
||||
"from gpytorch.means import Mean\n",
|
||||
"import seaborn as sns\n",
|
||||
"import time\n",
|
||||
"import copy\n",
|
||||
"import sys\n",
|
||||
"import os\n",
|
||||
"\n",
|
||||
"sns.set_style('white')\n",
|
||||
"# style.use('whitegrid')\n",
|
||||
"palette = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"sns.set(palette = palette, font_scale=2.0, style=\"white\", rc={\"lines.linewidth\": 2.0})"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "3744a8e8-cb18-43ed-81c1-54aed694e292",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def LoadDat(dat, sym, start_idx, ntrain, ntest):\n",
|
||||
" px = torch.FloatTensor(dat[dat.symbol == sym].close_price.to_numpy())\n",
|
||||
" train_y = px[start_idx:ntrain+start_idx].squeeze()\n",
|
||||
" test_y = px[start_idx + ntrain:start_idx + ntrain+ntest].squeeze()\n",
|
||||
" return train_y, test_y\n",
|
||||
"\n",
|
||||
"def LoadSims(SPDR, sym, kernel, mean, k=100):\n",
|
||||
" fpath = \"./saved-outputs/\" + SPDR + \"/\"\n",
|
||||
" fname = sym + \"_\" + kernel + \"_\" + mean + str(k) + \".pt\"\n",
|
||||
" \n",
|
||||
" return torch.load(fpath + fname)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"id": "19d2a38f-0bb2-427f-bd7a-f9423b60e448",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"['XOM' 'CVX' 'EOG' 'COP' 'SLB']\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"hypers = torch.load(\"./saved-outputs/metadata.pt\")\n",
|
||||
"ntrain = hypers['ntrain']\n",
|
||||
"ntest = hypers['ntest']\n",
|
||||
"start_idxs = hypers['start_idxs']\n",
|
||||
"\n",
|
||||
"SPDR = \"XLE\"\n",
|
||||
"dat = pd.read_pickle(dpath + SPDR + \".pkl\")\n",
|
||||
"syms = dat.symbol.unique()\n",
|
||||
"print(syms)\n",
|
||||
"\n",
|
||||
"train_x = torch.arange(ntrain) * 1./252\n",
|
||||
"test_x = torch.arange(ntest) * 1./252 + train_x[-1] + train_x[1]\n",
|
||||
"percentiles = np.linspace(0.05, 0.95, 19)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "4257adde-7106-429c-b418-973dc0345424",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Examples"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"id": "6ccf46b5-183d-400f-ace6-200a3d9d7844",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"['AAPL' 'MSFT' 'NVDA' 'V' 'PYPL']\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"hypers = torch.load(\"./saved-outputs/metadata.pt\")\n",
|
||||
"ntrain = hypers['ntrain']\n",
|
||||
"ntest = hypers['ntest']\n",
|
||||
"start_idxs = hypers['start_idxs']\n",
|
||||
"\n",
|
||||
"SPDR = \"XLK\"\n",
|
||||
"dat = pd.read_pickle(dpath + SPDR + \".pkl\")\n",
|
||||
"syms = dat.symbol.unique()\n",
|
||||
"print(syms)\n",
|
||||
"train_x = torch.arange(ntrain) * 1./252\n",
|
||||
"test_x = torch.arange(ntest) * 1./252 + train_x[-1] + train_x[1]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 5,
|
||||
"id": "9704d083-ce47-4c5e-a5e6-029fa1effbd7",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def ECDF(sample_pxs, true_px): \n",
|
||||
" return (torch.sum(sample_pxs < true_px, 0)/sample_pxs.shape[0])\n",
|
||||
" \n",
|
||||
"def Calibration(pcts, percentile=0.95):\n",
|
||||
" in_band = np.where((pcts < percentile))[0].shape[0]\n",
|
||||
" return in_band/pcts.shape[0]\n",
|
||||
"\n",
|
||||
"def GetCalibration(kernel, mean, k=100, horizon=np.arange(75,100), logger=[], exp=True):\n",
|
||||
" pcts = torch.zeros(len(syms), len(start_idxs), horizon.shape[0])\n",
|
||||
" for sym_idx, sym in enumerate(syms):\n",
|
||||
" \n",
|
||||
" fpath = \"./saved-outputs/\" + SPDR + \"/\"\n",
|
||||
" fname = sym + \"_\" + kernel + \"_\" + mean + str(k) + \".pt\"\n",
|
||||
" if os.path.exists(fpath + fname):\n",
|
||||
" for idx, start_idx in enumerate(start_idxs):\n",
|
||||
" train_y, test_y = LoadDat(dat, sym, start_idx, ntrain, ntest)\n",
|
||||
" preds = LoadSims(SPDR, sym, kernel, mean, k=k)[idx, :, horizon]\n",
|
||||
" if exp:\n",
|
||||
" preds = preds.exp()\n",
|
||||
" \n",
|
||||
" \n",
|
||||
" pcts = pcts.flatten()\n",
|
||||
" percentiles = np.linspace(0.05, 0.95, 19)\n",
|
||||
" log_name = kernel\n",
|
||||
" for pct in percentiles:\n",
|
||||
" clb = Calibration(pcts, pct)\n",
|
||||
" logger.append([clb, np.round(pct, 2), log_name, mean, k])\n",
|
||||
" \n",
|
||||
" return logger"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 20,
|
||||
"id": "4ff19402-e7d7-49c0-890c-5b136265f3ff",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"['AMT' 'PLD' 'CCI' 'EQIX' 'PSA']\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"SPDR = \"XLRE\"\n",
|
||||
"dat = pd.read_pickle(dpath + SPDR + \".pkl\")\n",
|
||||
"syms = dat.symbol.unique()\n",
|
||||
"print(syms)\n",
|
||||
"horizon = np.arange(75, 100)\n",
|
||||
"train_x = torch.arange(ntrain) * 1./252\n",
|
||||
"test_x = torch.arange(ntest) * 1./252 + train_x[-1] + train_x[1]\n",
|
||||
"\n",
|
||||
"logger = []\n",
|
||||
"logger = GetCalibration('matern', 'ewma', 100, logger=logger, horizon=horizon)\n",
|
||||
"logger = GetCalibration('matern', 'ewma', 200, logger=logger, horizon=horizon)\n",
|
||||
"logger = GetCalibration('matern', 'ewma', 400, logger=logger, horizon=horizon)\n",
|
||||
"logger = GetCalibration('matern', 'dewma', 100, logger=logger, horizon=horizon)\n",
|
||||
"logger = GetCalibration('matern', 'dewma', 200, logger=logger, horizon=horizon)\n",
|
||||
"logger = GetCalibration('matern', 'dewma', 400, logger=logger, horizon=horizon)\n",
|
||||
"logger = GetCalibration('matern', 'tewma', 100, logger=logger, horizon=horizon)\n",
|
||||
"logger = GetCalibration('matern', 'tewma', 200, logger=logger, horizon=horizon)\n",
|
||||
"logger = GetCalibration('matern', 'tewma', 400, logger=logger, horizon=horizon)\n",
|
||||
"logger = GetCalibration('matern', 'constant', 100, logger=logger, horizon=horizon, exp=False)\n",
|
||||
"xlre_df = pd.DataFrame(logger)\n",
|
||||
"xlre_df.columns = [\"Calibration\", \"Percentile\", \"Type\", 'Mean', \"k\"]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 22,
|
||||
"id": "49f482b5-af3b-4b77-b04d-82716d5b7afb",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"pd.to_pickle(xlre_df, \"./new_matern_calib.pkl\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "7b846b87-3132-42e1-9bff-f9d117176543",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Get NLL"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 94,
|
||||
"id": "63958d47-67f8-4298-94ab-3ff353b4f9ff",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def GetNLL(kernel, mean, k=100, horizon=np.arange(75,100), logger=[], exp=True):\n",
|
||||
" N = 0\n",
|
||||
" nll = 0.\n",
|
||||
" for spdr_idx, spdr in enumerate(SPDRS):\n",
|
||||
" dat = pd.read_pickle(dpath + spdr + \".pkl\")\n",
|
||||
" syms = dat.symbol.unique()\n",
|
||||
" for sym_idx, sym in enumerate(syms):\n",
|
||||
" fpath = \"./saved-outputs/\" + spdr + \"/\"\n",
|
||||
" fname = sym + \"_\" + kernel + \"_\" + mean + str(k) + \".pt\"\n",
|
||||
" if os.path.exists(fpath + fname):\n",
|
||||
" for idx, start_idx in enumerate(start_idxs):\n",
|
||||
" train_y, test_y = LoadDat(dat, sym, start_idx, ntrain, ntest)\n",
|
||||
" test_y = test_y[horizon]\n",
|
||||
" preds = LoadSims(spdr, sym, kernel, mean, k=k)[idx, :, horizon]\n",
|
||||
" if exp:\n",
|
||||
" preds = preds.exp() \n",
|
||||
" try:\n",
|
||||
" nll -= torch.distributions.Normal(preds.mean(0), preds.std(0)).log_prob(test_y).sum()\n",
|
||||
" N += test_y.numel()\n",
|
||||
" except:\n",
|
||||
" pass\n",
|
||||
" # print(\"Failed:\", spdr, sym, idx)\n",
|
||||
"\n",
|
||||
" if N >= 0:\n",
|
||||
" logger.append([nll.item(), N, kernel, mean, k])\n",
|
||||
" return logger"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 105,
|
||||
"id": "8bc987c8-982a-4f1c-bd40-5d444e58e7e0",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"SPDRS = [\"XLRE\", \"XLY\", \"XLF\", \"XLE\", \"XLK\"]\n",
|
||||
"horizon = np.arange(75, 100)\n",
|
||||
"train_x = torch.arange(ntrain) * 1./252\n",
|
||||
"test_x = torch.arange(ntest) * 1./252 + train_x[-1] + train_x[1]\n",
|
||||
"\n",
|
||||
"logger = []\n",
|
||||
"logger = GetNLL('matern', 'ewma', 100, logger=logger, horizon=horizon)\n",
|
||||
"logger = GetNLL('matern', 'ewma', 200, logger=logger, horizon=horizon)\n",
|
||||
"logger = GetNLL('matern', 'ewma', 400, logger=logger, horizon=horizon)\n",
|
||||
"logger = GetNLL('matern', 'dewma', 100, logger=logger, horizon=horizon)\n",
|
||||
"logger = GetNLL('matern', 'dewma', 200, logger=logger, horizon=horizon)\n",
|
||||
"logger = GetNLL('matern', 'dewma', 400, logger=logger, horizon=horizon)\n",
|
||||
"logger = GetNLL('matern', 'tewma', 100, logger=logger, horizon=horizon)\n",
|
||||
"logger = GetNLL('matern', 'tewma', 200, logger=logger, horizon=horizon)\n",
|
||||
"logger = GetNLL('matern', 'tewma', 400, logger=logger, horizon=horizon)\n",
|
||||
"logger = GetNLL('matern', 'constant', 100, logger=logger, horizon=horizon, exp=False)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 106,
|
||||
"id": "dee5af9f-d8d1-4576-bc42-eba43d846e8e",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df = pd.DataFrame(logger)\n",
|
||||
"df.columns = [\"NLL\", \"N\", \"Kernel\", \"Mean\", \"K\"]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 107,
|
||||
"id": "84d27556-5456-433b-945a-8107613c7d0f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df['Mean_NLL'] = df[\"NLL\"]/df['N']"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 108,
|
||||
"id": "0d562366-4268-4bd2-a4ff-2e62797cbdb7",
|
||||
"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>NLL</th>\n",
|
||||
" <th>N</th>\n",
|
||||
" <th>Kernel</th>\n",
|
||||
" <th>Mean</th>\n",
|
||||
" <th>K</th>\n",
|
||||
" <th>Mean_NLL</th>\n",
|
||||
" </tr>\n",
|
||||
" </thead>\n",
|
||||
" <tbody>\n",
|
||||
" <tr>\n",
|
||||
" <th>0</th>\n",
|
||||
" <td>6.726755e+04</td>\n",
|
||||
" <td>4350</td>\n",
|
||||
" <td>matern</td>\n",
|
||||
" <td>ewma</td>\n",
|
||||
" <td>100</td>\n",
|
||||
" <td>15.463804</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>1</th>\n",
|
||||
" <td>5.704342e+04</td>\n",
|
||||
" <td>4200</td>\n",
|
||||
" <td>matern</td>\n",
|
||||
" <td>ewma</td>\n",
|
||||
" <td>200</td>\n",
|
||||
" <td>13.581766</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>2</th>\n",
|
||||
" <td>6.180194e+04</td>\n",
|
||||
" <td>6300</td>\n",
|
||||
" <td>matern</td>\n",
|
||||
" <td>ewma</td>\n",
|
||||
" <td>400</td>\n",
|
||||
" <td>9.809832</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>3</th>\n",
|
||||
" <td>5.138177e+05</td>\n",
|
||||
" <td>4800</td>\n",
|
||||
" <td>matern</td>\n",
|
||||
" <td>dewma</td>\n",
|
||||
" <td>100</td>\n",
|
||||
" <td>107.045358</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>4</th>\n",
|
||||
" <td>1.109547e+05</td>\n",
|
||||
" <td>4200</td>\n",
|
||||
" <td>matern</td>\n",
|
||||
" <td>dewma</td>\n",
|
||||
" <td>200</td>\n",
|
||||
" <td>26.417781</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>5</th>\n",
|
||||
" <td>7.489109e+04</td>\n",
|
||||
" <td>6450</td>\n",
|
||||
" <td>matern</td>\n",
|
||||
" <td>dewma</td>\n",
|
||||
" <td>400</td>\n",
|
||||
" <td>11.611021</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>6</th>\n",
|
||||
" <td>3.585653e+06</td>\n",
|
||||
" <td>5250</td>\n",
|
||||
" <td>matern</td>\n",
|
||||
" <td>tewma</td>\n",
|
||||
" <td>100</td>\n",
|
||||
" <td>682.981476</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>7</th>\n",
|
||||
" <td>1.216625e+05</td>\n",
|
||||
" <td>3900</td>\n",
|
||||
" <td>matern</td>\n",
|
||||
" <td>tewma</td>\n",
|
||||
" <td>200</td>\n",
|
||||
" <td>31.195507</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>8</th>\n",
|
||||
" <td>7.496420e+04</td>\n",
|
||||
" <td>6450</td>\n",
|
||||
" <td>matern</td>\n",
|
||||
" <td>tewma</td>\n",
|
||||
" <td>400</td>\n",
|
||||
" <td>11.622357</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>9</th>\n",
|
||||
" <td>5.738497e+04</td>\n",
|
||||
" <td>7500</td>\n",
|
||||
" <td>matern</td>\n",
|
||||
" <td>constant</td>\n",
|
||||
" <td>100</td>\n",
|
||||
" <td>7.651330</td>\n",
|
||||
" </tr>\n",
|
||||
" </tbody>\n",
|
||||
"</table>\n",
|
||||
"</div>"
|
||||
],
|
||||
"text/plain": [
|
||||
" NLL N Kernel Mean K Mean_NLL\n",
|
||||
"0 6.726755e+04 4350 matern ewma 100 15.463804\n",
|
||||
"1 5.704342e+04 4200 matern ewma 200 13.581766\n",
|
||||
"2 6.180194e+04 6300 matern ewma 400 9.809832\n",
|
||||
"3 5.138177e+05 4800 matern dewma 100 107.045358\n",
|
||||
"4 1.109547e+05 4200 matern dewma 200 26.417781\n",
|
||||
"5 7.489109e+04 6450 matern dewma 400 11.611021\n",
|
||||
"6 3.585653e+06 5250 matern tewma 100 682.981476\n",
|
||||
"7 1.216625e+05 3900 matern tewma 200 31.195507\n",
|
||||
"8 7.496420e+04 6450 matern tewma 400 11.622357\n",
|
||||
"9 5.738497e+04 7500 matern constant 100 7.651330"
|
||||
]
|
||||
},
|
||||
"execution_count": 108,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"df"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 110,
|
||||
"id": "733dcdd7-bc8a-4530-9ac5-eaa2e5f86085",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"pd.to_pickle(df, \"./matern_nll.pkl\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 56,
|
||||
"id": "330cbdba-bf55-4bf6-814a-09d10a53fd04",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"mean = torch.tensor([1., 5.])\n",
|
||||
"std = torch.tensor([1., 1.])\n",
|
||||
"nrml = torch.distributions.Normal(mean, std)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 57,
|
||||
"id": "1582602c-8015-476a-b01a-a05d8aaf4a4a",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"tensor([0.3989, 0.3989])"
|
||||
]
|
||||
},
|
||||
"execution_count": 57,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"nrml.log_prob(torch.tensor([1., 5.])).exp()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 68,
|
||||
"id": "0f47d060-7f64-4537-b18f-7c5772df4bea",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"preds, y = GetNLL('matern', 'tewma', 400, logger=logger, horizon=horizon)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 69,
|
||||
"id": "e2403da1-c3b7-4d9d-8ae2-c00babfcfdcd",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"torch.Size([1000, 25])"
|
||||
]
|
||||
},
|
||||
"execution_count": 69,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"preds.shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 73,
|
||||
"id": "ff68035a-39a5-44e1-bf6a-ecd9b0837bc9",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"nll = 0."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 78,
|
||||
"id": "012b1b96-1afb-4f0a-8ab8-f96e92beacaf",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"nll -= torch.distributions.Normal(preds.mean(0), preds.std(0)).log_prob(y).sum()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 79,
|
||||
"id": "5a9241b1-a686-40fe-a378-cc219333f9bb",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"tensor(242.3718)"
|
||||
]
|
||||
},
|
||||
"execution_count": 79,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"nll"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "5cc3b61d-a2eb-4a33-9dfa-1049c9992345",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.8"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,408 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "9754a692-edb8-41bf-8159-fd82edcb9195",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Warning no robinhood utils.\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import pandas as pd\n",
|
||||
"import torch\n",
|
||||
"from torch import nn\n",
|
||||
"import seaborn as sns\n",
|
||||
"import time\n",
|
||||
"import copy\n",
|
||||
"import sys\n",
|
||||
"from torch.utils.data import DataLoader\n",
|
||||
"from voltron.data import make_ticker_list, GetStockHistory\n",
|
||||
"sns.set_style('white')\n",
|
||||
"# style.use('whitegrid')\n",
|
||||
"palette = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"sns.set(palette = palette, font_scale=2.0, style=\"white\", rc={\"lines.linewidth\": 2.0})"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "dbc24c1a-f499-4b8f-8544-6d4a66f48488",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Dataset Setup"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "62f91be1-38aa-4781-8b65-c5caf9752045",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from torch.utils.data import Dataset\n",
|
||||
"\n",
|
||||
"class SequenceDataset(Dataset):\n",
|
||||
" def __init__(self, data, sequence_length=5):\n",
|
||||
" self.sequence_length = sequence_length\n",
|
||||
" self.X = data.float()\n",
|
||||
"\n",
|
||||
" def __len__(self):\n",
|
||||
" return self.X.shape[0]-1\n",
|
||||
"\n",
|
||||
" def __getitem__(self, i): \n",
|
||||
" if i >= self.sequence_length - 1:\n",
|
||||
" i_start = i - self.sequence_length + 1\n",
|
||||
" x = self.X[i_start:(i + 1)]\n",
|
||||
" else:\n",
|
||||
" padding = self.X[0].repeat(self.sequence_length - i - 1, 1).squeeze(-1)\n",
|
||||
" x = self.X[0:(i + 1)]\n",
|
||||
" x = torch.cat((padding, x), 0)\n",
|
||||
" \n",
|
||||
" return x.unsqueeze(0), self.X[i+1]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 239,
|
||||
"id": "566c9386-c297-430f-b4ef-a0f3ba7aec37",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"tckr = \"JPM\"\n",
|
||||
"ntrain = 400\n",
|
||||
"lookback = 1\n",
|
||||
"data = GetStockHistory(tckr, end_date=\"2021-12-07\", history=ntrain + lookback).Close.to_numpy()\n",
|
||||
"data = torch.FloatTensor(data).log()\n",
|
||||
"\n",
|
||||
"data = (data - data.mean())/data.std()\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"# xin = torch.linspace(0, 5*np.pi, 250)\n",
|
||||
"# data = torch.sin(xin) + 0.2 * torch.randn(xin.shape)\n",
|
||||
"\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 269,
|
||||
"id": "5dece53d-785f-4ca9-9cbf-03e16de71e9b",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"seq_len = 25\n",
|
||||
"dset = SequenceDataset(data, seq_len)\n",
|
||||
"\n",
|
||||
"trgts = []\n",
|
||||
"for i in range(len(dset)):\n",
|
||||
" trgts.append(dset[i][1].item())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 270,
|
||||
"id": "2317ba11-4f62-443f-9ad8-63f9a0996ee9",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"train_loader = DataLoader(dset, batch_size=20, shuffle=True)\n",
|
||||
"X, y = next(iter(train_loader))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 271,
|
||||
"id": "ffecfde0-6938-4dee-b1b0-6627012e3329",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"class ShallowRegressionLSTM(nn.Module):\n",
|
||||
" def __init__(self, input_size, hidden_units=128):\n",
|
||||
" super().__init__()\n",
|
||||
" self.input_size = input_size # this is the number of features\n",
|
||||
" self.hidden_units = hidden_units\n",
|
||||
" self.num_layers = 5\n",
|
||||
"\n",
|
||||
" self.lstm = nn.LSTM(\n",
|
||||
" input_size=input_size,\n",
|
||||
" hidden_size=hidden_units,\n",
|
||||
" batch_first=True,\n",
|
||||
" num_layers=self.num_layers\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
" self.linear = nn.Linear(in_features=self.hidden_units, out_features=2)\n",
|
||||
" \n",
|
||||
"\n",
|
||||
" def forward(self, x):\n",
|
||||
" batch_size = x.shape[0]\n",
|
||||
" h0 = torch.zeros(self.num_layers, batch_size, self.hidden_units).requires_grad_()\n",
|
||||
" c0 = torch.zeros(self.num_layers, batch_size, self.hidden_units).requires_grad_()\n",
|
||||
"\n",
|
||||
" _, (hn, _) = self.lstm(x, (h0, c0))\n",
|
||||
" out = self.linear(hn[0]) # First dim of Hn is num_layers, which is set to 1 above.\n",
|
||||
" output = torch.zeros_like(out)\n",
|
||||
" output[:, 0] = out[:, 0]\n",
|
||||
" output[:, 1] = torch.exp(out[:, 1])\n",
|
||||
" return output\n",
|
||||
" \n",
|
||||
"from torch.autograd import Variable \n",
|
||||
"class LSTM1(nn.Module):\n",
|
||||
" def __init__(self, num_classes, seq_len, hidden_size, num_layers):\n",
|
||||
" super(LSTM1, self).__init__()\n",
|
||||
" self.num_classes = num_classes #number of classes\n",
|
||||
" self.num_layers = num_layers #number of layers\n",
|
||||
" self.input_size = seq_len #input size\n",
|
||||
" self.hidden_size = hidden_size #hidden state\n",
|
||||
"\n",
|
||||
" self.lstm = nn.LSTM(input_size=seq_len, hidden_size=hidden_size,\n",
|
||||
" num_layers=num_layers, batch_first=True) #lstm\n",
|
||||
" self.fc_1 = nn.Linear(hidden_size, 128) #fully connected 1\n",
|
||||
" self.fc = nn.Linear(128, num_classes) #fully connected last layer\n",
|
||||
"\n",
|
||||
" self.relu = nn.ReLU()\n",
|
||||
" self.softplus = nn.Softplus()\n",
|
||||
" \n",
|
||||
" def forward(self,x):\n",
|
||||
" h_0 = Variable(torch.zeros(self.num_layers, x.size(0), self.hidden_size)) #hidden state\n",
|
||||
" c_0 = Variable(torch.zeros(self.num_layers, x.size(0), self.hidden_size)) #internal state\n",
|
||||
" # Propagate input through LSTM\n",
|
||||
" output, (hn, cn) = self.lstm(x, (h_0, c_0)) #lstm with input, hidden, and internal state\n",
|
||||
"\n",
|
||||
" hn = hn[self.num_layers-1]\n",
|
||||
" hn = hn.view(-1, self.hidden_size) #reshaping the data for Dense layer next\n",
|
||||
" out = self.relu(hn)\n",
|
||||
" out = self.fc_1(out) #first Dense\n",
|
||||
" out = self.relu(out) #relu\n",
|
||||
" out = self.fc(out) #Final Output\n",
|
||||
" \n",
|
||||
" output = torch.zeros_like(out)\n",
|
||||
" output[:, 0] = out[:, 0]\n",
|
||||
" output[:, 1] = self.softplus(out[:, 1])\n",
|
||||
" return output"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 272,
|
||||
"id": "3bc83b15-3b0f-4f21-803b-29e354cb558f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# model = ShallowRegressionLSTM(seq_len)\n",
|
||||
"model = LSTM1(2, seq_len, 128, 1)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 273,
|
||||
"id": "41547213-7851-4ad1-ac04-03cab5e0a91c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def NLL(targets, outputs):\n",
|
||||
" dist = torch.distributions.Normal(outputs[:, 0], outputs[:, 1])\n",
|
||||
" return -dist.log_prob(targets).sum()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 274,
|
||||
"id": "e6958f51-5cd0-414a-b419-3ce55f12b486",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def train_model(data_loader, model, loss_function, optimizer, epochs=200):\n",
|
||||
" num_batches = len(data_loader)\n",
|
||||
" total_loss = 0\n",
|
||||
" model.train()\n",
|
||||
" for epoch in range(epochs):\n",
|
||||
" for X, y in data_loader:\n",
|
||||
" output = model(X)\n",
|
||||
" loss = loss_function(y, output)\n",
|
||||
"\n",
|
||||
" optimizer.zero_grad()\n",
|
||||
" loss.backward()\n",
|
||||
" optimizer.step()\n",
|
||||
"\n",
|
||||
" total_loss += loss.item()\n",
|
||||
"\n",
|
||||
" if epoch%10 == 0:\n",
|
||||
" avg_loss = total_loss / num_batches\n",
|
||||
" print(f\"Train loss: {avg_loss}, Epoch: {epoch}\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 275,
|
||||
"id": "92484e45-11c9-4fe1-bc5b-87b0f27d5a2c",
|
||||
"metadata": {
|
||||
"tags": []
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Train loss: 14.94985544681549, Epoch: 0\n",
|
||||
"Train loss: 2.123787060379982, Epoch: 10\n",
|
||||
"Train loss: -118.20948788821697, Epoch: 20\n",
|
||||
"Train loss: -225.282495072484, Epoch: 30\n",
|
||||
"Train loss: -393.8007117182016, Epoch: 40\n",
|
||||
"Train loss: -546.9331771463155, Epoch: 50\n",
|
||||
"Train loss: -724.232698109746, Epoch: 60\n",
|
||||
"Train loss: -916.3166337996721, Epoch: 70\n",
|
||||
"Train loss: -1073.7028237611055, Epoch: 80\n",
|
||||
"Train loss: -1267.8471668988466, Epoch: 90\n",
|
||||
"Train loss: -1434.5194520920516, Epoch: 100\n",
|
||||
"Train loss: -1630.7973297566175, Epoch: 110\n",
|
||||
"Train loss: -1831.4267205685378, Epoch: 120\n",
|
||||
"Train loss: -2038.7122355431318, Epoch: 130\n",
|
||||
"Train loss: -2246.2204612165688, Epoch: 140\n",
|
||||
"Train loss: -2457.420981016755, Epoch: 150\n",
|
||||
"Train loss: -2669.7478996187447, Epoch: 160\n",
|
||||
"Train loss: -2884.8198515325785, Epoch: 170\n",
|
||||
"Train loss: -3108.3130769401787, Epoch: 180\n",
|
||||
"Train loss: -3337.9945754677055, Epoch: 190\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"optimizer = torch.optim.Adam(model.parameters(), lr=0.01)\n",
|
||||
"train_model(train_loader, model, NLL, optimizer)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 276,
|
||||
"id": "07eab392-37e2-4dae-bfa4-dd177e3d552f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"means = []\n",
|
||||
"vrs = []\n",
|
||||
"for X, y in dset:\n",
|
||||
" output = model(X.unsqueeze(0))\n",
|
||||
" means = means + list(output[:, 0].detach().numpy())\n",
|
||||
" vrs = vrs + list(output[:, 1].detach().numpy())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 277,
|
||||
"id": "b10866df-b554-44a6-bf73-b98abb2e6400",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"<matplotlib.collections.PathCollection at 0x7fa59d357370>"
|
||||
]
|
||||
},
|
||||
"execution_count": 277,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYIAAAEFCAYAAADuT+DpAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/MnkTPAAAACXBIWXMAAAsTAAALEwEAmpwYAABsMElEQVR4nO29eXxdV3nv/Vt7PJOOZsmTPMay4zF2BpyQYAIJIW1DKVMCobS0lFvgbfsylHuhvL2XhLZMbSGhgUJoS5sAKSUlDcRJEy44TuIMTpzYluMhii3Lg2zpSEdn3tNa7x9r731m6WiWddb388kn9jn77LP3tvQ8az3D7yGMMQaBQCAQ1C3SXF+AQCAQCOYW4QgEAoGgzhGOQCAQCOoc4QgEAoGgzhGOQCAQCOocZa4voBZyuRwOHTqE9vZ2yLI815cjEAgEFwWO42BwcBCbNm1CIBCoetxF4QgOHTqE22+/fa4vQyAQCC5K7r//flxxxRVV378oHEF7ezsAfjOLFi2a46sRCASCi4OBgQHcfvvtvg2txkXhCLxw0KJFi7Bs2bI5vhqBQCC4uBgvpC6SxQKBQFDnCEcgEAgEdY5wBAKBQFDnCEcgEAgEdY5wBAKBQFDnCEcgEAgEdY5wBAKBQDAN2JTBoRfneBfhCAQCgWCKMMZwLm6gL5ZFzqJzfTkT5qJoKBMIBIK5xqEMhAASIUWvmzbFuVEDpsOgEIKBhIFFUR2xtImwJsNhQFNQgSyRKmeee4QjEAgEgnEwLIrT8RyCqoQlTcXibYmcDcth0GUCQggMm+JMPAfGgIzhAAAsh6KzQQMh89MZCEcgEAiqwhgDZZjXq9mZZiRtIZa2IBEgY1I4lPnPgzKGRNaG6joBAFAkAocx6AqPvDPGkMw6oNRERJegKTIypoOILhd9zjtf6Y5jNhCOQCAQVIQxhtGsjZGMjWXNOlS5/lKKqZyNWNqCKhNI7mo/Z1GEda7dkzEdOAxQC4y3LBHIyP+dEAJdAXKWg7TpALAAAMNpC5GA7O8UGGPoH86hOaQgGlRn9T7r719WIBCMSzJn42Qsi8GkBZsyDCYtMHZxVsRMBMYYchaFTRkMm2IgYUKRiL9KlwgwmrUBAImshXOj/P3xIIRAlSUEFAm6TBBQJGgyQTLrIOMml23KYNoMQ6nZf9ZiRyAQ1CHVQj6eIRxMmpAIQUDlK9W06cCwGQLqwg4R5WyK08M5gBDoCr/XwmekSAQZ04FpU4zmHCgSqckRFOKFggghUGRgKGki2BKAafNkNGWAQwFlFmdwiR2BQFCHxLM2To/k/JWn7TBQxmA5DGfiBoC8ASSEuCtha86ud7ZI5RwQiUCTCbIWLTPynhEfzdowLAp5in5RkQgsh2E4bSFnOf7rpjO7JahiRyAQ1BkO5YaHMsB0GCTCcGo4B5kQhHS+NizNB6gSQTLnoC3C5n3iOGdR5CwHTaGJxdkdypDI2VDdUFBIrbwkVyTih4emowpIkwniWRuSe24eIqIIabO3JRA7AoGgDqCM+av/nEXBGEAAZE0HGcMBYwBjwGjWqRjq8AzecNqCYc/vhqlY2sRwmsfZHcqQNZ3xPwSeGGasvE+gFIkAINyATweEEEgAGPguTJG4Y5jNPIFwBALBAocyHu5JuKvYtOGAgK8+h1IWhtIWN0BuErPail+VCUYyNoZScxciylkUqZxd9f20YSNjUlDG4/2xtIUzcQNWDaGWjEVRy2aHEAJdlqa1J0CVJWjuLkyWCGyHIZmzkTUdjKRn/nkLRyAQLHDShoOs6WAobYG61TCyRCBLBKpEIJPaEp4SIQgoBFnTgT1HmjqjWQtnR01krfJVPmO8ukmVCCQJOJ8wMZq1QQgwMGrCsOmYWkCWMzc1/JVQJYILKd6/MJKZ+Soi4QgEggVOzqKQCQFjgGFTmHZ+5Su5DqFWvFWwMQd6OowxpA0HEgFG0naZUbcpg00ZJAKokgSHMmgygSZLMB2K/uEczsRzuJAwK+oB2U5tO4LZQJJ4uChtOHx3M8PPWzgCgWCBk7Oob+xThgNGpp7kzFVYkc80psNAwWPzGdPB6XiuaKVs2vzP3r1psuSv8DWZ1+2bNsNIxkLKsGE5FBnT4aW0lJfTzidUWUJQ5SbameEdgagaEggWMIzxChRV5iv/pOHwrOQUkCVeWjmbMMaQM/m1E7fGP2dTZK18dU3adDCWeyOEQJYA4jaFjWZtUAosbdahyvlS2fkEv56Z91BiRyAQzBEOZRMKsUwmTmxTBgbXCBJ+DmmKv/Uy4SEmOs715CwK28kfQxmrGNuvhVPDOQxnbMiFUg4EGM3Y/rmTORvKOJU8isRDRQSAQghAuHroXOU85gvCEQgEc0TasHEhadZ0rE15rX9igk1dhaWThHAjqE7RExA333Bu1Kh6DGUM5+I59I/kfIeRMSnOxs0JG13mNrrxJHf+dUUiSJsOLIci61YK1ZrsVWWJ50dcp+Y4whEIBII5IGNSZC2naNVcjVjKhOlQDE5QhyaZc2YkAarJxDW+5dfCXEVOh/Fdjxe7T+Z4XD6entg9eH4jrMlFoRvvzyMZG/GMNakuX09ILmPRMcNKCx3hCASCWcS0Kc7FDeQsBxnTAUCKEq+WU25cLYcimXOgy5Jb+TO2EXUor66xHFpRJmE68IywVeLETJvi9IiBwRQv4wQAw3L8ip+gKmEky3dCtTqD0u8oRJN5l2/anNx9SgSwbIaM4cz7jumZRCSLBYJZwnIohtMWEjkbhkPhMB7nzlgUkQBfSZ+NGwjrsp8AHU5bUGQCgnziMGM6UGWCrOWUrZIBLo18dtRAUJX8/MDM3ROD7loRyhjOjhpwaH5Ii8QYEjkHulv9IhECXeY7lcZgbSJ2Nq2eRyFubwNjbFL3SQgBA4PpUASUya2LYz29OLt7H8xEGlo0jCU7r0DrxjWTOtdcIRyBQDAL2A7XmrcpQ1CVkLOZL1HgxfENm4dRDJvPAAAAMJ7s9YyUIhHecWo5SBsOljUHyjRpUoYDVeIhj6nmA8bDsPiAFQBIZvmkrkKDqkgEOdvBYCr/Gc+h5SwHAVUCZcx3dN4MBMNm6GhQQQgPQY1n4qfi7DSZgIFM+Byxnl70P74XTs7EsGHjQDyH7ZYD8+HdSJ0+jxU3XTPpa5pthCMQCKYJL6RTmrB0KMOFpAHKuEEnhCDoroQZYzAdromTMmwQCQjIUkHYhICyvKGTCK+nNx3GNWkydpEjcKi7Y5AIVDKzTkAmxG90silDLG1Bq6DWqckSciUxeJkQZAyKxiBD/7ABTSFYFNWQMXkeBAB0hQuw1VINNBUIIVUdTaGxBwA5qKPrhh0AgL5dT4HZ3In/4PVh9KZM/OLMKD6/qRPYfwRD+4/k79f93HzdKQhHIBBMEyNpCwxAW0Tjf89YCKkyUoaNtOFAV8r1abzVcdqw+cjDAulnj0IbSAiBJDFIICCEN3Z5TiOWtmDajAvKzUI9vCxxPZ9E1kIyxztg1QoGWyIEhLCSKV5A1nZgOQyWQ2E6QN9wDpbD/GfgD2ghpOizs0XfY88UGXMAcLIG+h7ZA0lTwWwHhkPx1cMXcN7VP8o6DD/pi2NjUwAxw8HSoIrVEQ3NAE4+vBv9Tzw7Lx2CcAQCwTTAGPObtVrDXOpgKGUhrDnIWRRaBSfgIROCwRSXha7F4BWGeyx3aAzABeEYmz5VzPHwQjnnk3ye71jfq5fIWnufTRncgGoyAaXwcwtAfh7CXExGi/X0+k6AMob7TozgVNrEDYsbsKMtDCfLS2cPxHO+E7iyNYT9wxkciOdwIJ7zzxWQCP5sfTu6whp3JLueAgDfGcyHHINwBALBNGA5zC8DNV3lSAIgbXL9mrHq22UJACWoIn8/LufiBjRF4sZ4kgnPqcCdwOS+N+GWt0qEQKpy/3PR7dv/xLMAeKjtp/1xPB/LAACeOJfEjrawf9yBkSwA7gQ+uLIZo6aDY8ni/oocZfjZ6VH8ybp2AACzHZx8eLf/fmGIyUykyxzFbCAcgUAwRRhjGEqZIMTT9LeRyNnQ5NoSkN7IwskgSzyHkLUoQur0OgFKGR564iVsuGQJ1q1eXPGYWu+xGpbDK4zmmsJVuRzQ/JzAA6fieGYw7R83kLORsSlCigSLMhwe5Sv/31oahSwRvHVRBMeSBq7rCOMtnQ1wGMPXDl/A0YSBr/Scx462ME6kDEiE4C0PPI6uaNB3AknLgUUZWgD0P75XOAKB4GLCtBnSJoUuc6PsSR/PxkpWcaekqNLkyidLKYyL77mQwgN9cYQVCd//83dhxRUbyo6fyneqUnEifK6I9fQWrco9J3BgJIu9rhO4fWUz9g6l8XrKxMm0iQ2NARxL5GBQhmUhFa1uDe3GpiDu2LIITZrs7wJvW9mMH50cQX/GQv+puP+9z8cy+O1ljbhxcQMA4K6jgxjI2vjkpe1Y7V7XbDkD0VAmEEyRZC5v+CXCf6lKq2dmmul0AmcyJr50cAAP9MUBAGmb4t5/fhR9jz0z5e8oRHKH4cw1Z3fv853AiZSB8zkLZzMW/rk3BgbgbYsbcHV7GKvcIoAjozmcyZj49vEYAGBLU7DofC26UhQKvLI1hDu3Lsa17Tyk1BVS0e46jkfPJpC0HGRsinNZGwzAT9znfnb3vhm862LEjkAgmCIpw/GTvIQQaMrcG7dqVEtMeslRxhh+fDKOgZwNXSJYGlLxesrEM4Np/Ma+wwBwUdXH14KZ4Kv+YcPG3x8ZLJKjXhnWcMvSKABga1MQvxxI4cXhDKyCg65sDY37HWFFwm0rm/HOrkYE3HzKPceGcHg0h6cupLGuUfePPZe1QBnzr2s2EI5AIJgCjjsMZbYqdaaCt+LP2BT/8noMy0Oj+K3kk+h79CkwywFlDH996DwGcjY0ieDOrYsRUiR87fB59KUtHIjnoO0/gsiyznlX/jgVtGgYZiKNVxO5spkElzUH/d3WqoiGNl3GkOFgjxsy+pN1bWgP1G5Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"xx = torch.arange(data.shape[0])\n",
|
||||
"\n",
|
||||
"plt.plot(xx[1:], means)\n",
|
||||
"plt.fill_between(xx[1:], means - 2*np.sqrt(vrs), means + 2*np.sqrt(vrs), color=palette[1], alpha=0.5)\n",
|
||||
"plt.scatter(xx, data, color=palette[5])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "6be11181-0698-429c-9dc7-a0f4b6c99584",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Rollouts"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 267,
|
||||
"id": "81bb90d7-e790-4e31-ba87-fe61449eb9ce",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"nroll = 50\n",
|
||||
"roll_len = 100\n",
|
||||
"xin, xout = dset[len(dset)-1]\n",
|
||||
"xx = torch.cat((xin[0, 1:], xout.unsqueeze(0)))\n",
|
||||
"xx = xx.repeat(nroll, 1).unsqueeze(1)\n",
|
||||
"roll_pxs = torch.zeros(nroll, roll_len)\n",
|
||||
"with torch.no_grad():\n",
|
||||
" for idx in range(roll_len):\n",
|
||||
" out = model(xx)\n",
|
||||
" smpl = torch.normal(out[:, 0], out[:, 1])\n",
|
||||
" roll_pxs[:, idx] = smpl\n",
|
||||
" xx = torch.cat((xx[..., 1:], smpl.unsqueeze(-1).unsqueeze(-1)), -1)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 268,
|
||||
"id": "1aec9c78-083b-4161-af02-4ef9e6590858",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAX8AAAEFCAYAAAAL/efAAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/MnkTPAAAACXBIWXMAAAsTAAALEwEAmpwYAABLQElEQVR4nO3de3xcdZ34/9eZM/dMMrk0lyZN06ZtStMrKWBQsCK3roAXEIpU2WVl/e667rq7qLuy7E+BXffL7uoq+FUUdFcEEYEqolCkUEqFtkDT0qalCb0kTdrmfp3J3Of8/pg5JzOZSZq0SdNk3s/Hw4fNzJmZk055n895f96f90fRNE1DCCFERjFN9wkIIYQ49yT4CyFEBpLgL4QQGUiCvxBCZCAJ/kIIkYHM030C4+H3+6mvr6ewsBBVVaf7dIQQYkaIRCJ0dnayYsUK7HZ70nMzIvjX19ezcePG6T4NIYSYkZ544gkuuuiipMdmRPAvLCwEYr9ASUnJNJ+NEELMDG1tbWzcuNGIoYlmRPDXUz0lJSXMmzdvms9GCCFmlnTpcpnwFUKIDCTBXwghMpAEfyGEyEAS/IUQIgNJ8BdCiAwkwV8IITKQBH8hhIjr6uoiEAic9fuEQiE6Ojom4YymjgR/IUTGOHbsGF1dXaM+f/ToUbq7u8/ovTVNIxKJ0NzcTFNTE+++++6ZnuY5IcFfCJEx3nnnHXp7e9M+19LSQiAQwOv1ntF7nzp1ir1799LV1UVjYyPt7e1A7C6gubmZUChEMBgc13t1dXXR399/RucxXhL8hRAzXk9PT8pjoVCIgYEB42ev14vD4cDn8/Haa6/h8/mSjt+zZw99fX1JwX+8wRqgvb2d/v5+hoaG8Hq9hMNhAAYGBnjvvfc4cuQIR44cSXndwMAAx48fByAQCNDS0kJdXR3Nzc3j/uwzIcFfCDHjvfjii3i9Xurr643HmpqaOHz4MBAL7O+88w4XXnghQ0NDNDY20tfXZxzb0dFBR0cHzc3N9Pb20tbWBsAzzzxj/FnTNLxeLy+//LLxmE7TNLZu3UokEqG/v59oNIrb7cbn83H06FF6e3vZu3cvp06d4sSJEwQCATRNY+fOndTX1/P222+jaRoDAwM0NDQwODhopJ9aWlqm5O9sRvT2EUKIsTgcDhobG2loaODkyZMUFxfT39+P2Wzm6NGj7Ny5E6fTic/nIxQKMX/+fLq7uyksLMRsNvPSSy/R3t6OyWTiwIEDdHd3U1paSk5ODg0NDZSUlDAwMEBdXR379u1j//793HDDDSxZsoTt27fT3NyMx+Phvffew2azkZWVRUFBAb29vdTV1REMBunp6aG7u5tIJEJlZSXZ2dns2rULRVGw2+2cPHmSV199lVOnTmG1Wjl58iQXXngh77//PuXl5ZP+dybBXwgxo0UiERRFoaOjg7a2Ntra2mhubmbBggVEo1H27t3LsmXL2LdvHwAej4fCwkK2b99OOBxm9erVtLa2Eg6HMZvN9PT0YDKZOHToEAUFBbjdbt555x1KS0vp7u7GZrOhaRrbt29nyZIlHDp0iI6ODjRNo6+vj9zcXLxeL263m56eHrq6unA4HKiqysDAgHE34na7CQQCWK1W+vv7efPNN2lqaiIajZKVlUVfXx8vvvgiOTk5U/L3JsFfCDEjHTlyhEWLFnHy5EkjXRIOh7HZbPT19dHe3s68efMYGBhg0aJFDA4OAjA4OIjf7wdg9+7dzJ8/38jth8NhTCYT4XCYSCRCT08PHo+HU6dOcd111xk5fZPJxMDAANFolM7OTiO/DxiPNzU10drais1mw+v1YjabCQaDdHZ24nK5UFWVaDSKx+MhEolw/Phx4zGfz0d2djZtbW14PJ4p+fuT4C+EmJG2b9+O3+9n+/btRCIRAoEAiqKgaRqhUIihoSGi0SihUIgdO3ZgMpnw+/1Eo1EikQiqqqIoCi+88AJms9kI4KqqGpU2kUgEn8+H1Wpl7969BINBwuGwEaR37NhBJBIxzkkv9wSwWq10d3djsViIRqOYTMNTrK2trXR0dBAKhYhGo0DsjkRRFCA2OW2327HZbMaFarJJ8BdCzEitra04nU66urqMwKoHcT2Fc+LECbq7uzGZTGiaZlwgAoEAOTk5tLW1UVlZmRTAw+EwmqYBsWDudDoZGhriyJEj2Gw2wuEwoVCI/Px8tm7dahw7ks/nIxqNMjQ0BMSqjwAURSEYDKIoihH4dYnv1dvbi8lkwuVyTd5fWgKp9hFCzDh6sG5qakoKojabzXhOr+AJBoOoqkooFELTNGMk7vF4CIVCNDQ0JAXdkcFcD96KojA0NGQ87/V6iUQixs/6qF3X39+fdFHR/2w2x8bcp1tJHI1GCYfDOJ3OCfzNjJ8EfyHEjLNnzx4cDoeR5tEDamJ+PBwO43A40DTNSJ3oKSH9eSApHTOWxDsCRVGS8vz6e4/HeNYO2Gw248/6YrHJJmkfIcSUaalr5ODmnfj6PDhyXVSvr6W8puqs37e1tRWr1UpPTw92uz1p1K0HYX3CFkgK1ImjcSAl9QIYF5VEiT+nS9mMh8lkGtfrEu8KxntRmfC5TMm7CiEy3t5N29j9yy34+jwMhiJ0dw2w+5db+OOPnjvr9+7u7jZy4ZqmJQXUxJH8mQTodIF/JE3Tzigon8lr9LuaySYjfyHEWRs5wi++oIKmnQcA8IQi3L+/jZAGGypyqeUEv/naD4BYoK34QDVrblw3oc/z+/1GINXLMnUjq2/GS1EUFEUhNzc3bbuIRGc6Gtc0DZvNNma+f+TFJ93m65NBgr8Q4qy01DVS9/SrDAXDDAQjFGmDRuDXNI3nWvsZisSC2ePHenn51CBWk8LqPAdr851o8WMncgEIhUJGHn9kGudsRKNRysvLjeA/njSNnnIaTyooMZef+FpN04zS05FzCfocxWSTtI8Q4qzs3bSNaDjCDxq6uL++naeP9xnP7e/zs6NrKOn4dn+YlqEQvzsxwM+OxoJs084DtNQ1jvszvV5vSg5+NOMdOVssFkwmE2azGavVCsTaRqSbEB45x6CnnkZ+VnZ2dtLPJpMpKZgnvreiKFit1nFPQJ8tGfkLIc7Y3k3biARD7O7xccwbq2J5t9fHLRV5ABzoH16gtDrPwbu9sU6aORYTA6EoTd4gf/N2KzfNdxN9+lUAymuqjDTSUO8gzrzspIniUChEOBzGYrEY761X/IwcNcPwnUFiOsVisSQFYf21iqJw4MAB4/30VJBOVVWjnUTixUe/Q7BYLEmjd5fLRTQaxev1GoE/8e7A6XQaK4/10k6d/hllZWXj/DYmRoK/EOKMtNQ10rTzAJ5QhCeahnvk94eiBCJRrCaFxoFYbvvPKvO5wG0jFI1ySUEWFxU4+dnRHt7uHkIDnjnez54eH18Ov0x30yladjfwu6YeXu/w8LcXhAk++xoQuzC88847QGo6JDFIJ14I9EVaFouFSCSSkibSF1L5/f6k+QO9lFS/I9A/T78AJH6ufiGw2WxG8FdVFZfLhdvt5tChQ0nnZzKZyM3NJT8/n9LSUhobG3E6ncYdjcViMS5GU3UnIGkfIabBG3VH+M+fvow/ODX53KnUUtfIS996jN2/3ALAjq4hQlGNqmwbJfbYeLLDH+aYJ0hnIIzLbOLCfAcus8oXqwq5qCC2aOnm+bncuaiAuY7Ya454gjxwoINXt9QRCYX5/ckBBsNRftDYRSQU5uDmnQBJbZsTA2riqHnJkiVJxyiKQiQSwW63A7ELh8lkwuFwUFZWht/vNwJ1NBo1WkAEAgFsNhuFhYVALLjPmTMHGE4nqaqK0+nEZDJhMpmwWq2YzWZyc3Nxu91cdtllZGdno2kaeXl5WCwWFixYgNPpZN68eRQUFGC1Wlm5ciWqqnLBBRcwf/58nE4niqIk7UkwmST4C3EOtdQ18vz/9yif++pP+d7Pt/KFP//OhHLd062lrpG9z76Gr8+Dpmm0+0P8/kSsD85Hil0UO2KpmCZvkP+J5/MvyneipsnJO80m1uQ7+KflxXyxag4Wk8IJX4jftibvYNUbjI2yfX2xBVyjBcPENMzChQuNEbPX6zVG0CUlJcYFw2Kx4Ha7ufbaa1FVlTvvvJPs7GyjFYTJZKKwsJCCggIKCwux2+2sWbMGi8WCxWIx8vOKopCVlUVubi6hUIi5c+dy2WWXsXz5cjRNw+fzsXr1agAWLFiA1WolOzub7OxsAoEAbW1trFy5krKyMvLz87nmmmsIBAKUlpYaF62pIMFfiHPkjz96jv99+Hn+5a1mgtFYoPrjqQHe/MUfZswF4ODmnURCYbacGuRv3znB/fvbCWtQ7baxMtdOpSs2Ufp86wC9wQgWBa4tzR7zPVVFodpt595VJQAcHgzS4U++IwpFNSNoJzY60wO+1WpFVVVj9L1q1SojMANGPr6srIzy8nJsNhsulwu73U5JSQl5eXnYbDacTic2m41oNEplZSWBQICKigrC4TD5+fnk5ubidDqZP3++cRHQP6OiooKioiLmz5+PzWajsrISk8nE888/j8ViMd4/NzeXwcFB5s2bx7Fjx1AUheLiYsrKyoy7A6/XS21tLXl5eRQVFU3CN5fqnAX/TZs2sXTpUiNfJ8RsFo1GGfQOB6m9m7bRcbiVX7f00xUYHskFoxr7e/3seXbrdJzmhOmj7z29sVy9bllObJVtTZ4DgKFIbFLz4oIssi3jq7bJsahcnB9LCT3fmjy6HwhFjNYMenoncQSvV+hkZ2eTl5eHyWQyKm304/Sfr732WqLRKEuWLMFisaCqqpEmcrvd2Gw2LBYLH/vYx/D5fNTU1BAOh1m4cCENDQ2YzWbcbrcx4Wwymejv7zc2cTGbzdhsNvLy8jh8+DCqqnLw4EFWrFjBkiVLqKioIDs7mzVr1hAKhbj11ltxuVz4fD6WLVsGwBVXXMG8efNYvHgxNTU14/x2JuacBP89e/Zw//33n4uPEmLatdQLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"full_x = torch.arange(data.shape[0])\n",
|
||||
"test_x = torch.arange(data.shape[0], data.shape[0] + roll_len)\n",
|
||||
"plt.plot(full_x[1:], means)\n",
|
||||
"plt.scatter(full_x, data, color=palette[5])\n",
|
||||
"plt.plot(test_x, roll_pxs[:20, :].T.detach(), color='gray', alpha=1., lw=0.5)\n",
|
||||
"plt.show()"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.12"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,854 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "92e5b691-e01a-438a-b2b9-543c5fa3f8b2",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Warning no robinhood utils.\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import pickle as pkl\n",
|
||||
"import pandas as pd\n",
|
||||
"import numpy as np\n",
|
||||
"import torch\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns\n",
|
||||
"import os\n",
|
||||
"from voltron.data import make_ticker_list, GetStockHistory\n",
|
||||
"\n",
|
||||
"sns.set_style('white')\n",
|
||||
"palette = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"sns.set(palette = palette, font_scale=2.0, style=\"white\", rc={\"lines.linewidth\": 4.0})"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 30,
|
||||
"id": "1f59e48e-eb80-48b2-8ce5-91596ad136f2",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def ECDF(sample_pxs, true_px): \n",
|
||||
" return (torch.sum(sample_pxs < true_px, 0)/sample_pxs.shape[0])\n",
|
||||
" \n",
|
||||
"def Calibration(pcts, percentile=0.95):\n",
|
||||
" in_band = np.where((pcts < percentile))[0].shape[0]\n",
|
||||
" return in_band/pcts.shape[0]\n",
|
||||
"\n",
|
||||
"def GetNLL(model, mean='ewma', k=100, horizon=np.arange(75,100), \n",
|
||||
" logger=[], exp=True, fdir=\"./saved-outputs/\"):\n",
|
||||
" \n",
|
||||
"\n",
|
||||
" ntrain = 400\n",
|
||||
" n_test_times = 20\n",
|
||||
" ntest = 100\n",
|
||||
" nll = 0.\n",
|
||||
" N = 0\n",
|
||||
" nlls = torch.tensor([])\n",
|
||||
" for tckr in ticker_list:\n",
|
||||
" data = None\n",
|
||||
" try:\n",
|
||||
" data = GetStockHistory(tckr, history=1000, end_date=end_date)\n",
|
||||
" except:\n",
|
||||
" print(\"failed\", tckr)\n",
|
||||
" \n",
|
||||
" if data is not None:\n",
|
||||
"\n",
|
||||
" for idx, date in enumerate(data.index):\n",
|
||||
" fpath = fdir + tckr + \"/\"\n",
|
||||
" fname = model + \"_\"\n",
|
||||
" if model in ['volt', 'matern', 'sm']:\n",
|
||||
" fname += mean + str(k) + \"_\"\n",
|
||||
"\n",
|
||||
" fname += str(date.date()) + \".pt\"\n",
|
||||
"# print(fpath + fname)\n",
|
||||
" if os.path.exists(fpath + fname): \n",
|
||||
" preds = torch.load(fpath + fname) \n",
|
||||
" if isinstance(preds, tuple):\n",
|
||||
" preds = preds[0]\n",
|
||||
" if preds.shape[-1] == 100:\n",
|
||||
" preds = preds[:, horizon]\n",
|
||||
"\n",
|
||||
" test_y = torch.tensor(data.iloc[idx:idx+100].Close.to_numpy())\n",
|
||||
" if test_y.shape[0] == 100:\n",
|
||||
" if exp:\n",
|
||||
" preds = preds.exp()\n",
|
||||
"\n",
|
||||
" try:\n",
|
||||
" curr = torch.distributions.Normal(preds.mean(0), preds.std(0)).log_prob(test_y[horizon])\n",
|
||||
" if curr.mean().abs() < 500:\n",
|
||||
" nlls = torch.cat((curr, nlls))\n",
|
||||
" except:\n",
|
||||
" pass\n",
|
||||
"\n",
|
||||
"\n",
|
||||
" if nlls.numel() > 0:\n",
|
||||
" logger.append([-nlls.sum().item(), -nlls.mean().item(), nlls.std().item(), model, mean, k])\n",
|
||||
" \n",
|
||||
" return logger"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 25,
|
||||
"id": "5fe6b4d6-b49e-4e06-8874-9be43b4ee4bd",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"data_path = \"../../voltron/data/\"\n",
|
||||
"ticker_list = make_ticker_list(data_path + \"nasdaq100.txt\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 27,
|
||||
"id": "d9ed9129-be2d-47e8-9866-5a708c27d225",
|
||||
"metadata": {
|
||||
"collapsed": true,
|
||||
"jupyter": {
|
||||
"outputs_hidden": true
|
||||
},
|
||||
"tags": []
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- ALXN: No data found, symbol may be delisted\n",
|
||||
"failed ALXN\n",
|
||||
"failed CA\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- CELG: No data found, symbol may be delisted\n",
|
||||
"failed CELG\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- CTRP: No data found, symbol may be delisted\n",
|
||||
"failed CTRP\n",
|
||||
"failed ESRX\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- LVNTA: No data found for this date range, symbol may be delisted\n",
|
||||
"failed LVNTA\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- QVCA: No data found for this date range, symbol may be delisted\n",
|
||||
"failed QVCA\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- LMCA: No data found for this date range, symbol may be delisted\n",
|
||||
"failed LMCA\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- LMCK: No data found, symbol may be delisted\n",
|
||||
"failed LMCK\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- LLTC: No data found for this date range, symbol may be delisted\n",
|
||||
"failed LLTC\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- MXIM: No data found, symbol may be delisted\n",
|
||||
"failed MXIM\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- MYL: No data found, symbol may be delisted\n",
|
||||
"failed MYL\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- SYMC: No data found, symbol may be delisted\n",
|
||||
"failed SYMC\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- PCLN: No data found for this date range, symbol may be delisted\n",
|
||||
"failed PCLN\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- VIAB: No data found, symbol may be delisted\n",
|
||||
"failed VIAB\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- WFM: No data found for this date range, symbol may be delisted\n",
|
||||
"failed WFM\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- YHOO: No data found for this date range, symbol may be delisted\n",
|
||||
"failed YHOO\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"log = []\n",
|
||||
"end_date = \"2022-01-20\"\n",
|
||||
"for k in [400]:\n",
|
||||
" for mean in ['ewma']:\n",
|
||||
" log = GetNLL('volt', mean=mean, k=k, horizon=np.arange(75,100), \n",
|
||||
" logger=log, exp=True, fdir=\"../trading/saved-outputs/\")\n",
|
||||
" \n",
|
||||
"# log = GetNLL('matern', mean='constant', k=100, horizon=np.arange(75,100), \n",
|
||||
"# logger=log, exp=True)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 31,
|
||||
"id": "c00697ab-5884-4c70-9843-1bb49a52716f",
|
||||
"metadata": {
|
||||
"tags": []
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- ALXN: No data found, symbol may be delisted\n",
|
||||
"failed ALXN\n",
|
||||
"failed CA\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- CELG: No data found, symbol may be delisted\n",
|
||||
"failed CELG\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- CTRP: No data found, symbol may be delisted\n",
|
||||
"failed CTRP\n",
|
||||
"failed ESRX\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- LVNTA: No data found for this date range, symbol may be delisted\n",
|
||||
"failed LVNTA\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- QVCA: No data found for this date range, symbol may be delisted\n",
|
||||
"failed QVCA\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- LMCA: No data found for this date range, symbol may be delisted\n",
|
||||
"failed LMCA\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- LMCK: No data found, symbol may be delisted\n",
|
||||
"failed LMCK\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- LLTC: No data found for this date range, symbol may be delisted\n",
|
||||
"failed LLTC\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- MXIM: No data found, symbol may be delisted\n",
|
||||
"failed MXIM\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- MYL: No data found, symbol may be delisted\n",
|
||||
"failed MYL\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- SYMC: No data found, symbol may be delisted\n",
|
||||
"failed SYMC\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- PCLN: No data found for this date range, symbol may be delisted\n",
|
||||
"failed PCLN\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- VIAB: No data found, symbol may be delisted\n",
|
||||
"failed VIAB\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- WFM: No data found for this date range, symbol may be delisted\n",
|
||||
"failed WFM\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- YHOO: No data found for this date range, symbol may be delisted\n",
|
||||
"failed YHOO\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- ALXN: No data found, symbol may be delisted\n",
|
||||
"failed ALXN\n",
|
||||
"failed CA\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- CELG: No data found, symbol may be delisted\n",
|
||||
"failed CELG\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- CTRP: No data found, symbol may be delisted\n",
|
||||
"failed CTRP\n",
|
||||
"failed ESRX\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- LVNTA: No data found for this date range, symbol may be delisted\n",
|
||||
"failed LVNTA\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- QVCA: No data found for this date range, symbol may be delisted\n",
|
||||
"failed QVCA\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- LMCA: No data found for this date range, symbol may be delisted\n",
|
||||
"failed LMCA\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- LMCK: No data found, symbol may be delisted\n",
|
||||
"failed LMCK\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- LLTC: No data found for this date range, symbol may be delisted\n",
|
||||
"failed LLTC\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- MXIM: No data found, symbol may be delisted\n",
|
||||
"failed MXIM\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- MYL: No data found, symbol may be delisted\n",
|
||||
"failed MYL\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- SYMC: No data found, symbol may be delisted\n",
|
||||
"failed SYMC\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- PCLN: No data found for this date range, symbol may be delisted\n",
|
||||
"failed PCLN\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- VIAB: No data found, symbol may be delisted\n",
|
||||
"failed VIAB\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- WFM: No data found for this date range, symbol may be delisted\n",
|
||||
"failed WFM\n",
|
||||
"\n",
|
||||
"1 Failed download:\n",
|
||||
"- YHOO: No data found for this date range, symbol may be delisted\n",
|
||||
"failed YHOO\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"log = []\n",
|
||||
"end_date = \"2022-01-20\"\n",
|
||||
"log = GetNLL('sm', mean='constant', k=100, horizon=np.arange(75,100), \n",
|
||||
" logger=[], exp=True)\n",
|
||||
"log = GetNLL('sm', mean='ewma', k=400, horizon=np.arange(75,100), \n",
|
||||
" logger=log, exp=True)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 32,
|
||||
"id": "89054cf9-fb1c-45a1-aa8b-8884713f4c07",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df = pd.DataFrame(log)\n",
|
||||
"df.columns = [\"Mean_NLL\", \"NLL\", \"Std_NLL\", \"Model\", \"Mean\", \"k\"]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 33,
|
||||
"id": "95915269-4bbf-4949-a944-37e076db8889",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"pd.to_pickle(df, \"./sm_nll.pkl\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 34,
|
||||
"id": "e2a5b10f-61a9-4876-a53c-385380955e28",
|
||||
"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>Mean_NLL</th>\n",
|
||||
" <th>NLL</th>\n",
|
||||
" <th>Std_NLL</th>\n",
|
||||
" <th>Model</th>\n",
|
||||
" <th>Mean</th>\n",
|
||||
" <th>k</th>\n",
|
||||
" </tr>\n",
|
||||
" </thead>\n",
|
||||
" <tbody>\n",
|
||||
" <tr>\n",
|
||||
" <th>0</th>\n",
|
||||
" <td>3.201143e+06</td>\n",
|
||||
" <td>80.430728</td>\n",
|
||||
" <td>113.825740</td>\n",
|
||||
" <td>sm</td>\n",
|
||||
" <td>constant</td>\n",
|
||||
" <td>100</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>1</th>\n",
|
||||
" <td>1.134694e+06</td>\n",
|
||||
" <td>147.842929</td>\n",
|
||||
" <td>161.222031</td>\n",
|
||||
" <td>sm</td>\n",
|
||||
" <td>ewma</td>\n",
|
||||
" <td>400</td>\n",
|
||||
" </tr>\n",
|
||||
" </tbody>\n",
|
||||
"</table>\n",
|
||||
"</div>"
|
||||
],
|
||||
"text/plain": [
|
||||
" Mean_NLL NLL Std_NLL Model Mean k\n",
|
||||
"0 3.201143e+06 80.430728 113.825740 sm constant 100\n",
|
||||
"1 1.134694e+06 147.842929 161.222031 sm ewma 400"
|
||||
]
|
||||
},
|
||||
"execution_count": 34,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"df"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "eb9db52e-d56a-47a8-86f5-05e43437c635",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## NLL Plotter"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 48,
|
||||
"id": "38b65cc2-ed27-41a6-ae83-5396a40b2e0d",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"nll = pd.read_pickle(\"./volt_matern_const_nll.pkl\")\n",
|
||||
"nll = pd.concat((nll, pd.read_pickle(\"volt_const_nll.pkl\")))\n",
|
||||
"mat_nll = pd.read_pickle(\"matern_nll.pkl\")\n",
|
||||
"mat_nll[mat_nll[\"Mean\"] != 'constant']\n",
|
||||
"mat_nll.columns = nll.columns\n",
|
||||
"nll = pd.concat((mat_nll, nll))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 49,
|
||||
"id": "7952ac2d-0036-4ce4-bcd9-f7b1cfe1ca68",
|
||||
"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>Mean_NLL</th>\n",
|
||||
" <th>NLL</th>\n",
|
||||
" <th>Std_NLL</th>\n",
|
||||
" <th>Model</th>\n",
|
||||
" <th>Mean</th>\n",
|
||||
" <th>k</th>\n",
|
||||
" </tr>\n",
|
||||
" </thead>\n",
|
||||
" <tbody>\n",
|
||||
" <tr>\n",
|
||||
" <th>0</th>\n",
|
||||
" <td>6.726756e+04</td>\n",
|
||||
" <td>15.463807</td>\n",
|
||||
" <td>82.817116</td>\n",
|
||||
" <td>matern</td>\n",
|
||||
" <td>ewma</td>\n",
|
||||
" <td>100</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>1</th>\n",
|
||||
" <td>5.704343e+04</td>\n",
|
||||
" <td>13.581769</td>\n",
|
||||
" <td>62.178253</td>\n",
|
||||
" <td>matern</td>\n",
|
||||
" <td>ewma</td>\n",
|
||||
" <td>200</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>2</th>\n",
|
||||
" <td>6.180195e+04</td>\n",
|
||||
" <td>9.809834</td>\n",
|
||||
" <td>21.427580</td>\n",
|
||||
" <td>matern</td>\n",
|
||||
" <td>ewma</td>\n",
|
||||
" <td>400</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>3</th>\n",
|
||||
" <td>5.138178e+05</td>\n",
|
||||
" <td>107.045380</td>\n",
|
||||
" <td>558.429443</td>\n",
|
||||
" <td>matern</td>\n",
|
||||
" <td>dewma</td>\n",
|
||||
" <td>100</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>4</th>\n",
|
||||
" <td>1.109547e+05</td>\n",
|
||||
" <td>26.417788</td>\n",
|
||||
" <td>160.354996</td>\n",
|
||||
" <td>matern</td>\n",
|
||||
" <td>dewma</td>\n",
|
||||
" <td>200</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>5</th>\n",
|
||||
" <td>7.489107e+04</td>\n",
|
||||
" <td>11.611018</td>\n",
|
||||
" <td>42.626156</td>\n",
|
||||
" <td>matern</td>\n",
|
||||
" <td>dewma</td>\n",
|
||||
" <td>400</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>6</th>\n",
|
||||
" <td>3.585652e+06</td>\n",
|
||||
" <td>682.981384</td>\n",
|
||||
" <td>2605.378174</td>\n",
|
||||
" <td>matern</td>\n",
|
||||
" <td>tewma</td>\n",
|
||||
" <td>100</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>7</th>\n",
|
||||
" <td>1.216624e+05</td>\n",
|
||||
" <td>31.195498</td>\n",
|
||||
" <td>169.649368</td>\n",
|
||||
" <td>matern</td>\n",
|
||||
" <td>tewma</td>\n",
|
||||
" <td>200</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>8</th>\n",
|
||||
" <td>7.496416e+04</td>\n",
|
||||
" <td>11.622351</td>\n",
|
||||
" <td>42.427406</td>\n",
|
||||
" <td>matern</td>\n",
|
||||
" <td>tewma</td>\n",
|
||||
" <td>400</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>0</th>\n",
|
||||
" <td>4.061322e+03</td>\n",
|
||||
" <td>7.735851</td>\n",
|
||||
" <td>4.734124</td>\n",
|
||||
" <td>matern</td>\n",
|
||||
" <td>constant</td>\n",
|
||||
" <td>100</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>0</th>\n",
|
||||
" <td>1.173272e+03</td>\n",
|
||||
" <td>4.693086</td>\n",
|
||||
" <td>0.389815</td>\n",
|
||||
" <td>volt</td>\n",
|
||||
" <td>constant</td>\n",
|
||||
" <td>100</td>\n",
|
||||
" </tr>\n",
|
||||
" </tbody>\n",
|
||||
"</table>\n",
|
||||
"</div>"
|
||||
],
|
||||
"text/plain": [
|
||||
" Mean_NLL NLL Std_NLL Model Mean k\n",
|
||||
"0 6.726756e+04 15.463807 82.817116 matern ewma 100\n",
|
||||
"1 5.704343e+04 13.581769 62.178253 matern ewma 200\n",
|
||||
"2 6.180195e+04 9.809834 21.427580 matern ewma 400\n",
|
||||
"3 5.138178e+05 107.045380 558.429443 matern dewma 100\n",
|
||||
"4 1.109547e+05 26.417788 160.354996 matern dewma 200\n",
|
||||
"5 7.489107e+04 11.611018 42.626156 matern dewma 400\n",
|
||||
"6 3.585652e+06 682.981384 2605.378174 matern tewma 100\n",
|
||||
"7 1.216624e+05 31.195498 169.649368 matern tewma 200\n",
|
||||
"8 7.496416e+04 11.622351 42.427406 matern tewma 400\n",
|
||||
"0 4.061322e+03 7.735851 4.734124 matern constant 100\n",
|
||||
"0 1.173272e+03 4.693086 0.389815 volt constant 100"
|
||||
]
|
||||
},
|
||||
"execution_count": 49,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"nll"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 52,
|
||||
"id": "76d6abf1-8385-4450-9c0e-4e92f897bff8",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/home/greg_b/miniconda3/envs/rpp/lib/python3.8/site-packages/pandas/core/indexing.py:1720: SettingWithCopyWarning: \n",
|
||||
"A value is trying to be set on a copy of a slice from a DataFrame.\n",
|
||||
"Try using .loc[row_indexer,col_indexer] = value instead\n",
|
||||
"\n",
|
||||
"See the caveats in the documentation: https://pandas.pydata.org/pandas-docs/stable/user_guide/indexing.html#returning-a-view-versus-a-copy\n",
|
||||
" self._setitem_single_column(loc, value, pi)\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"temp_df = nll[(nll['Mean'].isin(['constant', 'ewma']))]\n",
|
||||
"temp_df.loc[(temp_df['Mean']=='constant') & (temp_df['Model']=='volt'), 'k'] = 400\n",
|
||||
"temp_df.loc[(temp_df['Mean']=='constant') & (temp_df['Model']=='matern'), 'k'] = 400\n",
|
||||
"\n",
|
||||
"temp_df = temp_df[temp_df['k']==400]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 53,
|
||||
"id": "6582a178-a65f-4f24-a20e-8b8b17e56e5a",
|
||||
"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>Mean_NLL</th>\n",
|
||||
" <th>NLL</th>\n",
|
||||
" <th>Std_NLL</th>\n",
|
||||
" <th>Model</th>\n",
|
||||
" <th>Mean</th>\n",
|
||||
" <th>k</th>\n",
|
||||
" </tr>\n",
|
||||
" </thead>\n",
|
||||
" <tbody>\n",
|
||||
" <tr>\n",
|
||||
" <th>2</th>\n",
|
||||
" <td>61801.953125</td>\n",
|
||||
" <td>9.809834</td>\n",
|
||||
" <td>21.427580</td>\n",
|
||||
" <td>matern</td>\n",
|
||||
" <td>ewma</td>\n",
|
||||
" <td>400</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>0</th>\n",
|
||||
" <td>4061.321999</td>\n",
|
||||
" <td>7.735851</td>\n",
|
||||
" <td>4.734124</td>\n",
|
||||
" <td>matern</td>\n",
|
||||
" <td>constant</td>\n",
|
||||
" <td>400</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>0</th>\n",
|
||||
" <td>1173.271529</td>\n",
|
||||
" <td>4.693086</td>\n",
|
||||
" <td>0.389815</td>\n",
|
||||
" <td>volt</td>\n",
|
||||
" <td>constant</td>\n",
|
||||
" <td>400</td>\n",
|
||||
" </tr>\n",
|
||||
" </tbody>\n",
|
||||
"</table>\n",
|
||||
"</div>"
|
||||
],
|
||||
"text/plain": [
|
||||
" Mean_NLL NLL Std_NLL Model Mean k\n",
|
||||
"2 61801.953125 9.809834 21.427580 matern ewma 400\n",
|
||||
"0 4061.321999 7.735851 4.734124 matern constant 400\n",
|
||||
"0 1173.271529 4.693086 0.389815 volt constant 400"
|
||||
]
|
||||
},
|
||||
"execution_count": 53,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"temp_df"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 51,
|
||||
"id": "20d9682b-5e98-4793-8a65-42532cc476d9",
|
||||
"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>Mean_NLL</th>\n",
|
||||
" <th>NLL</th>\n",
|
||||
" <th>Std_NLL</th>\n",
|
||||
" <th>Model</th>\n",
|
||||
" <th>Mean</th>\n",
|
||||
" <th>k</th>\n",
|
||||
" </tr>\n",
|
||||
" </thead>\n",
|
||||
" <tbody>\n",
|
||||
" <tr>\n",
|
||||
" <th>2</th>\n",
|
||||
" <td>61801.953125</td>\n",
|
||||
" <td>9.809834</td>\n",
|
||||
" <td>21.427580</td>\n",
|
||||
" <td>matern</td>\n",
|
||||
" <td>ewma</td>\n",
|
||||
" <td>400</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>0</th>\n",
|
||||
" <td>4061.321999</td>\n",
|
||||
" <td>7.735851</td>\n",
|
||||
" <td>4.734124</td>\n",
|
||||
" <td>matern</td>\n",
|
||||
" <td>constant</td>\n",
|
||||
" <td>400</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>0</th>\n",
|
||||
" <td>1173.271529</td>\n",
|
||||
" <td>4.693086</td>\n",
|
||||
" <td>0.389815</td>\n",
|
||||
" <td>volt</td>\n",
|
||||
" <td>constant</td>\n",
|
||||
" <td>400</td>\n",
|
||||
" </tr>\n",
|
||||
" </tbody>\n",
|
||||
"</table>\n",
|
||||
"</div>"
|
||||
],
|
||||
"text/plain": [
|
||||
" Mean_NLL NLL Std_NLL Model Mean k\n",
|
||||
"2 61801.953125 9.809834 21.427580 matern ewma 400\n",
|
||||
"0 4061.321999 7.735851 4.734124 matern constant 400\n",
|
||||
"0 1173.271529 4.693086 0.389815 volt constant 400"
|
||||
]
|
||||
},
|
||||
"execution_count": 51,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"temp_df"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "2cb0ed8b-e4d7-4b78-b676-96d7123ad643",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.8"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,339 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "92e5b691-e01a-438a-b2b9-543c5fa3f8b2",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Warning no robinhood utils.\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import pickle as pkl\n",
|
||||
"import pandas as pd\n",
|
||||
"import numpy as np\n",
|
||||
"import torch\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns\n",
|
||||
"import os\n",
|
||||
"from voltron.data import make_ticker_list, GetStockHistory\n",
|
||||
"\n",
|
||||
"sns.set_style('white')\n",
|
||||
"palette = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"sns.set(palette = palette, font_scale=2.0, style=\"white\", rc={\"lines.linewidth\": 4.0})"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "a519062e-3992-4d40-9bf1-990c80990fc0",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"ename": "FileNotFoundError",
|
||||
"evalue": "[Errno 2] No such file or directory: './matern_calib.pkl'",
|
||||
"output_type": "error",
|
||||
"traceback": [
|
||||
"\u001b[0;31m---------------------------------------------------------------------------\u001b[0m",
|
||||
"\u001b[0;31mFileNotFoundError\u001b[0m Traceback (most recent call last)",
|
||||
"\u001b[0;32m<ipython-input-2-18f65b61756d>\u001b[0m in \u001b[0;36m<module>\u001b[0;34m\u001b[0m\n\u001b[0;32m----> 1\u001b[0;31m \u001b[0mmatern_df\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mpd\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mread_pickle\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m\"./matern_calib.pkl\"\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 2\u001b[0m \u001b[0mvolt_df\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mpd\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mread_pickle\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m\"./volt_calib.pkl\"\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 3\u001b[0m \u001b[0mmatern_df\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mcolumns\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mvolt_df\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mcolumns\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/pandas/io/pickle.py\u001b[0m in \u001b[0;36mread_pickle\u001b[0;34m(filepath_or_buffer, compression, storage_options)\u001b[0m\n\u001b[1;32m 183\u001b[0m \"\"\"\n\u001b[1;32m 184\u001b[0m \u001b[0mexcs_to_catch\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0;34m(\u001b[0m\u001b[0mAttributeError\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mImportError\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mModuleNotFoundError\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mTypeError\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 185\u001b[0;31m with get_handle(\n\u001b[0m\u001b[1;32m 186\u001b[0m \u001b[0mfilepath_or_buffer\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 187\u001b[0m \u001b[0;34m\"rb\"\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/pandas/io/common.py\u001b[0m in \u001b[0;36mget_handle\u001b[0;34m(path_or_buf, mode, encoding, compression, memory_map, is_text, errors, storage_options)\u001b[0m\n\u001b[1;32m 649\u001b[0m \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 650\u001b[0m \u001b[0;31m# Binary mode\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 651\u001b[0;31m \u001b[0mhandle\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mopen\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mhandle\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mioargs\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mmode\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 652\u001b[0m \u001b[0mhandles\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mappend\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mhandle\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 653\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;31mFileNotFoundError\u001b[0m: [Errno 2] No such file or directory: './matern_calib.pkl'"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"matern_df = pd.read_pickle(\"./matern_calib.pkl\")\n",
|
||||
"volt_df = pd.read_pickle(\"./volt_calib.pkl\")\n",
|
||||
"matern_df.columns = volt_df.columns"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 6,
|
||||
"id": "37d94b48-6e05-4a59-91a3-11291bfc7f4d",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"ename": "NameError",
|
||||
"evalue": "name 'matern_df' 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-6-5b36a26c82e9>\u001b[0m in \u001b[0;36m<module>\u001b[0;34m\u001b[0m\n\u001b[0;32m----> 1\u001b[0;31m \u001b[0mmatern_df\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mMean\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0munique\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m",
|
||||
"\u001b[0;31mNameError\u001b[0m: name 'matern_df' is not defined"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"matern_df.Mean.unique()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 7,
|
||||
"id": "1f59e48e-eb80-48b2-8ce5-91596ad136f2",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def ECDF(sample_pxs, true_px): \n",
|
||||
" return (torch.sum(sample_pxs < true_px, 0)/sample_pxs.shape[0])\n",
|
||||
" \n",
|
||||
"def Calibration(pcts, percentile=0.95):\n",
|
||||
" in_band = np.where((pcts < percentile))[0].shape[0]\n",
|
||||
" return in_band/pcts.shape[0]\n",
|
||||
"\n",
|
||||
"def GetCalibration(model, horizon=np.arange(75,100), \n",
|
||||
" logger=[], exp=True):\n",
|
||||
" \n",
|
||||
"\n",
|
||||
" ntrain = 400\n",
|
||||
" n_test_times = 20\n",
|
||||
" ntest = 100\n",
|
||||
" pcts = torch.tensor([])\n",
|
||||
" for tckr in ticker_list:\n",
|
||||
" data = GetStockHistory(tckr, history=1000, end_date=end_date)\n",
|
||||
" for idx, date in enumerate(data.index):\n",
|
||||
"\n",
|
||||
" fpath = \"./saved-outputs/\"+ tckr + \"/\"\n",
|
||||
" fname = model + \"_\"\n",
|
||||
" if model == 'volt':\n",
|
||||
" fname += \"constant\"\n",
|
||||
"\n",
|
||||
" fname += str(date.date()) + \".pt\"\n",
|
||||
"\n",
|
||||
"\n",
|
||||
" if os.path.exists(fpath + fname): \n",
|
||||
" preds = torch.load(fpath + fname)\n",
|
||||
" if isinstance(preds, tuple):\n",
|
||||
" preds = preds[0]\n",
|
||||
" \n",
|
||||
" if preds.shape[-1] == 100:\n",
|
||||
" preds = preds[:, horizon]\n",
|
||||
"\n",
|
||||
" test_y = torch.tensor(data.iloc[idx:idx+100].Close.to_numpy())\n",
|
||||
" if test_y.shape[0] == 100:\n",
|
||||
" if exp:\n",
|
||||
" preds = preds.exp()\n",
|
||||
" pcts = torch.cat((pcts, ECDF(preds, test_y[horizon])))\n",
|
||||
" \n",
|
||||
" if pcts.numel() == 0:\n",
|
||||
" return logger\n",
|
||||
" \n",
|
||||
" pcts = pcts.flatten().numpy()\n",
|
||||
" percentiles = np.linspace(0.05, 0.95, 19)\n",
|
||||
" for pct in percentiles:\n",
|
||||
" clb = Calibration(pcts, pct)\n",
|
||||
" logger.append([clb, np.round(pct, 2), model, \"Constant\", 100])\n",
|
||||
" \n",
|
||||
" return logger"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 8,
|
||||
"id": "5fe6b4d6-b49e-4e06-8874-9be43b4ee4bd",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"data_path = \"../../voltron/data/\"\n",
|
||||
"ticker_list = make_ticker_list(data_path + \"test_tickers.txt\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 9,
|
||||
"id": "da419f67-e0d9-4008-b038-96fd341e10f3",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"end_date = \"2022-01-13\""
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 10,
|
||||
"id": "d9ed9129-be2d-47e8-9866-5a708c27d225",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"log = []\n",
|
||||
"for model in ['volt', 'lstm']:\n",
|
||||
" log = GetCalibration(model, horizon=np.arange(75,100), \n",
|
||||
" logger=log, exp=True)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 11,
|
||||
"id": "89054cf9-fb1c-45a1-aa8b-8884713f4c07",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"ename": "ValueError",
|
||||
"evalue": "Length mismatch: Expected axis has 0 elements, new values have 5 elements",
|
||||
"output_type": "error",
|
||||
"traceback": [
|
||||
"\u001b[0;31m---------------------------------------------------------------------------\u001b[0m",
|
||||
"\u001b[0;31mValueError\u001b[0m Traceback (most recent call last)",
|
||||
"\u001b[0;32m<ipython-input-11-e23ecfc68e6a>\u001b[0m in \u001b[0;36m<module>\u001b[0;34m\u001b[0m\n\u001b[1;32m 1\u001b[0m \u001b[0mdf\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mpd\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mDataFrame\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mlog\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m----> 2\u001b[0;31m \u001b[0mdf\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mcolumns\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0;34m[\u001b[0m\u001b[0;34m'Calibration'\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m'Percentile'\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m\"Model\"\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m\"Mean\"\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m\"k\"\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/pandas/core/generic.py\u001b[0m in \u001b[0;36m__setattr__\u001b[0;34m(self, name, value)\u001b[0m\n\u001b[1;32m 5476\u001b[0m \u001b[0;32mtry\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 5477\u001b[0m \u001b[0mobject\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m__getattribute__\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mname\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 5478\u001b[0;31m \u001b[0;32mreturn\u001b[0m \u001b[0mobject\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m__setattr__\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mname\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mvalue\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 5479\u001b[0m \u001b[0;32mexcept\u001b[0m \u001b[0mAttributeError\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 5480\u001b[0m \u001b[0;32mpass\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32mpandas/_libs/properties.pyx\u001b[0m in \u001b[0;36mpandas._libs.properties.AxisProperty.__set__\u001b[0;34m()\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/pandas/core/generic.py\u001b[0m in \u001b[0;36m_set_axis\u001b[0;34m(self, axis, labels)\u001b[0m\n\u001b[1;32m 668\u001b[0m \u001b[0;32mdef\u001b[0m \u001b[0m_set_axis\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0maxis\u001b[0m\u001b[0;34m:\u001b[0m \u001b[0mint\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mlabels\u001b[0m\u001b[0;34m:\u001b[0m \u001b[0mIndex\u001b[0m\u001b[0;34m)\u001b[0m \u001b[0;34m->\u001b[0m \u001b[0;32mNone\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 669\u001b[0m \u001b[0mlabels\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mensure_index\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mlabels\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 670\u001b[0;31m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_mgr\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mset_axis\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0maxis\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mlabels\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 671\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_clear_item_cache\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 672\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/pandas/core/internals/managers.py\u001b[0m in \u001b[0;36mset_axis\u001b[0;34m(self, axis, new_labels)\u001b[0m\n\u001b[1;32m 218\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 219\u001b[0m \u001b[0;32mif\u001b[0m \u001b[0mnew_len\u001b[0m \u001b[0;34m!=\u001b[0m \u001b[0mold_len\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 220\u001b[0;31m raise ValueError(\n\u001b[0m\u001b[1;32m 221\u001b[0m \u001b[0;34mf\"Length mismatch: Expected axis has {old_len} elements, new \"\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 222\u001b[0m \u001b[0;34mf\"values have {new_len} elements\"\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;31mValueError\u001b[0m: Length mismatch: Expected axis has 0 elements, new values have 5 elements"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"df = pd.DataFrame(log)\n",
|
||||
"df.columns = ['Calibration', 'Percentile', \"Model\", \"Mean\", \"k\"]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 15,
|
||||
"id": "185272ab-5724-4dc8-add6-42f054832940",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df = pd.concat([df, matern_df, volt_df])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 16,
|
||||
"id": "81e78fa5-dfcb-4c71-9b09-094be0e58b36",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"mat_df = df[(df['Model'] == 'matern') & (df['Mean'] == 'tewma') & (df['k'] == 400)]\n",
|
||||
"lstm_df = df[df['Model']=='lstm']\n",
|
||||
"volt_df = df[(df['Model'] == 'volt') & (df['Mean'] == 'ewma') & (df['k']==100)]\n",
|
||||
"plt_df = pd.concat([lstm_df, mat_df, volt_df])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 27,
|
||||
"id": "3b8bf1a9-42f3-4a5a-aafb-3385cd32ea9e",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"mat_df = df[(df['Model'] == 'matern') & (df['Mean'] == 'constant')]\n",
|
||||
"volt_df = df[(df['Model'] == 'volt') & (df['Mean'] == 'Constant')]\n",
|
||||
"const_df = pd.concat([mat_df, volt_df])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 41,
|
||||
"id": "bf42d9c5-0bc3-46a8-a1ab-14b87fd8851b",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAABPoAAAIRCAYAAADTKdPXAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAABcSAAAXEgFnn9JSAAEAAElEQVR4nOzdd3hTZfsH8O/Jbrp3gQKFliGj7A0qU+EHshRxgDhwgeBCxdetr7gnr4ATxYlQQFBBEGTvVUaBDgp00b2SZp/fH6WxadI2SdOWlu/nurhIznlWmrQ5ufM8zy2IoiiCiIiIiIiIiIiImjRJYw+AiIiIiIiIiIiI6o6BPiIiIiIiIiIiomaAgT4iIiIiIiIiIqJmgIE+IiIiIiIiIiKiZoCBPiIiIiIiIiIiomaAgT4iIiIiIiIiIqJmgIE+IiIiIiIiIiKiZoCBPiIiIiIiIiIiomaAgT4iIiIiIiIiIqJmgIE+IiIiIiIiIiKiZoCBPiIiIiIiIiIiomaAgT4iIiIiIiIiIqJmgIE+IiIiIiIiIiKiZoCBPiIiIiIiIiIiomaAgT4iIiIiIiIiIqJmQNbYAyAiau46depkc3/u3Ll47LHHGmk01460tDSMHDnS5tiiRYswZcqURhoRNUeffvopFi9ebHPs7Nmz9VZvxowZOHDggPV+//79sWLFCidHS/WJzw0RERFdDRjoI2omTCYTkpKSkJKSguLiYhQXF8NiscDLywtqtRoRERFo1aoVIiMjoVAoGnu4RHSNuXz5MlJSUpCeno7i4mLodDqoVCr4+vrC398fbdq0QYcOHSCVSht7qERERERETRYDfURNmMFgwObNm7F69WocPnwYOp2u1jpyuRwdOnRA9+7d0a9fPwwZMgRBQUENMFpqSkaMGIH09HSnyspkMvj4+MDX1xctWrRA165dERsbi+HDh8PLy6ueR0pXK4vFgp07d2LTpk3YtWsXLl++XGsdLy8vdOnSBTfeeCMmTJiAFi1aNMBIiYiIiIiaDwb6iJqov//+G6+99hqysrJcqmc0GnH69GmcPn0av/zyCyQSCe666y688MILtdblsiRyxGQyobCwEIWFhbh06ZL1NeLr64uJEydi7ty5CAwMbORRUkMRRRFr1qzB0qVLceHCBZfqlpWV4fDhwzh8+DA++OADDBgwAHPmzEH//v3rabR0LXvuueewZs0a6/1WrVph69atjTgiIiIiorpjMg6iJkYURbzyyit49NFHXQ7yOWKxWJCRkeGBkRHZKikpwffff4/x48dj+/btjT0cagAXL17E9OnTsXDhQpeDfFWJooh9+/ZhxowZePDBB5GWluahURIRERERNV+c0UfUxLz88sv45ZdfHJ5r2bIlBg4ciJiYGAQFBcHLywtarRbFxcVITU3FqVOncObMGRgMhgYeNTUHnTt3dnjcaDSiuLgYOTk5Ds/n5uZizpw5+PzzzzF48OD6HCI1ot27d2PevHkoLS11eF6hUKB3796IjY1FUFAQAgMDoVQqodFokJGRgcTERBw8eBAFBQV2dbdv344DBw4gMjKyvh8GEREREVGTxkAfUROyZcsWh0G+rl27YsGCBRg4cCAEQaixjbKyMuzcuRObN2/Gli1boNVq62u41MysW7euxvP5+fn4559/8PXXXyMxMdHmnNFoxGOPPYa//voLwcHB9TlMq8jISKeymFLd/fPPP5g7dy6MRqPduejoaMydO9epPRstFgsOHDiAX3/9FRs3boTJZKqvITe6xx57jNm3mxluZUFERERXAy7dJWoiRFHEm2++aXd8zJgx+PnnnzFo0KBag3xA+Wb3Y8aMwbvvvosdO3Zg4cKFaNu2bX0Mma4xQUFBmDJlCuLi4nDbbbfZnS8tLcXixYsbYWRUnxISEvDEE0/YBfnkcjleeuklrF+/HuPGjXMqMYtEIsHAgQPx/vvv448//sDw4cPra9hERERERM0SA31ETcSRI0fssqCGh4dj0aJFUCgUbrXp6+uLWbNm4dlnn/XEEIkAlC/RfO211zBkyBC7c2vWrOHS8WZEr9fjySeftJsZrFarsWzZMtx1112QSqVutd22bVssXboUb7/9NtRqtSeGS0RERETU7DHQR9RE7Nixw+7Y5MmT4ePj0wijIaqZRCLB008/bXe8IqsqNQ9Lly5FSkqK3fGPPvrIYaDXHZMmTcJPP/2EiIgIj7RHRERERNSccY8+oibCUWbcbt26NcJIGsb58+eRnJyMvLw8FBYWwsvLC8HBwYiIiEBsbCzkcnm99V1WVoYTJ04gJycHBQUFKCkpgUqlgq+vL6KiohAdHY3AwMB667+56NKlC1q1amU3E/XUqVMYNGiQ2+2mpqbizJkzyMrKglarhVwuR2hoKCZNmlTHETvHYDDg5MmTyMrKQmFhIYqLi6FQKODj44PWrVsjJiYGoaGhHuvvwoULSEpKQn5+PgoKCqBQKBAQEICIiAj07NkTKpXKY325Ii8vD8uXL7c7fscdd+CGG27waF/VJYKpzeXLl5GSkoK0tDSUlpZCp9PBx8cH/v7+aNmyJbp37w6lUunRsV4tLl26ZH2d6nQ6BAUFITw8HD179kRAQEC991/X39Nr+bmrLCMjAwkJCTa//0FBQQgLC2uw33+j0Yj4+HgkJyejoKAAMpkMQUFBiIqKQmxsrNuzdomIiKh+MNBH1ETk5+fbHXNmz6u66tSpU7XnDhw4UOP5Cn///bdT2TIvX76ML7/8Elu3bkVaWlq15by9vTFo0CDMnDkTAwYMqLVdZ+j1emsCgGPHjjlMKlBBEAR06tQJN9xwA6ZMmYKoqCiPjKE6BoMBzz//PNavX29zPCIiAp9//rlTz0Fj6dChg12gz9FrGbB/rc2dO9earECr1eL777/HypUrcenSJYf1qwYQ0tLSMHLkSJtjixYtwpQpU1x5CAAAs9mM9evXY/369Th8+DDKyspqLB8VFYXrr78ekydPRpcuXVzu79KlS1i+fDm2b99e7eMFAKVSib59++Kee+7xeHCtNitXrrRbsuvr64sFCxY06Dgqy8/Px5YtW7Bnzx4cPHgQubm5NZaXy+Xo2bMn7rrrLtx0002QSBpuocOnn35qt2dlXZPHiKKIuLg4LF++HOfOnXNYRi6XY+DAgXjwwQfRv39/l/vw9O9phYZ67kaMGGH3N6lCenq6U39Pv/vuO4fvPTNmzMCBAwes9/v37+9Wgo7i4mJ8/fXX2Lx5M5KSkqotp1Qq0a9fP0yfPh2jR492uZ+4uDgsXLjQ5ljl9+ucnBx8/vnnWLNmDUpKShy24efnh8mTJ+ORRx7hF2BERERXCQb6iJoIR/vwOZrl1xSZzWYsXrwY33zzTa0BFADQaDTYsmULtmzZghtuuAGvvPIKWrZs6Xb/P/30E/73v/8hJyfHqfKiKOLMmTM4c+YMli1bhk8++QQ33XST2/3XpKioCHPnzrX58AgA1113HZYtW4bw8PB66ddTHC0tr+4DY3WOHz+Oxx9/vNFe75s2bcL777+PCxcuOF0nNTUVqamp+O677/Dcc8/h3nvvdapeaWkpPvjgA6xcubLGYHMFvV6P3bt3Y/fu3ejTpw/ee++9Ov0uuCIuLs7u2KRJk+Dt7d0g/Vf11FNPuZyp12g04uDBgzh48CCio6Px8ccfo0OHDvU4yvqTn5+Pxx57DIcOHaqxnNFoxM6dO7Fr1y5MnToVL774okdmhdXl9/Raf+4qW7FiBT799FMUFRXVWlav12PXrl3YtWsXevXqhVdffdVjX/xs2rQJL7zwAoqLi2ssV1xcjG+//Rbr1q3DsmXL0LNnT4/0T0RERO7jHn1ETYSjpYB//vlnI4zEs8rKyjBnzhx89tlnTgX5qtq+fTtuv/12nDlzxuW6er0eTz/9NF555RWng3yOaDQat+vW5NKlS5g+fbpdkG/YsGH4/vvvr/ogH1AeuKrK19fX6foHDx7EjBkzGiXIZ7FY8M4772DevHkuBfmqcvQzcCQ9PR133HEHfvjhB6eCfFUdPnwYt912G44dO+ZyXVclJSXh4sWLdsdvv/32eu+7OkePHnUpUFRVcnIypk2bhj179nhwVA2jqKgId911V61BvspEUcSqVavw8MMPQ6fT1an/uv6eXsvPXQWz2YyXXnoJb7zxhlNBvqqOHj2KO++8E3v37q3zWH766SfMnz+/1iBfZYWFhbj33nuRkJBQ5/6JiIiobjijj6iJ6NWrF3755RebY3v27MGKFSswY8aMeuu38t5YFy9etFmqp1ar0aZNm1rbqG4/PYvFgkcffdThhzNvb28MHz4csbGxCA0NRWlpKVJTU7Flyxa7oEt2djbuvvturF69Gm3btnXqcRmNRtx///04ePCg3TmJRIKuXbti0KBBaNGiBQICAmAwGFBYWIizZ88iPj6+xuVUnhAfH4+HH34YeXl5NsenTZuGl19+GTJZ0/jznZiYaHcsKCjIqbo5OTmYO3cu9Hq99VhsbCyGDBmCVq1awdvbG9nZ2UhOTsbGjRs9NuYKCxYswIYNGxye69ixIwYPHow2bdogMDAQRqMRRUVFSEpKwsmTJ3H69GmIouh0X+np6Zg2bZrD5YqxsbHo3bs32rVrBz8/PxiNRuTk5ODo0aPYsWOHTRbj3NxcPPTQQ4iLi0OrVq1cf9BO2r9/v92xkJCQq2ZGlVQqRZcuXdChQwe0a9cOgYGB1pmGFX9Ljh8/jiNHjsBisVjrabVaPPHEE1i7di1atGjRWMN32TPPPGOTFKVFixYYPXo0oqOj4efnh9zcXJw8eRJ///23XeB57969eOKJJ7BkyRK3+vb072l9P3fR0dHWLxsyMzNtgmpyuRzR0dG1jrE+skC/+OKLWL16td1xpVKJoUOHol+/fggNDYVOp0N6ejr+/vtvu6XepaWlmD17Nr799lv06dPHrXHs2LEDr7/+uvXvl6+vL4YMGYJevXohODgYFosF6enp+OeLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 900x450 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"from matplotlib.lines import Line2D\n",
|
||||
"fig, ax = plt.subplots(1,1,dpi=150, figsize=(6, 3))\n",
|
||||
"\n",
|
||||
"percentiles = np.linspace(0.05, 0.95, 19)\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"pal = [palette[5], palette[7]]\n",
|
||||
"# pal = [palette[5]]\n",
|
||||
"sns.lineplot(x='Percentile', y=\"Calibration\", hue='Model', data=const_df, ax=ax, alpha=0.2,\n",
|
||||
" palette=pal, legend=True)\n",
|
||||
"sns.scatterplot(x='Percentile', y=\"Calibration\", hue='Model', data=const_df, ax=ax, s=120, legend=False, zorder=4,\n",
|
||||
" palette=pal, alpha=0.35)\n",
|
||||
"\n",
|
||||
"pal = [ palette[0], palette[4], palette[6]]\n",
|
||||
"sns.lineplot(x='Percentile', y=\"Calibration\", hue='Model', data=plt_df, ax=ax, alpha=0.5,\n",
|
||||
" palette=pal, legend=True)\n",
|
||||
"sns.scatterplot(x='Percentile', y=\"Calibration\", hue='Model', data=plt_df, ax=ax, s=120, legend=False, zorder=4,\n",
|
||||
" palette=pal)\n",
|
||||
"x = np.linspace(0.05,0.95)\n",
|
||||
"y = np.linspace(0, len(percentiles))\n",
|
||||
"ax.plot(x, x, color=\"gray\", lw=1., ls=\"--\")\n",
|
||||
"ax.set_title(\"Stock Price Calibration\")\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"plt.tick_params(labelsize=16)\n",
|
||||
"sns.despine()\n",
|
||||
"\n",
|
||||
"custom_lines = [Line2D([0], [0], color=palette[0], lw=2),\n",
|
||||
" Line2D([0], [0], color=palette[4], lw=2),\n",
|
||||
" Line2D([0], [0], color=palette[6], lw=2),\n",
|
||||
" Line2D([0], [0], color=palette[5], lw=2),\n",
|
||||
" Line2D([0], [0], color=palette[7], lw=2)]\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"plt.legend(custom_lines, ['LSTM', r\"Matérn + Magpie\", \"Volt + Magpie\", \"Matérn + Constant\",\n",
|
||||
" \"Volt + Constant\"],\n",
|
||||
" fontsize=14, frameon=False, bbox_to_anchor=(1., 0.85))\n",
|
||||
"# ax.legend(fontsize=14, bbox_to_anchor=(1., 0.75))\n",
|
||||
"# plt.label(\"Percentile\")\n",
|
||||
"plt.savefig(\"./stock_calibration.pdf\", bbox_inches=\"tight\")\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "957ddda1-e78b-4bb7-9003-886f18e424f7",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "d9fa07e1-f9e5-4309-9311-160c53b43d0d",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.8"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,722 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "9754a692-edb8-41bf-8159-fd82edcb9195",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Warning no robinhood utils.\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import pandas as pd\n",
|
||||
"import torch\n",
|
||||
"from torch import nn\n",
|
||||
"import seaborn as sns\n",
|
||||
"import time\n",
|
||||
"import copy\n",
|
||||
"import sys\n",
|
||||
"from torch.utils.data import DataLoader\n",
|
||||
"from voltron.data import make_ticker_list, GetStockHistory\n",
|
||||
"sns.set_style('white')\n",
|
||||
"# style.use('whitegrid')\n",
|
||||
"palette = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"sns.set(palette = palette, font_scale=2.0, style=\"white\", rc={\"lines.linewidth\": 2.0})"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "dbc24c1a-f499-4b8f-8544-6d4a66f48488",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Dataset Setup"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 29,
|
||||
"id": "62f91be1-38aa-4781-8b65-c5caf9752045",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from torch.utils.data import Dataset\n",
|
||||
"\n",
|
||||
"class SequenceDataset(Dataset):\n",
|
||||
" def __init__(self, data, sequence_length=5):\n",
|
||||
" self.sequence_length = sequence_length\n",
|
||||
" self.X = data.float()\n",
|
||||
"\n",
|
||||
" def __len__(self):\n",
|
||||
" return self.X.shape[0]-1\n",
|
||||
"\n",
|
||||
" def __getitem__(self, i): \n",
|
||||
" if i >= self.sequence_length - 1:\n",
|
||||
" i_start = i - self.sequence_length + 1\n",
|
||||
" x = self.X[i_start:(i + 1)]\n",
|
||||
" else:\n",
|
||||
" padding = self.X[0].repeat(self.sequence_length - i - 1, 1).squeeze(-1)\n",
|
||||
" x = self.X[0:(i + 1)]\n",
|
||||
" x = torch.cat((padding, x), 0)\n",
|
||||
" \n",
|
||||
" return x.unsqueeze(0), self.X[i+1]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 32,
|
||||
"id": "566c9386-c297-430f-b4ef-a0f3ba7aec37",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"tckr = \"JPM\"\n",
|
||||
"ntrain = 400\n",
|
||||
"lookback = 1\n",
|
||||
"pxs = GetStockHistory(tckr, end_date=\"2021-12-07\", history=ntrain + lookback).Close.to_numpy()\n",
|
||||
"data = torch.FloatTensor(pxs).log()\n",
|
||||
"\n",
|
||||
"data = (data - data.mean())/data.std()\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"# xin = torch.linspace(0, 5*np.pi, 250)\n",
|
||||
"# data = torch.sin(xin) + 0.2 * torch.randn(xin.shape)\n",
|
||||
"\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 33,
|
||||
"id": "5dece53d-785f-4ca9-9cbf-03e16de71e9b",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"seq_len = 25\n",
|
||||
"dset = SequenceDataset(data, seq_len)\n",
|
||||
"\n",
|
||||
"trgts = []\n",
|
||||
"for i in range(len(dset)):\n",
|
||||
" trgts.append(dset[i][1].item())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 34,
|
||||
"id": "2317ba11-4f62-443f-9ad8-63f9a0996ee9",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"train_loader = DataLoader(dset, batch_size=20, shuffle=True)\n",
|
||||
"X, y = next(iter(train_loader))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 35,
|
||||
"id": "ffecfde0-6938-4dee-b1b0-6627012e3329",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"class ShallowRegressionLSTM(nn.Module):\n",
|
||||
" def __init__(self, input_size, hidden_units=128):\n",
|
||||
" super().__init__()\n",
|
||||
" self.input_size = input_size # this is the number of features\n",
|
||||
" self.hidden_units = hidden_units\n",
|
||||
" self.num_layers = 5\n",
|
||||
"\n",
|
||||
" self.lstm = nn.LSTM(\n",
|
||||
" input_size=input_size,\n",
|
||||
" hidden_size=hidden_units,\n",
|
||||
" batch_first=True,\n",
|
||||
" num_layers=self.num_layers\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
" self.linear = nn.Linear(in_features=self.hidden_units, out_features=2)\n",
|
||||
" \n",
|
||||
"\n",
|
||||
" def forward(self, x):\n",
|
||||
" batch_size = x.shape[0]\n",
|
||||
" h0 = torch.zeros(self.num_layers, batch_size, self.hidden_units).requires_grad_()\n",
|
||||
" c0 = torch.zeros(self.num_layers, batch_size, self.hidden_units).requires_grad_()\n",
|
||||
"\n",
|
||||
" _, (hn, _) = self.lstm(x, (h0, c0))\n",
|
||||
" out = self.linear(hn[0]) # First dim of Hn is num_layers, which is set to 1 above.\n",
|
||||
" output = torch.zeros_like(out)\n",
|
||||
" output[:, 0] = out[:, 0]\n",
|
||||
" output[:, 1] = torch.exp(out[:, 1])\n",
|
||||
" return output\n",
|
||||
" \n",
|
||||
"from torch.autograd import Variable \n",
|
||||
"class LSTM1(nn.Module):\n",
|
||||
" def __init__(self, num_classes, seq_len, hidden_size, num_layers):\n",
|
||||
" super(LSTM1, self).__init__()\n",
|
||||
" self.num_classes = num_classes #number of classes\n",
|
||||
" self.num_layers = num_layers #number of layers\n",
|
||||
" self.input_size = seq_len #input size\n",
|
||||
" self.hidden_size = hidden_size #hidden state\n",
|
||||
"\n",
|
||||
" self.lstm = nn.LSTM(input_size=seq_len, hidden_size=hidden_size,\n",
|
||||
" num_layers=num_layers, batch_first=True) #lstm\n",
|
||||
" self.fc_1 = nn.Linear(hidden_size, 128) #fully connected 1\n",
|
||||
" self.fc = nn.Linear(128, num_classes) #fully connected last layer\n",
|
||||
"\n",
|
||||
" self.relu = nn.ReLU()\n",
|
||||
" self.softplus = nn.Softplus()\n",
|
||||
" \n",
|
||||
" def forward(self,x):\n",
|
||||
" h_0 = Variable(torch.zeros(self.num_layers, x.size(0), self.hidden_size)) #hidden state\n",
|
||||
" c_0 = Variable(torch.zeros(self.num_layers, x.size(0), self.hidden_size)) #internal state\n",
|
||||
" # Propagate input through LSTM\n",
|
||||
" output, (hn, cn) = self.lstm(x, (h_0, c_0)) #lstm with input, hidden, and internal state\n",
|
||||
"\n",
|
||||
" hn = hn[self.num_layers-1]\n",
|
||||
" hn = hn.view(-1, self.hidden_size) #reshaping the data for Dense layer next\n",
|
||||
" out = self.relu(hn)\n",
|
||||
" out = self.fc_1(out) #first Dense\n",
|
||||
" out = self.relu(out) #relu\n",
|
||||
" out = self.fc(out) #Final Output\n",
|
||||
" \n",
|
||||
" output = torch.zeros_like(out)\n",
|
||||
" output[:, 0] = out[:, 0]\n",
|
||||
" output[:, 1] = self.softplus(out[:, 1])\n",
|
||||
" return output"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 36,
|
||||
"id": "3bc83b15-3b0f-4f21-803b-29e354cb558f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# model = ShallowRegressionLSTM(seq_len)\n",
|
||||
"model = LSTM1(2, seq_len, 128, 1)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 37,
|
||||
"id": "41547213-7851-4ad1-ac04-03cab5e0a91c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def NLL(targets, outputs):\n",
|
||||
" dist = torch.distributions.Normal(outputs[:, 0], outputs[:, 1])\n",
|
||||
" return -dist.log_prob(targets).sum()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 38,
|
||||
"id": "e6958f51-5cd0-414a-b419-3ce55f12b486",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def train_model(data_loader, model, loss_function, optimizer, epochs=200):\n",
|
||||
" num_batches = len(data_loader)\n",
|
||||
" total_loss = 0\n",
|
||||
" model.train()\n",
|
||||
" for epoch in range(epochs):\n",
|
||||
" for X, y in data_loader:\n",
|
||||
" output = model(X)\n",
|
||||
" loss = loss_function(y, output)\n",
|
||||
"\n",
|
||||
" optimizer.zero_grad()\n",
|
||||
" loss.backward()\n",
|
||||
" optimizer.step()\n",
|
||||
"\n",
|
||||
" total_loss += loss.item()\n",
|
||||
"\n",
|
||||
" if epoch%10 == 0:\n",
|
||||
" avg_loss = total_loss / num_batches\n",
|
||||
" print(f\"Train loss: {avg_loss}, Epoch: {epoch}\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 39,
|
||||
"id": "92484e45-11c9-4fe1-bc5b-87b0f27d5a2c",
|
||||
"metadata": {
|
||||
"tags": []
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Train loss: 9.58264362514019, Epoch: 0\n",
|
||||
"Train loss: -7.665388387441635, Epoch: 10\n",
|
||||
"Train loss: -96.57648058235645, Epoch: 20\n",
|
||||
"Train loss: -222.45588295161724, Epoch: 30\n",
|
||||
"Train loss: -393.4549728780985, Epoch: 40\n",
|
||||
"Train loss: -557.4210761398077, Epoch: 50\n",
|
||||
"Train loss: -733.938936367631, Epoch: 60\n",
|
||||
"Train loss: -916.8391911536455, Epoch: 70\n",
|
||||
"Train loss: -1094.614848962426, Epoch: 80\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"ERROR:root:Internal Python error in the inspect module.\n",
|
||||
"Below is the traceback from this internal error.\n",
|
||||
"\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Traceback (most recent call last):\n",
|
||||
" File \"/home/greg_b/miniconda3/lib/python3.8/site-packages/IPython/core/interactiveshell.py\", line 3444, in run_code\n",
|
||||
" exec(code_obj, self.user_global_ns, self.user_ns)\n",
|
||||
" File \"<ipython-input-39-1bf22962a03c>\", line 2, in <module>\n",
|
||||
" train_model(train_loader, model, NLL, optimizer)\n",
|
||||
" File \"<ipython-input-38-1b396ffa6e98>\", line 11, in train_model\n",
|
||||
" loss.backward()\n",
|
||||
" File \"/home/greg_b/miniconda3/lib/python3.8/site-packages/torch/_tensor.py\", line 307, in backward\n",
|
||||
" torch.autograd.backward(self, gradient, retain_graph, create_graph, inputs=inputs)\n",
|
||||
" File \"/home/greg_b/miniconda3/lib/python3.8/site-packages/torch/autograd/__init__.py\", line 154, in backward\n",
|
||||
" Variable._execution_engine.run_backward(\n",
|
||||
"KeyboardInterrupt\n",
|
||||
"\n",
|
||||
"During handling of the above exception, another exception occurred:\n",
|
||||
"\n",
|
||||
"Traceback (most recent call last):\n",
|
||||
" File \"/home/greg_b/miniconda3/lib/python3.8/site-packages/IPython/core/interactiveshell.py\", line 2064, in showtraceback\n",
|
||||
" stb = value._render_traceback_()\n",
|
||||
"AttributeError: 'KeyboardInterrupt' object has no attribute '_render_traceback_'\n",
|
||||
"\n",
|
||||
"During handling of the above exception, another exception occurred:\n",
|
||||
"\n",
|
||||
"Traceback (most recent call last):\n",
|
||||
" File \"/home/greg_b/miniconda3/lib/python3.8/site-packages/IPython/core/ultratb.py\", line 1101, in get_records\n",
|
||||
" return _fixed_getinnerframes(etb, number_of_lines_of_context, tb_offset)\n",
|
||||
" File \"/home/greg_b/miniconda3/lib/python3.8/site-packages/IPython/core/ultratb.py\", line 248, in wrapped\n",
|
||||
" return f(*args, **kwargs)\n",
|
||||
" File \"/home/greg_b/miniconda3/lib/python3.8/site-packages/IPython/core/ultratb.py\", line 281, in _fixed_getinnerframes\n",
|
||||
" records = fix_frame_records_filenames(inspect.getinnerframes(etb, context))\n",
|
||||
" File \"/home/greg_b/miniconda3/lib/python3.8/inspect.py\", line 1503, in getinnerframes\n",
|
||||
" frameinfo = (tb.tb_frame,) + getframeinfo(tb, context)\n",
|
||||
" File \"/home/greg_b/miniconda3/lib/python3.8/inspect.py\", line 1461, in getframeinfo\n",
|
||||
" filename = getsourcefile(frame) or getfile(frame)\n",
|
||||
" File \"/home/greg_b/miniconda3/lib/python3.8/inspect.py\", line 708, in getsourcefile\n",
|
||||
" if getattr(getmodule(object, filename), '__loader__', None) is not None:\n",
|
||||
" File \"/home/greg_b/miniconda3/lib/python3.8/inspect.py\", line 744, in getmodule\n",
|
||||
" for modname, module in sys.modules.copy().items():\n",
|
||||
"KeyboardInterrupt\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"ename": "TypeError",
|
||||
"evalue": "object of type 'NoneType' has no len()",
|
||||
"output_type": "error",
|
||||
"traceback": [
|
||||
"\u001b[0;31m---------------------------------------------------------------------------\u001b[0m",
|
||||
"\u001b[0;31mKeyboardInterrupt\u001b[0m Traceback (most recent call last)",
|
||||
" \u001b[0;31m[... skipping hidden 1 frame]\u001b[0m\n",
|
||||
"\u001b[0;32m<ipython-input-39-1bf22962a03c>\u001b[0m in \u001b[0;36m<module>\u001b[0;34m\u001b[0m\n\u001b[1;32m 1\u001b[0m \u001b[0moptimizer\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mtorch\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0moptim\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mAdam\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mmodel\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mparameters\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mlr\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0;36m0.01\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m----> 2\u001b[0;31m \u001b[0mtrain_model\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mtrain_loader\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mmodel\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mNLL\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0moptimizer\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m",
|
||||
"\u001b[0;32m<ipython-input-38-1b396ffa6e98>\u001b[0m in \u001b[0;36mtrain_model\u001b[0;34m(data_loader, model, loss_function, optimizer, epochs)\u001b[0m\n\u001b[1;32m 10\u001b[0m \u001b[0moptimizer\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mzero_grad\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 11\u001b[0;31m \u001b[0mloss\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mbackward\u001b[0m\u001b[0;34m(\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 12\u001b[0m \u001b[0moptimizer\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mstep\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/lib/python3.8/site-packages/torch/_tensor.py\u001b[0m in \u001b[0;36mbackward\u001b[0;34m(self, gradient, retain_graph, create_graph, inputs)\u001b[0m\n\u001b[1;32m 306\u001b[0m inputs=inputs)\n\u001b[0;32m--> 307\u001b[0;31m \u001b[0mtorch\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mautograd\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mbackward\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mgradient\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mretain_graph\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mcreate_graph\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0minputs\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0minputs\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 308\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/lib/python3.8/site-packages/torch/autograd/__init__.py\u001b[0m in \u001b[0;36mbackward\u001b[0;34m(tensors, grad_tensors, retain_graph, create_graph, grad_variables, inputs)\u001b[0m\n\u001b[1;32m 153\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 154\u001b[0;31m Variable._execution_engine.run_backward(\n\u001b[0m\u001b[1;32m 155\u001b[0m \u001b[0mtensors\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mgrad_tensors_\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mretain_graph\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mcreate_graph\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0minputs\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;31mKeyboardInterrupt\u001b[0m: ",
|
||||
"\nDuring handling of the above exception, another exception occurred:\n",
|
||||
"\u001b[0;31mAttributeError\u001b[0m Traceback (most recent call last)",
|
||||
"\u001b[0;32m~/miniconda3/lib/python3.8/site-packages/IPython/core/interactiveshell.py\u001b[0m in \u001b[0;36mshowtraceback\u001b[0;34m(self, exc_tuple, filename, tb_offset, exception_only, running_compiled_code)\u001b[0m\n\u001b[1;32m 2063\u001b[0m \u001b[0;31m# in the engines. This should return a list of strings.\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 2064\u001b[0;31m \u001b[0mstb\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mvalue\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_render_traceback_\u001b[0m\u001b[0;34m(\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 2065\u001b[0m \u001b[0;32mexcept\u001b[0m \u001b[0mException\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;31mAttributeError\u001b[0m: 'KeyboardInterrupt' object has no attribute '_render_traceback_'",
|
||||
"\nDuring handling of the above exception, another exception occurred:\n",
|
||||
"\u001b[0;31mTypeError\u001b[0m Traceback (most recent call last)",
|
||||
" \u001b[0;31m[... skipping hidden 1 frame]\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/lib/python3.8/site-packages/IPython/core/interactiveshell.py\u001b[0m in \u001b[0;36mshowtraceback\u001b[0;34m(self, exc_tuple, filename, tb_offset, exception_only, running_compiled_code)\u001b[0m\n\u001b[1;32m 2064\u001b[0m \u001b[0mstb\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mvalue\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_render_traceback_\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 2065\u001b[0m \u001b[0;32mexcept\u001b[0m \u001b[0mException\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 2066\u001b[0;31m stb = self.InteractiveTB.structured_traceback(etype,\n\u001b[0m\u001b[1;32m 2067\u001b[0m value, tb, tb_offset=tb_offset)\n\u001b[1;32m 2068\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/lib/python3.8/site-packages/IPython/core/ultratb.py\u001b[0m in \u001b[0;36mstructured_traceback\u001b[0;34m(self, etype, value, tb, tb_offset, number_of_lines_of_context)\u001b[0m\n\u001b[1;32m 1365\u001b[0m \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 1366\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mtb\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mtb\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1367\u001b[0;31m return FormattedTB.structured_traceback(\n\u001b[0m\u001b[1;32m 1368\u001b[0m self, etype, value, tb, tb_offset, number_of_lines_of_context)\n\u001b[1;32m 1369\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/lib/python3.8/site-packages/IPython/core/ultratb.py\u001b[0m in \u001b[0;36mstructured_traceback\u001b[0;34m(self, etype, value, tb, tb_offset, number_of_lines_of_context)\u001b[0m\n\u001b[1;32m 1265\u001b[0m \u001b[0;32mif\u001b[0m \u001b[0mmode\u001b[0m \u001b[0;32min\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mverbose_modes\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 1266\u001b[0m \u001b[0;31m# Verbose modes need a full traceback\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1267\u001b[0;31m return VerboseTB.structured_traceback(\n\u001b[0m\u001b[1;32m 1268\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0metype\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mvalue\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mtb\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mtb_offset\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mnumber_of_lines_of_context\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 1269\u001b[0m )\n",
|
||||
"\u001b[0;32m~/miniconda3/lib/python3.8/site-packages/IPython/core/ultratb.py\u001b[0m in \u001b[0;36mstructured_traceback\u001b[0;34m(self, etype, evalue, etb, tb_offset, number_of_lines_of_context)\u001b[0m\n\u001b[1;32m 1122\u001b[0m \u001b[0;34m\"\"\"Return a nice text document describing the traceback.\"\"\"\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 1123\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1124\u001b[0;31m formatted_exception = self.format_exception_as_a_whole(etype, evalue, etb, number_of_lines_of_context,\n\u001b[0m\u001b[1;32m 1125\u001b[0m tb_offset)\n\u001b[1;32m 1126\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/lib/python3.8/site-packages/IPython/core/ultratb.py\u001b[0m in \u001b[0;36mformat_exception_as_a_whole\u001b[0;34m(self, etype, evalue, etb, number_of_lines_of_context, tb_offset)\u001b[0m\n\u001b[1;32m 1080\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 1081\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1082\u001b[0;31m \u001b[0mlast_unique\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mrecursion_repeat\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mfind_recursion\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0morig_etype\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mevalue\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mrecords\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 1083\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 1084\u001b[0m \u001b[0mframes\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mformat_records\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mrecords\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mlast_unique\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mrecursion_repeat\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/lib/python3.8/site-packages/IPython/core/ultratb.py\u001b[0m in \u001b[0;36mfind_recursion\u001b[0;34m(etype, value, records)\u001b[0m\n\u001b[1;32m 380\u001b[0m \u001b[0;31m# first frame (from in to out) that looks different.\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 381\u001b[0m \u001b[0;32mif\u001b[0m \u001b[0;32mnot\u001b[0m \u001b[0mis_recursion_error\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0metype\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mvalue\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mrecords\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 382\u001b[0;31m \u001b[0;32mreturn\u001b[0m \u001b[0mlen\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mrecords\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;36m0\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m 383\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 384\u001b[0m \u001b[0;31m# Select filename, lineno, func_name to track frames with\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;31mTypeError\u001b[0m: object of type 'NoneType' has no len()"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"optimizer = torch.optim.Adam(model.parameters(), lr=0.01)\n",
|
||||
"train_model(train_loader, model, NLL, optimizer)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 40,
|
||||
"id": "07eab392-37e2-4dae-bfa4-dd177e3d552f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"means = []\n",
|
||||
"vrs = []\n",
|
||||
"for X, y in dset:\n",
|
||||
" output = model(X.unsqueeze(0))\n",
|
||||
" means = means + list(output[:, 0].detach().numpy())\n",
|
||||
" vrs = vrs + list(output[:, 1].detach().numpy())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 41,
|
||||
"id": "b10866df-b554-44a6-bf73-b98abb2e6400",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"<matplotlib.collections.PathCollection at 0x7f0b7c18d910>"
|
||||
]
|
||||
},
|
||||
"execution_count": 41,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYwAAAEFCAYAAADwhtBaAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAABpqklEQVR4nO39d5xdV33vjb/XbqdNL+rNkiXbkgyWMWAbgxNicAmQELghYAIJuUDim4QbQgKEPAnce/klEG4S8vBwUzGEnlBjwPY1TeDIYMu2bEtGki1ZdSxp+plTd1u/P3aZU2fO9JFnvV8vXmhO2/ssz6zP+nYhpZQoFAqFQjEN2lLfgEKhUCguDpRgKBQKhaIllGAoFAqFoiWUYCgUCoWiJZRgKBQKhaIljKW+gblSKpU4ePAg/f396Lq+1LejUCgUFwWe5zE4OMju3btJJpMtveeiF4yDBw9y++23L/VtKBQKxUXJ5z//ea655pqWXnvRC0Z/fz8QfOk1a9Ys8d0oFArFxcG5c+e4/fbb4z20FS56wYjcUGvWrGHDhg1LfDcKhUJxcTETV74KeisUCoWiJZRgKBQKhaIllGAoFAqFoiWUYCgUCoWiJZRgKBQKhaIllGAoFAqFoiWUYCgUCsUC40uJ6138o4cu+joMhUKhWO4M5RzGiy49aYOejIkQYqlvaVYowVAoFIo5UnZ9fClJmdVFcFJKBsbLFMo+liEYybtYhobrSzQBng8ZSydhXhzOHiUYCoVCMQeKtsfAeBkpYVNPEsuY3PyLjk/B9kkYAiEEuiY5l7URSCSBlTFWcNjUm8LQlr/VcXHImkKhWHZIKcmVPXx58fvmZ0vB9jg7VkYLXUzZkhs/J6VktOCiC2IXlKlrWLogYegkDY2koSGB0yMlBidsHM8nV/Yo2h6ev/zWVVkYCoViRkgpKdg+2ZLLRMkjbWms70pctH752WK7PuezNromMDSBJmC86NKdNtE1EVoXHgm9el20mnWydA3Pl4wXXcaLLhIQgK4J1nclYovF8XykpMqCWWyUYCgUipbxfMlI3mGs4KAJQdIQFG2fsXCjfK5TsD3Kro8mYCTvIpGYWrCBa0LgS0mu7JI0dQbGyxiaaElIdU2g17ikbM/n/ITNhlCML0zYFG2fzb1JTH1pREMJhkKhqEJKyUTJpS1hoFVsYlJKzoyWcDxJwtDijdDSYTTv0Jky6k7PzyU8X/LseJnIUySARM1p39AE2aKHE6bQziUuYWqCkhNYcu0Jg6LjIwliJmZKCYZCoVgG5Msez2ZtetOSvnaL4ZzNWNElZWrYniRZs0lqmsBxJfmyR3tyeW8pUkocT87KrVNyfHwJSUPDl1HIuhpdQMn1Kbs+pj438RRCYOlBSq4mBMjg83Nln47UnD561qigt0KhiJFSMlJwsDSN0aLLcM5mpOBihKdds8mJWdcEo3kHdxkGaivJlT1OjZQoOT5Syvj/W2Gi5BJ9fU00djUJEcQzgpjG3K0tTQikhAsTNpoWWCwF28Px/Dl/9qzuZ0muqlAolg3ZohO6mnwcT2K7El0DSxeMFFx0Ldi4TF2r87NH6AJsX3J2tNTyBjzfRALQrKLa9SWDORuAgfESgzmH0yMlzmftaTdg1w8ywpoJZiVGg3jEXLB0gQBMLXADCmC04E73tgVheduPCoViQQk2UQdPSkbzLqYRbHTRxpQ0Wtv4hBAkdEHZ9Sk4Phmr9Slu80Wu7HEua2NogjUdFqmaexjJO/h+EHdwfUm26GIZgrztURjx6EwFVdjR96mkaHsNH18MhBBV7i1TF2SLLm0JnaLtLWrluLIwFIoVTNH2kBKSusZ4KUjrnEugVhOQLS7+6VfKIHvL0AJf/5mxMvny5H34YSA/2ngNTZAwNDQhsHQNXQhG8g5nx8qM5B2klHi+jK2lkus3jFksBUIEKbznsjZDOYd82Vu0ayvBUCiWIW6UkbPAMYFc2UMLC8sMIbBdyVy8KZGPfaHvu5aCE7jTDE1g6AJTE5yfsOPit0LZw5f1NRARmha43EqOz2jBZWC8zPGhIiP5QHRKtl+VMbbUGFqQwmtogYW0WCjBUCiWEYWyy3jRIV9yGS+4jJemP637UjKUszk3Xp5RdbDnSwq2F1sUhi5Imdqc3BsiDNIOh6f0qa5tu9VxA7fiRD8Tio7H+fEyesV965rA94NgseP5XMjZ08Yf9NDqMDRBOQzwT5RdpAzudY5JT/NK4ALUWGwPmYphKBTLhKLjMTBux5W+liEYL7h0pYwpN/HhnMNY0QUJaculI9VaAV2u5CLl/PvlLV0wVnTpSBkkGsRApJQMjJUpuT7ruxKkLR3Pl5waKdGR1OlMGS0XvEHgApOSOP5SeR/5shdbUXqLO76uCXREIBRekC4sWZr4xXJDCYZCsQwYyTuM5B10IeJTo64FQeSgbkBQsF18n7CJXdCaQtc0xosuCV3gySB7pj3ZXGACX39QoZydY7yiGUHAXFK0vbrCtlzJZSTvYHsSXQQFfylTI1d28fyg99JowWVNh9VSTUfUpqTR9xBCkDCCjX82m330PS7knEU/yS9XlGAoFEuI70vGig7DeRdLb5y7X7A9TF0wlHMoOX6wkQlC942Im9vpBCfiXNmj6PikTY22mk235PqMFBwEIATzmv5ZiS6C031XRbuQsaLD4ISDIQLXj5SSguMzOOFQcr24fsH1gr5KjQTDD1NnU2E7cNcPgtNTWQ9zsQxMXeB4ctZFeMOHjjGwdz92No/VkWHdjdfQu2vbrO9nqVGCoVAsAr6UlB0/TvXMFh3yZR9dC9w3zcTC0ARjYeFc2QmqrCs3QM+fDFIHQhK1z4aJEqQsPRYFKSXjBReNhW9gp2tQDGdEaCKwlIYmnKrvGfjhYbzkBi64cFPWNeJ6Ctv1GCm4ZCwdISBve+TLHp0pk3KYuSRYOHeREAKrxdTiSoYPHeP0fQ/glez4MTub58Rdexl6/CiXvfHW+bzNRUMJhkKxCBRtn4HxMhu6EyQMjeG8Gweok0Zzf72uCUquz7NZG1Ovf12thWBVNKUruz5FxydjBdk/IwUnmM2wCNFbEZhAlJ0gu+jceBlN1GcpCSHQhZx8T/z/MhCZnIPrBVZF8FzQmiNXciG4RCw0i02tKOipBBtvuhaAk3ffj3QbZy/lTj7Lw3/5qfjn6H0Xg+WhBEOhWATKbnBizhbduHCstidTM5Kh+2amp2hNwMBYie60STY8xScaiM5C8mzWxvclQlSLWSXNOq9mS24Qv2lwz7M59c8nw4eOceJbPwoUK8Qrljn5nR+jWWaVWJQ8nx+ez7GjPcHW9kTdZ0XvA5a9aCjBUCgWgaLtY+lBK/Bc2cOaYexgNpt80MpDMlpw0DTRdMNeKExdxJXVM8XQBNmSizmDbKnF5PR3f1IlFhHS8/GK5arHvnM2y/fP5wB48yXdXNuXafi+E9/6EbC8RUMJhkKxwEgpKblB4NoO+xxpi+RG0cIA81JsuZoQaLPsEKJrgpSYW03IQjF86FidKOQcj5ShVdWCAPzoQi4WC4AvnRjlkjaLCcen09QxBHQnwm1YSk7ctZfcmfNsvvn6qustl8C5EgyFYoEpuX5FvYNc9NnNF+uMiuUiFpUbtp60qgLZAF88Mcp/DuZZmzJ4/67V8XqXPJ+vnhoDYE93CtuXHBov8T+fOF/1/tdv6uTnVrfHPw89epjSyDiXvfHWOteXnc0vqSWiBEOhWGDyZY/oiD8b98xyYzmdeBea4UPHqgLYtWLx0HCB/xzMA/Bs0eV03mFzmwXA0WwZT0J/wuA3t/XwyEiRQ+Olumt862yWF/dlSFW4DHMnn+WRv/p0EIiqdX1Jyen7HlCCoVA815BSkit5GMvktDwXak+7Zwo2Xt7GbuBGea4wsHd/02yn03mbzz4zAgTnAQkcHC/GghGJw4v60mhC8LyuJFsyFr0Jnddu7KLo+Xz55ChPT9j889PD3Laug41pk1HHY3XSRHo+NGkTVStci4USDIViAXF9ievLJUv9nC+GDx3jxF1745+P58p8/PAgvoT/emkvz3/0MMBzTjTsbL7pc3cPZPElXN+f4fldSf7PU8McHCvxi+s78XzJY6NFAK7sSgJBlth7dq6K39+Fzms3dPF3RwY5ki1zJDuIIcCVsC5lcsu6dq7uSS/sF5whSjAUigWkaF88fYhO3ruPoQNHCAMu9F11GZtvvr5KLPKuz788PczRicmg75dOjHJFR4KhRw/TtmH1c8o9ZXVk6kQj63g8OFTg8bEShoBXre8gpWtYmuB0wWHM9jieK5NzfVYnDdZP0dtrc5vFH+1cxXfOZnlktIgbep8Gig6fOT7C6qTJ+rTJsYkyP7qQ49L2BC9d1Yaeqk/PXQyUYCgUC0jR8efULnyxOHnvPoZCKwEAKRl69DBDjx2BsMDQk5J/OznK0YkyCU1wTW+aYxNlzpVcHhgqcOPqNgb27n9OCca6G6+JxfJbZ8fZN5gn60x22X3jlm46zCAV7LKOBE+MlfjPwRzfPRdkRl3fnwmKE5MWumVWxX1yZ84zdOAIa1Imv7mth5dO2ORcj23tCb58cozHRovcO5DlbZf2cvdAlsPZMg+PFNnVleLqV1+7+IuBEgyFYkEpOn5dquVyZOjAkfjfR7MlRmyPF/emERUdyL9wYpSHR4oI4N1X9LM+bbF/uMCnj4/w6GggGFO5cJaS2Qbqe3dt49Q9/4lTdvjeuRxORfv4XkvnxRU1FVd0JnlirMTdAxNAICA/v7oNoWtsfMV1ddfr3bUtduEd+eLdbD/5bPzc6zd18vhokcfGimQdj7MFJ37ujO3zCpVWq1A8t/B8ietdJPGLMJCdczz+7sgQAD2Wzra2BA8OFxAiyAgCeNOWbtang8Durs4kmoBj4em4zdAZPnRs0a2MqQSh1nqys3lO3n0/MH1q6vChY/iux9mig+NLLE2wpzvFT4cL3LKuo+q1l2Ssqp+v7EqhCcHLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"xx = torch.arange(data.shape[0])\n",
|
||||
"\n",
|
||||
"plt.plot(xx[1:], means)\n",
|
||||
"plt.fill_between(xx[1:], means - 2*np.sqrt(vrs), means + 2*np.sqrt(vrs), color=palette[1], alpha=0.5)\n",
|
||||
"plt.scatter(xx, data, color=palette[5])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "6be11181-0698-429c-9dc7-a0f4b6c99584",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Rollouts"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 42,
|
||||
"id": "81bb90d7-e790-4e31-ba87-fe61449eb9ce",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"nroll = 50\n",
|
||||
"roll_len = 100\n",
|
||||
"xin, xout = dset[len(dset)-1]\n",
|
||||
"xx = torch.cat((xin[0, 1:], xout.unsqueeze(0)))\n",
|
||||
"xx = xx.repeat(nroll, 1).unsqueeze(1)\n",
|
||||
"roll_pxs = torch.zeros(nroll, roll_len)\n",
|
||||
"with torch.no_grad():\n",
|
||||
" for idx in range(roll_len):\n",
|
||||
" out = model(xx)\n",
|
||||
" smpl = torch.normal(out[:, 0], out[:, 1])\n",
|
||||
" roll_pxs[:, idx] = smpl\n",
|
||||
" xx = torch.cat((xx[..., 1:], smpl.unsqueeze(-1).unsqueeze(-1)), -1)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 43,
|
||||
"id": "1aec9c78-083b-4161-af02-4ef9e6590858",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYcAAAEFCAYAAAAIZiutAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAABZ9UlEQVR4nO29d3wc9Z3//5ztVVr1bss2lhum2ARsDgIJJDHpCZeeu8uFSy4hR658c5dy973Aldwdxy+58r1CkguBhJAGoRy9xIADxmBhXIQt27LktWW1VdveZn5/zM5oV7uyLVvFlt7Px4MH1uzszme1q89r3l3RNE1DEARBEPKwzPcCBEEQhHMPEQdBEAShCBEHQRAEoQgRB0EQBKEIEQdBEAShCNt8L+BsSSQS7N27l5qaGqxW63wvRxAE4bwgm80yODjIhRdeiMvlKnr8vBeHvXv38qlPfWq+lyEIgnBecu+993LZZZcVHT/vxaGmpgbQ32B9ff08r0YQBOH8oK+vj0996lPmHjqZ814cDFdSfX09zc3N87waQRCE84up3PESkBYEQRCKEHEQBEEQihBxEARBEIoQcRAEQRCKEHEQBEEQijjvs5UEQTh/CbZ3svuhF0nHkwA4PC7Wv/8qWja0zfPKBBEHQRDmnMmiYJCKJWj/xXMA55VAaJqGoiinPO/YsWPnTcq9iIMgCLPKrgeep+eVDnMDrVreyMjRfrLpTMnztazKzp8+Q/vPnmXpFWu55MPXzPGKp+bNN9+krKyMbDZLNptl2bJlRCIR2tvbeetb32qed+LECRoaGoqe/8tf/pJbbrnlvGj1IzEHQRBmjV0PPE/39n0YAyc1TWPo8PECYfh5zwj/tK+f/kS64LmaptG9fR+7Hnh+Ttc8GU3TOHjwIAAHDx7k4MGDPPLII3R0dBCLxTh27BhjY2MAqKrKjh07ePLJJ4teR1VV4vE4R44cmdP1nykiDoIgzBrdr+w76eOxjMoLA1GCsTTfOxhCLTG1uHv7PoLtnbO1xFMSjUZpb29nfHyc48ePE41Gueaaa/B4PLz44ovs2bPHdClFo1H27dtHKBQiHo8XvM6bb74JQFdX15y/hzNB3EqCIMwKwfZOOMWE+gePjZn/7ktk2DuaYHWZE5tFwZLnw991/1ZAj0ME2zvpeGI78dEI7oCPtVs2zVp8QlVVTpw4QTwe59ChQ4RCIfx+P+FwGIfDQW9vL4qisGbNGlRVZWRkhOHhYZYsWcJ9993H7//+76MoCpqmsW3bNhwOBxaLhVQqhcPhOO1YhUE2m0VVVUZHR6fsiTRTiDgIgjAr7Hl4m/nvw+Ekvzw6SoPbzu8sq0BRFEaSGV4ajAJQ5bASSmX57qGQ+Zz3NJVxQ2MZANl0hp0/fYZQ9wmCOw+Ybqn4aKRAOGaa3t5eHnroIVatWsXg4KC5odfU1LBv3z4sFgvpdJrOzk5OnDiB0+nkggsuwGq1kk6nicVieL1e0uk0AwMDLF26FLfbzYEDBwgEAmzdupVrrrmGJUuWnHItmqbx6KOPUlZWRiwW493vfveMv998xK0kCMKMEWzv5Mlv3cODf/GfpGIJ8/hDx8YIxtLsCMUYSmb1c2N6jKHJbedjrRVFr/Xo8XFGU9mCY93b9xUFsrPpDLsfenGm3woAo6OjXHbZZTQ0NNDb24vH4yEUCmG1WvnMZz5DOp0mmUwyPDxMd3c3R48eZf369TgcDjweD0eOHGF4eJhoNIrT6aSurg5FUejo6OCpp54iEonw+OOPA7Bnzx5GRkZKrmNgYIDx8XH6+vpQVZXq6mp2795txnJmAxEHQTiHyN9cn/zWPfPqa58uwfZOdt2/lfhopOD4oXCSrkjK/PlYTP/38bguDqvKnLR6HRjOlSuqPdhzP+wbK/TbT0U6npyV39Xo6Chr1qyhoaGBwcFBc6hYNpvF7/cD+h29qqpomkY0GuXnP/85dXV1OBwOfvOb3/DEE0+wdetWAoEAFRUVHDlyhOPHj3P11VdTVlZGKKRbS8ePH2f//v1mcNvghRde4J577qG/vx+Px8PBgwex2+38+te/5sCBAzP+ng1EHAThHCDY3smDX/svdv70GQaHxni2L0xoaJz2nz973ghExxPbS6anPnyscLM7lrMYjhuWg8eOx2bhT1bX8Odra/mdZZX89tIAAA8FxzgcTnIsliKVVU95/dPh6aefPq3z3njjDY4fP86DDz5IdXU1qqrS2NhIIBDA7XaTSCSwWq04nU4sFguqqpoisWPHDtLpNIODg/T29tLf34+iKFRUVHDixAn8fj8rVqwgmUxitVoZGBggGo3y4osv0tHRQX9/PyMjI6TTaXbt2kU6nWb//v3U19cTCoXo6OhgbGzMzKKaDeY85tDV1WVG+Pfu3Ut3dzeapvGv//qvbNmyZa6XIwjzTrC9k50/fcb8+ec9o7QPxzkYTvKFldXsfujF86IgbLLFADCeznIkksKmwMdbK/jxkRGOxdJomsaRiF4At9TrAGCF32k+7/IqL1v7I5yIZ/jO/kEAyu0W3tNUzpU13tO+vkE4HGZsbAyPx8P+/ft529vehtVqnTIYnM1m6e/vJx6PMzIywvHjxykvL6empoZkMklzczPxeBy73U44HEZVVdxuN6qqUlZWhqIo7N+/H03T8Hq9+P1+MxDt9/tpbm5mcHCQ48ePA/DKK68QDodJJBK8/vrrZDIZBgcHGR0dxWq1oqoqwWCQtrY2MpkMoVAITdM4ceLEaXwyZ8aci8N9993HPffcM9eXFYRzglKZNsYdb1bT+N9j47QP666UvaMJBhMZZjcnZeZwB3xFG/QLAxE0YFWZixU+ffMPxlKMpLKMpVU8VoVaV/E2ZLco/NnqWv79wCBHcxbGWFrl5z0jrA+4cFktdIwlaHTbqck93x3wTbm2rq4uDhw4wIkTJ1i2bBnPPPMMiqLwzne+s0ggEokE//u//0sgEOAtb3kLzz77LE899RSVlZV4vV5SqRSVlZWMjo6STqexWCyUl5czPj6Ooiik02lGRkZMSyIUClFbW0tNTQ3t7e0oisKRI0ewWq3U1tYyMDDAnj17qK2txWazMTo6SiKRoK2tjZdeeolEIoGmaYyPjzM2Nobb7SYcDgO622u2mHO3UltbGzfddBPf+c53ePrpp7n88svnegmCMC/seuB5dv70GXMDNTJtjJ9fGozydF+44Dm7R3WhOB9cS2u3bAIgq2q0D8f44eEQT/Tq7+e6eh9VTisui8J4WuXR4+MALPc5sSiKuUG7Az42fvx6Nn78erxuO3+yppaPLQ3wByuqaPHYyWiwIxTjpcEo3zsU4rY9fUQyWax2m3n9UoRCIS6//HIuuugiMpkMkUiE5uZment7i87dt28f1113HW63m66uLlwuF6FQiPLycjweD+FwGK/XyxtvvIGmaVRWVnLBBReQzWax2WxkMhnS6bT5nqxWKx0dHXR0dDA4OIimaQwPD5uv6XA4yGQypkvKYrGwZ88edu7cicfjIZPJkM3qgfnu7m5sNht2u33mPrgpmHPL4SMf+chcX1IQ5p1geyfd24sLwsyUzKzKMyf0jfSDzeVUOKzc1TXMm2MJrqv30/HE9nPetdSyoY2dP3uGHaEY93ZPZN00e+y0lbkAPb5wOJLilVAMgHc1+lGsFjZ85O1F769lQxu7HngeR+73ltE0fpj7nVQ7J7auYc3CNTdeS8uGNrq6uli+fDmgu5Jee+01rr32WmKxGEeOHKG2tpaxsTG8Xi9Llixh//799Pb2ctFFF+F06pbN6OgoFRUVjI+Ps3//fiorKwF9FHEqlTI3/q6uLjKZDPF4HLfbjd1u5z3veQ9bt24lHA6zdOlSuru7SaX0APzQ0BA2m414PE42m6W3t5clS5ZQUVFBf38/qVTKFIeWlhaOHDmC3W7HYrHgcrnwer2EQiEzEyqRSJRs0TFTSJ2DIMwB+cHS0VSWR46NcX2Dnwa3fgf4+PFxQqksHquFq2u9RDN68DWY88+fzJ8+m0yn4CzY3onFZjUDzgZe24SDotZl43Auc8lrs7DM52TpW9ZM+ZqXfPgaU1TbyvTN+3AkVVBbN5JRzed3dHSY4tDe3s7hw4dJpVLE43F6e3tZsWIFXq+XNWvW4PP5iEQijI2NUVFRQUVFBVVVVVgs+nrtdjvxeByLxUJzczPV1dXs2LGD8fFxMpkM4+PjZuZSIpFg8+bN7N6923T59PT0YLPZzLt+ox+T0VcpmUxy+PBhACwWC/F4HKvVitfrZWxsjGQySTKZ5NJLLyUYDBKLxcwAuCE4VVVVp/oIzxjJVhKEOSB/c7/3yDCvhGJ892CItKrxWijGrhHdffSRpQGcVgsVDitem4VoRmUkl+s/166lyamphhus1DqMc9V0lsGkbg2t9DtR0IvZDGry4gt1uX/37+856TocHt3qKLNbaXDbSKsaB8YnurmeGIma/x4cHCSTyXD06FF27NiB1Wpl7969hMNhYrEYiqJgt9tpbGw03T6RSISuri527txZUDcQj8fxer3YbDYuu+wybDYbbrcbl8vF2NiYufFns1mGh4cZHh7m+PHjpNO6OKqqisPhKHo/RlW0USORSqXIZrMkk0nT5ZTvZhoZGWHdunUArFq1yrQaFEUxhWg2EHEQhDnACJZqmsabuY1tMJnhvw8O8cOuYYZzArCuXN8IFUWhxaNbFcad+Ommak6XqWor9jy8jWw6w3g6y7aBCFlVI5vOlFxHfhrrYEL//0eWBPiXjU0s901kIdXkuYPqXfr7O5VVtP79V5n/XuV3FT0+nMoWCFZXVxcvvPACmqaZPv5YLEZdXR3t7e0MDQ0xPq7HPBRFwWKxMDQ0REdHB+Pj42b6aSQSYc2aNdjtdpYsWcLRo0ex2+0sW7aMaDSKx+NBVVVcLhfd3d309PQQi8VMgXE4HDQ2Nhat19j0DQKBABaLxRQawKx1cDqdhEIhqqqqiMfjrF271kyLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"full_x = torch.arange(data.shape[0])\n",
|
||||
"test_x = torch.arange(data.shape[0], data.shape[0] + roll_len)\n",
|
||||
"plt.plot(full_x[1:], means)\n",
|
||||
"plt.scatter(full_x, data, color=palette[5])\n",
|
||||
"plt.plot(test_x, roll_pxs[:20, :].T.detach(), color='gray', alpha=1., lw=0.5)\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "0b7e9075-855e-4c81-8334-8507718152f1",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Just playing"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 44,
|
||||
"id": "83db2178-1011-4e7b-9070-8bb5955f943f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import torch\n",
|
||||
"import gpytorch\n",
|
||||
"import sys\n",
|
||||
"\n",
|
||||
"sys.path.append(\"../\")\n",
|
||||
"from voltron.likelihoods import VolatilityGaussianLikelihood\n",
|
||||
"from voltron.models import SingleTaskVariationalGP\n",
|
||||
"from voltron.kernels import BMKernel, VolatilityKernel, FBMKernel\n",
|
||||
"from voltron.models import BMGP, VoltronGP, MaternGP, SMGP, VoltMagpie\n",
|
||||
"from voltron.means import LogLinearMean, EWMAMean, DEWMAMean, TEWMAMean\n",
|
||||
"from gpytorch.kernels import ScaleKernel, RBFKernel, MaternKernel\n",
|
||||
"from gpytorch.means import ConstantMean\n",
|
||||
"from voltron.train_utils import LearnGPCV, TrainVolModel, TrainVoltMagpieModel\n",
|
||||
"from voltron.train_utils import TrainBasicModel, TrainVoltModel\n",
|
||||
"from voltron.rollout_utils import GeneratePrediction\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 45,
|
||||
"id": "c9c02c1e-c7ec-4719-9178-cacbb7c628a4",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"train_y = torch.FloatTensor(pxs).log()\n",
|
||||
"train_x = (torch.arange(train_y.numel())/252.)[1:]\n",
|
||||
"test_x = torch.arange(100)/252. + train_x[-1]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 46,
|
||||
"id": "dbe354ea-dc9d-479f-8025-48a1ca9cae22",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"vol = LearnGPCV(train_x, train_y, train_iters=200,\n",
|
||||
" printing=False)\n",
|
||||
"vmod, vlh = TrainVolModel(train_x, vol, \n",
|
||||
" train_iters=200, printing=False)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 47,
|
||||
"id": "b89d9a80-cf7f-4bf7-9cbf-59ee9a0563f6",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"voltron_lh = gpytorch.likelihoods.GaussianLikelihood()\n",
|
||||
"voltron = VoltronGP(train_x, train_y[1:], voltron_lh, vol)\n",
|
||||
"voltron.mean_module = ConstantMean()\n",
|
||||
"\n",
|
||||
"voltron.likelihood.raw_noise.data = torch.tensor([1e-5])\n",
|
||||
"voltron.vol_lh = vlh\n",
|
||||
"voltron.vol_model = vmod"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 48,
|
||||
"id": "50591bec-1240-4a8e-81b9-46b8295e642b",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Iter 1/200 - Loss: 7.323\n",
|
||||
"Iter 51/200 - Loss: 1.576\n",
|
||||
"Iter 101/200 - Loss: 0.906\n",
|
||||
"Iter 151/200 - Loss: -0.836\n",
|
||||
"Iter 201/200 - Loss: -1.960\n",
|
||||
"Iter 251/200 - Loss: -1.965\n",
|
||||
"Iter 301/200 - Loss: -1.965\n",
|
||||
"Iter 351/200 - Loss: -1.965\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"voltron.train();\n",
|
||||
"voltron_lh.train();\n",
|
||||
"voltron.vol_lh.train();\n",
|
||||
"voltron.vol_model.train();\n",
|
||||
"\n",
|
||||
"# Use the adam optimizer\n",
|
||||
"optimizer = torch.optim.Adam([\n",
|
||||
" {'params': voltron.parameters()}, # Includes GaussianLikelihood parameters\n",
|
||||
"], lr=0.1)\n",
|
||||
"\n",
|
||||
"# \"Loss\" for GPs - the marginal log likelihood\n",
|
||||
"mll = gpytorch.mlls.ExactMarginalLogLikelihood(voltron_lh, voltron)\n",
|
||||
"\n",
|
||||
"print_every = 50\n",
|
||||
"for i in range(400):\n",
|
||||
" # Zero gradients from previous iteration\n",
|
||||
" optimizer.zero_grad()\n",
|
||||
" # Output from model\n",
|
||||
" output = voltron(train_x)\n",
|
||||
" # Calc loss and backprop gradients\n",
|
||||
" loss = -mll(output, train_y[1:])\n",
|
||||
" loss.backward()\n",
|
||||
" if True:\n",
|
||||
" if i % print_every == 0:\n",
|
||||
" print('Iter %d/%d - Loss: %.3f' % (i + 1, 200, loss.item()))\n",
|
||||
" optimizer.step()\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "2bbbe667-d72d-4ece-898e-727ca1069d86",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 49,
|
||||
"id": "de189faa-0fea-4f7e-b7b5-2d8a795ffceb",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"voltron.eval();\n",
|
||||
"voltron_lh.eval();"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 50,
|
||||
"id": "91d4bf03-5636-45f1-8fdb-3cd6b30bf61f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"pred_vols = vmod(test_x).sample(torch.Size((10,)))\n",
|
||||
"samples = GeneratePrediction(train_x, train_y[1:], test_x, pred_vols.exp(), voltron)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 51,
|
||||
"id": "79430340-faf5-4693-91cb-53556ad9276b",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAY8AAAEFCAYAAAAbsWtZAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAABJj0lEQVR4nO3dd3hUZdrA4d9M2qT3ShJIgCT0jhRBBUR0AQuWZRHs6Opadj+liCtgWyyA2FhURBYB6wK6IEgRRDpSQ0IJJCSk996mfH9MZjKTQjLp5bmvy8uTc868553hzHnm7QqdTqdDCCGEsICytTMghBCi/ZHgIYQQwmISPIQQQlhMgocQQgiLSfAQQghhMevWzkBTKikpITIyEm9vb6ysrFo7O0II0S5oNBrS09Pp27cvKpWqXq/pUMEjMjKSGTNmtHY2hBCiXVq/fj1Dhw6t17kdKnh4e3sD+g/Az8+vlXMjhBDtQ0pKCjNmzDA+Q+ujQwUPQ1WVn58fgYGBrZwbIYRoXyyp7pcGcyGEEBaT4CGEEMJiEjyEEEJYTIKHEEIIi0nwEEIIYTEJHkKINkGt0XAlIaO1syHqSYKHEKJNWP3DQcbOWkbEnxYTFZPc2tkRdZDgIYRodTsPRvP6yp8BKCgq5e3VvwCQkV1AbGJma2ZN1KJDDRIUQrRtOp2OlV/v58q1DF7722Qc7G0BeGTBOrPzdh++wIg/v0NSei5arY4tHz3FkD7BrZFlUQspeQghWsyWPWd469PtfL3tONt/jwKgpKy8xnOvpeag1epXyT58JrbF8ijqR4KHEKJFaLVavvrpqPHvlIw8AC7Hpxv32atsanztvz7dwdOvbaS0TN28mRT1JsFDCNEiNmw9zuHTlSWIlIxcAKKvpALg4erAzs+fMx6/59aBbPrgSePfP/56lt2Hz1dLd8ue0xw9G9dMuRa1keAhhGgRe49dBGBw7yCgsuRxIVYfPB6+eyTdungaz3dQ2RIe4muWxtuf/8LBU1eMfyem5vDM699wz3OfotZomjX/wpwEDyFEszl06goJKdkcPHWF7fv1bRx/vkO/XkRqZj4A52NTAIioCBSG4w/fPQIXJxVfvjWLYH93AC4nZDDjpTUkpuYAkJyea7zW6fOJzf+GhJEEDyHaqLSsfNZsOkR2XlFrZ6VB4pOzuO/vnzNy+ru8/fkvxv0jB4YC+pLHoVNX+PWIvkQSEaJfg+ed/7uLs5tfMf49YWQEBze8xP9WPo2nmyPlag1Pv/41pWVqUjLzjOkeMimRdGQbo84ydO2nnMtIa9V8SPAQoo3JzCngT3/9hOH3v80/P/iJF9/9b2tnqUFir1WOz/jjXDwAi/82mQBvV+xsrUlMzWHpl7uN53QN8ABAqVTi7upQLb2BEYH89MlfcXdx4I9z8ew8GE1aRekFID27oLneSqtTa7UkF+Tz/YUoZm3dxOm0FF769Ze6X9iMJHgI0cb8cuA8p89fQ63RArDj9yjOXmx/VTKZOYXV9v35jiHY2VobSx+GBvSP//kAVlZ1P46C/T14evpYAHYcMA8e7amEduqHvez76AfUtXRTruqV33YTvHI503/83rhvZ9wVen76Afdu+oZr+XnXeXXzkOAhRCspLC7lk42/GevvDc5cvAbA3/5yE0/efyMAH2/Y19LZa7T07Pxq+xzt7QAYPSjUbP+Q3vUfAHjbjb0BfVC9bDIXVk5ecUOy2SJ0Oh0ZFSUjjVpD3JEosuNTeXvFl8xZu4GycjX7/rWOTY8vIS0qrtpr1507Y/z7vVsmMiJAv1JqXG4OW2Iu8NqBlr8/ZIS5EK1k9sIN7Dt2idMXrrFq0V8A2HPkAut+1I+FGDO0B/5erqz69neOVVT7tCdpWbVXI/l6uhi3FQoFvl4utZ5bVWigFyMGhHD4dCw/7z9n3N8SJY+sqylc/v0M/abeiMq5etVabT7//gCLP9nGvxdN58Yefsb9fdLLIT2HdSu/JueznwC4sucEYX8ZT++/3UU2Gs6mp5JWVIivgyOXn3weO2trXOzsOJx0zZjOmrMn6evtw3NDbmi6N1sHCR5CtILSMjX7jl0CYN+xSxQUlbL8P3s4YjIOYmBEIPZ2NtirbEjNyCM7rwh3l/o/sFpbxnWCh6ebo3Hb19MZG+v6r50NcOe4/mZjRgBy8hsWPHKTMijMysO/dwgKpeK65/72sb79SVOm5oaHb0ehqPl8tUbDe1/sori0nKenj2XxJ9sAeGrRRh4K9WCYp/m/oza6slrSxsmemO/28cemvaS5q4gNcIBhvjzSfxB21vpH9kN9B+Jmp2J811D+sWcHayNPMX/fLu4P74Ofk1ODPgdLSbWVEK0gPjnLuG2lVLDww/+x6pv9nDqv/zX5+esP4mhvh1KpJKybvgvr+SsprZJXg6S0HB6c+2W1h3ZtDNVWL8+ehMrOhpdnTzIeMw0eXXzdLM7LqIGh1fY1pNoq5rdT/Pr+txz9z3bi/ziPVqNFp9PVeG5ucmUHgJToOHa+sRZNmZrC9Bzifjtl9rof95zlow37WP3DQYbcu8QsnY1x2eSVm49JcS7TYm1vx8TlzxI86Qa6TR1NaO8e9CpUMOF4On/972VmFDmi0+koLylDqVBwd1gvXOzs+Pz2qTwzeDi2SivKtC031kVKHkK0AtOeSLkFJXyz/Q+z4xEmg+Miuvly+vw1Ll1NMzY0t4ZnXv+GY5FXOXY2jgvbFtV5flzFbLjjR4TzxH2jzUoXpsGjW4BntdfWJTTIC1cnFbkFJcZ9uQXFaLValErz38QleYVY2dpgo7I17suMS+birydIjb5q3Hfyu185+d2v2DmocHJU0fXG/pTkFKByc8K3Tyh7l34NSgVKTSleOafIKQtj+yufkHxsG3s1nmRb+eHUtydKO1uOn7tKVcN6+HEsJoUyrY6XTyVze14mI3xd8OwbgsrTlW4zxnP5wFlAX5XnEuKPS4g/1/acoJ+TE7/MWcmWZ5ah9HXlqe//hZOvO4nHL+Ds78mDcWrm3foAfi6uFn+WDSXBQ4hWEJt4/UWPAv3cjNshgfqHa1xiVi1nNx2NRsv76/bQM9iHqeP6mx07Fql/IBYWl9WZTm5BMfHJ2djZWtM92AtrK/NqKQ9Xk+DRxcPifCoUCjzcHI3Bw8VRRV5hCTn5xWZp5yZlsO+jH1AoFPS/awxOPYP4aP0+rp64SGlBMV52VjwwtjdpFxP071+n42RyDvlXkvBauRnQEeycwtfuPfG01/CUyzUOZ5SQaqVmstUWjmgCeYdJFFlVBKbIyrYpD62G4TmZbHf3xqu8jNvzMpnY15c3I/Uj6n928eT3/HIezi0h3FWFLqey2i3dSou3Rh8EQ4f5Y6csIWvQMOboEkjxUKFZuZ6gsymknI4BwMHfk6gdR7lz2bO4Bnpb/Hk2hAQPIVqIVqtl0cdbSUjJJq/ioVf11zPA0D7BZg9bw5QdcS2wrsXGrcdYvnYPAFNu6Wes04+JrxyQZmdb92PjXMViThGhftUCR9U0TEshlnB1sjduR4T6cvTsVT78ai8FxaU8fNsgyi9eJe5oNFRUJ5387le2F8P/Iq+ZpVN2LpmBAZ7sP3qJgxlF6EOjLQP9A0iwVqCy8iFR5wolcKLEhRScQQs/qntxyD8Qja0ChzgNXRXZXHapGKtSBnMiAinT+dM/PgNnVPj5WaGxKucVq2I+LLEi18aWfGsbPryYwRsRDgRYF1Nk34WfVcU89ecpxK/ZQWjCd3jl6ntaLR74N6556NtK3tem8n+nL5HmaUdcTy/u9g/Fwcqa4rxCXJHgIUSHsv5/x/jiv4fM9s26cwQfrt8LwOP3juau8QPoEWz+5TcEjxPR8eh0ulobaZvCDztPGbfjk7ONA/c2bqusVistU3PkTCwqOxsGhAfWmI5hptyIbr41HjcVEujVoLxGhPoZ24gev3c0R89e5bPvD+Bmo2BG0QZy1CGgCyFHpcTe2hrnzGQOni8ElPiprEFpRUpRKT+diuOnU9XTP2Vb0aitqwxSKTgbt4+6dCF3oL7E4RVuT45CSZZS/7dbWTGjDy5CpdWQqHKhS0ke6AfS00uhoo+NG5utevCVwwDy+9rwUU4Sbo4lRDupCbKzYfjWrxnhksP6vEgAyhRKtjiWA/r0M13tePVJfZdltdKKzeTwpEMAfwoPatBn2RASPIRoIYa5nUyNuyEcR3tbNm47zt0TBtT4MDYEj4zsQj777gCzK8Z+NDWtVkvU5crlX89eSjQGj5ir5lNhTHv+M2ysrTj89Uv4eDjz+4nL9O7uh6ebvqePofHa4zqlii/fmsXpC9e4cXD3BuX35dm3kZlTwEN3jiCsm49xf6BjCos8R/NbeRnhyljGlZ1j8JVMUuI1FAXejb1tEf8YNZC0nt1xT0nmdFwppy+dIzkznYkOEDyyC4d+vERysZICnS3dVWW4uwfT29uZ3WpbrsSnEjowkAPOKYB+kF+clRrDgx0gx9ae+wY/wNcnv9MHDhMO1iX0sE7hL3YFfHbDIMqtlZx2r+y+a/ikD6vcGDvyUW7Iz6BbWREZdo6EFWTwl6SzLAq7BbXSvES3qiiJxzPSGOjr36DP01ISPIRoIaaT+Bl4ujvytxk387cZN9f6OicHOyaMjGDXofMcOn2l2YJHfHK2WXvGoVOxTL6pHwDJ6foH4EN33sDaLUcAKFdr+HLTYfy9XXj5/R8JD/Fl9xfPA5BToA8ers721GbCyAgmjIxocH49XB1Z8+YsQB/4VHY2lJSWo/HqztYifU+v4zhznBEQgv4/NIAdD6VFQloks/uOwqq3Lb9p1BDqziWAvCKe6l/KishtFPYbT+ygKVy7qsF3RBfu6TWc5IJi7ti0msLycnp5ejGj0ImFxXHY6RQ8l+/MTlUxf9iVc9rFj5tGzsZfDTPTUxnW6wbsQwLw+fyvXPII4v4+kylWm4ww1+pAqSAsPYOL3vrSWKLKhf+qKsfA3JV2iUevnUSLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(train_x, train_y[1:])\n",
|
||||
"plt.plot(test_x, samples.T.detach());"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "7faf7644-ea71-442c-a5d1-15f197d0517c",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Check Volt + Constant Preds"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 52,
|
||||
"id": "088cecfb-8a26-4d08-bc87-906bab10004c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"preds = torch.load(\"./saved-outputs/AMZN/volt_constant2020-01-09.pt\").cpu()\n",
|
||||
"# preds = torch.load(\"../trading/saved-outputs/AMZN/volt_dewma100_2019-12-18.pt\").cpu()\n",
|
||||
"\n",
|
||||
"tckr = \"AMZN\"\n",
|
||||
"ntrain = 400\n",
|
||||
"lookback = 1\n",
|
||||
"data = GetStockHistory(tckr, end_date=\"2019-12-18\", \n",
|
||||
" history=ntrain + lookback).Close.to_numpy()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 53,
|
||||
"id": "df65a48d-ad1a-41a1-8122-d1c2d6087103",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"train_x = torch.arange(data.shape[0])\n",
|
||||
"test_x = torch.arange(preds.shape[-1]) + train_x[-1]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 54,
|
||||
"id": "74fc106c-b377-4724-8391-6534499de38c",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAZ4AAAEFCAYAAADT3YGPAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAABijUlEQVR4nO2deZgcZbX/v9V79/Qy+5pZk5msE7KRkASEiwghkusliBCDkUVQUS9eBQHxSuL1Khe5bNcFjJBgQvJD0bCIEBEQEEImELJNJjPZZt969ume3rt+f9S8b1f1MntPzyTn8zw89FRXVddMd/pb57znfI8giqIIgiAIgpgkVIm+AIIgCOL8goSHIAiCmFRIeAiCIIhJhYSHIAiCmFRIeAiCIIhJRZPoC5jquN1uHDt2DBkZGVCr1Ym+HIIgiGlBIBCA3W7HggULYDAYFM+R8AzDsWPHsHHjxkRfBkEQxLTk+eefx7JlyxTbSHiGISMjA4D0x8vOzk7w1RAEQUwPWltbsXHjRv4dKmdEwuPz+fDxxx/j3XffxcGDB9Hc3Iyenh6kpKRg8eLF2LhxI1asWBFx3H333Yc9e/bEPG9xcTHeeOONmM+/+uqr2L17N6qrqxEMBlFcXIzrrrsOGzZsgEoVe3lqrMdFg6XXsrOzMWPGjFEdSxAEcb4TbYliRMJz4MAB3HLLLQCkCGD+/PkwGo04ffo09u7di7179+LOO+/EXXfdFfX4JUuWoLCwMGJ7NCVkbNmyBbt27YJer8fKlSuh0Wiwb98+/OQnP8G+ffvwxBNPRP2FxnocQRAEMTmMSHgEQcBVV12FTZs2ReTq/vrXv+Luu+/Gr3/9a6xYsQIXXXRRxPHXX3891q9fP+KL2rt3L3bt2oWMjAzs3LkTRUVFAICOjg5s2rQJb775Jnbu3ImvfvWrE3IcQRAEMXmMKO+0cuVKPPnkkxGiAwBr167FtddeCwB45ZVXJuSinn76aQDA3XffzcUDANLT07F582YAwNatWxEMBifkOIIgCGLymJA+nnnz5gEA2traxn2u1tZWVFZWQqvVYs2aNRHPL1++HFlZWbDb7Th06NC4jyMIgiAmlwmpaqutrQUQe81m//79qK6uxsDAANLS0rB06VKsXr066kL/8ePHAQClpaURtd+M8vJytLW1oaqqCkuWLBnXcQRBEMTkMm7hsdvtvHLtyiuvjLrPSy+9FLFt1qxZePTRRzF79mzF9sbGRgBAbm5uzNfMyclR7Due4wiCIM4nRFGEy+WC0WiEIAgJuYZxpdr8fj/uuece9Pf3Y+XKlbj88ssVz8+ZMwc/+tGP8Nprr+HTTz/F+++/j6effhpz5szBqVOncMstt0Sk5wYGBgAARqMx5usmJSUBAJxO57iPIwiCOJ9ob29HRUUF6uvrE3YN4xKeBx98EPv27UNOTg5+8YtfRDx/88034ytf+QpmzZoFk8mEzMxMXHbZZfjjH/+IRYsWobOzkxcEMNhcutEq8ViPIwiCOJ84ffo0AODs2bMJu4YxC89Pf/pTvPjii8jIyMD27duH7MkJR6fT4Y477gAAvPvuu4rnWFTCIphosIiF7Tue4wiCIKYboiiiuroaTU1Nozqmo6ND0cfo9/vh9XrR398fj8uMyZjWeB566CHs2LEDqamp2L59u6J0eaSUlJQAiKyEy8vLAwA0NzfHPLa1tVWx73iOIwiCmG60t7ejpaUFwMi/z+x2Oy/CYjgcDhw+fBiiKGL58uUwmUwTfq3RGHXE8/DDD2Pbtm1ITk7Gtm3bMGvWrDG9cE9PD4DI6IOVZp88eRJutzvqsUePHgUAzJ07d9zHEQRBTDc6OjpGfQz7zpUzMDDAlylcLtd4L2vEjEp4HnnkETzzzDOw2WzYtm0b5syZM+YXfv311wEACxYsUGzPycnB/Pnz4fP5ovq4VVRUoLW1FRkZGVi8ePG4jyMIgphueDwe/nikDfHRrMJqamr4Y/n6uCiK6OzshM/nG8dVxmbEwvP4449j69atsFqtePbZZ3mEEYuqqiq88847CAQCiu1+vx/btm3Djh07AEgFCOGw9Z9HHnkEdXV1fHtnZye2bNkCALj99tsj+oDGehxBEMR0wu/3R3081P4NDQ0jPmd7ezuOHj2Kw4cPj/0ih2BEazxvvfUWfvOb3wAACgoKsHPnzqj7lZSU8C//pqYmfOtb30JycjKKioqQlZUFp9OJmpoatLe3Q6VS4e6778Yll1wScZ41a9Zgw4YN2L17N9atW4dVq1Zxs0+Hw4ErrrgCN91004QdRxAEMZ2QRyLhN/fROHnypOLn/Pz8CCGSn4el8hwOx3guMyYjEp7e3l7++NixYzh27FjU/ZYvX86FZ/bs2di0aROOHj2KpqYmHD9+HIIgIDs7G+vXr8fGjRsj0mxyNm/ejKVLl+L5559HRUUFgsEgSkpKhh1vMNbjCIIgpgOiKI464uns7OSPrVYriouLoVKpFJkhufDE+3tyRMKzfv36UblLA5KiPvDAA2O6KMa6deuwbt26STuOIAhiqhMIBHhBADAy4TGZTOjr6wMgNfarVCpotVrFPvLzyM8fCAQmfJQM3f4TBEFME3p7e3HgwAHFtpEIj1w4mJelRqOMO+QRj7x4IVaV8Hgg4SEIgpgmnDx5UiEKAFBZWTlsI6nP54MoisjIyOCRT3jEIxcer9fLH8djjAwJD0EQxDQhVsorvHggHL/fD5fLhebmZhw6dAhOpzMi4pFHTkzc5s6dC7PZPM6rjoSEhyAIYpoQLhbBYHDE5dSBQIAXDfT19cWMeAKBAILBIARBQGZmZlz8L0l4CIIgpglMZJKTk1FcXIzGxkY0NjYOWVLNquCCwSAXHo/HE3ONh5Vqa7XauJkuT8ggOIIgCCL+MOGZOXMmDAYDX3+J5jAgiiIEQeDHCILAhcTr9UKr1cJqtcLtdsPr9fJzyIUnXlDEQxAEMU1gUYlGo4FareYVaqIoKkqgT506hX379sHn83EhkUcvHo8HgiBg8eLFWLp0KQCQ8BAEQRCRsOhFo9HA7/cjKysLgiDwdRlGY2MjvF4v2tvbufuAXEhY8YAgCHw7q3ybDOGhVBtBEMQ0QO5YoFar4XK5IAgCVCoVLzIIr3oTRZHP2tHpdHy7vCRbpVJBrVYjEAjg5MmTfLQMRTwEQRDnOSyiUalUUKlUXIRUKhUCgUDM6rZoEY/P51MUJDBRks8zCy8+mEhIeAiCIKYooiiirq4OnZ2dqK6uhsPh4IIRLjxsu3ytJxgMRl3jAZSOBNGiG6vVOrG/jAxKtREEQUxRuru7cfbsWXR0dMBisaCjo4ObNssFhUU8TqcTjY2NvHTa7/dzgWIRk9FohMvlwoEDB3DJJZdArVYr0nAAsGzZsrg0jjJIeAiCIBIMK30Oh4nGwMAAr2DT6/U4cuQIj2zUajV3Jjh16hT6+vrQ1dWFzMxMLjxsfYil6Rgulwtmszki4jEajfH6VQFQqo0gCCKhtLe34/3334fdbo/6fDAY5P8BgNlsRldXF7q7uwGA9/N0dXVhYGAAwWAQAwMDAEJrOX6/H4IgQK/XIy0tjZ9bXp7d3d2Nrq4uiKIY97EIJDwEQRAJhKXGKisrIww55amyQCCAnJwcHvkwjEYjAoEAent7I5pKPR4PRFFEIBCAIAgwmUwoLCxUnJ/1APX29qKvr4/vG09IeAiCIBKIvAQ6fOKnXHhEUYRer496PCun9vl8XHjcbjd3opav76jVamRkZPDzHzlyBMePH+fni8cYhHBIeAiCIBKIvKdGPu0ZgKJazeVyRZRMWywWaDQaCIKAYDCoaCSVn5c9Zms3rFS6paUF3d3digiHhIcgCOIcZyjh8Xg8fGy1KIo8ggGA0tJSLFy4EBqNBiqVCj09PQBC0Y28rJpVwLGIiUVZ7Bj5mo7X68WRI0fiMoeHQcJDEASRINi4AkZ/fz/q6+tRVVWF/v5+HDt2DEAoRSaPeFg1msfjgUqlQl9fH0RR5IJhs9l4hMNegwlPeHOoXHhUKhW6urrw8ccfo6KiQjEUbqKgcmqCIIgEwSIYk8kEt9sNj8eDM2fOAJAiDyYixcXFqK+vVwiPvPeGCYfP5+PbRVGEzWaDy+UaVnhYqs1ms/HiBVYZN9TIhbFCEQ9BEESCaGlpAQBkZ2dHNGx2d3cjGAzCZDIhJyeHV6cBklAwgcnOzlYIDxMXURQxc+ZMZGVl8WICdky4pxs73mQyKcqtLRZLXHp6SHgIgiAShNPpBACkpqYiKSkp4vlgMIikpCReJh0MBjFjxgyUl5dz8SgtLYXNZgMgRUnyiEej0fBqOK1WywVGHvEwo1FAEia5VY5chCYSEh6CIIgEwQoL9Hp9RGQRCATg8/mg0WhgsVi48KSlpSE1NZXvp1arkZycDADcoZqlzgKBAFwuFwBwcQKUabqysjIUFBTwMddysYmXXxut8RAEQSQAVligUqmg0WhgMpn4c2q1Gv39/QgGg0hJSYHBYOCNnuFpMiBk8ul0OhXWOKy3BwBKSkr4/larFWVlZRBFEVlZWcjOzubO1DabDXq9Hh6Ph4SHIAjiXEIe7QiCoIh4SkpKUFFRAYfDgebmZj6kDVCWSTPkqbNgMAi1Wg21Wg2fz8cLEuR+bIIgIDc3V3GOwsJCiKIInU6HZcuWIRgMxm00AqXaCIIgEoBceABpYT8vLw8zZ85EdnY230+lUqGjo4P/XF1dHXGu0tJSGAwG5OTkAJAiJpVKBa/Xq5haOhTFxcU8KtJqtVFdEiYKingIgiASQLjwCIKA0tJS/nx6ejp6enqgUqlQXV0NnU4Xc+CbyWRCfn4+T6upVCoIgsDXd+TrPlMBingIgiASABOJ8JEEwWAQn376Kex2OwRB4Gs6qampCpdq+Xk6OzuxZMkSvk2r1UIQBN6LE89pomNhal0NQRDEeUIs4Tl16hQOHz7Mf2alzuz/gUBAMb/nk08+gdvtRnl5OT+GRVELine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(train_x, data)\n",
|
||||
"plt.plot(test_x, preds[:10, :].T.exp(), color='gray', alpha=0.5);\n",
|
||||
"# plt.ylim(1000, 5000)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "b94a72c1-92ab-4d5f-8f3d-b9498fd9ed15",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,494 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "e0bf2b78",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Warning no robinhood utils.\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import voltron \n",
|
||||
"import seaborn as sns\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import torch\n",
|
||||
"import pandas as pd\n",
|
||||
"\n",
|
||||
"import argparse\n",
|
||||
"import datetime\n",
|
||||
"import gpytorch\n",
|
||||
"from botorch.models import SingleTaskGP\n",
|
||||
"from botorch.optim.fit import fit_gpytorch_torch\n",
|
||||
"from gpytorch.likelihoods import GaussianLikelihood\n",
|
||||
"from gpytorch.mlls import ExactMarginalLogLikelihood\n",
|
||||
"from gpytorch.means import ConstantMean, LinearMean\n",
|
||||
"from gpytorch.kernels import SpectralMixtureKernel, MaternKernel, RBFKernel, ScaleKernel\n",
|
||||
"from voltron.means import EWMAMean, DEWMAMean, TEWMAMean, MeanRevertingEMAMean\n",
|
||||
"from voltron.train_utils import LearnGPCV, TrainVolModel, TrainVoltMagpieModel, TrainBasicModel\n",
|
||||
"from voltron.models import VoltMagpie\n",
|
||||
"from voltron.means import LogLinearMean\n",
|
||||
"from voltron.data import make_ticker_list, DataGetter, GetStockHistory\n",
|
||||
"\n",
|
||||
"import copy\n",
|
||||
"\n",
|
||||
"sns.set_style('white')\n",
|
||||
"# style.use('whitegrid')\n",
|
||||
"palette = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"sns.set(palette = palette, font_scale=2.0, style=\"white\", rc={\"lines.linewidth\": 2.0})"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 109,
|
||||
"id": "a1d5a6c3-6d89-4ba2-ba06-439b3fc23886",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"dpath = \"/home/greg_b/DATA/autoformer/exchange_rate/exchange_rate.csv\"\n",
|
||||
"dat = pd.read_csv(dpath)\n",
|
||||
"dat = dat.drop(['4', '5'], 1)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 111,
|
||||
"id": "638d226b-75ee-409f-8b0d-2e5bf1d541da",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAA0EAAAIjCAYAAADFthA8AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAABcSAAAXEgFnn9JSAAEAAElEQVR4nOzdd3hkZd0+8PucM2d6yaT3Ldnel957RxEpFkQFBMH3BVHxVdSfDSygomJ5UVH0FZReFOl1F1hYyu6yfbMtm94nyfSZU35/zCabackkmUnZuT/XxXWRkzlnzmySmXOf53m+X0HXdR1ERERERER5QpzqEyAiIiIiIppMDEFERERERJRXGIKIiIiIiCivMAQREREREVFeYQgiIiIiIqK8whBERERERER5hSGIiIiIiIjyCkMQERERERHlFYYgIiIiIiLKKwxBRERERESUVxiCiIiIiIgorzAEERERERFRXjFM9QlQ5k4//XT09vbCZDKhurp6qk+HiIiIiGjCmpubEQ6HUVhYiNdee21SnpMhaAbp7e1FKBRCKBRCf3//VJ8OEREREVHW9Pb2TtpzMQTNICaTCaFQCGazGXV1dVN9OkREREREE7Z3716EQiGYTKZJe06GoBmkuroa/f39qKurwxNPPDHVp0NERERENGGXXHIJtm3bNqnLPXIegtrb21FfX4+2trahKVwulwvFxcVYvnw5SktLc30KY7Z3715s27YNXV1diEQisNlsqK2txapVq1BQUDDVp0dERERERBOQ9RDU3d2NV199FW+//TbWr1+Pnp6eER9fW1uLyy67DJ/4xCfgdruzfToZi0ajePDBB3H//fejsbEx5WMkScKJJ56I66+/HkcdddQknyEREREREWVD1kLQli1b8POf/xzvvfceNE3LeL/Gxkb88pe/xJ/+9Cf8v//3//Dxj388W6eUsYaGBtx0002or68f8XGqqmLt2rVYu3YtrrzyStx6662QZXmSzpKIiIiIiLIha32Ctm7divXr148pAA3n8/lw66234rbbbsvWKWVkz549+NSnPjVqAEr0wAMP4Oabb4aiKDk6MyIiIiIiyoWcrgmaNWsWTjjhBBxzzDGoq6tDUVERTCYTurq6sHHjRjzyyCPYtGlT3D7/+Mc/UFhYiBtvvDGXpwYgFry+9KUvwePxxG1fuXIlrrrqKixduhQulwttbW14+eWX8cADD6Cvr2/oca+88gp++ctf4hvf+EbOz5WIiIiIiLIj6yHIYDDgvPPOwyc/+Ukcc8wxKR/jcDgwd+5cXHrppXjooYfwox/9CNFodOj799xzD8477zzMmzcv26cX53e/+13S+p+rrroKt956KwRBGNpWUFCAxYsX4/LLL8e1116L3bt3D33vvvvuw0UXXYRFixbl9FyJiIiIiCg7sjYdThRFXHjhhXjmmWdw1113pQ1AiT71qU/hhz/8Ydw2RVHw+9//PlunllJHRwf++c9/xm0766yz8K1vfSsuAA1XXl6Oe++9FzabbWibruu4++67c3quRERERESUPVkLQZdddhl++ctfYvbs2WPe99JLL00KTWvWrEEkEsnS2SX785//jHA4PPS12WzG97///VH3q6iowJe//OW4ba+++ip27tyZ9XMkIiIiIqLsy1oIkiRpQvtffPHFcV/7/X7s2rVrQsdMR9d1vPDCC3HbzjvvvIx7Fl122WWwWq1x25577rmsnR8REREREeVO1kLQRKVaU9PV1ZWT59qyZQs6OjritiWGsJHY7XacddZZcdteffXVbJwaERERERHl2LQJQWazOWlbMBjMyXOtXbs27mtZlnHEEUeM6RiJ0/fq6+vR2to64XMjIiIiIqLcmjYhKFWAKCwszMlzJfYEWrp0KUwm05iOkSo0jbXXEBERERERTb5pE4LefffdpG2zZs3KyXPt27cv7uu5c+eO+RizZ8+GwRBfYTzxuERERERENP1MixCkqiqefvrpuG0LFixAZWVlTp6roaEhbtt4nkeSJJSUlMRtYwgiIiIiIpr+pkUIevTRR9HW1ha37YILLsjJc3k8nrjGrECs/894VFRUxH2dWGyBiIiIiIimnykPQe3t7fjFL34Rt62goACf+cxncvJ8gUAgaZvdbh/XsRL3S3VsIiLKDlXT4Qsr0HR9qk+FiIhmOMPoD8kdRVFwyy23wOv1xm2/5ZZb4HQ6c/KcqYJKqsp0mUjcjyGIiCg3QlEVW1r8UDQdsiTAZTFA1wGHWUK50whBEKb6FImIaAaZ0hD0k5/8BO+//37ctlNOOQWf+MQncvacqYLKWCvDpduPIYiIKDeaPGEoWmwEKKrq6PbFpjX3+KPo8UexuNwGSWQQIiKizEzZdLj7778f//jHP+K2lZeX484778zp8+opplGM9w5iqmMREVF26bqOgaCS9vvekIp3GwbgDaV/DBER0XBTEoKeffZZ/OQnP4nb5nA48Mc//jFnvYEG2Wy2pG2hUGhcx4pEInFfW63WcR2HiIjSCys6IuroN532d+emwTYRER1+Jj0ErV27Ft/4xjegadrQNrPZjD/84Q9YtGhRzp8/VVAZbwhK3I8hiIgo+zId4fFHNISi2ugPJCKivDepIej999/Hl7/85bgS1bIs4+6778ZRRx01KeeQKqj4fL5xHStxP4YgIqLs6/JFR3/QQb3+zB9LRET5a9IKI2zbtg033HADgsFD0xVEUcSdd96J0047bbJOA263G7IsxwWx9vb2cR0rcb/S0tIJnRsRHb4G1xBOtIpZf1BBS18YogDUFpphNUojPr69P4xufxQOk4QqtxmGGVQ8IBTVsLPdj2Ca0R2jJMBqlNA3bL1Qjz+KyoLxFbshIqL8MSkhaPfu3bjmmmuSSmH/8Ic/xIUXXjgZpzBEkiTMmjULe/bsGdrW2to65uOoqorOzs64bXV1dRM+PyI6fEQUDbIkoNMbxYHeECQBqCuxoMAqj+t47f1h7O85NA03EPFjVY0D4sFgFYyoCEY1WI0SzLIIjz869HhvSIU3rM6YKmq6rmN/TzApAEkCsKrGAVXTYZZF9AWVuBDkC6vwR1TYRgmHRESU33IeghobG3H11Vejr68vbvs3v/nNnJbCHkldXV1cCNq3b9+Yj3HgwAEoSvw89blz50743Iho5tN1HXu6gkNlnAepAPZ1B7G6xjDmEaFmTwhNnnDctrCio7UvjGq3Gb3+KHZ1HCrTb5HFpADhDalo6489frryBKLY2xVENE0hhFlFFhgNh2ZyuywGyJIQ9/iOgQjmFltyfq5ERDRz5XRNUFtbG6666ip0dXXFbb/ppptwzTXX5PKpR7RgwYK4r7dt24ZwOJzm0al98MEHox6XiPLTQEhNCkCDwoqO7W1+7OrwwxfObMF/OKqh2ZP6ParJE8aOdj/2dMb3KUs3haxnGq+ZCUZU1HcE0gagCpcRpY74UTRREFDqMMZt6/VH2cKAiIhGlLMQ1N3djauuugotLS1x26+++mrceOONuXrajJxyyilxX0ejUWzcuHFMx3jvvffivl6wYAEqKysnfG5ENPONVs1sIKSi169ga4sf3pAy6gV7f0jBSI/oCyjIoII0ACAY0aBNs4Cg6Tp6/VFsavZBS3NqhTYDZhdZUo6gldjjg1FUzaykNhER5a+chKD+/n5cffXVaGhoiNv+yU9+ErfeemsunnJMli9fjrKysrhtTz31VMb7+3w+vPTSS3HbzjjjjGycGhEdBtKNwiTSAWxt9WNXRwDBiIpuXwShqAZfSIkLKr6QmrVz0xELQtOFpuvY0eaPm8qXSrHdmPZ7ZlmElJCN9nYF0dYfhsIwREREKWQ9BPn9flx77bWor6+P237RRRfhBz/4QbafblwEQcC5554bt+35559PmraXzhNPPIFAIP4D+7zzzsva+RHRzDbWkOEJKNjU7MPuziA2NnmxpdWPDw544QnEpq4lTpsrsctwWw1It6rIbpJQ4TJCTkwGB/kj2QtVE9XQE8LAKCHPZhRRaE2/hFUQBNhM8YUQ+oMKGnpCqO8McGocERElyWoICofDuOGGG7B58+a47eeccw7uuOMOiGL2B54WLlwY999nP/vZjPa79tprYTIdKqMaDAZx++23j7pfe3s77r777rhtZ5xxBhYvXjy2EyeinOnyRvD+gQFsbPRiIKRA0/RJuxDWdR3B6MRDhqLp2NkeQEtfGP6EUFXiMGJRuQ1HzXbCnSIczCm2YHaRBUfWOnDULEfSYwLhQ+enaDra+sNo8YQQUSZ3hMgbUtAxEBnxMWZZxKJy26iFJBJD0KD+oJLxyBwREeWPrKUSRVFw88034913343bfsopp+Cuu+6CJE2vcqVlZWW44oor4ra98MILuPPOO9NeLHV0dOC6666La5IqCAJuvvnmnJ4rEWVOUXXs645VFwspGra1+rG+YQDr9w9gZ7sfoRxfEEcUPWldy5IK27iP19gbStpmP3jBbxAFLCyzYl6pBS6LARZZxNxiy9D3BUGALIlJAcE3bCRoX1cQDT0hNHrC2N7mh5ZuUc4oVE1HsyeExt7Mw1RitbtBxXYZi8qtqCuxYGWVPa4aXDr2NCEIiK2ZIiIiGi5rJbJ/+9vf4rXXXos/uMGAWbNm4Te/+c24jrl06VKcf/752Ti9lG688Ua8/PLLaGpqGtp23333YcOGDbjqqquwbNkyOJ1OtLW14eWXX8YDDzwAj8cTd4yrr74aixYtytk5EtHY+CNqysX1OmLTznxhH46ocUDMUa+cQMIokCQATrMEu0mCLzzxESKbSYrr8yMIAkrsRpSMsGbGntAzxxdSoag6VF2PqxYXjGpY3zAAo0GASRJR5jSixJH+uIOiaixsDo649PqjWFltH3H0xhOIoj8YH05sJglVBSYUWsdeQtxhSv9x5gmwgSoREcXLWgjq6OhI2qYoCu6///5xH/PjH/94TkOQ3W7HPffcgyuvvDKuj9GmTZvwla98ZdT9zzjLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 900x600 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.figure(dpi=150)\n",
|
||||
"plt.plot(dat.iloc[:600, 1:]);\n",
|
||||
"# plt.axvline(0.8 * dat.shape[0])\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 48,
|
||||
"id": "4f2f1338-5874-4ea7-9ea0-401201dbe4c7",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import torch\n",
|
||||
"import gpytorch\n",
|
||||
"from gpytorch.utils.cholesky import psd_safe_cholesky\n",
|
||||
"from gpytorch.utils.cholesky import psd_safe_cholesky\n",
|
||||
"\n",
|
||||
"def GeneratePrediction(train_x, train_y, test_x, pred_vol, model, latent_mean=None, theta=0.5):\n",
|
||||
" vol = model.log_vol_path.exp()\n",
|
||||
" if model.train_x.ndim != test_x.ndim:\n",
|
||||
" test_x_for_stack = test_x.unsqueeze(0).repeat(model.train_x.shape[0], 1)\n",
|
||||
" else:\n",
|
||||
" test_x_for_stack = test_x\n",
|
||||
" if vol.ndim == 1:\n",
|
||||
" vol_for_stack = vol.unsqueeze(0).repeat(pred_vol.shape[0], 1)\n",
|
||||
" else:\n",
|
||||
" vol_for_stack = vol\n",
|
||||
"\n",
|
||||
" full_x = torch.cat((model.train_x, test_x_for_stack),dim=-1)\n",
|
||||
" # print(\"vol stack = \", vol_for_stack.shape)\n",
|
||||
" # print(\"pred_vol = \", pred_vol.shape)\n",
|
||||
" full_vol = torch.cat((vol_for_stack, pred_vol),dim=-1)\n",
|
||||
"\n",
|
||||
" test_x.repeat(2, test_x.numel())\n",
|
||||
"\n",
|
||||
" idx_cut = model.train_x.shape[-1]\n",
|
||||
" \n",
|
||||
" cov_mat = model.covar_module(full_x.unsqueeze(-1), full_vol.unsqueeze(-1)).evaluate()\n",
|
||||
" K_tr = cov_mat[..., :idx_cut, :idx_cut]\n",
|
||||
" K_tr_te = cov_mat[..., :idx_cut, idx_cut:]\n",
|
||||
" K_te = cov_mat[..., idx_cut:, idx_cut:]\n",
|
||||
"\n",
|
||||
" train_mean = model.mean_module(model.train_x)\n",
|
||||
" train_diffs = model.train_y.unsqueeze(-1) - train_mean.unsqueeze(-1)\n",
|
||||
" # use psd cholesky if you must evaluate\n",
|
||||
" K_tr_chol = psd_safe_cholesky(K_tr, jitter=1e-4)\n",
|
||||
" pred_mean = K_tr_te.transpose(-1, -2).matmul(torch.cholesky_solve(train_diffs, K_tr_chol))\n",
|
||||
" # print(voltron.mean_module(test_x).detach().T.shape)\n",
|
||||
" # print(pred_mean.shape)\n",
|
||||
" pred_mean += model.mean_module(test_x).detach().T.unsqueeze(-1)\n",
|
||||
" \n",
|
||||
" if latent_mean is not None:\n",
|
||||
" pred_mean -= theta * (pred_mean - latent_mean)\n",
|
||||
"\n",
|
||||
" pred_cov = K_te - K_tr_te.transpose(-1, -2).matmul(torch.cholesky_solve(K_tr_te, K_tr_chol))\n",
|
||||
"\n",
|
||||
" pred_cov_L = psd_safe_cholesky(pred_cov, jitter=1e-4)\n",
|
||||
" samples = torch.randn(*cov_mat.shape[:-2], test_x.shape[0], 1).to(test_x.device)\n",
|
||||
" samples = pred_cov_L @ samples\n",
|
||||
"\n",
|
||||
" if pred_mean.ndim == 1:\n",
|
||||
" return samples + pred_mean.unsqueeze(-1)\n",
|
||||
" else:\n",
|
||||
" return (samples + pred_mean).squeeze(-1)\n",
|
||||
" \n",
|
||||
"\n",
|
||||
"\n",
|
||||
"def Rollouts(train_x, train_y, test_x, model, nsample=50, method = \"volt\", theta=0.5):\n",
|
||||
" if method != \"volt\":\n",
|
||||
" return nonvol_rollouts(train_x, train_y, test_x, model, nsample=nsample)\n",
|
||||
" latent_mean = train_y.log().mean()\n",
|
||||
" ntest = test_x.numel()\n",
|
||||
" samples = torch.zeros(nsample, ntest)\n",
|
||||
" pred_vol = model.vol_model(test_x).sample(torch.Size((nsample, ))).exp()\n",
|
||||
" samples[:, 0] = GeneratePrediction(train_x, train_y, \n",
|
||||
" test_x[0].unsqueeze(0), \n",
|
||||
" pred_vol[:, 0].unsqueeze(1),\n",
|
||||
" model, latent_mean, theta).squeeze()\n",
|
||||
" train_stack_y = train_y[1:].log().repeat(nsample, 1)\n",
|
||||
" train_stack_vol = model.log_vol_path.repeat(nsample, 1)\n",
|
||||
" \n",
|
||||
" for idx in range(1, ntest):\n",
|
||||
" stack_y = torch.cat((train_stack_y, \n",
|
||||
" samples[:, :idx].to(train_stack_y.device)), -1)\n",
|
||||
" stack_vol = torch.cat((train_stack_vol, \n",
|
||||
" pred_vol[:, :idx].to(train_stack_vol.device).log()), -1)\n",
|
||||
"\n",
|
||||
" rolling_x = torch.cat((train_x, test_x[:idx]))\n",
|
||||
" model.mean_module.train_y = stack_y\n",
|
||||
" model.mean_module.train_x = rolling_x\n",
|
||||
" \n",
|
||||
" model.train_x = rolling_x\n",
|
||||
" model.train_y = stack_y\n",
|
||||
" model.log_vol_path = stack_vol\n",
|
||||
" samples[:, idx] = GeneratePrediction(train_x, train_y, \n",
|
||||
" test_x[idx].unsqueeze(0), \n",
|
||||
" pred_vol[:, idx].unsqueeze(-1),\n",
|
||||
" model, latent_mean, theta).squeeze()\n",
|
||||
" return samples, pred_vol\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 94,
|
||||
"id": "cc2d8d8a-18b0-49d3-bd39-04defa131b62",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"dat_idx = 2\n",
|
||||
"ntrain = 500\n",
|
||||
"ntest = 150\n",
|
||||
"nstart = 0\n",
|
||||
"ys = dat.iloc[nstart:, dat_idx].to_numpy()\n",
|
||||
"\n",
|
||||
"train_x = torch.arange(ntrain-1).float()/365\n",
|
||||
"train_y = torch.FloatTensor(ys[:ntrain])\n",
|
||||
"test_x = torch.arange(ntrain, ntrain + ntest).float()/365\n",
|
||||
"test_y = torch.FloatTensor(ys[ntrain:ntrain+ntest])\n",
|
||||
"\n",
|
||||
"# if torch.cuda.is_available():\n",
|
||||
"# use_cuda = True,\n",
|
||||
"# train_x, train_y = train_x.cuda(), train_y.cuda()\n",
|
||||
"# test_x, test_y = test_x.cuda(), test_y.cuda()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 95,
|
||||
"id": "a12f5f71-fd26-440b-a820-b5db15a58a05",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAaUAAAEhCAYAAADf879gAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAABkYklEQVR4nO2dd3xT5ffHPzc73btQ9rAFKcieX1CGgiCIAxFREBC+iiI/FQQHCA5AFGSIfJG9FQcgoqACKigbZK8KhZa2tHSPNPP+/kjuzV5t0qTJeb9evkzufXJzGm7yec55znMOw7IsC4IgCILwAwS+NoAgCIIgOEiUCIIgCL+BRIkgCILwG0iUCIIgCL+BRIkgCILwG0iUqoFGo0FmZiY0Go2vTSEIgggISJSqQU5ODvr27YucnBxfm0IQBBEQkCgRBEEQfoPI1wbYQq1W48SJE/jjjz9w6tQpZGVloaioCNHR0WjXrh1GjhyJLl26VOnau3btwtatW3HlyhXodDo0adIETzzxBEaMGAGBgDSaIAjClzD+WNHh77//xpgxYwAA8fHxaNWqFeRyOf79919cvXoVADBx4kRMnjzZrevOnj0bW7ZsgVQqRbdu3SASiXD48GGUl5fjwQcfxOLFiyEUCl2+XmZmJvr27Yt9+/ahfv36btlCEARBWOOXnhLDMOjfvz9GjRqFjh07mp376aefMGXKFHzxxRfo0qULunbt6tI19+7diy1btiA+Ph6bNm1C48aNAQB3797FqFGj8Ouvv2LTpk0YPXq0p/8cgiAIwkX8Ml7VrVs3LFmyxEqQAGDgwIF47LHHAAA//PCDy9dcsWIFAGDKlCm8IAFAXFwcZs2aBQBYuXIldDpd1Q0nCIIgqoVfipIz7r33XgDAnTt3XBqfk5ODCxcuQCwWY8CAAVbnO3fujMTEROTl5eGff/7xpKkEQRCEG9RKUUpPTwegX29yhYsXLwIA7rnnHshkMptjWrduDQC4dOlS9Q0kCIIgqkStE6W8vDxs374dAPDQQw+59JrMzEwAQFJSkt0xdevWNRtLEN4kv6gMhSUVvjaDIPyOWiVKGo0GU6dORWlpKbp164Y+ffq49LqKCv2XXy6X2x0TGhoKACgvL6++oQThAJVagz5jFqP1ox/i7JXb8MMEWILwGbVKlN577z0cPnwYdevWxSeffOLy67gvPcMw3jKNIFymoLgC+UX6yc/AF5dh9x/nfWwRQbgGy7Jen0TVGlH68MMP8e233yI+Ph7r1q1zeT0JMHpBnMdkC85D4sYShLcoKVOYPV/+9UEfWUIQrsOyLM5kluFSjnfDzrVClObNm4eNGzciJiYG69atM0vpdoV69eoBALKysuyO4erXcWMJwluUlivNnp+5nIlKldpH1hCEa2h0LBRqHYoVGmh13vOW/F6U5s+fj7Vr1yIqKgpr165F8+bN3b4Gl0J+7do1VFZW2hxz7tw5AEDLli2rbixBuEBJufU9+Puxaz6whCBcxzRqV6TQ4E6JyiuhPL8WpU8//RSrV69GZGQk1q5dixYtWlTpOnXr1kWrVq2gVquxZ88eq/PHjh1DTk4O4uPj0a5du+qaTRAOsQzfAUBufokPLCEI1zGVn6t3KnD9rgJlSq3H38dvRWnRokVYuXIlIiIisGbNGt7bccSCBQswYMAALFiwwOrchAkTAOiF7ubNm/zx/Px8zJ49GwAwfvx4KspKeJ3SMmtPKa+wzAeWEITr6Gx4RSKh55PH/LL23b59+7B8+XIAQMOGDbFp0yab45o2bcqLDaDfw3Tjxg3k5eVZjR0wYABGjBiBrVu3YvDgwejevTtfkLWsrAz9+vXDs88+650/iCBMKDGsKUklIvRo1wz7j17BXRIlws+x1CQGgEzk+Um8X4pScXEx//j8+fM4f952ymznzp3NRMkZs2bNQocOHbB582YcO3YMOp0OTZs2pdYVRI1SUKzP9Jz8XG80rR+H/UevkKdE+D2WoiQWMl7ZZuOXovT444/j8ccfd/t18+bNw7x58xyOGTx4MAYPHlxV0wii2lzPuAsAaFIvFnHRYQCAuwUkSoR/Yxm8Ewq8s+/TL0WJIAIVlmVx9uptAMA9jRMgNvTvIk+J8HcsM+28JUoUryIIAMu2/oFuIz5BbkEpf2zvoYvo8ORcnL6U4bH3+fngBeTcLUGoXIKm9eMQF2PwlEiUCD/HMnwn9cJ6ElBFUWJZFr/88gvee+89/Pe//7VqjFdRUYHjx4/jxIkTHjGSIKoDy7JYvHE/9h25YvO8TqfD3C/3IiOnED/sP8sff3H2VtzJL8UzU9d4zJazV/Re0pP920MiFiEiVAaJWIhyhQqKSpXH3ocgPI2pJoVJhWgUa7vjQnVxO3yXnp6OSZMmIS0tzW5NOalUinfffRe3bt3Ct99+i1atWnnGWoKoAgdPpuGTNb8BADIPzLE6/8/l2/xjnY7F5Rs5yC8qh0wqglqjtarAUB3uGmretWySCED/3YmLDkNWbjHyCsvQsG6Mx96LIDwJlxIeKRfh3rreK8fmlqdUXFyMMWPG4Nq1a0hJScHkyZMRFhZmNU4oFGLEiBG8R0UQvuL7X//BM1PXOhyz78hl/nHO3RKMe3cThr++2qNiBABfbP0TX/2kjx7ERBm/1PFcskMhVagn/BcufOftutZuidKaNWuQnZ2NXr164dtvv8VLL71kt2ke11bi77//rr6VBFFF5q7c63TMlXRjB+OjZ2/gZlaB1ZjqllNhWRZzvjRWE+Gy7gAgIkzfUqXURvkhgvAXuK+AtxMR3Lr+/v37wTAMpk2bBpHIceSvYcOGkEgkuHXrVrUMJIjq0LJpHbPn81f/Aq1WZ3aszMQjOnPlNmyhUlevnIppAgUAxEUZRSk0RKK3o8KznhlBeBIWNdMCyC1RyszMhEwmQ7NmzVwaHxISQk3zCJ8ikQjNni/Z9Dt2miQzAECpQQxSDOs8AFA3PhI92jXln1coqpeEkHbTvMpIfIyJKMmkAIByBYkS4b/4ZfgOALRa12aMKpUKZWVl1J+I8CkKhXVLiDsWxU/LDaI0/YWH+GN14iLw9cIXUDc+Uj+mmoJx9Fy62fPwUGPYO4w8JaIWwIuSl9/HLVGqX78+1Go10tPTnY79888/odFoXPaqCMIbKJTWHs5HK/aYeT6cp9TqniT+WFGpvpFZqFwvGOXV8JTSbuVi4bp9/PMnHjSvRB8aYvCUKiglnPBfdP4YvnvggQfAsizWrHG8b6OgoAAff/wxGIZB3759q2UgQVQHRaXt5nmTPvoapy9lQKfTocyQYBAeIsXXC8chKkKOF4f3AuAZUfr1b2N23+TneuOTqY+ZnQ8ziNLHq3/B3/9cr/L7EIQ38cvw3ZgxYxAZGYlvvvkGc+fORXZ2ttn5/Px8bN26FUOHDkVGRgYSEhIwYsQIjxpMEO5QYWdD6t6/LmHwxOVY+e1fvOCEyiXo0a4Zzu14FyMf6QQACDGIUkU1wne5+cYkhwlP/QcSsXmSUKhcyj9+6rVVVX4fgvAmNZV959bm2ZiYGCxbtgwvvfQSNmzYgA0bNvDnunTpgpISfayeZVlERkZi2bJlCAkJ8azFRFCi0+mqVMVdoXTcZvyD5T8D0AsSd33T8AQnGNUK32XokxxWvj8SkYb0b1M4T4kg/BluU4RfeUoA0LFjR+zcuRODBg2CSCQCy7JgWRbFxcVgWRZCoRADBw7E999/j9TUVG/YTAQZKrUGDzy/CBPf3+r2azlPafr4/ogMtxYEDnvCwIXvyhRKFBZX4OK/2TbHOeLfW3pRat4w3vZ7GBIdCMIXcL/hrowDvL+mVKUq4UlJSfj000/x0Ucf4dy5c8jLywPLsoiNjUVqaipl3BEeJf12Pq5n3MX1jLvo260FHut7n8teE7emNO7xbnjlmfux/Ks/8dGKPVbj2t/b0ObrI8L0WXIlZZXoP2EpsnKLsX/tZCQ3TrQ53ur9lWpk5BRBKBCgUZLtEkL1E6NduhZBeINruQqUVmrQtkG4w8rfuhpaU6pW6wqpVIqOHTt6yhaCsIlpG+bJc75Bdm4xXhn5gNPXabU6KFUaAIBMKgYAxMeE8+ebN4xHmsGLefrhDjavEROpn2AVFJcjK1fffPLQqX9titKxc+m4djOPX48CgBuZd8GyLBrVj7VaS+Kw3OBLEDVJfrl+4lZaqUFUiNjuOD5852V73ArfvfXWW5g7d67L4+fPn4+3337bbaMIwhTLDLpte0+59LpKlf51MqmYDzn07nwPf/6RB1rzj5Ob2PZ8oiP0a6IFxRX8sV/+umTVaoJlWTz+6peYtmC7WYiPD901sB26A/Rt0R/s3oJ/rlJrHP9hBOEhTMN2zvojaQyVUERCP0oJ3759O3bv3u3y+D179mD79u1uG0UQplgmK8il9mdzpqhU+o3eUonRQ4mNCsP+tZPx9YJxeLxfW/54vYRIm9fgPKXbd4r4Y4dO/Yv+45fyogfArF5eSZmxhh3niTWzs57E8eXskQgP1a9rFZcqHI4lCE+hdaOko1KjH+ytPkocXu886+1FMSLwUVhkvtlL87ZEafA4pBZhs+TGiXz47bNpTyIqQm53jSrWUM2bExeOO/mlOH81Cx1TGwGAmXdkKkr/Glqf20ty4BCLhIiJDEVpuRJlChUcjyYIz6AxqQOpcyJQKo1+rEToXVHy2tV1Oh3y8/Mhl9vPeCIIV+A8pfb3NtA/t7Mh1hIuDCYRC+2OGTagPR7s3tLu+ZhIffgu/Xa+1blTF40daYtMvJuSMuPjQkPYz7TWnT1CZNXfE0UQ7qA2cZV0DjLwWJaF0iBKPvWUysrK+L1HHDqdDtnZ2XZTCFmWRWlpKXbs2AGlUokWLVrYHEcQrsKJEBdKc7b3iIMXJUnVAwIRNvYVcdzMMgpVUYlRiEwFqswgMK7sRfJE9QiCcAeNmSjZH6fU6KBjAZGAgZcdJceitG7dOixbtszsWGFhId8ryRWGDRtLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(train_x, train_y[1:])\n",
|
||||
"plt.plot(test_x, test_y)\n",
|
||||
"plt.ylabel(\"Exchange Rate\")\n",
|
||||
"plt.xlabel(\"Time Index\")\n",
|
||||
"sns.despine()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 96,
|
||||
"id": "cf3dc1fe",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"if torch.cuda.is_available():\n",
|
||||
" use_cuda = True,\n",
|
||||
" train_x, train_y = train_x.cuda(), train_y.cuda()\n",
|
||||
" test_x, test_y = test_x.cuda(), test_y.cuda()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 97,
|
||||
"id": "8cbb21ae-dd4e-4346-8ec3-46552602a276",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/home/greg_b/miniconda3/lib/python3.8/site-packages/gpytorch/utils/cholesky.py:38: NumericalWarning: A not p.d., added jitter of 1.0e-06 to the diagonal\n",
|
||||
" warnings.warn(\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Iter 1/200 - Loss: 12.911\n",
|
||||
"Iter 51/200 - Loss: -0.538\n",
|
||||
"Iter 101/200 - Loss: -0.584\n",
|
||||
"Iter 151/200 - Loss: -0.585\n",
|
||||
"Iter 1/1000 - Loss: 1.020\n",
|
||||
"Iter 51/1000 - Loss: 0.836\n",
|
||||
"Iter 101/1000 - Loss: 0.647\n",
|
||||
"Iter 151/1000 - Loss: 0.466\n",
|
||||
"Iter 201/1000 - Loss: 0.305\n",
|
||||
"Iter 251/1000 - Loss: 0.175\n",
|
||||
"Iter 301/1000 - Loss: 0.083\n",
|
||||
"Iter 351/1000 - Loss: 0.025\n",
|
||||
"Iter 401/1000 - Loss: -0.007\n",
|
||||
"Iter 451/1000 - Loss: -0.023\n",
|
||||
"Iter 501/1000 - Loss: -0.031\n",
|
||||
"Iter 551/1000 - Loss: -0.036\n",
|
||||
"Iter 601/1000 - Loss: -0.039\n",
|
||||
"Iter 651/1000 - Loss: -0.042\n",
|
||||
"Iter 701/1000 - Loss: -0.044\n",
|
||||
"Iter 751/1000 - Loss: -0.045\n",
|
||||
"Iter 801/1000 - Loss: -0.046\n",
|
||||
"Iter 851/1000 - Loss: -0.047\n",
|
||||
"Iter 901/1000 - Loss: -0.048\n",
|
||||
"Iter 951/1000 - Loss: -0.049\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"with gpytorch.settings.max_cholesky_size(2000):\n",
|
||||
" vol = LearnGPCV(train_x, train_y, train_iters=200,\n",
|
||||
" printing=True)\n",
|
||||
" vmod, vlh = TrainVolModel(train_x, vol, \n",
|
||||
" train_iters=1000, printing=True)\n",
|
||||
" "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 98,
|
||||
"id": "17dd53e9-cf5c-4756-afbc-3be069a24ff3",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAA5sAAAFXCAYAAAAs3lr2AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAADEtklEQVR4nOzdd3xTdfcH8E/adKR7l7Jp2WWVPR0sEQXZPlVkCiqoyCOCyoMMFfkhUxkispciyFJW2SCFsqEtlDK66N5N2zTz98fNvUmatE3SpEna8369nteT3twk36pNcu4533N4CoVCAUIIIYQQQgghxITsLL0AQgghhBBCCCG1DwWbhBBCCCGEEEJMjoJNQgghhBBCCCEmR8EmIYQQQgghhBCTo2CTEEIIIYQQQojJUbBZDVKpFCkpKZBKpZZeCiGEEGI16POREEIIQMFmtaSnp2PAgAFIT0+39FIIIYQQq0Gfj4QQQgAKNgkhhBBCCCGEmAEFm4QQQgghhBBCTI6CTUIIIYQQQgghJse39AJ0kUgkuHnzJi5evIjbt28jNTUV+fn58Pb2RlhYGN5991306NHDqOc+duwY9u3bh7i4OMjlcjRr1gyjR49GeHg47Owo9iaEEEIIIYQQU7DKYPPGjRuYPHkyAMDf3x+hoaEQCAR4+vQpTp06hVOnTmHGjBmYNWuWQc+7ePFi7N27F05OTujVqxf4fD4iIyOxZMkSREZGYu3atbC3tzfHr0QIIYQQQgghdYpVBps8Hg+vvfYaJkyYgK5du2rcd/z4ccyZMwcbNmxAjx490LNnT72e89SpU9i7dy/8/f2xe/duNG3aFACQnZ2NCRMmICIiArt378bEiRNN/esQQgghhBBCSJ1jlXWjvXr1wk8//aQVaALA0KFDMXLkSADA0aNH9X7OTZs2AQDmzJnDBZoA4Ofnh0WLFgEANm/eDLlcbvzCCSGEEEIIIYQAsNJgsypt27YFAGRkZOh1fnp6OmJiYuDg4IAhQ4Zo3d+9e3cEBgYiKysLd+/eNeVSCSGEEEIIIaROsslgMyEhAQCzn1MfsbGxAIAWLVrA2dlZ5znt27cHADx8+LD6CySEEEIIIYSQOs7mgs2srCwcOnQIADB48GC9HpOSkgIAqF+/foXnBAUFaZxLCDGfklKxpZdACCGEED0VFYugUCgsvQxig2wq2JRKpfjiiy9QVFSEXr16oX///no9rqSkBAAgEAgqPMfV1RUAUFxcXP2FEkIqFPUgAS2HLsJPu89beimEEEIIqULUgwS0eXMJvtt00tJLITbIpoLNhQsXIjIyEkFBQfjxxx/1fhx7JYbH45lraYQQPR05dx8AsHxLBF6dtBp/Rdy17IIIIYQQUqH1ey8CADb9cdnCKyG2yGaCze+++w4HDhyAv78/tm/frvd+TUCVtWQznLqwGU32XEKIebi5OHG34xOz8OnS/RZcDSGEEEIqcvZaHM5ei7P0MogNs4lgc9myZdi1axd8fHywfft2jdEl+mjQoAEAIDU1tcJz0tPTNc4lhJhHfqH2RR+xRGqBlRBCCCGkMhO/2mHpJRAbZ/XB5vLly7Ft2zZ4eXlh27ZtaN68ucHPwY5KiY+Ph0gk0nnOgwcPAABt2rQxfrGEkCrlF5VqHftw0T4LrIQQQgghhJiTVQebK1aswJYtW+Dp6Ylt27ahdevWRj1PUFAQQkNDIZFIcPKk9ubmqKgopKenw9/fH2FhYdVdNiGkEnk6Mpunrz5EVm6RBVZDCCGEEH1JZTJLL4HYGKsNNtesWYPNmzfDw8MDW7du5bKTlVm5ciWGDBmClStXat03ffp0AEwAm5iYyB3PycnB4sWLAQDTpk2DnZ3V/iMhpFbIL9TObAJARg4Fm4QQQog1KBNL8dq0n7WOF+ioTiKkMnxLL0CXs2fPYuPGjQCAxo0bY/fu3TrPCw4O5oJIgJnB+fz5c2RlZWmdO2TIEISHh2Pfvn0YNmwYevfuDT6fj8jISAiFQgwcOBDjx483zy9ECOHkF2lmNju1boi7j1KQkVOEdi0stChCCCGEcG48SEDMkzSt43mFJfD1crPAioitsspgs6CggLsdHR2N6Ohoned1795dI9isyqJFi9ClSxfs2bMHUVFRkMvlCA4OxujRoxEeHk5ZTUJqgHpm09fLFS2aBODuoxRkUmaTEEIIsQoSmVzn8dyCiic7EKKLVQabo0aNwqhRowx+3LJly7Bs2bJKzxk2bBiGDRtm7NIIIdVQJpaiRCQG394O0UcXwI7Hw9rd5wEAWXkUbBJCCCHWICe/WOdxCjaJoawy2CSE1E45+UIAgJeHgJu36e/NlOPQnk1CCCHEOmTnCXUepz2bxFBUN0oIqTEPnzLzbFs0CeCOBfq6AwCV0RJCCCFWIouCTWIiFGwSQmrMg/hUAED7Fg24YwHKYDMjp9AiayKEEEKIpuzcCoJNIQWbxDAUbBJSy8jlujf163u/uUhlMkRcfQgAaNeiPne8QYAXAOBFZoGuhxFCCCGkBsUnZiIhNUfnfRRsEkNRsElILXIzOhEtXl+E7Ycidd6/9NeTaD/ieySl5dbwyoA/T97BvbgXcODbo2fHptzxQD8P2NnxkJlTBLFEWuPrIoQQQggj7nkGXp20BrdiknTeX1AkquEVEVtHwSYhVizhRU6Fb/i6jJ+3HWViKf730zGd92/YdwkFRaX44ddTplqi3mKeMiW008b2QX1lNhMAHPj2qOfnAYVCgbQsKqUlhBBCLCXy7jONn28f/AqfTeiPfl2aA6A9m8RwFGwSYsX6jl+Jtz7+Ra/9jMWlZRCWlAEAXJwdKz330fN0k6zPEBnZTAOg9moltKyGgV4AgJSMvJpcEiGEEELUODs5cLft7ezg5+WKOZMH4rMJrwIACqmMlhiIgk1CrFDs0zTcilVlNPXp1Hr634caPysUCpSKxHj3i21Y+utJFApVpS9JaTUX1CkUCkikMi5gZhsCqWMznakZtG+TEEIIsRT2ojUA2NvzYGfHhAqe7gIAuvdsPkvOxq6j1yGTWaYnBLFuNGeTECs0ZPo6yOUK7mdRWdV7Gf++GM3dLhGJkVdYgofP0nHxZjwu3ozHnYfJ3P1lYinEEikcHcz/FvD78Vv4YsVf3M+Bvh5a59TzY45RR1pCCCHEcvIKS7jbPPC4255uymBTx57NlyasAgAInBww5rXOZl4hsTUUbBJiZaQymUagCQBFJVVvyE9K1Wz6c+nmE8jUOs9G3n2ucX9xiRiOnuZ/C1APNAHVXE11XsorpoXF1HiAEEIIsZS8ghKdx7lgs5Iy2idJWWZZE7FtVEZLiJUpLhVrHXuSmIXr95/rOFuFDdTC2jQCwAR5j55lAACaNvAFj8fTOF+fALa6yo9Z8XRzhkDHflJ3V2dmTRRsEkIIIRajntns37MVd1vg7AAnRz7KxFJk5+mewekiqLxfBKmbKNgkxMoUq+2XYC3ZeByjZ21GzJPUCh/Hbtpf/eVotAmuh1KRBBt/vwQAeG94dyya+Qbs7FQBp7BY+3VMrfxVTl0ltIAq2FTfV0pIXXHs2DG888476NKlC8LCwjBq1Cjs2bPH4Jm4z549w44dOzBnzhwMGTIErVu3RqtWrXDy5MkqH5ueno5vv/0Wr732Gjp06ID27dtj8ODB+Oabb5CcnFzl4wkhtQMbbPbq1AzLPx/JHefxeOjVKRgA0GnUUjxJytR6bFXNCUndRMEmIVZGV2aTFR2fpvO4TCZHUXEZeDweghv64dtPh2ncH+jrgamje+Ph39+gW7smAGomsxnzRHO9nUMb6zzP000ZbFJmk9Qxixcvxpw5cxAdHY2uXbuid+/eSEhIwJIlS/Dpp59CJpPp/Vz79u3D0qVLcezYMTx//hwKhaLqBwGIjY3FsGHDsHv3bohEIvTt2xf9+vWDSCTCH3/8geHDh+P27dvG/oqEEBvCBpsLPhwKb08XjfuGvhTK3Y64+ggANJoCOTvS7jxbk5MvxIFTtyESS8z2GhRsEmJl2GCzXYv6eHtIF437ZBVkOoqU2VB3FyfY2dmhR4emaBzkzd3P7pN0FThxWURzZzaLS8tw4PQdjWM92jfVea67MtgsoswmqUNOnTqFvXv3wt/fH0ePHsWmTZuwfv16nD59GiEhIYiIiMDu3bv1fr6WLVti6tSpWL16NSIiItC9e3e9HrdkyRIUFhZi3LhxOHPmDDZs2IANGzbg7NmzGD16NEpKSrBo0SIjf0tCiK3IKyxBTn4xAGgFmgAwckBH7naucm+nekWShLrR2pwN+y7hs2UHcOpKrNleg4JNQqwMW0br5uKENiH1NO6bu+IQ/oq4q/UYtoTWQxm08Xg8vPOm6otmoJ+qfNXd1Yl5jJmziN/9cgIXb8QDABz49ghu5IfX1a6KqqMyWlIXbdq0CQAwZ84cNG3alDvu5+fHBXebN2/Wu5x27NixmDt3LoYOHYrGjXVXEZRXVlaGO3eYi0KffvopHBxUM/YcHBwwa9YsAEBcXBxKS2m+HiG1VeyTNLR/6zukZzNd4f283bTOETg7YvkcprQ2VxmU5hWp9niKJVV3zifWobRMgtgnaUjLYkbO6VkIYxTKdxNiZdjMppuLExc8qvt06X682r0lPN2duflXbJDmoewWBwAfvt0X9+NSkJUr1MhyurkwwaZQx95QU9p1NIq7feLXmWjdrF6F53oog83kjDxs/P0SRg8OQ4CPdtdaQmqL9PR0xMTEwMHBAUOGDNG6v3v37ggMDERGRgbu3r2Lzp3NM07Azs4OfD4fUqlUZ9kt21jMxcUFzs7a70eEkNphx5Fr3G0fTxcInBx0nufj6QoAyClggs2CItVFKLFY/7J/Ylkff/eHRjZT1/dNU6HMJiFWhg0CXQWOXMavvA4jv0fjAf/DjCX7IJPJuVbk6m8WfHt7/Lr4XRz6+QPw7e25425c51fzBZvq8zID/TwQ3NCv0vPZdZeKJPh+00nMWPK72dZLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 1080x360 with 2 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(1, 2, figsize=(15, 5))\n",
|
||||
"ax[0].plot(train_x.cpu(), train_y[1:].cpu())\n",
|
||||
"# ax[0].plot(test_x, test_y, color=palette[7])\n",
|
||||
"ax[0].set_ylabel(\"Exchange Rate\")\n",
|
||||
"ax[0].set_xlabel(\"Time\")\n",
|
||||
"\n",
|
||||
"ax[1].plot(train_x.cpu(), vol.cpu())\n",
|
||||
"ax[1].set_xlabel(\"Time\")\n",
|
||||
"ax[1].set_ylabel(\"Volatility\")\n",
|
||||
"sns.despine()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 99,
|
||||
"id": "a8ec2eee-30f2-489a-a56e-ed4df28d1f62",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Iter 1/500 - Loss: 0.741\n",
|
||||
"Iter 51/500 - Loss: -1.585\n",
|
||||
"Iter 101/500 - Loss: -2.979\n",
|
||||
"Iter 151/500 - Loss: -3.109\n",
|
||||
"Iter 201/500 - Loss: -3.131\n",
|
||||
"Iter 251/500 - Loss: -3.140\n",
|
||||
"Iter 301/500 - Loss: -3.144\n",
|
||||
"Iter 351/500 - Loss: -3.147\n",
|
||||
"Iter 401/500 - Loss: -3.149\n",
|
||||
"Iter 451/500 - Loss: -3.150\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"voltron, lh = TrainVoltMagpieModel(train_x, train_y[1:], \n",
|
||||
" vmod, vlh, vol,\n",
|
||||
" printing=True, \n",
|
||||
" train_iters=500,\n",
|
||||
" k=200, mean_func=\"ewma\", theta=0.0)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 100,
|
||||
"id": "befae0e8",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"vmod.eval();\n",
|
||||
"voltron.eval();\n",
|
||||
"voltron.vol_model.eval();\n",
|
||||
" \n",
|
||||
"with torch.no_grad():\n",
|
||||
" full_samples, pred_vol = Rollouts(train_x, train_y, test_x, voltron, \n",
|
||||
" nsample=50, theta=0.0)\n",
|
||||
" torch.cuda.empty_cache()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 101,
|
||||
"id": "efd66509-02a9-45e6-be29-9d38e3db916a",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"torch.Size([50, 150])"
|
||||
]
|
||||
},
|
||||
"execution_count": 101,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"full_samples.shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 102,
|
||||
"id": "5249b0b1",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAA5sAAAFXCAYAAAAs3lr2AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAADiZUlEQVR4nOzdd3hcxdU/8O/dvqtV77JsS+62XLFcwVQDhoRiagyEhPrS8hJ+IRBCEgNvMIQaAoQEMJjqhGYDoRgDobgXWe6WJVuyurRqK21v9/fHekb3bpF2pV0V+3yehwdp9+7dkWxrdOacOSOIoiiCEEIIIYQQQgiJIcVgD4AQQgghhBBCyImHgk1CCCGEEEIIITFHwSYhhBBCCCGEkJijYJMQQgghhBBCSMxRsEkIIYSQmPJ4PKitrYXH4xnsoRBCCBlEFGz2A02mhBBCSLDGxkacc845aGxsHOyhEEIIGUQUbPYDTaaEEEIIIYQQEhoFm4QQQgghhBBCYo6CTUIIIYQQQgghMUfBJiGEEEIIIYSQmKNgkxBCCCGEEEJIzFGwSQghhBBCCCEk5ijYJIQQQgghhBAScxRsEkIIIYQQQgiJOQo2CSGEEEIIIYTEHAWbhBBCCCGEEEJijoJNQgghcSGKIjwez2APgxBCCDnhDdX5loJNQgghcVFWVoYNGzbAYrEM9lAIIYSQE1ZzczM2bNiA+vr6wR5KEAo2CSGExEVjYyMADMnJjxBCCDlRHDx4EABw+PDhQR5JMAo2CSGExJXP5xvsIRBCCCEnLFEUAQBKpXKQRxKMgk1CCCFxxSZBQgghhMSPRqMZ7CEEUQ32AAghhJzYKLMZH59++ilWr16NsrIy+Hw+FBYW4vLLL8eyZcugUES/luxwOPDWW2/hyy+/xLFjx+B2u5Geno6pU6fiF7/4BWbPnh2Hr4IQQkisULBJCCHkpEPBZuw9/PDDePfdd6HVarFgwQKoVCps3rwZjzzyCDZv3oznnnsuqnKqmpoa3HTTTTh27BjS09MxZ84caDQa1NXV4dtvv8WkSZMo2CSEkCFIOscOxTJaCjYJIYTEFQWbsbVu3Tq8++67yMzMxNtvv42CggIAQEtLC66//nqsX78eb7/9Nn7xi19EdD+bzYYbb7wR1dXVuOOOO3DHHXdArVbz59vb29HR0RGHr4QQQkh/eb1e/vFQnG9pzyYhhJC4oj2bsfXPf/4TAHDvvffyQBMAMjIy8NBDDwEAXnnllYh/6XjppZdQXV2NSy+9FHfffbcs0ASA1NRUFBYWxmTshBBCYksabEo/Hioo2CSEEBJXQ3GldbhqbGzE/v37oVarsWTJkqDn586di+zsbJhMJpSWlvZ6P5fLhffeew8AcOutt8Z6uIQQQuLM4/Hwj4fifEtltIQQQuJqKE5+w9WBAwcAAOPHj4dOpwt5zbRp09DU1ISDBw/ilFNO6fF++/fvR0dHB3JzczF27FiUlJTgu+++Q0dHBzIyMrBo0SLMmjUr5l8HIYSQ/vP5fLJtDlarFU1NTcjOzh68QQWgYJMQQkhcUbAZO7W1tQCAvLy8sNfk5ubKru0JOwB89OjR+N3vfoc1a9bInn/xxRdx/vnn44knnggb3BJCCBkcZWVlaGpqkj128OBBCjYJIYScPOx2O9xud9BeQBI9m80GANDr9WGvSUhIAOBf4e6N2WwGAOzYsQNerxc33ngjli1bhpSUFGzfvh0PP/ww1q1bh4SEBDz22GMx+AoIIYTESmCgORTRnk1CCCExJ20KJIoi2traBnE0Jw72fRUEISb3Y1lnj8eDK664Avfffz9GjRqFpKQknHPOOXjxxRchCALWrl2LmpqamLwnIYSQ+Orq6pLt5RxMFGwSQgiJucAOtCwjR/qHZS17+n6yjCa7NpL7AcBVV10V9Py0adNQVFQEn8+HrVu3RjtcQgghg2Dnzp3Yvn37YA8DAAWbhBBC4iAw2HQ4HIM0khPLiBEjAAD19fVhr2lsbJRdG8n9ACA/Pz/kNezxlpaWiMdJCCFkcDmdzsEeAgAKNgkhhMRBYFMgu90+SCM5sUyZMgUAUF5eHjaA37t3LwBg8uTJvd6vqKiIf9ze3h7yGva4wWCIaqyEEEIIBZuEEEJixmazYevWrUFNC7q6uqiUNgZyc3NRVFQEt9uNL7/8Muj5bdu2obGxEZmZmREdWZKdnY0ZM2YAALZs2RL0vNls5setTJ06tZ+jJ4QQEi8ajWawhxASBZuEEEJi5siRI7Db7aioqAAAaLVaZGdnQxRFKsOMkVtvvRUA8NRTT+HYsWP88dbWVjz88MMAgFtuuQUKRfcU//TTT2PJkiV4+umng+532223AfAfc3Lw4EH+uNPpxEMPPYSuri4UFRXReZuEEDKETZkyZUgeUUVHnxBCCImZwL2aCoUCycnJaGpqosxmjCxZsgTLli3D6tWrcdFFF2HhwoVQqVTYvHkzLBYLFi9ejOuuu072GpPJhMrKSphMpqD7nX322bjxxhvx2muv4corr8SMGTOQkpKCPXv2oLm5GdnZ2XjmmWdi1gGXEEJI7Gk0GqjV6iHXI4GCTUIIITEjzaYB/iM62F6/xsZGTJgwIegaEr2HHnoIs2fPxjvvvINt27bB5/NhzJgxuPzyy7Fs2bKov8f3338/TjnlFLz11ls4ePAg7HY78vLycMMNN+DWW29FWlpanL4SQgghfRG4uKvRaKBSDb3QbuiNiBBCyLAVGOQoFApZY5nGxkbk5eUN9LBOSBdddBEuuuiiiK59/PHH8fjjj/d4zbnnnotzzz03FkMjhBASZ4GN+JRKJdRqtewxURQHvSqFlpcJIYTETOCkJgiCbLV1qLRiJ4QQQoYzr9fLP54xYwYEQQgKNgMD0sFAwSYhhJCYCZzYWKZz9OjRAOSTIyGEEEL6xuVyAfAfS5WamgoAQWW0FGwSQgg5oXg8HtnnLNhk/x8KEx8hhBAy3LFKIWkHWspsEkIIOaEFZi4Dg03KbBJCCCH9x7rOarVa/lhgZnMozLkUbBJCCImZcMGmUqkEMDRWWQkhhJDhjgWb0sxmYLB56NAh1NXVDei4AlGwSQghJGYCy2hZkElltIQQQkjssD2b0swmm3OZzs5OlJeXD+i4AlGwSQghJGbCNQhiE2C4kh6XyxUUqBJCCCEkNDafSgPMcGcsB57JOZAo2CSEEBIz4YLNnjKbPp8PmzZtwoYNG+I/QEIIIWSYaW1txZ49e3DkyBH+GJtPpQFmYGaTGczFXAo2CSGExEzg6mkkZbSsFIgQQgghcqIoYu/evWhra0NNTQ1sNhuA0MFmuMym2+2O/0DDUPV+CSGEENI7URT7VEYrDVBFUYQgCHEcJSGEEDJ8sEZAjM1mw/79+2G1WgFEltl0uVwwGAzxG2QPhmSw6Xa7sWPHDnz//fcoKSlBfX09Ojo6kJqailmzZuHaa6/FvHnzBu1+hBBCIsMCx54ym9IAlIJNQgghpJvdbpd9fuDAAdlcKp0zKbMZoe3bt+OGG24AAGRmZqKoqAh6vR5HjhzBunXrsG7dOtxxxx24++67B+V+hBBCgklLetjHkQSb0sd8Pl/YyZIQQgg52QQGm+EqiIDwmU0KNgMIgoDzzz8f119/PYqLi2XPff7557j33nvx97//HfPmzcP8+fMH/H6EEEKCsXJYabDJ9FRGK30s0qNRPB7PoJYFEUIIIQOB7dEMRxpshqsMGsxgc0guHy9YsAB/+9vfggJDALjwwguxdOlSAMAnn3wyKPcjhBASLDCbKSXNbAY2EepLsLljxw5s27aN71khhBBCTkQWi6XH5yMJNqkbbZSmTJkCAGhqahqS9yOEkJORNLMZSBAEPgkGBpSBZbSRYA0T2tvbg56rqKjAjh07gkqPCCGEkOFEFMWogs1wKNiMUlVVFQD//suheD9CCDkZ9ZTZBLpLaQMDSmlmM1SZbSTvKVVbWwuLxYLy8vKo7kUIIYQMJR6Pp9d5kYLNGDOZTFizZg0A4Lzzzhty9yOEkJNVqDO/VKru1gDhmgT1pYw2kuujDVwJIYSQoSSSObGnYJMl0gYz2BySDYLC8Xg8+O1vf4uuri4sWLAAZ5999pC6HyGEnKxaWlqwb98+AP7M5uTJk9Ha2ors7Gx+TbgmQX0pow13feB+UEIIIWS4imTRNLCaiDXpmz17NkRRhMlkomAzUsuXL8fmzZuRm5uLJ598csjdjxBCTlaHDh3iHysUCmRnZ8sCTfY4EN/MpnRCpcwmIYSQ4YzNYwaDAWlpaaitrZU9L+2HwMyfPx8OhwOJiYm8ky2V0Ubgz3/+Mz744ANkZmZi1apV/d5fGev7EULIyUxaLhtuz2asgk1p9jKW+z8JIYSQoYTNcSqVCllZWUHPhyqh1Wg0SEpK4q8DKNjs1eOPP4633noLaWlpWLVqFQoKCobU/Qgh5GSn1Wr5x+H2j0jLaG02G59Eoy2jlQabgROo9HO73Y6GhgbqSksIIWRYYoumSqWSz6FSvTUHkgabg7XNZMgHm0888QRef/11pKSk4PXXX8e4ceOG1P0IIYTIM5vhAkY2KZpMJmzbto2X3kYbbEqvCQw2pdlMu92Or7/+GiUlJRF8BYQQQsjQIYoi9uzZA8A/f/Yl2FQoFFAoFBBFMeptKrEypIPNp556CitXrkRycjJef/11TJo0aUjdjxBCiJ806Ost2KyvrwcANDc3B10fbbAZWCorHYfNZoPX64Xb7R60SZYQQgjpC6vVyj9WKpVhz7DuzWCX0g7ZYPOvf/0rXnnlFSQlJeG1117DlClTen3N008/jSVLluDpp5+Oyf0IIYREJpK9kqFWZYGe92CGIr0msCzI4/HA4/GgsbERXV1d/PGysrJe70sIIYQMFdJAUhCEkHPocAg2h2Q32m+++QYvvfQSAGDUqFF4++23Q143ZswY3Hrrrfxzk8mEyspKmEymmNyPEEJIZCIJNgNXZUM1DIok2AzVUEgLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 1080x360 with 2 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt_idx = 3\n",
|
||||
"fig, ax = plt.subplots(1, 2, figsize=(15, 5))\n",
|
||||
"ax[0].plot(train_x.cpu(), train_y[1:].cpu())\n",
|
||||
"ax[0].plot(test_x.cpu(), full_samples[:10].T.exp().cpu(), color='gray', alpha=0.5)\n",
|
||||
"# ax[0].plot(test_x, test_y, color=palette[7])\n",
|
||||
"ax[0].set_ylabel(\"Wind Speed\")\n",
|
||||
"ax[0].set_xlabel(\"Time\")\n",
|
||||
"\n",
|
||||
"ax[1].plot(train_x.cpu(), vol.cpu())\n",
|
||||
"ax[1].plot(test_x.cpu(), pred_vol[:5].T.cpu(), color='gray', alpha=0.5)\n",
|
||||
"# ax[1].plot(test_x, pred_vol.mean(0))\n",
|
||||
"ax[1].set_xlabel(\"Time\")\n",
|
||||
"ax[1].set_ylabel(\"Volatility\")\n",
|
||||
"sns.despine()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 108,
|
||||
"id": "5811eb04-d29d-4c37-bb49-c3a124250e56",
|
||||
"metadata": {
|
||||
"tags": []
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAA5sAAAFXCAYAAAAs3lr2AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAAD01ElEQVR4nOydd3xb1d3/P/dqy7K87XjEibP3hgz2DrQBElbDLFDSlrbw9Fd2n5bxUEjZtIynhUCAQAqFJwHaQBiBQiCT7GU7ieM95CnJ2tL9/aGc43s1bFmWPJLv+/XihXTvueceOba+93O+S5AkSQJBEARBEARBEARBJBBxoBdAEARBEARBEARBnHiQ2CQIgiAIgiAIgiASDolNgiAIgiAIgiAIIuGQ2CQIgiAIgiAIgiASDolNgiAIgiAIgiAIIuGQ2OwDPp8PNTU18Pl8A70UgiAIghg0kH0kCIIgABKbfaKhoQHnnXceGhoaBnopBEEQBDFoIPtIEARBACQ2CYIgCIIgCIIgiCRAYpMgCIIgCIIgCIJIOCQ2CYIgCIIgCIIgiIRDYpMgCIIgCIIgCIJIOCQ2CYIgCIIgCIIgiIRDYpMgCIIgCIIgCIJIOCQ2CYIgCIIgCIIgiIRDYpMgCIIgCIIgCIJIOCQ2CYIgCIIgCIIgiIRDYpMgCIJICpIkwefzDfQyCIIgCOKEZ7DaWxKbBEEQRFIoLS3Fxo0bYbfbB3opBEEQBHHC0tTUhI0bN6Kurm6glxIGiU2CIAgiKTQ0NADAoDR+BEEQBHGicPDgQQBAWVnZAK8kHBKbBEEQRFIJBAIDvQSCIAiCOGGRJAkAoFKpBngl4agHegEEQRDEiQ0zgkRi+fjjj7F69WqUlpYiEAigpKQEV1xxBZYuXQpR7P1essvlwltvvYVPP/0UlZWV8Hq9yMrKwpQpU3DTTTdh9uzZSfgUBEEQRKLQ6XQDvYQwSGwSBEEQSYU8m4nn4YcfxjvvvAOdTof58+dDrVZj06ZNeOSRR7Bp0yY8//zzvdrhrq6uxq233orKykpkZWXhlFNOgVarRW1tLTZs2IAJEyaQ2CQIghiEyDd0tVrtAK4kMiQ2CYIgiKRCYjOxrF+/Hu+88w5ycnKwatUqjBw5EgDQ3NyMG2+8EZ9//jlWrVqFm266Kab5HA4HbrnlFlRVVeH222/H7bffDo1Gw8+3tbWhvb09CZ+EIAiC6CtysSkIwgCuJDKUs0kQBEEkFRKbieVvf/sbAOCuu+7iQhMAsrOz8dBDDwEAXnnllZh/7i+//DKqqqpw+eWX484771QITQDIyMhASUlJQtZOEARBJBZ5y5PBaG9JbBIEQRBJhXI2E0dDQwP2798PjUaDhQsXhp0/9dRTkZeXB4vFgl27dvU4n8fjwXvvvQcAWLZsWaKXSxAEQSQZv9/PXw9GsUlhtARBEERSGYzGb6hy4MABAMDYsWOh1+sjjpk6dSoaGxtx8OBBzJo1q9v59u/fj/b2duTn52P06NHYsWMHvv76a7S3tyM7OxtnnHEGZs6cmfDPQRAEQSQGuWfTZrPB5/NBrR48Em/wrIQgCII4ISGxmThqamoAAAUFBVHH5OfnK8Z2B+vJNmLECNx3331Ys2aN4vyLL76Iiy66CE888URUcUsQBEEMDDabDfv27VMcO3DgAKZNmzZAKwqHxCZBEASRVEhsJg6HwwEAMBgMUcekpKQAADo7O3ucr6OjAwCwfft2+P1+3HLLLVi6dCnS09Oxbds2PPzww1i/fj1SUlLw+OOPJ+ATEARBEIli586dYTa2tbV1gFYTGcrZJAiCIJKK0+mE1+sd6GWcELD810RVHGQPKT6fD1deeSXuvfdeFBcXw2w247zzzsOLL74IQRCwdu1aVFdXJ+SeBEEQRGIYCpu5JDYJgiCIhCMvCiRJ0qDbaR2qMK8l83BGgnk02dhY5gOAq6++Ouz81KlTMXnyZAQCAWzZsqW3yyUIgiAGgJ07d6K+vn6glwGAxCZBEASRBEIr0HYnjojYKSwsBADU1dVFHdPQ0KAYG8t8AFBUVBRxDDve3Nwc8zoJgiCIgaOjowOlpaUDvQwAJDYJgiCIJBAqNl0u1wCt5MRi0qRJAIDy8vKoP9O9e/cCACZOnNjjfJMnT+av29raIo5hx41GY6/WShAEQRAkNgmCIIiEE5pH4nQ6B2glJxb5+fmYPHkyvF4vPv3007DzW7duRUNDA3JycmJqWZKXl4fp06cDADZv3hx2vqOjg7dbmTJlSh9XTxAEQZxskNgkCIIgEobD4cCWLVvQ2NioOG6z2SiUNkEsW7YMAPDUU0+hsrKSH29pacHDDz8MALjtttsgil0m/umnn8bChQvx9NNPh833i1/8AkCwzcnBgwf5cbfbjYceegg2mw2TJ0+mfpsEQRCDmEmTJiWseFwiodYnBEEQRMI4cuQInE4nDh8+DADQ6XRIT09HY2MjmpubUVxcPMArHPosXLgQS5cuxerVq7Fo0SIsWLAAarUamzZtgt1ux/nnn4/rr79ecY3FYkFFRQUsFkvYfOeeey5uueUWvPbaa7jqqqswffp0pKenY8+ePWhqakJeXh6eeeaZQfkQQxAEQQQxGo1ISUmB3W4f6KUoILFJEARBJIzQXE1RFJGWlobGxkbybCaQhx56CLNnz8bbb7+NrVu3IhAIYNSoUbjiiiuwdOlShVczFu69917MmjULb731Fg4ePAin04mCggLcfPPNWLZsGTIzM5P0SQiCIIh4CLW3Op0OGo1mgFYTHRKbBEEQRMIIFTmCIPDCMg0NDRg3blyvhRARmUWLFmHRokUxjV2+fDmWL1/e7ZgLLrgAF1xwQSKWRhAEQSSZ0NoIarUaarVS2kmSNOBRKWTxCYIgiIQRKiRFUVRUMWVtOQiCIAiCiB+/389fT5gwAYIghHk2Q72fAwGJTYIgCCJhhO6gCoIArVbLd1vdbvdALIsgCIIgTii8Xi+AYK7msGHDACBMbMoF6UBBYpMgCIJIGKFhPczTOWLECACDw/ARBEEQxFCH9VrW6/X8WGgYbahNHghIbBIEQRAJw+fzKd4zscn+PxgMH0EQBEEMdZjY1Ol0/Bh5NgmCIIgTmlDDFio2B4PhIwiCIIihTiyezV27dqGioqJf1xUKiU2CIAgiYUQTmyqVCgB5NgmCIAgiEbAaCHLPJrO1DI/Hg8rKyn5dVygkNgmCIIiEERpGywwfhdESBEEQROJg9lQuMEPFJmMgq9JSn02CIAgiYUQrEMQMYLQwWo/HA1EUw0KACIIgCOJkp6KiApWVlUhNTcWsWbMgCAK3t/KWY9H6WHu9Xmi12n5Zayjk2SQIgiASRjSx2Z1nMxAI4Pvvv8fGjRuTv0CCIAiCGEL4/X4eCmuz2dDZ2YnOzk54PB4ASoEZzbPJxg4EtIVMEARBJIzQUJ1YwmgH0ggSBEEQxGDG6XQq3jc3N+PYsWP8fSxik/XkHAjIs0kQBEEkBEmS4gqjlQvUgcwrIQiCIIjBRqjYlAtNILYwWvJshuD1erF9+3b85z//wY4dO1BXV4f29nZkZGRg5syZuO666zB37twBm48gCIKIDUEQAHTv2ZQLUEmS+DUEQRAEcbLjcDi6PT/YPZuDUmxu27YNN998MwAgJycHkydPhsFgwJEjR7B+/XqsX78et99+O+68884BmY8gCIIIR16sgL2ORWzKjwUCgag7swRBEARxsmG327s9L7eZ0TZrQyvF9yeDUmwKgoCLLroIN954I+bMmaM4t27dOtx111146aWXMHfuXMybN6/f5yMIgiDCYSGwcrHJ6C6MVn4s1tYoPp8PHo8HRqMx3uUSBEEQxKDHZrN1ez6WDVrK2Qxh/vz5+Mtf/hImDAHgkksuweLFiwEAH3300YDMRxAEQYQT6s2UI/dshuZlxiM2t2/fjq1bt6KzszPe5RIEQRDEoMbv98PlcnU7JhaxOZCezUEpNnti0qRJAIDGxsZBOR9BEMTJiNyzGYogCFyEhgrK0DDaWGDGt62tLezc4cOHsX379rCiCgRBEAQxlIjWm1oOic0kwKow5eTkDMr5CIIgTka682wCXaG0oYJSbkxjMayR7imnpqYGdrsd5eXlvZqLIAiCIAYTfRGbo0ePxrRp0wBQzmavsFgsWLNmDQDgwgsvHHTzEQRBnKzICwQx1OouMxOtSFA8YbSxjO+tcCUIgiCIwQSzYyqVCoIghIlGedQQY+rUqWhtbUVRURFPNaGczRjx+Xy4++67YbPZMH/+fJx77rmDaj6CIIiTlebmZmzbtg1A0PhNnDgRubm5yMvL42OiFQmKJ4w22njq00kQBEGcKDAbZzQauZdSTiSvZlZWFsaOHQtBEKDRaACQZzNmHnzwQWzatAn5+fl48sknB918BEEQJyuHDh3ir0VRRF5enkJosuNAcj2bcoNKnk2CIAhiKCP3bEYSlj3la7LoIsrZjIFHH30U77//PnJycrBy5co+51cmej6CIIiTGXm4bLSczUSJTbn3MpH5nwRBEAQxWJAkiW/kqlQqHh0kpyexKYoiBEFAIBDo9WZuohgSYnP58uV46623kJmZiZUrV2LkyJGDaj6CIIiTHZ1Ox19HM37yMFqHw8ENX2/DaOViM3S3Vv7e6XSivr6eqtISBEEQQw6n0wm32w0gutiMtrkrP882gwcqb3PQi80nnngCr7/+OtLT0/H6669jzJgxg2o+giAIQunZjCYYmQi1WCzYunUr37HtrdiUjwkVm3JvptPpxBdffIEdO3bE8AkIgiAIYvAg31gVBCFqW7GeGOi8zUEtNp966imsWLECaWlpeP311zFhwoRBNR9BEAQRRG7EehKbdXV1AICmpqaw8b0Vm6GhsvJ1OBwO+P1+eL3eAQsfIgiCIIh4kItNv98fl2cTGPi8zUErNp977jm88sorMJvNeO211zBp0qQer3n66aexcOFCPP300wmZjyAIgoiNWHIlIxlLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 1080x360 with 2 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"\n",
|
||||
"fig, ax = plt.subplots(1, 2, figsize=(15, 5))\n",
|
||||
"ax[0].plot(train_x.cpu(), train_y[1:].cpu())\n",
|
||||
"ax[0].plot(test_x.cpu(), full_samples[:10, :].T.exp().cpu(), color='gray', alpha=0.5)\n",
|
||||
"ax[0].plot(test_x.cpu(), test_y.cpu(), color=palette[7])\n",
|
||||
"# ax[0].plot(test_x.cpu(), full_samples.exp().mean(0).cpu())\n",
|
||||
"ax[0].set_ylabel(\"Exchange\")\n",
|
||||
"ax[0].set_xlabel(\"Time\")\n",
|
||||
"\n",
|
||||
"ax[1].plot(train_x.cpu(), vol.cpu())\n",
|
||||
"ax[1].plot(test_x.cpu(), pred_vol[:10].T.cpu(), color='gray', alpha=0.5)\n",
|
||||
"# ax[1].plot(test_x.cpu(), pred_vol.mean(0).cpu())\n",
|
||||
"ax[1].set_xlabel(\"Time\")\n",
|
||||
"ax[1].set_ylabel(\"Volatility\")\n",
|
||||
"sns.despine()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "17ca27d2-7b36-405a-b7a5-5dfb94becd59",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,221 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "e0bf2b78",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Warning no robinhood utils.\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import voltron \n",
|
||||
"import seaborn as sns\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import torch\n",
|
||||
"import pandas as pd\n",
|
||||
"\n",
|
||||
"import argparse\n",
|
||||
"import datetime\n",
|
||||
"import gpytorch\n",
|
||||
"from voltron.train_utils import LearnGPCV, TrainVolModel\n",
|
||||
"from voltron.models import VoltMagpie\n",
|
||||
"\n",
|
||||
"import copy\n",
|
||||
"\n",
|
||||
"sns.set_style('white')\n",
|
||||
"# style.use('whitegrid')\n",
|
||||
"palette = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"sns.set(palette = palette, font_scale=2.0, style=\"white\", rc={\"lines.linewidth\": 2.0})"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "a1d5a6c3-6d89-4ba2-ba06-439b3fc23886",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"dpath = \"/home/greg_b/DATA/autoformer/exchange_rate/exchange_rate.csv\"\n",
|
||||
"dat = pd.read_csv(dpath)\n",
|
||||
"dat = dat.drop(['4', '5'], 1)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"id": "638d226b-75ee-409f-8b0d-2e5bf1d541da",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAA0EAAAIjCAYAAADFthA8AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAABcSAAAXEgFnn9JSAAEAAElEQVR4nOzdd3hkZd0+8PucM2d6yaT3Ldnel957RxEpFkQFBMH3BVHxVdSfDSygomJ5UVH0FZReFOl1F1hYyu6yfbMtm94nyfSZU35/zCabackkmUnZuT/XxXWRkzlnzmySmXOf53m+X0HXdR1ERERERER5QpzqEyAiIiIiIppMDEFERERERJRXGIKIiIiIiCivMAQREREREVFeYQgiIiIiIqK8whBERERERER5hSGIiIiIiIjyCkMQERERERHlFYYgIiIiIiLKKwxBRERERESUVxiCiIiIiIgorzAEERERERFRXjFM9QlQ5k4//XT09vbCZDKhurp6qk+HiIiIiGjCmpubEQ6HUVhYiNdee21SnpMhaAbp7e1FKBRCKBRCf3//VJ8OEREREVHW9Pb2TtpzMQTNICaTCaFQCGazGXV1dVN9OkREREREE7Z3716EQiGYTKZJe06GoBmkuroa/f39qKurwxNPPDHVp0NERERENGGXXHIJtm3bNqnLPXIegtrb21FfX4+2trahKVwulwvFxcVYvnw5SktLc30KY7Z3715s27YNXV1diEQisNlsqK2txapVq1BQUDDVp0dERERERBOQ9RDU3d2NV199FW+//TbWr1+Pnp6eER9fW1uLyy67DJ/4xCfgdruzfToZi0ajePDBB3H//fejsbEx5WMkScKJJ56I66+/HkcdddQknyEREREREWVD1kLQli1b8POf/xzvvfceNE3LeL/Gxkb88pe/xJ/+9Cf8v//3//Dxj388W6eUsYaGBtx0002or68f8XGqqmLt2rVYu3YtrrzyStx6662QZXmSzpKIiIiIiLIha32Ctm7divXr148pAA3n8/lw66234rbbbsvWKWVkz549+NSnPjVqAEr0wAMP4Oabb4aiKDk6MyIiIiIiyoWcrgmaNWsWTjjhBBxzzDGoq6tDUVERTCYTurq6sHHjRjzyyCPYtGlT3D7/+Mc/UFhYiBtvvDGXpwYgFry+9KUvwePxxG1fuXIlrrrqKixduhQulwttbW14+eWX8cADD6Cvr2/oca+88gp++ctf4hvf+EbOz5WIiIiIiLIj6yHIYDDgvPPOwyc/+Ukcc8wxKR/jcDgwd+5cXHrppXjooYfwox/9CNFodOj799xzD8477zzMmzcv26cX53e/+13S+p+rrroKt956KwRBGNpWUFCAxYsX4/LLL8e1116L3bt3D33vvvvuw0UXXYRFixbl9FyJiIiIiCg7sjYdThRFXHjhhXjmmWdw1113pQ1AiT71qU/hhz/8Ydw2RVHw+9//PlunllJHRwf++c9/xm0766yz8K1vfSsuAA1XXl6Oe++9FzabbWibruu4++67c3quRERERESUPVkLQZdddhl++ctfYvbs2WPe99JLL00KTWvWrEEkEsnS2SX785//jHA4PPS12WzG97///VH3q6iowJe//OW4ba+++ip27tyZ9XMkIiIiIqLsy1oIkiRpQvtffPHFcV/7/X7s2rVrQsdMR9d1vPDCC3HbzjvvvIx7Fl122WWwWq1x25577rmsnR8REREREeVO1kLQRKVaU9PV1ZWT59qyZQs6OjritiWGsJHY7XacddZZcdteffXVbJwaERERERHl2LQJQWazOWlbMBjMyXOtXbs27mtZlnHEEUeM6RiJ0/fq6+vR2to64XMjIiIiIqLcmjYhKFWAKCwszMlzJfYEWrp0KUwm05iOkSo0jbXXEBERERERTb5pE4LefffdpG2zZs3KyXPt27cv7uu5c+eO+RizZ8+GwRBfYTzxuERERERENP1MixCkqiqefvrpuG0LFixAZWVlTp6roaEhbtt4nkeSJJSUlMRtYwgiIiIiIpr+pkUIevTRR9HW1ha37YILLsjJc3k8nrjGrECs/894VFRUxH2dWGyBiIiIiIimnykPQe3t7fjFL34Rt62goACf+cxncvJ8gUAgaZvdbh/XsRL3S3VsIiLKDlXT4Qsr0HR9qk+FiIhmOMPoD8kdRVFwyy23wOv1xm2/5ZZb4HQ6c/KcqYJKqsp0mUjcjyGIiCg3QlEVW1r8UDQdsiTAZTFA1wGHWUK50whBEKb6FImIaAaZ0hD0k5/8BO+//37ctlNOOQWf+MQncvacqYLKWCvDpduPIYiIKDeaPGEoWmwEKKrq6PbFpjX3+KPo8UexuNwGSWQQIiKizEzZdLj7778f//jHP+K2lZeX484778zp8+opplGM9w5iqmMREVF26bqOgaCS9vvekIp3GwbgDaV/DBER0XBTEoKeffZZ/OQnP4nb5nA48Mc//jFnvYEG2Wy2pG2hUGhcx4pEInFfW63WcR2HiIjSCys6IuroN532d+emwTYRER1+Jj0ErV27Ft/4xjegadrQNrPZjD/84Q9YtGhRzp8/VVAZbwhK3I8hiIgo+zId4fFHNISi2ugPJCKivDepIej999/Hl7/85bgS1bIs4+6778ZRRx01KeeQKqj4fL5xHStxP4YgIqLs6/JFR3/QQb3+zB9LRET5a9IKI2zbtg033HADgsFD0xVEUcSdd96J0047bbJOA263G7IsxwWx9vb2cR0rcb/S0tIJnRsRHb4G1xBOtIpZf1BBS18YogDUFpphNUojPr69P4xufxQOk4QqtxmGGVQ8IBTVsLPdj2Ca0R2jJMBqlNA3bL1Qjz+KyoLxFbshIqL8MSkhaPfu3bjmmmuSSmH/8Ic/xIUXXjgZpzBEkiTMmjULe/bsGdrW2to65uOoqorOzs64bXV1dRM+PyI6fEQUDbIkoNMbxYHeECQBqCuxoMAqj+t47f1h7O85NA03EPFjVY0D4sFgFYyoCEY1WI0SzLIIjz869HhvSIU3rM6YKmq6rmN/TzApAEkCsKrGAVXTYZZF9AWVuBDkC6vwR1TYRgmHRESU33IeghobG3H11Vejr68vbvs3v/nNnJbCHkldXV1cCNq3b9+Yj3HgwAEoSvw89blz50743Iho5tN1HXu6gkNlnAepAPZ1B7G6xjDmEaFmTwhNnnDctrCio7UvjGq3Gb3+KHZ1HCrTb5HFpADhDalo6489frryBKLY2xVENE0hhFlFFhgNh2ZyuywGyJIQ9/iOgQjmFltyfq5ERDRz5XRNUFtbG6666ip0dXXFbb/ppptwzTXX5PKpR7RgwYK4r7dt24ZwOJzm0al98MEHox6XiPLTQEhNCkCDwoqO7W1+7OrwwxfObMF/OKqh2ZP6ParJE8aOdj/2dMb3KUs3haxnGq+ZCUZU1HcE0gagCpcRpY74UTRREFDqMMZt6/VH2cKAiIhGlLMQ1N3djauuugotLS1x26+++mrceOONuXrajJxyyilxX0ejUWzcuHFMx3jvvffivl6wYAEqKysnfG5ENPONVs1sIKSi169ga4sf3pAy6gV7f0jBSI/oCyjIoII0ACAY0aBNs4Cg6Tp6/VFsavZBS3NqhTYDZhdZUo6gldjjg1FUzaykNhER5a+chKD+/n5cffXVaGhoiNv+yU9+ErfeemsunnJMli9fjrKysrhtTz31VMb7+3w+vPTSS3HbzjjjjGycGhEdBtKNwiTSAWxt9WNXRwDBiIpuXwShqAZfSIkLKr6QmrVz0xELQtOFpuvY0eaPm8qXSrHdmPZ7ZlmElJCN9nYF0dYfhsIwREREKWQ9BPn9flx77bWor6+P237RRRfhBz/4QbafblwEQcC5554bt+35559PmraXzhNPPIFAIP4D+7zzzsva+RHRzDbWkOEJKNjU7MPuziA2NnmxpdWPDw544QnEpq4lTpsrsctwWw1It6rIbpJQ4TJCTkwGB/kj2QtVE9XQE8LAKCHPZhRRaE2/hFUQBNhM8YUQ+oMKGnpCqO8McGocERElyWoICofDuOGGG7B58+a47eeccw7uuOMOiGL2B54WLlwY999nP/vZjPa79tprYTIdKqMaDAZx++23j7pfe3s77r777rhtZ5xxBhYvXjy2EyeinOnyRvD+gQFsbPRiIKRA0/RJuxDWdR3B6MRDhqLp2NkeQEtfGP6EUFXiMGJRuQ1HzXbCnSIczCm2YHaRBUfWOnDULEfSYwLhQ+enaDra+sNo8YQQUSZ3hMgbUtAxEBnxMWZZxKJy26iFJBJD0KD+oJLxyBwREeWPrKUSRVFw88034913343bfsopp+Cuu+6CJE2vcqVlZWW44oor4ra98MILuPPOO9NeLHV0dOC6666La5IqCAJuvvnmnJ4rEWVOUXXs645VFwspGra1+rG+YQDr9w9gZ7sfoRxfEEcUPWldy5IK27iP19gbStpmP3jBbxAFLCyzYl6pBS6LARZZxNxiy9D3BUGALIlJAcE3bCRoX1cQDT0hNHrC2N7mh5ZuUc4oVE1HsyeExt7Mw1RitbtBxXYZi8qtqCuxYGWVPa4aXDr2NCEIiK2ZIiIiGi5rJbJ/+9vf4rXXXos/uMGAWbNm4Te/+c24jrl06VKcf/752Ti9lG688Ua8/PLLaGpqGtp23333YcOGDbjqqquwbNkyOJ1OtLW14eWXX8YDDzwAj8cTd4yrr74aixYtytk5EtHY+CNqysX1OmLTznxhH46ocUDMUa+cQMIokCQATrMEu0mCLzzxESKbSYrr8yMIAkrsRpSMsGbGntAzxxdSoag6VF2PqxYXjGpY3zAAo0GASRJR5jSixJH+uIOiaixsDo649PqjWFltH3H0xhOIoj8YH05sJglVBSYUWsdeQtxhSv9x5gmwgSoREcXLWgjq6OhI2qYoCu6///5xH/PjH/94TkOQ3W7HPffcgyuvvDKuj9GmTZvwla98ZdT9zzjLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 900x600 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.figure(dpi=150)\n",
|
||||
"plt.plot(dat.iloc[:600, 1:]);\n",
|
||||
"# plt.axvline(0.8 * dat.shape[0])\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"id": "cc2d8d8a-18b0-49d3-bd39-04defa131b62",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"dat_idx = 2\n",
|
||||
"ntrain = 500\n",
|
||||
"ntest = 150\n",
|
||||
"nstart = 0\n",
|
||||
"ys = dat.iloc[nstart:, dat_idx].to_numpy()\n",
|
||||
"\n",
|
||||
"train_x = torch.arange(ntrain-1).float()/365\n",
|
||||
"train_y = torch.FloatTensor(ys[:ntrain])\n",
|
||||
"test_x = torch.arange(ntrain, ntrain + ntest).float()/365\n",
|
||||
"test_y = torch.FloatTensor(ys[ntrain:ntrain+ntest])\n",
|
||||
"\n",
|
||||
"# if torch.cuda.is_available():\n",
|
||||
"# use_cuda = True,\n",
|
||||
"# train_x, train_y = train_x.cuda(), train_y.cuda()\n",
|
||||
"# test_x, test_y = test_x.cuda(), test_y.cuda()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 5,
|
||||
"id": "a12f5f71-fd26-440b-a820-b5db15a58a05",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAaUAAAEhCAYAAADf879gAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAABkYklEQVR4nO2dd3xT5ffHPzc73btQ9rAFKcieX1CGgiCIAxFREBC+iiI/FQQHCA5AFGSIfJG9FQcgoqACKigbZK8KhZa2tHSPNPP+/kjuzV5t0qTJeb9evkzufXJzGm7yec55znMOw7IsC4IgCILwAwS+NoAgCIIgOEiUCIIgCL+BRIkgCILwG0iUCIIgCL+BRIkgCILwG0iUqoFGo0FmZiY0Go2vTSEIgggISJSqQU5ODvr27YucnBxfm0IQBBEQkCgRBEEQfoPI1wbYQq1W48SJE/jjjz9w6tQpZGVloaioCNHR0WjXrh1GjhyJLl26VOnau3btwtatW3HlyhXodDo0adIETzzxBEaMGAGBgDSaIAjClzD+WNHh77//xpgxYwAA8fHxaNWqFeRyOf79919cvXoVADBx4kRMnjzZrevOnj0bW7ZsgVQqRbdu3SASiXD48GGUl5fjwQcfxOLFiyEUCl2+XmZmJvr27Yt9+/ahfv36btlCEARBWOOXnhLDMOjfvz9GjRqFjh07mp376aefMGXKFHzxxRfo0qULunbt6tI19+7diy1btiA+Ph6bNm1C48aNAQB3797FqFGj8Ouvv2LTpk0YPXq0p/8cgiAIwkX8Ml7VrVs3LFmyxEqQAGDgwIF47LHHAAA//PCDy9dcsWIFAGDKlCm8IAFAXFwcZs2aBQBYuXIldDpd1Q0nCIIgqoVfipIz7r33XgDAnTt3XBqfk5ODCxcuQCwWY8CAAVbnO3fujMTEROTl5eGff/7xpKkEQRCEG9RKUUpPTwegX29yhYsXLwIA7rnnHshkMptjWrduDQC4dOlS9Q0kCIIgqkStE6W8vDxs374dAPDQQw+59JrMzEwAQFJSkt0xdevWNRtLEN4kv6gMhSUVvjaDIPyOWiVKGo0GU6dORWlpKbp164Y+ffq49LqKCv2XXy6X2x0TGhoKACgvL6++oQThAJVagz5jFqP1ox/i7JXb8MMEWILwGbVKlN577z0cPnwYdevWxSeffOLy67gvPcMw3jKNIFymoLgC+UX6yc/AF5dh9x/nfWwRQbgGy7Jen0TVGlH68MMP8e233yI+Ph7r1q1zeT0JMHpBnMdkC85D4sYShLcoKVOYPV/+9UEfWUIQrsOyLM5kluFSjnfDzrVClObNm4eNGzciJiYG69atM0vpdoV69eoBALKysuyO4erXcWMJwluUlivNnp+5nIlKldpH1hCEa2h0LBRqHYoVGmh13vOW/F6U5s+fj7Vr1yIqKgpr165F8+bN3b4Gl0J+7do1VFZW2hxz7tw5AEDLli2rbixBuEBJufU9+Puxaz6whCBcxzRqV6TQ4E6JyiuhPL8WpU8//RSrV69GZGQk1q5dixYtWlTpOnXr1kWrVq2gVquxZ88eq/PHjh1DTk4O4uPj0a5du+qaTRAOsQzfAUBufokPLCEI1zGVn6t3KnD9rgJlSq3H38dvRWnRokVYuXIlIiIisGbNGt7bccSCBQswYMAALFiwwOrchAkTAOiF7ubNm/zx/Px8zJ49GwAwfvx4KspKeJ3SMmtPKa+wzAeWEITr6Gx4RSKh55PH/LL23b59+7B8+XIAQMOGDbFp0yab45o2bcqLDaDfw3Tjxg3k5eVZjR0wYABGjBiBrVu3YvDgwejevTtfkLWsrAz9+vXDs88+650/iCBMKDGsKUklIvRo1wz7j17BXRIlws+x1CQGgEzk+Um8X4pScXEx//j8+fM4f952ymznzp3NRMkZs2bNQocOHbB582YcO3YMOp0OTZs2pdYVRI1SUKzP9Jz8XG80rR+H/UevkKdE+D2WoiQWMl7ZZuOXovT444/j8ccfd/t18+bNw7x58xyOGTx4MAYPHlxV0wii2lzPuAsAaFIvFnHRYQCAuwUkSoR/Yxm8Ewq8s+/TL0WJIAIVlmVx9uptAMA9jRMgNvTvIk+J8HcsM+28JUoUryIIAMu2/oFuIz5BbkEpf2zvoYvo8ORcnL6U4bH3+fngBeTcLUGoXIKm9eMQF2PwlEiUCD/HMnwn9cJ6ElBFUWJZFr/88gvee+89/Pe//7VqjFdRUYHjx4/jxIkTHjGSIKoDy7JYvHE/9h25YvO8TqfD3C/3IiOnED/sP8sff3H2VtzJL8UzU9d4zJazV/Re0pP920MiFiEiVAaJWIhyhQqKSpXH3ocgPI2pJoVJhWgUa7vjQnVxO3yXnp6OSZMmIS0tzW5NOalUinfffRe3bt3Ct99+i1atWnnGWoKoAgdPpuGTNb8BADIPzLE6/8/l2/xjnY7F5Rs5yC8qh0wqglqjtarAUB3uGmretWySCED/3YmLDkNWbjHyCsvQsG6Mx96LIDwJlxIeKRfh3rreK8fmlqdUXFyMMWPG4Nq1a0hJScHkyZMRFhZmNU4oFGLEiBG8R0UQvuL7X//BM1PXOhyz78hl/nHO3RKMe3cThr++2qNiBABfbP0TX/2kjx7ERBm/1PFcskMhVagn/BcufOftutZuidKaNWuQnZ2NXr164dtvv8VLL71kt2ke11bi77//rr6VBFFF5q7c63TMlXRjB+OjZ2/gZlaB1ZjqllNhWRZzvjRWE+Gy7gAgIkzfUqXURvkhgvAXuK+AtxMR3Lr+/v37wTAMpk2bBpHIceSvYcOGkEgkuHXrVrUMJIjq0LJpHbPn81f/Aq1WZ3aszMQjOnPlNmyhUlevnIppAgUAxEUZRSk0RKK3o8KznhlBeBIWNdMCyC1RyszMhEwmQ7NmzVwaHxISQk3zCJ8ikQjNni/Z9Dt2miQzAECpQQxSDOs8AFA3PhI92jXln1coqpeEkHbTvMpIfIyJKMmkAIByBYkS4b/4ZfgOALRa12aMKpUKZWVl1J+I8CkKhXVLiDsWxU/LDaI0/YWH+GN14iLw9cIXUDc+Uj+mmoJx9Fy62fPwUGPYO4w8JaIWwIuSl9/HLVGqX78+1Go10tPTnY79888/odFoXPaqCMIbKJTWHs5HK/aYeT6cp9TqniT+WFGpvpFZqFwvGOXV8JTSbuVi4bp9/PMnHjSvRB8aYvCUKiglnPBfdP4YvnvggQfAsizWrHG8b6OgoAAff/wxGIZB3759q2UgQVQHRaXt5nmTPvoapy9lQKfTocyQYBAeIsXXC8chKkKOF4f3AuAZUfr1b2N23+TneuOTqY+ZnQ8ziNLHq3/B3/9cr/L7EIQ38cvw3ZgxYxAZGYlvvvkGc+fORXZ2ttn5/Px8bN26FUOHDkVGRgYSEhIwYsQIjxpMEO5QYWdD6t6/LmHwxOVY+e1fvOCEyiXo0a4Zzu14FyMf6QQACDGIUkU1wne5+cYkhwlP/QcSsXmSUKhcyj9+6rVVVX4fgvAmNZV959bm2ZiYGCxbtgwvvfQSNmzYgA0bNvDnunTpgpISfayeZVlERkZi2bJlCAkJ8azFRFCi0+mqVMVdoXTcZvyD5T8D0AsSd33T8AQnGNUK32XokxxWvj8SkYb0b1M4T4kg/BluU4RfeUoA0LFjR+zcuRODBg2CSCQCy7JgWRbFxcVgWRZCoRADBw7E999/j9TUVG/YTAQZKrUGDzy/CBPf3+r2azlPafr4/ogMtxYEDnvCwIXvyhRKFBZX4OK/2TbHOeLfW3pRat4w3vZ7GBIdCMIXcL/hrowDvL+mVKUq4UlJSfj000/x0Ucf4dy5c8jLywPLsoiNjUVqaipl3BEeJf12Pq5n3MX1jLvo260FHut7n8teE7emNO7xbnjlmfux/Ks/8dGKPVbj2t/b0ObrI8L0WXIlZZXoP2EpsnKLsX/tZCQ3TrQ53ur9lWpk5BRBKBCgUZLtEkL1E6NduhZBeINruQqUVmrQtkG4w8rfuhpaU6pW6wqpVIqOHTt6yhaCsIlpG+bJc75Bdm4xXhn5gNPXabU6KFUaAIBMKgYAxMeE8+ebN4xHmsGLefrhDjavEROpn2AVFJcjK1fffPLQqX9titKxc+m4djOPX48CgBuZd8GyLBrVj7VaS+Kw3OBLEDVJfrl+4lZaqUFUiNjuOD5852V73ArfvfXWW5g7d67L4+fPn4+3337bbaMIwhTLDLpte0+59LpKlf51MqmYDzn07nwPf/6RB1rzj5Ob2PZ8oiP0a6IFxRX8sV/+umTVaoJlWTz+6peYtmC7WYiPD901sB26A/Rt0R/s3oJ/rlJrHP9hBOEhTMN2zvojaQyVUERCP0oJ3759O3bv3u3y+D179mD79u1uG0UQplgmK8il9mdzpqhU+o3eUonRQ4mNCsP+tZPx9YJxeLxfW/54vYRIm9fgPKXbd4r4Y4dO/Yv+45fyogfArF5eSZmxhh3niTWzs57E8eXskQgP1a9rFZcqHI4lCE+hdaOko1KjH+ytPkocXu886+1FMSLwUVhkvtlL87ZEafA4pBZhs+TGiXz47bNpTyIqQm53jSrWUM2bExeOO/mlOH81Cx1TGwGAmXdkKkr/Glqf20ty4BCLhIiJDEVpuRJlChUcjyYIz6AxqQOpcyJQKo1+rEToXVHy2tV1Oh3y8/Mhl9vPeCIIV+A8pfb3NtA/t7Mh1hIuDCYRC+2OGTagPR7s3tLu+ZhIffgu/Xa+1blTF40daYtMvJuSMuPjQkPYz7TWnT1CZNXfE0UQ7qA2cZV0DjLwWJaF0iBKPvWUysrK+L1HHDqdDtnZ2XZTCFmWRWlpKXbs2AGlUokWLVrYHEcQrsKJEBdKc7b3iIMXJUnVAwIRNvYVcdzMMgpVUYlRiEwFqswgMK7sRfJE9QiCcAeNmSjZH6fU6KBjAZGAgZcdJceitG7dOixbtszsWGFhId8ryRWGDRtLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(train_x, train_y[1:])\n",
|
||||
"plt.plot(test_x, test_y)\n",
|
||||
"plt.ylabel(\"Exchange Rate\")\n",
|
||||
"plt.xlabel(\"Time Index\")\n",
|
||||
"sns.despine()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 6,
|
||||
"id": "cf3dc1fe",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"if torch.cuda.is_available():\n",
|
||||
" use_cuda = True,\n",
|
||||
" train_x, train_y = train_x.cuda(), train_y.cuda()\n",
|
||||
" test_x, test_y = test_x.cuda(), test_y.cuda()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 9,
|
||||
"id": "8cbb21ae-dd4e-4346-8ec3-46552602a276",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/home/greg_b/miniconda3/lib/python3.8/site-packages/gpytorch/utils/cholesky.py:38: NumericalWarning: A not p.d., added jitter of 1.0e-06 to the diagonal\n",
|
||||
" warnings.warn(\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Iter 1/200 - Loss: 12.911\n",
|
||||
"Iter 51/200 - Loss: -0.538\n",
|
||||
"Iter 101/200 - Loss: -0.584\n",
|
||||
"Iter 151/200 - Loss: -0.585\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"with gpytorch.settings.max_cholesky_size(2000):\n",
|
||||
" vol = LearnGPCV(train_x, train_y, train_iters=200,\n",
|
||||
" printing=True)\n",
|
||||
"# vmod, vlh = TrainVolModel(train_x, vol, \n",
|
||||
"# train_iters=1000, printing=True)\n",
|
||||
" "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 10,
|
||||
"id": "17dd53e9-cf5c-4756-afbc-3be069a24ff3",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAA5sAAAFXCAYAAAAs3lr2AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAADKdUlEQVR4nOzdd3hTZfsH8G/apHtP2tIyym4ZZW+RoRUE2YpsEPTngJdXBJVXGQoigoiAqAiylyBLmYrssgQKLVBauuneTZs08/dHck6SJm2TNGnS9v5cl9eVnJxz8lTa5NznuZ/75sjlcjkIIYQQQgghhBATsrH0AAghhBBCCCGENDwUbBJCCCGEEEIIMTkKNgkhhBBCCCGEmBwFm4QQQgghhBBCTI6CTUIIIYQQQgghJkfBZi1IJBKkp6dDIpFYeiiEEEKI1aDvR0IIIQAFm7WSlZWFIUOGICsry9JDIYQQQqwGfT8SQggBKNgkhBBCCCGEEGIGFGwSQgghhBBCCDE5rqUHQAghhDRmJ0+exP79+xEXFweZTIYWLVpg3LhxmDRpEmxs9L8nnJiYiCtXruDhw4eIiYlBcnIy5HI5NmzYgMjIyGqPzcrKwtatW3H16lVkZmZCLpcjICAAvXv3xpw5cxAcHFzbH5MQQkgjZJXBplgsxp07d3Dp0iXcvXsXGRkZKCoqgqenJyIiIjB58mT06tXLqHOb6kudEEIIqa3ly5dj3759sLe3R58+fcDlchEVFYUVK1YgKioKGzZsgK2trV7n2r9/P3bt2mXwGB49eoTp06ejpKQETZo0Qf/+/QEAMTExOHjwIE6ePIlt27aha9euBp+bEEJI42aVwebt27cxc+ZMAICvry/CwsLg6OiIZ8+e4ezZszh79izeffddzJ8/36DzmvJLnRBCCKmNs2fPYt++ffD19cWePXvQvHlzAEBeXh6mTZuG8+fPY8+ePZg+fbpe52vTpg1mz56N8PBwhIeHY8mSJbh161aNx61YsQIlJSWYOHEiPv/8c/B4PACKG79Lly7FkSNHsGzZMpw4ccLon5UQQkjjZJXBJofDwcsvv4xp06ahe/fuGq+dOnUKCxcuxA8//IBevXqhd+/eep3T1F/qhBBCSG389NNPAICFCxey30kA4OPjg2XLlmHq1KnYunUrpk6dqlfmzYQJEwweQ0VFBe7duwcAmDdvHhtoAgCPx8P8+fNx5MgRxMXFQSAQwNHR0eD3IIQQ0nhZZd5onz598P3332sFmgAwfPhwjBkzBgAMusta05c6AGzduhUymcz4gRNCCCF6yMrKQmxsLHg8ns71lD179oS/vz9yc3Nx//59s43DxsYGXK7ivrNcLtd6ncPhAACcnJzg4OBgtnEQQghpmKwy2KxJhw4dAADZ2dl67W8tX+qEEEIIoFgnCQCtW7euMojr2LEjAODx48dmGwePx2MzhDZu3AixWMy+JhaL8d133wEAxo0bxwaehBBCiL6sMo22JsnJyQAU6zn1oe+XenZ2Nh4/fkxFEAghhJhVeno6ACAwMLDKfQICAjT2NZdly5bhrbfewqFDh3D58mWEh4cDAB4+fIiSkhJMmzYNixYtMusYCCGENEz1LtjMzc3F0aNHAQAvvfSSXsdY05c6IQQoF4jg5Ghn6WEQYjHl5eUAUO0aSGdnZwBAWVmZWccSHByM/fv3Y/Hixbh8+TKysrLY18LDw9GjRw+NtZyEkMZFJJZALgfs7epd2ECsQL1Ko5VIJPjoo49QWlqKPn36YPDgwXodZ01f6oQ0drceJqPN8GX4fs8/lh4KIRbDrI+0htTUu3fvYuTIkUhNTcUPP/yAGzduICoqCps3b0ZJSQk++OADbNq0ydLDJIRYgFgixaDp6/HK3E0oF4gsPRxSD9WrYHPp0qWIiopCQEAAvvnmG72Ps6YvdUIau+MXHgAA1mw7jxdnrMfv5+9bdkCEWABzg5O5GaoLc/OT2dccSkpK8N5776GsrAy//PILhgwZAk9PT3h5eWHo0KH45Zdf4ODggC1btrBLWAghjUdOfilSMwvxNCUH735xAAIhBZzEMPUm2Pzyyy9x+PBh+Pr6YseOHXqv1wSs50udEAK4ONmzj+NTcjFv1SELjoYQywgKCgIAZGRkVLkPk87K7GsOFy9eREFBATp37ozg4GCt15s1a4ZOnTpBIpHo1bOTENKw5BSUso//inqCTfsuWXA0pD6qF8Hm6tWrsXv3bnh5eWHHjh0arUv0YS1f6oQQoKhE+6aPSCyxwEgIsRymqnp8fDyEQqHOfR4+fAgAaN++vdnGkZmZCQBwdXWtch83NzcAQFFRkdnGQQixTnmFfACAk4OizsKPB6/gTkyKJYdE6hmrDzbXrFmDX3/9FR4eHvj111/RqlUrg89hLV/qhBCgqFSgte2dZfstMBJCLCcgIABhYWEQi8U4c+aM1uu3bt1CVlYWfH19ERERYbZx+Pn5AQBiY2M12p4wxGIxYmNjAQBNmzY12zgIIdYpp0ARbI4c1BHjhkWgQiTBxr0XLTsoUq9YdbC5du1abNu2De7u7vj111/Rrl07o85jLV/qhBCgUMfM5rnrj5GrlqpDSGMwd+5cAIrvupQU1UxBfn4+li9fDgCYM2cObGxUX9Xr1q1DZGQk1q1bZ5IxDBw4EI6OjsjIyMBXX30FkUi1HkskEuHLL79EZmYm3N3dMWDAAJO8JyGk/mC+m329XDB1VE8AQH4RFdMk+rPaGsbfffcdtm7dCjc3N2zfvp2dnazOunXrcP78eQwbNgwffvihxmtz587F/PnzsXbtWkRERKBZs2YAqv9SJ4SYXlGJ9swmAGTnl8LXq+pUPkIamsjISEyaNAn79+/HyJEj0bdvX3C5XERFRYHP52Po0KGYMmWKxjG5ublISkpCbm6u1vliY2PZ7zMASEhIAACsX78e27dvZ7cfOqRaJ+3t7Y2lS5diyZIl2Lt3L86fP4+wsDAAQExMDHJzc2FnZ4dVq1ZVm2pLCGmYcpVptL5ernB1VvSqLxNUWHJIpJ6xymDz77//xpYtWwAAISEh2LNnj879WrZsyd4ZBqr/EjbmS50QYnpFpZozm13aNcX9J+nIzi9FeGsLDYoQC1m2bBm6deuGvXv34tatW5DJZGjZsiXGjRuHSZMmGXQDlM/nIzo6Wmt7TVVkx4wZgzZt2mDnzp24c+cOrl27BgDw9/fH+PHjMXPmTKOWsBBC6j92ZtPTBa7OigJ/pWUUbBL9WWWwWVxczD6OiYlBTEyMzv169uypEWzWxJRf6oQQ46jPbHp7OKN1Mz/cf5KOnHxKoyWN08iRIzFy5Ei99l29ejVWr16t87VevXohLi7OqDGEhYVhzZo1Rh1LCGmY5HI5niRlAwAC/dzh4kQzm8RwVhlsjh07FmPHjjX4uOq+hBmGfKkTQkyrQiRBuVAErq0NYk58BhsOBxv2/AMAyC2kYJMQQgixFrEJmUhMy4O3hzO6tG8KDhT96vnlIshkMpqkIXqh3xJCSJ3JL1Ks/fBwc4SLkz2cHO3g6+kCQLFmkxBCCCGWJZfLkfQ8H1fvPgMAvNyvA7i2trC1tYGTgx3kcjnKhdrVqwnRxSpnNgkhDdPjZ4p+tq2b+bHb/L0VRUcojZYQQgixvF+PRuHzjX/A2VHRW7NViC/7mquzPcqFIvDLK+DiZG+pIZJ6hGY2CSF15mF8BgCgY+sgdpufMtjMzi+xyJgIIYQQovL5xj8AAGUCRSukJj5u7GvOygCTX667bz0hlVGwSUgDI5PJavW6uUikUpy//hgAEN46kN0e5OcBAHieU6zrMEIIIYTUkRwdPa8DfN3Zx65OVJGWGIaCTUIakDsxKWj9yjLsOBql8/VVP59Bx9ErkZpZUMcjA347cw/Rcc/B49qid+fm7HZ/HzfY2HCQk18KkVhS5+MihBBCiEJiWp7WtgBf7ZnNsnIKNol+KNgkxIolP8/Hv7Gpeu8/ZfEOVIgk+N/3J3W+/sP+yyguFeCrn8+aaoh6i32mSKGdM6EfApWzmQDA49qiiY8b5HI5MnMplZYQQgixFF1BJLPcBVCb2aRgk+iJgk1CrFj/Kevw2vs/6rWesUxQAb7yw9/Jwa7afZ8kZZlkfIbIzlOk5nRUS6FlNPX3AACkZxfW5ZAIIYQQooav7KFpx7Nlt9nxVPVEmV6b5649xr3HaZBKLbM0h9QfVI2WECv06FkmBBWqsuI5+aXw93ar5gjFB786uVwOYYUYb322F2GtA/D+m4PY11Iz6y6ok8vlkEhlbMCsfoeUoZjpTEFGNq3bJIQQQiyFWYv52uDO4HA0aywAgI+nMwDgt7N38dvZu/j4rZfw/uRBdT1MUo9QsEmIFYqcuwkymZx9LqyoeS3jH5di2MflQhEKS8rxODELl+7E49KdeNx7nMa+XiGSQCSWaNytNJcDp/7FR2t/Z5/rCpqZSndUkZYQQgixnDLlzKa7qyOWvTdC6/W3Xx+Anw5dZZ+v/uUcuoaFoG+XlnU2RmK8p8nZiH7yHONfjgCHw6mT96Q0WkKsjEQq1Qg0AaBUjxLjqRmaRX8u30lAVp4qeIu6n6Txelm5qBaj1J96oAmo+mqq83B1BACUlFEpdUIIIcRSmOU4VfXQ9PNyxbwpgzS2TVzwi7mHRUwkcu4mLPj6MP64+LDO3pOCTUKsDNPXSl1CSi5uPkjSsbcKE6hFtA8GoAjyniRmAwCaB3lr3cHSJ4CtrcptVtxdHOCoYz2pq7NiDUgpBZuEEEKIxTA3oqsKNgGgRZCP1rbCknKzjYmYhlQqg0gsBQBc/jehzt6Xgk1CrIyuSnArtpzCuPlbEZuQUeVxJXwBAGD9x+PQvmUTCIRibDlwGQAwdVRPLHtvBGxsVAEnvw56ZCWk5mo8r2rdKRNslvAp2CSEEELM4fr9RJy5+gj/+/4E3vhwG0RiCU5djsH1+4nsPuzMpmPVhQaHvxCGXp2aa2xLTNdumUKsS4pLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 1080x360 with 2 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(1, 2, figsize=(15, 5))\n",
|
||||
"ax[0].plot(train_x.cpu(), train_y[1:].cpu())\n",
|
||||
"# ax[0].plot(test_x, test_y, color=palette[7])\n",
|
||||
"ax[0].set_ylabel(\"Exchange Rate\")\n",
|
||||
"ax[0].set_xlabel(\"Time\")\n",
|
||||
"\n",
|
||||
"ax[1].plot(train_x.cpu(), vol.cpu())\n",
|
||||
"ax[1].set_xlabel(\"Time\")\n",
|
||||
"ax[1].set_ylabel(\"Volatility\")\n",
|
||||
"sns.despine()"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
Vendored
BIN
Binary file not shown.
@@ -0,0 +1,271 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "93e12f63",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"[pyKeOps]: Warning, no cuda detected. Switching to cpu only.\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns\n",
|
||||
"import torch\n",
|
||||
"import gpytorch\n",
|
||||
"import os\n",
|
||||
"import pickle5 as pickle\n",
|
||||
"\n",
|
||||
"sns.set_style('white')\n",
|
||||
"# style.use('whitegrid')\n",
|
||||
"palette = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"sns.set(palette = palette, font_scale=2.0, style=\"white\", rc={\"lines.linewidth\": 4.0})\n",
|
||||
"\n",
|
||||
"# import joypy\n",
|
||||
"import pandas as pd"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "dea9272c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"mt_preds = torch.load(\"finance_mt_preds_3.pt\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"id": "442cc7f8",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"8 8\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"mt_x_list = mt_preds[\"x_paths\"]\n",
|
||||
"\n",
|
||||
"train_splits = list(range(252, 1259, int(252/2)))\n",
|
||||
"eval_splits = list(range(252+int(252/2), 1259, int(252/2))) + [1259]\n",
|
||||
"print(len(train_splits), len(eval_splits))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"id": "c80581c1",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"train_splits.insert(0,0)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 5,
|
||||
"id": "cc856c43",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ind_preds = torch.load(\"finance_ind_preds.pt\", map_location=\"cpu\")\n",
|
||||
"ind_x_list = ind_preds[\"x_paths\"]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 6,
|
||||
"id": "b3eda070",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"times = np.arange(0, int(252/2), 20)\n",
|
||||
"times[-1] -= 2"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 7,
|
||||
"id": "e8163cb1",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"full_data = ind_preds[\"y\"]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 8,
|
||||
"id": "3b5acbf4",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"8"
|
||||
]
|
||||
},
|
||||
"execution_count": 8,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"len(eval_splits)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 9,
|
||||
"id": "a6465c5e",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"\n",
|
||||
"\n",
|
||||
"mt_quantiles_list = []\n",
|
||||
"ind_quantiles_list = []\n",
|
||||
"for i in range(len(eval_splits)):\n",
|
||||
" \n",
|
||||
" torch.cuda.empty_cache()\n",
|
||||
" \n",
|
||||
" train_start = train_splits[i]\n",
|
||||
" train_end = train_splits[i+1]\n",
|
||||
" test_end = eval_splits[i]\n",
|
||||
" test_y = full_data[:, train_end:test_end]\n",
|
||||
" test_y = test_y[:, times].unsqueeze(-2).cpu()\n",
|
||||
" \n",
|
||||
" ind_quantiles_list.append((ind_x_list[i][..., times] > test_y).sum(1))\n",
|
||||
" mt_quantiles_list.append((mt_x_list[i][..., times] > test_y).sum(1))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 10,
|
||||
"id": "6555f54e",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ind_quantiles = torch.stack(ind_quantiles_list) / 100\n",
|
||||
"mt_quantiles = torch.stack(mt_quantiles_list) / 100"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 11,
|
||||
"id": "3a90fb3e",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ind_quantiles = ind_quantiles.reshape(-1, 7)\n",
|
||||
"mt_quantiles = mt_quantiles.reshape(-1, 7)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 12,
|
||||
"id": "03761ff9",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"array([ 1, 21, 41, 61, 81, 101, 119])"
|
||||
]
|
||||
},
|
||||
"execution_count": 12,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"times+1"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 14,
|
||||
"id": "be5f3902",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/Users/wesleymaddox/anaconda3/lib/python3.7/site-packages/ipykernel_launcher.py:14: UserWarning: Tight layout not applied. tight_layout cannot make axes height small enough to accommodate all axes decorations\n",
|
||||
" \n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAnkAAAKNCAYAAABcCn2uAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAABa10lEQVR4nO3dd3hUZd7/8c9kUiCEklDFIEoLAuIKitJRkLUhiugqqKD+aEvQVRbrFuVZuy5WrGBB3BUQpVgBEUXBgqBGlKZBAiYoBFImmcnMnN8fPMlDTICEzNwz58z7dV17PY/Tvt/DGcIn9zn3fbssy7IEAAAAR4mLdAMAAAAIPUIeAACAAxHyAAAAHIiQBwAA4ECEPAAAAAci5B3E7/crJydHfr8/0q0AAADUCSHvILm5uRo8eLByc3Mj3QoAAECdEPIAAAAciJAHAADgQIQ8AAAAByLkAQAAOBAhDwAAwIEIeQAAAA5EyAMAAHAgQh4AAIADEfIAAAAciJAHAADgQPGRbgAAAESn/AKPijxeI7VSkpOU2ijZSK1YQcgDAADVKvJ4Ne/ddUZqXXZOT0JeiHG5FgAAwIEIeQAAAA5EyAMAAHAgQh4AAIADEfIAAAAciJAHAADgQIQ8AAAAByLkAQAAOBAhDwAAwIHY8QIAAERcMGhpR25+2OvE0vZphDwAABBxJV6f3lqVFfY6sbR9GpdrAQAAHIiQBwAA4ECEPAAAAAfinjwgBgUK98ryFBqp5UpuKHfDNCO1cPT4TthH8a+7FSjeb6RWfXd9I3UQHoQ8IAZZnkKVLHvVSK36Z4+S+Ac96vGdsI9A8X6tfeIBI7V6TpxqpA7Cg5AHADAqUBaQJ3uLkVruBo3VoHkLI7WcKCkhTkM6phipleb2GakTSwh5AACjgl6P1j71qJFaZ2TeLBHyjprLV6q8N2YZqZXOqGHIEfIAAEDMSLFK5M8rMFIr0vefEvIAAEDMcHuLVbJ6gZFakb7/lJAHwBHyCzwq8niN1DK5LZKpWa9WmTPvh7IsGdkqS4qt7bJgD7YKeb17967zZ7hcLn366ach6AZANCnyeDXv3XVGapncFsnUrNd6/YaHvUYk+AMBR34vgJqwVchLT0/Xt99+W6fPcLlcIeoGQDRJsUqMzQJMsUokpRqpBcQKUzN5462AAmGvEh1sFfLmzZunGTNm6Nlnn5XL5dJNN92kk08+OdJtAYgCbm+xsVmAJ2TebKQOEEtMzeTtNOmGsNeIFrYKeeXBLjExUU888YTmzJmjyy67TI0bN450awAAAFHFViGvXGZmpjZu3KiVK1fqvvvu07333hvpliDn3vgOAIAd2TLkSdL06dM1dOhQLVq0SGPGjFHnzp0j3VLMc+qN7wAA2JFtQ16zZs30+OOPa+vWrSotLY10OxA3vtuJzxdQQZGZvzduX8C+P2gOIcHtkj9vu5Fa/tJSI+cqMWiFvYbTBYOWkeVaUgKcK9SMrX/29u3bV3379o10G/hf3PhuH/5AQN9t3WWk1hkBB85j85WqZNV8I6Xiew8zcq76n01wqKsSr09vrcoKe52x/dLDXgPOYOuQByD6mVqMltENRFqa22fkakaSKxj2GnAGQh6AsDK1GC2jG4g0t89j5GpGhwlTwl4DzkDIAwA4lqkFdiVG2BB9CHkAAMcytcCuxAgbok9cpBsAAABA6DGSFyGmFg5m0WD7MLqYNJMU8HuWjC2rk8pyLYARhLwIMbVwMIsG24fJxaSZpIDfC1pBY8vq9B1MyANMIOQBAIDYYXDUOtKLwRPyAABAzDA5ah3pxeAJeRFiaguwtGChCrJ/C3sdSXIH/UbqmBYo3CvLUxj2OilWQthrIES4fw2ADRDyIsTUFmAdJkzR2mceD3sdSerr0OUDLE+hSpa9GvY67n4jw14DocH9awDsgJAHIKxMLUbLQrQAUBkhD0BYmVqMloVoAaAyQh4QJRLcLrZfAgCEDCEPiBZsvwQACCG2NQMAAHAgRvKAI/D5AkaWy2CpDABAKBHygCPwBwJGlstgqQwAQCgR8mBL8XEuFWRvMVLLqYs8AwCcjZAHW7K8HhZ5BgDgMJh4AQAA4ECEPAAAAAci5AEAADgQ9+QdJBAISJJyc3PDXqsoN097POFflmNX3m4jdahlnzpOreXEYzJZy4nH5NRaTjwmk7WMHlNungri64e9TqtWrRQfXzXSuSzLYt2G//Xll19q9OjRkW4DAACgxlasWKH09PQqjxPyDlJaWqqsrCw1b95cbrc70u0AAAAcESN5AAAAMYSJFwAAAA5EyAMAAHAgQh4AAIADEfIAAAAciJAHAADgQIQ8AAAAByLkAQAAOBAhDwAAwIEIeQAAAA5EyAMAAHAgQt5B/H6/cnJy5Pf7I90KAABAnRDyDpKbm6vBgwcrNzc30q0AAADUCSEPAADAgQh5AAAADkTIAwAAcCBCHgAAgAMR8gAAAByIkAcAAOBAhDwAAAAHIuQBAAA4ECEPAADAgQh5AAAADkTIAwAAcCBCHgAAgAMR8gAAAByIkAcAAOBA8ZFuAIB5gcK9sjyFRmq5khvK3TDNSC0AwP8h5AExyPIUqmTZq0Zq1T97lETIAwDjCHkAAKPyCzwq8niN1EpJTlJqo2QjtYBoQ8gDABhV5PFq3rvrjNS67JyehDzELCZeAAAAOBAjeQ7HZREAAGITIc/huCwCAEBs4nItAACAAxHyAAAAHIiQBwAA4EC2uCdv1apVIfmcgQMHhuRzAAAAop0tQt6ECRPkcrnq9Bkul0sbN24MUUeIJaa2ADO5/ZfPF1BBUamRWm5fwB4/aGKcyZn4Pp/fSB0g1tniZ+8dd9yhGTNmyOPxSJJat24d4Y4QS0xtAWZy+y9/IKDvtu4yUuuMQMBIHdSNyZn45w/sZqQOEOtsEfKuuuoq9ejRQ+PGjVN+fr5GjBihzMzMSLeF3wkGLe3IzTdSizX5AAA4PFuEPEnq2rWrnn76aV155ZWaOXOmevfurZ49e0a6LRykxOvTW6uyjNRiTT4AAA7PVrNru3fvrmnTpikYDOrOO+9UMBiMdEsAAABRyTYjeeVGjx6tBQsWaNOmTVq6dKkuvPDCSLeECEixSuTPKzBSyyrzGamDujE1QUYyO0kGAI6W7UJeXFyc5s+fL6/Xq8TExEi3gwhxe4tVsnqBkVr1+g03Ugd1Y2qCjGR2kgwAHC3bhTxJSkxMJODVUIpVoiEdU4zUOjbeXK14KyDmbCIWmBqhTLESwl6jXJrbZ+xnRYpVIinVSC0g2tgy5BUWFsrv9ys1tWZ/cffs2SOv1xuTS6+4vcXKe2OWkVodJkwxVqvTpBuM1AEizdQIpbvfyLDXqKjl8xj7WXFC5s1G6gDRyFYh75VXXtGLL76onTt3SpLS0tI0bNgwjR8/Xmlph750MmXKFG3YsIHFkBHVfL6AfjG1BE3AMlJHkixLRpbWaVrG2C6qSnC75M/bbqQW92oi2tgm5E2bNk1Lly6VZf3fP0579uzRSy+9pMWLF+uhhx5Snz59Dvn+g98HRCN/IKB5K8wsRju2X7qROtL/HpeBRXbHDWqruu2LA0fylapk1XwjpbhXE9HGFiFv0aJFWrJkiRo0aKCbbrpJZ555pkpKSrR8+XLNmjVLe/fu1fjx43Xvvfdq2LBhkW63Rop/3a1A8f6w13EH2T4IAIBYZIuQt2DBArlcLj3wwAMaPHhwxePt27fXiBEjNHnyZH3zzTe69dZb5Xa7dd5550Ww25oJFO/X2iceCHudvhOmhL0GAACIPrYIeT/88IOaNm1aKeCVa968uV566SWNHz9eX3zxhW6++WY1bNhQ/fv3j0CnMMaSCopKjZSKDwSN1Eo2eJ+cE1mWVGjoO+H2Bezxw7MWEtwuYzNek1wsZA+YYIufU6WlpTr22GMP+Xz9+vX1zDPPaMyYMfr22291/fXX66WXXlL37t0NdgmTglZQ323dZaRW38FmavUcwj98deEPBIx9J84IOHCSh6/U6Ex8AOFni23NmjdvruzsbHm93kO+Jjk5Wc8884zatGmjkpISjRs3Tlu2bDHYJQAAQPSwxUje6aefrjfffFP33HOP7rrrrkO+Li0tTbNnz9YVV1yh3377TWPHjtUjjzxirlGgDpIS4rhcBuCI8gs8KvIcetAjlFKSk5TaKNlILYSeLULeuHHj9Pbbb2vevHnKysrS0KFDNWTIELVv377Ka9u0aaPnn39e1113nfbs2aOxY8cqKSkpAl0DtePichmAGijyeI0sSyRJl53Tk5BnY7YIee3atdOMGTM0bdo0fffdd9q4caNatmxZbciTpM6dO2vu3LmaOHGisrOz5fF45HKxghYQCaZGKBmdRLUMTtJKLAtIBhZeTguamyTDtnD2ZouQJ0lnnXWW3n//fc2bN09ffPGF2rVrd9jXH3/88XrzzTf11FNP6dVXX1VRUZGhTgEczNQIJaOTqI7JSVr9vR75PlwS9jrxvYexLRxqxDYhT5KaNm2qSZMmadKkSTV6fb169XTjjTdq8uTJ2rx5c5i7A4DQ8/kCRkaiUoMs4YOqTG1LKHH/XzjYKuQdrcTERHXr1i3SbQBArZlaGqbvYEIeqjK1LaHE/X/hEBMhDwAARLdg0DIyahhLI4aEPAAAEHElXp/eWpUV9jqxNGJoi8WQAQAAUDuM5AFALRm9GZ09jRFBJhdpPza+xEitWFoWhpAHALVk8mb0sf3SjdQBqmN6kXYTtWJpWRgu1wIAADgQIQ8AAMCBuFwLALVk8j4ltmuzEUNbqLFwdd3E0gLPhDwAqCXT9ynBHkxtocbC1XUTSws8E/IAAEDMMDkSH+mZvIQ8AAAQM0yOxEd6Ji8TLwAAAByIkAcAAOBAhDwAAAAHIuQBAAA4ECEPAADAgQh5AAAADsQSKgcJBAKSpNzc3LDXKsrN0x5P+FdG35W320gdatmnjlNrOfGYTNZy4jE5tZYTj8lkLaPHlJungvjLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 720x720 with 7 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(7, 1, figsize = (10, 10), sharex=True, sharey=True)\n",
|
||||
"\n",
|
||||
"for i in range(7):\n",
|
||||
" ax[i].hist(ind_quantiles[:, i].numpy(), alpha = 0.5, bins = np.arange(0, 1, 0.05), label=\"Volt\")\n",
|
||||
" ax[i].hist(mt_quantiles[:, i].numpy(), color = palette[-2], bins = np.arange(0, 1, 0.05), \n",
|
||||
" alpha = 0.5, label=\"MT-Volt\")\n",
|
||||
" ax[i].set_ylim((0, 30))\n",
|
||||
" ax[i].set_yticks(ticks=[])\n",
|
||||
" ax[i].set_ylabel(str(times[i]+1))\n",
|
||||
"ax[-1].legend(loc=\"lower center\", bbox_to_anchor=(0.55, -1.6), ncol = 2)\n",
|
||||
"plt.text(0.45, 0.05, \"Quantile\", fontsize=25, transform=plt.gcf().transFigure)\n",
|
||||
"plt.text(0.04, 0.45, \"Time\", fontsize=25, transform=plt.gcf().transFigure, rotation=90)\n",
|
||||
"sns.despine()\n",
|
||||
"plt.tight_layout()\n",
|
||||
"plt.savefig(\"lookahead_predictions_mt.pdf\", bbox_inches = \"tight\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "69a5c15e",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.7.7"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,565 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 71,
|
||||
"id": "93e12f63",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns\n",
|
||||
"import torch\n",
|
||||
"import gpytorch\n",
|
||||
"import os\n",
|
||||
"import pickle5 as pickle\n",
|
||||
"\n",
|
||||
"sns.set_style('white')\n",
|
||||
"# style.use('whitegrid')\n",
|
||||
"palette = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\",\n",
|
||||
" \"#be97c6\", \"#6e4176\"]\n",
|
||||
"sns.set(palette = palette, font_scale=2.0, style=\"white\", rc={\"lines.linewidth\": 4.0})\n",
|
||||
"\n",
|
||||
"# import joypy\n",
|
||||
"import pandas as pd"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 72,
|
||||
"id": "d19933cb",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAjwAAABECAYAAACF4e8fAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAAC3klEQVR4nO3aq24UYRjH4XfZQjblIIAJB8EVrECsQKFQJGjE3gDXUo8FtRJJSNAoqECMAIXjkAUEhXaBDoMqEMoICF+/8PI8cj7zF5PMbzIz6vu+DwCAxA7VHgAAUJrgAQDSEzwAQHqCBwBIb23oYLVaRdu20TRNjMfjg9wEAPBbuq6L5XIZ0+k0JpPJvvPB4GnbNubzedFxAAB/02KxiNlstu/6YPA0TRMRETunLkc/Xi+3rKKbGzdqTyjq+r3btScUc+tazntyz9Prj2tPKOrC1Uu1JxR18clG7QlFnbyyVXtCMY9Gd2pPKOr+3fu1JxSz6nbi4esH3/rlZ4PBs/cZqx+vR792tMy6ypoz52tPKGr3+LHaE4o5dTbnPbnneBypPaGok+snak8o6tzhL7UnFHX6xG7tCcU8O/Trh2UW60mf5z8a+g3HT8sAQHqCBwBIT/AAAOkJHgAgPcEDAKQneACA9AQPAJCe4AEA0hM8AEB6ggcASE/wAADpCR4AID3BAwCkJ3gAgPQEDwCQnuABANITPABAeoIHAEhP8AAA6QkeACA9wQMApCd4AID0BA8AkJ7gAQDSEzwAQHqCBwBIT/AAAOkJHgAgPcEDAKQneACA9AQPAJCe4AEA0hM8AEB6ggcASE/wAADpCR4AID3BAwCkJ3gAgPQEDwCQnuABANITPABAeoIHAEhP8AAA6QkeACA9wQMApCd4AID0BA8AkJ7gAQDSEzwAQHqCBwBIT/AAAOmtDR10XRcREaNu+8DGHLTlq+e1JxS1tvW+9oRi3rz8UntCUVvxqfaEot5uv6s9oagXn3O/S358N/jo+Oe9Hi1rTyhqe/dD7QnFrLqdiPjeLz8b9X3f/+pgc3Mz5vN5uWUAAH/ZYrGI2Wy27/pg8KxWq2jbNpqmifF4XHwgAMCf6roulstlTKfTmEwm+84HgwcAIIvcH5oBAELwAAD/AcEDAKQneACA9L4CY8xw8/mSdGgAAAAASUVORK5CYII=\n",
|
||||
"text/plain": [
|
||||
"<Figure size 720x72 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"sns.palplot(palette)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 33,
|
||||
"id": "dea9272c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"mt_preds = torch.load(\"finance_mt_preds_3.pt\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 34,
|
||||
"id": "442cc7f8",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"8 8\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"mt_x_list = mt_preds[\"x_paths\"]\n",
|
||||
"\n",
|
||||
"train_splits = list(range(252, 1259, int(252/2)))\n",
|
||||
"eval_splits = list(range(252+int(252/2), 1259, int(252/2))) + [1259]\n",
|
||||
"print(len(train_splits), len(eval_splits))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 35,
|
||||
"id": "c80581c1",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"train_splits.insert(0,0)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 36,
|
||||
"id": "cc856c43",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ind_preds = torch.load(\"finance_ind_preds.pt\", map_location=\"cpu\")\n",
|
||||
"ind_x_list = ind_preds[\"x_paths\"]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 37,
|
||||
"id": "b3eda070",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"times = np.arange(0, int(252/2), 20)\n",
|
||||
"times[-1] -= 2"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 38,
|
||||
"id": "249c9687",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"7"
|
||||
]
|
||||
},
|
||||
"execution_count": 38,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"len(times)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 39,
|
||||
"id": "e8163cb1",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"full_data = ind_preds[\"y\"]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 40,
|
||||
"id": "3b5acbf4",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"8"
|
||||
]
|
||||
},
|
||||
"execution_count": 40,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"len(eval_splits)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 41,
|
||||
"id": "69a5c15e",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"\n",
|
||||
"mt_quantiles_list = []\n",
|
||||
"ind_quantiles_list = []\n",
|
||||
"for i in range(len(eval_splits)):\n",
|
||||
" \n",
|
||||
" torch.cuda.empty_cache()\n",
|
||||
" \n",
|
||||
" train_start = train_splits[i]\n",
|
||||
" train_end = train_splits[i+1]\n",
|
||||
" test_end = eval_splits[i]\n",
|
||||
" test_y = full_data[:, train_end:test_end]\n",
|
||||
" test_y = test_y[:, times].unsqueeze(-2).cpu()\n",
|
||||
" \n",
|
||||
" ind_quantiles_list.append((ind_x_list[i][..., times] > test_y).sum(1))\n",
|
||||
" mt_quantiles_list.append((mt_x_list[i][..., times] > test_y).sum(1))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 42,
|
||||
"id": "b7fe0c10",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ind_quantiles = torch.stack(ind_quantiles_list) / 100\n",
|
||||
"mt_quantiles = torch.stack(mt_quantiles_list) / 100\n",
|
||||
"\n",
|
||||
"ind_quantiles = ind_quantiles.reshape(-1, 7)\n",
|
||||
"mt_quantiles = mt_quantiles.reshape(-1, 7)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 43,
|
||||
"id": "d40297ee",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"<AxesSubplot:ylabel='Density'>"
|
||||
]
|
||||
},
|
||||
"execution_count": 43,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAakAAAEFCAYAAABZ8hjOAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAA/tklEQVR4nO3dd1QU59cH8O8uHREBQaQqiIuKoBAVEAt2YwNLYi+oscduRI3lp7GXqLHGbsSoGOwFFRQLViyIIoqgUgRREaQuuzvvH76sDp1lcXaX+zknJ2fulL07jHN3Zp55Hh7DMAwIIYQQBcTnOgFCCCGkOFSkCCGEKCwqUoQQQhQWFSlCCCEKi4oUIYQQhaXOdQLKICcnBxERETAxMYGamhrX6RBCiFIQi8VISUlB48aNoa2tLdM2qEiVQUREBAYPHsx1GoQQopT8/PzQrFkzmdalIlUGJiYmAL7s6Nq1a3OcDSGEKIekpCQMHjxYeg6VBRWpMsi/xVe7dm1YWlpynA0hhCiXijwmoYYThBBCFBYVKUIIIQqLihQhhBCFRUWKEEKIwqKGE6RI2TlCxCWl4tPnbIjFEmhpasC8Vg2Y1qwOHo/HdXqEkCqCihQBAIjFElwNi8aFG09xLewlXid+RFGjuOhoa8DazAhOAgu0bGoL96Y2sKxtyEHGhJCqgIpUFSfME+Hg6bvYdvga4pM/lbp8dk4eomKTERWbDP/A+wAAO2sT9PB0RE9PR9jbmFZyxoSQqoSKVBV282EMZq05hlcJHyq0neg3KVi/Pxjr9wfDvm4t9PB0hFf7JrC1MpZTpoSQqoqKVBUkFkuwYucFbD10tdhleDwerGobwNhQDxrqasjIykXc21SkZ+aUuO2oV+8QtTcIa/cGoUkDS/Tp2BS92jnCxKi6vL8GIaQKoCJVxWRk5WLC4kMIvh1VaJ6Wpjp6tXNCz3aOcHOyga6OJms+wzBITc9C5Msk3HwUi5sPYxD25A1EYkmRn/XoWTwePYvH4i1n0eqHeujTqSm6tmqEajpalfLdCCGqh4pUFZKekYOhvnsR9uRNoXmDezTH9BEdYFpTv9j1eTwejGpUg4dLPXi41AMApKZnIfD6U5y68hjXw15CLClcsMQSCULuvkDI3RfQ0dZAF49G6N2xKVr/UA+aGnQIEkKKR2eIKiI7R4ghs/fg/tM4VtzUWB9/zfsZLZvayrRdQ31dDOjWDAO6NUNqWhbOXnuC40EPcevRqyJbB2bn5OF40CMcD3oEPV0ttP7BDh3c7NHOVVBigSSEVE1UpKoAsViCSUuPFCpQTgIL7F46FLWN5VMcDGvoYnCP5hjcozkS333C8aBwHLv0EJExSUUun5GVi3PXnuDctScAgEb1zNDcsQ6aOVijeeM6sDA1oHeyCKniqEhVAat2X0Tg9aesmEsjKxxY6QN9PdkGIiuNeS0DTBjYBhMGtkFkTBKOXXqIY5ce4W1KWrHrPH35Fk9fvsW+47cAfLnKc6xvjgY2prC3MUUD29qoZ2VMtwgJqULoX7uKC7oVhc0HQ1ixBra1K7VAFdTQtjYajukK39GdcTv8FQIuPcSl0GdISc0ocb3k9+lIfp+OSzefSWPqanzUszKWFi37ul8KmLWZIfh86uWLEFVDRUqFJX9Ix9Tl/qyYac3q2L982HcrUN/i8/lwb2oL96a2kEgkiHjxFsG3oxB0KwoPn8UX+QyrIJFY8qWZ+6t3OHn5sTSuo60BQZ1aX4qXTW04N7SCo705tDU1KvMrEUIqGRUpFcUwDOasO4HU9CxpTI3Px5YFA2Bey4C7xP4fn8+Hk70FnOwtMHVYe6SmZSHs6Rvci3iNe0/e4OGzeOTk5pV5e9k5eXgUlYBHUQnSmKaGGhwFFmjlUg/tXe3RtIEl1NToaosQZUJFSkUdD3qEC6GRrNiskR3h6mTDUUYlM6yhi47uDdDRvQEAIE8kxvNXyYiMScazmCREvUrGs5jkEp9pFSTMEyPsyRuEPXmDDf9choG+Djq3bIg+HZvCvaktFSxClAAVKRWUnpGD/205w4q5NLLC+AFtOMqo/DTU1eBgZw4HO3NWPC0jW9p34DPp/5PwKT271G1+Ss/GkfP3ceT8fZga66NfZ2cM93aDuUmNyvoahJAKoiKlgtbtC8L71EzptJaGOtbN7qsSVw419HTQwrEuWjjWlcYYhkFKasaXghWThEdRCbgX8brEDnOT36dj88EQbDt0DT08G2PMz63QxN6y8r8AIaRcqEipmOevkrEn4CYrNmFQG9hZ1+Ioo8rH4/FQy6g6ahlVR+sf7KTxtylpuHH/JS7feY4rd18g7XPhqy2xRIITweE4ERyOLh4NMXNkJzS0rf090yeElICKlIpZuesiq2siq9qGmDCwLYcZccfMpAb6dXFBvy4uEInFuHE/BgEXH+LctSfIyhEWWj7wRiQuhD6Dd4cmmDumC8zoNiAhnFP++z9EKuzJm0Iv7f4+7kfoaFEzbHU1NbRtXh8b5v6EhwFzsWKaF+ysTQotxzAMjl16CM/hf2Lb4WvIE4k5yJYQko+KlApZsTOQNe3c0Ard2jhwlI3i0tXRxJBergjeMwX7VwxH0waFn0VlZgvxx7Zz6DL6LzyIjCtiK4SQ74GKlIq4HR6Lmw9jWbE5Y7pQ33cl4PP5aO9qj1NbxmP3H0PRoIhnUc9fv4PXpG1YsSMQuUIRB1kSUrVRkVIRG/+5wppu/YOdzD2bVzU8Hg+dPRoi8O9JWDqlF2oU6I1DImGw6WAIuo3bXGxnuYSQykENJ1TAw2fxCLn3ghWbMrQdR9koLzU1PoZ7u6GHZ2Ms3X4eR87fZ82Pik1Gj/Fb8L+J3TG4Z4sqc5WalS3ErfBY3H8ah6jYZCS++4TPWbkQiyWooacDMxN9NLCtjWaN68C9qQ11RUXkioqUCth59AZr2tWpLtyaKGbPEsqgpoEe1s3uhx6ejvhtzTEkvU+XzssViuD75wlcv/8Sq2b24aQPxO8hTyTGhRuR+O/CA4TcfYHcvOJvdYY/T0DgjS+9m+hoa6BHW0cM83KFc0Or75UuUWF0u0/JJX9Ix+krj1mxiYOqZpNzeWvvao+gPVPQt5NzoXmnQyLQdYzqNarIzhFid0AoWg1Zi7GLDuJCaGSJBarw+nnwD7yPnhO2YvBvexDxIrESsyVVARUpJXfg1B2IxF/fi7K1MoZn8/ocZqRaaujpYMPcn7Bh7k/Q1dZkzXvzNhV9Jv+N3QGhZerBXZEJ80TY4X8drgNWYcFfp5FQQm8dZRVy9wW6jduM+RtPISu78HtphJQF3e5TYrlCEQ6cvMOK+fR2p3GVKkHfTs5wbmCF8Yv/xZPot9J4nkiMBX+dxp3wV1g9qw+qV1Ou238Mw+DSzWdYvPUsYuM/FLucVW1DtG1eH00bWMLGsiaMalQDn8fDp8/ZiIl7j3tPXiPoVhTr1ijwpdHJnmM3EXLvBTb/3h+OAovK/kpExVCRUmKnQx6zBg7U09XCT11cOMxItdlaGePE5nH4Y9s57D12izXvdEgEnrx8i+2LBqFRPTOOMiyfxHefMOfPEwi6FVXkfC0NdfzU1QUDuzeDk8Ci2IYiPzhY46euLhCLJbhy9wW2/BuC2+GvWMvExL2H96TtWDWrd5G3TwkpDhUpJVawj76fu7pAT1eLo2yqBm1NDfwxuRfcm9hi5ur/8DkzVzovNv4Dek7YiqVTemFAt2YcZlkyiUSCg6fv4o/t55GRlVtovp6uFkZ4u2FU35YwMape5u2qqfHRwc0e7V0FCL79HPM3nsSbt6nS+bl5IkxZ5o/4pFRMGdpeLt+FqD66L6Sk7j+Nw8Nn8azYCG93jrKperq3bYyz2ycVumrKFYowc3UApq88iuwi+gfkWmzCB/SfsQu+f54oVKD4fB6G9GyB6wdmwPeXLuUqUN/i8Xjo4GaPS7umYHCP5oXmr959CSt2BCr9czzyfVCRUlIHz9xlTbdzFcDWypijbKomG4uaOLF5HAZ2L3zVdOT8ffSYsBUv36RwkFlhYrEE249cQ6dRGwv1TAIAzRvXwfm/J2HFdG8YG+rJ5TN1dTSxckZvrJ/zE7Q02DdtNh0MwaLNZ6hQkVJRkVJC2TnCQs3Oh3u5cZRN1aajpYHVM/tgvW8/aBfoyDcqNhndxm3GieBHHGX3xdPot+g1aRuWbD2HnNw81jxdbU38Mbkn/tvwS6U9S+vX2RmH1o5C9WrsW9G7/gvF4i1nK+UzieqgIqWEzl9/yrpVU8uoOjxbULNzLvXr4oLTW8cX6lk9M1uIiUsOY+KSQ0hNy/quOWXn5mH5jkD8OHYzHhW4NQwAbZvVR/CeKRjxHVqENnesg0NrR8FAX4cV33H0RqGX0Qn5FhUpJXQ08AFrunfHJlBXU+MoG5KvgU1tnNk2AV7tnQrNOxEcjg4jNxQaSqWyXAuLRseRG7D5YAhrfDEAqKGnjXWz++LAqhGwrG34XfIBgCb2lvD/8xcYG1Zjxf+35SzOhER8tzyIcqEipWQSU9JwNSyaFetHzc4VRjUdLWz6vT+WTe0FTQ32D4d3Hz9j1PwDGOa7DzFx7yvl82Pi32PMQj8MnLkbrxM/FprftXUjBO+dip+7/sBJ34MNbWvDb5UPqul8fTGaYRhMXnoEdx+//u75EMVHRUrJBFx8yHrY7FjfnIY7VzA8Hg/DvNxweuuEIp/zBN+OQoeRG/C/zWcKvfwqq8R3n/D7xpNoP2I9zl59Umh+bWN97FwyBDsXD4FpTX25fKasHOzM8ff/BkNd7evpJzdPhDEL/ZD8QT77g6gOKlJKhGEYHA1k98xNV1GKq1E9M5zeOh5Th7WHWoFnPnkiMXYcvYGWg1Zj5uoAPIst/xAgDMPg/tM4TFxyCO4D12DvsVusLrKALwVzuLcbLu+diq6tGlXo+8hT2+b1sWpmb1YsJTUD4//3L42GTFjoZV4l8vBZPKK/adKsrsaHd4fCzz+I4tDUUMdLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"sns.kdeplot(ind_quantiles[..., 1], cut=0)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 44,
|
||||
"id": "1596ad48",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import pandas as pd"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 45,
|
||||
"id": "bcdf58f4",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ind_df = pd.DataFrame(\n",
|
||||
" {\"quantiles\": ind_quantiles.reshape(-1).numpy(), \n",
|
||||
" \"date\": np.repeat([np.array(times)+1],30*8)}\n",
|
||||
")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 46,
|
||||
"id": "319cc5d0",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ind_df[\"date\"] = ind_df[\"date\"] #/ (1259/5) #+ 0.7"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 47,
|
||||
"id": "9caee0ba",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df = pd.DataFrame(\n",
|
||||
" {\"quantiles\": mt_quantiles.reshape(-1).numpy(), \n",
|
||||
" \"date\": np.repeat([np.array(times)+1],30*8)}\n",
|
||||
")\n",
|
||||
"df[\"date\"] = df[\"date\"] #/ (1259/5)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 48,
|
||||
"id": "5c2ff989",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"0 1\n",
|
||||
"1 1\n",
|
||||
"2 1\n",
|
||||
"3 1\n",
|
||||
"4 1\n",
|
||||
" ... \n",
|
||||
"1675 119\n",
|
||||
"1676 119\n",
|
||||
"1677 119\n",
|
||||
"1678 119\n",
|
||||
"1679 119\n",
|
||||
"Name: date, Length: 1680, dtype: int64"
|
||||
]
|
||||
},
|
||||
"execution_count": 48,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"df[\"date\"]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 49,
|
||||
"id": "553d91a8",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ind_df[\"model\"] = \"Independent\"\n",
|
||||
"df[\"model\"] = \"Multitask\"\n",
|
||||
"full_df = pd.concat((ind_df, df))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 50,
|
||||
"id": "f01cd1a0",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"Text(0, 0.5, 'Quantile of True Value')"
|
||||
]
|
||||
},
|
||||
"execution_count": 50,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAu8AAAFWCAYAAADZrnNxAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAEAAElEQVR4nOx9d7wcZd39mZntu7ekAiEEJJGEGEKVUEINSHlBIIiiQKhSBKIICoi/F6Qo+Ip0RYKCSBUJUgWkhpoICVVSSUhPbnL71inP74+Z7zO7s+XutN29Yc/no+Hu7sw8uzs7c57znO/5CowxhiaaaKKJJppoookmmmii4SHWewBNNNFEE0000UQTTTTRRHVokvcmmmiiiSaaaKKJJpoYJGiS9yaaaKKJJppoookmmhgkaJL3JppoookmmmiiiSaaGCRokvcmmmiiiSaaaKKJJpoYJAjUewCDAZlMBp9++ilGjBgBSZLqPZwmmmiiiSaaaKKJJrZQqKqKjo4OTJo0CZFIpOj5JnmvAp9++ilOOeWUeg+jiSaaaKKJJppooomvCB566CHstddeRY83yXsVGDFiBAD9Q9x6663rPJommmiiiSaaaKKJJrZUrF+/Hqeccgrnn1Y0yXsVIKvM1ltvjdGjR9d5NE000UQTTTTRRBNNbOkoZ9VuFqw20UQTTTTRRBNNNNHEIEGTvDfRRBNNNNFEE0000cQgwaAg77Nnz8b48ePx/vvv29puw4YN+N///V9MmzYNkydPxhFHHIG77roLuVzOp5E20UQTTTTRRBNNNNGEf2h48r5gwQJcd911trdbv349vvvd7+Kxxx5Da2srDj74YCSTSdx+++04++yzIcuyD6NtookmmmiiiSaaaKIJ/9DQ5P2ll17C2WefjVQqZXvba665BuvXr8ePf/xjPPnkk7j99tvx0ksvYb/99sO8efPwt7/9zYcRN9FEE0000UQTTTTRhH9oSPK+fv16/PznP8fFF18MTdMwfPhwW9t/8cUXeP311zFmzBicf/75/PFYLIYbbrgBkiThwQcf9HrYTTTRRBNNNNFEE0004SsakrzfeuuteOqppzBp0iQ89thj2HHHHW1t/9Zbb4ExhkMOOQSiWPgWR40ahYkTJ2LNmjVYunSpl8NuookmmmiiiSaaaKIJX9GQOe877rgjbrrpJnz7298uIt/VgEj517/+9bL7/+STT7B48WKMGzfO1Vi9QEdHBxYuXIhMJgMAWLJkCf7xj38AAL7zne8UvI9wOIzx48djq622qsnY0uk0PvroI/T19ZUdlyiKGDNmTNnPe7Bg9erVWLp0Kf773/8WvM+ddtoJI0aMwC677OLofHSLnp4efPrpp/joo4+KPv9EIoHddtsN0Wi05uNyA8YYFi9ejNWrV0PTNP54pXOfEA6HsfPOO5dtXuF2XIsWLcKaNWsKxlXt2IYPH47JkyeXzeZtBHR1deHzzz9HMpkseq6a9wgA8XgcEydORHt7u2fjkmUZH374Ibq7ux2PKxAI4Otf//qg6seRyWTw0Ucfobe3t+T79PN8Lwfr77Pc5x8IBDBu3Dhst912NRtbE03UEuvXr8eiRYuQy+UgCAK/xteDC+SjIcn7ueee62r7jRs3AgBGjhxZ8nm6CG7atMnVcbxAKpXCU089VZCA88gjj/Ab2KOPPoqTTz65YJtFixbh+9//PhKJhO/je+mll7BmzZoBx7Vo0SJIkmR7laRR0NnZiWeeeQaMsZLvc/Hixchms9h7771rOi7GGJ555hn09vaW/fw3bNiA448/vqbjcoulS5filVdeKXp8oHOfsGjRIpx22mmIRCKejmvx4sV49dVXSz5XzdiWLFmCZDKJqVOnejouryDLMp599ln09fWVfL7azx8Ali9fju9973sIBLy5jcyZMweLFi1yPa6FCxfie9/7HoYOHerJuPzGv//9b6xatQpA+ffp1/leDkuWLCn4fVb6/BcuXIiTTjqJ21uz2Sw6OzvR19cHVVVrMt4mmvADmqYhmUwiGo1ygSyXy+Hjjz9GOByuej+SJKGlpQVDhw61tV0lNCR5d4t0Og0AZS909LiTQlivsXz5cuRyOfT2pbCxsxcACpSnrq4uLP1yPf97+JAWtLfGsXTpUuy2226+ji2Xy2HNmjVgTEO2f03RuDJ9+g1HCsYRjAzFihUrBi15X716NRhj6O5NFr3P1es2Y/Q2w7Bq1aqak/fu7m709vZCU+WSn384sS3Wrl0LWZYRDAZrOjY3+OKLLwAA6zu60Z/K8McrnfuErYe3IRHXvzOvV86WL18OAMilN0FT0gXPlTv/CYIYQDi+DVasWNGw5H3NmjXo6+tDMp3Fuo1dRc9X8/kDwKiRQwD0YO3atRgzZozrcTHGsHz5cjDGkE2uBVjhqsdAnz0hEG5DINSKL7/8clCQd8YYv/YsW7mh5Oe/zYh2xGP6JH377bevybhoMpHLbIYmp8p+/vR5r1q1CsOHD0c2m8XKlSsxZMgQ7LDDDggGgxAEoSZjdgrGGFRVRSqVwvr1+vm+9dZbc7ImiiJEUazL+1BVFZqmIZ1OF42tnuP6qiCTyaC3txeMaWCaAkEQIIhBBAKBqq8vjDHIsoze3l6sXLkSY8aM8YTAb5HknZYzyp3UjLGCf+sJWiXo6klic1dpNSz/cUkU0N4aR0dHh+9j6+npAQBoag5yrqfoeXpMYwqCkaH89YMRNOHrT2aKntvc1YfR2wzjr6klyNqgacW9CeRcD0LaVhCkEFKpFNra2mo9PMfYvHkzAGBTZy+S6Wz515X4TcQiISTiUXR2dno+LjqHlVwPVKXy913qNxGKbY2+vj4oiuKZIu0lNmzYAADo6S1/vclHudfEo2HEomFs2LDBE/Le19eHXC4HxhTI2eJJhRWlPnsAgCAiEGrl51ejI5vN6uRR00p+1pu7+tDeEkM8FuG2ylqAhC011wdF7i96nj5/QQoiEGrl18bOzk4MGTLEdshEvSDLuijCGENnZydfKVi7dm0BOQuFQmhra6spUe7r6yv4XEuNrR7j+iqBc0SmgTEVgAABQVvcURAEhEIh/pvo7OzENtts43psjXd38QCxWAwAyl7sslmdLDSCT5hIeDKtj3XYyOIv9Ws7TUKytxsb169GMqWPnUi/nyC1hamlydVuk7bHh59+yZ8v5VUdLKBzRSmxzCsbj9WTvDOtdF8CpimAFEIymRw05F1RFPT09IAxhnRGn5SM2XECpEAA8999ueC1X9tpEv/vzRvXore7EymD7HtN0Bhj/EauGef0N3cfi3BIv0w++Pc3C14/dcp4/t9vzdXtHkzLgQlh9PT0YNiwYZ6OzwsQeafVjrYhwzB0hHnNqfT5A+Z3QNuTGugW9F1qinnNzv98K332ALChowdLvlgPTc0U7K/RQfciRdHKvkZRtYLX1gJ0PdQJSzG2HtGG9R09+vUH5rWxr68PO+ywQ03G6AVSqZT+ezfUd4KueOsETRD0VehcLueZ5WEgqKrKP1PGtKKxMaZBEATkcjnIsoxQKFSTcX3VUETSWZnHq0RraytWrFjRJO/lQF73cp52IszlPPG1gizL2Lx5MzSNcfVxux12KnrdsBHbYMjQkejYsAapTBaqpqGnpweZTMZXDySR8VKqLwBM2nkMPv7vSmiaAsb0i43fY/ILdKGUleKblapq0DQNsizXXFElBYxuklZomgwJKFl82Kjo6uoCYwyZbA4aYwhHIhi5TemCt2F5xDISjaO3ey5SBuH3mqD19/cbN0YFjGkIhwLYeadty75+xx3MovFlKzZg3YZuaGoOotSY5J0xho0bN4Ixxsn3dl8bj0g0Xnab/M8fAGLxFny24F0kje1pf26VP1pFoUnT+HHbFHy+VlifG7XNUIO860p2V1cXNE2re1HZQKBap0recHquluSdjlXuujNhp2118s4Kx6aq6qCy7ymK/v6y2WJxhISFUCiAYECqqX+fxsWYCk0tvv9qahaiGIQgBqAoyqAn77IsI51OF4QEiKKIWCxW1xVMczys4F+n5D0YDHp2HjX2lc0hqBK+XBTksmXLAAA77VRMlGuJTZs2GepjFprGEInGEAiW/hGKkoRoLAHGgFSN1PeuLn35WiujvIuigNZE1HiNfoEZrOo7kXelBHkHTPWrlkvXQL7yXvomSo8PJvJORI1IeDRWXeF1NBqHIEAn/ZqG3t5eTzslk2WGGedya0v1K3P0Wvqt9Pb2ejYur9Dd3W0odQpkWUUgEEA4ErO1j0g0BikQQDanICcryGaznrxX81pjrAi02hxXOIhIOGh4U2VohsDR6OCkVyuvvNNz+aEGfmMg5T0Y1AkV04onFoPFwsEY4+SsEhljhgJfS/LOSWNFksgKXztIoWkauru7kclk+ApHLpdDJpPhK6H1ArdY0995jzsZl5e/jS2SvB9wwAEAgFdffbXoxF67di0+//xzbLvttnWPiTSXsPULXyzRWvH18ZZW4/WZgu39Alfey5B3AGg1brLaILfOcNtMOfJuPF4v8q6Vs80w/fFGKL6uFqSYp9M6GYlUSd5FSUI4EgNjQMZQyrz0vefXeABAIl49eW+hSayxStWIxJEm++b1xr5XVhBExOP6dSjp4XXIqry3t5VfDSgH2ob2QROCRoapvFcg7zW2zWiaZtQfMDBWelxkJaPna7kq4BU0TeMkrBQPC0eM3zSrPXk3j1VhUlGHcfkBfbWTQdMYMlmZ/48mV/V8fyZBz/8eGmPSNOjJ+9q1a7Fs2bKCm/h2222HAw44AMuXL8dtt93GH0+lUvjlL38JVVVx5pln1mO4BaCbHt0EEy2VPcvxRFvB6/1U3kv5f0uhzaI4DoYbZilw72mZm2i9yXtZz7uqP97fX1xU1qig3yotS8eqJO8AEInpBI1Uey/PN1KQiYC3JKq3fyXi+mvp+2hE5Z2uF3T9IDHALmKJFmM/+m/GbfF8qWtNW4s95R0A2loLr0WLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 864x360 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(figsize=(12,5))\n",
|
||||
"violins = sns.violinplot(x='date', y='quantiles', data=full_df, hue=\"model\", cut=0, ax=ax,\n",
|
||||
" palette=[palette[0], palette[-2]])\n",
|
||||
"for violin in violins.collections[::2]:\n",
|
||||
" violin.set_alpha(0.5)\n",
|
||||
" # violin.set_edgecolor(palette[-2])\n",
|
||||
" \n",
|
||||
"plt.xlabel(\"Time Step Lookahead\")\n",
|
||||
"plt.ylabel(\"Quantile of True Value\")\n",
|
||||
"# violins = sns.violinplot(x='date', y='quantiles', data=df, color=palette[0], cut=0,ax=ax)\n",
|
||||
"# for violin in violins.collections[::2]:\n",
|
||||
"# violin.set_alpha(0.5)\n",
|
||||
"# violin.set_edgecolor(palette[1])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 51,
|
||||
"id": "6a1b967e",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def is_covered(vec, test_pt, sd_scale=1.96):\n",
|
||||
" vec_mean = vec.mean(1)\n",
|
||||
" vec_std = vec.std(1)\n",
|
||||
" return ((vec_mean - sd_scale * vec_std) <= test_pt) * ((vec_mean + sd_scale * vec_std) >= test_pt)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 52,
|
||||
"id": "ee8084c2",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Automatic pdb calling has been turned OFF\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"%pdb"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 53,
|
||||
"id": "84bf8119",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"\n",
|
||||
"mt_sd_list = []\n",
|
||||
"ind_sd_list = []\n",
|
||||
"for i in range(len(eval_splits)):\n",
|
||||
" \n",
|
||||
" torch.cuda.empty_cache()\n",
|
||||
" \n",
|
||||
" train_start = train_splits[i]\n",
|
||||
" train_end = train_splits[i+1]\n",
|
||||
" test_end = eval_splits[i]\n",
|
||||
" test_y = full_data[:, train_end:test_end]\n",
|
||||
" test_y = test_y[:, times]#.unsqueeze(-2).cpu()\n",
|
||||
" \n",
|
||||
" mt_sd_list.append(is_covered(mt_x_list[i][..., times], test_y))\n",
|
||||
" ind_sd_list.append(is_covered(ind_x_list[i][..., times], test_y))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 54,
|
||||
"id": "56b23fbf",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"mt_sds = torch.stack(mt_sd_list)\n",
|
||||
"ind_sds = torch.stack(ind_sd_list)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 76,
|
||||
"id": "e23516f8",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAhwAAAFWCAYAAAAi1UTrAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAACbIklEQVR4nOzdd3xUVfr48c+5M5NJJpPeSOgQQlfp6CIqxY5dUbHL6uqiuPb16/52xd3vV1d3RdDVdXFlxbLYxQIoWEClKIKI9BIgpPdMMpl2z++PIZMMk04mMyHn/Xrx0px7Z+aZ5Cbz3FOeI6SUEkVRFEVRlCDSQh2AoiiKoignPpVwKIqiKIoSdCrhUBRFURQl6FTCoSiKoihK0KmEQ1EURVGUoFMJR5C43W5ycnJwu92hDkVRFEVRQk4lHEGSn5/P1KlTyc/PD3UoiqIoihJyKuFQFEVRFCXoVMKhKIqiKErQqYRDURRFUZSgUwmHoiiKoihBpxIORVEURVGCzhjqABSlPTylBTjWL0e3laPFJGCecC6GxLRQh6UoYa2ytIpt322npsqOJSaKEacNIzYxJtRhKd2ESjiULkWvKKF8/hwcm75AaAak24kwRiBfeBDzmCnE3/McWlxSqMNUlLBiq6jmzafeYef3u9E0gdvtwWg08O7CZQwZl8U1D1yBNS461GEqJzg1pKJ0GXpFCUV3TcbxwypwOZCOGvC4vf91OXD8sIqiuyajV5SEOlRFCRu2imqeuu1ZdmzcidvlxulwoXt0nA4XbpebHRt38dRtz2KrqA51qMoJTiUcSpdRPn8OelkhuF2Nn+B2oZcVUj7/rs4NTFHC2JtPvUNVmQ2PW2/0uMftoarMxptPvdPJkSndjUo4lC7BU1qAY9MXTScbddwuHJtW4ykt6JzAFCWMVZZWsfP73XjcnmbP87g97Px+N5WlVZ0UmdIdqYRD6RIc65cjNEPrThYS+2fPI6vzkLLxuzpF6Q62fbcdTROtOlfTBNu+2x7kiJTuTE0aVboEj60M6apt3cluF3ruj+i//BOMkYiYfhDbHxHTH6JSEKJ1f4AVpaurqbLjdjXfu1HH7fZgt9mDHJHSnamEQwl70m1H1OwBgwZ6K3osDAa0SLP3/921yLKdULYTCWCKRsT2h5h+3v+aE1UCopyQpJQIXaAZNPRW/N4YjQairFGdEJnSXamEQwlrsjoXfe9bRPSOBl227kG6JGLwgMaPuaqRJdugZJs3AYmIRcT2P5qE9EeY4zoqdEUJGd2jU3yglLT0NKRs3e+NrktGnDYsyJEp3VnYJhzfffcdL774Irt27cLlcjF8+HBuu+02Tj/99FY/x/fff8+//vUvtmzZgsvlon///lxxxRVcddVVGI2Nv/VVq1bx6quvsn37dhwOBwMGDOCqq67i2muvVXfCnUhKiSzahDy0HHQPhphoIjL74tybDZ5m7taMRswjx2HsMxZZlQ3uFrqInZXI4p+QxT95v45MRNT1fsT2R5isHfWWFKVTOGucFO4rwVXrJjommr5ZfcjedRC9md8bzaAxZFyWKgKmBJWQrU1/O9F7773H73//eyIiIpg4cSK6rrNhwwZcLhfz5s1j5syZLT7Hm2++ybx589B1nX79+jFgwAB2795NTk4OkyZNYuHChVgsFr/HPPPMM7z44ouYTCYmTJiAlJIffvgBh8PBzTffzMMPP9zq95CTk8PUqVNZvXo1vXr1avP3oDuTHhfy4Mf1ScBRerWd0peWotuqG1+tYjShJaSSsnANWlyS986uJt+beFQeQFYdBI+jbcFEpXh7QGL6Q0xfhMnS8mMUJURsxdUUHyxDNugNtFfbeWPBUmpsNY0ujdUMGharhbuevp0eA1M7M1ylmwm7hKOwsJCpU6diNpt54403yMrKAmDr1q3cfPPNuFwuPv/8c9LSmi5jvX//fmbMmIHb7eYPf/gD1113HQAej4enn36af//739xyyy089NBDvsesW7eOm266ibS0NF555RUGDhwIwL59+5g1axZlZWW8//77DBvWui5HlXC0j6wtQd+7FGoKA46JtAnI2DFULLgHx6bVCE1Dul0Iowmp65jHTCX+noVNVhqVUofqXGTlAW8CYjsEurv1wQkBUWkNhmD6Igzm9r5VRekwui4pPVxOVaEt4FiExUR0ciRvL/zQr9KoZtCQuqRvVh+mXzGF6Nhoeo5IwxgRth3fShcXdlfWa6+9htPp5Pbbb/clGwAnnXQSs2fPZv78+SxdupS77767yed47733cLvdXHjhhb5kA8BgMHD//fezdu1alixZwq9//WsSExMBePHFFwF44oknfMkGwMCBA5k9ezavvfYa27Zta3XCobSdLN2OfuDDwF4IQwRav4sQSSMASPzjG966HBtWePdSscZjnngehoTm786E0MDaC2HtBRmnI3U3VB+pT0Cqc0BvZkZ/XY9JTT4yfx0IDRGdUT8B1doHYTAd77dBUdrE5XBTtK8ER7Uz4FhMSjSJvePRDBq//vONvr1U7DY7EZFmEhOSsFi8E0V1j07JwTJSM5PV8LESFGGXcKxduxaAadOmBRybPn068+fPZ82aNc0mHLt37wZgypQpAccMBgNjx45lz549rFu3jgsuuICysjI2btxIVlYWp512WsBjZs+ezezZs9v7lpQWSN2DzFnl/RA/VlQKWuZViKgUv2ZDYhqW8248rtcVmtHbSxHTF3qeifS4wHbIm4BUZSOrj3iTjKYDR9pywJaDzPsGNAMiupd37kdsf4ju6X0NRQmSmnI7RQdK0Y8ZKhGaIKlvAjHJ/vujxCbGcNqFE3xfVxRUUXqovMHz1VJdWoM1Se2ronS8sPprKKVk7969aJrGgAGBqwz69euHpmns3bvXu+SriSy8bglYdHTjvzR1E0b37dsHwM6dO9F1nZNOOgmA9evX8/XXX1NVVUVmZiYXX3wxCQkJx/3+lEDSWYnc9w6y6lDAMZE0AtHvIoQholNiEQYTxA1ExHl7uKS71peAyMoDUJPf/BPoHu88kaqDyCNfgWZCWHs3SEAyvL0sinKcpJSU51ZSnlsZcMwUaSRlYBJmS8u/N7GpVmrK7NRW1fcqlhwqJyo2EoOplYX2FKWVwirhqKiowOl0kpiYSERE4C+L0WgkISGBkpISqqursVobX0HQv39/1q5dyw8//MCZZ57pd0xKyY8//ghAaWkpAIcOeT/sEhISuOeee1i+fLnfY1544QWef/55xo4de7xvUWlAVh5A3/cOuI7ZNEozIHqfg0gdF9KuXWGMhPgsRLx3aE+6arzJRNXRBMRe1PwT6C5k5X6o3O9dgmswe3tT6oZgLD1U17XSZh6Xh6L9pdgrAwvhWRKiSO6XiMHYusRWCG9PSO72At9EU92tU3K4nNQBatdlpWOFVcJht3uXMEZFNV18JjIyEqDZhOOSSy5hyZIlLF68mLFjx/qSDiklzz//PL/88gsATqd3zLOqyrt/wNKlS3G73Tz++ONMmzaN6upq/vOf/7BkyRJ++9vf8tFHH5GaqmZxHy8pJTLvG+SRLwKHLMxxaAOv9M6zCDPCZIHEoYjEoQBIZ5V36KWuB8RR1vwTeBzI8t1QvtubgBij6qugxvaHSDV2rjSv1uagcF8JHucxc40EJPSMI65HTJuvoYgoE/EZsZTlVPjaqktqqEm0YIlXhcCUjhNWCYemtZyVt2ZRzfDhw5k7dy7z58/n9ttvZ9iwYWRkZLBnzx5yc3OZOXMmS5cu9Q2tOBze7sTKykr+9re/ceGFFwKQmJjIo48+SmFhIStXruT111/nd7/73XG8Q0W67cj973s/eI8h4gYiBlzeZZaeiogYSBqJSBoJgHSU+ycgzsDubj9uO7JsB5TtOFoF1YqI7eddghvbH8wJKgFRAO/fvcpCG6WHy+GYP4EGk4GUAYlExUa2+/nj0mKoLq3BWVO/3Lz4YBk9reZW95YoSkvCKuGoq4tRlwA0pu5Yc70gAHfccQeZmZm8/PLL7Nixg5ycHMaOHcszzzzD/v37Wbp0KbGxsX6vGxsb60s2Grr66qtZuXIlGzZsaNf7UrxkdR763rcCewKEQGScgciY3Oo5DnWz7Wuq7Fhiohhx2rCQFy0S5ngwn4JIPsWbGDtK6yegVh4IHDo6lsuGLNnmrYQKYI5DxPSvL8XeTBXU8v37OPz2q3gqyzDEJtD7yhuIHzCwyfOVrkP36BRnl1FdWhNwLDLGTMqAJIwRxzffQmiC5H6J5O4o8CU0HqeHsiMVJPdV89eUjhFWCYfVasVisVBWVobb7Q6oBup2uykrK8NsNvuSheZMnz6d6dOnB7SvWrUKgPT0dADfhNCePXs2+jwZGRkAlJW10GWuNEpKCUU/oh9aHlj3whiFNvByRFxmq57LVlHNm0+941dPwGg08O7CZQwZl8U1D1yBNS70M+yFEBCZhIhMgtSx3u+BvQhZdXQJblU2uFvYjM5RgXRsQRZv8X4dmVhfhCy2H8JkpfLwYfY9cA3J5TuIFQKD9OARBqq+fI4D8UMZ+NSbxPbuHey3qwSJ0+6icF8xLntgvZi4HjEk9IrrsF4wc3QEcT1iqMir36K+qtBGdELUcfWeKEqdsEo4hBBkZmaydetWsrOzycz0/xA6cOAAuq771edoTGlpKbt27SI9PZ1+/foFHF+/fj0AI0d6u8IHDx4MQFFR45MAi4uLAXw1O5TW81YN/aT+Q7MBYe2JGHhVq/cvsVVU89Rtz1JVVuVXMdF5tGTzjo27eOq2Z3ngpblhkXQ0JIQASyrCkgppE7xFyGoKjvaA1FVBDayj4Ke2FFlbiizcBECVLZLCp58jRa/BII5+PwQYcIOA5PLt5N9+GvzzO5V0dEG2kmqKs/2rhgJoRo3kfolEJ3T8/Ir49Fhqyuy4ausTnJKDZWQMS0MzqKEV5fiEVcIBcPrpp7N161ZWrVoVkHDU9UycccYZzT7HL7/8wuzZs5kLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 576x360 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.figure(figsize=(8, 5))\n",
|
||||
"plt.plot(times+1, mt_sds.float().mean((0,1)), color = palette[-2], label=\"MT Volt\", alpha = 0.5)\n",
|
||||
"plt.plot(times+1, ind_sds.float().mean((0,1)), color = palette[7], label=\"Volt\", alpha = 0.5)\n",
|
||||
"plt.scatter(times+1, mt_sds.float().mean((0,1)), color = palette[-1], s=120, zorder=4)\n",
|
||||
"plt.scatter(times+1, ind_sds.float().mean((0,1)), color = palette[6], s=120, zorder=4)\n",
|
||||
"plt.xlabel(\"Time Step Lookahead\")\n",
|
||||
"plt.ylabel(\"95% Calibration\")\n",
|
||||
"plt.legend()\n",
|
||||
"plt.axhline(0.95, color = palette[1], linestyle=\"--\")\n",
|
||||
"sns.despine()\n",
|
||||
"plt.savefig(\"mt_calibration.pdf\", bbox_inches = \"tight\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 134,
|
||||
"id": "705742fd",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"mt_sd_list = []\n",
|
||||
"ind_sd_list = []\n",
|
||||
"for i in range(len(eval_splits)):\n",
|
||||
" \n",
|
||||
" torch.cuda.empty_cache()\n",
|
||||
" \n",
|
||||
" train_start = train_splits[i]\n",
|
||||
" train_end = train_splits[i+1]\n",
|
||||
" test_end = eval_splits[i]\n",
|
||||
" test_y = full_data[:, train_end:test_end]\n",
|
||||
" test_y = test_y[:, times]#.unsqueeze(-2).cpu()\n",
|
||||
" \n",
|
||||
" mt_sd_list.append(is_covered(mt_x_list[i][..., times], test_y, 1.15))\n",
|
||||
" ind_sd_list.append(is_covered(ind_x_list[i][..., times], test_y, 1.15))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 135,
|
||||
"id": "43562c95",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"mt_sds = torch.stack(mt_sd_list)\n",
|
||||
"ind_sds = torch.stack(ind_sd_list)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 136,
|
||||
"id": "6aeb0a62",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"<matplotlib.lines.Line2D at 0x7fcc9b4464d0>"
|
||||
]
|
||||
},
|
||||
"execution_count": 136,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAawAAAEgCAYAAADoukcPAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAB7u0lEQVR4nO3dd1gTWRcH4N8kIYHQmw1BUAzYBXvDhrp27L27a9e1Y1sVdde6uvYudrEiIvbeALF3RAEBRemd1Pn+4CM4hgCBQAjc93l8djlTckJCTubOLRRN0zQIgiAIopRjaToBgiAIgigIUrAIgiAIrUAKFkEQBKEVSMEiCIIgtAIpWARBEIRW4Gg6gbIqMzMTr1+/hqWlJdhstqbTIQiCKPWkUiliYmJQt25d6OrqKmwnBauYvH79GsOGDdN0GgRBEFrn6NGjaNy4sUKcFKxiYmlpCSDrF1+pUiUNZ0MQBFH6RUdHY9iwYfLPz1+RglVMspsBK1WqhKpVq2o4G4IgCO2h7DYK6XRBEARBaAVSsAiCIAitQJoECaIAEsOjEbTXF+/O34MoLRNcfV3U6t0Gjcf3gEk1co+SIEoCucIiiHyE3n6Gg93m4pXXDYhSMwCahig1A6+8buBgt7kIvf1M0ykSRLlAChZB5CExPBo+U/6FJEMImUTK2CaTSCHJEMJnyr9IDI/WUIYEUX6QgkUQeQja6wuZWJLnPjKxBEH7LpZQRgRRfpGCRRB5eHf+nsKV1a9kEineed8toYwIovwiBYsg8iBKy1TrfgRBFB7pJUgQeeDq62Z1tCjAftpMKBQiPj4eKSkpkErzvqIkiIJgs9kwNDSEmZkZeDyeWs5JChZB5KFW79Z4cfRanvuwOGzUcnMpoYzUTygU4suXLzA1NYWtrS10dHRAUZSm0yK0GE3TEIvFSE5OxpcvX2BjY6OWokWaBAkiDybVKue7D8ViofG47iWQTfGIj4+HqakpLCwswOVySbEiioyiKHC5XFhYWMDU1BTx8fFqOS8pWAShhEQoxrNDl/LdT9/SBEZVK5RARsUjJSUFRkZGmk6DKKOMjIyQkpKilnORgkUQSrw4cgXJkTE5AQrQ4esCv1yBJEfFaHUvQalUCh0dHU2nQZRROjo6arsvSgoWQeQiMzkN/tvOMmJOI3/D9NeHMPuTF2q5tWFse/CvFyRCUUmmqFakGZAoLup8b5GCRRC5eLzzPDITU+U/cw300Hxqf/nPrWYNApub02cp5Vscnh28XKI5EkR5QwoWQfwi+Wssnh5gzlzR5I9e4Jvn3OcxrloBDUf8xtgnYPs5ZPxU5AiCUC9SsAjiFw83nYREKJb/rF/BFI3GKvYCbDa5D3iGfPnPwuQ0BGw/q7AfQZQ1NE1r5HFJwSKIn8R++IK3Z+8wYi3/HJDV2eIXeqaGaDq5DyP2/NBlJEfFKOxLaKezZ8/CwcEBDg4OaN26NWQyWZ77X758Wb6/u7s7ACAgIEAeK+i/gICAXM///PlzODg4oF69ekhOTs43/+vXr8PBwQF9+vTJd99fbdmyBQ4ODti+fTsjfvHiRcyZM0fl86kDGThMED+5t/YYaFnOt0ezGlao27+90v2dRnXFs0OXkfotDgAgFUnw4F8vdN0wtdhzJUpWTEwMnjx5giZNmijd59IlxWEQFhYW6NmzJyOWkZGB69evA4DCtuxjctOwYUPY29sjJCQEV65cwYABA/LM2cfHBwDQv3//PPcrqKdPn2LWrFlo2rSpWs6nKlKwCOL/Ivzf4POtp4xYm3lDweKwlR6jo8tFq5mDcGVezrfQt9730GhcD1SobVtcqRIlzMjICMnJybhy5YrSgpWeno47d+5AR0cHYnFOk3KNGjWwfv16xr6RkZHygvXrtvz07dsXa9euha+vb54FKykpCbdu3QKXy0WPHj1Uegxl8rvCLG6kSZAgkNUmf3f1EUasSiMH1HBtnO+xtfu4wMLB5ueT4e6ao+pOscxIDI/G9SV7saX+KGyoMQhb6o/C9SV7S/WaYq1btwaPx8PVq1eV3r+5desWMjIy0KZNm1y3q4ubmxs4HA4CAwPx48cPpftdunQJIpEInTt3hrGxcbHmVFJIwSIIAMF+jxD98hMj5uI+vEBjSFhsFtrMG8qIhd97gfD7L9WaY1mgras38/l8uLi44Pv373j2LPcc/fz8wOfz0a5du2LNxdzcHC4uLpDJZPDz81O6n7e3NwCgX79+8phEIsGRI0fQt29fNGzYEE5OTujfvz+OHj0KiSTvdd/c3d0xbNgwAEBgYCDjPl1JIQWLKPekIgnurz/OiNXs0hRWjRwKfA67dk6wbl6HEbu75ihoDTehlCbavnpz165dAQBXrlxR2Jaamop79+6hQ4cO0NUt/pn7s4vQhQsXct0eERGBZ8+ewcrKCi1atACQNcnxmDFjsGLFCoSFhaF58+Zo1qwZPn36BA8PD0yYMAEikfLB705OTmjdujWArKLZs2dPODk5qfmZ5Y3cwyLKvZcnriMx/Lv8Z4rNQuu5Q/M4QhFFUXCZPwxH+yyUx368CcX7Cw9Rq3drteVa0r48eo0bf+1D/KeoEnk8SYYQ+9pPL/J5zGpYoaPHONi0qKuGrLK0a9cOurq6uHr1KhYsWMDYdv36dQiFQnTt2hVpaWlqe8y8crGwsMDr168RFhYGW1tbxvbsq6u+ffvKWwn+/fdfBAYGwsnJCdu3b4eZmRkAIC4uDhMmTMD9+/exefNmpT0ABw0ahBo1auD+/fu53pcrCSpfYb18+RITJ05Es2bNULduXdSqVUvpv9q1axdHzgShNsKUdDzacpoRqz+oI8yqV1H5XJUa2MOhR0tG7P6G44wxXdrm+qI9JVas1Cn+UxSuL9qj1nPq6+vDxcUFX79+xcuXzObeS5cuwdDQEC4uJbPMDIfDQa9evQAAvr6+Ctt9fHzAYrHk3dkzMzNx4sQJcDgcbNy4UV6sgKyrpY0bN4LNZuPo0aMQCoUl8hwKQ6WC9fr1a4wYMQJ37txBUlISJBIJaJpW+k/TPUoIIj9Bey8gIy5nPIsOn4cWM/LuKpyX1rMHg6WT06swOTIGL44oNiER2im7WfDy5ZxpuJKSkvDgwQO4urqCy+WWWC7ZXdV/LVhPnjzBly9f0KJFC1hZWQHI+uzOzMxEgwYNULmy4pI51tbWqFevHtLT0/Hq1aviT76QVGoS3LFjB4RCIezt7fHHH3/Azs6uRNprCaI4pP5IQNBe5h97o3E9oW9pUuhzmlSrhAZDOzHmFfTfdhZ1BrSHrpF+oc+rKa6rfseNpfsQH6JdV1lm9lbouHyc2s/7c7PgvHnzAADXrl2DWCxGt27d1P54ealRowYaNGiAFy9e4PXr16hbN6v58/z58wCYY6+yexNmF7DcVK1aFc+fP0dsbGwxZl00KhWsoKAg8Hg8eHp6Kh3YRhDa4tF/pyDJyGn+4Jsbo8nvioM4VdV8an+8OXMnqxccgMzEVDzeeV6hJ6E2sGlRF2OublTLua4v2YtXXjcUOlz8jMVho95gV7h6qL/YqAOfz0fbtm1x5coVvHnzBnXq1MGlS5dgYmKCli1b5n8CNevXrx9evHgBX19f1K1bFyKRCJcvX4aJiQlcXV3l+2V3xc+r12v2EiAleZWoKpWaBDMzM1GjRg1SrAitF/cpCq9O3mTEWkzvD66BXpHPzTc3QpM/ejFiTw9cRMr/Z8MorxqP7wGWTt7fkVk6nFK/evNvv2VNenz16lUkJCTA398fXbp0AYdT8n3YunfvDj09Pfj5+YGmady+fRtJSUno2bMno/BUqJC1wGhERITSc2VvK82f7yoVLBsbmzwHqhGEtri/7hhoac49VlPbyqg3uKPazt9obHfoVzCV/ywRivFw00m1nV8bmVSrhF7bZoGjx1OYPYTFYYOjx0OvbbNgUq2ShjIsmPbt20NPTw9XrlzBjRs3IJFI5Pe2SpqBgQE6deqE79+/4/nz5/J7az+PvQKAunXrQk9PDy9fvsTXr18VzvPlyxe8ffsWhoaGcHR0VPp4ml43TaWC1atXL8TGxjJuOBKEtokKeo+Qq48ZsdZzhoCdz7d/VejwddHyT2bnjTdnbiP2wxe1PYY2smvnhFF+61BvsGvW1SxFgWugh3qDXTHKbx3s2pXsuJ7C0NPTg4uLC0JDQ7Fv3z5YWFigWbNmGssnuzhdvHgRt2/fRp06dVCrVi3GPnp6ehg4cCAkEglmzZqFhIQE+bb4+HjMmjULMpkMAwcOzLNJkMfjAYDalrxXlUp/oWPHjkVAQAAWLlyIqKgouLi4oGLFinkur62nV/QmFoJQF5qmceeXKZgqN6yJml3V/4FTt397PNl3Ud4tnJbRuLf2GPrsK9nZAUobk2qV4OoxrtTepyqIrl274sqVK/j8+TOGDx8OFktzczA0a9YM1tbWOHHiBMRiscLVVbZZs2bh7du3ePz4MVxdXeVzIgYGBiItLQ2tW7fGn3/+medjWVlZgcPh4N27dxg7diyaNGmCSZMmqfspKaVSwXJzc4NUKkV6ejrWr1+f78AxiqLw9u3bIiVIEOoUcvUxvj0NZsRc3IcVS1MHi8NGm3lDcX7COnns862niAh4C+tmZIyiNmvXrh34fD7S09NLvHfgryiKQp8+fbB582bweDylE93q6upi//79OHbsGHx8fPDo0SPo6OhAIBCgX79+6NevX76F19TUFCtWrMDWrVsRGBgIsVhcogWLolVYiSuvtk1l3r9/r/IxZUFkZCQ6duyIGzduoGrVqppOh0DW9D+ev81GwuecNvzqHRuhz575xfaYNE3jxMC/8PXJB3msUgN7DD27SuP3A7K9e/dOoQmJINSpoO+x/D43VbrCunHjhiq7E0Sp8urkTUaxolgUXOYNK9bHpCgKbReMwPH+i+Wx6BchCL7kD4duLYr1sQmirFGpYOU16IwgSjNRWiYe/XeKEavTvz3Maxb/1W8VZwFqdmmLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(times, mt_sds.float().mean((0,1)), color = palette[4], marker = \".\", markersize=20, label=\"MT Volt\")\n",
|
||||
"plt.plot(times, ind_sds.float().mean((0,1)), color = palette[-2], marker = \".\", markersize=20, label=\"Volt\")\n",
|
||||
"plt.xlabel(\"Time Step Lookahead\")\n",
|
||||
"plt.ylabel(\"75% Calibration\")\n",
|
||||
"plt.legend()\n",
|
||||
"plt.axhline(0.75, color = palette[1], linestyle=\"--\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "ec351f4d",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.8"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,866 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Warning no robinhood utils.\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns\n",
|
||||
"import torch\n",
|
||||
"import gpytorch\n",
|
||||
"import os\n",
|
||||
"import pandas as pd\n",
|
||||
"# import pickle5 as pickle\n",
|
||||
"\n",
|
||||
"import sys\n",
|
||||
"sys.path.append(\"../\")\n",
|
||||
"from voltron.likelihoods import VolatilityGaussianLikelihood\n",
|
||||
"from voltron.kernels import BMKernel, VolatilityKernel\n",
|
||||
"from voltron.models import BMGP, VoltronGP\n",
|
||||
"from gpytorch.kernels import ScaleKernel, RBFKernel, MaternKernel"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"sns.set_style('white')\n",
|
||||
"# style.use('whitegrid')\n",
|
||||
"palette = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"sns.set(palette = palette, font_scale=2.0, style=\"white\", rc={\"lines.linewidth\": 4.0})"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"<torch._C.Generator at 0x7f2bf8c04650>"
|
||||
]
|
||||
},
|
||||
"execution_count": 4,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"torch.random.manual_seed(400)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"# Header"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 5,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ntest = 200\n",
|
||||
"ntrain = 200\n",
|
||||
"span = \"5year\"\n",
|
||||
"interval = 'day'\n",
|
||||
"T = 5."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 6,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# with open(\"../../spdr-data/XLF.pkl\", \"rb\") as handle:\n",
|
||||
"# raw_data = pickle.load(handle)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 7,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"raw_data = pd.read_pickle(\"../../spdr-data/XLF.pkl\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 8,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"tckrs = np.unique(raw_data[\"symbol\"])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Data Setup"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 9,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"data_list = []\n",
|
||||
"for i, tckr in enumerate(tckrs):\n",
|
||||
" data = raw_data[raw_data[\"symbol\"] == tckr]\n",
|
||||
"\n",
|
||||
" y = torch.FloatTensor(data['close_price'].to_numpy())\n",
|
||||
" data_list.append(y)\n",
|
||||
"\n",
|
||||
"ts = torch.linspace(0, T, data.shape[0]) + 1\n",
|
||||
"dt = ts[1] - ts[0]\n",
|
||||
"full_data = torch.stack(data_list)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 10,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"full_returns = torch.log(full_data[..., 1:] / full_data[..., :-1])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 11,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"device = torch.device(\"cpu\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 12,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# torch.cuda.set_device(device)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Now apply GCPV"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 13,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Automatic pdb calling has been turned ON\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"%pdb"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 24,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from voltron.models import MultitaskVariationalGP\n",
|
||||
"from gpytorch.priors import LKJCovariancePrior, SmoothedBoxPrior\n",
|
||||
"\n",
|
||||
"def get_and_fit_mtgpcv(train_x, log_returns):\n",
|
||||
" likelihood = VolatilityGaussianLikelihood(batch_shape=[log_returns.shape[0]], param=\"exp\")\n",
|
||||
"\n",
|
||||
" # corresponds to ICM\n",
|
||||
" model = MultitaskVariationalGP(\n",
|
||||
" inducing_points=train_x, \n",
|
||||
" covar_module=BMKernel().to(device), learn_inducing_locations=False,\n",
|
||||
" num_tasks = log_returns.shape[0], \n",
|
||||
" prior=LKJCovariancePrior(eta=5.0, n=log_returns.shape[0], sd_prior=SmoothedBoxPrior(0.05, 1.0))\n",
|
||||
" )\n",
|
||||
" model = model.to(device)\n",
|
||||
" model.initialize_variational_parameters(likelihood=likelihood, x=train_x, y=log_returns.t())\n",
|
||||
" \n",
|
||||
" model = model.to(device)\n",
|
||||
" likelihood = likelihood.to(device)\n",
|
||||
"\n",
|
||||
" # this is for running the notebook in our testing framework\n",
|
||||
" import os\n",
|
||||
" smoke_test = ('CI' in os.environ)\n",
|
||||
"# training_iterations = 2 if smoke_test else 300\n",
|
||||
" training_iterations = 1\n",
|
||||
"\n",
|
||||
"\n",
|
||||
" # Find optimal model hyperparameters\n",
|
||||
" model.train()\n",
|
||||
" likelihood.train()\n",
|
||||
"\n",
|
||||
" # Use the adam optimizer\n",
|
||||
" optimizer = torch.optim.Adam([\n",
|
||||
" {\"params\": model.parameters()}, \n",
|
||||
" ], lr=0.01)\n",
|
||||
"\n",
|
||||
" # \"Loss\" for GPs - the marginal log likelihood\n",
|
||||
" # num_data refers to the number of training datapoints\n",
|
||||
" mll = gpytorch.mlls.VariationalELBO(likelihood, model, train_x.shape[0])\n",
|
||||
" \n",
|
||||
" batched_train_x = train_x#[:-1]\n",
|
||||
" \n",
|
||||
" print_every = 50\n",
|
||||
" for i in range(training_iterations):\n",
|
||||
" # Zero backpropped gradients from previous iteration\n",
|
||||
" optimizer.zero_grad()\n",
|
||||
" # Get predictive output\n",
|
||||
" output = model(batched_train_x)\n",
|
||||
" # Calc loss and backprop gradients\n",
|
||||
" loss = -mll(output, log_returns.t())\n",
|
||||
" loss.backward()\n",
|
||||
" if i % print_every == 0:\n",
|
||||
" print('Iter %d/%d - Loss: %.3f' % (i + 1, training_iterations, loss.item()))\n",
|
||||
" optimizer.step()\n",
|
||||
" \n",
|
||||
" model.eval();\n",
|
||||
" likelihood.eval();\n",
|
||||
" predictive = model(train_x)\n",
|
||||
" # pred_scale = likelihood(predictive).scale.mean(0).detach()\n",
|
||||
" samples = likelihood(predictive).scale.detach()\n",
|
||||
" return samples.mean(0) / dt**0.5"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 25,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from voltron.models import MultitaskBMGP\n",
|
||||
"from gpytorch.likelihoods import MultitaskGaussianLikelihood\n",
|
||||
"\n",
|
||||
"def get_and_fit_volmodel(train_x, pred_scale):\n",
|
||||
" prior = LKJCovariancePrior(eta=5.0, n=len(tckrs), sd_prior=SmoothedBoxPrior(0.05, 1.0))\n",
|
||||
"\n",
|
||||
" vol_lh = MultitaskGaussianLikelihood(num_tasks=pred_scale.shape[-1])\n",
|
||||
" vol_lh.noise.data = torch.tensor([1e-6])\n",
|
||||
" vol_model = MultitaskBMGP(train_x, pred_scale.log(), vol_lh, prior=prior).to(device)\n",
|
||||
"\n",
|
||||
" optimizer = torch.optim.Adam([\n",
|
||||
" {'params': vol_model.parameters()}, # Includes GaussianLikelihood parameters\n",
|
||||
" ], lr=0.01)\n",
|
||||
"\n",
|
||||
" # \"Loss\" for GPs - the marginal log likelihood\n",
|
||||
" mll = gpytorch.mlls.ExactMarginalLogLikelihood(vol_lh, vol_model)\n",
|
||||
"\n",
|
||||
" for i in range(500):\n",
|
||||
" # Zero gradients from previous iteration\n",
|
||||
" optimizer.zero_grad()\n",
|
||||
" # Output from model\n",
|
||||
" output = vol_model(train_x)\n",
|
||||
" # Calc loss and backprop gradients\n",
|
||||
" loss = -mll(output, pred_scale.log())\n",
|
||||
" loss.backward()\n",
|
||||
" if i % 50 == 0:\n",
|
||||
" print(loss.item(), vol_model.covar_module.data_covar_module.raw_vol.item())\n",
|
||||
" optimizer.step()\n",
|
||||
"\n",
|
||||
" print((vol_model.covar_module.data_covar_module.raw_vol))\n",
|
||||
" \n",
|
||||
" return vol_model, vol_lh"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 26,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from voltron.means import LogLinearMean\n",
|
||||
"\n",
|
||||
"def get_and_fit_datamodel(train_x, train_y, pred_scale, vol_model, vol_lh):\n",
|
||||
" \n",
|
||||
" voltron_lh = gpytorch.likelihoods.GaussianLikelihood().to(device)\n",
|
||||
" voltron = VoltronGP(train_x, train_y.log(), voltron_lh, pred_scale.t()).to(device)\n",
|
||||
" # voltron.mean_module = gpytorch.means.LinearMean(1, batch_shape=torch.Size((6,)))\n",
|
||||
" voltron.mean_module = LogLinearMean(input_size=1, batch_shape=torch.Size((pred_scale.shape[1],))).to(device)\n",
|
||||
" voltron.mean_module.initialize_from_data(train_x, train_y.log())\n",
|
||||
" voltron.likelihood.raw_noise.data = torch.tensor([1e-6]).to(device)\n",
|
||||
" voltron.vol_lh = vol_lh.to(device)\n",
|
||||
" voltron.vol_model = vol_model.to(device)\n",
|
||||
"\n",
|
||||
" grad_flags = [False, True, True, True, *[False] * len(list(vol_model.named_parameters()))]\n",
|
||||
"\n",
|
||||
" for idx, (n, p) in enumerate(voltron.named_parameters()):\n",
|
||||
" # print(n)\n",
|
||||
" p.requires_grad = grad_flags[idx]\n",
|
||||
"\n",
|
||||
" voltron.train();\n",
|
||||
" voltron_lh.train();\n",
|
||||
" voltron.vol_lh.train();\n",
|
||||
" voltron.vol_model.train();\n",
|
||||
"\n",
|
||||
" # Use the adam optimizer\n",
|
||||
" optimizer = torch.optim.Adam([\n",
|
||||
" {'params': voltron.parameters()}, # Includes GaussianLikelihood parameters\n",
|
||||
" ], lr=0.1)\n",
|
||||
"\n",
|
||||
" # \"Loss\" for GPs - the marginal log likelihood\n",
|
||||
" mll = gpytorch.mlls.ExactMarginalLogLikelihood(voltron_lh, voltron)\n",
|
||||
"\n",
|
||||
" for i in range(500):\n",
|
||||
" # Zero gradients from previous iteration\n",
|
||||
" optimizer.zero_grad()\n",
|
||||
" # Output from model\n",
|
||||
" output = voltron(train_x)\n",
|
||||
" # Calc loss and backprop gradients\n",
|
||||
" loss = -mll(output, train_y.log()).sum()\n",
|
||||
" loss.backward()\n",
|
||||
" if i % 50 == 0:\n",
|
||||
" print(loss.item())\n",
|
||||
" optimizer.step()\n",
|
||||
" return voltron"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 27,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def predict(test_x, voltron, nvol=10, npx=10, num_tasks=len(tckrs)):\n",
|
||||
" vol_paths = torch.zeros(num_tasks, nvol, test_x.shape[0])\n",
|
||||
" px_paths = torch.zeros(num_tasks, npx*nvol, test_x.shape[0])\n",
|
||||
"\n",
|
||||
" voltron.vol_model.eval();\n",
|
||||
" voltron.eval();\n",
|
||||
"\n",
|
||||
" for vidx in range(nvol):\n",
|
||||
" print(vidx)\n",
|
||||
" vol_pred = voltron.vol_model(test_x).sample().exp()\n",
|
||||
" vol_paths[:, vidx, :] = vol_pred.detach().t().cpu()\n",
|
||||
"\n",
|
||||
" px_pred = voltron.GeneratePrediction(test_x, vol_pred.t(), npx)\n",
|
||||
" px_pred = px_pred.exp()\n",
|
||||
" px_paths[:, vidx*npx:(vidx*npx+npx), :] = px_pred.detach().transpose(-1, -2).cpu()\n",
|
||||
" return px_paths, vol_paths"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 28,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"8 8\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"train_splits = list(range(252, data.shape[0], int(252/2)))\n",
|
||||
"eval_splits = list(range(252+int(252/2), data.shape[0], int(252/2))) + [-1]\n",
|
||||
"print(len(train_splits), len(eval_splits))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 29,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"train_splits.insert(0, 0)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 30,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"[0, 252, 378, 504, 630, 756, 882, 1008, 1134]"
|
||||
]
|
||||
},
|
||||
"execution_count": 30,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"train_splits"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 31,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ts = ts.to(device)\n",
|
||||
"full_returns = full_returns.to(device)\n",
|
||||
"full_data = full_data.to(device)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 32,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ts = ts[1:]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 33,
|
||||
"metadata": {
|
||||
"collapsed": true,
|
||||
"jupyter": {
|
||||
"outputs_hidden": true
|
||||
},
|
||||
"tags": []
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Now running 0 252\n",
|
||||
"Iter 1/1 - Loss: 164.411\n",
|
||||
"1.1449265480041504 -1.3862943649291992\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/home/greg_b/miniconda3/envs/rpp/lib/python3.8/site-packages/gpytorch/utils/cholesky.py:40: NumericalWarning: A not p.d., added jitter of 1.0e-06 to the diagonal\n",
|
||||
" warnings.warn(\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"0.928978681564331 -1.0333786010742188\n",
|
||||
"0.709404468536377 -1.0833930969238281\n",
|
||||
"0.4720747470855713 -1.344559669494629\n",
|
||||
"0.22222964465618134 -1.66531503200531\n",
|
||||
"-0.0330703966319561 -1.9696232080459595\n",
|
||||
"-0.28592270612716675 -2.22033953666687\n",
|
||||
"-0.5303149819374084 -2.389354944229126\n",
|
||||
"-0.7619909644126892 -2.475860834121704\n",
|
||||
"-0.9778973460197449 -2.4871413707733154\n",
|
||||
"Parameter containing:\n",
|
||||
"tensor([-2.4095], requires_grad=True)\n",
|
||||
"22.37977409362793\n",
|
||||
"22.356121063232422\n",
|
||||
"22.344392776489258\n",
|
||||
"22.33585548400879\n",
|
||||
"22.329547882080078\n",
|
||||
"22.324800491333008\n",
|
||||
"22.321184158325195\n",
|
||||
"22.31840705871582\n",
|
||||
"22.31624984741211\n",
|
||||
"22.314552307128906\n",
|
||||
"0\n",
|
||||
"1\n",
|
||||
"2\n",
|
||||
"3\n",
|
||||
"4\n",
|
||||
"5\n",
|
||||
"6\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"ename": "KeyboardInterrupt",
|
||||
"evalue": "",
|
||||
"output_type": "error",
|
||||
"traceback": [
|
||||
"\u001b[0;31m---------------------------------------------------------------------------\u001b[0m",
|
||||
"\u001b[0;31mKeyboardInterrupt\u001b[0m Traceback (most recent call last)",
|
||||
"\u001b[0;32m<ipython-input-33-b50faf988533>\u001b[0m in \u001b[0;36m<module>\u001b[0;34m\u001b[0m\n\u001b[1;32m 19\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 20\u001b[0m \u001b[0mest_covar\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mmodel\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mvol_model\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mcovar_module\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mtask_covar_module\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mcovar_matrix\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mevaluate\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mcpu\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 21\u001b[0;31m \u001b[0mx_paths\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mv_paths\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mpredict\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mts\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0mtrain_end\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0mtest_end\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mmodel\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 22\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 23\u001b[0m \u001b[0mcovar_list\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mappend\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mest_covar\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m<ipython-input-27-e74218f72c37>\u001b[0m in \u001b[0;36mpredict\u001b[0;34m(test_x, voltron, nvol, npx, num_tasks)\u001b[0m\n\u001b[1;32m 8\u001b[0m \u001b[0;32mfor\u001b[0m \u001b[0mvidx\u001b[0m \u001b[0;32min\u001b[0m \u001b[0mrange\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mnvol\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 9\u001b[0m \u001b[0mprint\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mvidx\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 10\u001b[0;31m \u001b[0mvol_pred\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mvoltron\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mvol_model\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mtest_x\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0msample\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mexp\u001b[0m\u001b[0;34m(\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 11\u001b[0m \u001b[0mvol_paths\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mvidx\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m:\u001b[0m\u001b[0;34m]\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mvol_pred\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mdetach\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mt\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mcpu\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 12\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/gpytorch/distributions/multivariate_normal.py\u001b[0m in \u001b[0;36msample\u001b[0;34m(self, sample_shape, base_samples)\u001b[0m\n\u001b[1;32m 220\u001b[0m \u001b[0;32mdef\u001b[0m \u001b[0msample\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0msample_shape\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mtorch\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mSize\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mbase_samples\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0;32mNone\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 221\u001b[0m \u001b[0;32mwith\u001b[0m \u001b[0mtorch\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mno_grad\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 222\u001b[0;31m \u001b[0;32mreturn\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mrsample\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0msample_shape\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0msample_shape\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mbase_samples\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mbase_samples\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 223\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 224\u001b[0m \u001b[0;32mdef\u001b[0m \u001b[0mto_data_independent_dist\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/gpytorch/distributions/multitask_multivariate_normal.py\u001b[0m in \u001b[0;36mrsample\u001b[0;34m(self, sample_shape, base_samples)\u001b[0m\n\u001b[1;32m 237\u001b[0m \u001b[0mbase_samples\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mbase_samples\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mview\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0msample_shape\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m*\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mloc\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mshape\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 238\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 239\u001b[0;31m \u001b[0msamples\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0msuper\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mrsample\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0msample_shape\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0msample_shape\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mbase_samples\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mbase_samples\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 240\u001b[0m \u001b[0;32mif\u001b[0m \u001b[0;32mnot\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_interleaved\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 241\u001b[0m \u001b[0;31m# flip shape of last two dimensions\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/gpytorch/distributions/multivariate_normal.py\u001b[0m in \u001b[0;36mrsample\u001b[0;34m(self, sample_shape, base_samples)\u001b[0m\n\u001b[1;32m 179\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 180\u001b[0m \u001b[0;31m# Get samples\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 181\u001b[0;31m \u001b[0mres\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mcovar\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mzero_mean_mvn_samples\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mnum_samples\u001b[0m\u001b[0;34m)\u001b[0m \u001b[0;34m+\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mloc\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0munsqueeze\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;36m0\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 182\u001b[0m \u001b[0mres\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mres\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mview\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0msample_shape\u001b[0m \u001b[0;34m+\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mloc\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mshape\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 183\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/gpytorch/lazy/lazy_tensor.py\u001b[0m in \u001b[0;36mzero_mean_mvn_samples\u001b[0;34m(self, num_samples)\u001b[0m\n\u001b[1;32m 2115\u001b[0m \u001b[0mcovar_root\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mevaluate\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0msqrt\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 2116\u001b[0m \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 2117\u001b[0;31m \u001b[0mcovar_root\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mroot_decomposition\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mroot\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m 2118\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 2119\u001b[0m base_samples = torch.randn(\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/gpytorch/utils/memoize.py\u001b[0m in \u001b[0;36mg\u001b[0;34m(self, *args, **kwargs)\u001b[0m\n\u001b[1;32m 57\u001b[0m \u001b[0mkwargs_pkl\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mpickle\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mdumps\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 58\u001b[0m \u001b[0;32mif\u001b[0m \u001b[0;32mnot\u001b[0m \u001b[0m_is_in_cache\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mcache_name\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mkwargs_pkl\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mkwargs_pkl\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 59\u001b[0;31m \u001b[0;32mreturn\u001b[0m \u001b[0m_add_to_cache\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mcache_name\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mmethod\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mkwargs_pkl\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mkwargs_pkl\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 60\u001b[0m \u001b[0;32mreturn\u001b[0m \u001b[0m_get_from_cache\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mcache_name\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mkwargs_pkl\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mkwargs_pkl\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 61\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/gpytorch/lazy/lazy_tensor.py\u001b[0m in \u001b[0;36mroot_decomposition\u001b[0;34m(self, method)\u001b[0m\n\u001b[1;32m 1757\u001b[0m \u001b[0mroot\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mU\u001b[0m \u001b[0;34m*\u001b[0m \u001b[0mS\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0msqrt\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0munsqueeze\u001b[0m\u001b[0;34m(\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 1758\u001b[0m \u001b[0;32melif\u001b[0m \u001b[0mmethod\u001b[0m \u001b[0;34m==\u001b[0m \u001b[0;34m\"lanczos\"\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1759\u001b[0;31m \u001b[0mroot\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_root_decomposition\u001b[0m\u001b[0;34m(\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 1760\u001b[0m \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 1761\u001b[0m \u001b[0;32mraise\u001b[0m \u001b[0mRuntimeError\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34mf\"Unknown root decomposition method '{method}'\"\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/gpytorch/lazy/lazy_tensor.py\u001b[0m in \u001b[0;36m_root_decomposition\u001b[0;34m(self)\u001b[0m\n\u001b[1;32m 614\u001b[0m \u001b[0;34m(\u001b[0m\u001b[0mTensor\u001b[0m \u001b[0;32mor\u001b[0m \u001b[0mLazyTensor\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m \u001b[0mThe\u001b[0m \u001b[0mroot\u001b[0m \u001b[0mof\u001b[0m \u001b[0mthe\u001b[0m \u001b[0mroot\u001b[0m \u001b[0mdecomposition\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 615\u001b[0m \"\"\"\n\u001b[0;32m--> 616\u001b[0;31m res, _ = RootDecomposition.apply(\n\u001b[0m\u001b[1;32m 617\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mrepresentation_tree\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 618\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_root_decomposition_size\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/gpytorch/functions/_root_decomposition.py\u001b[0m in \u001b[0;36mforward\u001b[0;34m(ctx, representation_tree, max_iter, dtype, device, batch_shape, matrix_shape, root, inverse, initial_vectors, *matrix_args)\u001b[0m\n\u001b[1;32m 46\u001b[0m \u001b[0mmatmul_closure\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mlazy_tsr\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_matmul\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 47\u001b[0m \u001b[0;31m# Do lanczos\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 48\u001b[0;31m q_mat, t_mat = lanczos.lanczos_tridiag(\n\u001b[0m\u001b[1;32m 49\u001b[0m \u001b[0mmatmul_closure\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 50\u001b[0m \u001b[0mctx\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mmax_iter\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/gpytorch/utils/lanczos.py\u001b[0m in \u001b[0;36mlanczos_tridiag\u001b[0;34m(matmul_closure, max_iter, dtype, device, matrix_shape, batch_shape, init_vecs, num_init_vecs, tol)\u001b[0m\n\u001b[1;32m 98\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 99\u001b[0m \u001b[0;31m# Compute next alpha value\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 100\u001b[0;31m \u001b[0mr_vec\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mmatmul_closure\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mq_curr_vec\u001b[0m\u001b[0;34m)\u001b[0m \u001b[0;34m-\u001b[0m \u001b[0mq_prev_vec\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mmul\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mbeta_prev\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 101\u001b[0m \u001b[0malpha_curr\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mq_curr_vec\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mmul\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mr_vec\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0msum\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mdim_dimension\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mkeepdim\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0;32mTrue\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 102\u001b[0m \u001b[0;31m# Copy over to t_mat\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/gpytorch/lazy/sum_lazy_tensor.py\u001b[0m in \u001b[0;36m_matmul\u001b[0;34m(self, rhs)\u001b[0m\n\u001b[1;32m 38\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 39\u001b[0m \u001b[0;32mdef\u001b[0m \u001b[0m_matmul\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mrhs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 40\u001b[0;31m \u001b[0;32mreturn\u001b[0m \u001b[0msum\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mlazy_tensor\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_matmul\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mrhs\u001b[0m\u001b[0;34m)\u001b[0m \u001b[0;32mfor\u001b[0m \u001b[0mlazy_tensor\u001b[0m \u001b[0;32min\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mlazy_tensors\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 41\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 42\u001b[0m \u001b[0;32mdef\u001b[0m \u001b[0m_mul_constant\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mother\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/gpytorch/lazy/sum_lazy_tensor.py\u001b[0m in \u001b[0;36m<genexpr>\u001b[0;34m(.0)\u001b[0m\n\u001b[1;32m 38\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 39\u001b[0m \u001b[0;32mdef\u001b[0m \u001b[0m_matmul\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mrhs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 40\u001b[0;31m \u001b[0;32mreturn\u001b[0m \u001b[0msum\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mlazy_tensor\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_matmul\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mrhs\u001b[0m\u001b[0;34m)\u001b[0m \u001b[0;32mfor\u001b[0m \u001b[0mlazy_tensor\u001b[0m \u001b[0;32min\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mlazy_tensors\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 41\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 42\u001b[0m \u001b[0;32mdef\u001b[0m \u001b[0m_mul_constant\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mother\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/gpytorch/lazy/interpolated_lazy_tensor.py\u001b[0m in \u001b[0;36m_matmul\u001b[0;34m(self, rhs)\u001b[0m\n\u001b[1;32m 179\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 180\u001b[0m \u001b[0;31m# base_lazy_tensor * right_interp^T * rhs\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 181\u001b[0;31m \u001b[0mbase_res\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mbase_lazy_tensor\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_matmul\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mright_interp_res\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 182\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 183\u001b[0m \u001b[0;31m# left_interp * base_lazy_tensor * right_interp^T * rhs\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/gpytorch/lazy/kronecker_product_lazy_tensor.py\u001b[0m in \u001b[0;36m_matmul\u001b[0;34m(self, rhs)\u001b[0m\n\u001b[1;32m 235\u001b[0m \u001b[0mrhs\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mrhs\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0munsqueeze\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m-\u001b[0m\u001b[0;36m1\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 236\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 237\u001b[0;31m \u001b[0mres\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0m_matmul\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mlazy_tensors\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mshape\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mrhs\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mcontiguous\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\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 238\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 239\u001b[0m \u001b[0;32mif\u001b[0m \u001b[0mis_vec\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/gpytorch/lazy/kronecker_product_lazy_tensor.py\u001b[0m in \u001b[0;36m_matmul\u001b[0;34m(lazy_tensors, kp_shape, rhs)\u001b[0m\n\u001b[1;32m 41\u001b[0m \u001b[0mfactor\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mlazy_tensor\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_matmul\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mres\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 42\u001b[0m \u001b[0mfactor\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mfactor\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mview\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0moutput_batch_shape\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mlazy_tensor\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0msize\u001b[0m\u001b[0;34m(\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[0;36m1\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mnum_cols\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mtranspose\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m-\u001b[0m\u001b[0;36m3\u001b[0m\u001b[0;34m,\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[0;32m---> 43\u001b[0;31m \u001b[0mres\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mfactor\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mreshape\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0moutput_batch_shape\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m-\u001b[0m\u001b[0;36m1\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mnum_cols\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 44\u001b[0m \u001b[0;32mreturn\u001b[0m \u001b[0mres\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 45\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;31mKeyboardInterrupt\u001b[0m: "
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"> \u001b[0;32m/home/greg_b/miniconda3/envs/rpp/lib/python3.8/site-packages/gpytorch/lazy/kronecker_product_lazy_tensor.py\u001b[0m(43)\u001b[0;36m_matmul\u001b[0;34m()\u001b[0m\n",
|
||||
"\u001b[0;32m 41 \u001b[0;31m \u001b[0mfactor\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mlazy_tensor\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_matmul\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mres\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0m\u001b[0;32m 42 \u001b[0;31m \u001b[0mfactor\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mfactor\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mview\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0moutput_batch_shape\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mlazy_tensor\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0msize\u001b[0m\u001b[0;34m(\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[0;36m1\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mnum_cols\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mtranspose\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m-\u001b[0m\u001b[0;36m3\u001b[0m\u001b[0;34m,\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[0m\u001b[0;32m---> 43 \u001b[0;31m \u001b[0mres\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mfactor\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mreshape\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0moutput_batch_shape\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m-\u001b[0m\u001b[0;36m1\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mnum_cols\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0m\u001b[0;32m 44 \u001b[0;31m \u001b[0;32mreturn\u001b[0m \u001b[0mres\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0m\u001b[0;32m 45 \u001b[0;31m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0m\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdin",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"ipdb> q\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"covar_list = []\n",
|
||||
"x_list = []\n",
|
||||
"v_list = []\n",
|
||||
"# for i in range(len(eval_splits)):\n",
|
||||
"i = 0\n",
|
||||
"\n",
|
||||
"train_start = train_splits[i]\n",
|
||||
"train_end = train_splits[i+1]\n",
|
||||
"test_end = eval_splits[i]\n",
|
||||
"\n",
|
||||
"print(\"Now running \", train_start, train_end)\n",
|
||||
"torch.cuda.empty_cache()\n",
|
||||
"\n",
|
||||
"pred_scale = get_and_fit_mtgpcv(ts[train_start:train_end], full_returns[:,train_start:train_end])\n",
|
||||
"vol_model, vol_lh = get_and_fit_volmodel(ts[train_start:train_end], pred_scale)\n",
|
||||
"model = get_and_fit_datamodel(\n",
|
||||
" ts[train_start:train_end], full_data[:,train_start:train_end], pred_scale, vol_model, vol_lh\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"est_covar = model.vol_model.covar_module.task_covar_module.covar_matrix.evaluate().cpu()\n",
|
||||
"x_paths, v_paths = predict(ts[train_end:test_end], model)\n",
|
||||
"\n",
|
||||
"covar_list.append(est_covar)\n",
|
||||
"x_list.append(x_paths)\n",
|
||||
"v_list.append(v_paths)\n",
|
||||
"\n",
|
||||
"del model"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 34,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"torch.Size([252, 30])"
|
||||
]
|
||||
},
|
||||
"execution_count": 34,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"pred_scale.shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 33,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"tensor([[-0.0142, 0.0045, 0.0072, ..., 0.0000, 0.0010, 0.0114],\n",
|
||||
" [-0.0046, 0.0045, -0.0065, ..., 0.0019, 0.0009, 0.0099],\n",
|
||||
" [-0.0046, -0.0043, -0.0042, ..., 0.0022, 0.0008, -0.0038],\n",
|
||||
" ...,\n",
|
||||
" [-0.0025, 0.0002, -0.0009, ..., 0.0079, -0.0026, 0.0027],\n",
|
||||
" [-0.0121, -0.0007, -0.0030, ..., 0.0070, 0.0006, -0.0050],\n",
|
||||
" [-0.0127, -0.0009, -0.0047, ..., 0.0085, 0.0009, -0.0081]])"
|
||||
]
|
||||
},
|
||||
"execution_count": 33,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"full_returns[:,train_start:train_end]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 22,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"quantiles = torch.stack(\n",
|
||||
" [(x_paths[..., -1] > full_data[:, tend-1].unsqueeze(-1).cpu()\n",
|
||||
" ).sum(1)/100 for x_paths, tend in zip(x_list, eval_splits)]\n",
|
||||
")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 23,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"torch.save(f=\"finance_mt_preds_3.pt\", obj={\"quantiles\": quantiles, \"x_paths\": x_list, \"v_paths\": v_list, \"covar\": covar_list})"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 24,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import pandas as pd"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 25,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df = pd.DataFrame({\"quantiles\": quantiles.reshape(-1).numpy(), \"date\": np.repeat([np.array(eval_splits)],30)})"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 26,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df[\"date\"] = df[\"date\"] / (1259/5)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 27,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAApQAAAFXCAYAAAAcW967AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAADOaUlEQVR4nOz9eZxcVZ3/j7/OXWrpfU2ns5CFpLMB2SAQBESWMaL4AREVCbjMgOPKzEd+jCsDfBxl3EZG/TCgo6NE8OOoIH6VTRhBNBIISRqyrx2ydKc7vVbXfs/5/XHuObV3V1dVV91bnOfjwaND1b23T3VV3fu67+X1JowxBoVCoVAoFAqFokC0Si9AoVAoFAqFQuFulKBUKBQKhUKhUBSFEpQKhUKhUCgUiqJQglKhUCgUCoVCURRKUFYx8Xgcx44dQzwer/RSFAqFQqGYdtR1r3IYlV6AYvo4fvw4/uZv/gY/+9nPMHPmzEovR6FQKBSKaaW3txc33ngjnn76acybN6/Sy3lToQRlFdPf3w8AuPHGGyu8EoVCoVAoykd/f78SlGVGCcoqpr29HQBUhFKhUCgUbwpEhFJc/xTlQwnKKkbXdQDAzJkzMWfOnAqvRqFQKBSK8iCuf4ryoZpyFAqFQqFQKBRFoQSlQqFQKBQKhaIolKBUKBQKhUKhUBSFEpQKhUKhUCgUiqJQglKhUCgUCoVCURRKUCoUCoVCoVAoikIJSoVCoVAoFApFUSgfSoVCoVA4EsYY4ge7YfUdBWMUWk0DzGXnQfPXVXppCoUiDSUoFQqFQuFI4gd3IPzXJ1Ies069gZoNN1doRQqFIhcq5a1QKBQKRxI/cRgAEBs7iejpA2DUgjVwHCwSqvDKFApFOkpQKhQKhcKRsPA4AICGBhEPDoBZMfvxYCWXpVAosqAEpUKhUCgcCQtxQclo3P4ZBQDQcKBia1IoFNlRglKhUCgUjkREKJkVtX/aEUpbaCoUCuegBKVCoVAoHAeLRcFiETBKwajFHxTCMqQilAqF01CCUqFQKBSOgwXH+E87zQ0kRSjt5xQKhXNQglKhUCgUjoOOjwIAWDxZUEZSnlMoFM5BCUqFQqFQOA4WHOE/bRHJ/22nvJWgVCgchxKUCoVCoXAcdMwWlPFw4jH73zQwXIklKRSKCVCTchSKSWCUIvLqc4gf7AYAGGeeA+/ay0EIqfDKFAp7POGBHbAGe2HMXgRjzqJKL6kk0MAQgPQIZQyMUiASBItGQDzeSi1PoVCkoSKUCsUkWKeOIrbnZd5xGosgtudlWH09lV6WQgEAiO3divBLTyC2fxtCf/xvWKfeqPSSSgId44KSJkUogUTEko4Nln1NCoUiN0pQKhSTQEf4hSse6EM80JfymEJRaWK7twBICK/YgR2VXE5JYIyBjQ6CMYDF+JhF0lDPnxOCclR9BxUKJ6EEpUIxGcKqhFEwRlMeUygqDR0XzSv2FJnhgUoupySw4Bjv7qYxMGqBGDpIYxMAgMa5wKSjpyu4QoVCkY6qoVQoJkHaljAr8zGFwikwxn+AVXghxUNHuCimdnQSNbVATQ2ARMSyGoSzQlFNKEGpUExGzI5GUgsASX1MoXAKInpeBdDhfgAAiwcBAKS2FqS2lj8XC6Zso1AonIFKeSsUk8BivMuUp7x5lFJFKBWOw45QErjffYAOneI/bfGI2qQIZTwExhhoYAgspr6HCoVTUIJSoZgEedGilh2lTIhMhcIpyJudKkh5W0JQRpMilIYB4vfzhp1YCGBMRSkVCgehBKVCMRkyQhmXF21ElaBUOAzmfiEJAMyKg44McOFop7xRV5fyk8bGAQDWYF8llqhQKLKgBKVCMQksZvvgUQug8dTHFAqnkNQ05mbo0CmAUZ7aphTE5wMxTQAAEYIyOm5v21uxdSoUilRUU45CMQkswsUjs8UkADAVoVQ4DGFp5fYaSus0F4lCNKK+Xj5H7H8zEaEcOFnexSkUipyoCKVCMQksmhCUTNRQRlSEUuEwbEHp9hpKOnCc/4wGACREJAApLml0nDfmjPSrxhyFwiEoQalQTABjDCwS4uVpNM6NlhnAIkGwKqlZU1QH4mbH/RFKHnWUgrKhQT5HTBOkxs8dF2JBgDFYgyrtrVA4ASUoFYqJiEV45IfFubhkjNeqMSqbdRQKR1AFNZQsHAQdPQ1GuWAkBCkpbwAg9VxgCsFJ+4+Xe5kKhSILSlAqFBPAwrzLNKV+ksZSnlMonIAsx3BxytsaOAEAoLEAv3mrrQPR9dSNGoSgHLP3UYJSoXACSlAqFBNAQzwKwpJmd4t/s/B4RdakUGSlCiKUlh1tpJHMdLeANDba24zJfVT5iUJReZSgVCgmQIpGmjRq0f43DSlBqagsQkgxlohQuhmr/w0AiegjbPGYQm0tiKGDxiNgVpTXM48OlnGVCoUiG8o2SKGYAGaLRmYlOkllhNKOXrodxhhiu15CvGcP4PPDu/IS6K2dlV6WIh+kiKSJWd4uFZbMsmCdPgnGEtFHkkVQEkJ42ntwCDQSgF7TAqv/DWiNreVeskKhSEJFKBWKCWBBfmFLTXlzcVktgtI6cQiRbf8Da/AkrBOHEHrh1yqF6BYsu7aXUQA09TGXQQdPAlYcLB4Eo3EQrxfE58u6LWngQtOKjPKfp46VbZ0KhSI7SlAqFBNAg/yCxaxER7cQlOI5tyNsWuKBPjArDjY+WnX1odbQKUR3v4x475FKL6WkMCkomTQ2Zy4VlFafne4W0cmmLOluG1lHKRpzTr0xzatTKBSToVLeCsUEJCKUySlvO0I5Xh2CEnE7hR+PgDELBIZ8rBqwBk4g+PQmmQr2XXAVzEUrK7yqEhG3R4GypJR33KWC8lSqoERDbkGJhgYQjYDFxsFoHDQwDBocg1ZTn3sfhUIxragIpUIxAUI0sniyoOTRSiE23Q6z/TQZS55VXj3TR2IHdgDUAo2FAADRPa9UeEWlgyWnvGWE0n03A4xSWAPHwRhgRfl3jjQ15dye6DpQV8/rLW0/ShWlVCgqixKUCkUOGKWgwTHeQZuc8o5H+YVsfBSM0gqusDQIQQlqcVEJgMWqZ7SkmKQiZkPT4VOVXE5pEeKRUW4GzsDrEF1WA0tHBsCiYTArAhaPgpgGUFMz4T4Z9kFKUCoUFUUJSoUiByw0xqM+NArGGB/7Zpo8vUij/CIecn+UUs4qZ3EZoUS0CqcAMfeL/3SYKFcQHpTMnY05oqmG2k02aGzk3dwTICKYYh9LTcxRKCqKEpQKRQ5oYJj/jNviyu8D7K5T8ZjYxtVEbEFJ43IiEIuEKrmi0iKidVVg/J1BSpd30k+X1cAm/CeFofkE9ZMCOTFnHIxR0OFT8uZIoVCUHyUoFYocsLFh/jPOL1LE5wfx++zHIinbuBkWsUdIWlUqKG1YFQpKWetql17IkoW4u2pgrX4RocztP5kOMU2Q2houJqPjAGNydKNCoSg/SlAqFDmgUlDaEUqfD/D57ce4yHR7hJIxJif+MBqTNXnVZhsEwLWG3xNiZU95Mxd1etPxUW5VReOgsSCIRoD6/Lq1RSRTNub0Kz9KhaJSKEGpUOSABoYAAMyy02h+P/8v6TE6NlSRtZWMWBSwYmDUsv+zx0oGq8O0PZnqjFAmmnJSfrqo09sasOd326IQ9fUgWp6XpvTGHBWhVCgqhvKhVChyIMQitTueiS0mkx9zu6CkaT6bIlXKqsS0PYWqjlDakUnqvpQ3HeDG+jRi10/WN+S9L5F1lHxfevokb6CbpKFHoVCUHhWhVCiywBgDGx3klkFxu56wpkZamYiUN9/GXRYtyTA5CcgWkrY9UrVMAUpG1IdWE1I4CrEsorAxF0UoB21BGbMjlA35C0r4/SCGAWZFuZ1XNAw2OjgNq1QoFJOhIpQFcujQIfzpT3/Ca6+9htdffx1HjhwBYwz33XcfNmzYUPBxf/vb3+KRRx7B3r17QSnFggULcN111+GGG26Alm8aSFE0LDjGL9aUp4OJoQOmCQAghg4Wj0tTaRYKgLh0QoesExVm7Ra3SEJwDMyyuIF0lcBoHIwBhKB6olhpKW8ZqXRJhJIxBnq61zYo53W7JM/6SQD8PayvB4aGQGMB6EYLrMFeaI2t07VkhUKRAyUoC+SRRx7BT3/605Ie8+6778bDDz8Mr9eL9evXwzAMbN68Gffccw82b96M++67D3oVXeCdDB3jUQ4aF/WTNQkB4q8BxsZA4yHoej3o6GnXjnxjok5URFwZA7OiIMQLFhgGqYYLc/o0GaLxxwyzsusqAbJEgaVFKN0iKMeG+GuwomBWjBua29Zc+ULq68GGhkCj49D9LaCDfcCCFdO0YoVCkQsV8iqQrq4u/O3f/i3+7d/+Dc888wzWrVtX1PGeeuopPPzww2hvb8fjjz+OBx54AN///vfx9NNP48wzz8QzzzyDTZs2lWj1ismgI1xQinQ3SZraQWr8Kc+Jbd0ItdODNGkyjnxdY+59XcmkiC6X2urkQr4OISRFDaVLUt5yilHMdhWoq5965Liuzj4Gt7+yhvpKtj6FQpE/KkJZINdff31Jj/fAAw8AAG6//XbMnz9fPt7W1oa77roLN910E37wgx/gpptuUqnvMkBHTwMAmD3/GUkNOfDbdZSyMce9wouODABIqhOF/Zp9TaDDA8CcxZVaWukQFjqUgjEKArjO+DsntnCUzTguS3nT4X7+0xaDpLY28RxjeGH7Cby6tx9Hennz2PyZ9VizpB2XrJoFzRaexBaUzD4GHR4o2/oVCkUCJSgdQG9vL3bu3AnTNLPWX65btw4dHR3o6+vD9u3bsWbNmgqs8s2FEJQi5Z18oRONOVRE8uxt3QaLRkADw2CMypQ3AFBbRFsj/ZVaWslgjKVG8WQEzx2CazISr03YBomUtzsEMx0RN262ub79PRsai+DBx3di15FUF4UdB09jx8HTeGlXH2599wo013t5Y46mgcYjXFiHA6DhcWi+WigLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 576x360 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(figsize=(8,5))\n",
|
||||
"violins = sns.violinplot(x='date', y='quantiles', data=df, color=palette[-1])\n",
|
||||
"for violin in violins.collections[::2]:\n",
|
||||
" violin.set_alpha(0.5)\n",
|
||||
" violin.set_edgecolor(palette[-2])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 28,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"<Figure size 576x720 with 0 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAK0AAAaaCAYAAAC4AP1sAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAAAl+klEQVR4nO3df5CdZX338W+yuwmEyA9hYalBTacQrTpgx4Gm2k7rYnHs09EW0WHSQKcWaNVWsR3aqvP4C9Rx0JRKbTG0qECwZah2sI4ds/oHA5S4Y/ixsSPaijwbORAQE8I+y+Zs9vmDJ1tC9iSbzTnZ88m+XjP84eE+97mC771yn3PtdZ9FU1NTUwVBFs/3AOBgiZY4oiWOaIkjWuK0jLbZbNbo6Gg1m83DOR44oJbRNhqNGhwcrEajcTjHAwfk8oA4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIlzhEX7cSu5rw8l8Ond74H0G5L+nrrgsvXz+m5t667pM2joROOuJmWI59oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1ridG20NhnSStdubJzrBkWbE498XTvTQiuiJY5oiSNa4oiWOKIljmiJI1riiJY4on0O97bN0LXLuPPBvW0zmGmJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJU5Ho7VDlU7o6G5cu1vpBJcHxBEtcURLHNESR7TEES1xREsc0RJHtMQRLXFESxzREke0xBEtcURLHNESR7TEES1xREsc0bbJfHxF6ULdOOprRtvkUDdxzuW5C3Xzp5mWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiVOy+02k5OTVVXVaDQO6QUmxrbP6Xmjo6Nzeu5cn5f43NHR0Tm9XoqBgYHq7d030UVTU1NTMz1heHi41qxZ0/GBQStDQ0O1YsWKfR5vGe34+HiNjIxUf39/9fT0dHyA8HwHPdNCt/JGjDiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiVOy2ibzWaNjo5Ws7kwv+ma7tUy2kajUYODg4d8WyRoN5cHxBEtcURLHNESR7TEES1xREsc0RJHtMQRLXFESxzREke0xBEtcURLHNESR7TEES1xREsc0RJHtMQRLXFESxzREke0xBEtcURLHNESR7TEES1xREsc0RJHtMQRLXFESxzREke0xBEtcURLHNESR7TEES1xREsc0RJHtMQRLXFESxzREke0xBEtcURLHNESR7TEES1xREsc0RJHtMQRLXFESxzREke0xBEtcURLHNESR7TEES1xREsc0RJHtMQRLXFESxzREke0xBEtcURLHNESR7TEES1xREsc0RJHtMQRLXFESxzRLnATu5oR53yu3o6ena63pK+3Lrh8fVvPeeu6S9p6vucz0xJHtMQRLXFESxzREke0xBEtcURLHNESR7TEES1xREsc0RJHtMQRLXFESxzREke0xBEtcRZ8tIkb+xa6Bb+xMXFj30K34Gda8oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiVOVLQ2DFIVtrHRJkSqwmZaqBItgURLHNESR7TEES1xREsc0RJHtMQRbQhL2P8jahl3IevEEnZV5jK2mZY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOB2L1rIjndKxZVw7Z+kUlwfEES1xREsc0RJHtMQRLXFESxzREke0HWA1sLNsbOwAq4GdZaYljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuK03CM2OTlZVVWNRmPOJ58Y2z7n585kdHTUOdusU2Nth4GBgert3TfRRVNTU1MzPWF4eLjWrFnTlheHuRgaGqoVK1bs83jLaMfHx2tkZKT6+/urp6en4wOE5zvomRa6lTdixBEtcURLHNESR7TEES1xREsc0RJHtMQRLXFESxzREke0xBEtcURLHNESR7TEES1xREsc0RJHtMQRLXFaRttsNmt0dLSaTV8DT3dpGW2j0ajBwcFDui0SdILLA+KIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1r2MrGrOS/PPRi9h+VViLGkr7cuuHz9nJ5767pL2jyamZlpiSNa4oiWOKIljmiJI1riiJY4oiXOgol2vlZ6ElaY0iyYFbH5WulJWGFKs2BmWo4coiWOaIkjWuKIljiiJY5oiSNa4oiWOKI9Ah3py78LZhl3ITnSl47NtMQRLXFESxzREke0xBEtcURLHNESR7TEEW0XO9KXY+fKMm4Xm+tybMJS7KEw0xJHtMQRLXFESxzREke0xBEtcURLHNESR7TEES1xREsc0RJHtMQRLXFESxzREke0xBEtcUQ7CzYYdhcbG2fhSL/faxozLXFESxzREke0xBEtcURLHNESR7TEES1xREsc0RJHtMQRLXFESxzREke0xBEtcURLnJbbbSYnJ6uqqtFoHLbBdNrE2PY5PW90dDTqufM53nYaGBio3t59E100NTU1NdMThoeHa82aNW0dBByMoaGhWrFixT6Pt4x2fHy8RkZGqr+/v3p6ejo+QHi+g55poVt5I0Yc0RJHtMQRLXFESxzREke0xBEtcURLHNESR7TEES1xREsc0RJHtMQRLXFESxzREke0xBEtcURLHNESp2W0zWazRkdHq9n0Ddx0l5bRNhqNGhwcPKJui8SRweUBcURLHNESR7TEES1xREsc0RJHtMQRLXFESxzREke0xBEtcURLHNESR7TEES1xREsc0RJHtMQRLXFESxzREke0xBEtcURLHNESR7TEES1xREsc0RJHtMQRLXFESxzREke0xBEtcURLHNESR7TEES1xREsc0RJHtMQRLXFESxzREke0xBEtcURLHNESR7TEES1xREsc0RJHtMQRLXFESxzREke0xBEtcURLHNESR7TEES1xREsc0RJHtMQRLXFESxzREke0xBEtcURLHNESR7TEES1xREsc0RJHtMQRLXFESxzREmfBRDuxq9nV52P2eud7AIfLkr7euuDy9W07363rLmnbuTg4C2am5cghWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOG2JthOb/GwcpJW2bGxs96bBKhsHac3lAXFESxzREke0xBEtcURLHNESR7TEES1xREsc0RJHtMQRLXFESxzREke0xBEtcURLHNESR7TEEW2X8DWos7dgvma02/ka1Nkz0xJHtMQRLXFESxzREke0xBEtcURLHNESR7TEES1xREsc0RJHtMQRLXFESxzREke0xBEtcbo22m7fmLcQx9ctf+au3djY7Rv9Ftr4qrpns2TXzrTQimiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiVOy7smTk5OVlVVo9GY1Ykmxra3Z0T/3+joaFvP6XztOefhNDAwUL29+ya6aGpqamqmJwwPD9eaNWs6PjBoZWhoqFasWLHP4y2jHR8fr5GRkerv76+enp6ODxCe76BnWuhW3ogRR7TEES1xREsc0RJHtMQRLXFESxzREke0xBEtcURLHNESR7TEES1xREsc0RJHtMQRLXFESxzREke0xGkZbbPZrNHR0Wo2m4dzPHBALaNtNBo1ODg469siweHi8oA4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4CzbaiV3NyHNT1TvfA5gvS/p664LL13fk3Leuu6Qj5+VZC3amJZdoiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5oiSNa4oiWOKIljmiJI1riiJY4oiWOaIkjWuKIljiiJY5Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 216x1728 with 8 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.figure(figsize = (8, 10))\n",
|
||||
"ridge_plot = sns.FacetGrid(df, row=\"date\")\n",
|
||||
"# Use map function to make density pslot in each element of the grid.\n",
|
||||
"ridge_plot.map(sns.histplot, \"quantiles\", clip_on=False)\n",
|
||||
"ridge_plot.set_titles(\"\")\n",
|
||||
"ridge_plot.set(yticks=[], ylabel=\"\")\n",
|
||||
"plt.tight_layout()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 29,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAWcAAAENCAYAAADT16SxAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAAAzBElEQVR4nO3de3hU5bU/8O+eScgNkoAZcwE1RIlAkrYBDA2e2qpgU3pSPVjbQ4GotdBzEOqvBS+ttSaUtqhBihZbDo9AhSTWo9JK1SDFHktrJFKLJCGEa1CIEwYkQBJCktn790c6qXH2eudCdpgh38/z5Hl01r7NnmFl591rr1czDMMAERGFFNvFPgAiIvLG5ExEFIKYnImIQhCTMxFRCGJyJiIKQUzOREQhKOJiHwARUaAOHTqE7du3o6amBrW1tWhsbIRhGFi5ciUKCgqC3u7mzZtRUVGBhoYG6LqO0aNH4/bbb8fMmTNhsw3stSyTMxGFnYqKCjz33HP9us2SkhKUl5cjKioK+fn5iIiIQFVVFZYsWYKqqiqsXLkSdru9X/epwuRMRGEnMzMT99xzD7Kzs5GdnY2HH34Y1dXVQW9vy5YtKC8vh8PhwMaNG5Geng4AOHHiBIqKirB161Zs3LgRd955Zz+9A9+YnIko7Nxxxx39ur3Vq1cDABYvXtybmAEgKSkJxcXFmDNnDtasWYM5c+YM2PDGgA2ibN68Gd/61rcwceJE5ObmYsaMGSgrK4Ou6wN1CEREXpxOJ+rq6hAZGWk6Xp2Xl4fk5GS4XC7s2rVrwI5rQJJzSUkJFi9ejNraWkyaNAlTpkxBY2MjlixZgu9973twu90DcRhERF727NkDABgzZgyio6NNl8nJyQEA1NfXD9hxWT6sEYpjOUREHkePHgUApKWlicukpqb2WXYgWJ6c+3ssp6OjA7W1tXA4HAN655RosHC73XC5XMjOzhavJP3R0tKC1tZWv5Y1DAOapnm9Hh8fj/j4+KCPwR/t7e0AgJiYGHGZuLg4AEBbW5ulx/JJliZnf8dympubsWvXLkyYMMHnNmtrazFr1iwrDpeIPqGsrAyTJk0Kat2WlhZMm3ozzpz1LzlHRUXh/PnzXq8vWLAACxcuDOoY/OXpmmz2y+FisjQ5+zuW09zcjPr6er+Ss8PhAAD8dtUTSLk8ySv+1y8/oVy/TZOvzrMSTokxV8tQMZYY26HcZ2eXfIUfG9Mpxna3DhdjXT6+SBn6OTEWHSmP8ds0+QbtkKhu5T6PnhkmxprskWIsykdL8Wu0djEWGSG/l51u+YrL199o1w05Lca63fLaEXb5/EUOUZ+/hhb5886MbxFjVe0jxNiUoSeV+zRzwujGg0Zz77+1YLS2tuLM2VY8t+oxJDu8/51+UrPrBIrufRBlZWVISUnpE7P6qhn411Wx5wrajOeK2bPsQLA0OVsxluMZyki5PAkjU5O94pfpUcr1hyiGTpI1OYEYxhAxNgLqG5rnDTk5D9XkxJRoyO+lC+rk7DDkY4ox5CShSs5RPvZ5TnHu22zyuY32kZwdkI93iCKWoDh/vpLz5ZCPt9tQJGcofrn5+IXqUhyv6rupep+q9ZQM9MuwYXLSCIxM8ZHkjZ5zlpKSglGjRl3wPgM1cuRIAEBTU5O4jNPp7LPsQLA0OYfqWA4RDRBd7/nxtcxFNH78eADA/v370dHRYfpXfk1NDQBg3LhxA3ZclpbShepYDhENDAMGDENX/+DizpSXmpqKrKwsdHV1obKy0iteXV0Np9MJh8OB3NzcATsuS5NzqI7lENEAcXf79zMAli9fjoKCAixfvtwrNm/ePABAaWkpjhw50vv6yZMnUVJSAgCYO3fugDY/snRYw8qxnL9++QnT8eWptT9Trvduzv1i7J0W+cbFZYox3GPticp9RirGVNvPyH9VZEL+pWb3cbXh1OUSqI5OOaYr/sjp6lT/BZSMLjE20XZGjKlusAHAh+5Y+ZgUfxHnR8r71H1crNUoPtMIxefZrThHXfI9WgBAOryrFTzePXOZGJscIb/PKsV3GoDpXYQW7Tww9JhyPb/pOqD7eMgsiGGNurq63qQJAAcOHAAArFixAmvXru19/YUXXuj9b5fLhcOHD8Plcnltr6CgADNnzkRFRQUKCwsxZcqU3sZHra2tmDp1KmbPnh3wcV4IS5NzqI7lENEAMfTeG37KZQLU2tqK999/3+v1xsbGgLflUVxcjIkTJ6KsrAzV1dXQdR0ZGRmXZstQz1hOXV0dKisrcdttt/WJX6yxHCIaIIYfNwSDSM6TJ09GQ0NDQOssW7YMy5YtUy5TWFiIwsLCgI/HCpb/KgjFsRwiGhg+bwb+84e8Wf74diiO5RDRANENP0rpLm61RqgakH7OoTaWQ0QDxN3V8+NrGfIyYM32Q2ksh4gGiEU3BAcDzoRCRNYx/BjW8PH4/mAVtsm5TbOZ9slQ1TEDwHU1cmOkVZMWi7FCQ661rfLRECjJkE/zPsjFr1OvkBsxRSaov/C7dl0pxj6MkNftVNRPH4fcpAkA7lIcUsb1Z+V9utTvZft+uenUcU2uoS2YJjek0jvUn9nGv8qf91BFb40ziuNxGXIdMwA8HCnH13XJn8u/3yjXwz/+Z3V/DLtJpfN5Tf05B4RXzkEL2+RMRGEgDHprhComZyKyjKF3w9DVN/wMfWAe3w43TM5EZB2LHkIZDJicicg6HHMOGpMzEVnHosZHgwGTMxFZh1fOQQvb5JyVcMp0Ch5V209AXS63fmepGPvgpv8WY18r/pJyn0ZtrRizf+3bYqyw8Gkx5uyS57gDgDc/94EYi70xQ4xpiQlyLPs65T5f/fprYuwblc1ibHSs93Rjn1R5mzwPXsRnrxVjCQu3i7G4IepZpZvW/JsYMz76SIxpaXLrW1uOvE0AeHT6s2Js1e9uE2MjvvSAGDv13nrlPnHeewaiY8c/xvTvqVvv+o3VGkEL2+RMRGFAd/tupu9r2GOQYnImIuvwyjloTM5EZBnDcMNQzCLkWYa8MTkTkXXYMjRoTM5EZB1WawSNyZmIrMMnBIMWtsnZ1TIUhjHE63XVLNmAurucqlzuyjd/LcYOXr9Auc/z5+XTHPfCU2KsSHeIsa4IOQYAB3bLncpi98qlaXa798zEHlHRe5X7HGYkirHioRPEWIyPP2v3v9IqH1PlfjG2IvlLYszu4y/pxkd2irHubrkrXUSEPGt1VMzbyn3mnZfLQI9+e60YK036ghg79s1fKPdplhebjX7sdeH2o1rDzTFnM2GbnIkoDHBYI2hMzkRkHZbSBY3JmYiswzHnoDE5E5F1DMOPYQ2W0plhciYi67i7/bghyGb7Zpicicg6HHMOWtgm58TYDoyAdwnOsfZE5XqqyVhV3eVU5XJX/+1Xyn2ee3i+GIuYKncqq7z/H2LspC5PXgoAz14hx+MnxYgxLV4uNbRlXq3cZ+1DH8rH45a75F0RmajcbmHeGTEWOfZyMfbU6noxFmuPUu6z6J4xYkw/Lk+8a0uWy+G0MZnKff7u+3Vi7CsPyaWIqxZvE2Pf/unXlfs0Ok2+J6dagSflksDA+DGsoZhUeDAL2+RMRGGAV85BY3ImIuswOQeNyZmIrGMYvqsxWK1hismZiKzjdgPdfHw7GEzORGQdPr4dNCZnIrIOx5yDxuRMRNbhmHPQwjY5d3bZcd6we70e6eODTjLkt6yaJVvV9lNVxwwAMT97Rt7u0z8SY2mQZ4gealN/dO1n5KuR6CNyC0573DkxFuFWX+HYDbl2OCNqhBhL1tQ1xx1Ouc5ZGyK3OB0TLc/qHa2pz1/3/iYx5j7VKcYizsr15fbz8noAEG/IrUj1+n1iLFPxPvWa3cp9Gl1d3uu0nVeuExDOhBK0sE3ORBQG2PgoaEzORGQdtxuGr2oMVmuYYnImIuvwhmDQmJyJyDpsGRo0Jmciso5h+L7hx+RsismZiKzDYY2ghW1yjo3pxFDN+zdu+xlNud4+yKVi9q99W4ypZslWtf0E1OVyUQt/LsZ2r7tXjDV3yeVlAJCQKX+00Z+/SoxpCfFyLOc65T47f/uGGPvrmQNiLCNWLgUDgLhsuaQw4jNya89tL/9B3makvE0AiLzxDnmfH30kxrSRaWLMljVFuc9jazeIMfv0W8XYtqeK5fVulcs4AQCd3v8e7K6PgT/K7VYD4nb7vuHHG4KmwjY5E1EYMPyoc76AYY3NmzejoqICDQ0N0HUdo0ePxu23346ZM2fCZpPrxj/toYcewqZNm8T46NGjUVlZGfRxBoPJmYiso/sx5hzkQyglJSUoLy9HVFQU8vPzERERgaqqKixZsgRVVVVYuXIl7HbvB9VUJkyYgKuu8v7L0uFwBHWMF4LJmYisY1Hjoy1btqC8vBwOhwMbN25Eeno6AODEiRMoKirC1q1bsXHjRtx5550BbfeOO+7AjBkzAj4eK/h/3U9EFCjPlbOvnwCtXr0aALB48eLexAwASUlJKC4uBgCsWbMGehjfbGRyJiLLGLru108gnE4n6urqEBkZiYKCAq94Xl4ekpOT4XK5sGvXrn56JwOPwxpEZB1d912NEWBy3rNnDwBgzJgxiI42r7rJyclBc3Mz6uvrMWGCPDnup+3YsQMNDQ1ob2/HZZddhokTJ+L6668P6OZifwnb5Ly7dTgSTbqgZaJdud7UK+SZkwsLnxZjRbp8Q0A1Szag7i6nKpf743urxFj7onnKfW74c6oYO7DfuxOZR4dxXIwd0/9Xuc8HFe+zvihdjHUfk7vkAcCK11PEWNPrTjF24oF8eaPn1J3X5v94rxgbpuhod9rYI8aa3NXKff4mQY59/VvPi7GP7pOTz/TpTyr3Gal5J50urQuQJ2gPjAU3BI8ePQoASEuTyxZTU1P7LOuv3//+916vXXPNNXjyySdx7bXXBrStCxW2yZmIwkAAD6E4nd6/aOPj4xEf37f2vr295wIsJkb+DRIXFwcAaGtr8+sLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 2 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAWMAAAENCAYAAADaPATLAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAAA0w0lEQVR4nO2de3BU95Xnv7el1hMk8Wj0AGwk80ZsVogIhMcZx0Ow4kRxDEPtMALZVBZSi814K5EzduzYEuPZUWwwJjtjh2WNiNGjyusMLjO2wYRlKDxRkFlHiV4WT8kGIaklECAJgdT37h/tlhHd59e3hS50o+/H1VW4z+93z69vd3/163PPPUczDMMAIYSQO4rtTi+AEEIIxZgQQoICijEhhAQBFGNCCAkCKMaEEBIEUIwJISQICL/TCyCEkEA5ffo0jhw5gpqaGtTW1qKpqQmGYWDbtm3IyckZ9nH37t2LiooKNDY2Qtd1pKamYsWKFVi1ahVsNnnvOtx5N0IxJoSEHBUVFXj77bdH9JhFRUUoLy9HZGQksrOzER4ejsrKSmzatAmVlZXYtm0bwsLCRmzezVCMCSEhx8yZM/GjH/0I6enpSE9Px/PPP4+qqqphH2///v0oLy+Hw+FAaWkppk2bBgDo6OhAfn4+Dhw4gNLSUjz++OMjMs8XjBkTQkKOlStX4mc/+xkeeeQR3HPPPbd8vO3btwMACgoKBgUVACZOnIjCwkIAwI4dO6Dr+ojM88VtE+O9e/fib//2b5GZmYmMjAwsX74cZWVlphZJCCFW0drairq6Otjtdp/x5qysLCQmJsLpdKK6uvqW50ncFjEuKipCQUEBamtrsXDhQixZsgRNTU3YtGkT/u7v/g4ul+t2LIMQQryor68HAMyYMQNRUVE+x8yfPx8A0NDQcMvzJCyPGY9kTIUQQkaas2fPAgBSUlLEMcnJyUPG3so8CcvF2F9MZc2aNdixYwfWrFljKgWkr68PtbW1cDgcpq5QEkICw+Vywel0Ij09XdzxmaGrqwvd3d2mxhqGAU3TvJ6Pi4tDXFzcsNdght7eXgBAdHS0OCY2NhYA0NPTc8vzJCwVY7Mxlba2NlRXV2PBggV+j1lbW4u8vDwrlksIuYGysjIsXLhwWHO7urrwnaV/hctXzIlxZGQkrl275vX8U089hY0bNw5rDWbxVBH29cfAinkSloqx2ZhKW1sbGhoaTImxw+EAAPzmX15F0qSJXva+Lb9Qzu8/3yfabFHySe07L9ti0yOVPl3tV0Wb3i/P6z4fIdrGfUP9Aag5OFa0zc+R/0qf/Fh+LanfvKT0aYuUf6lodtl27GC88rjffET+Qn+8T56bEXdBtI2d5P3Fv5HINHm3s+/AONH23IU/iLY/fFN91d8+Uf5laIuT3xfXRcVnOsau9AmX9wX0tmsD+HFDx+B3bTh0d3fj8pVuvP0vv0Siw/t7OsSfswP5T/49ysrKkJSUNMRm9a4Y+Hr36tnp+sKzs/WMvZV5EpaK8UjHVAAMhiaSJk3E5OREL/vVaPWH77pdVr+wCPnL0Bsmi99YPz4HImSfOuTa/lc0+e2ZoPjDAQAtuizkKdHXRdslQzEvQv1xUYpxhGwbr6v/mKVEy8KZYMhzEzX5fYkPV180jlK8pyqfAwPy+5kcrj5/EZEKMVasx9UzIM+L8vMVH5CzmUYiDJg4cTwmJ/kRdcO9hqSkJEyZMuWWfQbK5MmTAQAtLS3imNbW1iFjb2WehKXZFCMdUyGEhBi6bu5xB5k7dy4A4MSJE+jr8/0ro6amBgAwZ86cW54nYakYj3RMhRASWhgwYBi6+qH4dXg7SE5Oxrx589Df3499+/Z52auqqtDa2gqHw4GMjIxbnidhqRiPdEyFEBJiuAbMPW4DW7ZsQU5ODrZs2eJlW79+PQBg8+bNaG5uHny+s7MTRUVFAIB169Z5ZXwNd54vLI0Zj3RM5Ub6tvzCZ3w4+h/+RTnPXnNINsbKF4MiT9eItrBF31f7bDkuG6/Jf6hiTtbLPv/yMaXPzOubRFvkM78UbXPthaItfLX6qrYWJV801OxynHVx54vK40b81PvL4+HBoz8TbRPyZ8vrSZuh9Gmb+4Bo+3bl34u2qC455j72Jz9U+tSS02TbOPn7oXc0izYtdrzSJ1ze1zMi2juBDYXqeWbRdUD3c1PXMMIUdXV1g2IHACdPngQAbN26FTt37hx8/p133hn8t9PpxJkzZ+B0Or2Ol5OTg1WrVqGiogK5ublYsmTJYMGf7u5uLF26FKtXrx6xeb6wVIxvjqn4yqgIJKZCCAkxDH3wAp1yTIB0d3fjT3/6k9fzTU1NAR/LQ2FhITIzM1FWVoaqqirouo60tDS/pTCHO+9mLBVjT0ylrq4O+/btww9/+MMh9kBjKoSQEMMwcYFuGGK8aNEiNDY2BjSnuLgYxcXFyjG5ubnIzc0NeD3DnXcjltemGMmYCiEktPB78e6rB7kNt0OPZEyFEBJi6Ib/nbF+Z7MpgoXbUlx+pGIqhJAQw9Xv8yKh1xhy+zp9jERMhRASYlh0Ae9uhG2XCCHWYZgIUxgMUwAhLMb95/t81plQ5hEDCJ//bfmYh8pEm9H8hWjTk2uVPo3zTbLxqqKI0Jdyfrb2RZ3S57VzciJ9RPOfRVv/l1dEW9gXfq5eR8m3vRuKPOPec+oaCNFfyOf3Uqfsc9wX8vmz+am7oCtyzrsuxIg2l0J4jObTSp/GdbkGh3a5U57X2SofNFbO/Qbg84YL48Jl9ZxA4M7YNCErxoSQEMBM7Qm2XgNAMSaEWIihD8BQ1Yn9agyhGBNCrMSimz7uRijGhBDrYMzYNBRjQoh1WFQo6G6EYkwIsQ7ujE0TsmJsi9J8t0lSpCQB6vQ1+7flRqf97a/JB/XjU4seI9qMMPkt0GLltC3EqHuDhcUocjcV67XFKFonxfipOR2tSKMKl0tLhkerd06aYr12u2JurOIcjZHfE/dcRTnVSEWbI1UjBT8+NdV6FevR+hRdcvx8Nn3WEpazLQOH2RSmCVkxJoSEALrLf/F4f2GMUQLFmBBiHdwZm4ZiTAixDMNwwTDUO19/9tECxZgQYh0soWkaijEhxDqYTWEaijEhxDp4B55pQlaM+85r6A3zTiNSdXEG1NXXVOlr9v/yE3newd1qn+fl6mG4dl00uc7I1bi0yeqqbd1nvTtne4g+I8+92ix/MSLOnFT6hI+Gs4PY5dS2y22KFD4AcYr39MJluYJaUpN3F2Cz2CLkKnMdCp+qqm366Sa1T0UFP1xoV9jkim5+U/gGfGQ6XJI7lgeMy0Q2hYsxYyCExZgQEgIwTGEaijEhxDqY2mYaijEhxDoYMzYNxZgQYh2GYSJMwdQ2gGJMCLES14CJC3gsLg9QjAkhVsKYsWlCVoxj0yMxNto7fSts0feV85TNQxUVrlTpa/a/WqP0OXDiqGzsk9OItAn1oi1s0SNKnwnf+FQxVz5HYyqPiTZb1kNKn6qqbVq4nCo2ce4nysOGL/6BaJsy5d9Fm33hHHk9981U+rTNfUC0TU0+ItrCr8hV78IWLVb61JLuk23jU0Sb4WyW540Zr/RpuLxbImntnQB+p5xnHhNhCjBMAYSwGBNCQgDujE1DMSaEWAfF2DQUY0KIdRiG/2wJZlMAoBgTQqzE5fJ9y/XNYwjFmBBiIbwd2jQUY0KIdTBmbBqKMSHEOhgzNk3IirGr/SoGIrxzJO0tx5XzjPNNok3ZxVlRBlOZRwwgfMYi0db/cYnss1Uum6ifP6H02d8ml+aMOPe5aHOdvyLaws+fVvpElNw92lCU0LzWruioDCC6pVG09VyS85fHtcolNG3R6rKderx8jrovy6VCVSU0jZYvlT5VO0Tj6mV5Xud5eV6sorwm4PPuN6PzknpOILDTh2lCVowJISEACwWZhmJMCLEOlwuGv2wJZlMAoBgTQqyEF/BMQzEmhFgHS2iahmJMCLEOw/B/gY5iDIBiTAixEovDFHv37kVFRQUaGxuh6zpSU1OxYsUKrFq1CjabzdQxjh49ivz8fFNjDx06hJSUryvoPfvss9izZ484PjU1Ffv27TN17JAVY70f0H2V3rvmp7OtogOvEaY4HYouzqoymIA6fc2+bK3scttzCp89Sp+uPkW62PU+0aRfV6RX9Sm6FwNQJqj5KNU4aLru50ujeK0DA4q5qvfs2jU/PuX3VOlThR+fxnXZrinWo3pftHA/X3FftypfU7/PAeFy+b9AN8wLeEVFRSgvL0dkZCSys7MRHh6OyspKbNq0CZWVldi2bRvCwuSSph4mTpyIxx57TLT/+c9/xqlTp3DPPfcgOTnZ55gFCxbg3nvv9Xre4XCYfj0hK8aEkBDAMJFnPIwwxf79+1FeXg6Hw4HS0lJMmzYNANDR0YH8/HwcOHAApaWlePzxx/0e67777kNxcbFo/973vgcAWLFiBTTN95Zj5cqVWL58ecCv40aG+SeeEEJMoBvmHgGyfft2AEBBQcGgEAPuXW5hYSEAYMeOHdBvMVPjj3/8I06ePImwsDDl7nkkoBgTQqzDUyjI3yMAWltbUVdXB7vdjpycHC97VlYWEhMT4XQ6UV1dfUvL/+1vfwsAeOCBB5CYmHhLx/IHwxSEEOsws/MNcGdcX+9uRzZjxgxERfm+NX3+/Ploa2tDQ0MDFixYENDxPVy9ehUffvghAOCv//qvlWOPHj2KxsZG9Pb2YsKECcjMzMT9999v+iIiQDEmhFiIoesw/IQK/Nlv5uzZswAwJKvhZjwX2jxjh8O+ffvQ09ODCRMm4MEHH1SOfe+997yemz59Ol577TXMmjXLlD+KMSHEOnTdf7bEV2Lc2trqZYqLi0NcXNyQ53p73Zkl0YpiT7Gx7qJVPT3qrCMVnhDFo48+Crvdu/kxAMyePRsvvPACsrOzkZKSgu7ubtTX12Pr1q34/PPPsXbtWuzZs8dUiCNkxbj7fASuaN7Ljzkpd1QLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 2 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAWcAAAENCAYAAADT16SxAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAAAyGklEQVR4nO3de3xU5bU38N9M7gkEUMYkgBpyIAgJVi6GBj/at4I2VaMUih4MoFShFkHbGs/xUi2heAqSqNFiy5sjIOSiFMW3qQpaVGprJFqLJgHCNVFIEycpwVwYSGbv9484kTiznj0zZIcZ8vv2k8/Hztp7PzuTyWLn2Wuvx6Lrug4iIgoo1nN9AkRE5I7JmYgoADE5ExEFICZnIqIAxORMRBSAmJyJiAJQ6Lk+ASIiXx0+fBjvv/8+KioqUFlZiZqaGui6jvz8fGRkZPh93NLSUpSUlKC6uhqapmHkyJGYNWsW5syZA6u1b69lmZyJKOiUlJRg48aNvXrMnJwcFBcXIyIiAunp6QgNDUVZWRmWL1+OsrIy5OfnIyQkpFfHVGFyJqKgk5ycjLvuugupqalITU3Fo48+ivLycr+Pt337dhQXF8Nms6GwsBCJiYkAgMbGRsyfPx9vv/02CgsLcccdd/TSd2CMyZmIgs7s2bN79Xhr164FAGRnZ3cnZgAYOnQoli1bhnnz5qGgoADz5s3rs+mNPptEKS0txe23345JkyZhwoQJmDlzJoqKiqBpWl+dAhGRm/r6elRVVSEsLMzjfHVaWhri4uJgt9uxe/fuPjuvPknOOTk5yM7ORmVlJSZPnoypU6eipqYGy5cvx3333Qen09kXp0FE5GbPnj0AgNGjRyMyMtLjNuPHjwcA7N27t8/Oy/RpjUCcyyEicjl69CgAYNiwYeI2CQkJPbbtC6Yn596ey3E4HKisrITNZuvTO6dE/YXT6YTdbkdqaqp4JemN5uZmtLa2erWtruuwWCxur8fGxiI2Ntbvc/BGe3s7ACAqKkrcJiYmBgDQ1tZm6rmcydTk7O1cTkNDA3bv3o2JEycaHrOyshJZWVlmnC4RnaGoqAiTJ0/2a9/m5mZcN30avmrxLjlHRETg1KlTbq8vWbIES5cu9escvOXqmuzpH4dzydTk7O1cTkNDA/bu3etVcrbZbACAF9esRvxFQ93if/vBauX+Fw9sEWMOR5gYi4zsEGO1LQOVY6qo/laIhTxmM+RzBYBLFd/nF36erxXq1t+dMOfDfUmM/At+tG2AGIuBfC/D6FxPKn4ywyLaxVhYuDym0fuu+mw2tkSLsZOQ/4K0hTmUY4aHuZ9vIzrxq9Bj3b9r/mhtbcVXLa3YuGYV4mzuv6dnarA3Yv69/42ioiLEx8f3iJl91Qx8c1XsuoL2xHXF7Nq2L5ianM2Yy3FNZcRfNBTDE+Lc4hdqEcr94yB/WNt1OeHJvxpAi8GYKqrkPFjxL7lVD1ceV/V9tinOV5WyjJJzh0nJ+SK4X1G5nFR8LwMsnWKsQ1efa7si4V2k+EczXPETVb3vgPpnBk3+ebcpfo1tuvpme4TifHtj2jBu6AUYHm+Q5PWuiq34+HiMGDHirMf01fDhwwEAdXV14jb19fU9tu0LpibnQJ3LIaI+omldX0bbnEPjxo0DABw4cAAOh8PjX/kVFRUAgLFjx/bZeZlaSheoczlE1Dd06NB1Tf1l8BeZ2RISEpCSkoKOjg5s27bNLV5eXo76+nrYbDZMmDChz87L1OQcqHM5RNRHnJ3effWBvLw8ZGRkIC8vzy22aNEiAEBubi5qa2u7X29qakJOTg4AYOHChX3a/MjUaQ0z53L+9oPVHueXp1c+odzv8pT/FGN3R44RYxtOHBRj5ZtvU46JztNyrF2+CXRg6btibML/na4c8oqsdWLsn2vS5R1VvyjHjyvHtIwcJQfD5akto1/Oa+/cLMZ2PCXfRD7xh7+JsejR6huq4bfeLMYW/VT+uew6XivGPn1hmnLM9EWvirEdafK0X1TWtWJs9X8fUI75jqPB7bVOS2fvZQZNAzSDh8z8mNaoqqrqTpoAcPBg1+/n008/jXXrvvnsb978zWfHbrfjyJEjsNvtbsfLyMjAnDlzUFJSgszMTEydOrW78VFrayumT5+OuXPn+nyeZ8PU5ByoczlE1Ed0rfuGn3IbH7W2tuLTTz91e72mpsbnY7ksW7YMkyZNQlFREcrLy6FpGpKSks7PlqGuuZyqqips27YNM2bM6BE/V3M5RNRHdC9uCPqRnKdMmYLq6mqf9lm5ciVWrlyp3CYzMxOZmZk+n48ZTP+nIBDncoiobxjeDPz6i9yZ/vh2IM7lEFEf0XQvSunObbVGoOqTfs6BNpdDRH3E2dH1ZbQNuemzZvuBNJdDRH3EpBuC/QFXQiEi8+heTGvonNbwJGiT88UDWzz2IlDVMQPAZ1UvibG2XywUY/c+8lsxNvCKecoxQ61yj4LoMLnfQnG0XMEydYa6wVPTZ8ViLDntp2LMapGnmC4MVzfuOdjyphg71Sn/6RpiMK1lf0e+wz7+plVi7L6ocWJs3xeK2nMAm15+Wow1PDNDjIVMe0iMXTxhvnLMz3c+JcZ+e0uhGHvxg61irGKO+vmB//rx7W6vHWs8gRsf/YNyP6/xytlvQZuciSgIBEFvjUDF5ExEptG1Tuia+oafrvXN49vBhsmZiMxj0kMo/QGTMxGZh3POfmNyJiLzmNT4qD9gciYi8/DK2W9Bm5wdjjCPy0qp2n4C6nK5mKcL5P3uv1uMPZxwjXJM1VIDoYpotIcFL10esCnafgJwLP+VGPvZgPFiTFXUFm6wtNOJCy8RY9pZNFR35P5ejC2KlrsZJp2SbzTZrOqWoXFD5ff3q8J/iLGo9907pbncN3iSckzHSrk88opTCfKYMalirPGvXyrHHHTsRbfXTp3uxRt0rNbwW9AmZyIKAprTuJm+0bRHP8XkTETm4ZWz35icicg0uu6EbrACuFG8v2JyJiLzsGWo35icicg8rNbwG5MzEZmHTwj6LWiTc2RkB6I9vK5aJRtQd5dTlcvF5P+vGFt9sbz6MQCEhii60oXKXenSIi8XY/lN5coxH3nkBTFW8MbPxViIoivdkLAByjEPn6gXYw5FQ3XVmADw0P2Pi7H1s54TY/dGyGWV1SHqfg/FTf8UYw8+/EMxFvL9H4ux50vvUY6Znb1CjFWWbRFjL7btFWOLb7hIOWbY7Cy31yKaTgCPyZ93nzi9qNZwcs7Zk6BNzkQUBDit4TcmZyIyD0vp/MbkTETm4Zyz35icicg8uu7FtAZL6TxhciYi8zg7vbghyGb7njA5E5F5OOfst6BNzrUtA9GiuZehlW++TbmfajFWVXc5VbncV1+8oxxTP9kix9pPiLGDNz4pxho+lBeqBYCYMbeIsdYPfiefj+oqpvGockxrcpoYs0TE+DcmgEsUP7OaNx8TY21PyIuURl5uU46Zu1BeUPWH18gd/3YulT9/bbs3KseMUyy8+/msRDGWvfRRMfbw7D8qx9y62f2zYLHqiLhQuZsPvJjWOIuOheezoE3ORBQEeOXsNyZnIjIPk7PfmJyJyDy6blyNwWoNj5icicg8TifQyce3/cHkTETm4ePbfmNyJiLzcM7Zb0zORGQezjn77fxLzp2nleFQq9y+U7lKtqLtp6qOGQAsUQPFmHZCvTqyqFNemRtQf5/KuuIOxXEN5gb1DsV7H6JY7VrRThQAwhTfC1Rjqi7IjFbfULwPIYpPilXV/tSgnjvEKu+rq85XcVx1M1bhc2LtxStZroTit/MvORNR4GDjI78xOROReZxO6EbVGKzW8IjJmYjMwxuCfmNyJiLzsGWo35icicg8um58w+8sknNpaSlKSkpQXV0NTdMwcuRIzJo1C3PmzIFVcYP12x566CFs3bpVjI8cORLbtm3z+zz9weRMROYxcVojJycHxcXFiIiIQHp6OkJDQ1FWVobly5ejrKwM+fn5CFFUWXkyceJEXHrppW6v22zqLoZmCNrkbIVQJtSuLmuLDpNXuw5VlEipVslWtf0E1OVyIfH/IR9XcUGht6nHVH2fUO2rKMvS21uVY1ram+V9VeVyBiVmqvceijJGTXFY3aEuuVT9TCMt8q9NTHikX8cEDL7PTsWHQfHzDNNVBaJAVEi422u61YlOOJT7ec3pNL7h58cNwe3bt6O4uBg2mw2FhYVITEwEADQ2NmL+/Pl4++23UVhYiDvuuMOn486ePRszZ870+XzM4P11PxGRr3T9m6tn6cuPaY21a9cCALKzs7sTMwAMHToUy5YtAwAUFBRAC+KbjUzORGQeTffuywf19fWoqqpCWFgYMjIy3OJpaWmIi4uD3W7H7t27e+kb6XtBO61BREHAhMZHe/bsAQCMHj0akZGep5HGjx+PhoYG7N27FxMnTvT62Lt27UJ1dTXa29tx4YUXYtKkSbjqqqt8urnYW5icicg83lwZfx2vr693C8XGxiI2NrbHa0ePdi2XNmzYMPGQCQkJPbb11muvveb22qhRo/DUU09hzJgxPh3rbDE5E5FpdE2DbjDv64pnZWW5xZYsWYKlS5f2eK29vR0AEBUVJR4zJqZrzcq2tjavzvOyyy7Dr371K6Snp2PYsGFobW3Fnj178PTTT2Pfvn1YsGABtm7diri4OK+O1xuYnInIPJpmXI3xdXIuKipCfHx8j9C3r5oBQP/6BqLFoq5E8cWdd97Z4/9HR0fjoosuwtSpUzFv3jzs3r0ba9euxeOPP95rYxoJ2uQciw4M9vDDObD0XeV+xdETxFj0KbkTWVrk5WJMtUq2EdWN6jG7nhVj1VPuUx63JFqeZzvwk9eMTssj3aAsKyTkIzGm+j0yulm/xjpKjB24Z7sYO91xgXw+e9WDRpT+Xoz94nS0GFsSM0mMHbpdvfr2C+EpYuzg+/JnM/zDP4mxGxxyaR8ATEei22tNOIVfY69yP6/5MK0RHx+PESNLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 2 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAWcAAAENCAYAAADT16SxAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAAAqy0lEQVR4nO3df3AU530/8PfqB5IAS0A4SwJiC40RPyRmInBEhSee+caQqu6odfG4HQqS7SYwHRviTo3n29RNIlHakW0RB6dOh2ECBPSj8dihU5oEStNvHU8io7iJHEmAjA2iBXryQZCDJISl2+f7h3Jny7f7efb2bk970vs1oxl7n93nnts7fVg9+9nPYyilFIiIyFcypnoAREQUi8GZiMiHGJyJiHyIwZmIyIcYnImIfIjBmYjIh7KmegBERPG6cOEC3njjDXR3d6Onpwf9/f1QSmHfvn2oqalx3e/x48fR3t6Ovr4+mKaJpUuX4uGHH8bmzZuRkZHaa1kGZyJKO+3t7Thy5EhS+2xsbERbWxtycnJQXV2NrKwsdHR0YPfu3ejo6MC+ffuQmZmZ1NeUMDgTUdopKyvDF7/4RVRUVKCiogLPPvssOjs7Xfd38uRJtLW1IRAIoKWlBSUlJQCAa9euob6+HqdOnUJLSwseffTRJL0DPQZnIko7jzzySFL7279/PwBg165d0cAMAAsXLkRDQwPq6upw4MAB1NXVpWx6I2WTKMePH8ef/umfYu3ataisrMSmTZvQ2toK0zRTNQQiohjBYBC9vb3Izs62nK+uqqpCYWEhQqEQurq6UjaulATnxsZG7Nq1Cz09Pbj33nuxfv169Pf3Y/fu3fjyl7+McDicimEQEcU4c+YMAGDZsmXIzc213Gf16tUAgLNnz6ZsXJ5Pa/hxLoeIKOLy5csAgEWLFtnuU1xcPGnfVPA8OCd7Lmd0dBQ9PT0IBAIpvXNKNFOEw2GEQiFUVFTYXkk6MTg4iKGhIUf7KqVgGEbM9vz8fOTn57segxMjIyMAgLy8PNt95syZAwAYHh72dCwf52lwdjqXMzAwgK6uLqxZs0bbZ09PD7Zs2eLFcInoY1pbW3Hvvfe6OnZwcBAbNzyA39x0FpxzcnJw+/btmO07duzAzp07XY3BqUjVZKt/HKaSp8HZ6VzOwMAAzp496yg4BwIBAMB3X34BRXcujGlfvk6+i3vup622beELv7Rty7yr3L5T3RV/5iz7NmV/Q/SlTd+ybdv5ve3yawrEL2G2/ZWSGr0p95t7h32j8D5haM6fcKyRaf8VNkftr3KM7Bz5NXVjcsHQfE/Cv75if+ws+6s6Y/Y82zYlnAMAMLKzY7YFQ9fx2F98Lfq75sbQ0BB+c3MIR15+DoWB2N/TjxsIXUP9k/8Xra2tKCoqmtTm9VUz8NFVceQK2krkijmybyp4Gpy9mMuJTGUU3bkQi4sLY9rHx+W1A6yOiQj/psD+dYvutO9UF5yzhEAgBJ45yv64xUXuf3EMKfAIQUCNyH/iGrPtz9+UBOdb9v+YSMHO0Zhc0AbnrFH7Y3Psg4Ixd4FtmxLOAQAY2fYXDsmYNixcuED/Xf3t51tUVIQlS5Yk/JrxWrx4MQDg6tWrtvsEg8FJ+6aCp8HZr3M5RJQipjnxo9tnCq1atQoAcP78eYyOjlr+ld/d3Q0AWLlyZcrG5WkqnV/ncogoNRQUlDLlH0ztSnnFxcUoLy/H2NgYTpw4EdPe2dmJYDCIQCCAysrKlI3L0+Ds17kcIkqR8LiznxTYu3cvampqsHfv3pi27dsn7uE0Nzfj0qVL0e3Xr19HY2MjAGDbtm0pLX7k6bSGl3M5y9c9Yjm/PHr1DfG4vMX327bdnW8/r/w/N685H9wnhE37h2ykvyourF5h21aw/CHX4zGFNX2lseZkCTc2Adwe/9C2TXqfujWG3R6bLcxHjwvvU8erNZELcu0vUD4UAtitsdgshwjdZ2Z1HrKyDHx6iWZO3inTBHTn2sW0Rm9vbzRoAsC7774LAHjxxRdx8ODB6PZXXnkl+t+hUAgXL15EKBSK6a+mpgabN29Ge3s7amtrsX79+mjho6GhIWzYsAFbt26Ne5yJ8DQ4+3Uuh4hSRJnyzeDIPnEaGhrC22+/HbO9v78/7r4iGhoasHbtWrS2tqKzsxOmaaK0tHR6lgyNzOX09vbixIkTeOihhya1T9VcDhGliHJwQ9BFcF63bh36+vriOqapqQlNTU3iPrW1taitrY17PF7w/J8CP87lEFFqaG8G/vaHYnn++LYf53KIKEVM5SCVbmqzNfwqJfWc/TaXQ0QpEh6b+NHtQzFSVmzfT3M5RJQiHt0QnAm4EgoReUc5mNbwKDUx3aVtcD7301bLOhlSHjMA3LryE9u28e7/Z9uWeY+76lwAxHoV0p90u9f/vW3bB30vyK8p/EJI9SjE2hpDvxZfUqrxIP7pmqGp4SDlyWbGFu6JULeFwke62hpuJVCTw7z+P/bduq2tMfKB+JqGRR70leD7qKlLUiU4Xjm7lrbBmYjSQBrU1vArBmci8owyx6FM+YafMlPz+Ha6YXAmIu949BDKTMDgTETe4ZyzawzOROQdjwofzQQMzkTkHV45u5a2wTl84ZeWy0pJZT8BOV0ua/X/sT+u65TzwX2SsBSQVMt2rrIvlWn2vSm/psv0M3GsmiWPkCesIShdHemeEHV77If2yz5BU0rTtQSedlU33rdtM2YJS53NEZYH031mFkuoha/L6XdxYbaGa2kbnIkoDZhhfTH9BGprT2cMzkTkHV45u8bgTESeUSoMpeQrY137TMXgTETeYclQ1xicicg7zNZwjcGZiLzDJwRdS9vgnHlXOTKLYtPmdKtkS9XlpHS5rM9stG0zh26Ir2kI6WlKuJM9YvynbVtGqWbNRekLL6TSGbNiF+GNdjmsqXAmpHRJ71Oskge4TqVTo1NQlS6BVDrz2mXbNkNYmduYO9+2TfeZITs2lS5jbuzq1K6FHWRrhDnnbCVtgzMRpQFOa7jG4ExE3mEqnWsMzkTkHc45u8bgTETeUcrBtAZT6awwOBORd8LjDm4Isti+FQZnIvIO55xdS9/gnJGRUNpSvKR0uQwhlQnQLDSawIKgIq/6dWlK3qdu4ViJF98t3TmQXtPt+dO9D8t+k/neHUxrgNMaVtI3OBOR//HK2TUGZyLyDoOzawzOROQdpfTZGMzWsMTgTETeCYeBcT6+7QaDMxF5h49vu8bgTETe4ZyzawzOROQdzjm7lr7BOXOW9crBusUipVKRQmlPseynkMcMAEaOfblHNf6hbZt0PaEreamEPxWlEp1Sv2rMfqzaYxMoGer6WOnJM49KhhqJ5EdbfJ+jLEp7Rl8zkc/M4nttZAurs8eLK6G4lr7BmYj8j4WPXGNwJiLvhMNQumwMZmtYYnAmIu/whqBrDM5E5B2WDHWNwZmIvKOU/oYfg7MlBmci8g6nNVxL3+Bs8+SRYRjyceExoc0+9UpM59KUc5TS5Yws+xQ9qVdpPDoKwrFSv9K5g2ZMQoqj9rrJ5bHiZ6Z5L24plUCZUvG7ad8mnnfd+7RKRUxmsAyH9Tf8eEPQUvoGZyLyP+UgzzmBaY3jx4+jvb0dfX19ME0TS5cuxcMPP4zNmzcjI46c87/6q7/CsWPHbNuXLl2KEydOuB6nGwzOROQd08Gcs8uHUBobG9HW1oacnBxUV1cjKysLHR0d2L17Nzo6OrBv3z5kZsb3l8yaNWtw9913x2wPBAKuxpgIBmci8o5HhY9OnjyJtrY2BAIBtLS0oKSkBABw7do11NfX49SpU2hpacGjjz4aV7+PPPIINm3aFPd4vOCvtYyIaHqJXDnrfuK0f/9+AMCuXbuigRkAFi5ciIaGBgDAgQMHYKbxzUYGZyLyjDJNRz/xCAaD6O3tRXZ2NmpqamLaq6qqUFhYiFAohK6uriS9k9TjtAYRecc09dkYcQbnM2fOAACWLVuG3Nxcy31Wr16NgYEBnD17FmvWrHHc9+nTp9HX14eRkRF86lOfwtq1a3HffffFdXMxWdI2OL+06VuYo2IrdV1YvUI8bvf6v7dtm6vs0/BGjP90PLZPkr560kf+1bf+1rbtb+/9qviaUkKh9EekNFbdrRXpV1B6n1794TkVr5mIHOFTGxc+Nem86z4zq16HjduAUCAvLh7cELx8+TIAYNGiRbb7FBcXT9rXqX/+53+O2XbPPffgG9/4BpYvXx5XX4lK2+BMRGkgjodQgsFgTFN+fj7y8/MnbRsZGQEA5OXZl0qdM2eiTO/wsFzON2LFihX4m7/5G1RXV2PRokUYGhrCmTNn8OKLL+LcuXN4/PHHcezYMRQWFjrqLxkYnInIO3E8vr1ly5aYph07dmDnzp2f2H1if+0DZ3F47LHHJv3/7Nmzceedd2L9+vWoq6tDV1cX9u/fj6997WtJe00dBmci8k4chY9aW1tRVFQ0qemTV83AR1fFkStoK5Er5si+bs2aNQvbt2/HE088gddffz2hvuLF4ExE3oljzrmoqAhLlizRdrl48WIAwNWrV233iUyRRPZNRGlpKQBgYGAg4b7iweBMRJ5R42GocTlbQ9f+SatWrQIAnD9/HqOjo5YZG93d3QCAlStXxtW3lcHBQQCJX4XHi3nOROQd5eABlDhraxQXF6O8vBxjY2OW9S46OzsRDAYRCARQWVmZ8Fv40Y9+BACoqKhIuK94pO2V887vbcfiotjn3QuWPyQe90HfC7ZtZt+btm0Zpe4/ZLcLn0rpclKaHQCYv7b/k8+yEtlviWMd+UB8TWN2gf2xUuU0XQ6pcLdfWuBVXHjXqwVeNRUKJeZ1Ie1LWCQ4Y+582zbdZwaLqohXBkL41z97Rj7OKY8e396+fTueeuopNDc3o7KyMloP4/r162hsbAQAbNu2bVJ+8t69e3Hq1Cls3LgRTz/9dHT72bNnEQwGcf/990+qxTE+Po6jR4/i6NGjAGJvGnoLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 2 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAWcAAAENCAYAAADT16SxAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAAAqpklEQVR4nO3dfXBV1bk/8O/OCyGgAZRjEqAaciUIhHuLaGhw9N5RtKmdtBbH9kcDwTeYjoL+puK01rYmDHduUKLGVjtcRqCYl5ar0ilXBa3tWKdGIrXRvEBEIfEHeMKBEjWJ0eTs9fsjnkg4ez97n52zTs4h30/nzNiz9ss6KycPO2s/+1mGUkqBiIjiStJod4CIiMIxOBMRxSEGZyKiOMTgTEQUhxiciYjiEIMzEVEcShntDhARRerw4cN4/fXX0dTUhObmZrS3t0MphaqqKhQVFXk+7u7du1FXV4e2tjaYpomZM2fi5ptvxrJly5CUFNtrWQZnIko4dXV12LFjR1SPWV5ejtraWqSlpaGwsBApKSmor6/H+vXrUV9fj6qqKiQnJ0f1nBIGZyJKOHl5ebjjjjuQn5+P/Px8PPjgg2hoaPB8vL1796K2thY+nw/V1dXIyckBAJw8eRKlpaV45ZVXUF1djZUrV0bpEzhjcCaihHPLLbdE9XibN28GAKxbt24oMAPA1KlTUVZWhhUrVmDLli1YsWJFzKY3YjaJsnv3bvzwhz/EwoULsWDBAixduhQ1NTUwTTNWXSAiCuP3+9HS0oLU1FTL+eqCggJkZmYiEAigsbExZv2KSXAuLy/HunXr0NzcjCuuuAKLFy9Ge3s71q9fj3vuuQfBYDAW3SAiCtPa2goAmDVrFsaPH2+5zfz58wEABw4ciFm/tE9rxONcDhFRyNGjRwEA06ZNs90mOzt72LaxoD04R3sup6+vD83NzfD5fDG9c0o0VgSDQQQCAeTn59teSbrR1dWF7u5uV9sqpWAYRtj7GRkZyMjI8NwHN3p7ewEA6enptttMnDgRANDT06O1L2fSGpzdzuV0dnaisbERl19+ueMxm5ubUVJSoqO7RHSGmpoaXHHFFZ727erqwvVLrsMnn7oLzmlpafj888/D3l+zZg3Wrl3rqQ9uhaomW/3jMJq0Bme3czmdnZ04cOCAq+Ds8/kAAL998hFkXTQ1rH3gtZ3yAT79xL5t8hTbptv/6x+2bVt/PEc+Z0qqfVuS8IUwvN8SWF3RaNv23z/5V2/n7HW4akgTrrKUcOPX6aZwqjB+Qn/XVDTZtv36p/Plc0r3QXT9xSaNkcQUSrJ76GvnJ724fdufh37XvOju7sYnn3Zjx5MbkekL/z0ddr7ASZTe/RPU1NQgKytrWJvuq2bgq6vi0BW0ldAVc2jbWNAanHXM5YSmMrIumorp2Zlh7f2THQbP6Ldvm2K/b4qyH6ppkybI50wdZ98mBecRpOykKvuANk0aI+mcqQ7rMoy3/7NQW3AW+iuOgfCzBgAMCME5RVNw9pq5JAXnEfQ1GtOGmVMvwPQshyD/5XcjKysLM2bMGPE5IzV9+nQAwPHjx2238fv9w7aNBa3BOV7ncogoRkzT+R+dUU6nnTt3LgDg0KFD6Ovrs/wrv6lp8K+wOXMc/kqOIq2pdPE6l0NEsaGgoJQpvzC6K+VlZ2dj3rx56O/vx549e8LaGxoa4Pf74fP5sGDBgpj1S2twjte5HCKKkeCAu1cMVFZWoqioCJWVlWFtq1evBgBs2rQJHR0dQ++fOnUK5eXlAIBVq1bFtPiR1mkNnXM5A6/ttJxfTi26Xdwv+P+abduSsvNs214/sdW2Lfnqh8RzGmnCPzzJwo8gyfuc3xv/rLZtS/n3h4X+2M/Tmqf94jmTzrO/oaqkX0CHG2HGOOFGo9DffZ/Y3xxOuWajeE4xYKSmyft6ZXp8GEsaP2F8bHf5qBOo+l9vfTmbaTp/Lg/TGi0tLUNBEwDef/99AMBjjz2GrVu/+l3dufOr70AgEMCRI0cQCATCjldUVIRly5ahrq4OxcXFWLx48VDho+7ubixZsgTLly+PuJ8joTU4x+tcDhHFiDKds1A8ZKl0d3fjnXfeCXu/vb094mOFlJWVYeHChaipqUFDQwNM00Rubu65WTI0NJfT0tKCPXv24KabbhrWPlpzOUQUI8rFDUEPwXnRokVoa2uLaJ+KigpUVFSI2xQXF6O4uDji/uig/Z+CeJzLIaLYcLwZ+OWLwml/fDse53KIKEZM5SKVbnSzNeJVTOo5x9tcDhHFSLB/8OW0DYWJWbH9eJrLIaIY0XRDcCzgSihEpI9yMa2hOK1hJXGD86efWNbJkPKYASD5a/m2bcFj9nd/fRMn2baZgQ7bNgAwxgt5zoaQyzyCwkdTx9sXjDEDH3o6p/r4hHhOs7fLft8R1NYwxNok9nm8F6adb3/KE+3iOcU/tVM05Tl7vYJUQh6xMD52zJP/9NYPK7xy9ixxgzMRxb8EqK0RrxiciUgbZQ5AmfINP2XG5vHtRMPgTET6aHoIZSxgcCYifTjn7BmDMxHpo6nw0VjA4ExE+vDK2bPEDc6Tp1guKyWV/QTkdLnk6bNt2zq7T9u2JWVdKp7TSBOWsRJKOhojeHKyU0hrS8qeZX9OoYSpmW6fmgYAhlAyVCrB6VRbwRhnv5KO1N8Tn31s2yb9rAFA9YcvNuqmPyMhllUVd7QfP2l87CQbnd76YYXZGp4lbnAmovhnBp2L6XutY32OY3AmIn145ewZgzMRaaNUEEp6gvHLbSgcgzMR6cOSoZ4xOBORPszW8IzBmYj04ROCniVscL79v/6BFBXefWmVbECuLiely/Udf922bcolS8RzfiHcrTaFL6YaQSnF3tb/sW2bcPG1ns45SaquB6D7iz7btmSh2p1hGOJx+z2OX8/f7b8L6TP+QzxnktDfoKbsAqdxsN0P9vtJ42MnJcVAzsXneepLmKCLbI0g55ytJGxwJqIEwGkNzxiciUgfptJ5xuBMRPpwztkzBmci0kcpF9MaTKWzwuBMRPoEB1zcEGSxfSsMzkSkD+ecPUvY4Lz1x3MwbVJ4tbfkqx8S95MWY5Wqy0npcqc7/iSfs8tv3ygt4prkffHXGf9WYtvW0/qs/Y5ClTwlfQ4AxgT7NEXpT1vltMDruPH2jcIY/cvC223beg7uEs8JoSodNFWlExeVlb4LUV6M9pg/gG/d+n8j3s+ai2kNcFrDSsIGZyJKALxy9ozBmYj0YXD2jMGZiPRRyjkbg9kalhiciUifYBAY4OPbXjA4E5E+fHzbMwZnItKHc86eMTgTkT6cc/YscYNzSiqQOi7sbSNNLmtpCGUvpVWyxbKfDvm/SZOzbNvU5z32OzrkMkv6hNxXaQzEPGeHsYV0XOHqyHBcfVvIcxbG6AvT/mcmrogOQAmrVutafVvMV5a+C9ITdqmR5zkbad0R72OLK6F4lrjBmYjiHwsfecbgTET6BINQTtkYzNawxOBMRPrwhqBnDM5EpI/mkqG7d+9GXV0d2traYJomZs6ciZtvvhnLli1DUpL7ezY//elPsWuXfb2VmTNnYs+ePZ776QWDMxHpo5TzDT+Pwbm8vBy1tbVIS0tDYWEhUlJSUF9fj/Xr16O+vh5VVVVIThYKRlm4/PLLcckll4S97/P5PPVxJBiciUgfTdMae/fuRW1tLXw+H6qrq5GTkwMAOHnyJEpLS/HKK6+guroaK1eujOi4t9xyC5YuXRpxf3RI3OCcZAy+ziakQAEADOFfUiGNTFzF2CHlTUqXk1L/1AiKkIsrdwufU/oshsOfiYYw9grCZ3H63ZTGV0r98zoGcEjv85ji6DR+WhLKvPR1BCmcYYJB5xt+Hm4Ibt68GQCwbt26ocAMAFOnTkVZWRlWrFiBLVu2YMWKFRFNb8STxOw1ESUGpb66erZ7RTit4ff70dLSgtTUVBQVFYW1FxQUIDMzE4FAAI2NjVH6ILGXuFfORBT/TBdzzhE+hNLa2goAmDVrFsaPt35Aaf78+ejs7MSBAwdw+eWXuz72vn370NbWht7eXlx44YVYuHAhrrrqqlG5+mZwJiJ9NBQ+Onr0KABg2rRptttkZ2cP29atP/zhD2HvXXrppXj00Ucxe/bsiI41UgzORKRPBFfOfn94GYSMjAxkZGQMe6+3txcAkJ5u/xj9xImD93J6eoTyCGe47LLL8POf/xyFhYWYNm0auru70draisceewwHDx7Ebbfdhl27diEzM9PV8aKBwZmItFGm6bhGZKi9pCR83cs1a9Zg7dq1w7f/co7aMCwSAjy69dZbh/3/CRMm4KKLLsLixYuxYsUKNDY2YvPmzfjlL38ZtXM6YXAmIn1M0zkb48vgXFNTg6ys4UXCzr5qBr66Kg5dQVsJXTGHtvVq3LhxWL16Ne666y689tprIzpWpBI3OBtJgNUkvbRKcWg/2yb7NjEtawTnlNLlpNQ0p3k602NivzgGTilW0thKY+B0r0UYX6m/XsfgywN7OueIeE1hE3bz0teofr4IpjWysrIwY8YMx0NOnz4dAHD8+HHbbUJTJKFtRyI3NxcA0NnZOeJjRSJxgzMRxT8ND6HMnTsXAHDo0CH09fVZZmw0NTUBAObMmRPRsa10dXUBGPlVeKSY50xE+oQe35ZeEf6Fk52djXnz5qG/v9+y3kVDQwP8fj98Ph8WLFgw4o/w0ksvAQDy8/NHfKxIMDgTkT6hwkfiK/Lpp9WrVwMANm3ahI6OjqH3T506hfLycgDAqlWrhuUnV1ZWoqioCJWVlcOOdeDAAfzlL39B8Ky58YGBAWzbtg3PPPMMgPCbhrpxWoOI9NHwEAoAFBUVYdmyZairq0NxcTEWL148VPiou7sbS5YswfLly4ftEwgEcOTIEQQCgWHvHzt2DHfffTcmT56MnJwcZGZmoqenB++99x5OnDiBpKQkrFu3DldffXXE/RwJBmci0kYNBKEG5GwNp3Y7ZWVlWLhLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 2 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAWMAAAENCAYAAADaPATLAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAAAvNklEQVR4nO3de3AU150v8G+PXkhgIRCjJw9J5o3YWiEvWPJ6kzg4kZ3SzQZC7VWEZLO5UHWx2VTtyrv24rUlKlur2MhEqdpNUcSIa/S468LB16wNFOUilLPWInu9SvSyABtpDURikC3DjB5I0+f+MR4ZMXNO94ymxYz0/aRUFffp032mJX5qnf7172hCCAEiIrqnbPd6AERExGBMRBQWGIyJiMIAgzERURhgMCYiCgMMxkREYSD6Xg+AiChQn376Kd577z20tbWhvb0dPT09EEKgtrYWRUVFQR/3xIkTaGpqQnd3N3RdR3Z2NrZu3YqSkhLYbPJ712D73YnBmIgiTlNTE1577bWQHrOqqgqNjY2Ii4tDQUEBoqOj0dzcjH379qG5uRm1tbWIiooKWb+7MRgTUcRZuXIlfvzjHyM3Nxe5ubnYu3cvWlpagj7e6dOn0djYCLvdjvr6emRlZQEAbty4gfLycpw5cwb19fV44oknQtLPH84ZE1HE2bZtG/72b/8Wjz/+OJYuXTrl4x08eBAAUFFRMRFQAWDRokWorKwEABw6dAi6roeknz/TFoxPnDiBH/3oR8jPz0deXh62bNmChoYGU4MkIrJKX18fOjo6EBMT43e+eePGjUhNTYXD4UBra+uU+8lMSzCuqqpCRUUF2tvb8cADD6CwsBA9PT3Yt28f/uqv/gput3s6hkFE5KOzsxMAsGLFCsyZM8fvPuvXrwcAdHV1TbmfjOVzxqGcUyEiCrUrV64AADIyMqT7pKenT9p3Kv1kLA/GRnMqZWVlOHToEMrKykylgIyMjKC9vR12u93UE0oiCozb7YbD4UBubq70js+MwcFBOJ1OU/sKIaBpms/2xMREJCYmBj0GM4aGhgAA8fHx0n3mzp0LAHC5XFPuJ2NpMDY7p9Lf34/W1lZs2LDB8Jjt7e0oLS21YrhEdIeGhgY88MADQfUdHBzEo5u/jZu3zAXjuLg4jI6O+mx/+umnsWfPnqDGYJa3irC/XwZW9JOxNBibnVPp7+9HV1eXqWBst9sBAP/nn19GWsoin3bd9YWyvxiR/3Boc+Ypek5hel1xx99TUiNtW1r9qLRNS16sPKWWIL+bELfk10h89Bv5MfP+TH3OpFT5cYcG5W0D15THtS1ZK23T3/1X+XgKH5O2fflcrfKcSa88J20Tg33ytjHfgOJlS71feU594DNp259uf1naNjI+Jm376IS8HwD8qrzJZ5tLu41Tcd0T/9aC4XQ6cfOWE6/988+Qavf9d3qnfscNlD/1d2hoaEBaWtqkNqvvioGv7169d7r+eO9svftOpZ+MpcE41HMqACamJtJSFiEz3fcfv35L/ZHEcJy0TYtXfOM1a4LxEGKkbZnJ8+XDSUlWnlKbmyRtE/Hy3+R6ovxPLlvKQvU5F8r/8QqnfEpJaOo/4Wxp8uO6leOVB4GEKPXPycJUxWeJuS1vuz0sH4/icwCAbpPfKOhu+fdM9fw7M1UdCOcJ+b+HUEwDpi5aiEyDzw3hyahKS0vD4sXqmwwrZGZmAgCuXZPfFPT19U3adyr9ZCzNpgj1nAoRRRhdN/d1D61d6/nL6+LFixgZGfG7T1tbGwBgzZo1U+4nY2kwDvWcChFFFgEBIXT1F+7tym/p6elYt24dxsbGcOrUKZ/2lpYW9PX1wW63Iy8vb8r9ZCwNxqGeUyGiCOMeN/c1DWpqalBUVISaGt/nNLt27QIA7N+/H729vRPbBwYGUFVVBQDYuXOnT8ZXsP38sXTOONRzKnfSXV/4nR+23aeeS9Xd8ocdWqwijWdKc8byubfoaMWfaIrxaLHyqR9Pu7yvUPWNk88hqsZjOCZFmzA6bpziF3WMfM5ddQ20KPXdWNDXT8i/n4bfsxh5u4Yg/7o0OKe/n+qQ/h2r64Bu8FJXENMUHR0dE8EOAC5dugQAOHDgAA4fPjyx/fXXX5/4/w6HA5cvX4bD4fA5XlFREUpKStDU1ITi4mIUFhZOFPxxOp3YvHkztm/fHrJ+/lgajO+eU/GXURHInAoRRRihK39BTewTIKfTid/97nc+23t6egI+lldlZSXy8/PR0NCAlpYW6LqOnJwcw1KYwfa7m6XB2Dun0tHRgVOnTuHP//zPJ7UHOqdCRBFGmHhAF0Qw3rRpE7q7uwPqU11djerqauU+xcXFKC4uDng8wfa7k+W1KUI5p0JEkcXw4d1XXzQNr0OHck6FiCKMLozvjPV7m00RLqaluHyo5lSIKMK4xzxfRvvQ9K30EYo5FSKKMBY9wJuJuOwSEVlHmJimEJymACI4GIsRp986E6o8YgCwJaVJ23RFARir8oxv31a8/z8if1lGDN8yOKdivCM35W3D8toKMDinakzK8Q6rK3sp+0peQwUAMSp/xd49qv5+Bv1ZFLUpjL5nYkRxziDfUhPDiu81gHE/ScWKMhhBDIB3xmZFbDAmoghgpvYEl14DwGBMRBYS+jiErv5rVejT8zp0uGMwJiLrWPTSx0zEYExE1uGcsWkMxkRkHYsKBc1EDMZEZB3eGZsWscFYmzPP7zJJyjKYUKevqdLeVMvpGFKktsXGKu4a5iRIm7T4+5SnVLWLMfmyQVCsyoIpnFP1lpWIV609aHBcxerFqtKbUXHqAKA6p/LaKlIKDb9ncxTnDLKwpXIpMQDRfjLmQrrmOrMpTIvYYExEEUB3GxePN5rGmCUYjInIOrwzNo3BmIgsI4QbQqjvfI3aZwsGYyKyDktomsZgTETWYTaFaQzGRGQdvoFnWgQHY5v/SmpG1dUU7ar0NdXKvmJckSpmZkzB9DMqyK/oqwVbzF8zSHpSjXcqVe9UVJ9FkVJoOJwgx6tN5RpYsciCwTn9tYZ0FG4T2RRuzhkDER2MiSjscZrCNAZjIrIOU9tMYzAmIutwztg0BmMiso4QJqYpmNoGMBgTkZXc4yYe4LG4PMBgTERW4pyxaQzGJqnS17ToWHVf/uZXU6SgUaQzMU0R5GKrMw2DMRFZh3fGpjEYE5F1GIxNYzAmIusIYZwtwWwKAAzGRGQltxsY5+vQZjAYE5F1+Dq0aQzGRGQdzhmbxmBMRNbhnLFpkRuMbTb/JQeNclZV7ao2VelNgzxiLUp+mTVN8YOoKqloWI4xuM8CTb4KsVHpTVW7mEIucfAlP6dQDDLY66daHXoKZU+DZkVZzkBwpQ/TIjcYE1H4Y6Eg0xiMicg6bjeEUbYEsykAMBgTkZX4AM80BmMiso7FJTRPnDiBpqYmdHd3Q9d1ZGdnY+vWrSgpKYHN5Hz5+fPnUV5ebmrfs2fPIiMjY+K/n332WRw/fly6f3Z2Nk6dOmXq2AzGRGQdIYwf0AUZjKuqqtDY2Ii4uDgUFBQgOjoazc3N2LdvH5qbm1FbW4uoKOMHx4sWLcIPfvADafvvf/97fPLJJ1i6dCnS09P97rNhwwYsW7bMZ7vdbjf9eRiMicg6Fk1TnD59Go2NjbDb7aivr0dWVhYA4MaNGygvL8eZM2dQX1+PJ554wvBY999/P6qrq6Xt3/ve9wAAW7duhSbJNtq2bRu2bNkS8Oe4U8QG456SGgwhxmd7dLT6G3v7tvw3ZWysNQ8SVOlrOb/9Z2nbh3/0jLRtbpx6RWrVdbh9W/5t7x2dK21bFlejPGd8vHxMo6Pyc7pG1SVIFybWSds6v1wobVvxTy9K295wL1ae8y++8dfStqFh3587r3G3/E/j+feNKM/pdMVJ21xj8r7juvzn9uK3/kF5zneFb98xjCn7BMTtNn5AF8QDvIMHDwIAKioqJgIx4LnLraysRFlZGQ4dOoSysjLT0xX+/Nd//RcuXbqEqKgo5d1zKNzjJEQimtGE+PruWPYV4DRFX18fOjo6EBMTg6KiIp/2jRs3IjU1FQ6HA62trVMa/htvvAEAePjhh5GamjqlYxmJ2DtjIooAuok54wBf+ujs7AQArFixAnPmzPG7z/r169Hf34+uri5s2LAhoON7DQ8P45133gEA/PCHP1Tue/78eXR3d2NoaAjJycnIz8/HQw89FNBdOYMxEVkngEJBfX19Pk2JiYlITEyctO3KlSsAMCmr4W7eB23efYNx6tQpuFwuJCcn45vf/KZy3zfffNNn2/Lly/HKK69g1apVps7HYExE1gngzri0tNSn6emnn8aePXsmbRsaGgIAxMfHSw85d67n2YfL5QpktJN4pyi+//3vIybG/3OC1atX4/nnn0dBQQEyMjLgdDrR2dmJAwcO4OOPP8aOHTtw/PhxU1McDMZEZBmh6xAG2RLe9oaGBqSlpU1qu/uuGADEV3PMssyGUOjt7cUHH3wAQD1F8eSTT07674SEBKSkpKCwsBBlZWVobW3FwYMH8cILLxiek8GYiKyj68bZEl8F47S0NCxerM5yAb6+6/XeIfvjvSP27hso711xXl4e7r///oD7x8bGYteuXdi9ezfOnTtnqk/EBuOl1Y8iM3m+b0Os/wn9CSPybyDmJMjbplQBTN5Xlb72wO9flraN/+c76nPO9XNtvJxfSJuWN78vbYva9C31ORco/hRz3ZQ2ievqeT0te620LfPffi1ti/r296VtJU/J+wHAkkM/kjd+cV3aJEblKWhaZo7ynFBch/jyj6Vtw4qVy+8/JL8GAPAnf/kbn20ubRRXEfxc6yQWPMDLzMwEAFy7dk26j3f+2btvINxu98Qc8NatWwPu75WT4/l+9/f3m9o/YoMxEUUAC176WLvW88v54sWLGBkZ8ZtR0dbWBgBYs2ZNQMcGgN/+9rfo7+9HQkICHn/88YD7ew0ODgIwf3fOPGMiso73dWjVV4B5xunLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 2 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAWMAAAENCAYAAADaPATLAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAAAvjElEQVR4nO3df1BV550/8Pe9gIga8Bfyw1iVKFHBzvgjKrjdyTS0ZTPLTms2OzUKars4UyPpP7Rp1jYBm90liTahnU3WMAYbQbJOdu3UyUbH6TRqKhVTl3wREE2iTv1x8apBBcTAPc/3j5tLvN5zPufcwz16L7xfM8wk9znPOc899/Lh+JzP+TwupZQCERHdV+77PQAiImIwJiKKCgzGRERRgMGYiCgKMBgTEUUBBmMioigQf78HQEQUrs8++wyHDx9GS0sLTpw4gbNnz0IpherqahQWFtre7969e9HQ0ICOjg5omoaZM2fiiSeewMqVK+F2G1+72u13JwZjIoo5DQ0NePvttyO6z8rKSuzatQuJiYnIy8tDfHw8GhsbsXnzZjQ2NqK6uhpxcXER63c3BmMiijnZ2dn44Q9/iNzcXOTm5mLTpk1oamqyvb/9+/dj165dSE1NRV1dHWbMmAEAuHLlCkpKSnDgwAHU1dVhzZo1Eemnh3PGRBRznnzySfz0pz/F448/jq997WtD3t+2bdsAAOXl5YMBFQAmT56MiooKAEBNTQ00TYtIPz33LBjv3bsXTz31FBYtWoQFCxZgxYoVqK+vtzRIIiKneDwetLa2IiEhQXe+ecmSJUhLS4PX60Vzc/OQ+xm5J8G4srIS5eXlOHHiBBYvXoz8/HycPXsWmzdvxjPPPAOfz3cvhkFEFKKtrQ0AMHv2bIwePVp3m/nz5wMA2tvbh9zPiONzxpGcUyEiirTz588DADIzMw23ycjICNp2KP2MOB6MzeZUiouLUVNTg+LiYkspIH19fThx4gRSU1Mt3aEkovD4fD54vV7k5uYaXvFZ0dXVhe7ubkvbKqXgcrlCXk9OTkZycrLtMVjR29sLAEhKSjLcZuzYsQCAnp6eIfcz4mgwtjqn0tnZiebmZixcuNB0nydOnMCqVaucGC4R3aG+vh6LFy+21berqwvfKngMN25aC8aJiYm4fft2yOsbN25EWVmZrTFYFagirPfHwIl+RhwNxlbnVDo7O9He3m4pGKempgIAfvsfryB9yuSQ9mWPrhP7u21Ok/uU8Y3G275+sW+82/gKPjEuwbBNKjWtIJeh1oS+0nuRbqiOjh8lHvMLbcCwLUE4B/2afM9AGlNivPH5u3brpmHb5DEp4jGlc98nfN6jhc/z8z45ME0YPU5sN9J12/iqa3ziWLGv3vfI5VZ4YNJXv2t2dHd348bNbrz9Hy8hLTX09/ROnd4rKHn6WdTX1yM9PT2ozemrYuCrq9fAla6ewJVtYNuh9DPiaDCO9JwKgMGpifQpkzE1Iy2kXfPJf6Xs/hWTkj7M7j+6hbiphPFIZf81kyUBpGAs9fUJ71O55XMnxVRNGfc1S6iRxqQJ529gQPhjpsnvRTz3wvuUPk9pPFbGZGe/KsHkM9Pp6r9cURGZBkybPBFT002C+pcXB+np6XjwwQeHfMxwTZ06FQBw8eJFw208Hk/QtkPpZ8TRbIpIz6kQUYzRNGs/99G8efMAAKdPn0ZfX5/uNi0tLQCAuXPnDrmfEUeDcaTnVIgotigoKKXJPyZTbk7LyMhATk4O+vv7sW/fvpD2pqYmeDwepKamYsGCBUPuZ8TRYBzpORUiijG+AWs/98DWrVtRWFiIrVu3hrStX78eALBlyxacO3du8PWrV6+isrISAFBaWhqS8WW3nx5H54wjPadyp2WPrtOdH77wyXtiPzXwhXGjMBmovtD/ZwgAKOFGEQC4Rhv/oXElPSD2NWRy00t8n/2hd60HCf1cKVPkY/Z0GfcVzoHqM5miksY0boJhm++8caJ93PSvy8eUvgu9143HI9wY1C6dFg/pzpgtNBrP32qeT427pT8kHhM6NyMveC6jcNXTcj+rNM30u2pnmqK1tXUw2AHAJ598AgB49dVX8dZbbw2+vnv37sH/9nq9OHPmDLxeb8j+CgsLsXLlSjQ0NKCoqAj5+fmDBX+6u7tRUFCA1atXR6yfHkeD8d1zKnoZFeHMqRBRjFHa4A06cZswdXd34+OPPw55/ezZs2HvK6CiogKLFi1CfX09mpqaoGkasrKyTEth2u13N0eDcWBOpbW1Ffv27cN3v/vdoPZw51SIKMYoCzfobATjpUuXoqOjI6w+VVVVqKqqErcpKipCUVFR2OOx2+9OjtemiOScChHFFtObd1/+0D14HDqScypEFGM0ZX5lbJY0P0Lck+LykZpTIaIY4+vXvUkYsg3du5U+IjGnQkQxxqEbeMMRl10iIucoC9MU0rPnI0jMBmM33LpP9on5tQBcUrEb6S+01DYgF9BBnPFpdgltEuWSp3ZcwniVzcdPzcaqhCI5kNoSEm2Nx2y/0mctfg8A8fNW8cbjdUnvRToHZn0lQ/h+6YZBdwTDAq+MLYvZYExEMcBK7QkuvQaAwZiIHKS0AShNvkGnhNKrIwmDMRE5x6GHPoYjBmMicg7njC1jMCYi5zhUKGg4YjAmIufwytiy4ReMzf4KSx+8lC4m/fUeTn/ZnfrFENdkGsIvq1Tq0m6qIkzS/xxIGxxq36jFbArLhl8wJqLoofnMi8ebXUCNEAzGROQcXhlbxmBMRI5RygelTFalMWkfKRiMicg5LKFpGYMxETmH2RSWMRgTkXP4BJ5lMRuMfUrT/YylVZz9G9hLUZNWcdauXxYP6RKK5yubVdvMvuB2V4dWA0LbF7fkMQntSlpAQDim6XFHJRm33RbGY/ZeJOJ4QhfdtXxMqV06f9J5Nzmm0sl0UNLq4eHyWcim8HHOGIjhYExEMYDTFJYxGBORc5jaZhmDMRE5h3PGljEYE5FzlLIwTcHUNoDBmIic5BuwcAOPxeUBBmMichLnjC2L2WB829evmxGjbt2UO0qLhwpfCil9LS79IfGQvkun5THZYZraJqQn9Qtpb1Jqm5RGBpNzL139mKV8SWNyxRn3u9Vt3M/seyKcX7GvVPnP5JimY7LRz3Sfep9Ln/F5C5+FaQr9ZVFHnJgNxkQUA3hlbBmDMRE5h8HYMgZjInKOUubZEsymAMBgTERO8vmAAT4ObQWDMRE5h49DW8ZgTETO4ZyxZQzGROQczhlbFrPBON4dB7fOZ+gaPVbuaLNkpVQG0yyPOC5jtmGb1nvd1njMribEsp1Sbq50fkzyjKX8ZfFzkcpDAlD9xmOS9qtGjzHulyh/T8SVpaWcael9CuU1TftKpM/F7DPTey9m4wwHV/qwLGaDMRHFABYKsozBmIic4/NBmWVLMJsCAIMxETmJN/AsYzAmIuewhKZlDMZE5BylzG/QMRgDYDAmIic5PE2xd+9eNDQ0oKOjA5qmYebMmXjiiSewcuVKuE2ydO528eJFvPnmmzh8+DA6Ozsxbtw4zJ8/H2vXrsXy5ct1+/zsZz/Dnj17DPc5c+ZM7Nu3z9LxYzYYJ8YlQLlcIa9LqzgDgMtmapvtVZwhp6+5x6QIBxXKOJp9gX1CepLNlaPFsQKQRuQSUqzMzq0r3rjsqfR5220DAJdw7jXh3LuF/aox48VjSn0lSvhczPaptzq0K7HX1jh0+XzmN+hs3sCrrKzErl27kJiYiLy8PMTHx6OxsRGbN29GY2MjqqurERcnpHHe4eOPP0ZpaSmuX7+OqVOn4tFHH8Xly5fx4Ycf4tChQygvL0dpaalh/4ULF2L69Okhr6emplp+PzEbjIkoBigLecY2pin279+PXbt2ITU1FXV1dZgxYwYA4MqVKygpKcGBAwdQV1eHNWvWmO7r9u3beOaZZ3D9+nUUFxfjueeeGwzif/7zn/GjH/0IW7ZsweLFi7FgwQLdfTz55JNYsWJF2O/jTuFdxxMRhUNT1n7CtG3bNgBAeXn5YCAGgMmTJ6OiogIAUFNTI/4rJuDAgQPweDyYNm0ann322aCr6WXLlmHt2rUAgDfeeCPscYaDwZiInBMoFGT2EwaPx4PW1lYkJCSgsLAwpH3JkiVIS0uD1+tFc3Oz6f5aWloG+yUkJIS05+fnAwCOHDmC7u5IroISjNMUROQcK1e+YV4Zt7W1AQBmz56N0aP1743Mnz8fnZ2daG9vx8KFC8X99fb658gnTJig2x54vb+/H6dOndLd39GjR9HR0YHe3l5MmjQJixYtwvLly8O6ichgTESOUZpmerM50O7xeELakpOTkZycHPTa+fPnAQCZmZmG+8zIyAjaVjJx4kQAwF//+lfd9jtfP3/+vG4w/t3vfhfy2qxZs/CrX/0KDz/8sOkYAAZjInKSpplnS3wZjFetWhXStHHjRpSVlQW9FriSTUoyztAZO9ZfdKmnp8d0iMuWLcN//ud/4uDBg/B4PEhPTw9qf+eddwb/++5pijlz5uDnP/858vLykJmZie7ubrS1teHVV1/FyZMnsW7dOuzZswdpaWmm44jZYKyUGh654tJ8mbTSsJhIZrJfu3mdZnN7fKxV5lRBnGgutBPGNEV9fX1IILz7qhjw/+4DgEsntdWOvLw8PPLIIzh27Bh+8IMf4Be/+AXmz58Pr9eL7du344MPPkB8fDwGBgZCph0CN/cCxowZgylTpiA/Px/FxcVobm7Gtm3b8Pzzz5uOI2aDMRHFgDAe+khPT8eDDz5ousvAVW/gCllP4Io4sK2Z6upqlJWV4S9/+UtIgC0uLsaxY8dw8uRJpKTIufYBo0aNwvr167FhwwYcPHjQUh8GYyJyjgOPQ0+dOhWA/4k5I4H558C2ZiZNmoT6+nocOXIER48exeeff46JEyfiscceQ05ODhYvXgwAyM7OtjzOrKwsAEBnZ6el7RmMicg5DhQKmjdvHgDg9OnT6Ovr082oCKSLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 2 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAWcAAAENCAYAAADT16SxAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAAAmnklEQVR4nO3df3AU5eE/8PflByGiASxnEqA2ZCRKEmYawNDg6Gc+gja137QWxnbSkKCjMB0EnalxvrW2NWHoTNREja12GEagkB8tVXHKRwWp7Yc6NRKpRZMAAYSkA/HiwRhrEiK52+fzR7zocbvP7u3dc9nLvV+dm7H37D777JK8WZ599nlcQggBIiJylKSJbgAREYViOBMRORDDmYjIgRjOREQOxHAmInIghjMRkQOlTHQDiIjCdfr0abz11lvo6OhAZ2cnenp6IIRAY2MjSktLbde7d+9etLa2oru7G5qmYd68eVi1ahXKy8uRlBTbe1mGMxHFndbWVuzcuTOqddbW1qKlpQVpaWkoKSlBSkoK2trasGnTJrS1taGxsRHJyclRPaYMw5mI4k5eXh7uvfdeFBYWorCwEI8++ija29tt17d//360tLTA7XajqakJOTk5AIDz58+jqqoKBw4cQFNTE9asWROlMzDHcCaiuHPXXXdFtb4tW7YAAKqrq8eDGQBmzZqFmpoaVFZWYuvWraisrIxZ90bMOlH27t2LH//4x1i8eDGKioqwcuVKNDc3Q9O0WDWBiCiEx+NBV1cXUlNTdfuri4uLkZmZCa/XiyNHjsSsXTEJ59raWlRXV6OzsxNLlizBsmXL0NPTg02bNuGBBx6A3++PRTOIiEIcPXoUADB//nxMnTpVd5uFCxcCAI4dOxazdinv1nBiXw4RUcDZs2cBALNnzzbcJjs7O2jbWFAeztHuyxkZGUFnZyfcbndMn5wSJQq/3w+v14vCwkLDO0krBgYGMDg4aGlbIQRcLlfI9xkZGcjIyLDdBiuGh4cBAOnp6YbbTJs2DQAwNDSktC1fpTScrfbl9Pf348iRI1i0aJFpnZ2dnaioqFDRXCL6iubmZixZssTWvgMDA7htxXL85zNr4ZyWlobPP/885PsNGzZg48aNttpgVWDWZL2/HCaS0nC22pfT39+PY8eOWQpnt9sNAPj9c08i65pZ0WusCeG7ZFwoKwPg/+cbhmUpN6202yTbZOfiSpb8SLjUPKIQoyPScleq5O5NGD9QLvnv+wzL2v53m2m77FjwrR8ZlmmStgJA96E/GZbJrpH0+tjg+fg81tz/8Pjvmh2Dg4P4z2eD2Pnc48h0y39P+73nUXX//0dzczOysrKCylTfNQNf3hUH7qD1BO6YA9vGgtJwVtGXE+jKyLpmFuZkZ0bYQuuk4Twa+jf+V/lmXmlYlhrDcwhwXDhfuigtd00x/uemLJw1v/GdkKqfHZ/PeO0KzWRdC1mbZNdIen0iEI1uw8xZV2NOlknIf/FnmJWVhblz50Z8zHDNmTMHANDX12e4jcfjCdo2FpSGs1P7cogoRjRt7GO2zQTKz88HAJw8eRIjIyO6/8rv6OgAACxYsCBm7VI6lM6pfTlEFBsCAkJo8g8mdqW87OxsFBQUYHR0FPv27Qspb29vh8fjgdvtRlFRUczapTScndqXQ0Qx4vdZ+8RAQ0MDSktL0dDQEFK2bt06AEB9fT16e3vHv79w4QJqa2sBAGvXro3p5EdKuzWc2peT8DTJSz+yPmcHEnzDNDJ6ffYmDy7Domnyn7fANmHq6uoaD00AOHXqFADg6aefxrZtXz7s3b179/h/e71enDlzBl6vN6S+0tJSlJeXo7W1FWVlZVi2bNn4xEeDg4NYsWIFVq9eHXY7I6H0N9GpfTlEFCNCMw97G38ZDA4O4v333w/5vqenJ+y6AmpqarB48WI0Nzejvb0dmqYhNzd3ck4ZGujL6erqwr59+3DnnXcGlU9UXw4RxYiw8EDQRjgvXboU3d3dYe1TV1eHuro66TZlZWUoKysLuz0qKP+rwIl9OUQUG6YPA7/4UCjlHYxO7MshohjRhIWhdBM7WsOpYvL0x2l9OUQUI/7RsY/ZNhQiZo/mndSXQ0QxouiBYCKIr3FTRBRfhIVuDZPX2hMVwzkRJTlsqtUI2uNil1hk9OZLieYcKrxzto3hTETqxMHcGk7FcCYiZYTmg9DkD/yEFpvXt+MNw5mI1FH0EkoiYDgTkTrsc7aN4UxE6iia+CgRMJyJSB3eOdvGcCYidThawzaGMxGpo/nNJ9M36/ZIUAxnIlKHd862MZyJSBkh/BBCfmdsVp6oGM5EpA6nDLWN4UxE6nC0hm0MZyJSh28I2ha34Sx8lyB8l2J2PFfKFMMy7eJn8p0lT6tjeQ7jx7x00bhQcp6ms8fZfOoe0TWQTNSeKmmvquueIjmmX7ik+8raJG1vJLMM6vyZidER+/Vdzm9htIaffc564jaciSgOsFvDNoYzEanDoXS2MZyJSB32OdvGcCYidYSw0K3BoXR6GM5EpI7fZ+GBICfb18NwJiJ12OdsW/yGs+8SMPp5zA4nGy6XdNXXpPv6ZUOFYngO42RD+yS7uUwW/hR2+w7NroHkuEJyLrJhbaque0qy8TFdmnwonbRNsrIIFmTV+zMTPvmyUmEewUKfMrs19MRvOBOR8/HO2TaGMxGpw3C2jeFMROoIYT4ag6M1dDGciUgdvx/w8fVtOxjORKQOX9+2jeFMROqwz9k2hjMRqcM+Z9viNpz9/3wDvplXxvCAxv1m0nHMAFK/c59h2eibu2w3ybZRyTjWFMnYYLPxtHb/eeoz6XNMTbV1zEXTrjU+5NuvmDRKQnLMpTOvMywzuzrSNsn+zGTXx4zOuWifDNmvL6QyroRiV9yGMxHFAU58ZBvDmYjU8fshzEZjcLSGLoYzEanDB4K2MZyJSB3FU4bu3bsXra2t6O7uhqZpmDdvHlatWoXy8nIkJVmfc+RnP/sZ9uzZY1g+b9487Nu3z3Y77WA4E5E6Qpg/8LMZzrW1tWhpaUFaWhpKSkqQkpKCtrY2bNq0CW1tbWhsbESyZCIqPYsWLcI3vvGNkO/dbretNkaC4UxE6ijq1ti/fz9aWlrgdrvR1NSEnJwcAMD58+dRVVWFAwcOoKmpCWvWrAmr3rvuugsrV64Muz0qxG04p9y0EqnZmTE7nnT1Y5PpJ2XD5VKXV9ptkm2yc3ElS34kIpiaUka6GjgA15R0yc7Gv9j/+PRPhmWpy58xa5YtBz/eZlimmfzzPnX57wzLZNdIen1sSPmoH2j8n+hU5vebP/Cz8UBwy5YtAIDq6urxYAaAWbNmoaamBpWVldi6dSsqKyvD6t5wkvhsNRHFByG+vHs2+oTZreHxeNDV1YXU1FSUlpaGlBcXFyMzMxNerxdHjhyJ0onEXtzeORNRHNAs9DmH+RLK0aNHAQDz58/H1KlTdbdZuHAh+vv7cezYMSxatMhy3YcOHUJ3dzeGh4fxta99DYsXL8ZNN900IXffDGciUieMiY88Hk9IUUZGBjIyMoK+O3v2LABg9uzZhlVmZ2cHbWvVK6+8EvLdddddh6eeegrXX399WHVFiuFMROqEcedcUVERUrRhwwZs3Lgx6Lvh4WEAQHq6cV/7tGnTAABDQ9ZeRb/hhhvwi1/8AiUlJZg9ezYGBwdx9OhRPP300zh+/Djuuece7NmzB5mZsXvOxXAmImWEpkGYjMYIlDc3NyMrKyuo7PK7ZgAQX/RRu1wmazKG4e677w76/1dccQWuueYaLFu2DJWVlThy5Ai2bNmCX/3qV1E7phmGMxGpo2nmozG+COesrCzMnTvXtMrAXXHgDlpP4I45sK1dU6ZMwbp167B+/XocPHgworrCxXAmInUUPBCcM2cOAKCvr89wm0D/dWDbSOTm5gIA+vv7I64rHAxnIlJHwUso+fn5AICTJ09iZGREd8RGR0cHAGDBggVh1a1nYGAAQOR34eHiOGciUifw+rbsE+Y45+zsbBQUFGB0dFR3vov29nZ4PB643W4UFRVFfAqvv/46AKCwsDDiusLBcCYidQITH0k/4c+tsW7dOgBAfX09ent7x7+/cOECamtrAQBr164NGp/c0NCA0tJSNDQ0BNV17Ngx/O1vfwtZNMPn82H79u3YtWvsDd/LHxqqxm4NIlJHQZ8zAJSWlqK8vBytra0oKyvDsmXLxic+GhwcxIoVK7B69eqgfbxeL86cOQOv1xv0/blz53D//fdjxowZyMnJQWZmJoaGhnDixAl8/PHHSEpKQnV1NW6++eaw2xkJhjMRKSN8fgiTZcjMyo3U1NRg8eLFaG5uRnt7OzRNQ25ubthThl5//fWoqqpCR0cHzp07h6NHj8LlciErKwsrV65ERUVFzLs0AIYzEamkcMpQACgrK0NZWZmlbevq6lBXVxfy/de//nU8+uijttugCsM5EWmSOxXZrHQOZPaCA5nQe7U6mmv6hfH6NgWLr99EIoovGiz0OcekJXGH4UxEyghNQJiEs1l5omI4E5E6fj/gM3kwx9W3dVkO59OnT+Ott95CR0cHOjs70dPTAyEEGhsbdSe8/qpoLcJIRHFG0VC6RGA5nFtbW7Fz586wD6BiEUYiihMMZ9ssh3NeXh7uvfdeFBYWorCwEI8++ija29ul+6hahJGI4oMQYnyKT9k2FMpyON91111hV54IizASkYSAhXHOMWlJ3FGWiImyCGNcSko2/sQZV1KS4YcscCXpf6LFbNIjK90eCUrZT7DVRRiBsYlHiGjyET7N0odCKRtKp3IRRiKKEwLmL5nwxlmXsnBWsQgjEcUXvoRin7JwVrEIIxHFGQ6ls01ZOMdyEUYicigN5t0a7HLWpSycY70IIxE5jxAWujU4zlmXsnBWvQij8F2C8F0KLZBNhwnYHi4mLl00LvT75DuPjhrXq3cOVkRwnq6UKcbt+VzS/5+cGlmbjI5pdg1kf2Z+42ubKtnLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 2 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"for est_covar in covar_list:\n",
|
||||
" est_corr = torch.diag(est_covar.diag().sqrt().pow(-1.0)).matmul(est_covar).matmul(\n",
|
||||
" torch.diag(est_covar.diag().sqrt().pow(-1.0))\n",
|
||||
" )\n",
|
||||
" f = plt.imshow(est_corr)\n",
|
||||
" plt.colorbar(f)\n",
|
||||
" plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.8"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 4
|
||||
}
|
||||
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
@@ -0,0 +1,350 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "209634f4",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"[pyKeOps]: Warning, no cuda detected. Switching to cpu only.\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns\n",
|
||||
"import torch\n",
|
||||
"import gpytorch\n",
|
||||
"# from voltron.robinhood_utils import GetStockData\n",
|
||||
"import os\n",
|
||||
"# import robin_stocks.robinhood as r\n",
|
||||
"import pickle5 as pickle\n",
|
||||
"\n",
|
||||
"sns.set_style('white')\n",
|
||||
"# style.use('whitegrid')\n",
|
||||
"palette = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"sns.set(palette = palette, font_scale=2.0, style=\"white\", rc={\"lines.linewidth\": 4.0})"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 19,
|
||||
"id": "7949a79f",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAcwAAABECAYAAAAMTwWHAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAAClElEQVR4nO3arY4TURjH4XdpIc3CrthlwofgCioQFSgUigSNqMByLXgsqDo8N7ECMQJuAJY0ILZkKbDDIAgfTSh5BWcP2zyPnJMm/yaT/DLTbvV93wcA8Ffnag8AgLNAMAEgQTABIEEwASBhuO5guVxG27bRNE0MBoPT3AQAVXRdF/P5PMbjcYxGo5WztcFs2zam02nxcQDwv5nNZjGZTFaurQ1m0zQREfFx/3b0g+2yyyp5/Ohh7QlF3X/+tPaEop7c28z7MiLi1f0XtScUdePurdoTirr58lHtCUXt3VnUnlDM4WIYD57d+NnA360N5o/XsP1gO/rhxXLrKmquXK89oaiTnUu1JxS1f3Uz78uIiJ24UHtCUXvbu7UnFHXt/NfaE4q6vHtSe0Jxf/op0p9+ACBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIGG47qDruoiI2OqOT23MaZu/fV17QlHDxYfaE4p6d/i19oRiFvG59oSi3h8f1Z5Q1Jsvm/0s8ulobTrOvMPF9+/2o4G/2+r7vv/Thw4ODmI6nZZdBgD/odlsFpPJZOXa2mAul8to2zaaponBYHAqAwGgpq7rYj6fx3g8jtFotHK2NpgAwC+b/aIdAP4RwQSABMEEgATBBICEb0f+ZBgefApDAAAAAElFTkSuQmCC\n",
|
||||
"text/plain": [
|
||||
"<Figure size 576x72 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"sns.palplot(palette)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "cc68c104",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"mt_preds = torch.load(\"multitask_predictions.pt\")\n",
|
||||
"ind_preds = torch.load(\"ind_predictions.pt\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "8b069e6c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"id": "4bdc7816",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"mt_emp_percentile = (mt_preds[\"paths\"] > mt_preds[\"test_y\"].unsqueeze(-2)).sum(1) / 100\n",
|
||||
"ind_emp_percentile = (ind_preds[\"paths\"] > ind_preds[\"test_y\"].unsqueeze(-2)).sum(1) / 100"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 25,
|
||||
"id": "a2ee0ccf",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"torch.Size([10, 200])"
|
||||
]
|
||||
},
|
||||
"execution_count": 25,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"mt_emp_percentile.shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 52,
|
||||
"id": "b4b40349",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"bins = np.arange(0, 1, 0.05)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 57,
|
||||
"id": "16115784",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"reshaped_mt_emp = mt_emp_percentile[:5].reshape(-1)\n",
|
||||
"reshaped_ind_emp = ind_emp_percentile[:5].reshape(-1)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 58,
|
||||
"id": "31ba854e",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"mt_counts, _ = np.histogram(reshaped_mt_emp, bins=bins)\n",
|
||||
"ind_counts, _ = np.histogram(reshaped_ind_emp, bins=bins)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 60,
|
||||
"id": "e9a6963f",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"<matplotlib.legend.Legend at 0x7fb6f77a0710>"
|
||||
]
|
||||
},
|
||||
"execution_count": 60,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYsAAAEFCAYAAAASWssjAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAAsyElEQVR4nO3de3wM978/8NduruLSiETEXbEbIUhE3CpCxe1UlSI9IhLlSzjke1AaVceXVGiPtvmqFq1Lq7RpOaqlCK24hQYRdQ+pW4JECCKJ3Hbn94ffTq3dmJXMJpvN6/l49OHrM5+Z+cx8x7z3c5nPRyEIggAiIqLnUFZ1AYiIyPIxWBARkSQGCyIiksRgQUREkhgsiIhIkm1VF0BuhYWFOHv2LNzc3GBjY1PVxSEiqhY0Gg2ys7PRoUMHODo6Gmy3umBx9uxZhISEVHUxiIiqpU2bNsHPz88g3eqChZubG4AnF9yoUaMqLg0RUfWQmZmJkJAQ8R36LKsLFrqmp0aNGqFp06ZVXBoiouqlrOZ7dnATEZEkBgsiIpLEYEFERJIYLIiISBKDBRERSWKwICIiSQwWREQkyeq+s6Dqr0SjhUZb/v1tlICdDX8HEcmJwYIsjkYLZD8qNj1/STGK8h+gtOgxtFoNbBSAQqEwYwmJLJuNjQ3q1q0LFxcXODg4yHJMBguq1jQlxcjPuQ0XFxfU83CHra0t7GyUUCoZLKhmEgQBJSUlyM3NxY0bN9C8eXNZAgbr6lStFeU/gIuLCxo0aAA7OzvWKKjGUygUsLe3h6urK+rXr4+cnBxZjstgQdVaadFj1KtXr6qLQWSR6tWrh0ePHslyLAYLqta0Wg1sbdmaSmSMnZ0dNBqNLMdisKBqj01PRMbJ+W+DwYKIiCQxWBARkSQGCyKqVIIgmCUvmReDBVm9wuKSqi7CC5GzvFu3boVarYZarcYrr7wCrfb5n8bv3r1bzB8VFVWuc2ZkZECtViMoKEgv/dq1a5gwYQJu3rypl65Wq+Hl5aWXdufOHbzzzjs4fvx4ucpQHqGhoVCr1Thx4kSlnbM64TASsnqO9nZo2ve9qi6GyTISYsxy3OzsbCQnJ6Nr165l5tm1a5dZzg0AkydPxrVr10zKGxUVhcTERIwaNcps5aEXw2BBVAPUq1cPubm5iI+PLzNYFBQU4MCBA7Czs0NJify1sbJqNTt37jQYtSNVA6LKx2YoohrglVdegYODA/bs2VNmP0BCQgIeP36M3r17V2rZWrdujZdffrlSz0kvjsGCqAZwcnJCQEAAsrKykJKSYjTPzp074eTkhMDAQL10Xb/HvHnzDPbJzMyEWq1Gv379yjx3UlIS1Go1bty4AQB49dVXoVarxe1P91no+juOHj0KABg3bhzUajUyMjLE/KdPn8asWbPQt29fdOjQAT4+Pnj99dexcuVKFBUV6Z1bq9Xi66+/xsiRI+Hn5wcfHx8MGzYMK1euxOPHj59zx/7ef9asWVCr1XjrrbeQn58vuY+1YrAgqiEGDx4MAIiPjzfYlpeXh0OHDqFfv35wdHSU9byurq4YOnQonJycAAD9+/fH0KFDjeZ1cnLC0KFD4ebmBgDo2bOn3r47duxAcHAwdu7cicaNG6Nfv35Qq9W4dOkSYmNjMWvWLL3jLV26FEuWLEFGRgb8/PzQvXt3ZGZmIjY2FpMmTZIcbbVw4ULs2LEDHTt2xJo1a1C7du2K3o5qi30WRDVEYGAgHB0dsWfPHsydO1dv22+//YaioiIMHjxY9l/PrVu3xrJlyxAUFIQbN25g7ty5aNq0qdG8Li4uWLZsGcLDw5GdnY2IiAh069YNAFBcXIzo6GjY2tpi06ZN6Nixo7jfqVOnMHbsWOzduxdZWVlwd3fHrVu38M0336BVq1b4v//7P/FF//DhQ4wePRrHjh3DsWPHxOM/a9myZYiLi0P79u2xdu1a1KlTR9b7Ut2wZkFUQ9SuXRsBAQG4desWTp8+rbdt165dqFu3LgICAqqodNKys7PRu3dvTJgwQS9QAEDnzp3Fpq1bt24BAO7evQsAcHZ21qsRvPTSS4iOjkZMTAyaNWtm9FyrVq3CV199BU9PT6xbt46TVYI1C7JAmgqOhOFnXGUbPHgw9uzZg927d4sv3IcPHyIxMRGvvfYa7O3tq7iEZWvSpAmWLVuml6bRaJCRkYEzZ87g/v37ACCO5Grbti2cnZ2RkpKCkJAQDBkyBAEBAWjWrBn8/f3h7+9v9DxxcXHYvn07ACA2NhbOzs7mu6hqhMGCLI6NUolPvvndpLwje3og616uXloTt5fMUSyr8HRT1Jw5cwAAe/fuRUlJCYYMGVLFpZMmCAISEhKwbds2pKam4ubNm2Jw0A2/1fVD1KpVC7GxsZg5cyZOnDghfmzXqlUrDBgwAGPGjEGjRo0MzrF9+3bY2tqitLQUX375JZYsWVJJV2fZ2AxFVIM4OTmhT58+SE9Px7lz5wA8aYJydnZGz549X/h4ck1/beq5pkyZgilTpmDfvn2oX78+RowYgfnz52Pbtm1G+x569OiBffv24dNPP8WwYcPQsGFDXL16FatXr8bgwYMNmuMAoFmzZtiyZQucnZ2xdetWcWRWTcdgQVTDDBo0CACwZ88e3L9/H3/88QcGDhxY5rogSuWT14SxwCDXwjqm+OWXX5CQkIAOHTpg//79iIuLw6JFizB27Fi0a9euzLLUqlULQ4YMwUcffYRDhw7h559/RmBgIAoKCvDvf//bIH90dDTatWuH2bNnAwAWLFiAwsJCs15bdcBgQVTD9O3bF7Vq1UJ8fDx+//13lJaWisNqjdENW83OzjbYdurUKZPP+yJrKxjL++effwIARo0aBVdXV71t2dnZuHTpEoC/v/7etWsXgoKCsGrVKr28np6eeOeddwAAt2/fNjiPnZ0dAODNN9+Er68vrl+/jhUrVphcdmvFYEFUw9SqVQsBAQG4evUq1q5dC1dX1zKHjwKASqUC8OTjutTUVDH9ypUr+OKLL0w+r4ODA4An33SYmvfp2oKHhwcAYP/+/Xq1nKysLERGRop9F8XFxQCeDNm9ceMGNmzYgOvXr+sdf8eOHQAAb2/vMsugUCjwr3/9C7a2tli/fj0uXrwoWW5rxmBBVAPpahJXrlzBoEGDxKYmY1q2bIm+ffuipKQEo0aNQkREBN5++228/vrrUKlUqF+/vknnbNGiBQAgMjISkZGRzw0aurwLFy5EZGQkrl+/jjfeeAPOzs5ISEjAwIEDERkZiXHjxuHVV1/Fn3/+iVatWgH4uwakUqkQHh6Oe/fu4T/+4z8QFhaGyMhIDBkyBKtWrYKrqyumT5/+3DKr1WqMGzcOpaWlmDdvXqX20VgaBguiGigwMFBsXjJlFNSnn36KiIgIuLm54fDhw7h+/ToiIiLwxRdfwMbGxqRzzpkzB126dEFmZib++OMPvSk8njV58mQEBgbi0aNHSExMxLVr1+Du7o7vvvsOQUFBKCwsxIEDB5CZmYm+ffviu+++E7/eTkhIEI/z7rvvYsGCBWjXrh1Onz6Nffv2oaioCGPHjsW2bdvK/DjwadOnT4eHhwfOnj2LDRs2mHSt1kghWNnqIhkZGXj11Vfx+++/m/QgkOUpLNFi7grTpsoe2dMDzVvqT0LXxO0lKJV/t3kXFpfA0d5O1jKaU3UrL1m2CxcuoF27dpL5pN6drFmQ1atuL97qVl6qGRgsiIhIEoMFERFJYrAgIiJJDBZERCSJwYKIiCQxWBARkSQGCyIiksRgQUREkhgsiIhIEoMFERFJYrAgIiJJDBZERCSJwYKIqBqpqonCjS+6S2RFSjRaaLRVXQrT2SgBOxt5fsep1WoAMOuU/V5eXtBoNHqr6Fm7qrrmX3/9Ffv27cPHH39cqecFGCzIiIq+XOV82clBowWyHxVXdTFM5lbXHnamrSdENcjJkycxc+ZM+Pv7V8n5GSzIQEVfrtX9ZSdAACpS01cACiik8xG9AK22aqvHDBZEzxIAbQWChRIAYwVZG8tpKyCiSrF161ao1Wps3LgRJ06cQHh4OLp06QIfHx+Eh4fj+PHjRvc7ePAgQkND4efnh27dumHu3Lm4d+9emed5+PAh/vd//xcDBgyAt7c3unfvjsjISFy8eNEgb2hoKNRqNXJycvDvf/8bffr0QadOnTB06FBs2rSpzF/V27dvx5gxY+Dr64vOnTtj5MiR2Lx5s0EnsCVf86NHj7BmzRoMHjwY3t7e6N27NxYuXIicnBwxb1RUFEJCQgAAx44dg1qtRlRUVJnlMAcGC6Ia6vDhwxg3bhwyMjLQo0cPeHh44OjRoxg/fjzOnz+vl/f777/HpEmTkJycjA4dOsDHxwe7d+/G2LFjjY7OuXXrFt58802sWbMGpaWlCAgIQMuWLbFnzx6MGjUKCQkJRsv03nvv4YsvvkDTpk3Rs2dPZGRkYNGiRZg9e7ZB3nnz5uGdd97BhQsXxBfz1atX8f7772P27NlGy2WJ1xwVFYWPP/4Yzs7OCAgIQEFBAb777jtMnDhRzOPj44NXXnkFANCgQQMMHToUPj4+Ro9nLmyGIqqhEhISEBERgcjISNjY2EAQBLz77rv4+eefsXHjRsTExAAAbt++jaVLl8LR0RHr1q2Dr68vACAzMxNhYWFGf/W/8847SE9Pxz/+8Q/MmDEDNjZPOrESExMxZcoUzJ49G/Hx8WjQoIHefgcPHsSKFSsQFBQEALh58yZCQ0OxY8cOBAUFYdCgQQCAzZs3Y8uWLWjXrh1WrlwJDw8PAEBOTg4iIiKwfft2dO3aFcHBwRZ/zUlJSYiLi0OnTp0AAFlZWRg+fDjOnTuHEydOwM/PD8HBwWjdujUOHz6M1q1bY9myZSb9fywn1iyIaigPDw/885//FF9qCoUCY8aMAQCcOXNGzPfTTz+hsLAQ48aNE1+aANCoUSPMmzfP4LinTp1CcnIy2rdvj1mzZonHB4BevXoLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.hist(reshaped_mt_emp.numpy(), bins=bins, label = \"Multitask\")\n",
|
||||
"plt.hist(reshaped_ind_emp.numpy(), bins=bins, alpha = 0.5, label = \"Independent\")\n",
|
||||
"plt.legend()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 44,
|
||||
"id": "ee85c64f",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAakAAAEgCAYAAAAOk4xLAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAABZgklEQVR4nO3deXxU1dnA8d/syWQhKyGBBEggYRME2VRWQXEBQcB9F/f1da1Vq2+12NJqa2mt1oW2CPqySAVFBAFlEQRFFGQVQiAhZN9nn7n3/SNkYJJJSGCSTJLn+/n4aeecO3fODDfzzDn3POdoVFVVEUIIIYKQtrUbIIQQQtRHgpQQQoigJUFKCCFE0JIgJYQQImhJkBJCCBG0JEg1gtvtJicnB7fb3dpNEUKIDkWCVCPk5eUxYcIE8vLyWrspQgjRoUiQEkIIEbTaRJBatmwZGRkZfP/99016Xn5+Pi+++CITJkxg4MCBTJo0iTfffBOn09lMLRVCCBFIQR+kdu7cySuvvNLk5+Xl5XHdddexaNEiIiMjGTduHBaLhblz5zJr1ixcLlcztFYIIUQgBXWQWrNmDbNmzcJqtTb5uf/7v/9LXl4ejz32GP/973+ZO3cua9as4aKLLmL79u188MEHzdBiIYQQgRSUQSovL49nnnmGRx55BEVRiIuLa9LzMzMz+frrr0lJSeH+++/3lpvNZmbPno1Op2PBggWBbrYQQogAC8og9cYbb7B8+XIGDBjAokWLSE1NbdLzN2/ejKqqjB8/Hq3W9y0mJSXRr18/jh8/zqFDhwLZbCGEEAGmb+0G+JOamsqcOXO4+uqr6wSZxqgJPr179673/Lt37+bgwYP06tXrnNoqhBCBpHgUqvJKqMgtQnHVys1UVXA7wVNzT/3kY3sVGnsluBzN0qboi8YQ3rNns5z7TIIySN17773n9PyCggIAOnfu7Lc+Pj4egKKionN6HSGEOBuqqlKVX0rBnkyKDmRTnlNAeXYBFTkF2PJOkGjKo4u5hBCdE6POhUHrxqhzYdS60Glbfnel8g8gb+JT9HrqhRZ/7aAMUufKZrMBEBIS4re+pvxsJmQIIUR9FLeH4kPHKdh7hII9R6g4Xoi93IKjwoLLavce56i0QEURMaYKIo0WjDoXCTo3aQYbCd1L0GqCa5s/rQY86/4NEqQCo2aIUKPR+K2v2edR9nsUQpwrVVXJ353Jz0u+Yv+n3+CsqCTSaCU6pDoAdTbYCDPYCNGdys80RLsJiQ/uNBh3pwSKb/w9MR+/jKEwC3eY/5Gp5tYug5TZbAbAbrf7rXc4qsdtQ0NDW6xNQoi2ryK3iOytP3P8+/1U5ZWgKc/FXPYLBkseEQYb46NtRCZY0GuVZmuDghZFc+qrW0WHS2PCozHh0RgA/z/Om/w6Sjj2lEHk3fwG5sVz6fHC6wE5b1O1yyBVcy+qvntOhYWFPscJIToej8tN/u5Mjm39mcK9WdjLqrBXWnBW2jCoNiK1JURoytBoqgOO6vagWCsxat1E6pykhpYRZnCAier/AkjXtRcho67G2GcY2vAoNBFRaMM6oQ2PAlNovaNEgVZscXGwUxzxf/k/wiONLfKatbXLIFUzq6++KeaHDx8GID09vcXaJIQILJfNga2kEnt5Fc6SYjy26pETVQWPy4mjwoqj0oLHakXrqv7PU1WJpbAMS1EZVTknCPGUE26w0dlgx3hyckKIzonZ4GeWnAEI5OCLyYwhdQCG1PPQxiaiDa8OQvrU89CnZLRYIGpIbJiBGLOeMpuLBAlSgTN69GgA1q9fz1NPPeUzjT03N5d9+/bRtWtXmX4uRAtzO1y47f6nSStuBUeFBXtF9UQDR4XFO+nAWZiPtuAX9CVH0ZTnorcVEapWYdI5Merc6DQqulrnM5+pMU1bI6BJtNGdMaQNRN9zAPqkVHQJ3dHFJYLu5FeuTocurhsaXe1Wt7wKmxu3ohITZvBb36uzGW0rxss2H6Ryc3Ox2WxER0cTExMDQHJyMqNHj2bTpk389a9/5fHHHweqZ/O98MILeDwe7rzzztZsthBtkqqqqNZKUDw1BSi2StSqchRLOZzcc83tcFL0Sw6F+7Io3H8Ue1EJGlslOsVOqN5JuMFGmMFKiK7+yQNaqgNNuFbBrK8V2FrnR30dqt6INnUIpoEXYeqZgT6hO7rEnuhiElq7aWfkVlSOFdvJr3Si12qICNFh0NXNS9W1ZoSiHQSpX/3qV2zfvp2HH36YRx55xFv+0ksvceONN/L222+zfv16evbsyQ8//EBhYSFjxozhxhtvbMVWC9E2qE47tq+X4ti+BveJI3jyj6JaKhr1XCPQ9eR/hJ/8r63QG9Ak9Uab3A9NWJS32BgTgz4qBk14FPqEFAzpF6AxBviGVAsosbjILLLh8lTPcHYrKkeK7KQnnLH/2eLafJCqT3JyMkuWLGHu3Lls3LiRo0ePkpyczG233cbtt9+OXt9u37oQ50ypKMHy2XtYP30XpaywtZtzRh50qJpaf9MaDRqNBlWrQ9GHoujNqIYQdCYjepMRfZgZU7ee6Lt0R9c5GW2nODThUWgjotDFd0NjCJLuWgA53QpHim2UWOruMl5scWF1ejAbW38I8nQaVZKFzignJ4cJEyawbt06unXr1trNEaJZ2TYso/zvj6NWlbd2U7xUjQ53VDJKQm/0yRmYe/UhLKM/+vjE6tlvRv+J+6KaqqoUVrnIKrbjUep+5YcYtKTFhRIZGnw/3oOvRUKIVqFYKih78ykcXy1u8DiPosWtnrp34VZ0uBQDTo8ej7dcgyHEiKlTGCFR4ZiiIjFEx6HrFIM2IhpdQgq6Lt3RxSadmkxQHw3oYpPaZc+mJdhdCplFNsptdXtPAF2jTHSLMqFt5XtP9ZEgJUQHpXrcOLZ9geOnTVh//g7l6F60nroJ8C6PjkPl3ciu6kyV04xDNWGKCPPWhyfE0Ll/Tzr360FsegpRKZ2JSIxDZ5Svl9akqionyp1kl9rx03kizKgjLT6UMFNwDe/VJleREB2AqiiolaXV/9/jwr7xv1Qs/TsUH/ceU3tel0fR8HNxGgfLkjEldCH9hhGkXHQe3Yb1xRQZfDfYxSkWh4fDRTYsDk+dOq0GkqNDSOxkDIpcrDORICVEO6a6XVR89Beqlr2J1t74e0xljjC25g8ibuwEplw/ge6jBqL1Mz1ZBB+3orIntwqPn95Tp1AdqXGhhBiCu/d0OglSQrRTZd9uoPT1hwipymn07qYOj4Ej1h5oJ93PzLumEtm1GTNeRbPQazV0iw7haMmpoVudFnrEhhIfbmgTvafTSZASop2xHj5I1suP0Sn/W0LOsOWDR9FypCKRPEcCYQOHk3LlZVx0+UgZzmvjEjsZKbK4sDg8xIYZ6BEbglHfNnvCEqSEaCfcxfkcfeUxTPvXEK1R6iyG7VG0uJTqYR6Hx8Cxyi5U9ZxAxr1XMWLSCIzhsitAW+P2qOh1dXtGGo2GtLhQHG6l3uWO2goJUkK0caqqYlm9gJK5T2NW7X53ajhmT8Z+8T3ooqp3pTZ1CmPYxGFE90xs4daKQHC6FbKK7VicHgZ1Dfc7fTzMpAv6mXuNIUFKiDbMdXQfFf98DufOr/D3e9mihOO65GGGPPo4htC2t3yP8OUvKTenzEFKTPtNZpYgJUQboDpsOPdtR3VW3wz3FB7HtvYjXPu/83u8xR2Kfcj1ZLw4G0NYmN9jRNtid3lOJuX6TivPLXMQG2ZoF70mfyRICRHEVFXFtvZDKt55HrWq7IzHuxUthxnMkPc/IDIlqfkbKJrdmZJyzUYdbWzCXpNIkBIiSLlPZFH6l0dw797UqOPzLNHsUS5m6pK/EpEY28ytEy3B4vBwuNCGxdn2k3LPlgQpIYKM6nFTMv917EteQ6vWv99SjTxLDL+UdaPQ0IsbFr8iAaodUBSVnDIHx8v8bxBZnZRrJsTQNqeVN4UEKSFakeq048rah1Ka731c+PYrGEoO10nAdStaCqwxqICKhlJ7BJkVSVhcZjold2bm358gNq1ri78HEVgVNjeHi2zYXUqduraclHu2JEgJ0cI8hTlULfkrth82oOQeQqP6fhn5m6WXZ4lhe34/qlynkmxDYyPpNWMYvSeNIOXCAbKgaxunqipHiu3kVzj91seE6ekZG9pmk3LPllzVQrQQVVGwrnyfivdfAocV8JvS5MPh0XNAGUbolTczZdJwIhKrlynSaDWExXVCo+1YX1jtWX09I4NOQ2pcaJtPyj1bEqSEaAaq24Vj2xfYt65EObn6uKcwB/eRPY0+R7Y1CdNNL3LZPddJMOogUmJCKLW4cJ5cHTYhwkhKbAj6IN3rqSVIkBIigDzFJ7AsexPb+kWN3na9yhVChTMMVa3+InIqRpy9RjPk1dmywGsHo9dW95qySuxBu1NuS5NPQIgAsW9fQ9mf7m1UPpNL0fFTYW8caWNIm3IJqaMHYTBXrwhhDAuVdfTaMbtLIb/SSUq0ye8QX3SYgU5mPdoOMjHiTCRICdFEqqqi2i2olgpQq+faWVf9m6qPXmvU849XxfFdfj/SZl7F5Nn3dphZWh1d7aTcEL2WhEij32MlQJ0iQUqIRlDKi6n8zyvYv12FUlEMHvcZn+P06Mmq6MIJa5x3KK/SaabSFUbvScOZ+Mo9EqA6CH875R4tsRFt1ne42XpNJUFKiDNw7FhH2esPenOZzkRRYVdRLw6Udsej1l1PLfnC/lz5xqOy020HUJOUm1vmoPaKRh4FCiqddItuv4vDBoIEKSH8UBUF14Hvsa5ZiO2L/zT6eTa3kW9yB1Ls6Uxc/xQ69+9BeEIMNYurRXaNp+/VoySnqQNobFKuaJj8pYgOz3XkZ6xfzMeTm1ldoKq4svaiFJ+o9zluRYtL0XuH8RQ05Fli+bksnT43TeXaR2cS0im8JZovgoxbUTlWbCe/0n9SblvfKbelSZASHZLqdOD4cQOWT97CufOrRj1HUWFLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(bins[1:], np.cumsum(mt_counts, 0) / reshaped_mt_emp.shape[0], color = palette[4], label = \"MT-Vol\")\n",
|
||||
"plt.plot(bins[1:], np.cumsum(ind_counts, 0) / reshaped_mt_emp.shape[0], color = palette[-2], label = \"Vol\")\n",
|
||||
"plt.plot(bins[1:], bins[1:], color = palette[1], linestyle=\"--\")\n",
|
||||
"plt.xlabel(\"Probability\")\n",
|
||||
"plt.ylabel(\"CDF\")\n",
|
||||
"plt.legend()\n",
|
||||
"sns.despine()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 34,
|
||||
"id": "f22fd162",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def prep_ecdf(vec):\n",
|
||||
" vec = vec.reshape(-1)\n",
|
||||
" counts, _ = np.histogram(vec, bins=bins)\n",
|
||||
" return np.cumsum(counts, 0) / np.sum(counts)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 38,
|
||||
"id": "c7abc329",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"0\n",
|
||||
"1\n",
|
||||
"2\n",
|
||||
"3\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAABzcAAAE7CAYAAABQVQwKAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAADPBUlEQVR4nOz9eXgkaXUn+n8jIrdIZaYypUztqtLWGzTVCxiDcTc0m8fY2AZ7PBiugTZtBo8Bm+sZeox9wcbuAbzMGAbG/Ztu+OELZgZjsMdmgIG2zdpAQ3dRS68llVRV2lNSppRL5BIR7/1DVZIiIlNKSbnn9/M8PA+R61tQpaM3znvOkYQQAkRERERERERERERERERETU5u9AKIiIiIiIiIiIiIiIiIiCrB5CYRERERERERERERERERtQQmN4mIiIiIiIiIiIiIiIioJTC5SUREREREREREREREREQtgclNIiIiIiIiIiIiIiIiImoJrkYv4KhyuRzOnz+PWCwGRVEavRwiImoQwzAQj8dx8803w+fzNXo5LYNxlIiIAMbRo2IcJSIigHH0qBhHiYgIOF4cbdnk5vnz5/GGN7yh0csgIqIm8dd//dd43vOe1+hltAzGUSIi2otx9HAYR4mIaC/G0cNhHCUior2OEkdbNrkZi8UAbP+hBwYGGrwaIiJqlOXlZbzhDW/YiQtUGcZRIiICGEePinGUiIgAxtGjYhwlIiLgeHG0ZZOb11oWDAwMYGRkpMGrISKiRmMrm8NhHCUior0YRw+HcZSIiPZiHD0cxlEiItrrKHFUrsE6iIiIiIiIiIiIiIiIiIiqjslNIiIiIiIiIiIiIiIiImoJR05ufuELX8ANN9yAH/7wh4d638rKCt773vfiZS97GU6dOoWf+qmfwsc+9jEUCoWjLoWIiKilMIYSEREdHeMoERHR0TGOEhFROzhScvP06dP4oz/6o0O/b3l5Gb/8y7+Mz372swiFQnjJS16CTCaDj3zkI3jLW96CYrF4lOUQERG1DMZQIiKio2McJSIiOjrGUSIiaheHTm5+9atfxVve8hZks9lDf9kf/MEfYHl5Gb/1W7+Fv/u7v8NHPvIRfPWrX8VP/MRP4JFHHsGnPvWpQ38mERFRq2AMJSIiOjrGUSIioqNjHCUionZScXJzeXkZ7373u/GOd7wDpmkiGo0e6osuXryIr3/96zhx4gTe9ra37Tzu9/tx3333QVEUfPrTnz7UZxIREbUCxlAiIqKjYxwlIiI6OsZRIiJqR65KX/gXf/EX+F//63/h5ptvxn/6T/8Jf/zHf4y1tbWKv+jb3/42hBC46667IMvWnOrQ0BCe9axn4dy5c5iensbU1FTlfwKiNmSsL0HkDn+Srlp0LY/M+mbDvp86S/jW26F4vY1eRk0xhhKVZqYSMLc2avodjGnUzmSPB9233OaIDe2GcZQ6lRACxuplQNfr9p2GriO9kgCEqNt3EjWKuzuM0I03NXoZNcc4SkTHVdTySC/Xdu++HyEEsLEIYbAFdjPxDQ6j68SJhn1/xcnNiYkJfOhDH8LP/dzPHWnzPD09DQC47rrryn7+uXPn8MwzzzAQUscShTw2/uDfoHD6641eClHdXDE9cP36RzDyS69r9FJqhjGUyGnrv78Hmb/7b41eBlHLWzeDiP3nL6P72Tc3eik1wzhKnchYW8T6e34BxpVnGr0UorZVALDgGsHUZx6GOxhq9HJqhnGUiI7jqX/8Dr7yHz4Go1C/w1Z7eZU8Xjr6KCLedEO+n8rLCWAu9uO46a++3JADtxUnN9/61rce64tWV1cBAH19fSWfj8ViAHCok0NE7Sb3yFeY2KSO45ULSP7VfUAbJzcZQ4ms9JVLTGwSVUmXnML8f3kvuh/8QqOXUjOMo9SJsl/6BBObRHUQ1udx6f4PY+o//D+NXkrNMI4S0XF884OfblhiEwAmuxeY2GxSsgT0rH0fG9/+FqJ3vrj+31+vL9I0DQDg8/lKPn/t8aMMtSZqF/rFc41eAlHd6OHBnf9u+tr3lGw1MIZSu9Evnm/0Eoha3t44Knf3NnAlzY9xlFpRcYZ7Q6JaEZIEPTywc+3uG9zn1cQ4StS5tGQaqaX1hq4h4k019PvJyXT7YHRFAACGkOHqiTRkHRVXbh7XtbJUSZJKPi+uznMQnOtAHUxfuGi5liN9kNRA/b4/V7D2T5cAWVHq9v3UGYSsIPtT9yD7M29D6C/fCczPY/SPH2j0spoaYyi1G3u8kwLdkEPVTc4wplG7Mn0BpP/1vcg/96cQ+U//GoYrjBve/+FGL6upMY5SK9IXrbFS6RsBXJ6afZ8wBTYvr1gek12Mm9R+jJ5BpN54H4yeQYQ/8G9gTv0kbnrTPY1eVlNjHCXqXMm5Jcu14nEjOFjfg5Vhzw8s15rwwxT8HaVRiidvRuruD0Bem0fX/f8enlf+GkZuPtWQtdQtuen3+wEAuVyu5PP5fB4AoKpqvZZE1HT0xRnLdfg/fgLeUz9Zt+8/+9l/wtd+9/+3c33Dz/4EfvYjv12376f2pxUMTMc1ZPPG9vVv3Y9bRgJwK/Xvy95KGEOp3Ri2eBd43b9H4BffUdXvYEyjdrSp6ZiOZ1HQt28eFj/4Fdw81FX2ZiNtYxylViMMHcbynOWx6H97GHJX7bqdrD45h3/8mXfvXHef6Mc9//Rfa/Z9RI2wmipgbk2DcTUH537gh5iI8Wf/QRhHiTpXwpbcHH/Jrfj5+/9D3b5fCIGV134eYs+Pn5OfOQ2lp79ua6BtQghcSeQRT27/zDf6x9HzP86iP1S7w3cHqdvd5Gt92cv1X4/H45bXEXUaIQSMBevNXtdIfQexJ2atASsyPlTX76f2ZwogczWxCQBFQ2B2rfQGiXYxhlK70RemLdeu4erHO2dMY7sxan0F3dxJbAJAOm9g8ermkspjHKVWY6xcAfTizrUc6atpYhNg3KTOkC0YO4lNAFhJFZDIFsu/gQAwjhJ1MsfvB2P1/f3A3FiGyGV2riU1CDnCnzWNks5bZ6/OrWvIFY0yr669uiU3r7vuOgDA9PR0yednZraTOtdff329lkTUVMzEKoS2OxxZUgOQI/U9hWJvNRAZHyjzSqKj6fIqGIl4d67dioRowN3AFbUGxlBqN/a2tK7hyap/hyOm1XkTRlQL0YAbPV27zXdUt4xuP+PoQRhHqdXYO/ooQ9WPk3ZJx81L7gWp/ZyI+KC6d2+Fdqsu+D1sbXgQxlGizmVPbobrfPhJtxUCKcMT7FrTIJIkYTLmh7Lnf/5owNPQbnx1++Y77rgDAPDP//zPME3T8tzi4iKefPJJDA8PY2qqvpVqRM3C3qJPGRyv+w/rxNyy5Zo3gqkWhsNeBLwKevwu3DISQE8Xb8oehDGU2omZy8BcX9x9QJahDIxV/XvsMS3Mm7TUBiRJwkRUhVuRMNjtwanhAAJe3pQ9COMotRpHR5+hiZp/pzNuci9I7UeWJUz1+aHIwHivDzcN+OF1cUTKQRhHiTpXssH3ig3bDHJXHQ58UXlel4zxq/vRGwf8mIypUOTGJZtrEsEXFxcxMzODjY2NncdGR0dxxx13YHZ2Fh/+8Id3Hs9ms/j93/99GIaBu+++uxbLIWoJzhZ99f1hLUzTEbC4oaWjKhpm2bYEkiThpsEuXN/v56zNEhhDqd3ZNydK3ygkd3VnNJSKaTywQ60kWzCg7+2bt4dbkXHraBBjvSrkBm4kmxXjKLWDRuwN7TO1GDeplaVyOoQoHUcDXgW3nwhhoNvL6p8SGEeJ6BohRInfD+p7aNj5O1HtD3x1OiEEUjm97PPRgBu3jgYRaYIOQjW5q3zvvffiVa96Ff76r//a8vj73vc+xGIx3H///Xj1q1+Nd77znXjlK1+J73znO7jzzjvxK7/yK7VYDlFLcLboq++Jt9TSOozC7qwJXyQINRyo6xqoPWxkivjRlTSeWcnCLLOhdMkSN5JlMIZSu6tHvHPEtHAAaiRY9e8hqjYhBBaTeZxdSGNuXSv7OheTmmUxjlI7cBwEakRykzM3qQUZpsDsmobzixksbxXKvo5xtDzGUSK6Jru2iUJ6d0/i9nvR1Rep6xqcbWlZIV5LBd3EUytZnF/MYFMrneCUJKlp4mhdS2ZGR0fxuc99Dq997WuxsbGBr3/96+ju7sbv/M7v4KMf/ShcLtfBH0LUphxtaevQemgv54Botu+jw9FNgenVLJ5eyUI3BTIFEwuJfKOX1TYYQ6ld1CPeOWMab9BS88sVTTyxlMGljRyEAOLpItYzxYPfSBVhHKVWYr+RV+vKzfxWFtr61s617FYQGorW9DuJqi2V03F2Ib2T1Ly0kUO2ULqbEB0e4yhR5ynVsr7ehQpsS1s/6+kizsynkcxuJzWn49v3d5vZkSPPpz71qSM9Nzg4iA984ANH/VqitlXvDawd523ScWxqOqZXsyjYWujNJ/Po6XKji/PALBhDqZPVI97Zk5uct0nNTAiBeKqIuXUN9k60F+Maun0uuJTmOBnbLBhHqZ2JYgHG6mXLY67B2h58tVdtdo/2Q3bx93dqDaYQmE/ksZC0HqwVApiJa7h5qItdg2wYR4moEonZRct1vQthhGlCtyc363y/vBPohsDsuoa1tPVgbUEXuLyew0RMbdDKDsZjNURNoPQP6/qW2TsCFtsQUQUMU+DyRq5kyx8JwEjEC7+HczWJaFc92sqwtR61ioJu4uKahkTW2fJHkYGxqA8cT03UWYzlOcA0d67l6DAkn7+m3+noeMC4SS0iUzAwvZpFtmA6nvO6ZJzs9TGxSUR0REl7Icz4UF2/34jPA8XdgytSMAI51FPXNbS7ZLaI6biGov2ULYCgV8FQ2NOAVVWOyU2iJmCuLwKF3M61FOiGVOcf1vaAFWblJh0gldMxHdeQKzo3kqpbxlSfHwFWbBKRjb0tbS3ayjg2YYxp1ITW00VcXNNKtvrpVl2YjKnwupjZJOo0zg4HtR9Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 2304x360 with 4 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(1, 4, figsize = (32, 5))\n",
|
||||
"\n",
|
||||
"times = [0, 50, 100, -1]\n",
|
||||
"for i, time in enumerate(times):\n",
|
||||
" print(i)\n",
|
||||
" ax[i].plot(bins[1:], prep_ecdf(mt_emp_percentile[:, time]), color = palette[4], label = \"MT-Vol\")\n",
|
||||
" ax[i].plot(bins[1:], prep_ecdf(ind_emp_percentile[:, time]), color = palette[-2], label = \"Vol\")\n",
|
||||
" ax[i].plot(bins[1:], bins[1:], color = palette[1], linestyle=\"--\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 36,
|
||||
"id": "2b464404",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"[<matplotlib.lines.Line2D at 0x7fb6f6e5a7d0>]"
|
||||
]
|
||||
},
|
||||
"execution_count": 36,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAY4AAAEFCAYAAAD0cwBnAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAAj0klEQVR4nO3de3gTZd438G+aNimhLfRED9BCSxuV4wIeiorAI+ij14I+LAKCCNXaF1xAvXYV8VVBuqjs6+5ixV2UurhS1IoXyKUsi1YWQfAqLnIq1EJ6gEJpKT1hm7Zpknn/4Ek0pyZpk0wy+X7+0Tl05hcY8u3c99z3yARBEEBEROSiELELICKiwMLgICIitzA4iIjILQwOIiJyC4ODiIjcEip2Ab3V2dmJ0tJSxMfHQy6Xi10OEVFAMBgMaGhowKhRoxAeHt6rYwRscJSWlmLBggVil0FEFJC2bduGm2++uVc/G7DBER8fD+D6h09MTBS5GiKiwFBXV4cFCxaYv0N7o9fBsWPHDqxatcrt1Kqvr8fbb7+NQ4cOoaGhAUlJSZg5cyaeeOIJKBQKl49jap5KTEzEkCFD3K6fiCiY9aWJv1ed48eOHUNeXp7bP1dXV4c5c+agqKgIUVFRmDJlCtrb25Gfn4/HH38c3d3dvSmHiIh8yO3g+PLLL/H4449Dq9W6fbI1a9agrq4OTz31FHbu3In8/Hx8+eWXuP3223HkyBFs3brV7WMSEZFvuRwcdXV1eO6557B8+XIYjUbExcW5daLKykrs378fqampWLJkiXm9SqXCunXrIJfLUVhY6NYxiYjI91wOjg0bNmDXrl0YNWoUioqKkJ6e7taJvv32WwiCgKlTpyIkxPK0ycnJGDFiBC5dugSNRuPWcYmIyLdc7hxPT0/H+vXrMXPmTJsvfleYAiEzM9Ph8U+dOoWzZ88iIyPD7eMTSVFrWwfOaC7DaOQk1mTpxvQExA6MEOXcLgdHbm5un0505coVAMCgQYPsbjc9Gnb16tU+nYdIKs5UXMZDT29Ga1un2KWQHwoJkeHNVQ/hf6b9yvfn9tWJOjo6AMDhSEXT+t50uhNJ0Zad3zE0yCGjUcA7n3wryrl9Fhym5i2ZTGZ3u+l9UnyvFNF1ZRV1YpdAfm5Ycowo5/XZyHGVSgXg+hxT9nR1dQEA+vXr56uSiPyWIAjQXGiwWHfbmGGQyzkvKV03KiMZv51/lyjn9llwmPo2HPVhNDQ0WOxHFMzqG39Cm7bLvNy/nwKfbnjC4R07kS/57NcX09NUjh63raioAACo1WpflUTktyqs7jYyUuMZGuQ3fBYckyZNAgDs27cPRqPRYlttbS3KysowePBgPopLBNg0Uw1P7f2EdESe5pXgqK2tRUVFBZqamszrUlJSMGnSJFRVVeHNN980r9dqtXjxxRdhMBiQnZ3tjXKIAs65C1csljNT2YRL/sMrwbFy5Urcf//92LZtm8X61atXIz4+Hps2bcKMGTOwYsUK3HPPPTh06BDuuusuPPzww94ohyjgVFyw7AscnureFD9E3uTTRzRSUlKwfft2zJo1C01NTdi/fz8GDBiA3/3ud9i4cSNCQwP29SBEHmXdVJXBpiryI73+pu5pJtuetiUlJeG1117r7WmJJK9N24XLDa3mZXlICIYmx4pYEZElPhRO5GcqayybqVKTo6FU8G6c/AeDg8jPsGOc/B2Dg8jPWI/h4KO45G8YHER+hh3j5O8YHER+xt6ocSJ/wuAg8iN6gwFVlxot1rGpivwNg4PIj1y43Axdt8G8HB8dgYGRnDGa/AuDg8iPsGOcAgEfDifykY5OHXbtO4m6q9cc7nOsrMZimf0b5I8YHEQ+krvmQ/y75KxbP8PgIH/EpioiH7ja3OZ2aAAMDvJPDA4iHzh3/orznawkDxqArF+leaEaor5hUxWRD1h3eo8YnoTpt9/ocP8BEf0w87/GIFwR5u3SiNzG4CDyAY3VxIX3TRqBZxbdLVI1RH3DpioiH7BuqmLfBQUyBgeRD9hMIzKUM95S4GJwEHlZR6cOF+tbzMsymQxpQ/hiJgpcDA4iL6u8aPVipqRodnpTQGNwEHkZp0knqWFwEHnZufMMDpIWBgeRl3HiQpIaBgeRl7GpiqSGwUHkRQaD0aZzPDOVj+JSYGNwEHnRpSst6NLpzcsxA1SIHqASsSKivmNwEHkRO8ZJihgcRF7EjnGSIgYHkRdZd4yzf4OkgMFB5EUVNbzjIOlhcBB5EWfFJSlicBB5SXOrFk2tWvOyMiwUQxIGilcQkYcwOIi8xLp/Iy0lDnI5/8lR4ONVTOQlth3jbKYiaWBwEHmJhh3jJFEMDiIv0bBjnCSKwUHkJTZNVUMZHCQNDA4iL+jUdaOmrtliXfqQOJGqIfIsBgeRF1RfaoTRKJiXhyQMRL9whYgVEXkOg4PIC6ybqdgxTlLC4CDyAs156/4NzlFF0sHgIPICm7f+pbB/g6SDwUHkBZxOnaSMwUHkYUaj0WbwH8dwkJQwOIg87HLDNXR0dpuXB0SEIy46QsSKiDyLwUHkYTb9G0MHQSaTiVQNkecxOIg8zLZjnM1UJC0MDiIPY8c4SR2Dg8jD2DFOUsfgIPIw68F/GZzckCQmVOwCiPxB6bla/P7/7UBlzdU+H0vbqTP/f1ioHKlJ0X0+JpE/YXAQAVj1l10oPVfr8eOmDY5FqFzu8eMSiYlNVRT0unR6nCy/5JVjj71xiFeOSyQmBgcFvepLjTAYjR4/7o3piXhq4VSPH5dIbGyqoqBnPe7ijnHp+Pu6hX06ZohMxvdvkGQxOCjoVVg9PqseloD+/ZQiVUPk/9hURUHPZqQ3x10Q9YjBQUHPeqQ3g4OoZwwOCmqCIPA1r0RuYnBQULt89RraO34esBfZX4mE2EgRKyLyfwwOCmo2ExKmxHMKdCInGBwU1NgxTuQ+BgcFNc35KxbLDA4i5xgcFNTsva2PiHrG4KCgprGaDTcjNU6kSogCB4ODgtZP7Z2ov3rNvBwqD8HQ5FgRKyIKDAwOCloVVncbQ5NjEBbKKdCJnGFwUNBixzhR7zA4KGixY5yodzg7LkmG0WhE1cVGdOq6Xdr/hNXLmzJSeMdB5AoGB0lCm7YLc54pwMmzvX+TH+eoInINm6pIEr7Yf6pPoQEAw/koLpFLGBwkCaWa2j79/KjMZAyI6OehaoikjU1VJAkVF2wfrXX1LX4piQPxf//Pfd4oi0iSGBwkCdZPSG1ZtxDqYQkiVUMkbWyqooDXpu3C5YZW87I8hCPAibyJwUEBr6LG8m4jNTkaSgVvpom8hcFBAc+6mSozlQP5iLyJwUEBz+YtfhyPQeRVDA4KeHyLH5FvMTgo4FnfcTA4iLyLwUEBTW8woPJio8U6NlUReReDgwLahcvN6NYbzMvx0REYGMkR4ETexOCggMaOcSLfc+th98OHD2PTpk0oLy9Hd3c3Ro4cidzcXEyaNMmln9fr9Rg3bhx0Op3d7QkJCThw4IA7JVGQY8c4ke+5HBw7duzAqlWroFAokJWVBaPRiJKSEuTk5GDt2rWYO3eu02NoNBrodDqkpqZi7NixNtsHDhzoVvFEmvMMDiJfcyk4rly5gtWrVyMyMhIffvgh1Go1AODkyZPIzs7GunXrMGXKFCQk9Dw3UFlZGQBg1qxZWLp0aR9LJwI0NQwOIl9zqY+jsLAQOp0OixcvNocGAIwZMwY5OTno6upCUVGR0+OcOXMGADBy5Mhelkv0M0EQ2FRFJAKXguPgwYMAgGnTptlsmz59OgC41DdhuuNgcJAnNLa0o/WnDvNyuDIMyYMGiFgRUXBw2lQlCAI0Gg1CQkKQnp5us33YsGEICQmBRqOBIAiQyWQOj1NWVob4+Hjs27cPRUVFqKiogFKpxO23345ly5bZPT6RI9Z3G8NT4hASwgcFibzN6b+y1tZW6HQ6DBw4EAqFwmZ7aGgooqOj0dHRgfb2dofHqampQVtbGxoaGvDyyy9DqVTitttug1KpxO7duzF79mwcPXq0b5+Ggorm/BWLZTZTEfmG0zuOjo7rTQH9+jkeVBUeHg4AaG9vR0REhN19TP0bCQkJeOedd3DTTTcBuP6I7p/+9Cf8/e9/xzPPPIOvvvoKSqVrb26j4KapsXzrH8dwEPmG0zsOV279BUFwus+9996L/fv3Y/v27ebQAK7fsTz77LMYOXIk6uvrUVxc7PRYRIC96dQZHES+4DQVVCoVAKCrq8vhPqZtPd2VyGQyJCUl2X1kNyQkBJMnTwYAlJaWOiuJCABHjROJxWlwREREQKVSobm5GXq93ma7Xq9Hc3MzlEoloqKiel1IXFwcAKCzs7PXx6Dg0dGpw8X6FvOyTCZD2pA48QoiCiJOg0MmkyEjIwMGgwHV1dU226uqqmA0Gi3Gd9izbds2PP300zh8+LDd7RcvXgQAJCYmulA2BbvKi1ctmkhTk6LRTxkmYkVEwcOlZxdNc1HZ638wrTM1NTlSU1ODPXv2YOfOnTbburq6sHfvXgDAHXfc4UpJFORsH8VlMxWRr7gUHLNmzYJSqcTmzZst+iBOnTqFgoIChIeHY/78+eb1tbW1qKioQFNTk3nd7NmzIZfL8fnnn5tDAgC6u7uRl5eHS5cu4a677sKoUaM88blI4jhinEg8Ls1VNWTIEKxcuRJr167FvHnzkJWVBUEQUFJSAr1ej/Xr1yM2Nta8/8qVK3HkyBEsW7YMy5cvBwBkZGTg+eefx6uvvooVK1Zg9OjRSE5OxokTJ1BXV4f09HS8/vrr3vmUJDk2b/0byuAg8hWXZ8ddsGABkpOTUVBQgKNHj0KhUGD8+PFYunQpJk6c6NIxHn30UWRmZqKgoAAnT55EeXk5kpOTsWTJEuTm5qJ///69/iAUXM5xVlwi0bj1Po6pU6di6tSpTvfbunWrw20TJ050OWiI7DEYjKi8aDn4j8FB5Duc2IcCzqUrLejS/fxoeHSUCjEDeLdLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(prep_ecdf(mt_emp_percentile[:, i]))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 80,
|
||||
"id": "679fdc4b",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAu8AAAGXCAYAAAADGr5gAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAABeA0lEQVR4nO3deVyUVf//8dfAAIJLiBqCuRFCLmma+4JappWp3WZ6l0tqRuZWlqZmd4ve2malpGZqVoqZ2e2+pLmg4prmiqaiqCCh5JaKCgzz+8Mf841YhGFgZuD9fDy+j++361znXJ+LL37mw5lznctgNpvNiIiIiIiIw3OxdwAiIiIiIpI7Kt5FRERERJyEincRERERESeh4l1ERERExEmoeBcRERERcRIq3kVEREREnITR3gGIFIbg4OBcn9unTx/Gjh0LQO/evdm9ezf/+c9/6NWrV0GFJyIidxEXF8ejjz6apz7Tpk2jXbt2BRSRiH2oeJdiJSgoiFKlSuV4TuXKlQspGhERsUadOnVwd3e/63ne3t4FH4xIITPoJU1SHKTPvM+dO5cmTZrkul98fDw3b96kQoUKlClTpqDCExGRu/j7zPuGDRu477777ByRiH1o5l0kB/7+/vYOQURERMRCD6yKiIiIiDgJFe8iOejduzfBwcGEh4dbji1evJjg4GDGjRvHpUuXGDduHG3atKFOnTq0bt2ad999lwsXLmQ5XmxsLBMmTKBTp040aNCAOnXq0LJlSwYPHszOnTsznT969GiCg4P5+eef+f333xk6dChNmzblwQcfpGPHjsyYMYPk5OQcr9W+fXvq1q1L48aNefHFF4mMjMzy/IsXL/LRRx/RoUMH6tatS6NGjXjhhRf4+eefrfjJiYg4jvzk7eTkZL799lueeeYZ6tevz0MPPcS//vUvvv76a27fvp3p/PTPjd9++43333+fBg0a0KBBA/r27UtaWhoAJpOJRYsW0a1bNxo0aEDjxo0ZMmQIJ06c4IsvviA4OJgvvvgCgI0bNxIcHEzjxo2zzfdLly4lODiYvn372uYHJg5Ny2ZErHThwgW6du1KQkIClSpVolq1apw4cYIffviBrVu3snTp0gzr5CMjIxk8eDC3bt2idOnSVKlShdu3bxMbG8v69evZsGEDkyZN4qmnnsp0rZ07dzJixAgAqlevjqenJ9HR0Xz++eccOHCAL7/8MsP527Zt49VXX+XatWt4eXkRGBhIYmIikZGRREZGMmHCBLp162Y5PyoqipdeeomLFy/i7u5O9erVuXnzJjt37mTnzp107dqViRMnYjAYCuinKSJS8PKat69cucJLL73EwYMHcXFxoXLlypQoUYJjx45x5MgRVq1axddff03ZsmUzXeujjz5i//79BAUFceXKFSpUqICLiwspKSm88cYbrF27FoCAgACMRiPr169n69atPPzwwxnGCQkJwcfHh0uXLhEZGckjjzyS6VrLly8HoEuXLrb8cYmD0sy7iJV++eUXPDw8+N///seGDRtYuXIlCxYswNPTk3PnzvHjjz9azk1OTmbMmDHcunWLvn37sn37dpYuXcqaNWuIiIigefPmmM1mpk+fnuW1FixYQIsWLYiIiGDFihVs2rTJsp3lxo0bOXjwoOXcS5cu8cYbb3Dt2jW6d+9OZGQkixcvZsuWLYwZMwaA9957j7i4OACuXbvG4MGDuXjxIs8++yw7duxg+fLl/PLLL3z//ffce++9LF68mO+++66gfpQiIoUiL3kb7nz7efDgQerXr8/atWtZt24dy5cvZ8OGDTRs2JCoqChLLv6n/fv3M3XqVFasWMHmzZst582dO5e1a9dStmxZvv/+e9asWcOKFStYvnw59957L9u2bcswjtFopGPHjgCsXLky03USExPZuXMnnp6ePPbYY7b4MYmDU/EuxUqfPn0IDg7O9n/Wr1+fp/E+/vhjateubfnvBg0aWJLs/v37LccPHz5MUlISvr6+vPnmmxm2OCtfvjyDBw8GICYmxvK16t95e3szZcoUypcvn+FeqlSpkulaP/74I5cvX6ZevXqMGzeOkiVLAmAwGOjbty9t2rQhJSWF1atXA7Bw4UL++OMPGjduzPjx4zNspfnwww/z3//+F4CZM2eSkpKSp5+PiEhBePTRR3PM5cHBwYwePTrLvrnN24cOHWLTpk14e3szffp0S74F8PPzIywsjJIlS7JhwwZ+//33TNepX7++pZh2cXHB29ub1NRUZs6cCcCECRMyzLIHBQXxxRdfZPkN59NPPw3cmaxJSkrK0LZq1SpMJhOPPvroXbdClqJBy2akWLnbPu952RPY29ubevXqZTpevXp1AK5fv2451qBBA/bu3cutW7dwdXXN1MfT0xOAtLQ0bt++bfnvdI0bN6ZEiRJZXuvs2bMZrhUREQFA165ds/wQeP/990lJSaFSpUrAnQ8DgCeffDLL80NCQrjnnnu4ePEiUVFRPPTQQ5nOEREpTLnZ571atWqZjuUlb2/YsAGA5s2b4+Pjk6lPuXLlaNq0KRs2bGDLli088MADGdqzypX79u2zLKHJavnLAw88wMMPP8yePXsyHK9Tpw41atTgxIkTbNiwgU6dOlnatGSm+FHxLsXK22+/nad93nNy7733Znk8vcg2mUxZtkVFRXHkyBHOnj3L2bNnOX78ODExMZZzspp59/X1zfFaf+8TGxsLQI0aNbLsU7FixQz/ffLkSQDmzZtn+RD4p/QZ95iYGBXvImJ3U6ZMsWqf97zk7fTcuGfPHp577rks+6UvP/x7Dk9XoUKFTMfSxwwKCsr2GaKaNWtmKt4BOnfuzKeffsqqVassxfupU6eIioqifPnytGjRIsvxpOhR8S5iJTc3tzyd/+uvv/LBBx8QFRVlOWYwGKhatSqdOnXKtnDOzbX+/q61K1euAFiWy9xN+kxT+odKTq5du5arMUVEHFFe8nZ6brxw4UK2O9Gkyyo3enh4ZDqWnp//+e3q32WXu7t06cLnn39OZGQkV65cwdvbm2XLlgHw1FNPZfmtrhRNKt5FCsHx48fp378/ycnJNGzYkC5duhAcHMz9999PqVKliImJybF4z4sSJUpw/fr1TOsis+Pp6cm1a9f43//+R506dWwSg4iIs0svsEeNGkX//v1tOuaNGzeyPSe7Nl9fX5o2bcr27dv55ZdfePbZZy3PLmnJTPGiB1ZFCsG8efNITk6mWbNmzJ07l+7du1OvXj3L+vuEhASbXSt9nWd2M+mbNm2iZ8+eTJ06FYCqVavmeD7Arl27OHnyZLZ7DIuIFDW5yY1Hjhzh6NGjGdbK5yQwMBCAEydOZPjG9O9OnDiRbf/0In39+vVER0dz9uxZatSoQa1atXJ1fSkaVLyLFIJz584BEBwcnOVXmz/99JPl/85qrXxetGzZEsDydeo/rVixgj179nD16lUA2rRpA9zZpSarD5M9e/bQp08fOnbsSHx8fL5iExFxFum5cd26dVy6dClT+7Vr1+jbty9PP/00a9asydWYDRs2xNvbmz///JMtW7Zkaj979iy//vprtv3bt2+Pl5cXO3bssFyzc+fOubq2FB0q3kUKQfps+OrVqzlz5ozl+NWrV5k4cWKGvXuzemNfXvTs2ZMyZcrw66+/MnHiRMt4ZrOZefPmsWrVKtzc3OjZsycAzz//PGXLlmXPnj289dZbGdZuHjp0iOHDhwN3tmbLavcGEZGiqEmTJjRq1Ii//vqLl19+OUPuPn/+PIMGDeLq1atUqFAhw+4vOfHw8ODFF18EYOzYsRw4cMDSdvbsWQYPHpzjBI6Xlxft27fn9u3bzJkzBxcXl1xfW4oOrXkXKQT9+vVjxYoVXLhwgSeffJKAgAAATp8+TXJyMg888AAJCQlcuXKFCxcuZLlLQW7de++9fPbZZwwdOpTvvvuOxYsXU7VqVf744w8uXryIq6sr48aNsxTi5cqV44svvmDQoEEsXryYVatWERgYyPXr1y0fVsHBwXzwwQf5/jmIiNjCq6++etetIgEaNWrE66+/bvV1Pv30U1588UUOHjxIhw4dCAwMxMXFhVOnTpGSkkKpUqWYNWtWllv5Zqd///7s2bOHzZs30717d+6//37c3Nw4ceIEXl5eVK1alTNnzmT7AOrTTz/N0qVLSUpKomnTpvj5+Vl9f+KcVLyLFILKlSuzdOlSvvjiC/bs2UNMTAwlSpTggQceoGPHjjz//POMHj2aVatWsWnTpgwvELFGq1atWLZsGTNnzmTbtm0cO3aMkiVL0q5dO0JDQzPtc9yoUSNWrFjB119/zZYtW4iOjgbubGfWoUMH+vXrl+vda0RECtrhw4dzdV7ZsmXzdR1fX18WLVrE/PnzWbNmjaVov/fee2nZsiWhoaF53rLSaDQyffp0wsPDWbx4MWfOnMHDw4PHHnuM4cOHM378eM6cOZPtHwRNmjTB19eX8+fP60HVYspgzu6JCREREREpVM8++ywHDx7kk08+yXI9e1JSEi1atMBsNhMZGam3qhZDWvMuIiIiUghu3LhBq1at6N27Nzdv3szUnpiYyNGjRwGy3UFmzZo1JCUl8fjjj6twL6ZUvIuIiIgUgpIlS3LPPfewe/duJk+enGGDgoSEBIYPH05KSgoNGza0bCsJcObMGeLi4tiyZQsff/wxAL169Sr0+MUxaNmMiIiISCHZvn07L7/8MsnJyZQuXZoqVapw8+ZNzp49S2pqKlWrVmXOnDkZ1tJ/8sknzJ492/LfzzzzDBMnTrRH+OIA9MBqLqSmppKQkEDFihUxGvUjExFxdMrb4qiaN2/OqlWrmDNnDr/++isxMTG4urpSo0YN2rdvT+/evSldunSGPrVr16ZUqVIYjUaeeuopRo0aZafoxRFo5j0X4uLiePTRR9mwYUOenyoXEZHCp7wtIkWV1ryLiIiIiDgJFe8iIiIiIk5CxbuIiIiIiJNQ8S4iIiIi4iRUvIuIiIiLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 864x360 with 2 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(1, 2, figsize = (12, 5))\n",
|
||||
"\n",
|
||||
"sns.histplot(mt_emp_percentile[:5].numpy().reshape(-1), ax=ax[0], bins=bins, label = \"Multitask\",\n",
|
||||
" color=palette[4], alpha = 0.5, stat=\"density\")\n",
|
||||
"sns.histplot(ind_emp_percentile[:5].numpy().reshape(-1), ax=ax[0], bins=bins, alpha = 0.5, \n",
|
||||
" label = \"Independent\", color = palette[-2], stat=\"density\")\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"sns.histplot(mt_emp_percentile[5:].numpy().reshape(-1), ax=ax[1], bins=bins, label = \"Multitask\",\n",
|
||||
" color=palette[4], alpha = 0.5, stat=\"density\")\n",
|
||||
"sns.histplot(ind_emp_percentile[5:].numpy().reshape(-1), ax=ax[1], bins=bins, alpha = 0.5, \n",
|
||||
" label = \"Independent\", color = palette[-2], stat=\"density\")\n",
|
||||
"ax[1].legend(ncol=2, bbox_to_anchor=(-0.1, -0.4), loc = \"lower center\")\n",
|
||||
"#plt.tight_layout()\n",
|
||||
"ax[0].set_xlabel(\"Quantile\")\n",
|
||||
"ax[1].set_xlabel(\"Quantile\")\n",
|
||||
"\n",
|
||||
"ax[0].set_title(\"Finance\")\n",
|
||||
"ax[1].set_title(\"Energy\")\n",
|
||||
"sns.despine()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "af87f0f2",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.8"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,805 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "fc81655d",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"[pyKeOps]: Warning, no cuda detected. Switching to cpu only.\n",
|
||||
"Warning no robinhood utils.\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns\n",
|
||||
"import torch\n",
|
||||
"import gpytorch\n",
|
||||
"# from voltron.robinhood_utils import GetStockData\n",
|
||||
"import os\n",
|
||||
"# import robin_stocks.robinhood as r\n",
|
||||
"import pickle5 as pickle\n",
|
||||
"\n",
|
||||
"import sys\n",
|
||||
"sys.path.append(\"../\")\n",
|
||||
"from voltron.likelihoods import VolatilityGaussianLikelihood\n",
|
||||
"from voltron.kernels import BMKernel, VolatilityKernel\n",
|
||||
"from voltron.models import BMGP, VoltronGP, SingleTaskVariationalGP\n",
|
||||
"from voltron.means import LogLinearMean\n",
|
||||
"from gpytorch.kernels import ScaleKernel, RBFKernel, MaternKernel"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "28cce2f6",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"sns.set_style('white')\n",
|
||||
"# style.use('whitegrid')\n",
|
||||
"palette = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"sns.set(palette = palette, font_scale=2.0, style=\"white\", rc={\"lines.linewidth\": 4.0})"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"id": "ed0d165c",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"<torch._C.Generator at 0x7f831eb14510>"
|
||||
]
|
||||
},
|
||||
"execution_count": 3,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"torch.random.manual_seed(200)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"id": "6c8c58d0",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ntest = 200\n",
|
||||
"ntrain = 200\n",
|
||||
"tckrs = [['BAC', 'GS', 'JPM', 'MS', 'WFC'], ['COP', 'CVX', 'EOG', 'SLB', 'XOM']]\n",
|
||||
"indexes = [\"XLF\", \"XLE\"]\n",
|
||||
"span = \"5year\"\n",
|
||||
"interval = 'day'\n",
|
||||
"T = 5."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 5,
|
||||
"id": "c263fffd",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"['BAC' 'BRK.B' 'C' 'GS' 'JPM' 'MS' 'WFC']\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"<Figure size 720x360 with 0 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAtoAAAE7CAYAAADjHiaJAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAC/iklEQVR4nOzdd3iUVfbA8e8kk957AgRIIAQIHUJHQJAmTVRQRAUri6hrRf3pKuq6rrqKqFgRFRUQRRGV3puU0AOEEkpCEkjvbcrvj8skmcykkpAA5/M885B5651kmDnvfc89V2M0Go0IIYQQQggh6pRNQzdACCGEEEKI65EE2kIIIYQQQtQDCbSFEEIIIYSoBxJoCyGEEEIIUQ+uu0Bbp9MRHx+PTqdr6KYIIYQQQogb2HUXaCclJTFkyBCSkpIauilCCCGEEOIGdt0F2kIIIYQQQjQGEmgLIYQQQghRDyTQFkIIIYQQoh5IoC2EEEIIIUQ9kEBbCCGEEEKIeiCBthBCCCGEEPVAAm0hhBBCCCHqgQTaQgghhBBC1AMJtIUQQgghhKgHEmgLIYQQQlxL9AVQnNPQrRDVIIG2EEIIIW5cRgOk7YOi9IZuSfWc/Ax+coVfg+DcTw3dGlEFbUM3oDEwGAykp6eTk5NDQUEBBoOhoZsk6pmNjQ2Ojo64urri5eWFjY1ccwohxA3HoIONw+HiBrBxgCEbwK9vQ7eqYukHYO9j6uJAlwN7Z0DwbWBj19AtExW44QNtnU5HXFwcWq0Wb29vnJ2dsbGxQaPRNHTTRD0xGo0YDAby8vLIyMggKyuL4OBgtNob/r+DEELcWOJ/VUE2gKEQ9syAkfuhMcYARoNqn7FMZ2BhKlzcCEHDGq5dolI3fDdeWloaDg4ONGvWDDc3N2xtbSXIvs5pNBpsbW1xc3OjWbNmODg4kJaW1tDNEkIIcbWd/tr8ecZBOPIGGI0N057KnFsMKTstl5//+eq3RVTbDR9oZ2Zm4uPjI8H1DUqj0eDj40NmZmZDN0UIIcTVcvh1WKSFxFVW1r0Kv4fAodcgJ/Zqt6xip760XKZ1BVuHq98WUW03/L1ynU6Hvb19QzdDNCB7e3t0Ol1DN0MIIcTVYChW/xr1FW+Tew6OzAaNLXR8pe7OrcuF1b2g6RhoeQ94dqjefrlxcGmz+bIub0P4k2DrqJ4XpsHZ70GXB60eBEe/umu3qLUbvkcbkN7sG5z8/YUQ4gZiYwcd/wW3JUG398HBp+JtQ6bU7bnjV0BmNBx9G/7qCBtHVr2P0ai2p0w6i3cPaD+rNMg+vxT+aAtRT8LBF2HbHY0z/eUGJIG2EEIIIW4sxTlw6gs4850aUGiNbx/V631xY+3PU5QBUf+EY++r5+cWma/3iFA90HG/woW/LPc3GmHTrXBynvnylmUuAOKWwbaJUJhcuuzSFlWhRDS4Gz51RAghhBA3GDtXCBgM6VGqrF/ES+DTU/UIx/2iAuyUnbAiTG0/MRe0ztU/vkEPsfPh4P9BYYrKpW56KyStMd8u7hcVROvzwac3NB1Vus5ohAOzIHGl+T4aW2hxl/pZX6AqkVgT/xt4d61+m0W9kEBbCCGEEDce//7qUVb/JSqP+s8IladtknsWPNpX/9jFGbB/lvoXVM3rP9qab6OxVcc1Sf0b8hLAuYl6nrQWjr1reezQB8ApQP1s6wiDVsL2uyD7hPl2R14H/5vUBYVGEhgaivzmRaWWLVtGeHi4xaNdu3b06tWLu+++m0WLFlU5yc/bb79NeHg47du35+LFi9U+/86dO3n22WcZMWIEXbp0oWvXrtx22218/PHHUilECCFE3dO6gGuo+bKcM1Xvl3EYirPVzw4+0On1yrdvPV31opd1YXnpz4G3QIdyAzFbPwI9PjJf5t0VRkRBk9GW59gwFNb0hfykqtsv6oUE2qJafHx8GDNmTMnjlltuITw8nJiYGF577TWeeuopjBUMvNDpdKxYsQIHBwf0ej0//1x1zc+MjAymT5/O1KlT+fPPP3FycqJ///506NCB8+fP89FHHzFq1ChOnjxZ1y9VCCHE9aTgEux9En5yh+MfgL6o6n0sAu1qlPnbegf84qdyqk99CaH3Q7NxFW/fdDQ0u818WXyZQFujUcH6Ldugy39hxD7o+bn1cn52rjBoherBLi91F6zuCbnnq34Nos5J6oiollatWvHee+9ZLE9PT+eee+5h1apVrFu3jltuucVim82bN5OSksIjjzzCt99+y9KlS/nHP/5R4bTnRUVF3HfffcTExDBgwABefvllWrZsWbI+KyuL//znPyxbtoypU6eybNkyAgIC6uy1CiGEqCNFGZC2D7y7gb1nw7Rh18Nw4Xf1876n1aPpWOjxMbgEW9/HJcT8eVU92vmJpakbCX9B4hpoMQn6L4Wd91sOgrRxgIBB4NJCVQkxubRVXQjYlik77NdPPaqj2QQ1ELK8vDg49j/o8WH1jiPqjPRoiyvi5eXFtGnTAFi7dq3VbZYtWwbA8OHDuemmm0hMTGTLFisfBJd98MEHxMTE0KdPHz799FOzIBvA3d2df//73/Tt25eUlBTmz59fNy9GCCFE3ck9r0rYbRgCK1pb9gpf2gonP63fNuQnlQbZZSWurrysX/ke7dzLgXZFJfPKz9jo3R3s3FUpwd7fWPY0Nxmp8qvd24JTk9Ll+jzYO6PiSihVaf0QeFRQmzttb+2OaY3RCAdegj/aQ9TT5tPCCzMSaIsrZupNzs3NtViXlpbG5s2b8fX1JSIigrFjxwKwePFiq8cqKChgyZIlaDQaXnzxRezs7KxuZ2Njw2OPPUaHDh1wdq7BSHAhhBBXx76nIS9e/VyYCof+pX42GlUKx/rBsHem9RkPayMtCrbeDtvugqzLaYUJf1rf1n9g5VVEXMv1aMf/BktcYHkL1VtdXvlA27dP6c+29qpn2yPi8nMn6PCy+lmjgYCbzfc9PR+WBcD6IZC8o+I2WqN1gZH7VZrJiCjzdZlH6662dtwyOPofyDoGMR/A0XcgZbcqmyjMSOqIuGLR0dEAdO7c2WLd77//TnFxMaNHj0aj0TBo0CA8PT3ZsmULiYmJBAUFmW2/adMmcnNzadOmDeHh4ZWet0ePHvzyyy9190KEEELUjUvbVOm6ss7+oGYzPPJvOPVZ6fLdj6iBhN0/VIFnbcT/DtsnqXJ3ANkxKti8sML69uFPVH4811aWy/R5kJcHm8fAgF/NS/GVD4j9+po/d/SH4XtUr7Jbm9KqIaAC7bPfm29v1MPFDaD5d+XttMZGqwZIGo3Q7nnwaKd6uT3a1f73a9Y2I5z70XyZKf3FwQeG77a8I3ADkx7tSmzcCO0uvy+vpUe7dqrt9Umv15OamsrSpUv57LPPCA4O5q677rLYzpQ2MmHCBEBNdz5mzJgKB0XGxqpbix06VHNaWiGEEI1H5nHY+7gKnq058u/LAwTLBXwnPoJFNvB7mOrRNRRDQUr1znnhT9g6oTTIBjVZS/J2SCyX0tjkVhj4p6ppXRkHX7DztL7OUATbJ0LB5Qli9IWWaRlle7RNtE7gP8A8yAYIvNlyWwCnIMuqJDWh0UDX/0LoVPDpoXq7a8JogMxjpWUO9UVqYpyfPVWPtjWFqXB4du3bfB2qdo+2Xq9n0aJF/Prrr8TGxqLX6wkODmbUqFE89NBDODiUjoJNTExk0KBBFR6rW7duLFpkPjAgKyuLzz//nHXr1pGYmIivry/Dhg1j5syZuLq61vyV1YFHH4VrsajF8eOq7SdOVL1tde3evbvCHmY/Pz+++eYb3N3dzZZHR0cTExNDRESE2b533HEHCxcuZOnSpcyYMQNbW9uSdcnJ6oPL19e37hovhBCi/p1ZCHv+oepQV+TUZ+DVRdWr/nua5bY5p2DXQ+oBEHIf9F5QcR3oS1tVDWmj3ny5R3s1KFGfV7rMuRkMXFG9Xl2NBtzCIG2P9fW6XDj/E7R5TA32NJSpZOLcXJ2rulxaqB70nNPmy9s+27D1r/c+ribT0dhAr/lQcFFN9V6V8z9Dj09UJRRRvUBbr9czY8YMNm3ahLOzM507d0ar1XLw4EHmzp3L5s2b+fbbb3FycgLg6NGjAISHh9OmTRuL44WEmOc+5eTkMGXKFGJiYggJCWHQoEFER0ezYMECtm7dyuLFi3Fzc7vS1yqugI+PD337lt4KMxqNZGdnc/z4cS5evMhdd93FvHnz6NSpU8k25XuzTdq2bUv79u05evQomzZtYsiQISXrTEG3Tqerz5cjhBCiLh2eDYdfq3wbz87QYy74DbgcyIarPO2itIr3OfOdGlhYNtVDX6QC9oS/1KDG8jr8CzrNht3/MF/eZHTNUicqC7QBsmIut/Eb8+XWerOrEnKv+e+v/YvQ9p81P05dSd1bOu270QAHXgRjcen64NvVAM+oJy331edB3M+qJ11UL9BeunQpmzZtIjw8nC+//LJk8FtaWhozZsxg//79zJs3j2eeeQaAY8eOAfDQQw+VDH6rzJw5c4iJiWHixInMnj0bGxsbdDodL730EsuXL2fOnDm88sorVR6nrn3+OTz2GFx+OdeMdu3gk0/q9pgVlfczGAx8+eWXvP/++zz66KOsW7cOFxcXioqK+OOPPwAVcK9ebf5hmJqqRlQvWbLELND29/cH1HvrhmPQQ2Y02HuBYwCcXwL23mrSgrKlnoQQorFpdhscfds8fQPUrIRh/4Ad90DGQTWBim9faP8CNBkBA/9Qy8r2PJd34AVoMgrcWqvBdksr6XjrPhfCH1d5xAl/mK9rOqZmr8ktrPL12SchaR2c+sJ8eUWpIJVp/yIYdOo7oNUDqsZ2Q/LurqaN110e3FhQbsKbC7+rmt7JO9R3VXmpeyTQvqxa9yR+/fVXAF566SWzesXe3t689tprAPz5Z+nIXlOPdkRERJXHzsrKYunSpbi6ujJr1qyS2sparZZXX30VDw8Pfv75Z/LyKvlPWE8GD4aLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 864x360 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"['COP' 'CVX' 'EOG' 'SLB' 'XOM']\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAtoAAAE7CAYAAADjHiaJAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAEAAElEQVR4nOydd1gU19fHv7v03kFERVGKCNi7xt5770k0atQYTbGkvNFo8jMajTHWqLH3FnvvsRdsKIoCggLSpdfdnfeP47I7u7OwwK4U7+d5eGDu3Jm5u+zOnHvuOd8j4jiOA4PBYDAYDAaDwdAp4rIeAIPBYDAYDAaDURlhhjaDwWAwGAwGg6EHmKHNYDAYDAaDwWDoAWZoMxgMBoPBYDAYeoAZ2gwGg8FgMBgMhh6odIa2RCJBVFQUJBJJWQ+FwWAwGAwGg/EBU+kM7djYWHTq1AmxsbFlPRQGg8FgMBgMxgdMpTO0GQwGg8FgMBiM8gAztBkMBoPBYDAYDD3ADG0Gg8FgMBgMBkMPMEObwWAwGAwGg8HQA8zQZjAYDAaDwWAw9AAztBkMBoPBYDAYDD3ADG0Gg8FgMBgMBkMPMEObwWAwGAwGg8HQA8zQZjAYDAaDwWAw9AAztBkMBoPBYDAYDD3ADG0Gg8FgMBgMBkMPMEObwWAwGAwGQxvSXgCRe4CchLIeCaOCYFjWA2AwGAwGg8EoV6Q+A860BPJTgKo9gaZrgKxo4HwHQJYLmFUFetwHTJ3LeqSMcg4ztBkMBoPBYDAAIPUpcOszIPGGoi3mBHCpJyA2IiMbALJjgOergYCfy2SYjIoDM7QZDAaDwWB82DyaA1i4AyZOfCNbTuoT9bYXKwH/OYCIReEyNMMMbQaDwWAwGB8uuUnA41+Kf5ypK5D9BjB1AYIXAW9OA1V7AL7fASKR7sfJqJAwQ5vBYDAYDMaHS/oL/raJI2DlCSTdATiJ8DENlwA+35BBfXMsEL6Z2hOuANZ1ger99TliRgWCrXcwGAwGg8H4cEl7zt92/gjoeh0Yngf0fgZY1OTvd2ihMLIBoEoX/v7oo3obKqPiwTzaDAaDwWAwPlzSVQxtKy/6LRIB1t5ArydA6FqS9TO2IwUS5dAQVeWRpFv6HS+jQsEMbQaDwWAwGB8umgxtOYbmgM/X9COEfVMAIgAcbacGA/lpgJG1rkfKqICw0BEGg8FgMBgfLqqhI9Zewv00YWwDWPsoNXBA0t1SD4tROWCGNoPBYDAYjA8TTla0R1sbHJrxt1n4COMdzNBmMBgMBoPxYRG+GTjbFrgyEJBmK9qNbEl1pLg4NudvM0Ob8Q4Wo81gMBgMBuPDQZIJPFsGpDxU32ftVTINbAcVQzvxFsBxTE+bwTzaDAaDwWAwKilhG4CoI/w2Qwug8yXArKp6f+d2JbuOrT9gYKrYzokFcuJKdi5GpYIZ2gwGg8FgMCofWdHA3WnAf/2AK4NoW07MKaDGEH5/M1eg7qySXUtsBNj48dtSHpXsXIxKBTO0GQwGg8FgVD4idgLSLPr79b/A2daA7F2lR5d2wOsD/P4ttgCmJYjPlmMbwN9mhjYDzNBmMBgMBoNRGYm/zN+uMwkQv0tNM3MFOp6jKpCWtcnIdu2ifo7ioGZoB5XufIxKAUuGZDAYDAaDUbHhOCA3ETB1om2ZFEi4wu/j1pu/be0NdFYxxkuDHfNolwlZMfTbXCDmvhzADG0Gg8FgMBgVl5x44EJXUhGxbww0XgHI8qg6oxxje8DGV7/jsPHnb6c9o1AVMTO19MrzFUDwQlqZcP4IqD0ecGpV1qMqgP33GQwGg8FgVFzuTlNI9SUHAmcFjCznjwCRnqNlTR0B3+8Bq9oURmLjy4zs90H8f/Q7I4x+qnRhhjaDwWAwGAxGqeE4oEonIO48hY5owvmj9zOeBgvez3UYhCQLSL7Db3NuWzZj0QAztBkMBoPBYFRMRCKgzgSq5nhloOZ+JdXHLg3Zb4C4i0BOApCbAFi401gZuiPpFiDLV2xbegDm1cpuPAIwQ5vBYDAYDEbFpvoAoOU24OkS9YqPFu6Abf33P6a3j4DroxTbLh2Zoa1r5GEjct7XykUxYIY2g8FgMBiMik+t0fQDAFFHgZA/AWkO0HAxIDZ4/+ORK6DIyYl//2OoyEhzgMg99Ns2AHBsrh5nr2poOzFDm8FgMBgMBkO/VOtDP2WJiYqhnZtQNuOoqPw3EHhzUrFt14DkGI2saTs7Vl0rvZzFZwOsYA2DwWAwGIyKhEwKZL4GOFlZj6RwVD3auYnlf8zlhYwIvpENAG8fAE+XKrZfbgU4qWLbxpck/soZzNBmMBgMBoNRccgIAw7XAPZaAifqA7cnl/WIhDEwBQytFNucFMh7W3bjqUioeqrlhK2j5Mf8DODBbP4+j88oObacwQxtBoPBYDAYFYf05/Rbmk3VF9Oelu14CkMtTpuFj2iFJkM7+w0QdRjIjOS3i42AWmP0P64SUGJD+99//4W3tzfu3r0ruP/y5cv47LPP0KxZM/j5+aFDhw6YM2cOYmNjBfunpaVh8eLF6NatGwICAtCxY0csXLgQGRkZJR0ig8FgMBiMykbac/62lVfZjEMbWJx2ydBkaDu1BtJCANt6gGNLRbtbP/VJTTmhRIb2/fv38csvv2jcv27dOkycOBHXr19HrVq18NFHlAW6Z88eDBgwAGFhYbz+GRkZGD16NP755x+IRCK0b98eIpEImzZtwrBhw5Cenl6SYTIYDAaDwahspKsY2tbl2NA2deZvM0O7aLKigIxw4X3JgUB2DMXpV+tHbWJjoN4P7298xaTYhvaZM2fw2WefISsrS3B/aGgoli1bBnNzc+zcuRN79uzB6tWrcebMGYwcORLJycn44Qf+G7Js2TKEhIRg6NChOHHiBJYvX47Tp0+jX79+BedjMBgMBoPBUDO0K5JHm0n8FU3cRf62U2vAfQT9Lc0BpFkk11h9EFDvR6BXMGDf8P2PU0u0NrRjY2Mxa9YsfPnll5DJZHB0dBTsd/jwYUilUowdOxYNGypeuJGREX744QfY29vjwYMHiI6OBkAhI/v27YOlpSVmz54NsZiGZGhoiLlz58LGxgb79+/XaNgzGAwGg8GoAMSeB+5MASJ2lk59oyKFjrAY7eKRmww8+onf5twOaLERaL0HaLMfaLqW2q3qAPV/BazKn9KIMlob2suWLcPhw4fh5+eHPXv2wMPDQ7CfkZERvL290bRpU8F91apRacz4eJrV3blzBzk5OWjRogUsLS15/S0sLNCyZUvk5OTgzp07audjMBgMBoNRAUi6A1zqAbxYQ9USL/ctmQJHTgKQHa3YFomp7HZ5xYSFjhSL9BdAvkq4cPWBpODiPhSoMQgwMC6bsZUQrQ1tDw8PLFq0CPv27YO3t7fGftOmTcORI0fQsmVLtX1ZWVkIDQ0FAFSpUgUACrY9PT01XhcAQkJCtB0qg8FgMBiM8sST30iWTU7MceB8ZwoFyE+j4iRJWjjUHqvkh1l5lW/Di1WHLB6OzYEeDxSJjt7TAfvGZTqk0qJ1ZciJEyeW+mLr169HVlYW/P394erqCgBISKDZnZOTcLaovD0pKanU12cwGAwGg/GeyQgHog6pt7+9B+wxA8yrA1mvgcQbQLebgIW7ok/yfSDsH1KayEuioiXK1B6vz5GXHqY6UnwsqlMFyOerAc9yqpFeDN6bjvbly5exdu1aiMVizJw5s6BdHnttZmYmeJypqSmvH4PBYDAYlY7cZODB98DNz4CwDcCh6kDUkaKPy0uhmOc3ZwCO0/swiw3HAQ9/AlDI2LJe0++cWOBSbyAvVbEv6RbwYjUQd17dyLaoCXhN1fGAdYxadUhmaBfAcUD8VeDVPvV9YiPAZ3r5Xq3QEq092qXh0qVLmDZtGqRSKb799ls0b968YJ88+VGkoZoP9+7GwZXHGwijfCDNAdJDAWtv+nIyGAxGRSIjArjYTaGmEb6Rfv/XD6jaC2j8l3rCV34a8Owv4NkfQP47w9R3NtBg4XsbtlZkvgSiDmrf39AckOUqtqsPAu5O5ZfallN/AWBgUvox6hNLD6DxCpL5M3UCTF3LekTlg5wE4Pbn9Nlw6wvUGFLWI9Ibevdo79+/H1988QVyc3PxxRdfqIWgmJubAwBycnIEj8/NpS+cJo834wMnMxI45guc8AcOuwMvd5C+JoPBYFQEchKBcx+pS9bJiTkOBP2s2OY44MVa4HBNIGiOwsgGgOBFwNWhQOItfY64eFh6AK13U9IiAJi5AV2uAQYqz3SxMcXjdrrI1542dQKqdOb3NTAn3WT34foduy4wtgO8p1Iin0sHwManrEdU9sjygXPtFBMwA9OyHY+e0atHe9myZVizZg1EIhG+//57fPrpp2p9nJ3pC5WYmCh4jqJiuBkfEKlPaemwShfA1JEyk6+NJI8JQKVZb4wG7n8LeIwF6s4ETOzLdMgMBuMDIPsNYGgJGFkV/9jH8xWhE0IYWQMNfgNSg4GIXcCTXws/36t99NNgIXm4ywPV+gJN/wbuzwLa7AOcWgLdbgGRewFZHmBWFagxGDB3Ez6+3g9AzdH0/ppWAewaVoqQgg+WhGtA2lPFtuqkq5KhF0Ob4zj83//9H/bv3w9jY2MsWrQIPXv2FOwrVxuRq4+oIq8iWZjSCeMDIPYCSUPJ8orumxMHBC8Env0JOLYArOsCshwg5Ql5Vay8AO9pgEMT/Y+bwWBUXjgZLX+H/QMYWgENFwOen1OMcXYMJfkZWQJvHwIhywCIgIZLFA6AjJdA6N/q53VoDogNyftbdwZJ4j1ZULyxPV9ZfgxtAKgzAag2gJwkAGDrTz/a4PyR/salCzgOOHkSuH4d6NoV+Kicj7esSb7L3445DkiyKGyoJCQkABIJ4Fo+w3L0YmgvXLgQ+/fvh6WlJdasWYNmzZpp7Nu0aVOYmprixo0byMrKKgglAYDMzEzcuHED5ubmaNy4Ysu7MIoLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 864x360 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.figure(figsize = (10, 5))\n",
|
||||
"idx = -2\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"colors = [\"blue\", \"orange\", \"red\", \"green\", \"purple\", \"maroon\"]\n",
|
||||
"\n",
|
||||
"train_y_list = []\n",
|
||||
"test_y_list = []\n",
|
||||
"for j, (tckr1, loc) in enumerate(zip(tckrs, indexes)):\n",
|
||||
" with open(\"../../spdr-data/\"+loc+\".pkl\", \"rb\") as handle:\n",
|
||||
" raw_data = pickle.load(handle)\n",
|
||||
" print(np.unique(raw_data[\"symbol\"]))\n",
|
||||
" plt.figure(figsize = (12, 5))\n",
|
||||
" for i, tckr in enumerate(tckr1):\n",
|
||||
" data = raw_data[raw_data[\"symbol\"] == tckr]\n",
|
||||
" ts = torch.linspace(0, T, data.shape[0]) + 1\n",
|
||||
" train_x = ts[:ntrain]\n",
|
||||
" test_x = ts[ntrain:(ntrain+ntest)]\n",
|
||||
"\n",
|
||||
" y = torch.FloatTensor(data['close_price'].to_numpy())\n",
|
||||
" train_y = y[:ntrain]\n",
|
||||
" test_y = y[ntrain:(ntrain+ntest)]\n",
|
||||
"\n",
|
||||
" dt = ts[1] - ts[0]\n",
|
||||
" train_y_list.append(train_y)\n",
|
||||
" test_y_list.append(test_y)\n",
|
||||
"\n",
|
||||
" plt.plot(train_x, train_y, label = tckr, color = colors[i])\n",
|
||||
" plt.plot(test_x, test_y, linestyle=\"--\", color = colors[i])\n",
|
||||
"\n",
|
||||
" plt.legend()\n",
|
||||
" sns.despine()\n",
|
||||
" plt.show()\n",
|
||||
"\n",
|
||||
" train_y = torch.stack(train_y_list)\n",
|
||||
" test_y = torch.stack(test_y_list)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 6,
|
||||
"id": "252b7b0b",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def get_and_fit_gpcv(x, log_returns):\n",
|
||||
" train_x = x[:-1]\n",
|
||||
" # prepare model\n",
|
||||
" likelihood = VolatilityGaussianLikelihood(param=\"exp\")\n",
|
||||
" # likelihood.raw_a.data -= 6.\n",
|
||||
" covar_module = BMKernel()\n",
|
||||
" model = SingleTaskVariationalGP(\n",
|
||||
" init_points=train_x.view(-1,1), \n",
|
||||
" likelihood=likelihood, \n",
|
||||
" use_piv_chol_init=False,\n",
|
||||
" mean_module = gpytorch.means.ConstantMean(), \n",
|
||||
" covar_module=covar_module, \n",
|
||||
" learn_inducing_locations=False,\n",
|
||||
" use_whitened_var_strat=False,\n",
|
||||
" )\n",
|
||||
" # model.mean_module.constant.data -= 4.\n",
|
||||
" model.initialize_variational_parameters(likelihood, train_x, y=log_returns)\n",
|
||||
" \n",
|
||||
" import os\n",
|
||||
" smoke_test = ('CI' in os.environ)\n",
|
||||
" training_iterations = 2 if smoke_test else 500\n",
|
||||
"\n",
|
||||
"\n",
|
||||
" # Find optimal model hyperparameters\n",
|
||||
" model.train()\n",
|
||||
" likelihood.train()\n",
|
||||
"\n",
|
||||
" # Use the adam optimizer\n",
|
||||
" # likelihood parameters should be taken acct of in the model\n",
|
||||
" optimizer = torch.optim.Adam([\n",
|
||||
" {\"params\": model.parameters()}, \n",
|
||||
" # {\"params\": likelihood.parameters(), \"lr\": 0.1}\n",
|
||||
" ], lr=0.01)\n",
|
||||
"\n",
|
||||
" # \"Loss\" for GPs - the marginal log likelihood\n",
|
||||
" # num_data refers to the number of training datapoints\n",
|
||||
" mll = gpytorch.mlls.VariationalELBO(likelihood, model, log_returns.numel())\n",
|
||||
" \n",
|
||||
" old_loss = 10000.\n",
|
||||
" print_every = 50\n",
|
||||
" for i in range(training_iterations):\n",
|
||||
" # Zero backpropped gradients from previous iteration\n",
|
||||
" optimizer.zero_grad()\n",
|
||||
" # Get predictive output\n",
|
||||
" with gpytorch.settings.num_gauss_hermite_locs(75):\n",
|
||||
" output = model(train_x)\n",
|
||||
" # Calc loss and backprop gradients\n",
|
||||
" loss = -mll(output, log_returns)\n",
|
||||
" loss.backward()\n",
|
||||
" if i % print_every == 0:\n",
|
||||
" print('Iter %d/%d - Loss: %.3f' % (i + 1, training_iterations, loss.item()))\n",
|
||||
" optimizer.step()\n",
|
||||
" if old_loss <= loss and i > 100:\n",
|
||||
" print(old_loss, loss)\n",
|
||||
" break\n",
|
||||
" else:\n",
|
||||
" old_loss = loss.item()\n",
|
||||
" \n",
|
||||
" model.eval();\n",
|
||||
" likelihood.eval();\n",
|
||||
" predictive = model(x)\n",
|
||||
" pred_scale = likelihood(predictive).scale.mean(0).detach()\n",
|
||||
" samples = likelihood(predictive).scale.detach()\n",
|
||||
" \n",
|
||||
" plt.plot(x, pred_scale, linewidth = 4)\n",
|
||||
" plt.plot(x, samples.t(), color = \"gray\", alpha = 0.3)\n",
|
||||
" # plt.ylim((0, 0.25))\n",
|
||||
" plt.show()\n",
|
||||
" \n",
|
||||
" # return scaled volatility prediction\n",
|
||||
" return pred_scale / dt**0.5\n",
|
||||
" \n",
|
||||
"\n",
|
||||
"def get_and_fit_vol_model(train_x, est_vol):\n",
|
||||
" vol_lh = gpytorch.likelihoods.GaussianLikelihood()\n",
|
||||
" vol_lh.noise.data = torch.tensor([1e-6])\n",
|
||||
" vol_model = BMGP(train_x, est_vol.log(), vol_lh)\n",
|
||||
"\n",
|
||||
" optimizer = torch.optim.Adam([\n",
|
||||
" {'params': vol_model.parameters()}, # Includes GaussianLikelihood parameters\n",
|
||||
" ], lr=0.01)\n",
|
||||
"\n",
|
||||
" # \"Loss\" for GPs - the marginal log likelihood\n",
|
||||
" mll = gpytorch.mlls.ExactMarginalLogLikelihood(vol_lh, vol_model)\n",
|
||||
" old_loss = 10000\n",
|
||||
" for i in range(500):\n",
|
||||
" # Zero gradients from previous iteration\n",
|
||||
" optimizer.zero_grad()\n",
|
||||
" # Output from model\n",
|
||||
" output = vol_model(train_x)\n",
|
||||
" # Calc loss and backprop gradients\n",
|
||||
" loss = -mll(output, est_vol.log())\n",
|
||||
" loss.backward()\n",
|
||||
" if i % 50 == 0:\n",
|
||||
" print(loss.item())\n",
|
||||
" optimizer.step()\n",
|
||||
"# if old_loss <= loss:\n",
|
||||
"# break\n",
|
||||
"# else:\n",
|
||||
"# old_loss = loss.item()\n",
|
||||
" \n",
|
||||
" return vol_model\n",
|
||||
"\n",
|
||||
"def get_and_fit_data_model(train_x, train_y, pred_vol, vol_model):\n",
|
||||
" voltron_lh = gpytorch.likelihoods.GaussianLikelihood()\n",
|
||||
" voltron = VoltronGP(train_x, train_y.log(), voltron_lh, pred_vol)\n",
|
||||
" # voltron.mean_module = gpytorch.means.LinearMean(1)\n",
|
||||
" voltron.mean_module = LogLinearMean(1)\n",
|
||||
" voltron.mean_module.initialize_from_data(train_x, train_y.log())\n",
|
||||
" voltron.likelihood.raw_noise.data = torch.tensor([1e-6])\n",
|
||||
" voltron.vol_lh = vol_model.likelihood\n",
|
||||
" voltron.vol_model = vol_model\n",
|
||||
"\n",
|
||||
" grad_flags = [False, True, True, True, False, False, False]\n",
|
||||
"\n",
|
||||
" for idx, p in enumerate(voltron.parameters()):\n",
|
||||
" p.requires_grad = grad_flags[idx]\n",
|
||||
"\n",
|
||||
" voltron.train();\n",
|
||||
" voltron_lh.train();\n",
|
||||
" voltron.vol_lh.train();\n",
|
||||
" voltron.vol_model.train();\n",
|
||||
"\n",
|
||||
" # Use the adam optimizer\n",
|
||||
" optimizer = torch.optim.Adam([\n",
|
||||
" {'params': voltron.parameters()}, # Includes GaussianLikelihood parameters\n",
|
||||
" ], lr=0.1)\n",
|
||||
"\n",
|
||||
" # \"Loss\" for GPs - the marginal log likelihood\n",
|
||||
" mll = gpytorch.mlls.ExactMarginalLogLikelihood(voltron_lh, voltron)\n",
|
||||
"\n",
|
||||
" for i in range(500):\n",
|
||||
" # Zero gradients from previous iteration\n",
|
||||
" optimizer.zero_grad()\n",
|
||||
" # Output from model\n",
|
||||
" output = voltron(train_x)\n",
|
||||
" # Calc loss and backprop gradients\n",
|
||||
" loss = -mll(output, train_y.log())\n",
|
||||
" loss.backward()\n",
|
||||
" # print(loss.item())\n",
|
||||
" optimizer.step()\n",
|
||||
" return voltron\n",
|
||||
"\n",
|
||||
"def predict_prices(test_x, voltron, nvol=10, npx=10):\n",
|
||||
" ntest = test_x.shape[0]\n",
|
||||
" vol_paths = torch.zeros(nvol, ntest)\n",
|
||||
" px_paths = torch.zeros(npx*nvol, ntest)\n",
|
||||
"\n",
|
||||
" voltron.vol_model.eval();\n",
|
||||
" voltron.eval();\n",
|
||||
"\n",
|
||||
" for vidx in range(nvol):\n",
|
||||
" vol_pred = voltron.vol_model(test_x).sample().exp()\n",
|
||||
" vol_paths[vidx, :] = vol_pred.detach()\n",
|
||||
"\n",
|
||||
" px_pred = voltron.GeneratePrediction(test_x, vol_pred, npx).exp()\n",
|
||||
" px_paths[vidx*npx:(vidx*npx+npx), :] = px_pred.detach().T\n",
|
||||
" return px_paths, vol_paths"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 7,
|
||||
"id": "9cc32071",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"torch.Size([10, 200])"
|
||||
]
|
||||
},
|
||||
"execution_count": 7,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"train_y.shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 8,
|
||||
"id": "e7c817ea",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"log_returns = torch.log(train_y[..., 1:]/train_y[..., :-1])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 9,
|
||||
"id": "426c01ed",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"tckrs = np.array(tckrs).reshape(-1)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 10,
|
||||
"id": "64ec0591",
|
||||
"metadata": {
|
||||
"collapsed": true,
|
||||
"jupyter": {
|
||||
"outputs_hidden": true
|
||||
},
|
||||
"tags": []
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"now running BAC\n",
|
||||
"Using gp-exp parameterization.\n",
|
||||
"Iter 1/500 - Loss: 13.602\n",
|
||||
"Iter 51/500 - Loss: -2.584\n",
|
||||
"Iter 101/500 - Loss: -2.706\n",
|
||||
"Iter 151/500 - Loss: -2.730\n",
|
||||
"Iter 201/500 - Loss: -2.738\n",
|
||||
"Iter 251/500 - Loss: -2.742\n",
|
||||
"-2.7423224449157715 tensor(-2.7423, grad_fn=<NegBackward>)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAaEAAAEFCAYAAABKJVg6AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAC2lklEQVR4nOz9eZxlV3Ufin/PcOd5qrmqqye1pJaEZiSBJgS2/LD9A+wYywrzEDAQHGKHBJzgxwvPj8T+PUPihNiQwYBtDDHG2BiwECAJCUlISOqWWmp1V9dct+rO83SG98eptWuffc+9Vd3q6kGc7+ejT6uq7j3j3nvttdZ3fZdkmqYJFy5cuHDh4jxAPt8X4MKFCxcufnbhGiEXLly4cHHe4BohFy5cuHBx3uAaIRcuXLhwcd7gGiEXLly4cHHeoJ7vC7iQ0W63cfToUWQyGSiKcr4vx4ULFy4uCui6jlwuhyuuuAJ+v3/oZ10jNARHjx7Fvffee74vw4ULFy4uSnz5y1/G9ddfP/QzrhEagkwmA8B6kGNjY+f5aly4cOHi4kA2m8W9997L1tBhcI3QEFAIbmxsDFNTU+f5aly4cOHi4sJO0hguMcGFCxcuXJw3uEbIhQsXLlycN7hGyIULFy5cnDe4RsiFCxcuXJw3uEboZQrDMNDtds/3Zbhw4cLFULjsuJchGo0GVlZWYJomEokERkZGzvcluXDhwoUjXE/oZYiNjQ1Qm6hSqYRer3eer8iFCxcunOEaoZcZer1eXxiu0+mcp6tx4cKFi+FwjdDLDI1Go+93bm7IhQsXFypcI3SRwjAMtNttGIZh+71rhFy4cHExwSUmXITQNA2Li4vo9XpQVRV79uyBqqowTRPNZrPv864RcuHCxYUK1xO6CFGr1RjZQNM0lMtlAHD0jADXCLlw4eLChWuELkLkcjnbz4VCAYBlhJyg6zp0Xd/163LhwoWL04VrhC5CyLLzaxvGgnO9IRcuXFyIcI3QRQhJkvp+Z5qma4RcuHBx0cE1QhchnPI+nU5nqKFxjZALFy4uRLhG6CKDpmmORqherzv+nuCqJrhw4eJChGuEdgm6rkPTtNP6/E68lUHHdKoP2sn3XLhw4eJ8wq0T2gVUKhVks1lIkoR0Oo1kMjn0861WCysrK9B1HeFwGJOTkwM/O8ijaTQatla6wWDQVjPksuNcuHBxIcL1hHYBxWIRgEUWKBQKQ8Nk9HkyEvV6Ha1Wa+Bnd+oJBQIB28+uEXLhwsWFCNcI7QJIwRqwSATb5WPq9brt52q1OvCzTscipQTe2AWDQdtnDMOwXZcLFy5cXAhwjdAuwOv12n4+m8w0JyPUarWg6zozNJIkwe/32+qJTNPc1iNz4cKFi3MN1wjtAl6qERrmsfBGSNd1FItFlEolAMD6+jrW1tYgSRJkWbbliOjzLly4cHEhwSUm7AK2M0IbGxsol8vw+XwYHx/v+/4wI8TnhIrFoqOB6/V60HUdiqL0GS0XLly4uJDgGqFdgMfjsf3MG4p2u808l3a73acDBww2FoZhsL+ZpulogAKBABRFQavVcj0hFy5cXPBww3G7ANET4r2RSqVi+5tISgAGM+B4I+JkUDweD6LRKADL8LlGyIULFxc6XCO0C1BV1UYK4AtXd8JQG2SE+N+LBkVRFIyMjDDD0263XSPkwoWLCx6uEdoFSJLUF5I7HdkcXdcdjRVvRERD5fP5bD+7RsiFCxcXA1wjtEsYRE7Yaa0OGYxer4d6vd4nA6TrOlR1K6XH/z99T1Tbdo2QCxcuLjS4xIRdgmiEstksCoUCKpVKXw2PEzRNg6ZpWFpagmEYUFUVkUiE/b3X68Hj8TDDJHo9dAwerhFy4cLFhQbXE9oliEYIsAxHrVbbVmwU2KoBogJTTdNspAZd120hP9ETovOJx3ThwoWLCwmuEdoliDkagmEYQ2V5CJqmoVar2X7HGyFN06CqKgu5OXlCotHRdR21Wg2FQsFmoHq9nquy7cKFi/MCNxy3S/B6vZAkyZYDOh3pnGw2i3K5jFgsxgwNHYvkeUgVgQpTt0O1WmW5qVKphH379iGfz6NUKkGSJIyPj9tCfi5cuHCx23A9oV2CJEl9ITneM3EiKHQ6HRQKBZRKJfR6PTQaDeTzeVtIjv9XlmVIkuQYiut0OqjVarbzlMtl9rMo+UOK3y5cuHBxLuF6QrsIn8+HTqfDfua9IMMwbN6LYRgoFArMSND3ut0ucrkcUqkU+x1PRjBNs88I1et1VCoVqKoKVVURDocZ0YE/LxkgQqfTgWmaKJfLaDQaiEQiiMViZ+txuHDhwkUfXCO0ixjmCYkhNDIATp/VNA0bGxtMtqdWqzGRUqCflEC5I8MwUK/XUa/XoaoqKpUKC+/5/X5IktQXxqvVatjY2ACw1SgvHA6/lMfgwoULFwPhhuN2ESI5gfeETpc+rWkaWq0WCoUC2u02O5aqqvD7/exzYi+jRqPBPttut7GxsYFKpYJ8Po92u913HlFWaGVlZeh1uXDhwsVLgWuEdhGiERqm/caz1WRZRiqVsnVHNQwDnU4HlUoFlUoF1WoV6+vrSCQStvOIhk7XdabA0Ov12Hl0XUe5XO4jSvAtwQk7oZS7cOHCxZngtMJxDz/8MD73uc/hhRdeQK/Xw+HDh/He974Xt956646Psb6+jj/+4z/Gj370I+RyOYyPj+OXf/mX8Z73vMextmZtbQ2f+9zn8OCDD2JjYwOBQABXXHEF3vGOd+C2225zPMe3vvUt/K//9b9w4sQJKIqCa665Bh/4wAdw1VVXnc7tvmSIYbJhnpCmafB6vUgmkyzU5vV60W63YZomY9ZR629ZllkDu1gsBk3TEAgEWCgNsAxbq9VCt9uFLMus8Z3TeYehWCwiFAqd0TNw4cKFi2HYsSf013/913jHO96Bn/70p7jqqqtwzTXX4Kc//Sne/e534ytf+cqOjpHNZvFrv/Zr+MpXvoJoNIo77rgDjUYDn/3sZ/Gud72rr7jyxIkTeMMb3oC//Mu/BADcdttt2Lt3Lx555BG85z3vwRe+8IW+c/yn//Sf8C/+xb/Aiy++iFe+8pW45JJL8MMf/hD33HMPfvjDH+70ds8KJEliqtbAcCPU6/Xg9XqhKArL9ciyzLwh0XPi8zqyLGNychKyLLNwmmmazGDpus7+XzRCO9G0a7VabmtwFy5c7Ap25AltbGzgE5/4BCKRCP78z/8cl1xyCQDgmWeewTve8Q586lOfwh133IHR0dGhx/m93/s9ZLNZfPjDH8Zv/uZvArDCPx/4wAfw8MMP44tf/CLe+c53ss9//OMfR7lcxrve9S78y3/5L1kS/ZFHHsF73/te/OEf/iFuv/12HDhwAABw9OhR/Of//J8xOTmJv/iLv2DX84Mf/AAf+MAH8LGPfQz33XefLcy12xgdHWUeDd87iDcqhmEgEAg41uiEQiHU63VmRAjkKQFbuR4Kr2mahmq1yjwlwzBY/kc0ZjspUjVNs0+rzoULFy7OBnbkCX3pS19Ct9vF29/+dmaAAOCqq67Cu9/9bnQ6nW29obm5OfzgBz/AzMwM3ve+97HfB4NBfOpTn4KiKPjSl77Efn/q1Ck89dRTmJqashkgALj55ptxzz33QNd1fPvb32a//x//438AAD70oQ/ZDOIdd9yBN77xjcjn8/jWt761k1s+a6D8zujoaF+X07GxMUxOTmJychKJRMJRT44PyfFQFIXlgjqdDqNbS5Jk81wobEce0Jl4QnS9Lly4cHG2sSMj9OCDDwIAXvva1/b97XWvex0A4IEHHhh6jIceegimaeLOO+/sW2wnJiZw+eWXY2VlBSdOnABg1bBcffXVuO222xzVAGZnZwHAlgN58MEHIUkSXvOa15zxde4WarWa7b5VVYXX62U1PIPglLPxeDysgyp9hoyELMt9hoU3Qnxfo2aziVKp1OdlDboOFy5cuDjb2Da+YpomTpw4AVmWsW/fvr6/z87OQpZlnDhxgoV/nEDG5eDBg45/37dvH44cOYLjx4/jwIEDuPbaa4d6V0eOHAEA5vEQ9XhsbMyxwJKu/fjx40PudvdQqVTg9XrZgh8MBlGr1RAIBGwFrSI0TcPIyAiazSZ6vR4URUEkErHRsqnNA2nJkRdkmqbNQJEx0nUd9XodmqZBlmUUi0WMjY0Nlf5xPSEXLlzsBrb1hCqVCrrdLuLxuCOLSlVVJBIJtFqtoVRe8lhGRkYc/57JZAAA+Xx+24uem5vDN7/5TUiSxDwcyrfQcQYd/3xI07TbbXQ6HQSDQfY7MkLkkQxCJBJBIBBAOp1mzfIURUEqlWKf0XWdeT+8AWq1Wuj1euz8xLKrVCp9nV6HGULA9YRcuHCxO9jWE6Kd+7BkPu3KG43GwOp6Og6/g3c6xrAFGbD0z/75P//n6PV6+JVf+RWWo9ruOil/st3xdwN0bT6fD5IkwefzQVEUpoQwaIEn9luv10M6nQZg5YgmJycRCoWwvr4OwDIQohEiD4j05Uiuh4gL5PXw9O9hcI2QCxcudgPbGqHtmq8BO+sWSscZFK7jd/CDUCwW8e53vxsvvvgiDh8+jH/7b/8t+9ug457JtZ5tkJchSRICgYCNZVYul22fjUajGBkZQafTgd/vx9zcHPtuJpPB3r174fV6bZ6L6AmZpsn+TkaIfk//8XD6nQg3HOfChYvdwLZGiEJIw8I19Ldh3hIdx0kqZifHWFhYwHvf+17Mz8/j8OHD+MIXvmD7LBVTDrrOnVzjboHaJwDWcxi24EejUSiKgkAggGq1ina7zZrXKYrCQqJ8qwheFYF+5tlxkiTZfmcYRl/u6Ew8oU6ng2w2C03TkE6nXbFTFy5cnDa2NULhcBjBYBClUoklv3lomoZSqQSfz2crzBRBuaBBOR/K6TjljJ566im8733vQ6lUwo033oj/+l//a1/Yb6fHH5Qz2k3wRsjn8yEYDDqGBRVFYcb6xIkTWFpaAgDLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"0.7800340056419373\n",
|
||||
"0.5812243819236755\n",
|
||||
"0.36129868030548096\n",
|
||||
"0.12041786313056946\n",
|
||||
"-0.12965255975723267\n",
|
||||
"-0.3816063702106476\n",
|
||||
"-0.6296417713165283\n",
|
||||
"-0.8692443370819092\n",
|
||||
"-1.0984598398208618\n",
|
||||
"-1.3174093961715698\n",
|
||||
"now running GS\n",
|
||||
"Using gp-exp parameterization.\n",
|
||||
"Iter 1/500 - Loss: 14.026\n",
|
||||
"Iter 51/500 - Loss: -2.744\n",
|
||||
"Iter 101/500 - Loss: -2.826\n",
|
||||
"Iter 151/500 - Loss: -2.837\n",
|
||||
"-2.839012861251831 tensor(-2.8390, grad_fn=<NegBackward>)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAaEAAAEFCAYAAABKJVg6AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAACnuElEQVR4nOy9eZxcVZk+/ty6te9LV+/d6XSSTkhCgCC7bAKjjooDjmBgZBFQGUB/ozPyHWRkhtHfDDP6G0H5DiOgo6AzCIOiDoqyL8GwhCUJkKTTS3qr7qqufd/u74/Ke/rcW7e6K0l3EuA8n48fya3qW3c557znfd/nfV5JURQFAgICAgIChwGGw30BAgICAgLvXwgjJCAgICBw2CCMkICAgIDAYYMwQgICAgIChw3CCAkICAgIHDYYD/cFHMnI5/PYvn07gsEgZFk+3JcjICAg8K5ApVJBOBzG+vXrYbVa5/2uMELzYPv27bj00ksP92UICAgIvCvx05/+FB/4wAfm/Y4wQvMgGAwCqD3I9vb2w3w1AgICAu8OhEIhXHrppWwNnQ/7ZYQ2b96Mu+66Czt37kSpVMK6devw+c9/HqeffnrT55iensadd96JF154AeFwGB0dHTj//PNxzTXXwGw2131/amoKd911F5577jnMzMzAZrNh/fr1uPLKK3HGGWfUff/OO+/EHXfc0fD3v/nNb+LTn/50U9dKIbj29nZ0d3c3eYcCAgICAgCaSmM0bYQefvhh/O3f/i3MZjNOPvlkVKtVbNmyBVdffTVuvfVWXHzxxQueIxQK4eKLL0YoFMLatWuxbt06bN26FXfccQf++Mc/4oc//CFMJhP7/uDgIC699FLE43F0dXXhjDPOQCQSwYsvvojNmzfja1/7Gq666irVb7z99tsAgPPOO083Ftnb29vsLQsICAgILDWUJjA9Pa2sX79eOf7445WdO3ey42+88YayceNG5eijj1ZCodCC5/nCF76gDAwMKHfeeSc7lslklCuuuEIZGBhQ7r33XtX3L7roImVgYEC57bbblHK5zI5v3rxZWb9+vXLUUUcpu3fvVv3N2WefrRx99NFKqVRq5tbmxdjYmDIwMKCMjY0d9LkEBAQE3i/Yn7WzKYr2/fffj2KxiCuuuAIDAwPs+IYNG3D11VejUCjggQcemPccQ0NDePrpp9Hb24svfvGL7Ljdbse3vvUtyLKM+++/nx0fHh7G66+/ju7ubnz1q19VuXWnnHIKNm3ahEqlgt/97nfseDKZxMTEBNasWQOjUaS7BAQEBI50NGWEnnvuOQDAueeeW/fZeeedBwB49tln5z3H888/D0VRcPbZZ8NgUP9sZ2cn1q5di4mJCQwODgIAYrEYjj32WJxxxhm6ccW+vj4AwMzMDDv21ltvAQDWrVvXzG0JCAgICBxmLOguKIqCwcFBGAwG9Pf3133e19cHg8GAwcFBKIoCSZJ0z0PGZdWqVbqf9/f3Y9u2bdi1axdWrlyJjRs3zutdbdu2DQDQ1tbGjlE+yGaz4cYbb8RLL72E2dlZ9PX14dOf/jQuvfTSOgN4pKBaraJcLuuSMwQEBATeq1hwRU4kEigWi/B6vboLpNFohM/nQy6XQyaTaXge8lhaW1t1PycqXyQSWfCih4aG8Otf/xqSJDFPDJjzhO69915s3rwZ69atw1FHHYWhoSF885vfxJe//GVUq9UFz3+oUSgUsGfPHgwPD2NiYuJwX46AgIDAIcOCRiiXywGoeReNQCy0+YwQnadR9Swdz2az815PPB7Hl770JZRKJVx44YWqHBV5QpdddhmefPJJfP/738cDDzyAhx56CJ2dnfj973+Pn/70p/Oe/3AgFAox45hOp5HP5w/zFQkICAgcGixohJoJXylN9MWj8zQK19E55jtXNBrF5z73OezevRvr1q3D3/3d36k+//nPf45f//rXuOmmm1RU7zVr1uDrX/86ABxxRqhUKtUZnUKhcJiuRkBAQODQYkELY7fbAcy/MNJn83lLdJ5Gu/yFzjE6OopNmzZhx44dWLduHe69996679rtdgwMDOgaujPPPBOyLGN4eHhBb+tQIpVK1R07EkOGAgICAkuBBY2Q0+mE3W5HLBZDuVyu+7xcLiMWi8FiscDtdjc8D+WCGuV8wuGw6ns8Xn/9dVx88cUYGRnBiSeeiJ/85Cfw+XwLXboKJpMJHo8HQGNDeDiQTCbrjlUqlcNwJQICAgKHHgsaIUmSsHLlSlQqFYyMjNR9Pjw8jGq1qsrN6IFYccSS02LPnj0AUHee5557DpdffjlisRg+9rGP4d5774XT6az7+4mJCdx00024+eabdc+fyWQQjUZhtVqZMTrcKBQKuh6mMEICAgLvFzTFVyZtuMcff7zuMzp25plnNnWOJ598si7cNDk5ibfffhtdXV1YuXIlO/7GG2/g+uuvRz6fx+WXX47vfOc7DSnMTqcTjzzyCB588EGMjo7Wff7II48AqBW6HiltGfRCcQB0PU76/vT09LwEEAEBAYF3E5oyQhdeeCEsFgvuvvtubN++nR3ftm0b7rnnHlitVlxyySXs+OTkJPbs2YNoNMqO9fT04PTTT8fw8DBuv/12djybzeLmm29GpVLBlVdeyY7n83l85StfQT6fx8UXX4ybbrqpIakBADweDz7ykY8AAL7+9a+rwlzbt2/H7bffDoPBoFJrWGqEw2Hs2bMHU1NTuoQLvVAcoO8JRaNRTE5OIh6PY3x8/IgKKQoICAgcKJrStunu7saNN96IW2+9FZ/5zGdw8sknQ1EUbNmyBeVyGbfddhsCgQD7PhWKXn/99bjhhhvY8VtuuQWbNm3CXXfdhSeffBLLly/H1q1bEQ6HccYZZ2DTpk3su7/4xS8wPj4OoEbL/uu//mvdazvhhBOYeOrXv/517NixAy+//DLOO+88HHfccchms3jllVdQrVZx00034dhjj93vh3QgoPAfUDM2ZrNZ9YxyuRxKpZLu32qNUCqVYjkzQjweF+0lBAQE3vVoWmDt0ksvRWdnJ+655x68+uqrMJvN2LhxI6699lqccsopTZ2jp6cHDz74IO644w48++yzGB0dRU9PDy677DJcfvnlKr23l156if33Y4891vgGjEZmhPx+Px566CH84Ac/wGOPPYbnn38edrsdp512Gq6++mqcdNJJzd7uQUMbaotEIiojFIvFGv4tb4QURakzQECtnmg+hQoBAQGBdwMkpZkin/cpxsfHcc455+CJJ57Y735CY2NjdVTwvr4+WCwWxGIxleadHohqXiwWMTw8rPud7u5uOByO/bouAQEBgaXG/qydR6aQ2nsAeuSC0dFRZLPZOs/GbDbXFQWTN9SIpAA0JjYICAgIvFsgjNASQFEU3XyPoigYGxtTkRQMBgO6urrqWk+QEWqUNwJqITlB5xYQEHg3QxihJUCxWGxKyggAfD4fzGZzHW28GSNUqVQQj8cP+DoFBAQEDjeEEVoCFIvFpr9L0kNaT4jCcNpwnPZ7kUgEU1NT8xorAQEBgSMVwggtAbRGaL4ur6Qe3qwnFAwG6/JHyWQSU1NTB3y9AgICAocLwggtAbRGKBAI6EoNmUwmZnwaGSGtJ2SxWHR183K53LwkBgEBAYEjEcIILQG0RshsNuuqg/O9lfSMkB7BwWQywe/361KzRUhOQEDg3QZhhJYAWo/EbDazVhY85jNCxHzjCQ6yLMNgMMBgMKC7u7vOsAkjJCAg8G6DMEJLAN642Gw2GI1GWCyWuu/xx/SICXv37lUd035HGCEBAYF3O5qW7RFoHu3t7YhGo1AUBX6/H0CtJYbX62WUapPJpPKO+E6wBL1Q3Hz/FkZIQEDg3QZhhJYAsiwjGAzWHQ8GgzAajSiVSvD5fCrdN2q6l0gkGp5XGCEBAYH3GoQROoQwGAwqEVMt2tvbUS6XG/YLEkZIQEDgvQaREzrC4PV6G37WjBESerQCAgLvJggjdIRBLzdE0JIbJEmqIysIb0hAQODdBGGEjjA0MkIGg0FXeUHb7ryRcKoQOhUQEDgSIXJCRxgMBgNMJlOdMbFYLLoN7BbKC+XzeUxMTKBSqcDr9aK1tXXxL1pAQEDgACE8oSMQet6Q1uNp9N1CoaD6dyQSQblchqIoiMVi+yWuKiAgILDUEEboCISewdErdtU7ns1moSgKkskkotFoHdOuEfNOQEBA4HBAhOOOQOh5Qo2MkN1uhyRJjBVXLBYRCoWQTCZ1vy+ICwICAkcShCd0BGJ/wnEGg6FOvqeRAQLqw3UCAgIChxPCCB2B0DM48/Uk0lPUbgRhhAQEBI4kCCN0BMJisajCb3r9g3jsjxGqVCqi75CAgMARA5ETOkLR09ODeDwOWZbh8Xjm/a7FYoHRaGzauBQKhXk9KwEBAYFDBbESHaGQZXlenTktzGbzfhmh/fGeBAQEBJYKIhz3HkEj4oIe0um00JgTEBA4IiCM0HsE88n9tLS0qI7lcjnMzs4eissSEBAQmBfCCL1H0MgTam1tRSAQqKNxR6NRUTMkICBw2CGM0HsEjTwhIiB0dnaqyAiKoiCbzR6SaxMQEBBoBGGE3iNoZITouNForGPZiZohAQGBww1hhN4jaNTqgT9mtVpVn+Xz+SW/LgEBAYH5IIzQewgGQ/3r5I9p9ecKhYJgyQkICBxWCCP0HoJevyEeJpMJsiyzf1erVUFOEBAQOKwQRug9hGZUEERITkBA4EiCMELvIfj9ftW/9RQX9EJyAgICAocLwgi9h2C32+HxeCBJEmw2G7xeb913hCckICBwJEFox73H0N7ejvb29oafaz0h0e5bQEDgcEJ4Qu8zaOuJyuWyYMgJCAgcNggj9D6DJEl1hkgw5AQEBA4XhBF6H0JrhERITkBA4HBBGKH3IYQnJCAgcKRAGKH3IYQREhAQOFIgjND7ENq2D8IICQgIHC4II/Q+hPCEBAQEjhQII/Q+hB4xQdC0BQQEDgeEEXoLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"0.7835856080055237\n",
|
||||
"0.5836868286132812\n",
|
||||
"0.36538517475128174\n",
|
||||
"0.12609244883060455\n",
|
||||
"-0.12323341518640518\n",
|
||||
"-0.37310081720352173\n",
|
||||
"-0.6181354522705078\n",
|
||||
"-0.8563370704650879\n",
|
||||
"-1.0892924070358276\n",
|
||||
"-1.3147377967834473\n",
|
||||
"now running JPM\n",
|
||||
"Using gp-exp parameterization.\n",
|
||||
"Iter 1/500 - Loss: 13.751\n",
|
||||
"Iter 51/500 - Loss: -2.817\n",
|
||||
"Iter 101/500 - Loss: -2.988\n",
|
||||
"Iter 151/500 - Loss: -3.023\n",
|
||||
"Iter 201/500 - Loss: -3.036\n",
|
||||
"Iter 251/500 - Loss: -3.044\n",
|
||||
"Iter 301/500 - Loss: -3.048\n",
|
||||
"-3.048539161682129 tensor(-3.0485, grad_fn=<NegBackward>)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAa0AAAEFCAYAAABQGbi0AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAC/FUlEQVR4nOz9eZhldXUvjH/2cOb5VJ1Tp8buri66oRtbRgENAiJRk6Diq1EkTIoaQ4yvv8TLK+F9yEW5ufokPyNerhgJuQnEewlEJBoHhpZZ5qmbhu6u6uquuc6pM8/D3vv9Y9f61nfvs09V9Vzd2Z/n4aHrDHs+a33XWp/1WYKmaRps2LBhw4aNEwDi8T4AGzZs2LBhY7WwnZYNGzZs2DhhYDstGzZs2LBxwsB2WjZs2LBh44SB7bRs2LBhw8YJA/l4H8DJhFqthp07dyIWi0GSpON9ODZs2LBxQkBRFKRSKZx++ulwu93LftZ2WkcQO3fuxFVXXXW8D8OGDRs2Tkj8y7/8C84555xlP2M7rSOIWCwGQL/wiUTiOB+NDRs2bJwYmJubw1VXXcVs6HKwndYRBKUEE4kEBgYGjvPR2LBhw8aJhdWUVWwihg0bNmzYOGFgOy0bNmzYsHHCwHZaNmzYsGHjhIHttGzYsGHDxgkD22nZAACoqop6vQ5b9N+GDRtrGQfFHnzuuedw1113Yffu3Wg2m9i6dSu++MUv4sILL1z1Nubn53HnnXfi2WefRSqVQm9vLz760Y/iC1/4ApxO57Lf1TQN1157LWZnZ/Hoo49afubOO+/EHXfc0XEb3/rWt/CpT33K8NqOHTtw5513YseOHahUKhgZGcE111yDyy+/fNXndSKjUqlgdnYWrVYLLpcLQ0NDEEV7PWPDho21h1U7rZ/85Cf4xje+AafTifPPPx+qquKFF17ADTfcgNtuuw2f/vSnV9zG3NwcPv3pT2Nubg5btmzB1q1b8eqrr+KOO+7A888/j3vuuQcOh6Pj97/zne/ghRdewNDQUMfPvP322wCAyy67zLKz2vzdZ599Fl/60pegqirOPfdceDwe/Pa3v8Vf/MVfYHR0FF/72tdWPK8TGalUCnv37oUoiggGg6jX6yiVSggGg8f70GzYsGGjDatyWslkErfeeisCgQB+/OMfY9OmTQCAN998E9dffz1uv/12XHzxxejp6Vl2O3/1V3+Fubk5fPWrX8Wf/MmfANBX+TfeeCOee+453Hvvvfjc5z7X9r16vY7bbrsNDz744IrHumvXLrhcLvzd3/0dZHn506vVavj6178OALjnnntw/vnnAwAmJiZw9dVX46677sJll12G008/fcX9nohQVRXvvPMOms0mAF1Kpbu7m/1tw4YNG2sNq8oB3XfffWg0GrjuuuuYwwKAbdu24YYbbkC9Xsf999+/7Db27duHJ554AkNDQ/jjP/5j9rrX68Xtt98OSZJw3333tX3vqaeewhVXXIEHH3wQg4ODy+6jUChgenoap5566ooOCwAefvhhpNNpXH755cxhAXo09ud//ucAgHvvvXfF7ZyoKJVKBgdVr9fRbDbRarWO41HZsGHDRmesymk9/fTTAIAPfvCDbe9ddtllAHTnshyeeeYZaJqGSy65pK1e0tfXhy1btmB6ehqjo6OG977whS9gfHwcV199NX74wx8uu49du3YBALZu3br8CS2CzuvSSy9te+8DH/gAJEla8bxOZNRqtbbXyuWy7bRs2LCxZrGi09I0DaOjoxBFEcPDw23vr1+/HqIoYnR0dFnmGTmjU045xfJ92vaePXsMr3/oQx/CT3/6U9xyyy1wuVzLHivVszweD2666SZccskl2LZtGz760Y/i3nvvhaqqhs/v3bsXAAzRI8Hv9yMejyOTyWBhYWHZ/R5LlEoljI2NYWxsDKVS6bC21Wg02l6rVCqWr9uwYcPGWsCKTiufz6PRaCAcDluy+2RZRiQSQbVaRblc7ridZDIJAIjH45bvk1Ci2UHccccd2Lx580qHCWAp0vqHf/gHPPfcc9i6dStOO+007Nu3D9/61rfw1a9+1eC4UqmUYd+rPabjiWQyiVarhVarhZmZGaRSKSSTyUOqQ1lFWpqmHbYztGHDho2jhRULP9VqFYAevXQCsfTK5TL8fv+y2+k0K4Ver1QqKx1SR1Ckdc011+C//Jf/wpiI77zzDr785S/jkUcewb/8y7/g6quvPmbHdCShqqrBOWmahkwmA0B3QMuxKq1Qr9c7vq5pGgRBOPSDtWHDho2jgBUjrdX066ymIZW208kQ0jYOp7n1X//1X/Gzn/0MN998s4E6f+qpp+Iv//IvAehjQwiSJEEQhKN6TEcS5vQmj2q1etC1qE5pQFVVoSjKQW3Lhg0bNo4FVvRIXq8XQOdVOf/ectEYbccqJbXabawEr9eLTZs2WTqhiy66CJIkYXx8nEVOHo8HmqYtG3Hwx368sZIjOdhaVKfzVhTFJmPYsGFjTWJFp+X3++H1epHNZi0NWavVQjabhcvlWrYhlWpZnepDVF/qVPM6XDgcDoRCIQBLjpP2RfvudEyrGUx2LLBcpAWgra6Vz+exf/9+zM7Otjk8TdM6OjnbadmwYWOtYkWnJQgCRkZGoCgK9u/f3/b++Pg4VFW1ZODxINagmdJOGBsbA2DN5FsNpqencfPNN+OWW26xfL9cLiOTycDtdjPnRcdE++ZRKpWQTCYRjUbR3d19SMd0pHEwkVaj0cDc3Bzq9ToKhQKrfRFarVZHJ2g7LRs2bKxVrKpPi7QFH3vssbb36LWLLrpoVdvYvn17m7GcmZnB22+/jf7+foyMjKzmkNrg9/vx8MMP44EHHsCBAwfa3n/44YcBABdccAGbjrnceW3fvh2Koqx4XscSBxNp5fN5w3uZTAaqqrL6XKPR6Lg9VVVtp2XDho01iVU5rU984hNwuVz40Y9+hJ07d7LXd+zYgbvvvhtutxuf/exn2eszMzMYGxszrO4HBwdx4YUXYnx8HN/73vfY65VKBbfccgsURcH1119/yCcSCoXw4Q9/GADwl3/5lygUCuy9nTt34nvf+x5EUTSocXzoQx9CV1cXHnroITz55JPs9cnJSfzt3/4tBEHAddddd8jHdKSxktMyR1pm7N27F/v27UOtVutYWwSWTx3asGHDxvHEqrQHBwYGcNNNN+G2227DZz7zGZx//vnQNA0vvPACWq0Wvv3tb6Orq4t9/qabbsKLL76IP/3TP8VXvvIV9vqtt96KK6+8EnfddRe2b9+ODRs24NVXX0UqlcL73/9+XHnllYd1Mn/5l3+Jt956Cy+99BIuu+wynHnmmahUKnj55ZehqipuvvlmnHHGGezzfr8f3/zmN/Fnf/Zn+NKXvoRzzz0XPp8Pzz//PKrVKr72ta/h1FNPPaxjOpJYKT3IR1qdHFyr1WpLFVphOadmw4YNG8cLq1Z5v+qqq9DX14e7774br7zyCpxOJ8466yx8+ctfxgUXXLCqbQwODuKBBx7AHXfcgaeeegoHDhzA4OAgrrnmGlx77bWr0gtcDtFoFA8++CD+/u//Hr/+9a/xzDPPwOv14n3vex9uuOEGnHfeeW3fufTSS3HvvffizjvvxBtvvAFN07B582Zcd911+MhHPnJYx3OkYXZEsVgM6XSavU5pPVmWl42UisVi27UWRdGw/VKphOnpafT29tpjSmzYsLFmIGhrpQnpJMDU1BQuvfRSPP744xgYGDji25+bmzPUqnp6epDL5QzU9UgkAlVV22paPDRNQ61WM0RcHo+HNVvTdrxeL3p7e+0xJTZs2DiqOBjbeXihjY1jCnOkJUkSnE6nwWlls9kVt1OtVg2pREmSEA6HUa/XDVEbsHx/ng0bNmwca9h5nxMI5pqWKIrLDs3shHK5bNAX9Hq9cLlciEQi7DVyWjaL0IYNG2sJttM6gWAVaa2kfG9Gq9Vqq3d5vV643W643W7Wk0YO0h4IacOGjbUE22mdQDA7LVEUEQgEmPSVIAiWSvw8zOK/TqcTPp8PgUCAbRNYclp2pGXDho21BLumdQLBKj0oCAIGBwfRarVY0/T4+Lils9E0rc1p+Xw+DA0NseiLtsGnB23Fdxs2bKwV2JHWGsNyZE6r9CCgR1gOhwOiKEIUxY409UajYXB8giCwzxIFnrZBn9M0zVZ8t2HDxpqBHWmtIWSzWaTTaUiShL6+PkO9ipdgApaiLCt4vV6MjIxAVVXMzMyw6IqPshRFgc/nYzPDaEyLpmmQJAnNZpNFWNT7ZcOGDRvHG3aktUbQaDSQTCahKAoajQbm5+cN71vVs5aDIAiQJAmRSIQ5N4qYSqUScrkcisUiSwvS5wHrFKENGzZsrAXYTmuNIJfLGf6uVqsGlp85RUeOZSX4/X6sX78e0WiUDXek3itBEFAsFtu2aXZaNoPQhg0bawW201oDUFXVIPBL4F872EiLh9PpRCAQgKqqBgckiiJLAwJLzoq2bUdaNmzYWGuwndYaQLFYtCQ78E7Lijl4MJBlua0uJssyq1nR38CS87Jp7zZs2FhrsJ3WGgCv+cej2Wyy9zoxB1cLIlnQdgRBYE7KTHe304M2bNhYq7Cd1hrAcg3BNCLkcCMtckC0HUEQ2DbIadmRlg0bNtY6bB7zGgBp/lEfFa8LSA7lcCMtckD0f54yv1KkZTcY27BhY63AdlprAIIgIBqNAtDp6FZOq1wuG75zKE7LnB4ksV1zpEUOjW8wVlX1oPdpw4YNG0cadnpwjcGcKmw0Gmg0Gm11L5/Pd1DbVRQFXq+X/S2KItuGOdKif/PRnZ0itGHDxlqAHWmtMTgcDkaaAHRnYe7h8ng8KwrjmqEoCtxuN7xeLxRFQTQaZTWtVqsFRVGgqqphHhdfR7OlnGzYsLEWYDutNQZK2/GNxebBjqFQ6KC3S43FpAhvHmnSaDQwNzfHpJ7q9TpcLherZdlOy4YNG2sBdnpwDWK5KEoQBPj9/oPLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"0.7912479639053345\n",
|
||||
"0.5877759456634521\n",
|
||||
"0.3701268136501312\n",
|
||||
"0.1320444494485855\n",
|
||||
"-0.11983421444892883\n",
|
||||
"-0.3745497763156891\n",
|
||||
"-0.624917209148407\n",
|
||||
"-0.8667067885398865\n",
|
||||
"-1.0983527898788452\n",
|
||||
"-1.321250557899475\n",
|
||||
"now running MS\n",
|
||||
"Using gp-exp parameterization.\n",
|
||||
"Iter 1/500 - Loss: 14.385\n",
|
||||
"Iter 51/500 - Loss: -2.662\n",
|
||||
"Iter 101/500 - Loss: -2.755\n",
|
||||
"Iter 151/500 - Loss: -2.767\n",
|
||||
"-2.768977165222168 tensor(-2.7690, grad_fn=<NegBackward>)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAaEAAAEFCAYAAABKJVg6AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAC16ElEQVR4nOz9d5hlZ3UljK9zc06Vc3eruyXUWCARJCEQEiBjG2Mb+GwhMAKBcCCYMbZ/zIBt+DFmxvbAN2MMM9gGe2xEkEm2sbGxCUpIKCO1Ynd1V3VXrro5p3PO98ft9dZ73nvurepWVwfprOfRo64bzj1x73fvvfbammmaJhw4cODAgYOzANfZ3gEHDhw4cPDcheOEHDhw4MDBWYPjhBw4cODAwVmD44QcOHDgwMFZg+OEHDhw4MDBWYPnbO/AuYx6vY7HHnsMQ0NDcLvdZ3t3HDhw4OC8gK7r2NjYwPOf/3wEAoG+n3WcUB889thjeOtb33q2d8OBAwcOzkt86Utfwotf/OK+n3GcUB8MDQ0B6JzI0dHRs7w3Dhw4cHB+YHV1FW9961uFDe0Hxwn1AVNwo6OjmJycPMt748CBAwfnF7ZTxnCICQ4cOHDg4KzBcUIOHDhw4OCswXFCDhw4cODgrMFxQg4cOHDg4KzBcULPMhiGgUajAV3Xz/auOHDgwMGWcNhxzxK0Wi2sr6+jUqnANE14PB5MTk7C7/ef7V1z4MCBg55wIqFnCdbX11Eul8HxUO12G4VC4SzvlQMHDhz0h+OEniWoVqtdr7Xb7bOwJw4cOHCwfThO6FkAXddhGEbX63avOXDgwMG5BMcJPQvQK+JxyAkOHDg41+E4oWcBejkhJxJy4MDBuQ7HCT0L4DghBw4cnK9wnNCzAI4TcuDAwfkKxwk9C9BqtWxfNwxDULYdOHDg4FyE44SeBehHxXaiIQcOHJzLcJzQswCOE3LgwMH5CscJPQvgOCEHDhycr3Cc0HkOwzD69gM5TsiBAwfnMhwntENoNBqo1+s7/jtbSfM4TsiBAwfnMhwV7R1ANpvFxsYGACCVSmFoaGjHfstxQg4cODif4URCOwBZvTqXy+2oI+hFzyYcJ+TAgYNzGY4T2mGYpolGo7Fj23eckAMHDs5nOE5oB6AOkttJJ6Ru2+v1Wv52REwdOHBwLsNxQjuAk3FC7XYbuVwO5XL5lH6r2Wxa/g4Gg5a/nUjIgQMH5zIcYsIOQHVCqqMgDMPA8ePHRUptZGQEiUQCAFAul5FOp+F2uzEyMgKfz2f5bqVSQalU6tp2IBBAsVi0/IYDBw4cnKtwIqEdwFaRkGmaqFQqyGazlpoOCQ2GYWB1dRWNRgPVahULCwsWFly9Xsfi4mLX+G6fzwePx7qucJyQAwcOzmU4kdAOwOPxwOVyCQeg6zra7TY8Hg9M08SxY8dsU3TsK2o2m5ZaTrvdxtLSEqanp6FpGkqlku3v+nw+uFzWdYXjhBw4cHAuw4mEdgCapvWMhiqVSt8akWEYWF5exvr6OiqVini9Xq+jWq2KbdjB7/c7TsiBAwfnFRwntENQazh0PHQkvZBOp5HP59FqtZDP51EsFpHP59FoNMQ2NE3r+ZuOE3LgwMH5BCcdt0NQI6FSqYRUKtWTpEDkcjlLKo6pt0qlAq/Xi2Qy2bM3yC4ScijaDhw4OJfhREI7hEgkYvm7Xq+L/7ZCL8dRKBSwsrJi+77b7XYiIQcOHJx3OKlI6O6778bnPvc5PP3002i1Wjhw4AB+7dd+Da94xSu2vY21tTV89rOfxY9+9CNsbGxgbGwMv/ALv4B3v/vdXSksAFhZWcHnPvc53HnnnVhfX0cwGMTzn/983HTTTbj66qttf+M73/kO/vZv/xazs7Nwu9249NJL8d73vheXXHLJyRzuM4LX60U4HLbUb9Lp9LYik16OwzAM5PN5uN1uy+sulwtDQ0PQNM3WCZmm2TOF58CBAwdnE9uOhL75zW/ipptuwsMPP4xLLrkEl156KR5++GHcfPPNuPXWW7e1jdXVVfzKr/wKbr31VsRiMVxzzTWoVCr49Kc/jXe9611daabZ2Vn80i/9Er761a8CAK6++mrs3r0b99xzD9797nfjC1/4Qtdv/Pmf/zl++7d/G4cPH8bll1+O/fv34/bbb8cNN9yA22+/fbuHe1rAnh8im81ua9x2L0el6zqazabFSYXDYVxwwQWIx+MAYOuInBHfDhw4OFexrUhofX0dH/3oRxGNRvHlL38Z+/fvBwA8+uijuOmmm/CJT3wC11xzDUZGRvpu52Mf+xhWV1fxgQ98AO95z3sAdAr1733ve3H33Xfji1/8It75zneKz3/kIx9BPp/Hu971LvzO7/yOiADuuece/Nqv/Ro+9alP4ZWvfCX27t0LAHjsscfwmc98BhMTE/jKV74i9ue2227De9/7Xnz4wx/G9773vS5VgZ1COByGx+MRPT4kFqj1IhW9nJBhGGi32/B6vcLR2KXgZHo4v6d+xoEDBw7OBWzLMt1yyy1oNpt4xzveIRwQAFxyySW4+eab0Wg0toyGjh49ittuuw3T09P4jd/4DfF6KBTCJz7xCbjdbtxyyy3i9bm5OfzkJz/B5OSkxQEBwJVXXokbbrgBuq7j3/7t38Trf/M3fwMAeP/7329xiNdccw3e8IY3IJ1O4zvf+c52Dvm0QNM0hMNh8Xez2bTQs8fGxjA+Pm75jmmafdNx7Xbb8r6qFQfgWUNOqNVqqFQqaLfbyGazWF9f70ns0HUdq6urmJubw/z8PNbW1rYcc+HAgYOzj205oTvvvBMA8JrXvKbrveuuuw4AcMcdd/Tdxl133QXTNHHttdd2Gcnx8XFcfPHFWFpawuzsLIAOS+yFL3whrr766q4aCADs2rULQCdKk/dT0zS86lWvOuX9PN2QnVC73bYQE/x+/0kpHJimiVarZXEqdk5IfW0rRt65iGw2i+PHj2NxcRFHjhzBxsYGcrkclpaWutKLhmFgaWkJhUJBOPp8Pi9mOjlw4ODcxZbpONM0MTs7C5fLhT179nS9v2vXLrhcLszOzvYtgNO57Nu3z/b9PXv24ODBgzh06BD27t2Lyy67rG90dfDgQQAQEc/6+joKhQJGR0dFfUTdPgAcOnSoz9GefoRCIQCbY7j5n8fjgc/n61qtbxW1tFoti6OyI3OoTmircQ/nIlRJIoJOJhAIiNfW1tZQq9W6Plsulx1ShgMH5zi2jIS4ukwkErYGz+PxIJlMitRJLzBiGR4etn2f00fT6fSWO3306FF8+9vfhqZpIsLhqrfXFFO+nslkttz+6YTb7UYgELA4l2w2C7fbDU3TuiKh7aTO5M+o3we6HdP5GAn1c5zy8bTbbYtgqwymLx04cHDuYksnxBVmv2I+V6X9nBC3I69g7baxlaJAPp/Hb/3Wb6HVauGNb3yjqFFttZ8kA2y1/Z1AOBy2GMNms4l8Pi9W6XK6UXVCbre753gGO1ICcP5HQqSV94LshLZysDs5y8mBAwfPHFs6oe2wqrZDAeZ2eqVGuI1+28pms3jnO9+Jw4cP48CBA/iDP/gD8d52Uy5ng64cjUZt0250nHI0I6fa4vE4RkdHu5wQHRVTfSrO90hoqwZb+Xi2ihwdJ+TAwbmNLT0MDV2/h5nv9YuWuJ1eigFbbePYsWO44YYb8Pjjj+PAgQP4whe+YPksCQC99nM7+7hT8Pv9XT1DbrcbmUwGtVrN4oRkZ+V2uy10bIJGupcT8nq9FqesMurOdWynLrbdzzpOyIGDcxtbOqFIJIJQKIRcLmebX+dkUL/fj1gs1nM7rAX1qvmwpmNXM/rJT36C66+/HvPz83jpS1+Kv/u7v0MymTyl7feqGe00fD4fUqmU+Nvr9aJareL48eOWmobqhILBYBc7kIa3l0O1qzUdOXLklKe3nmmoDrMf20+9J1XHfL5FgQ4cPNewpRPSNA179+6FruuYn5/ven9ubg6GYVj6h+xAVhxZciqOHDkCAF3bufPOO/H2t78duVwOr3vd6/CFL3yhS5cNAFKpFAYGBrC6umprbHtt/0yh2WwiEAgIhyI7FsMw0Gq1YJqmxai6XC6EQiHbSMhugJ0MNSVnGEZP3blzDeo+qrUvmXCgftbOCTmKEQ4cnLvYVp8QteG+973vdb3H1175ylduaxs/+MEPula6y8vLePLJJzExMSHUDwDgkUcewfve9z7U63W8/e1vx6c+9Slbhp78G7qu44c//OEp7+dOgLRsTdMQDAa7yAgej0dQuGWD6fV6hQGW02umafYkeMjfVWEYhi2V+WyBgqzqkD47coZ63efm5lCtVrs+6/V6LcdumqYTDTlwcA5jW07ojW98I/x+P/7qr/4Kjz32mHj94MGD+PznP49AIIC3vOUt4vXl5WUcOXIE2WxWvDY1NYVXvOIVmJubw5/92Z+J16vVKn7/938fuq7jpptuEq/X63V88IMfRL1ex/XXX48Pf/jDW5IPbrjhBmiahk9+8pNYWFgQr99222341re+haGhIfz8z//8dg75tMEwDCwuLoq/Q6FQV82GTklNLcmpODUa2qq21ctZs55imibK5TJyudxZYc+Vy2Wsrq6iWCxieXnZUitUFyl2TogNquq+233WqQs5cHDuYlvacZOTk/jQhz6Ej3/843jzm9+MK664AqZp4t5770W73caf/MmfYGBgQHz+Qx/6EO677z68733vw/vf/37x+kc/+lHccMMN+NznPocf/OAH2L17Nx566CFsbGzg6quvxg033CA++61vfUsY73w+j9/93d+13beXvOQluP766wEAL3zhC/Gud70Ln//85/H6178eV1xxBSqVCu6//354PB588pOf7BtJ7QTW1tYsBtbr9WJkZKTL0Ho8ni7Shs/ng9frhdvthtvtFqv+aDTa9zharVbPVF2z2USxWEQmkxERQiaTwe7du22VKXYKau0unU5jcnISgH0kZLdvhmF0nTM6IblLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"0.7798691391944885\n",
|
||||
"0.5808621048927307\n",
|
||||
"0.3612504303455353\n",
|
||||
"0.1207217425107956\n",
|
||||
"-0.12942864000797272\n",
|
||||
"-0.3821050226688385\n",
|
||||
"-0.6324726343154907\n",
|
||||
"-0.8769910335540771\n",
|
||||
"-1.1130443811416626\n",
|
||||
"-1.338484525680542\n",
|
||||
"now running WFC\n",
|
||||
"Using gp-exp parameterization.\n",
|
||||
"Iter 1/500 - Loss: 13.906\n",
|
||||
"Iter 51/500 - Loss: -2.750\n",
|
||||
"Iter 101/500 - Loss: -2.881\n",
|
||||
"Iter 151/500 - Loss: -2.900\n",
|
||||
"Iter 201/500 - Loss: -2.905\n",
|
||||
"-2.905431032180786 tensor(-2.9054, grad_fn=<NegBackward>)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAaEAAAEFCAYAAABKJVg6AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAACkcElEQVR4nOz9eZgk1XUmjL+RkZH7nlWZtXd1dXc1dEMDDWIVCCQwsqzFyB4hhIVAgCVZYjSSZfONxEgeZvRY8mi+MXj4DR4hazzClrBkLbbEJ2RAbELsWzdbb1Vda1Zl5b5FLhHx+yPr3LoRmVlVvVRXNdz3efqBisyMuLHdc88573mPZBiGAQEBAQEBgXWAbb0HICAgICDw9oUwQgICAgIC6wZhhAQEBAQE1g3CCAkICAgIrBuEERIQEBAQWDfY13sAGxmqqmLv3r3o7u6GLMvrPRwBAQGBkwKapiGZTOK0006Dy+Va9rvCCC2DvXv34tprr13vYQgICAiclPiHf/gHnHPOOct+RxihZdDd3Q2geSF7enrWeTQCAgICJwcSiQSuvfZaNocuB2GElgGF4Hp6ejAwMLDOoxEQEBA4ubCaNIYgJggICAgIrBuEERIQEBAQWDcIIyQgICAgsG4QRkhAQEBAYN0gjNAawjAMVKtVCKFyAQEBgfYQ7Lg1gq7rmJiYQLVahSzLiMfj8Pv96z0sAQEBgQ0F4QmtEbLZLKrVKoBm9fDMzAwymcw6j0pAQEBgY0EYoTVCoVBo2ZZMJqFp2jqMRkBAQGBj4ojCcU8++STuvvtuvPnmm6jX69i5cyf++I//GBdffPGq9zE3N4e77roLv/nNb5BMJtHb24sPfvCDuPnmm+FwOFq+Pzs7i7vvvhuPP/445ufn4Xa7cdppp+GGG27AJZdc0vYY999/P/7+7/8eBw4cgCzLOOuss/DZz34Wu3btOpLTPWrUajWoqtqy3TAMlMtlEZYTEBAQWMSqPaEf//jHuOGGG/Diiy9i165dOOuss/Diiy/ipptuwn333beqfSQSCXzkIx/Bfffdh0AggEsvvRSlUgl33nknbrzxRtTrddP3Dxw4gN///d/HD37wAwDAJZdcgs2bN+O3v/0tbr75ZnznO99pOcbf/M3f4Atf+AL279+P8847D6Ojo3j00UdxzTXX4NFHH13t6R4TisVix8/K5fIJGYOAgIDASQFjFZibmzNOO+004+yzzzbefPNNtv3ll182du/ebZx++ulGIpFYcT+f+tSnjNHRUeOuu+5i20qlknH99dcbo6Ojxne+8x3T9z/ykY8Yo6Ojxje/+U2j0Wiw7U8++aRx2mmnGaeeeqqxf/9+tn3Pnj3G6Oiocdlll5nG8+tf/9rYsWOHceGFFxrlcnk1p2wYhmFMTk4ao6OjxuTk5Kp/YxiGMT4+brzxxhtt/x08ePCI9iUgICBwsuFI5s5VeUL33nsvarUarr/+eoyOjrLtu3btwk033YRqtbqiN3To0CE88sgjGBoawqc//Wm23ePx4Otf/zpkWca9997Lto+NjeGll17CwMAA/vRP/9SkQXTBBRfgmmuugaZp+OUvf8m2f/e73wUA3HLLLYjH42z7pZdeiquuugoLCwu4//77V3PKxwRd1zt+Vq/XWzw+AQEBgbcrVmWEHn/8cQDA5Zdf3vLZFVdcAQB47LHHlt3HE088AcMwcNlll8FmMx+2r68PO3bswPT0NA4cOAAAyGQyOPPMM3HJJZe0FcEbHh4GAMzPz5vGKUkS3v3udx/1OI8HQqEQ+/9oNAqPx2P6XITkBAQEBJpYkZhgGAYOHDgAm82GkZGRls+Hh4dhs9lw4MABGIYBSZLa7oeMy7Zt29p+PjIygj179mDfvn3YunUrdu/evax3tWfPHgBgHs/8/DxyuRx6enoQDAbb7h8A9u3bt8zZHh+Ew2H4fD4YhgGHw4FUKmUyPOVyue0YBQQEBN5uWNETyuVyqNVqCIVCbdlrdrsd4XAYlUoFpVKp437IY4nFYm0/p74TCwsLKw760KFD+Nd//VdIksQ8nGQyadpPp/2nUqkV9388oCgKu15ut9v0WTvmnICAgMDbESsaoUqlAqB1IuVB7VuXM0K0n06tXmn7SqGqbDaLf//v/z3q9To+/OEPsxzVSuN0Op2r2v9awOVymTzEWq0m6oUEBAQEsAojZM3ftIOxCm002k+ncB3tY7l9pdNpfPKTn8T+/fuxc+dO/Kf/9J/YZ532ezRjPd6w2WwtXiSpKQgICAi8nbGihaGk+nKTJn22nLdE++kUilppH4cPH8Y111yDV199FTt37sR3vvMd03e9Xu+y41zNGNcSVg9QhOQEBAQEVmGEfD4fPB4PMpkMGo1Gy+eNRgOZTAZOpxOBQKDjfigX1CnnQzmddjmjl156CVdffTXGx8dx7rnn4v/+3/+LcDh8VPtfTc/ztYAwQgICAgKtWNEISZKErVu3QtM0jI+Pt3w+NjYGXddN9UPtQKw4YslZcfDgQQBo2c/jjz+OT3ziE8hkMvi93/s9fOc734HP52v5fSQSQTQaRSKRaKtY0Gn/JwrCCAkICAi0YlV1QqQN9+CDD7Z8Rtve9a53rWofDz/8cEsx58zMDF5//XX09/dj69atbPvLL7+Mz33uc1BVFZ/4xCfw3//7f2/L0OOPoWkafv3rXx/1ONcKTqfTlLeq1+ttPUsBAQGBtxNWZYQ+/OEPw+l04tvf/jb27t3Ltu/Zswf33HMPXC4XPvaxj7HtMzMzOHjwINLpNNs2ODiIiy++GGNjY7jjjjvY9nK5jNtuuw2apuGGG25g21VVxRe/+EWoqoqrr74aX/7yl1ckH1xzzTWQJAnf+ta3MDk5ybY/8sgj+MlPfoLu7m68//3vX80pH3dIksQYegShnCAgIPB2x6pUtAcGBnDrrbfi9ttvx0c/+lGcf/75MAwDTz/9NBqNBr75zW8iGo2y799666145pln8LnPfQ633HIL2/61r30N11xzDe6++248/PDD2Lx5M1544QUkk0lccskluOaaa9h3f/KTn2BqagpAk5b9pS99qe3Y3vGOd+Dqq68GAJx55pm48cYbcc899+ADH/gAzj//fJRKJTz77LOw2+341re+tawntdZQFMUUhqvX6+tGlBAQEBDYCFh1K4drr70WfX19uOeee/D888/D4XBg9+7d+MxnPoMLLrhgVfsYHBzED3/4Q9x555147LHHcPjwYQwODuK6667DJz7xCdjtS8N55pln2P8/8MADnU/AbmdGCAD+7M/+DFu3bsX3vvc9/Pa3v4XX68Wll16KW265BTt37lzt6a4JFEUx/S08IQEBgbc7JGM9CmdOEkxNTeE973kPHnroIQwMDBzz/rLZLObm5tjfwWAQPT09x7xfAQEBgY2EI5k7RWfVEwjhCQkICAiYIYzQCYQwQgICAgJmCCN0AtHOCIloqICAwNsZwgidQEiSZCJfAMIbEhAQeHtDGKETDBGSExAQEFiCMEInGMIICQgICCxBGKETDGGEBAQEBJYgjNAJhtUIFYvFFi09AQEBgbcLhBE6wbDK9NRqNVMBq4CAgMDbCcIInWA4HA4Eg0HTtnw+LzqtCggIvC0hjNA6IBaLtShqt+uBdLygaZowcgICAhsSwgitA2w2W0tn2EKhsCbHKpfLOHjwIMbHxzEzM7MmxxAQEBA4WggjtE7wer2m/kjVavWYmXK1Wq2lY+vs7CxTZSgUCqjVasd0DAEBAYHjCWGE1gl2u72FpHAsIblcLoexsTEcPnyYER1UVW3p3iqMkICAwEaCMELrCJ/PZ/r7WIwQ38U2m82iXq8jl8u1fE/QwQUEBDYShBFaR1iNUKVSOSojYRhGi4dTLBbb5pk0TTvi/QsICAisFYQRWkcoimJqN24YBiqVyhHvxxpyA5qeUTuDI4yQgIDARoIwQusMj8dj+rtUKh3xPqyEhlqthqmpKUxPT7d4Q8IICQgIbCQII7TOsBqhcrm87PcNw0CpVEK5XGasN6sRyufzzNjk83lTrkkYIQEBgY0E+8pfEVhLWI1QtVpFo9Fo6TtEmJ+fRzabBQCEw2HEYrEWI2T9O5fLweFwwOFwCCMkICCwoSA8oXWGLMtwuVymbZ3UDXRdNzHeMpkMGo2GKSek63pbcgN5Q8IICQgIbCQII7QBwJMTgPZEA6CZ67G2Ay8UCibPp5ORoX0KIyQgILCRIMJxGwCyLJv+7mQo2nlIuVzO5Pl0MmCNRgOGYQgjJCAgsKEgPKENAGv+p5MhaWeEVFU10bpzuRzS6TQKhYLJazIMA7qus/8KCAgIbAQII7QBcCyekK7rzAhVKhUUi0U4HA7UajVUKhXY7Xa2//UOyamqilQqhbm5uRVZgAICAm8PiHDcBsCxeEKaprGcUKlUgiRJ8Hq9sNvtsNvt6OrqQiqVgqZpaDQacDqd0DStpcPrWmN+fh6ZTIb9ncvlMDw83JIPExAQeHtBeEIbAKvxhBqNRsft9K9UKqFer0OSJLhcLjidTsiyDJvNxr7baf9riUajwWjlBMMwkM/nT+g4BAQENh6EEdoAWI0ntBxtu9FoIJlMIpfLMYMDNI2NYRjrboTasfqAlQtzBQQE3voQRmgDoJ0nZJ20O7Vg0DQNuq4zFW3rvhqNxroboU59klRVFWw9AYG3OYQR2gCQJGnFkJx1Ii8Wi1hYWGB1QuQp8Z4QYDZCZNzWwxNqh6MVbBUQEHjrQBATNgjsdrvJOFile/jPqtUqU07I5/OoVqvMc+KNkMvlgq7rbBsZoNUYoVKphLm5OUiShHg83iIvdCRYrmNsqVRqaWkhICDw9oHwhDYIVvKE+DxRtVqF0+lk3ysUCiiVStA0zdQyXFEUSJJkMkydCA48DMNAIpFAvV5HrVYztQg/GixnhEReSEDg7Q1hhDYIViIn8IajXq/D4XCwcJaqqigWi0ilUkin01BVFS6XC16vt2V/qzFCJKLK/1ZLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"0.7853082418441772\n",
|
||||
"0.5843412280082703\n",
|
||||
"0.3659437894821167\n",
|
||||
"0.12593048810958862\n",
|
||||
"-0.12586945295333862\n",
|
||||
"-0.3793729543685913\n",
|
||||
"-0.6292864084243774\n",
|
||||
"-0.8712751269340515\n",
|
||||
"-1.1031606197357178\n",
|
||||
"-1.3245419263839722\n",
|
||||
"now running COP\n",
|
||||
"Using gp-exp parameterization.\n",
|
||||
"Iter 1/500 - Loss: 13.151\n",
|
||||
"Iter 51/500 - Loss: -2.483\n",
|
||||
"Iter 101/500 - Loss: -2.550\n",
|
||||
"Iter 151/500 - Loss: -2.557\n",
|
||||
"-2.558645725250244 tensor(-2.5586, grad_fn=<NegBackward>)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAZUAAAEFCAYAAAArPXp4AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAACjUElEQVR4nOz9ebxkdX3nj7/OUvu+3Lr72hubIIgIKpuixDXjEhUNIApxH78z8RsmiQOOfknizGSi5EdiRkxiJEwYDYkaiRqCCAqCQiPN3rfvfuvWvu9V55zfH9Xvzz3n1Kl7q5vbTdP9eT4eJk1V3VOnTtX5vD/v7fUWNE3TwOFwOBzODiC+1CfA4XA4nJMHblQ4HA6Hs2Nwo8LhcDicHYMbFQ6Hw+HsGNyocDgcDmfHkF/qEzieNBoNPPXUUxgaGoIkSS/16XA4HM4Jj6IoSKfTOOuss+B0Ord9/SllVJ566il86EMfeqlPg8PhcF52/P3f/z3OP//8bV93ShmVoaEhAN2LMzIy8hKfDYfD4Zz4JBIJfOhDH2Lr53acUkaFQl4jIyOYmJh4ic+Gw+FwXj4MmjLgiXoOh8Ph7BjcqHA4HA5nx+BGhcPhcDg7BjcqHA6Hw9kxuFE5ClqtFjqdzkt9GhwOh3PCcUpVf+0EiUQCxWIRgiAgHA4jEolAEISX+rQ4HA7nhIB7KkdAp9NBsVgEAGiahmw2i42NDfCRNBwOh9OFG5UjoNVq9TxWLpdRq9VegrPhcDicEw9uVI6Adrtt+biVseFwOJxTEW5UjoB+RoUn7TkcDqcLNypHADcqHA6HszXcqBwB/YyKoijH+Uw4HA7nxIQblSOAeyocDoezNdyoDIimaX2NBzcqHA6H04UblQHp56UA3fAX71XhcDgcblQGZiujAnBvhcPhcABuVAbmaIyKpmlotVrci+FwOKcMXPtrAAqFApLJ5JavMVeAqaqKlZUVNJtNSJKESCSCYDDIdcI4HM5JDfdUtqHT6SCVSvU8bjYOZk+lVCqh2WwC6BqcVCqFRCJx7E6Uw+FwTgC4UdkGQRB6wleCIMDv9xseMxsVKz2wUqnEcy8cDuekhhuVbZAkCSMjI7DZbJBlGW63G+Pj43A4HIbXmcNf9Xrd8nhcJ4zD4ZzM8JzKAAQCAQQCAcNjqqoa/lvvgbTb7b4eSavVgtvt3vmT5HA4nBMA7qkcJbJstMd6I9JoNPr+HfdUOBzOyQw3KkfJVkalX+gLAEveczgczskID38NiKIoSKfTaLVaCAQC8Pl8ALpGQhAEQzXYVkaFeyocDudkhhuVASkUCmyUcL1ehyRJyGQyzPPw+/1QVRWCIGzpjXQ6HSiKAkmSjst5czgczvGEG5UBqVarhv9eX183hLwqlQpKpRI0TTOUIMuyDFEUDR5Kq9WCy+U69ifN4XA4xxluVAbEKqQlSRIrJVZVFWtra7DZbIbXOBwOblQ4HM4pA0/UD4CmaZbyKqJovHxWw7ocDgfsdrvhMZ5X4XA4JyvcUxmAftL2ZqNi7l0B0NMkCWwvTsnhcDgvV7inMgBWRkAQhJ4mxn5GxZyU5+OHORzOyQo3KgNg7o632+2Ym5vD8PCw4XGzUREEAXa7nRsVDodzysCNygCYPRWPxwNZlrfV/6L+FW5UOBzOqQLPqQyA2VOhbnpJknoqwKxeZ+6+pxyNIAgol8vI5/Ow2+0YGhri/SscDudlDfdUBsDsqVDZsCRJ8Hq97HGzUQmFQgC6Hos+qa9pGlRVRbvdxsbGBur1OorFIrLZ7LH6CBwOh3Nc4EZlAMxGRZZl1Go1FItF2O12uN1uSJIEj8fDkvdut9ugbGwVAsvlcoaqsnw+fww/BYfD4Rx7ePhrAMzhr1arhUQiAUVRkM1mMTQ0hFAoBEmSMDk5CVVVe8qNZVk2GCdFUXq69DkcDuflDvdUtkFVVUNiXRAElEolAN0+FVVVkc1mUavV0Gq1oGlaj0EBej2VdrvN+1U4HM5JBzcq22CVpKdRwZQr6XQ6yOfzSCaTfcUkzUalXC5bvo5XhnE4nJczRxT+euihh/C1r30Nzz//PNrtNs4880z8zu/8Di6++OKBj5FMJnHbbbfh5z//OdLpNEZHR/HOd74TN9xwQ4+cCQAkEgn8+Z//OR544AHk83lEo1Fccskl+MQnPoHR0dEjOf2jQpZlw5x6s1wLeStANwFfLBbhdDp7jmM2KpVKxfL9Op0OrwDjcDgvWwb2VO6++25cd9112L9/P84++2yce+652L9/P66//nrcddddAx0jkUjgfe97H+666y74/X5cdtllqFaruPXWW/HRj360JxyUTCbx3ve+F9/5znfgdrtx2WWXwev14q677sI73vEOvPDCC0f2aY8CURQxMjICh8MBj8fTYzDMoa5+Ux/NZcX96DeGmMPhcF4ODLTSpVIp3HzzzfD5fLjzzjuxd+9eAMCTTz6J6667Drfccgsuu+yyng5zM1/4wheQSCTw2c9+Fp/85CcBALVaDZ/61Kfw0EMP4Vvf+hY+8pGPsNd//vOfRzqdxnXXXYff+73fYwv4X/7lX+IrX/kKvvCFL+DOO+88qg9+JPj9fvj9fgDA0tKS4TmzUdGHvyqVCsrlMqsOGwQe/uJwOC9nBvJU7rjjDrRaLXz4wx9mBgUAzj77bFx//fVoNpvbeisLCwu4//77MTU1hY9//OPscbfbjVtuuQWSJOGOO+5gj+dyOfz85z9HMBjE5z73OcPi/bGPfQxutxuPPfYYy28cDxRF6cmZ6MuGgc25K6VSCevr6yiVSkgkElvOrdfDPRUOh/NyZiCj8uCDDwIArrjiip7n3vSmNwEAHnjggS2P8bOf/QyapuHyyy/v2d2PjY3hjDPOwPr6Oubn5wEA4XAYDz30EO68886e0BFVTomiaFlpdawwGxSHw4GRkRHDY8ViEblcDhsbG4bHtxoxrIcbFQ6H83Jm2xVZ0zTMz89DFEXMzc31PD8zMwNRFDE/P28pD0+QsdizZ4/l83RsfZ4kGAxi165dhtc1Gg186UtfQrvdxpvf/GbLpPixwjwHxeFwIBgMGoyeqqp47rnnev62n7EIh8OG/+bhLw6H83Jm25xKsVhEq9VCOBy2rM6SZRmhUAjZbBbVatUgW6InlUoBAGKxmOXzQ0NDAIBMJmP5/IMPPohvfvObOHDgAAqFAt7whjfglltu2e70dxSzUSEF4lgshng8zh5vNBrodDoGY9PpdHqGfVE3fi6XM7zuxdBsNlEul+F0Ovt+FxwOh3Os2NZTobDNVuNvyVvYqkOcjtPPs6DH++VIHn30UTz44IMoFArseMvLy1uf/A7TTwPM4/H0KBZTDkVVVVSrVZTL5R6j5HK5ekJ7L8aotNttrKysIJvNYn19vW/ZMofD4RwrtjUqg+Qstgp7mY9jNZZXf4x+x7r22mvx5JNP4t5778VHPvIR/OIXv8DVV1/dU411LLHyVABYyuBT/qVUKqFQKKBUKvVoe/n9/h01Ktls1iBqWSwWj/pYHA6HczRsazFIILFfp7j+ua28GTpOvyqo7Y4RjUbhcDgwOTmJG2+8ER/4wAdQrVbx13/919t9hB1B07QeT4WMis1mszQq5KXoj0FGw+PxwOVyQRRFg6FVVdVygqQVqqoin88jn89DVdUeI/JSeCrVahXr6+tIpVIDfw4Oh3PysK1R8Xq9cLvdyOfzlrtokihxOBysl8MKyqX0y5mk02nD67bjHe94BwDgmWeeGej1L5Z2u23womRZZt6XzWaDzWbrkbc3G1BaZEdGRjA2NtZ3iNeg3koikUAqlUIqlcLa2lrP8/28wmNFp9NBPB5HpVJBPp835Io4HM6pwbZGRRAE7N69G4qiWIaaFhcXoaqqoX/FCqr6oiowM4cOHQIAdpxnn30WN910E775zW9avp68hONVgtsv9AVsSrmYvRWrhsdarQZFUQwGyGqI13aoqmrQD7MqWT7eRqVQKBi8Ez4fhsM59RioyYO0ve69996e5+ixSy+9dKBj3HfffT1hkXg8jmeffRbj4+PYvXs3gO4iedddd+Eb3/iGpZov9cWceeaZg3yEF81WRoUS9vrQndvthsPh6MlJdTqdHg+G/r7few1yPlYcSShtJziejagcDufEZCCj8u53vxsOhwNf//rX8dRTT7HHDxw4gNtvvx1OpxMf/OAH2ePxeByHDh0yhD8mJydx8cUXY3FxEV/96lfZ47VaDZ///OehKAquu+469vi5556LvXv3IplM4o/+6I8MHslPfvIT/NVf/RVkWcbVV199dJ/8CBnEU3G5XCx3ROW8Ho/H8HftdrvHuzKXam+VvzqS1wDHt5mSS/lzOJyBtL8mJiZw44034otf/CI+8IEP4MILL4SmaXjkkUfQ6XTw5S9/GZFIhL3+xhtvxKOPPopPf/rT+MxnPsMev/nmm3HVVVfha1/7Gu677z7Mzs7i8ccfRzqdxiWXXIKrrrqKvVYQBPzP//k/cc011+DOO+/ET3/6U5x++umIx+N45plnIMsyvvSlL+G0007bwcvRn62MiiiKCAQCKBQKCAaD8Pv9LPTl9/shiiJLonc6nZ6F3hw22ylPhd7Pqr9op1EUxdKAWQ0s43A4Jy8DS99/6EMfwtjYGG6//XY89thjsNvtOO+88/CJT3wCF1100UDHmJycxLe//W3ceuuteOCBB7C8vIzJyUlcc801uPbaa3tyC/v27cM///M/4y//8i/x05/+FPfffz/8fj+uvPJK3HDDDXjFK15xZJ/2RbCVUQGA4eFheL1e5rGsrKywMJfb7TZUZpnzH/3KkbfiRPNU+snQHC+jxuFwTgyOaJ7K5Zdfjssvv3zb133rW9/q+9zo6Cj++I//eOD3HB0dxRe/+MWBX38sUBSlZ/qjlZS9PtQ1OjqK1dVVdDodVuVF1WOtVsuwg7fZbIbnO50OFEXZUtn4RDEqpBTQz6hw2RkO59SCz6gfACsvZbvLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"0.7717570662498474\n",
|
||||
"0.5751380324363708\n",
|
||||
"0.3550450801849365\n",
|
||||
"0.11887390166521072\n",
|
||||
"-0.12673339247703552\n",
|
||||
"-0.37664085626602173\n",
|
||||
"-0.626334011554718\n",
|
||||
"-0.872161328792572\n",
|
||||
"-1.1113146543502808\n",
|
||||
"-1.3413406610488892\n",
|
||||
"now running CVX\n",
|
||||
"Using gp-exp parameterization.\n",
|
||||
"Iter 1/500 - Loss: 12.512\n",
|
||||
"Iter 51/500 - Loss: -3.124\n",
|
||||
"Iter 101/500 - Loss: -3.187\n",
|
||||
"Iter 151/500 - Loss: -3.194\n",
|
||||
"-3.194143533706665 tensor(-3.1941, grad_fn=<NegBackward>)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAaEAAAEFCAYAAABKJVg6AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAACnTklEQVR4nOz9aZhkV3UmCr9niHmOyIzIubIGlYYSmpgkQEICZEEbwyc8YMBMRmIw5uH6utu0gQtu2txuu/31Z2TTH7YB2w/CXAxmsEFtQIAQQqABjaWhxqycY57niHPO/RG1du5z4kRkVGVlZZa03+fRo8oYTpxh7732Wutd75IMwzAgICAgICCwA5B3+gQEBAQEBJ6/EEZIQEBAQGDHIIyQgICAgMCOQRghAQEBAYEdgzBCAgICAgI7BnWnT2A3o9ls4vDhwxgfH4eiKDt9OgICAgIXBDRNQyaTweWXXw632z30s8IIDcHhw4fxtre9badPQ0BAQOCCxJe//GW86EUvGvoZYYSGYHx8HEDvRk5MTOzw2QgICAhcGEgmk3jb297G1tBhEEZoCCgENzExgZmZmR0+GwEBAYELC6OkMQQxQUBAQEBgxyCMkICAgIDAjkEYIQEBAQGBHYMwQgICAgICOwZhhLYRhmGg1WpB1/WdPhUBAQGBXQnBjtsm6LqOpaUltFotqKqKubk5OByOnT4tAQEBgV0F4QltE8rlMlqtFgCg2+2iVCrt8BkJCAgI7D4II7RNyOVyQ/8WEBAQEBBGaNsgGtYKCAgIbA5hhLYJgowgICAgsDmEEdomCE9IQEBAYHMII7QN0DTN9nXhHQkICAiYIYzQNqDb7Z7R6wICAgLPVwgjtA3odDq2rw/ykAQEBASerxBGaBswyAgJT0hAQEDADGGEtgHCCAkICAiMBmGEtgEiHCcgICAwGoQR2gYIT0hAQEBgNJyRgOn999+Pz33uczhy5Ag6nQ4OHTqE9773vbj++utHPkYqlcJnP/tZ/OxnP0Mmk8Hk5CTe8IY34Pbbb4fT6Rz6XcMw8M53vhPr6+v4wQ9+YPuZz372s7jjjjsGHuNP//RP8Zu/+Zsjn+/ZYJDHI4yQgICAgBkjG6FvfOMb+OM//mM4nU5ce+210HUdDzzwAG677TZ86lOfwpvf/OZNj5FMJvHmN78ZyWQSl112GQ4dOoRHHnkEd9xxB37xi1/gi1/84lCl6T//8z/HAw88gLm5uYGfeeaZZwAAN998M9xud9/7w757ruDxeFCpVPpeF0ZIQEBAwIyRjFA6ncYnP/lJBAIB/NM//RMOHjwIAHjiiSfw7ne/G5/+9Kdx4403IpFIDD3On/zJnyCZTOLDH/4wfu/3fg8AUK/X8cEPfhD3338/vvSlL+F3f/d3+77XarXwqU99Cl//+tc3Pdenn34aLpcLf/mXfwlV3ZlOFRMTE3C73eh0OigWi+x1kRMSEBAQMGOknNCdd96JdruNd73rXcwAAcAVV1yB2267Da1WC1/96leHHuPkyZO45557MDc3h/e///3sda/Xi09/+tNQFAV33nln3/fuvfde3Hrrrfj617+O2dnZob9RLpexurqKSy65ZMcMEADIsoxoNIp4PG56vdvtCjkfAQEBAQ4jGaGf/vSnAIDXvOY1fe/dfPPNAHrGYhjuu+8+GIaBm266CbJs/tmpqSlcdtllWF1dxfHjx03v3X777VhYWMDb3/52/M3f/M3Q33j66acBAIcOHRp+QecJkiRBURTTa4NIC1bouo5SqYRqtbodpyYgICCwK7CpETIMA8ePH4csy9i3b1/f+/Pz85BlGcePHx+6yyfjctFFF9m+T8c+evSo6fVbbrkF3/rWt/Dxj38cLpdr6LlSPsjj8eAjH/kIbrrpJlxxxRV4wxvegC996Us7ot1mzUvVarWRvreysoJkMonV1VXRi0hAQOA5i02NUKlUQrvdRjgctmWvqaqKSCSCRqMxdIFNp9MA0BeiIoyPjwMAstms6fU77rgDF1988WanCWDDE/rCF76A+++/H4cOHcKll16KkydP4k//9E/x4Q9/+LwbIq/Xa/q7Xq+zf7daLVSr1b5zarVaaDQa7G/rPREQEBB4rmDTxAkthh6PZ+BnaLdfq9Xg9/uHHseOsca/zi/SZwryhN7xjnfgj/7ojxjT7tlnn8UHPvABfP/738eXv/xlvP3tbz/r3zhT+Hw+ZDIZ9ne9XodhGKhUKlhfXwfQu7ezs7OQJAkA0G63z9v5CQgICOwkNvWErPkbO4ySbKfj0EI76BhbSdz/8z//M/7t3/4NH/3oR01U70suuQQf+9jHAABf/vKXz/r4ZwOn02nKC+m6jmaziXw+z15rNBomSredtyYIDQICAs9FbGphKJzUarUGfobeG+Yt0XGazeZZH2MzeL1eHDx40NbQvfKVr4SiKFhYWNiSt3WmkCQJPp/P9FqlUum7n3zex468IHoRCQgIPBexqRHy+/3wer0oFAq2xZbdbheFQgEulwvBYHDgcSgXNCi/QSGrQTmjrcLhcCAUCgEYbAi3C9a8UKFQ6PtMu91mxtHuPosaIwEBgeciNjVCkiThwIED0DQNp06d6nt/YWEBuq6b6ofsQKw4KwWbcOLECQDY9DiDsLq6io9+9KP4+Mc/bvt+rVZDPp+H2+1mxuh8YTNWHyGdTkPXdVtPSBghAQGB5yJGqhMibbi777677z167ZWvfOVIx/jRj37UF1paW1vDM888g+npaRw4cGCUU+qD3+/Ht7/9bXzta1/D4uJi3/vf/va3AQDXXXddX+3OdmMzTTxCq9VCKpWy9YREOE5AQOC5iJGM0Jve9Ca4XC783d/9HQ4fPsxef/LJJ/H5z38ebrcbb33rW9nra2trOHHihCn5Pjs7i+uvvx4LCwv4zGc+w16v1+v4+Mc/Dk3T8O53v/usLyQUCuG1r30tAOBjH/sYyuUye+/w4cP4zGc+A1mWTWoN5wuyLI+s4FAqlWxzVsITEhAQeC5ipJVxZmYGH/nIR/CpT30Kv/3bv41rr70WhmHggQceQLfbxZ/92Z8hFouxz3/kIx/Bgw8+iN///d/Hhz70Ifb6Jz/5SbzlLW/B5z73OfzoRz/C3r178cgjjyCTyeCGG27AW97yli1dzMc+9jE89dRTeOihh3DzzTfj6quvRr1ex8MPPwxd1/HRj34UV1111ZZ+42zhcrlsPRxVVSHLMqNl67qOdrvdZ7SEERIQEHguYmSBtbe97W2YmprC5z//efzyl7+E0+nENddcgw984AO47rrrRjrG7Owsvva1r+GOO+7Avffei8XFRczOzuId73gH3vnOd25Z7y0ajeLrX/86/vZv/xbf+973cN9998Hr9eLlL385brvtNrz0pS/d0vG3AqfTaVvM6/f7oaoqI2xomtaXE2o2m1haWkKpVGLiqINgGAaazWYfNVxAQEBgN0IyRAHKQKysrODVr341fvjDH2JmZmZLxyoWi0ilUn2vT09PwzAMrK2tAQBTnhgbGwPQMyrJZBIejwfhcBhutxt79uwB0GPUZbNZGIaBsbExOJ1OLC4uotVqQVEUzM7OjkyKEBAQEDhXOJO1c+ekpp9nsCMnyLIMr9drCtN1u12TJ9Rut6HrOitWbTab0DQNiqIgmUwyJYp2u41YLMbqjzRNQy6Xw9TU1HZeloCAgMCWINp7nyfYGSGv1wtZluFwOFiBraZp0HWd5YDIqPDsuFarBcMwTPpy7XbbJA8EwLaxnoCAgMBugjBC5wl2+S4qYpUkiYXNyPiQd2RnhJrN5sDCYQEBAYELCcIInUcEAgH2b1mWTQoTZISIJdfpdBhTDuj3hEbtSyQgICCwmyFyQucR1K6i2+0iFouZ2GtOp5OF4oCeEeL15ayekFUKaBAofyQgICCwGyGM0HmEw+EYSBRwuVwm76bb7bLcD2A2QqQzV6/X4XK5hhqZTqcjjJCAgMCuhTBCuwS8ETIMA+VymRkiSZKYp0QGZXFxEfV6HZIkIR6PD6yxEnkiAQGB3QyRE9olUFUVmqaxhne1Wo2F4wzD6Ou2SgQGwzCGtqYQRkhAQGA3Q3hCuwi6rjMPyA58uI6X8RnW60kQGAQEBHYzhCe0S6BpWp8B4kNssiybuquOqiUnPCEBAYHdDOEJ7RIUi0WTYVEUxdRaXVVVliNSFGVkgySMkICAwG6GMEK7AO12G+l02uT5uFwueL1eJmyqKApUVUWxWES324VhGCaVBf5vHiIcJyAgsJshwnE7iG63i6WlJSwsLKBSqZio1OFwGH6/n/1N3s/Y2BgURenLAw3yhjqdjmgDISAgsGshjNAOodVqYXFxkTHeOp0OJEmCLMtQFAU+nw8ulwvBYJC9pus6FEUxGSdCt9tFtVpFuVzuMzrHjh0byqATEBAQ2CmIcNwOIZvN9qlnA73cj8PhYKE5h8OBUCgEn8/HQmutVstUvKooCvL5PMsTtVotps7QbreRz+dRLBaxf/9+9rqAgIDAboDwhHYIzWaT/ZtXzfZ4PHC5XKzjKtCT9IlEIjAMA4ZhMD05+o7P5zMRFdrtNjNYlUqFNcrL5/MiRyQgILCrIDyhHUC32+3zgsbGxqBpGtxuNzRNg8PhYAYjEAhAkiQ4HA50u13oug5ZltHtduFwOOD1elEul02/EY1GUalUmLEjMkO1WkUkEjl/FysgICAwBMII7QCspAJZlk3MuFAohMnJSaYdRx6Rw+FgRkVRFLTbbdbGW5IktFotZqQajYaJ6GAYBmPRCQgICOwWiHDcDoDCaQQrtdrpdEKSJLjdblN7bqfTyYyQLMtot9uMkCDLMiqVChqNBiRJQiqVMoX8ADADJSAgILBbIDyh8wRd15nqtdUTsnonvOFRFIWF5ZxOJxM0lWUZ9XodpVIJnU4H1WoV4XAYhmGYClt5WFuHCwgICOw0hBE6D9B1HadOnWI0bKvagaZppnCc2+1m/+ZVEwAwqjYdlxQU2u0Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"0.794218122959137\n",
|
||||
"0.5897232890129089\n",
|
||||
"0.37306469678878784\n",
|
||||
"0.13818427920341492\n",
|
||||
"-0.10908836126327515\n",
|
||||
"-0.3594636619091034\n",
|
||||
"-0.6069300174713135\n",
|
||||
"-0.8497869372367859\n",
|
||||
"-1.0859397649765015\n",
|
||||
"-1.3128845691680908\n",
|
||||
"now running EOG\n",
|
||||
"Using gp-exp parameterization.\n",
|
||||
"Iter 1/500 - Loss: 13.383\n",
|
||||
"Iter 51/500 - Loss: -2.652\n",
|
||||
"Iter 101/500 - Loss: -2.745\n",
|
||||
"Iter 151/500 - Loss: -2.762\n",
|
||||
"Iter 201/500 - Loss: -2.767\n",
|
||||
"-2.767080307006836 tensor(-2.7671, grad_fn=<NegBackward>)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAaEAAAEFCAYAAABKJVg6AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAACu70lEQVR4nOz9d5hkZ3Uuir97V865c5gsoVEWSoASIMAG44NsEIOMhISIIlyDDxxL+IjLAfuHr/HvIJCPjkH2NQiwQAYDPpgghOIIZWlmpEk9093Toaor56pdtcP9o2Z9/e1du6p7RtMTpP0+jx5NV9ip9l7rW2u9612CpmkaLFiwYMGChRMA8UQfgAULFixYePXCckIWLFiwYOGEwXJCFixYsGDhhMFyQhYsWLBg4YTBckIWLFiwYOGEwX6iD+BkRrPZxK5du5BIJGCz2U704ViwYMHCKQFFUZDJZHDmmWfC7Xb3/azlhPpg165duO666070YViwYMHCKYnvfe97eO1rX9v3M5YT6oNEIgGgcyGHhoZO8NFYsGDBwqmBVCqF6667jtnQfrCcUB9QCm5oaAhjY2Mn+GgsWLBg4dTCasoYFjHBggULFiycMFhOyIIFCxYsnDBYTsiCBQsWLJwwWE7IggULFiycMBwRMWH79u246667sHfvXrTbbWzduhUf/vCHcdlll616G0tLS7jzzjvx2GOPIZPJYHh4GO985zvxoQ99CE6ns+vzyWQSd911Fx555BGk02l4PB6ceeaZuPHGG3H55Zd3ff7OO+/EHXfc0XP/X/7yl/Hud7971cdr4eWh1WpBFEXY7RYHxoIFC91YtWX48Y9/jL/8y7+E0+nEJZdcAlVV8cQTT+Dmm2/Gl770JVx77bUrbiOVSuHaa69FKpXCGWecga1bt+LZZ5/FHXfcgd///vf4p3/6JzgcDvb5qakpXHfddSgWixgdHcXll1+ObDaLxx9/HNu3b8fnPvc5fPCDH9TtY/fu3QCAq6++2rRJamJiYrWnbOFlQNM0JJNJVCoViKKIkZER+Hy+E31YFixYONmgrQJLS0vamWeeqV1wwQXa3r172esvvPCCdv7552tnnXWWlkqlVtzORz7yEW3Lli3anXfeyV6r1WraBz7wAW3Lli3a3Xffrfv8e97zHm3Lli3aV7/6VU2WZfb69u3btTPPPFN7zWteo+3fv1/3nauuuko766yztHa7vZpT64u5uTlty5Yt2tzc3Mve1qsNmUxG27NnD/tvdnb2RB+SBQsWjhOOxHauqiZ0zz33oNVq4QMf+AC2bNnCXj/77LNx8803Q5Ik3HvvvX23cfDgQTz44IOYmJjARz/6Ufa61+vFV77yFdhsNtxzzz3s9enpaTz//PMYGxvDZz/7WR3f/NJLL8W2bdugKAp++ctfstfL5TIWFhZw+umnW+mfE4h6vY5cLqd7rdlsnqCjsWDBwsmMVTmhRx55BADw5je/ueu9q6++GgDw8MMP993Go48+Ck3TcNVVV0EU9bsdGRnBGWecgYWFBUxNTQEACoUCzj33XFx++eWmDU/r1q0DAKTTafbaSy+9BADYunXrak7LwhqhWCx2vaZpGhRFOf4HY8GChZMaK4YLmqZhamoKoihiw4YNXe+vW7cOoihiamoKmqZBEATT7ZBz2bx5s+n7GzZswM6dO7Fv3z5s2rQJ559/ft/oaufOnQCAwcFB9hrVgzweDz7/+c/jySefRC6Xw7p16/Dud78b1113XZcDtHDs0Wq1TF9vt9uWEKwFCxZ0WNEil0oltFothMNhU/aa3W5HJBJBo9FArVbruR2KWAYGBkzfJ42hbDa74kEfPHgQP//5zyEIAovEgOVI6O6778b27duxdetWvOY1r8HBgwfx5S9/GZ/+9KehquqK27fw8tBut01fz2QyqFQqx/loLFiwcDJjRSfUaDQAdKKLXiAWWj8nRNvpJetNr9fr9b7HUywW8alPfQrtdhvXXHONrkZFkdD111+PBx54AN/85jdx77334r777sPIyAh+/etf43vf+17f7Vt4eVAUpaejr9frWFxcXNVCw4IFC68OrOiEVpO+0jRt5R0d3k6vdB1to9+28vk8brrpJuzfvx9bt27FX/3VX+ne/+EPf4if//znuPXWW3VU79NPPx233XYbAFhOaI0hy/KKnymVSsfhSCxYsHAqYEUP4/V6AQCSJPX8DL3XL1qi7fRiSa20jdnZWWzbtg0vvvgitm7dirvvvrvrs16vF1u2bDF1dFdccQVsNhump6dXjLYsHD16peJ4rMZRWbBg4dWBFZ2Q3++H1+tFoVAwNR6yLKNQKMDlciEYDPbcDtWCeqViMpmM7nM8nn/+eVx77bWYmZnBRRddhO985zuIRCIrHboODocDoVAIgEUXXktYDsaCBQtHghWdkCAI2LRpExRFwczMTNf709PTUFVVV5sxA7HiiCVnxIEDBwCgazuPPPIIbrjhBhQKBbz97W/H3XffDb/f3/X9hYUF3HrrrfjCF75guv1arYZ8Pg+3282ckYVjj9VEQsDqUrgWLFh45WNVfGXShrv//vu73qPXrrjiilVt44EHHugqXC8uLmL37t0YHR3Fpk2b2OsvvPACPvGJT6DZbOKGG27A1772NVOGHtCJ2H7605/iRz/6EWZnZ7ve/+lPfwqg0+h6KtOET3Z2n+WELFiwcCRYlRO65ppr4HK58K1vfQu7du1ir+/cuRPf/va34Xa78b73vY+9vri4iAMHDiCfz7PXxsfHcdlll2F6ehpf//rX2ev1eh1f+MIXoCgKbrzxRvZ6s9nEZz7zGTSbTVx77bW49dZbe5IaACAUCuFtb3sbAOC2225DuVxm7+3atQtf//rXIYqiTq3hVEKr1cLs7Cz279+P+fn5k9YZrTYdd7IevwULFo4vVqVtMzY2hs9//vP40pe+hPe+97245JJLoGkannjiCciyjK9+9auIxWLs89Qo+olPfAKf/OQn2eu33347tm3bhrvuugsPPPAA1q9fj2effRaZTAaXX345tm3bxj77k5/8BPPz8wA6tOy/+Iu/MD22Cy+8kImn3nbbbXjxxRfx1FNP4eqrr8Z5552Her2Op59+Gqqq4tZbb8W55557xBfpRKPRaGBhYYEpDtRqNVSr1b41uBMFYyTk8XgYPZ+HFQlZsGABOAIV7euuuw4jIyP49re/jWeeeQZOpxPnn38+Pvaxj+HSSy9d1TbGx8fxox/9CHfccQcefvhhzM7OYnx8HNdffz1uuOEGnd7bk08+yf79q1/9qvcJ2O3MCUWjUdx33334x3/8R/zqV7/Co48+Cq/Xi9e//vW4+eabcfHFF6/2dE8aaJqGVCrVJXnTbDZPOiekaVpXJJRIJDA3N9fldKxIyIIFCwAgaNaStCfm5+fxpje9Cb/97W8xNjZ2Qo5BkiRTQojf78fo6OjxP6A+aLfbOHjwIPvbbrdj48aNUBSli5AyOTnZs3HZggULpzaOxHZaQmonEJIkoVqt9q2j9KKTr5YAcDxhPCZqGLbZbKxPjGBFQhYsWACOcLKqhWOHcrmMZDIJoEOD9/v9GBgY6BpB0atJ+GR0QkZnyp+LkVRiOSELFiwAViR0wsCPO9A0DZVKBfPz86a1HzOoqnrSNYb2c0JG+ScrC2zBggXAckInDGaRjCRJWFhYYH9rmtZXLulki4aOxAlZkZAFCxYAywkdd9RqNRSLxZ5RTKPRYNp2rVarr7E+2ZxQr5oQYKXjLFiwYA6rJnQckc/nmUZeP0iSBJvNZsqK49FreNyJgpWOs2DBwpHCckLHEblcblWfazabKBQKXa/bbDZdzehki4RWS0xQVRX5fB6CICAcDvdVwrBgwcIrG5YTWkMQBdvj8cDj8fRMQRmdCy85xCMcDusc2cnkhMwaVXtFQrlcDk6nE4qiQJIkDA0NHbfjtGDBwskFqya0Rmi325idnUU2m8Xc3FzfQW5mquBGxOPxLvXvkykdZ+aA+AiHnJCiKGi1WiwdZw24s2Dh1Q3LCa0Rstmsru7Rrxbkdrv7KnsnEgnEYrEuw64oykkzoK9XFKQoCubn5zE7O6sjZPDXxkhLt2DBwqsHlhNaIxhTav3YYA6Ho+eICmB52iw1tfI4WSKJXk5oaWkJtVoNmqahVquxv9PpNJLJJJrN5knX72TBgoXjB8sJrQGOlPnlcDjgcrlM3xMEQfeeMSVXLpcxMzNzwqfFmjmhZrOJSqUCYJmY0G630Ww2oaoqVFVFuVw+qWpbFixYOL6wnNAaoF+DqRnsdntPJ+R0OnVFfa/Xq+u/of0tLi6eUNqz0ZHY7XbdKHdyQhQR0rG22+2TqrZlwYKF4wvLCa0BzObn9IMoij3TcUalaUEQTMeTt9vt45rWqtVqSCaTKBQKpsw4+gyhlxMCOpT0arWKpaUlFjlZsGDh1QGLor0GOJrUmMfjgd1u7zLmZuMOIpEI6vV6Fymh1Wp1RUlrgXa7zQYOlstliKLIGG/kbIxkA3rdLForlUrM+RSLRYyPj3epbluwYOGVCSsSWgMciROiqEYQBIyPj+ucjiiKpvRtURQxPj7e9d7RprUkSTqiugwvvgoAe/bswfT0NDKZTM9orF8kZGziTafTqz4WCxYsnNqwIqE1wGoVAOx2OyKRCPvb6XRiYmIC5XIZkiQhGAx2jXbg4fF4UK1W2d9HWosCgFQqhVKpBEEQMDQ0tKpprbyzazabLO3WbrdRrVYxMDDQFQlRupFe70fRPprzsGDBwqkJywmtAQYGBpBMJntGBaFQCPF4HDabrcth9ar5mMFIZjjSSEiSJEbx1jQN+Xy+rxMql8uoVCqoVqss9WaUF6rX6/B4PF2OJBgMQtM05nyMTohP5VmwYOHVAysdtwbwer3YuHEjotGo6ftut7ur8fRoYCQzHKkTMpIA+kUg9XodyWQS1WoVuVwOlUoFiqJ09T9pmgaPx9N1LIFAwPSzBON7Vu+QBQuLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"0.7811691761016846\n",
|
||||
"0.5827669501304626\n",
|
||||
"0.3647322356700897\n",
|
||||
"0.12718917429447174\n",
|
||||
"-0.11822761595249176\n",
|
||||
"-0.3641211688518524\n",
|
||||
"-0.6068462133407593\n",
|
||||
"-0.8445554375648499\n",
|
||||
"-1.0750938653945923\n",
|
||||
"-1.2961429357528687\n",
|
||||
"now running SLB\n",
|
||||
"Using gp-exp parameterization.\n",
|
||||
"Iter 1/500 - Loss: 12.968\n",
|
||||
"Iter 51/500 - Loss: -2.964\n",
|
||||
"Iter 101/500 - Loss: -3.045\n",
|
||||
"Iter 151/500 - Loss: -3.057\n",
|
||||
"-3.059727430343628 tensor(-3.0597, grad_fn=<NegBackward>)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAaEAAAEFCAYAAABKJVg6AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAACXNklEQVR4nO39d5xkdZ3vj79OnVM5h67qOKFnmIEZGIKIoM4AArq6hgUDDFxBkgnQr+vuclVWdln5GVbvFYR72SvIurAqC6KuLmtAhCGDDGGGMDOdY4WunKtOnfP7o+b96XMqdNeE7p6Bz/Px8CFTVX3q1KlTn/fnnV5vQVVVFRwOh8PhrACGlT4BDofD4bx14UaIw+FwOCsGN0IcDofDWTG4EeJwOBzOisGNEIfD4XBWDGmlT+BIplQqYffu3ejq6oIoiit9OhwOh3NUUKvVEIvFcPzxx8NisSz4Wm6EFmD37t245JJLVvo0OBwO56jk3//933Hqqacu+BpuhBagq6sLQP1Cdnd3r/DZcDgcztFBOBzGJZdcwtbQheBGaAEoBNfd3Y3+/v4VPhsOh8M5uugkjcELEzgcDoezYnAjxOFwOJwVgxshDofD4awY3AhxOBwOZ8XgRmiZkWUZlUplpU+Dw+Fwjgh4ddwykcvlEI/HUSqVAAAejwehUGiFz4rD4XBWFu4JLQOVSgXT09PMAAFAKpWCLMsreFYcDoez8nAjtAzk8/mWj/OwHIfDeavDjdAyUK1WWz7OPSEOh/NWhxuhZaCdEarVast8JhwOh3NkwQsTlgBVVZHNZlGpVGAwGNqG47gnxOFw3upwI7QERKNRpFKpRV/HjRCHw3mrw8NxS0A7z6cRboQ4HM5bHW6ElgCHw9HR67gR4nA4b3W4EVoCurq60Nvbu+jruBHicDhvdbgRWgIEQYDT6YTNZlvwdYqiQFGUZTorDofDOfLgRmgJsVqti76Ge0McDuetDDdCS8hinhDAjRCHw3lrw43QEmKxWJoeM5vNun/zhlUOh/NWhhuhJcRgMMDlcrF/m0ymphAd94Q4HM5bGd6susQEg0EYjUYoigKfz4d0Oq17nhshDofzVoYboSVGFEUEAgH2b0nSX3JuhDgczlsZHo5bZrgR4nA4nHm4EVpmuBHicDicebgRWmYajRCvjuNwOG9luBFaZgwGAwRBYP+u1WpcNYHD4bxl4UZomREEgYfkOBwOZz/cCK0Aoijq/s1DchwO560KN0IrAPeEOBwOpw7vE1oB2hmhcrmMZDIJSZLg8/lgMPA9AofDeXPDjdAK0MoIybKMyclJFppTVRVdXV1tj5FKpZDP52G32+F2u3XFDhwOh3O0wLfaK0ArIxSNRnW5oYVGhBcKBUQiEeRyOUQiERQKhSU7Vw6Hw1lKuBFaARqNUKlUQjab1T1WLpfb/v3c3Jzu37FY7PCdHIfD4Swj3AitAI1GqFKpNL1mofBasVjU/Xshg8XhcDhHMgeUE3rqqadwxx13YM+ePahWq9i8eTM+/elPY+vWrR0fIxKJ4Pbbb8eTTz6JWCyGnp4efPjDH8bVV18Nk8nU9PrZ2VnccccdePzxxxGNRmG1WnH88cfj8ssvx7Zt21q+x0MPPYQf//jHGBoagiiKOPnkk3HNNddgy5YtB/Jxl4zGEu1WqKoKVVV5rofD4byp6dgTevDBB3H55ZfjxRdfxJYtW3DyySfjxRdfxFVXXYX77ruvo2OEw2F84hOfwH333QeXy4WzzjoL+Xwet956K6688kpUq1Xd64eGhvBXf/VX+NnPfgYA2LZtG9auXYunn34aV199Ne66666m9/jBD36AL33pS9i3bx/e8Y53YMOGDXjsscewfft2PPbYY51+3CVFFMWOjAvvH+JwOG92OvKEotEobrzxRjidTvzkJz/Bhg0bAACvvPIKLr/8ctx8880466yzEAqFFjzOP/zDPyAcDuOLX/wiPv/5zwOoJ9mvueYaPPXUU7jnnntwxRVXsNd/7WtfQyqVwpVXXokvf/nLzIN4+umn8elPfxrf+973cOaZZ2L9+vUAgN27d+O2225DX18ffvrTn7LzefTRR3HNNdfgq1/9Kh5++OGmwXLLjSAIEEVx0f4gLufD4XDe7HTkCd17772oVCr41Kc+xQwQAGzZsgVXXXUVyuXyot7QyMgIHn30UaxatQqf/exn2eM2mw0333wzRFHEvffeyx4fHR3FSy+9hP7+fp0BAoAzzjgD27dvR61Ww29/+1v2+N133w0AuO6663QG8ayzzsL555+Pubk5PPTQQ5185CXHaDQu+ppWnlA776hYLEJV1UM+Lw6Hw1lOOjJCjz/+OADg3HPPbXruvPPOAwDs2LFjwWM88cQTUFUVZ599dlMTZm9vLzZt2oTp6WkMDQ0BAJLJJE466SRs27atZQ5lzZo1AOpemvY8BUHAe97znoM+z+XC6XQu+ppWnlA772liYgIjIyNcfYHD4RxVLBqOU1UVQ0NDMBgMGBwcbHp+zZo1MBgMGBoaWjCRTsblmGOOafn84OAgdu3ahb1792L9+vU45ZRTFvSudu3aBQDM44lGo0in0+ju7obb7W55fADYu3fvAp92+fB6vQDq5dXtPJhWXk9j3kyLLMtIJpMLNrlyOBzOkcSinlA6nUalUoHH42lZvSZJErxeL4rF4oINluSxBIPBls/TwtnYA9OKkZER/PrXv4YgCMzDoV6ZdgswPR6Pxxc9/uFAVdVFczperxerVq2C1+tFMBhsMp4H4gkRiUTiwE+Ww+FwVohFPSHqSVkomW+xWADUu/wdDseCx6HXtjvGYt3/qVQKX/jCF1CtVvHRj36U5agWO0+z2dzR8Q8HuVwOs7OzANDSuGixWCzsszc2nbbyhHi4jcPhvJlY1BPqRESzk4Q4HadduI6OsdCxEokErrjiCuzbtw+bN2/G3//937PnOu2nWY7k/dzcHBRFgaIoC4bbGmnMfSmKglqthmKxyLyiTowQN1QcDudoYVFPyGazAVi4K5+eW8hbouOUSqWDOsb4+Dg+/elPY2xsDJs3b8Zdd92le63dbl/wPDs5x8OF1ujUajWUSqWO3rfR4BeLRYyOjqJWq8FkMmH16tUdGZhyudykysDhcDhHIou6OQ6HAzabDclksuUCSMlws9kMl8vV9jiUC2qX86FQVKuc0UsvvYQLL7wQY2NjOO200/Bv//ZvLLF/oMdfjqQ9hf6Idoa3kUZPqFgsspBcpVJBLpdbsDCB4DI+HA7naGFRIyQIAtavX49arYaxsbGm50dHR6Eoiq5/qBVUFUdVco0MDw8DQNNxHn/8cVx22WVIJpP4y7/8S9x1110t804+nw9+vx/hcBi5XK7j4y8FjXmvTo3QYqHPVCrVsSfE4XA4RwMd9QmRNtzDDz/c9Bw9duaZZ3Z0jEceeaSp6mtmZgavv/46+vr6mPoBALz88su49tprUSqVcNlll+F73/teywo97XvUajX86U9/OujzPBwcrBFaTFOuVCp1pKLAjRCHwzla6MgIXXDBBTCbzfjhD3+I3bt3s8d37dqFO++8ExaLBRdffDF7fGZmBsPDw7py4YGBAWzduhWjo6O45ZZb2OOFQgE33HADarUaLr/8cvZ4qVTCX//1X6NUKuHCCy/EV7/61UWLD7Zv3w5BEPDd734Xk5OT7PFHH30Uv/jFL9DV1YUPfvCDnXzkQ6LRCFUqlY6Mx2KeUKcFDpVKpUlpm8PhcI5EOspe9/f34/rrr8dNN92Eiy66CKeffjpUVcWzzz4LWZbx7W9/G36/n73++uuvx3PPPYdrr70W1113HXv8xhtvxPbt23HHHXfgkUcewdq1a7Fz507EYjFs27YN27dvZ6/9xS9+gampKQD1MNTf/M3ftDy3t7/97bjwwgsBACeddBKuvPJK3HnnnfjQhz6E008/Hfl8Hs8//zwkScJ3v/vdBT2pw4XBYIDJZNKNaCiVSqw4ox2dqGtr8fl8KJVKUFVVZ3RUVcXU1BTWrFnTkTwQh8PhrBQdl1Bdcskl6O3txZ133okXXngBJpMJp5xyCj73uc/hjDPO6OgYAwMDuP/++3Hrrbdix44dGB8fx8DAAC699FJcdtlluoqu5557jv337373u/YfQJKYEQKAv/3bv8X69etxzz334Omnn4bdbsdZZ52F6667Dps3b+704x4yFotFZ4QKhcKiRqiTcnhCEAT4fD5muNLpNMLhMHteURSkUimunsDhcI5oBJWrXrZlamoK55xzDv74xz+iv7//gP42lUohEomwfwuCgDVr1izqie3bt6+j0J3VasWqVat0j8ViMV0I1OFwoK+v74DOm8PhcA6VA1k7+WTVJcLpdOrCa6qq6oxSOzoNybWqEGx8jDetcjicIx1uhJYIURSbep4KhcKihqGTkJzBYGgpBdRowPhQPA6Hc6TDjdAS4nK5mhpXFyuf7qQZ1ev1tvSYGlUSuBHicDhHOtwILTGN5dqLGaHG4gWDwaDzjgwGQ5NahPY5bRk76ddxOBzOkQo3QktMoyekrZhrBWngAfVihoGBAfj9fgiCAEEQ0N3dvWDeqNEb4nkhDodzJMNVLpeYxmq4xTwht9sNRVFQLpfhcrnYqAe32w1BEBbNGYmiqAvp8ZAch8M5kuFGaIlp5QktNIGW+n8a6bRqjueFOBzO0QQPxy0xkiTpDIiiKEsaIms0Vjwcx+FwjmS4EVoGDjQkdyjwMm0Oh3M0wY3QMtAYksvlckgkEsjn84f9vXhhAofDOZrgOaFloNEIpdNp9t89PT0LDgM8ULgnxOFwjia4J7QMLDTaO5VKHdb34kaIw+EcTXAjtAyYzea2wqXFYvGwNpTycByHwzma4EZomVgo5HY4B9BxT4jD4RxNcCO0TDidzrbPFQqFw/Y+rYwQn9bB4XCOVLgRWiZMJlOTjhxxOI2QIAjcG+JwOEcN3AgtI6FQqOW47XK5fFgNRWNe6HCG+zgcDudLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"0.7867836952209473\n",
|
||||
"0.5850406289100647\n",
|
||||
"0.36788874864578247\n",
|
||||
"0.13060295581817627\n",
|
||||
"-0.11842338740825653\n",
|
||||
"-0.370065838098526\n",
|
||||
"-0.6204082369804382\n",
|
||||
"-0.867045521736145\n",
|
||||
"-1.1074048280715942\n",
|
||||
"-1.339046597480774\n",
|
||||
"now running XOM\n",
|
||||
"Using gp-exp parameterization.\n",
|
||||
"Iter 1/500 - Loss: 12.376\n",
|
||||
"Iter 51/500 - Loss: -3.199\n",
|
||||
"Iter 101/500 - Loss: -3.270\n",
|
||||
"Iter 151/500 - Loss: -3.279\n",
|
||||
"-3.2798614501953125 tensor(-3.2799, grad_fn=<NegBackward>)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAa0AAAEFCAYAAABQGbi0AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAACvAUlEQVR4nOz9ebglZXkujN81rXke9rx37+7eNNANzaygQQZFYwzxaFRAjigRjWP8knP8cfSYg4eEk89cnuMFxhMSiR4D0Y/AQROjRkVA5gaZmoaGnvbQe1xrr3mutarq90f18+63atXaezd0Nxup+7r66u411KqqVeu963me+7kfwTAMAy5cuHDhwsXrAOJrvQMuXLhw4cLFeuGSlgsXLly4eN3AJS0XLly4cPG6gUtaLly4cOHidQOXtFy4cOHCxesG8mu9A79NaDab2LNnD9LpNCRJeq13x4ULFy5eF9A0DdlsFqeddhp8Pt+qr3VJ6xhiz549uPrqq1/r3XDhwoWL1yX+6Z/+Ceeee+6qr3FJ6xginU4DME/8wMDAa7w3Lly4cPH6wOLiIq6++mq2hq4Gl7SOISglODAwgJGRkdd4b1y4cOHi9YX1lFVcIYYLFy5cuHjdwCUtFy5cuHDxuoFLWi5cuHDh4nUDl7RcuHDhwsXrBi5pbRBomoZWq/Va74YLFy5cbGi46sENgGazidnZWWiahmAw6CoPXbhw4aIH3EhrA6BUKkHTNABArVZDvV5f8z2tVguFQgHNZvN4754LFy5cbBi4kdYGgH0Op6qqCAQCPV+vqiqmp6dhGAYEQcDY2Nia1icuXLhw8dsAN9LaAFAUxfL/dru96utzuRwjOsMwUCqVjtu+uXDhwsVGgktaGwB20lJVddXXl8tly/+LxeKx3iUXLly42JBwSWsDYLVIq1qtIpvNotFoAAB0Xe96vyi6X6MLFy7eGHBrWhsATqRlGAaWlpZY6q9QKGB8fBydTqfr/YZhsPqWCxcuXPw2wyWtDQBZliGKIouidF3H/Pw8qtUqe41hGKjVao6RlmEYaLfb8Hg8J2yfXbhw4eK1gJtX2iCwR1s8YRE6nQ5LE9qxVh3MhQsXLn4b4JLWBoGdtJygqmpP0lpLcejChQsXvw1w04MbBOshLafoi+BGWi5cuHgjwI20NghebT3KJS0XLly8EeCS1gbBeiKt1eCSlgsXLt4IcNODGwRHG2n5fD60Wi3mjNHpdKBp2rrGVR8rtFotLCwsoNPpIJVKIRaLnbDPduHCxRsTbqS1QaAoChKJBPu/JElIp9M9I7BoNNpFdCfaPDebzaLVakHTNGQyGWb668KFCxfHC26ktYGQTqeRTCZhGAaLmKrVqqMy0OfzsWiL0Gw2EQwGj9v+qaqKarUKXdehKApqtRp7zjCM4/75Lly4cOGS1gaD3ZJJlru/IkEQ4PF44PP5LGa5xzPS6nQ6mJmZWTWasrvVu3DhwsWxhktaGxTlchnZbJaZ4/KjShRFgSiK8Pv9lvf06uE6FuBnfvWCmx504cLF8YZb09qA0HUdmUyGiSsKhYJlHInX6wVgijf4yEzTNEsqkVJ2x4JM1kOILmm5cOHieMONtDYgSNwAgNW2ms0mms0m/H4/G/goCAJ8Pp9l0nGj0YCiKDAMA7Ozs6jX6xBFESMjI12R2WowDAOVSgX1eh26rlvqV/R8p9OBJEmMOF3ScuHCxfGGS1obBIZhQNd1SJJk6bniJeyqqsLv97NIC0AXaZEwo9FosMd1XUcul8PIyMi694d3mHfa1+XlZaiqClEUkUqloCiKowO9CxcuXBxLuKS1AVCv17F79240Gg0MDg4imUyy5/j0H6X+eNKyy96JOCqViuVxe6S0GgzD6Bo0yaPRaDBipSgsFou5kZYLFy6OO1zS2gCYnp5mUdH8/DwAsNlYkiRBEAQ2fkSWZYui0N5MTMSx2mDIZrOJfD4PWZaRSqW6XkvzvOyglKB9UrJLWi5cuDhRcElrA8Deh9VoNJhaUBRFhMNhlMtlGIaBeDzOpO0+n6+LtJzmbRGoBjU3N8cisna7Db/fD0mSEIlEIAiCoyWUYRjIZrOrusm7pOXChYvjDZe0NgD4dB/QTTzhcJiRWKVSQTabBQAkk0lEIhHLa4k4nAik3W6j3W5bak/VapW5x6uqinQ67UhMjUYD7XYboigiFoshn8+z5ygqdEnLhQsXxxuu5H0DwK7qc4qWKE24vLzMHsvn811pPCIOJ1GEnbDsIOGFnbRCoRBz6YjFYvD7/ZZ9NgyDCUncBmMXLlwcT7iR1gbAekgLMMmk1WohFAoBWOnDsr/XMIyekdZqtS5N09DpdLpIi6I53qIpkUhA13UsLS1B13WmfNQ0zdHFw4ULFy6OBdxIawPA6/WyFBtgJS3qyQJMObudUKgPi0ARjxNpqaq6ZgpPVdWumpaiKI51LlEU2WfTPrspQhcuXBxPuLfEGwDVhop7frUbxWIVZ546gvPPnGDPhcNhFk0R6RiGwUiuXC53RU+NRgMLCwvQNM1SD2u32xZydIITMcqy3FOAIUkSc+6gXi17jc6FCxcujhXcSGsD4Nt3PYZ9k1mUa0089NQBzGcK7Dm/3w9ZlmEYhqU3isfS0pKlD2tpaQmtVovJ0+n1a9W0AFO+ztelJElyrFURMdndMNxIy4ULF8cTLmltADSaLQS8AiJ+AR4Z2HtwAYBpkuv3+6EoiqV3yk48oiiiVCoxcuH7qHiy63Q6a044tjch02fz8Pv9rK7mkpYLFy5OJFzS2gA4c1sfBAEQAPhkAQdnMtB1HcFgEKqqdtWU7MRAzceaprE/PPgoyW7tVKlUWA+YE5zqWR6Ph/WHuaTlwoWLEwmXtDYAzt0xBpO1TOKo1ZuYnM0gm81icnISjUbDQgadTsci0CDiaLfbUFW1KxKj9xqGgVwuB03TUCqV0Gw2EQqFUK1WLf6FPDweT1ek5fF4mELQFWK4cOHiRMIlrQ0ARRaweSACEQJkCWhrBnbvm2fPl0olC3FUKhXW5AusEAfJ1e3EwUdBmqZhfn4es7OzbLuyLPecl+X1eh3VhC5puXDh4rWAS1obAM1mE4PpECQJ0HQzTfjCgRXSEkWxywA3n88zU1s7afWKtDRNg67rKBQKLB1YLpfZKBO7q7sgCCxFycMp0nLTgy5cuDgRcElrA6BeryMW8kIRAZ8C+L3A3GIRLdUkH0mS0Gg0LGQkiiKazSZ0XbekB51IKxKJQFEU6LpuISBRFNFqtViq0T4wMhgMQhCEru0pisJqWvT3ak3NLly4cHGs4JLWBoAoilBkET6fDAGAVxYAGJicMy2bSHZOs7LoPdQ/xUdaTo3FhmEgEol0PU7v8/l8LNriJxSHw+Eux3dFUVhTMT8AcrWmZhcuXLg4Vjiq5uJHH30Ut956K15++WW0223s2LEDn/zkJ3HhhReuextLS0v41re+hUceeQTZbBaDg4P4gz/4A3ziE5/omg1lh2EY+OhHP4qFhQX88pe/dHxNs9nEd77zHfzsZz/DzMwMBEHAli1b8L73vQ9XX311VyPub37zG1x99dU9P/Pyyy/H17/+9XUf3yuBIAhoNpsI+xXoWgdtXYAoGDgwncUpmwfYPquqikAgAEEQGHGQfRKP1UiL7/GiRmNd1xEOh5HP51Gv1xEKhVhqkCcxwCQtAk9atB2++VnTNFQqFciyzCTyLly4cPFqsG7Suueee/ClL30JHo8H559/PnRdx65du3DdddfhxhtvxBVXXLHmNhYXF3HFFVdgcXER27dvx44dO/D000/jlltuweOPP47vfOc7lkXRjr/+67/Grl27MDY25vh8rVbDNddcgz179iAWi+G8885Du93Gc889h7/8y7/EY489hm9+85uWRf7FF18EAJx11lmOk33PPvvsNY/rWEAQBPi8CsrVBhQJEAVg/7Tp5s6TQKfTgaIoFmcLOxHb03k0bqRQKFhIyOv1oq+vD4ZhsInI5IiRTqe7pigD1qGTJIfnJfcUcYmiiMOHD7PoMJ1OI5FIHJuT5cKFizcs1kVamUwGN9xwA8LhML7//e9j27ZtAIDdu3fj2muvxU033YSLL74Y/f39q27nq1/9KhYXF/GFL3wBn/nMZwCY9ZzPfvazePTRR3H77bfjj/7oj7re12q1cOONN+Luu+9edfu33nor9uzZgwsuuAC33HILM3qdnZ3Fxz/+cfzqV7/CXXfdhSuvvJK9Z+/evQCAL37xizjnnHPWczqOOciI1u81CVsUzD+7983hn378BP7g0p0sQmo2mxavQt7Sif7vRFpPPfUUBEFgBEMDICndp+s6EokEGo0GotEo+vr6AGBV0uLFGHx/mKZpaDablnRmoVBwScuFCxevGuuqad1xxx1QVRUf+9jHGGEBwM6dO3Hdddeh1WrhzjvvXHUbhw4dwgMPPICxsTF86lOfYo8HAgHcdNNNkCQJd9xxR9f7HnzwQbzvfe/D3XffjdHR0VU/44c//CEA4C/+4i8sc6ZGRkbwxS9+EQDwk5/8xPKeF198EaIo4tRTT11128cTgUAAsizDo0gQRQmCYMrgAeCe+/fi1h/cxyJQVVVRLBaZVZMTafG2S6Iool6vM/f1eDyOaDSKZDLJBBher5e9LxgMWkjR3qPFR8L0bycFod1ZYy37KBcuXLhYD9ZFWg899BAA4B3veEfXc5dddhkAk1xWw8MPPwzDMHDJJZd0pbOGhoawfft2zM3N4cCBA5bnPvGJT2BychIf+chH8Hd/93c9t1+r1TA+Po6dO3c6ktvmzZsBmFEjQVVVHDx4EFu2bGGmsq8FfD4fZFmGIAjwe00xRsADxAMCRhISioU89k8tMCKiaCqfz0OSJIsRrr2eJUkSG/IImFGX1+tFKpWykM3S0hKWlpaQyWQsEdJ6Ii1Kt/KktZZd1FrQdR35fB7Ly8urTkt24cLFGwtrkpZhGDhw4ABEUcSWLVu6nh8fH4coijhLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"0.7942204475402832\n",
|
||||
"0.5890586972236633\n",
|
||||
"0.37192219495773315\n",
|
||||
"0.1361386924982071\n",
|
||||
"-0.11294876039028168\n",
|
||||
"-0.36603081226348877\n",
|
||||
"-0.6161594390869141\n",
|
||||
"-0.8605756163597107\n",
|
||||
"-1.0972329378128052\n",
|
||||
"-1.323974609375\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"all_gpcv_paths, all_gpcv_vols = [], []\n",
|
||||
"for i in range(10):\n",
|
||||
" print(\"now running \", tckrs[i])\n",
|
||||
" \n",
|
||||
" pred_vol = get_and_fit_gpcv(train_x, log_returns[i])\n",
|
||||
" vol_model = get_and_fit_vol_model(train_x, pred_vol)\n",
|
||||
" data_model = get_and_fit_data_model(train_x, train_y[i], pred_vol, vol_model)\n",
|
||||
" paths, vols = predict_prices(test_x, data_model)\n",
|
||||
" all_gpcv_paths.append(paths.detach())\n",
|
||||
" all_gpcv_vols.append(vols.detach())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 12,
|
||||
"id": "bf957ded",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"px_paths = torch.stack(all_gpcv_paths)\n",
|
||||
"vol_paths = torch.stack(all_gpcv_vols)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 13,
|
||||
"id": "b8e98f54",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"torch.save(f=\"ind_predictions.pt\", obj={\"paths\": px_paths, \"vol_paths\": vol_paths, \"test_y\": test_y})"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "dd5d90e0",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
@@ -0,0 +1,232 @@
|
||||
import numpy as np
|
||||
import matplotlib.pyplot as plt
|
||||
import seaborn as sns
|
||||
import torch
|
||||
import gpytorch
|
||||
# from voltron.robinhood_utils import GetStockData
|
||||
import os
|
||||
# import robin_stocks.robinhood as r
|
||||
import pickle5 as pickle
|
||||
|
||||
sns.set_style("whitegrid")
|
||||
sns.set_palette("bright")
|
||||
|
||||
sns.set(font_scale=2.0)
|
||||
sns.set_style('whitegrid')
|
||||
|
||||
import sys
|
||||
sys.path.append("../")
|
||||
from voltron.likelihoods import VolatilityGaussianLikelihood
|
||||
from voltron.models import SingleTaskVariationalGP as SingleTaskCopulaProcessModel
|
||||
from voltron.kernels import BMKernel, VolatilityKernel
|
||||
from voltron.models import BMGP, VoltronGP, MaternGP, SMGP
|
||||
from voltron.means import LogLinearMean
|
||||
from gpytorch.kernels import ScaleKernel, RBFKernel, MaternKernel
|
||||
|
||||
def get_and_fit_gpcv(x, log_returns, printing=False):
|
||||
train_x = x[:-1]
|
||||
dt = train_x[1]-train_x[0]
|
||||
# prepare model
|
||||
likelihood = VolatilityGaussianLikelihood(param="exp")
|
||||
# likelihood.raw_a.data -= 6.
|
||||
covar_module = BMKernel()
|
||||
model = SingleTaskCopulaProcessModel(
|
||||
init_points=train_x.view(-1,1),
|
||||
likelihood=likelihood,
|
||||
use_piv_chol_init=False,
|
||||
mean_module = gpytorch.means.ConstantMean(),
|
||||
covar_module=covar_module,
|
||||
learn_inducing_locations=False
|
||||
)
|
||||
model.mean_module.constant.data -= 4.
|
||||
# model.initialize_variational_parameters(likelihood, train_x, y=log_returns)
|
||||
|
||||
import os
|
||||
smoke_test = ('CI' in os.environ)
|
||||
training_iterations = 2 if smoke_test else 500
|
||||
|
||||
|
||||
# Find optimal model hyperparameters
|
||||
model.train()
|
||||
likelihood.train()
|
||||
|
||||
# Use the adam optimizer
|
||||
# likelihood parameters should be taken acct of in the model
|
||||
optimizer = torch.optim.Adam([
|
||||
{"params": model.parameters()},
|
||||
# {"params": likelihood.parameters(), "lr": 0.1}
|
||||
], lr=0.01)
|
||||
|
||||
# "Loss" for GPs - the marginal log likelihood
|
||||
# num_data refers to the number of training datapoints
|
||||
mll = gpytorch.mlls.VariationalELBO(likelihood, model, log_returns.numel())
|
||||
|
||||
old_loss = 10000.
|
||||
print_every = 50
|
||||
for i in range(training_iterations):
|
||||
# Zero backpropped gradients from previous iteration
|
||||
optimizer.zero_grad()
|
||||
# Get predictive output
|
||||
with gpytorch.settings.num_gauss_hermite_locs(75):
|
||||
output = model(train_x)
|
||||
# Calc loss and backprop gradients
|
||||
loss = -mll(output, log_returns)
|
||||
loss.backward()
|
||||
if printing:
|
||||
if i % print_every == 0:
|
||||
print('Iter %d/%d - Loss: %.3f' % (i + 1,
|
||||
training_iterations,
|
||||
loss.item()))
|
||||
optimizer.step()
|
||||
if old_loss <= loss and i > 100:
|
||||
if printing:
|
||||
print(old_loss, loss)
|
||||
break
|
||||
else:
|
||||
old_loss = loss.item()
|
||||
|
||||
model.eval();
|
||||
likelihood.eval();
|
||||
predictive = model(x)
|
||||
pred_scale = likelihood(predictive).scale.mean(0).detach()
|
||||
samples = likelihood(predictive).scale.detach()
|
||||
|
||||
# plt.plot(x, pred_scale, linewidth = 4)
|
||||
# plt.plot(x, samples.t(), color = "gray", alpha = 0.3)
|
||||
# # plt.ylim((0, 0.25))
|
||||
# plt.show()
|
||||
|
||||
# return scaled volatility prediction
|
||||
return pred_scale / dt**0.5
|
||||
|
||||
def get_and_fit_vol_model(train_x, est_vol):
|
||||
vol_lh = gpytorch.likelihoods.GaussianLikelihood()
|
||||
vol_lh.noise.data = torch.tensor([1e-6])
|
||||
vol_model = BMGP(train_x, est_vol.log(), vol_lh)
|
||||
|
||||
optimizer = torch.optim.Adam([
|
||||
{'params': vol_model.parameters()}, # Includes GaussianLikelihood parameters
|
||||
], lr=0.01)
|
||||
|
||||
# "Loss" for GPs - the marginal log likelihood
|
||||
mll = gpytorch.mlls.ExactMarginalLogLikelihood(vol_lh, vol_model)
|
||||
old_loss = 10000
|
||||
for i in range(500):
|
||||
# Zero gradients from previous iteration
|
||||
optimizer.zero_grad()
|
||||
# Output from model
|
||||
output = vol_model(train_x)
|
||||
# Calc loss and backprop gradients
|
||||
loss = -mll(output, est_vol.log())
|
||||
loss.backward()
|
||||
if i % 50 == 0:
|
||||
print(loss.item())
|
||||
optimizer.step()
|
||||
# if old_loss <= loss:
|
||||
# break
|
||||
# else:
|
||||
# old_loss = loss.item()
|
||||
|
||||
return vol_model
|
||||
|
||||
def get_and_fit_data_model(train_x, train_y, pred_vol, vol_model):
|
||||
voltron_lh = gpytorch.likelihoods.GaussianLikelihood()
|
||||
voltron = VoltronGP(train_x, train_y.log(), voltron_lh, pred_vol)
|
||||
# voltron.mean_module = gpytorch.means.LinearMean(1)
|
||||
voltron.mean_module = LogLinearMean(1)
|
||||
voltron.mean_module.initialize_from_data(train_x, train_y.log())
|
||||
voltron.likelihood.raw_noise.data = torch.tensor([1e-6])
|
||||
voltron.vol_lh = vol_model.likelihood
|
||||
voltron.vol_model = vol_model
|
||||
|
||||
grad_flags = [False, True, True, True, False, False, False]
|
||||
|
||||
for idx, p in enumerate(voltron.parameters()):
|
||||
p.requires_grad = grad_flags[idx]
|
||||
|
||||
voltron.train();
|
||||
voltron_lh.train();
|
||||
voltron.vol_lh.train();
|
||||
voltron.vol_model.train();
|
||||
|
||||
# Use the adam optimizer
|
||||
optimizer = torch.optim.Adam([
|
||||
{'params': voltron.parameters()}, # Includes GaussianLikelihood parameters
|
||||
], lr=0.1)
|
||||
|
||||
# "Loss" for GPs - the marginal log likelihood
|
||||
mll = gpytorch.mlls.ExactMarginalLogLikelihood(voltron_lh, voltron)
|
||||
|
||||
for i in range(500):
|
||||
# Zero gradients from previous iteration
|
||||
optimizer.zero_grad()
|
||||
# Output from model
|
||||
output = voltron(train_x)
|
||||
# Calc loss and backprop gradients
|
||||
loss = -mll(output, train_y.log())
|
||||
loss.backward()
|
||||
# print(loss.item())
|
||||
optimizer.step()
|
||||
return voltron
|
||||
|
||||
def predict_prices(test_x, voltron, nvol=10, npx=10):
|
||||
ntest = test_x.shape[0]
|
||||
vol_paths = torch.zeros(nvol, ntest)
|
||||
px_paths = torch.zeros(npx*nvol, ntest)
|
||||
|
||||
voltron.vol_model.eval();
|
||||
voltron.eval();
|
||||
|
||||
for vidx in range(nvol):
|
||||
vol_pred = voltron.vol_model(test_x).sample().exp()
|
||||
vol_paths[vidx, :] = vol_pred.detach()
|
||||
|
||||
px_pred = voltron.GeneratePrediction(test_x, vol_pred, npx).exp()
|
||||
px_paths[vidx*npx:(vidx*npx+npx), :] = px_pred.detach().T
|
||||
return px_paths
|
||||
|
||||
def get_and_fit_basic_model(train_x, train_y, cov="matern", mean="loglinear"):
|
||||
voltron_lh = gpytorch.likelihoods.GaussianLikelihood()
|
||||
# voltron = VoltronGP(train_x, train_y.log(), voltron_lh, pred_vol)
|
||||
if cov == "matern":
|
||||
model = MaternGP(train_x,
|
||||
train_y.log(), likelihood=voltron_lh)
|
||||
else:
|
||||
model = SMGP(train_x,
|
||||
train_y.log(), likelihood=voltron_lh)
|
||||
if mean == "loglinear":
|
||||
model.mean_module = LogLinearMean(1)
|
||||
model.mean_module.initialize_from_data(train_x, train_y.log())
|
||||
else:
|
||||
model.mean_module = gpytorch.means.ConstantMean()
|
||||
|
||||
|
||||
model.likelihood.raw_noise.data = torch.tensor([1e-6])
|
||||
|
||||
# Use the adam optimizer
|
||||
optimizer = torch.optim.Adam([
|
||||
{'params': model.parameters()}, # Includes GaussianLikelihood parameters
|
||||
], lr=0.1)
|
||||
|
||||
# "Loss" for GPs - the marginal log likelihood
|
||||
mll = gpytorch.mlls.ExactMarginalLogLikelihood(voltron_lh, model)
|
||||
|
||||
for i in range(500):
|
||||
# Zero gradients from previous iteration
|
||||
optimizer.zero_grad()
|
||||
# Output from model
|
||||
output = model(train_x)
|
||||
# Calc loss and backprop gradients
|
||||
loss = -mll(output, train_y.log())
|
||||
loss.backward()
|
||||
# print(loss.item())
|
||||
optimizer.step()
|
||||
return model, voltron_lh
|
||||
|
||||
def predict_basic_prices(test_x, voltron, voltron_lh, npath=1000):
|
||||
ntest = test_x.shape[0]
|
||||
voltron.eval();
|
||||
mod = voltron_lh(voltron(test_x))
|
||||
px_paths = mod.sample(torch.Size(((npath),))).exp().squeeze(-1)
|
||||
|
||||
return px_paths
|
||||
@@ -0,0 +1,71 @@
|
||||
import matplotlib.pyplot as plt
|
||||
import seaborn as sns
|
||||
import numpy as np
|
||||
import torch
|
||||
import pandas as pd
|
||||
|
||||
import gpytorch
|
||||
from torch.nn.functional import softplus
|
||||
from voltron.kernels import BMKernel, VolatilityKernel
|
||||
from voltron.models import BMGP, VoltronGP
|
||||
import argparse
|
||||
from torch.distributions import Beta
|
||||
from scipy.special import betainc
|
||||
from Trainers import *
|
||||
|
||||
def main(args):
|
||||
full_data = pd.read_pickle("../../spdr-data/" + args.SPDR + ".pkl")
|
||||
tckrs = full_data.symbol.unique()
|
||||
|
||||
for tckr in tckrs:
|
||||
data = full_data[full_data["symbol"] == tckr]
|
||||
|
||||
ts = torch.linspace(0, data.shape[0]/252., data.shape[0])
|
||||
# train_x = ts[:ntrain]
|
||||
# test_x = ts[ntrain:(ntrain+ntest)]
|
||||
|
||||
y = torch.FloatTensor(data['close_price'].to_numpy())
|
||||
log_returns = torch.log(y[1:]) - torch.log(y[:-1])
|
||||
dt = ts[1] - ts[0]
|
||||
|
||||
eval_times = list(range(100, ts.shape[0], 100)) #+ [ts.shape[0]]
|
||||
prob_of_increases = []
|
||||
|
||||
for i, time in enumerate(eval_times):
|
||||
print("now running time: ", time)
|
||||
with gpytorch.settings.max_cholesky_size(2000):
|
||||
data_model, data_lh = get_and_fit_basic_model(ts[:time], y[:time],
|
||||
cov=args.kernel, mean=args.mean)
|
||||
end_ind = -1 if i + 1 >= len(eval_times) else eval_times[i+1]
|
||||
paths = predict_basic_prices(ts[time:end_ind], data_model,
|
||||
data_lh).detach()
|
||||
# now we predict the probability of increase at time i + 1
|
||||
prob_of_increase = (paths[..., -1] > y[time]).sum() / paths.shape[-2]
|
||||
print("prob of stock increase: ", prob_of_increase.detach())
|
||||
|
||||
prob_of_increases.append(prob_of_increase.detach())
|
||||
|
||||
torch.save(obj=prob_of_increases, f="./outputs/" + args.kernel + "_" + tckr + ".pt")
|
||||
print(tckr, "Done")
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
parser.add_argument(
|
||||
"--SPDR",
|
||||
type=str,
|
||||
default="XLE",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--kernel",
|
||||
type=str,
|
||||
default="matern",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--mean",
|
||||
type=str,
|
||||
default="loglinear",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
main(args)
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,591 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"[pyKeOps]: Warning, no cuda detected. Switching to cpu only.\n",
|
||||
"Warning no robinhood utils.\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns\n",
|
||||
"import torch\n",
|
||||
"import gpytorch\n",
|
||||
"# from voltron.robinhood_utils import GetStockData\n",
|
||||
"import os\n",
|
||||
"# import robin_stocks.robinhood as r\n",
|
||||
"import pickle5 as pickle\n",
|
||||
"\n",
|
||||
"sns.set_style(\"whitegrid\")\n",
|
||||
"sns.set_palette(\"bright\")\n",
|
||||
"\n",
|
||||
"sns.set(font_scale=2.0)\n",
|
||||
"sns.set_style('whitegrid')\n",
|
||||
"\n",
|
||||
"import sys\n",
|
||||
"sys.path.append(\"../\")\n",
|
||||
"from voltron.likelihoods import VolatilityGaussianLikelihood\n",
|
||||
"from voltron.models import SingleTaskVariationalGP as SingleTaskCopulaProcessModel\n",
|
||||
"from voltron.kernels import BMKernel, VolatilityKernel\n",
|
||||
"from voltron.models import BMGP, VoltronGP\n",
|
||||
"from gpytorch.kernels import ScaleKernel, RBFKernel, MaternKernel"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# with open(\"../../stock_data.pkl\", \"rb\") as handle:\n",
|
||||
"# raw_data = pickle.load(handle)\n",
|
||||
" \n",
|
||||
"with open(\"./stock_data.pkl\", \"rb\") as handle:\n",
|
||||
" raw_data = pickle.load(handle)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"array(['AAPL', 'F', 'JPM', 'SBUX', 'TSLA', 'VIRT'], dtype=object)"
|
||||
]
|
||||
},
|
||||
"execution_count": 3,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"np.unique(raw_data[\"symbol\"])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"# Header"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ntest = 200\n",
|
||||
"ntrain = 200\n",
|
||||
"tckrs = ['TSLA', \"F\", \"JPM\", \"SBUX\", 'AAPL', \"VIRT\"]\n",
|
||||
"tckr = \"VIRT\"\n",
|
||||
"span = \"5year\"\n",
|
||||
"interval = 'day'\n",
|
||||
"T = 5."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Data Setup"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 5,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAX8AAAEdCAYAAADkeGc2AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAA5JklEQVR4nO3deWBU1dn48e/sk8m+L6wJSRAIIERQtIAoLmhVigpVKy4vovXVWq212tfWVqq/1upbBbXufVvFDYtalUURWQyyGHYIgYQ9Ifs62Wa59/fHkIGYBLJNZpL7fP5x5q7nHi/PnDz33HN0qqqqCCGE0BS9vwsghBCi90nwF0IIDZLgL4QQGiTBXwghNEiCvxBCaJAEfyGE0CCjvwsghC/Mnz+ftWvXcuGFF/KPf/yjQ/vcd999fPXVV1x77bX89a9/5dZbb2Xz5s3cd9993H///d7tFi1axIsvvtjucQwGAzabjYEDBzJlyhTuuusuQkNDvesvueQSCgoKOnU9EydO5O233+7UPkKciQR/0S/NnDmTtWvXsmnTJsrLy4mOjj7j9rW1taxduxaAn/zkJx06R0hICOnp6a2Wu1wuCgoKyMnJIScnh88//5yPPvqIqKgoADIyMoiPj2+xj91uZ//+/QCMHz++1THbOo8Q3SHBX/RL06dPJywsjJqaGlasWMEtt9xyxu1XrFiBw+EgMTGRCy64oEPnGDly5Blb419++SUPP/wwBQUFPPXUUzz33HMALFy4sNW2mzZtYu7cuQC89957HTq/EN0hOX/RL5nNZmbMmAHAsmXLzrr9Z599BsB1112HXt8z/ywuv/xy7r77bgBWrlyJ3W7vkeMK0RMk+It+a+bMmQBkZ2dTXFzc7nbFxcVs2bIF6HjKp6OmTJkCgNPp5MiRIz16bCG6Q4K/6LfGjx/PkCFDUFWV5cuXt7vd559/jqIojBs3jqFDh/ZoGXQ6nfezDKMlAokEf9GvXXfddQB88cUX7W7TnPKZNWtWj5//yy+/BCAoKIi0tLQeP74QXSXBX/RrM2fORKfTsXPnTo4dO9ZqfV5eHjk5OVitVq666qoeO6/T6WTJkiW8+eabAMydOxeLxdJjxxeiu6S3j+jXBgwYwIQJE9i8eTPLly9n/vz5LdY3t/qnT59OSEhIp469d+9ebrrpplbL6+rqOHbsGPX19YDnB+j09wSECAQS/EW/N3PmTDZv3syyZctaBH9VVbuV8rHb7WzdurXNdUlJScycOZNrrrmmzX77QvibpH1Ev3fFFVcQFBRETk4OBw8e9C7funUrBQUFJCQkMGnSpE4fd+LEieTm5pKbm8u+ffv47rvvuP/++zEajZSVlTFw4EAJ/CJgSfAX/V5ISAjTp08HWvb578m+/TqdjqioKO677z6efPJJHA4HzzzzTIeHlhCit0nwF5rQ3H+/Ofg7nU5WrFjRYl1Puf7667n22msB+Otf/9puakgIf5LgLzRh0qRJJCQkkJ+fT15eHhs3bqSyspJx48aRnJzc4+d7/PHHiY2Nxe1289hjj+FwOHr8HEJ0hwR/oQl6vZ5rrrkGgNWrV7Nq1Srg1FvAPS08PJzf/va3ABw+fJhXX33VJ+cRoqsk+AvNaE7vrFq1itWrV2OxWLj66qt9dr6rrrrK+yD5tdde4/Dhwz47lxCdJcFfaMawYcMYPXo0O3bsoKSkhOnTp7cYZ98Xfv/732MymXA4HDz55JM+PZcQnSHBX2jK6Q93e/pBb1tSUlK48847AcjKyuLzzz/3+TmF6AidKqNNCSGE5kjLXwghNEiCvxBCaJAEfyGE0CAJ/kIIoUEBH/xVVaWpqUlmQRJCiB4U8MHf4XCwe/fuLr8ev2fPnh4ukegMqX//kvr3r0Cu/4AP/t3V2Njo7yJomtS/f0n9+1cg13+/D/5CCCFak+AvhBAaJMFfCCE0SIK/EEJokAR/IYTQIAn+QgihQRL8hRBCgyT4CyFEgKne/DlHnr8TVXH77BwS/IUQIsCUf/UP3HXVuGsrfHYOCf5CCBFAFGeT97Ozsshn55HgL4QQAaSpYL/3s7Oy2GfnkeAvhBABxFle4P3sqvJdy9/osyMLIYTolMr1S6hc9z46kwVDSKRP0z4S/IUQIgCoqkrluvcBMEUmYAiOwFVd5rPzSdpHCCECgOu0Vr7qdqG3BqM01fnsfBL8hRAiADQWHvB+NkbEe4J/Y73PzidpHyGE8ANVVdHpdN7vjqKD6Ixm4q9/GEtSGlUbP0Vp8l3wl5a/EEL0ssN/u4OSj5/DUXbcu8xZcQJjZAK21EwMtjD0lmBUlwPF1bUpbM9Ggr8QQvQypb6GupzvOP7qL2kqPsyJ9xZQf+B7TJHx3m30FptnWx+lfiT4CyGE36hUf/cJDQe3A55ePs0M1mAAnz30leAvhBC9SHU7W3y371nv/Xx68Ndbm1v+EvyFEKLPUxyN3s/muMEt1umDQk99bm75S/AXQoi+Tz0Z/GOuvpfYax/wLrelnYct5Vzvd72lOe3jm5y/dPUUQohe1Nzy15utWOKHYhkwnKAhI4ma9rMW2/m65S/BXwghepHiaAA8wR9gwO1Pt7mdN/j7qOUvaR8hhOhFzWkfnTnojNvpjGZ0Ziuq2+WTckjLXwghetHpaZ8z0el0JP1sAcbwWJ+UQ4K/EEL0oh+mfc7Ekpjis3JI2kcIIXqRN+1jOnPax9ck+AshRC9SnCfTPpazt/x9SYK/EEL0IsXb8rf4tRwS/IUQohepjgZ0Jis6nX/DrwR/IYToRYqjsUMPe31Ngr8QQvQipakevcW/D3tBgr8QQvQqV2URxvA4fxdDgr8QQvQWVVVxVpzAFJXo76JI8BdCCF9y1ZRRs20V4JnBS2mqxxSd5OdSSfAXQogepaoK9Yd2oKoqACfe/xNly/6Ou6GWmuyVAJgipeUvhBB9ntJYR9WGj1EVNzXZKyl690nqczcD4KooAqDh0E4q138AEBAtfxnbRwghuqli7fvUfL8MY2QCjuLDALhqy0Cf4J22sXme3ogf3dhiukZ/keAvhBBd1FiwH2dlEW57JQDOikJUlwMApakBTuvR2XBkNzqjmcgps/1R1FYk+AshRBdVrH6bxoL93pZ8U+EBlEbP5Cuu6lIISvZu66oqwZI4zO9v9jaT4C+EEF3gbrDTeGwfqArOsuMANBXmede7qksh1tliH1PskF4t45kExk+QEEL0MQ0Ht4OqeL/bUjNx2yu9KSBXTSn6xtoW+1gSkgkU0vIXQoguqM/LRh8UStTUn6K3haO3BFGflw2AOW4wzooi9HUVLfaxJA7zR1HbJC1/IYToJFVxU5+/DVvqeMIyryRkxCTMp6V0wjJnoLocWI9+D+i8y83xQ3u/sO2Q4C+EEJ3kLCtAaaglKHmsd5khJML7OXjEJECHqewgptiB6G1hAOj9PIb/6STtI4QQneSsPAGAOXqAd5lOpyPiwlnobaEYgkKxDjqHxmM5WBKHET39dlS321/FbZMEfyGE6CRnZTEAxh+8rBU17Rbv57iZD7L/87cYMOHHGIJCe7V8HSFpHyGE6ARHeSH23evQW0MwBIW0u50xLJqG4ZcEVA+f00nLXwghOuHE4j/gri1v1erva6TlL4QQnaG4AFCdTX4uSPdI8BdCiE4wRsQDEH/9w34uSfdI8BdCiE5QHA3Yhp+PdeA5/i5Kt0jwF0KITlAdjQExAXt3SfAXQohOUBwN6E1Wfxej2yT4CyFEJyjS8hdCCG1R3U5wu9CZJfgLIYRmKE2NAOjNfT/t0+GXvNxuN4sXL+ajjz7i0KFDBAUFkZGRwdy5c7n44otbbX/o0CEWLVpEdnY2VVVVDB48mDlz5nDzzTej18tvjhCi71GcDQDo+0HLv8PB/7HHHuPTTz8lJCSESZMm4XQ62bx5M1lZWfziF7/gv//7v73b7tu3j1tuuQW73c748eMZPXo0mzZtYsGCBWzfvp1nn33WJxcjhBC+pJ5s+eu00vJftmwZn376KcnJybzzzjvExMQAcODAAW666SZefPFFrr76aoYOHYqqqjzyyCPY7XaeeeYZrrvuOgAqKiq4/fbb+eyzz7jsssu44oorfHdVQgjhA4qzOe3T91v+Hcq//Oc//wHg4Ycf9gZ+gLS0NK655hoURSErKwuArKwscnNzmThxojfwA0RFRfHEE08A8Pbbb/fYBQghRG9RmjSW9lm4cCGHDx9m6NChrdbV1dUBYDAYAFi/fj0A06dPb7VtZmYm0dHRZGdnY7fbCQlpf0Q8IYTwl5rtq6jZ8gUD7/pbi+WqwxP8+0Pap0Mtf7PZTHp6OmazucXyb775hhUrVmCz2bzBPi/PM3t9enp6m8dKTk5GURTy8/O7U24hhPCZsi/+jqPkKIrL0WK54tBgb59mjY2NPPLII+Tl5ZGfn09SUhLPPPOMNx1UUlICQGxsbJv7Ny8vKyvrapmFEKJXKHXV6MNPxTLF0X/SPp3uc1lYWMjKlStbtNxzc3O9nxsaPJVjtbb9y9i8vL6+vrOnFkKIXqEzmABw11W3WO6qKQODEf0ZJnHpKzrd8k9ISGDjxo3o9Xo2bNjAU089xYIFC6ivr2f+/PnePvw6na7N/VVVbfHfjtq9e3dni+qVnZ3d5X1F90n9+5fUf+eF643o3U727czGdeLUD0DwwRwM1jC2btve4WP5s/4zMzPbXdfp4G+z2bDZbADMmDGDxMREfvrTn/Lqq69y2223edc1Nja2uX9TU5P3OJ2RkZGBxWLpbHHJzs4+YwUI35L69y+p/645khWE29lASmIsoWMzUVWV8lX/R01xLkHDxpPWwToN5Prv9qu25557LoMHD8Zut3Ps2DHi4uKA9nP6paWlQPvPBIQQwt/0Jk9Dsznt4ziRT83mzwEw9fHpG5udNfirqsozzzzDgw8+iMvlanOb5l5ALpeLtLQ04FSvnx8e6+DBgxgMBoYNG9adcgshhM80p6Xd9Z7gX7tLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"idx = -2\n",
|
||||
"# data = GetStockData(tckr, span=span, interval=interval)\n",
|
||||
"data = raw_data[raw_data[\"symbol\"] == tckr]\n",
|
||||
"\n",
|
||||
"ts = torch.linspace(0, T, data.shape[0])\n",
|
||||
"train_x = ts[:ntrain]\n",
|
||||
"test_x = ts[ntrain:(ntrain+ntest)]\n",
|
||||
"\n",
|
||||
"y = torch.FloatTensor(data['close_price'].to_numpy())\n",
|
||||
"train_y = y[:ntrain]\n",
|
||||
"test_y = y[ntrain:(ntrain+ntest)]\n",
|
||||
"\n",
|
||||
"dt = ts[1] - ts[0]\n",
|
||||
"\n",
|
||||
"plt.plot(train_x, train_y)\n",
|
||||
"plt.plot(test_x, test_y)\n",
|
||||
"plt.title(tckr);\n",
|
||||
"sns.despine()\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 6,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"log_returns = torch.log(train_y[1:]/train_y[:-1])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 7,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAn4AAAGLCAYAAABZZ59XAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAA9hAAAPYQGoP6dpAADKTUlEQVR4nOydd7wU1d3/P7N99/YGXHpR9NoARcWosRfs8mjsRJ8kxhg1PiaPEZOfPdFYYnwsscSgqImigr0rahAVUJF2qVJvgdvL9p2d3x+z58yZ2dm+t7Hf9+vFi727szNnp53PfKukKIoCgiAIgiAIYo/HMtADIAiCIAiCIPoHEn4EQRAEQRAFAgk/giAIgiCIAoGEH0EQBEEQRIFAwo8gCIIgCKJAIOFHEARBEARRIJDwIwiCIAiCKBBI+BEEQRAEQRQIJPwIgiAIgiAKBBJ+BEEQBEEQBQIJP4IgCIIgiAKBhB9BEARBEESBYBvoARAEQRQiiqIgHA4jGo0O9FAIghjkWCwW2O12SJKU87pI+BEEQfQjPp8PXV1d6OnpgSzLAz0cgiCGCFarFSUlJSgrK4PH48l6PZKiKEoex0UQBEEkoKenBzt37oTdbkdpaSmKiopgsVjy8hRPEMSeiaIoiEaj8Hq96O7uRjgcxujRo1FSUpLV+kj4EQRB9AM+nw/btm1DaWkpRo4cSWKPIIiMURQFjY2N6O7uxrhx47Ky/FFyB0EQRD/Q1dUFu91Ooo8giKyRJAkjR46E3W5HV1dXVusg4UcQBNHHKIqCnp4elJaWkugjCCInJElCaWkpenp6kI3TloQfQRBEHxMOhyHLMoqKigZ6KARB7AF4PB7IsoxwOJzxd0n4EQRB9DGsZIvFQrdcgiByx2q1AkBW5aDoLkQQBNFPkJuXIIh8kMu9hIQfQRAEQRBEgUDCjyAIgiAIokAg4UcQBEEQBFEgkPAjCIIgCIIoEEj4EQRBEARBFAgk/AiCIAiCIAoE20APgCAIgiD6m8suuwxLly4FAHz88ccYPXr0AI+o79i5cydOOOGEtJZ1OBy8n/TUqVNx8cUXY8KECX08QrW7zebNm7HXXnv1+bYKHbL4EQRBEAQBAAiFQmhtbcXKlSsxb948nHnmmXjhhRf6dJvr1q3DRRddhH/84x99uh1ChSx+BEEQBFEg7L///vjTn/6U8HOv14tNmzZh4cKFWLFiBcLhMO666y6MHTsWRx99dJ+MadasWZBlGePHj++T9RN6SPgRBEEQRIFQVFSEurq6pMtMnz4dF1xwAe6880688MILiEajePDBB/tM+Mmy3CfrJcwhVy9BEARBEDokScJNN92EqqoqAMCaNWuwefPmAR4VkQ9I+BEEQRAEEYfD4cAhhxzC/966devADYbIG+TqJQiCIIgc6OrqwksvvYRPP/0UmzdvhtfrRWlpKSZPnoyTTjoJ5513HpxOZ9J1bNmyBc888wy+/PJLNDY2wu12Y/LkyZg1axZmzZqFp556Cg888AAAYP369f3xswDo3bChUCjhcuvWrcO//vUvfP3119i9ezcURcHw4cNx6KGH4qKLLsL+++8f95199tlH9/fChQuxcOFCAMDdd9+NWbNm6ZY7+OCD8e9//zvhGI4//ng0NDRg+PDh+Pzzz00/O++883DjjTfi7rvvxieffIJQKIThw4fjvPPOwy9+8QvcdNNNWLhwIcaOHYsPP/wQjY2NeOaZZ/D555+jubkZNpsNEydOxMyZM3HxxRcnPK6hUAivvPIKPvroI9TX16O7uxtFRUWora3FYYcdhgsuuGDAMphJ+BEEQRBElnz88ce4+eab0dnZqXu/ra0NX375Jb788kv84x//wEMPPYSDDjrIdB3vvPMObrzxRoTDYf5eOBzG8uXLsXz5crz11luYOnVqH/4Kc0KhEFasWMH/njRpUtwy0WgU9913H+bOnQtFUXSfbd26FVu3bsXLL7+Myy67DDfddBNstoGVHYFAAD/96U9RX1/P39u6dSvKysrilv3ss89www03oLe3V/f+999/j++//x6vvPIKnnvuOVRWVuo+37VrF37+859jw4YNuve7urrQ1dWFdevW4fnnn8evf/1rXHPNNXn8delBwo8gCGKQoSgKgqHCC3h3OqyQJGmgh5E2n376Ka677jpEIhFYLBacffbZOPnkk1FdXY3Gxka8/vrr+OSTT9DY2IjZs2fjxRdfxL777qtbx6JFi3DDDTdAURQ4nU5ceuml+PGPfwy73Y5vvvkGTz/9NJYsWYLly5f3++97+OGH0dbWBkC1uk2ePDlumdtuuw0vvfQSAGDy5Mm46KKL+G9kVsCNGzfiueeeQygUwh133MG/+9prrwEAzjnnHADAcccdh9/85jcAgNra2j75Te+++y5kWcapp56KCy+8EJFIBB9//DFOP/103XLt7e34zW9+g0gkgssuuwzHHXccPB4PVq9ejSeeeAItLS3YtGkT7r33Xtxzzz26786ZM4eLvv/6r//CSSedhOrqanR3d+Pbb7/Fs88+i56eHjz88MOYMmVKnyXNJGKPEn7RaBQLFy7Ea6+9hvXr18Pn86GmpgYHH3wwLrzwQhx66KF53+Y333yDSy+9FDU1NXGm5USsXbsWzzzzDJYtW4aWlhYUFxdjwoQJOOOMM3D++efD4XDkfZwEQQwNFEXB7x9ZjPqt7QM9lH6nbnwl/nLNUUNC/Hm9Xtx8882IRCKw2Wx49NFHceyxx/LPDzroIJx66ql44YUXcMcdd8Dv9+OGG27AW2+9BYtFDa8PBAK48847oSgK3G43nn32WUyZMoWv45BDDsGZZ56JSy+9FDt37uzz3xQMBtHV1YW1a9fi5ZdfxkcffQQAsNvtuPXWW+OWX7RoERd955xzDv70pz/pLHoHH3wwzj//fNxwww344IMP8NJLL+HUU0/Fj370IwCIyy4uLy9PmXGcK7Is49hjj8VDDz3E3zMTXr29vXA4HHjmmWcwffp0/v60adNwzDHH4IwzzkAwGMQ777yDW265BR6PBwDQ2NiIL774AgBw/vnn46677tKt98gjj8SRRx6Jiy66CADw4osvkvDLlp6eHlx99dW8EjujsbERjY2NePvtt3H55Zfjpptuyts229vbMWfOHESj0bS/M3fuXNx33326uImOjg50dHTg22+/xfz58/HEE09gxIgReRsnQRAEkV8WLFjArWE///nPdaJP5JJLLsHy5cvxzjvvYPPmzfjkk09w4oknAgA++OADNDQ0AACuvvpqnehj1NbW4s9//jNmz56dl3EvXbo0LrYuGdXV1fjLX/6iS/JgPPXUUwCAqqoq3H777aZuXLvdjjvuuAP/+c9/4Pf78cwzz3DhN1BccsklaS03a9YsnehjjB07FkcccQQ+/fRTBINBbN++nVs5W1pa+HLjxo0zXe/BBx+MX/7yl7BYLKZW1L5mjxB+iqLg+uuv56LvqKOOwkUXXYTq6mrU19fjqaeeQkNDA+bOnYvKykpceeWVOW+zs7MTV1xxBbZt25b2d958801uEh42bBiuuuoq7L///mhvb8f8+fOxaNEirFu3DldddRVeeumllMHABEHseUiShL9ccxS5egc5oofn4osvTrrsJZdcgnfeeYd/jwm/Tz75BIB6zM8777yE3z/88MMxadKkfiunUl5ejgMOOAAnnHACzj77bBQVFcUt093dje+++w6Aapl0uVwJ11dRUYGDDjoIX3/9NZYtW8atpANFuvGSySxxY8aM4a+9Xi9/PXbsWNjtdoTDYTzxxBOorKzEzJkzuUWQccMNN2Q26DyyRwi/N998E4sXLwagKvS7776bfzZ16lTMnDkTl1xyCTZt2oRHHnkEZ511Vk4Wtfr6elx33XXYvn172t/p7e3l1dKHDRuGV155BcOHD+efH3/88XjggQfw5JNPor6+Hs8//zx+9rOfZT1GgiCGLpIkweXcI27PeywbN24EoFrkxHu5GQcddBCsVitkWdYF/LMEg1GjRsUlCBiZMmVKXoSfsXOHLMvo7OzEBx98gFdeeQWyLKOmpga//OUvcdhhhyVcz9q1a7m364MPPkjbiujz+dDW1pZyn/UVHo8HpaWlaS2brH+zKIZFD15FRQUuvPBCPPfcc+jp6cHNN9+M2267DQcffDB+9KMf4YgjjsABBxzA3f0DwR5Rx2/u3LkAgOLiYvz+97+P+7y8vBy33347ADWGYd68eVltJxgM4sknn8QFF1zARV+6B2/BggXo6OgAAFx33XWmJ/3111/Pm2HPnTs3IxeykfXr1/dryj9BEEQhwe7nrMBxMhwOBxcb7HuA5hZMJfrS3U46sM4d7N8BBxyAo446CnfccQcee+wx2O12bNy4ET/96U+5ldIM8XdkijEDuj8pLi5Oe1m3253WcsZs5ptuugm/+MUvYLfbAajZ0V999RX++te/4vzzz8eRRx6JW265BZs2bUp/4HlkyD9S7tixA2vXrgWgZgSVl5ebLjd9+nRMmDABW7ZswXvvvYcbb7wxo+1s27YNl19+ORobGwEATqcTt956Kx599FEeo5GM999/H4Aa72DMHmJYrVbMmjULDzzwAFpaWrB8+fKkT1zJSFZvKVuCwSBWr16NAw44gNzQgxw6VkOHaDQKv98Pt9s9oFYAIjXisTJO9qlgy4vHmJVvyeUhP58ce+yx+OMf/4hbb70V0WgUN954IyorKzFjxoy4ZUUr1+zZs3nNvXRIFPuWK+J+9Hq9ptdUf4QS2Gw2/O53v8N///d/48MPP8SiRYuwdOlS7hJub2/HSy+9hFdeeQW33HILLrzwwj4fk258/bq1PuCbb77hr81OTpHDDjsMW7ZsQUNDA7Zv346xY8emvZ3m5mYu+g477DDccccdmDBhAh599NGU341EIvj+++8BqOZ6o69fRMw8XrJkSdbCry9gFzr1VRz80LEaOjBBkKmQIPof8ViVlZVh9+7dPMEjGcFLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 600x400 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(dpi=100)\n",
|
||||
"ax.plot(log_returns, label='Log Returns')\n",
|
||||
"\n",
|
||||
"ax.set_ylabel(\"Log Returns\")\n",
|
||||
"\n",
|
||||
"fig.legend()\n",
|
||||
"sns.despine()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Now apply GCPV"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 9,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"likelihood = VolatilityGaussianLikelihood()\n",
|
||||
"likelihood.raw_a.data -= 6\n",
|
||||
"\n",
|
||||
"# covar_module = MaternKernel(nu=0.5)\n",
|
||||
"covar_module = BMKernel()\n",
|
||||
"model = SingleTaskCopulaProcessModel(\n",
|
||||
" init_points=train_x[:-1].view(-1,1), likelihood=likelihood, use_piv_chol_init=False,\n",
|
||||
" mean_module = gpytorch.means.ConstantMean(), covar_module=covar_module, learn_inducing_locations=False\n",
|
||||
")\n",
|
||||
"# model.covar_module.base_kernel.raw_lengthscale.data = torch.tensor([[-5.]])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 10,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# this is for running the notebook in our testing framework\n",
|
||||
"import os\n",
|
||||
"smoke_test = ('CI' in os.environ)\n",
|
||||
"training_iterations = 2 if smoke_test else 500\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"# Find optimal model hyperparameters\n",
|
||||
"model.train()\n",
|
||||
"likelihood.train()\n",
|
||||
"\n",
|
||||
"# Use the adam optimizer\n",
|
||||
"optimizer = torch.optim.Adam([\n",
|
||||
" {\"params\": model.parameters()}, \n",
|
||||
" # {\"params\": likelihood.parameters(), \"lr\": 0.1}\n",
|
||||
"], lr=0.01)\n",
|
||||
"\n",
|
||||
"# \"Loss\" for GPs - the marginal log likelihood\n",
|
||||
"# num_data refers to the number of training datapoints\n",
|
||||
"mll = gpytorch.mlls.VariationalELBO(likelihood, model, train_y.numel())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 11,
|
||||
"metadata": {
|
||||
"tags": []
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Iter 1/500 - Loss: -2.493\n",
|
||||
"Iter 51/500 - Loss: -2.563\n",
|
||||
"Iter 101/500 - Loss: -2.565\n",
|
||||
"Iter 151/500 - Loss: -2.566\n",
|
||||
"Iter 201/500 - Loss: -2.566\n",
|
||||
"Iter 251/500 - Loss: -2.566\n",
|
||||
"Iter 301/500 - Loss: -2.566\n",
|
||||
"Iter 351/500 - Loss: -2.566\n",
|
||||
"Iter 401/500 - Loss: -2.566\n",
|
||||
"Iter 451/500 - Loss: -2.566\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"print_every = 50\n",
|
||||
"for i in range(training_iterations):\n",
|
||||
" # Zero backpropped gradients from previous iteration\n",
|
||||
" optimizer.zero_grad()\n",
|
||||
" # Get predictive output\n",
|
||||
" with gpytorch.settings.num_gauss_hermite_locs(75):\n",
|
||||
" output = model(train_x[:-1])\n",
|
||||
" # Calc loss and backprop gradients\n",
|
||||
" loss = -mll(output, log_returns)\n",
|
||||
" loss.backward()\n",
|
||||
" if i % print_every == 0:\n",
|
||||
" print('Iter %d/%d - Loss: %.3f' % (i + 1, training_iterations, loss.item()))\n",
|
||||
" optimizer.step()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 12,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"model.eval();\n",
|
||||
"likelihood.eval();\n",
|
||||
"predictive = model(train_x)\n",
|
||||
"pred_scale = likelihood(predictive, return_gaussian=False).scale.mean(0).detach()\n",
|
||||
"samples = likelihood(predictive, return_gaussian=False).scale.detach()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 16,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"[<matplotlib.lines.Line2D at 0x7fe730550a50>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7fe730550ad0>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7fe730542790>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7fe730550d90>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7fe730550e90>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7fe730550dd0>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7fe730550ed0>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7fe730550fd0>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7fe73055c310>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7fe730550f90>]"
|
||||
]
|
||||
},
|
||||
"execution_count": 16,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAbkAAAEKCAYAAACPCivzAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAADTD0lEQVR4nOy9eZxkdXU2/tStfd+ru7qql+nuGbZhcFgHlWERBI0rosgqGMUNTUw0vBp+YjDEqK+JkBBReTERjCKoaKKC4IAgowMMyywMM9P7Uvu+L7fq/v7onDO3qqu6q2fpmSH1fD5+Brurbte9de/3fM85z3kehSRJErrooosuuujidQjhaH+ALrrooosuujhS6Aa5LrrooosuXrfoBrkuuuiiiy5et+gGuS666KKLLl636Aa5LrrooosuXrfoBrkuuuiiiy5et1Ct5MVbt27FPffcg71796JareKUU07BTTfdhPPOO6/jY4TDYdx999149tlnEY1G4fV68a53vQsf/ehHodFolnyvJEn40Ic+hGAwiMcff7zla/L5PO6991785je/wfz8PBwOB84++2x87GMfw+joaMv37Ny5E3fffTd27tyJQqGA0dFRXH/99XjnO9952M+hiy666KKL1YOi0zm5n/3sZ/jCF74AjUaDTZs2oV6vY9u2bahWq7j99ttx5ZVXLnuMUCiEK6+8EqFQCCeffDL6+/vx4osvIhqN4uyzz8Z9990HtVrd9v1f+9rXcN9992FgYKBlkEulUrj22muxf/9+GAwGnHbaaajVanj55ZchCAL+6Z/+CW95y1sa3vPss8/iYx/7GOr1Os466yzo9Xr88Y9/RKlUwsc//nF89rOfPazn0EUXXXTRxSpC6gDhcFhav369dMYZZ0h79+7ln7/yyivS6aefLp166qlSKBRa9jgf+9jHpHXr1kl33303/yyfz0s33HCDtG7dOun//b//1/J9pVJJ+uIXvyitW7dOWrdunXTxxRe3fN3nPvc5ad26ddJ73/teaX5+nn++b98+afPmzdLGjRulcDjMPy8Wi9K5554rnXLKKdIf//hH/vn09LS0efNmad26ddLOnTsPyzl00UUXXXSx+uioJ/fAAw+gUqnghhtuwLp16/jnGzZswEc+8hGUy2U8+OCDSx5jYmICTz31FAYGBvDxj3+cf24wGHDHHXdAqVTigQceWPS+p59+Gu9973vx8MMPo7+/v+3xc7kcfv3rX0OpVOIb3/gG+vr6+Hdr167FLbfcwqVMwi9+8QvE43G8853vxKZNm/jnAwMD+Ou//msAwP3333/I59BFF1100cXRQUc9uWeeeQYAcPHFFy/63SWXXIJvfetbePrpp/GZz3ym7TH+8Ic/QJIkXHjhhRCExtja19eHk08+GTt37sTY2FhD7+yjH/0oBEHAddddh6uuugpvf/vbWx5/YmICoihizZo1GBkZWfT7c845B8BC0PziF7/YcF7NJUwAuOiii6BUKvH0008f8jm0Q71eRz6fh1qthkKhWPb1XXTRRRddLPAzqtUqjEbjorW4GcsGOUmSMDY2BkEQMDw8vOj3Q0NDEAQBY2NjkCSp7WI9NjYGYCGraoXh4WHs3LkT+/btawgQl156KT71qU/hhBNOwNzc3JKfEwCMRmPL3yuVSgDA9PQ0RFGESqXC/v37AaAhOyWYTCZ4PB4Eg0HEYjG4XK6DPod2yOfz2Ldv37Kv66KLLrroYjHWrVsHs9m85GuWDXLpdBqVSgUOh6Mlc1ClUsFutyMejyOfz8NkMrU8TiQSAQB4PJ6Wv3e73QCAWCzW8PO77rpruY8IABgcHIQgCBgfH0cqlYLNZmv4/fbt2wEsZE/pdBpOpxPRaLThb7f6TPIgd7Dn0A5EUFm3bt1BszJ37dqF9evXH9R7/7eie81Whu71Wjm612zlWMk1q1Qq2LdvX0ckv2WDXLFYBADo9fq2r9HpdACwZJCj49Br2x2jUCgs95Fawmaz4YILLsCWLVtwyy234Bvf+AYsFgsAYGZmBv/4j//Ir61UKgf1mQ73OVDWq9FooNVqO3pPKxzKe/+3onvNVobu9Vo5utds5VjpNeukzbNskFuu3gkcKBV2cpx2H4qO0cmx2uHLX/4y9u7di6eeegqXXHIJNmzYgHK5jFdeeQXr16+HTqfDvn37oFItnLZSqVyyxNr8mY7UOezatWtFr28GZalddI7uNVsZutdr5ehes5XjSFyzZYOcwWAAAJTL5bavod8tle3RcUql0kEfYzn09PTg4Ycfxr/+67/i8ccfxx//+Ef4/X588pOfxI033sjD3VTD1ev1yGQyKJfLLXcQ9Jnosx+pc1i/fv1B7/q2b9+OM84446De+78V3Wu2MnSv18rRvWYrx0quWblc7jg5WDbImUwmGAwGJJNJJmzIIYoikskktFotlwdbgfpY7fpV1B9r1+/qFA6HA1/60pfwpS99qeHnlUoFc3NzcDgcXFb0eDzIZDKIRqPw+/1tPxP12lbrHLrooosuujg8WLYWqVAoMDo6ilqthqmpqUW/n5ycRL1eb8lQlIMYicRQbMb4+DiA1kzHTrFr1y78/ve/b/m7F154AaIoNjQ26TPR35Yjl8shEonA4XDA5XKt2jl00UUXXXRx+NDRMDhpUz7xxBOLfkc/O//88zs6xpYtW1Cv1xt+FwgEsGfPHvh8vo6o9+3wla98BTfddBOmp6cX/e7HP/4xAOCyyy5b9JlandeWLVtQq9Uazms1zqGLLrrooovDh46C3OWXXw6tVovvfe97DXXQnTt34t5774VOp8PVV1/NPw8EAhgfH0cikeCf9ff347zzzsPk5CTuvPNO/nmhUMCtt96KWq2GG2+88ZBOhoa6v/GNb6BarfLPf/jDH+Kxxx7D4OAg3v3ud/PPL730UjidTvz85z9vyABnZ2fxzW9+EwqFAjfccMOqnkMXXXTRRReHDx0pnvj9ftxyyy24/fbb8cEPfhCbNm2CJEnYtm0bRFHE1772NTidTn79Lbfcgueeew4333wzPv3pT/PPb7vtNlx11VW45557sGXLFqxZs4bFjTdv3oyrrrrqkE7mhhtuwG9+8xs8/vjjuOyyy3DyySdjamoK+/btg8PhwN13393QUzSZTPjKV76Cz3zmM/jYxz6Gs846C0ajEX/6059QLBbx2c9+FieeeGLD3zjS59BFFweDWCyGarUKr9d7tD/KItTr9Y5Y2l10cSTQ8Z13zTXX4J577sFpp52G7du3Y9euXTj99NPx/e9/vyE7Wgr9/f146KGHcPnllyORSOCpp56C1WrFX//1X+Nf//VfF5FaVgqNRoN///d/xzXXXINKpYKnnnoK5XIZ1157LR555JGWSiVvectbcP/99+NNb3oT9uzZg+effx4nnHACvvWtbzXoU67WOXTRxcGgUCggn88f7Y+xCJlMBuPj4xBF8Wh/lC7+l6Jjq50uDj+IBtsdIVhdvB6v2djYGGq1GtauXXvYs6ZDuV6RSATJZBIul6uh2vN6x+vxHjvSOJgRgk7Wzm4NoYsujnPU63XUajUAOOYyJlIXSqfThyT0cLSQSqUa+vtdHH/oBrkuujhGUC6XOVitBPJF+FhbkKvVKpRKJarVKpdTK5XKMReMW6FarSIcDiOdTh/tj9LFIaAb5Lro4hhArVbD9PQ04vF4y9+LoohIJLJodAVoDGzHUvAgOxSr1QqVSoV0Oo16vY6ZmRnMzMwcVECPx+OYnJxsmxVms9kl1ZlWAlI2OpauaRcrRzfIddHFMYBsNgtJktqSR8LhMJLJZEvx72M1k6tWq5AkCRqNBlarFblcDslkErVaDaIoIhAIrKiEKUkSkskkKpUKstnsot+XSiUEAgHMz8+33Aw0o1gsLvm6bpB7faAb5Lro4hgALdqtSnn5fB65XA7AAScMOarVKgRBgEqlOqYWZAq4FOQUCgVisRg0Gg08Hg8KhULL82mHbDaLWq0GQRCQTCYX/T4SiXBplCT22iEWi2FmZgYTExMtjwUcCHIHk3F2ceygG+S66OIoQxRFFAoFFg6XZ2uZTAbBYBAajQZ6vb5tJqdWq6FWq4+pTI5IJ/TZyNDYZrPBYrFAoVA0ZK71en1RZpXL5TjbS6VS0Gg0cLvdKJVKnBXS64rFIlwuF+x2O1KpVNuyZTqdRjweh9lshk6nQyQSQTgcRqlU4k2CJEndTO51gm6Q66KLowwiNrhcLiiVSuTzedTrdQSDQQ5wfX19MBgMKJfLiwJBpVLhQNJuQRZFcdXZjfIME1gQT9fr9bBYLBAEAXq9viHIBYNBTE1N8Tlks1nMz88zw7FYLMJqtcJisUCr1SISiWBqaoqNkFUqFaxWK5xOJwRBaFBcIlQqFYTDYRiNRni9Xvj9fjgcDqRSKUxPT2NmZoZfV6/XoVarUavVjktmaBcL6Aa5Lro4iqjVakgmkzAajdBoNDAYDMjn85ienkYmk4HL5UJ/fz+0Wi30ej0kSUIikcDMzAxnbdVqFRqNBiqVivtgctTrdUxOTiKVSq3quVUqlQbHe71ej4GBASiVSgCA0WhEuVzmoFYsFlGtVjE3N4dqtcpuH/JyrclkgiAIGBwchM/nYxeUfD4Ps9kMhUIBpVIJq9WKbDbbkNlKkoRQKARBENDb28u+kG63GwMDA3A4HKhWqyiXy5zFmUwmSJLULVkex+gGuS66OIpIJBKo1Wps50R2VWq1Gj6fD06nkxdjvV4PhUKBeDyOYrGIRCLBGRplcq0WZMr+2vkgHk6Uy2UEAgFMTEygWCxCrVa3fS2VL/P5PKrVKmq1GiwWC6rVKiYnJ1GpVKDValEoFJDL5aDRaDhoKhQKmEwm6HQ6xONxSJLUYPVlt9sBgLO5Wq2G+fl5FItFeDyeRcpEer0eNpsNwAH1GKVSyd6Q3ZLl8YuuBlUXXRwlUBZC5TdgIXNo52IhCAIMBgNEUYRarUYmk+HfabVaDm60sJOLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(train_x, pred_scale)\n",
|
||||
"plt.plot(train_x, samples.t(), color = \"gray\", alpha = 0.3)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 17,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"[<matplotlib.lines.Line2D at 0x7fe72f445750>]"
|
||||
]
|
||||
},
|
||||
"execution_count": 17,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAZYAAAEFCAYAAADACsF7AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAACMPUlEQVR4nO29eZwcdZ3//6qrr7lnct/nBHKSBEi4hAABRRFFBMQLlC+L67G7ui7qYz0W1mt/7q4CurjionKJKAgIiIQbArmAkDuZ3MlkJnPP9F3H5/dH1afqU9XV10z3TM/k8/wnme7q6k9/qurz/rxvgRBCwOFwOBxOiRBHegAcDofDGVtwwcLhcDicksIFC4fD4XBKChcsHA6HwykpXLBwOBwOp6TIIz2A4cQwDMRiMSiKAkEQRno4HA6HMyoghEBVVVRVVUEU8+sjp5RgicVi2Lt370gPg8PhcEYlzc3NqKmpyXvcKSVYFEUBYE5OIBAo+vPbt2/H4sWLSz2sMQ2fs+Lg81U8fM6KYzDzlU6nsXfvXnsNzccpJVio+SsQCCAYDA7qHIP93KkMn7Pi4PNVPHzOimOw81WoC4E77zkcDodTUrhg4XA4HE5J4YKFw+FwOCWFCxYOh8PhlBQuWDgcDodTUrhg4XA4HE5J4YJlBLjnsffw7t6TIz0MDofDKQtcsIwAR9sHcKIrPtLD4HA4nLLABcswQwgBIYBh8MadHA5nbMIFyzBDBYpuGCM8Eg6HwykPXLAMM7olWLjGwuFwxipcsAwzBhcsHA5njMMFyzBjEGoK44KFw+GMTbhgGWa4xsLhcMY6XLAMM7rBNRYOhzO24YJlmOEaC4fDGetwwTLMcB8L51TmZHcc3f3JkR4Gp8xwwTLMcFMY51TmsZdb8Mz6gyM9DE6Z4YJlmOGmMM6pTFrVoao8OXiswwXLMGMLFsIFC+fUwyCEa+tlYt/RHryxtXWkhwGAC5ZhxzaF6XzXxjn1IKQyN1W7Dnbjz6/sH+lhDIl393bgtXePj/QwAHDBMuxw5z3nVMYgpCLNwC3Hekd9KwvDIBUjtLlgGWa4j4VzKmMYlSlYKlXgFYNeQXPLBcswY/CoMM4pTCXtqlkIqcxxFYNuVI7/iguWYYbevKP9Jj5VeGb9Qew62D3SwxgzVGovIlOTMgXMaEXXjYqZWy5YhhluChtdbNl1EnuOcMFSKio1KowOaRTLFRgVJLS5YBlmeILk6EIzDOg6v1alopJNYcDotiTohlEx4+eCZZjhGsvowtArcyEcrVSyKQwY3Rs+QzfbnleCOY8LlmHGDjfmu+CKh1hmm0pcCEcrlRp9NRY2fJXUnZYLlmGGZ96PHujzWQkP6ljBqNDoq8Fc60rTbiopR44LlmGmknYVnNwYhlkdQa/AhXC0QipUAyzWx3K0fQDf+9Wb6I+lyzmsoqBWkEqYXy5YhhnHlstLulQ6lfSgjhVMU9jwf69uELz89jGomu77vp0GUOC17u5PQtcJegYqpwUAXVMqQSPkgmWYcWy5leFk42RH49plyRmpWmGtHVH87a3D2He01/f9Yn0s1PKQVv0F1UhQSRGnXLAMM+xDVQHXn5ODsRApVEkQYkYtjUQB1nyCwPaxFCj06L2RTFeOYKmkAAS53F+wfv163HPPPdizZw9UVcWiRYtwyy234IILLijo85qmYfny5Uin/W2ZEydOxKuvvlrKIZcVdpEyDAOSKI3gaDi5oAtgJTyolYhuENz9h3dx6dkzsGhOU97ji128Swm9hmnNX6iRIh3f9N5IVZBgqST/bVkFy2OPPYZvfvObCAQCWL16NQzDwIYNG3DzzTfj9ttvx3XXXZf3HC0tLUin05gxYwaWLVuW8X59fX0ZRl4+2IuuGwTKCI6Fk5tKelArkVRaQ3t3HG1dscIEi8cMLAhCuYeY8d3ZNZbRbwqrJA27bILl5MmT+O53v4uamho89NBDaG5uBgC89957uOmmm/D9738fF110ESZOnJjzPLt27QIAXH311fjCF75QruEOG4ZLYxn5G+BUghCC9dtOYOVpExAK5L/1K8lmXYmkrU6QWoGmLdanSAgwjHLFdmxnFSxFpgHQ4yrJFFZJG6Gy+VgeeOABpNNp3HjjjbZQAYClS5fi5ptvRiqVwiOPPJL3PDt37gQALFq0qFxDHVZcPpYKuAFOJU72JPD06wex51BPQcdXUpRNJUIXaTWLecmL2784vHNKI9FSWdoi0+EU+kxqVsRgqgI1lkq4X8smWF577TUAwKWXXprx3tq1awGgIN8I1VjGjGDxmMI4w0cqrQEA1AJ32IWEG3f0JIY+sFFK2grdLVRjGUltnS62eTWWLOPauq8DbV0x5vjcGtBIYG+EKmBdKYtgIYSgpaUFoihizpw5Ge/PmjULoiiipaUlZ8gtIQS7du3C+PHj8eKLL+Kaa67B8uXLsXr1anz1q1/FgQMHyjH8ssI+g5VwA5xKUNNNoVFJ+WzWR9sH8N8Pv+1acE4l6HwWrrEw/x9uwVKgjyXbtX7y1QN4c9uJjPMlrc3KSEOY/KBK2LCWRbD09fUhnU6jvr4egUAg431ZltHQ0IBEIoFYLPtDefToUUSjUXR0dOA73/kOgsEgVq1ahWAwiKeffhrXXHMNtmzZUo6fUDZYNbUSboBTCWq20Ao2d+Q2hcWTqvVvZSwuw42tsRQoWMgI3vuOYPEfay6NhRCCZFpzaWZ0/JUSFebSBivAFFYW530iYZoHwuFw1mNCoRAAIBaLobq62vcY6l+ZOHEifvnLX+L0008HYIYg/+d//if+7//+D//0T/+E559/HsFgsODxbd++veBjvQxVkB06FMXAgDk/W997D/VVZY/4HnEqRfjvP5HEwMAA9u3bj5B6Iu/xx7vSGBgYANFivr/hSEcKAwMD2LFzJ3raMjdQg6VS5isfB9vN+Tx6PIUtW/JrbbGkjoGBAQDA2++8i0iwdPvafHPWYl37o8eSvmPt6OzGQFTHjp270NvuvpaaTtDfP4AjR9PYssUc/0HrOTbPN/LmUE0n9txu374Dncdz34/lvsfKsqqJYv4bppCs88svvxwvv/wyRFF0RY/Jsoyvf/3r2LBhA3bs2IF169bhgx/8YMHjW7x4cVGCiLJlyxasXLnS/nvHgS4AKCjUknI8vh9He9oAAAsXLsbkcVVFj2M04Z2zctHdn8STr+7HJy4/DUHFPzdI234C7x45gBkzZ2Dlyul5z1l1qBtvtexCbXXA9zeEDnRh04HdaG5egAUzG4f8G4Dhm69SQHa14+1DLWgaV4eVKxfnPb53IIW/bt0MwAziqasu/hn0o5A5I7vasfVICxqaarBy5dKM91/d+zZ0IYEFCzKv5UA8jZrNmzB+QgNWrlwIAGiNH8CR7hOoa6jCypVnlOR3DIVkSsMTmzcAABacdjrmTavPeuxg7rFUKlXUhrwsprBIJGIPJhv0vVxajSAImDx5sm9IsiiKuPDCCwEMTQMZCm+814rX3z1e1GfYOkmVoLKOFY62D2DvkV5092Wv3VSsjyVf+Cabl3EqQv0VBTvvRzAiMq8pjGS/ltTcxfqSbFNYhTjv9REMjPCjLIKluroakUgEPT090LRM+7Omaejp6UEwGERtbe2gv2fcuHEAgGRyZArBGUbxbVZdUWEjUNpirEKvQ65FrlgfS74om0oK7xwJig43HkE/gJ3QmK0IZY5rmfIRoJUWFaZXmI+lLIJFEATMmzcPuq7j0KFDGe8fPHgQhmG48lv8ePDBB/GP//iPWL9+ve/7x44dAwBMmjRpyGMeDIMSLNx5Xxa0Asqv0EWlUIFOcxWyXSfdjiQ6NTcItDzKYPJYRs557y8IcuWxUI1F0zPHXykJku4N68ivK2XLY6G1wNatW5fxHn2NmrKycfToUTz77LN4/PHHM95LpVJ47rnnAADnnXfeUIc7KHSDFL2osMdXgso6VqAPUy6NxTHdFFdoMNsOsJKK/o0ExZrC2GmsXFNYDsGiZT67qmpURJVy17pSAeMpm2C5+uqrEQwG8atf/crlA9m2bRvuvfdehEIh3HDDDfbrra2t2L9/P7q7u+3XrrnmGkiShKeeesoWIgCgqiruuOMOHD9+HO973/uweHF+x2E50A1jEKYw5v8VcAOMFeiDlet6UJNGwT4WqgVlEUSV1LFvJCi2pMtIJkjqTIKknyCwi1D6msKsxFofH4v5/shrLZXmYylbrOu0adNw22234fbbb8f111+P1atXgxCCDRs2QNM0/PjHP0ZTkxNNddttt2Hjxo340pe+hC9/+csAgHnz5uEb3/gGfvCDH+ArX/kKlixZgilTpmDr1q1oa2vDnDlz8KMf/ahcPyEvg+mHzjPvywPVWHKZAeyFsEAtU8+jseh5TGVjHWpaLNQUNrJ5LIY1BlMQKrLkeT+7xpJMZ2pmLsGS1guqPVdOKm1dKetsfPKTn8SUKVNw7733YsuWLQgEAlixYgW+8IUv4JxzzinoHJ/5zGcwf/583HvvvXjvvfewZ88eTJkyBbfeeituueUWVFWNXLjukH0sFWALHSs4Gkt+U1ih806vLSH+1XiLrYg71rBNYQUKlpF0MLPXKK36CRbrX59x+Zn8WK23EpIkK624bdnF7Jo1a7BmzZq8x91///1Z3zvnnHMKFkTDiTEIjUXXDUiiYGo7w/xwqZqOw20DOWPcRyuFaA9+0T05z+l5WCXJI1hOeR8LNYWRgsrgsxrLsJvCmEueVnVUhd0NK3JtEqjGombTWCrNFFYBJnbeQXII6AYpWuswCKAo5rRns92Xi/daOvF/T+7AQNy/adpohoYQ5xQs6eI0FiOPQ7SS+l+MBGyEVSEBESPqvGe+3K/ZFynQec8eR+VoJWgsleZj4YJlCBiDiAozDAJFNqfdz1FYTmhNq2Rq+B+Ev755CK9vLS6ZtBgMPX/yo23SKNTHwoaX+iycp7opjN2pq1nyQ1hGtmw+awrLHGsheSyLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(torch.randn(train_y.shape[0])*pred_scale / dt**0.5, alpha=0.75)\n",
|
||||
"plt.plot(log_returns)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Now Do Forward Predictions with Voltron"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 18,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"pred_scale = pred_scale / dt**0.5"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 19,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Parameter containing:\n",
|
||||
"tensor([-4.6074], requires_grad=True)\n",
|
||||
"tensor([0.2951])\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"vol_lh = gpytorch.likelihoods.GaussianLikelihood()\n",
|
||||
"vol_lh.noise.data = torch.tensor([1e-6])\n",
|
||||
"vol_model = BMGP(train_x, pred_scale.log(), vol_lh)\n",
|
||||
"\n",
|
||||
"optimizer = torch.optim.Adam([\n",
|
||||
" {'params': vol_model.parameters()}, # Includes GaussianLikelihood parameters\n",
|
||||
"], lr=0.01)\n",
|
||||
"\n",
|
||||
"# \"Loss\" for GPs - the marginal log likelihood\n",
|
||||
"mll = gpytorch.mlls.ExactMarginalLogLikelihood(vol_lh, vol_model)\n",
|
||||
"\n",
|
||||
"for i in range(500):\n",
|
||||
" # Zero gradients from previous iteration\n",
|
||||
" optimizer.zero_grad()\n",
|
||||
" # Output from model\n",
|
||||
" output = vol_model(train_x)\n",
|
||||
" # Calc loss and backprop gradients\n",
|
||||
" loss = -mll(output, pred_scale.log())\n",
|
||||
" loss.backward()\n",
|
||||
"# print(loss.item(), model.covar_module.vol.item())\n",
|
||||
" optimizer.step()\n",
|
||||
" \n",
|
||||
"print((vol_model.covar_module.raw_vol))\n",
|
||||
"print(vol_model.mean_module.constant.data.exp())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 20,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from voltron.means import LogLinearMean"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 21,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"voltron_lh = gpytorch.likelihoods.GaussianLikelihood()\n",
|
||||
"voltron = VoltronGP(train_x, train_y.log(), voltron_lh, pred_scale)\n",
|
||||
"voltron.mean_module = LogLinearMean(1)\n",
|
||||
"voltron.mean_module.initialize_from_data(train_x, train_y.log())\n",
|
||||
"voltron.likelihood.raw_noise.data = torch.tensor([1e-6])\n",
|
||||
"voltron.vol_lh = vol_lh\n",
|
||||
"voltron.vol_model = vol_model\n",
|
||||
"\n",
|
||||
"grad_flags = [False, True, True, True, False, False, False]\n",
|
||||
"\n",
|
||||
"for idx, p in enumerate(voltron.parameters()):\n",
|
||||
" p.requires_grad = grad_flags[idx]\n",
|
||||
" \n",
|
||||
"voltron.train();\n",
|
||||
"voltron_lh.train();\n",
|
||||
"voltron.vol_lh.train();\n",
|
||||
"voltron.vol_model.train();\n",
|
||||
"\n",
|
||||
"# Use the adam optimizer\n",
|
||||
"optimizer = torch.optim.Adam([\n",
|
||||
" {'params': voltron.parameters()}, # Includes GaussianLikelihood parameters\n",
|
||||
"], lr=0.1)\n",
|
||||
"\n",
|
||||
"# \"Loss\" for GPs - the marginal log likelihood\n",
|
||||
"mll = gpytorch.mlls.ExactMarginalLogLikelihood(voltron_lh, voltron)\n",
|
||||
"\n",
|
||||
"for i in range(500):\n",
|
||||
" # Zero gradients from previous iteration\n",
|
||||
" optimizer.zero_grad()\n",
|
||||
" # Output from model\n",
|
||||
" output = voltron(train_x)\n",
|
||||
" # Calc loss and backprop gradients\n",
|
||||
" loss = -mll(output, train_y.log())\n",
|
||||
" loss.backward()\n",
|
||||
" # print(loss.item())\n",
|
||||
" optimizer.step()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Predict"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 22,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Automatic pdb calling has been turned ON\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"%pdb"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 23,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"0\n",
|
||||
"1\n",
|
||||
"2\n",
|
||||
"3\n",
|
||||
"4\n",
|
||||
"5\n",
|
||||
"6\n",
|
||||
"7\n",
|
||||
"8\n",
|
||||
"9\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"nvol = 10\n",
|
||||
"npx = 10\n",
|
||||
"vol_paths = torch.zeros(nvol, ntest)\n",
|
||||
"px_paths = torch.zeros(npx*nvol, ntest)\n",
|
||||
"\n",
|
||||
"voltron.vol_model.eval();\n",
|
||||
"voltron.eval();\n",
|
||||
"\n",
|
||||
"for vidx in range(nvol):\n",
|
||||
" print(vidx)\n",
|
||||
" vol_pred = voltron.vol_model(test_x).sample().exp()\n",
|
||||
" vol_paths[vidx, :] = vol_pred.detach()\n",
|
||||
" \n",
|
||||
" px_pred = voltron.GeneratePrediction(test_x, vol_pred, npx).exp()\n",
|
||||
" px_paths[vidx*npx:(vidx*npx+npx), :] = px_pred.detach().T"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 24,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAA08AAAFTCAYAAADhtTfTAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAEAAElEQVR4nOy9eZgldXX//66qu++9b7MySzMzMBsOgmGPYnADMRGEGCEQXGI0USOJ/nh80C/fRCPffB2J8lWQJBpcAyhRAVkFxBkZZoZhlp6tp6enl3u7++5b3bpV9fvjcj5ddft29+2e7lnP63l86L5dt25V3bFOvT/nnPeRTNM0wTAMwzAMwzAMw0yJfLIPgGEYhmEYhmEY5nSAxRPDMAzDMAzDMEwdsHhiGIZhGIZhGIapAxZPDMMwDMMwDMMwdcDiiWEYhmEYhmEYpg5YPDEMwzAMwzAMw9QBiyfmrOeOO+5Ad3c3br311rrf88lPfhLd3d34+7//+xl/3j/8wz+gu7sbX/3qV2f83sk4dOiQ7fdjx46hu7sb3d3dyOVy4vWrrroK3d3deO6558RrjzzyCLq7u3H99ddP2G+pVEJ/f/+cHSfDMAxzcrHGh+r/nXvuuVi7di0uu+wyfPzjH8fTTz89o31PFU8Y5kyBxRNz1nPdddcBALZs2YKxsbFpt89kMnjhhRcAAO9///vn89Cm5ciRI7jlllvwjW98Y873/fLLL+M973mPOFeGYRjmzOK8887Dxo0bxf/Wr1+P5cuXo1gs4tlnn8Vf//Vf4+677z7Zh8kwpxSOk30ADHOyefvb345QKIR0Oo0nnngCN99885TbP/HEEyiVSujo6MBFF110go6yNo8//jheeeUVvPOd77S93tbWhl/96lcAAJ/PN+U+3vGOd2DdunXweDy21++//3709fXN7QEzDMMwpwzf+MY3sGDBggmva5qG++67D/fffz8efvhhXHrppbjqqqum3d9k8YRhziQ488Sc9bhcLlxzzTUAIATHVDz++OMAgGuvvRayfGr+X8jpdGLZsmVYtmwZJEmacttgMIhly5ahq6vrBB0dwzAMcyrjdDrxd3/3d9iwYQMA4OGHH67rfRxPmLOBU/PJj2FOMFS6t23bNkSj0Um3i0aj+MMf/gDg5JfsMQzDMMx8cuWVVwIAdu3adZKPhGFOHVg8MQyAjRs3YvHixTBNE7/+9a8n3e5//ud/YBgGNmzYgCVLlgAADMPAz372M9x888244IILcP755+Pqq6/GPffcM6UQq8W+ffvw//1//x/e+c53YsOGDTj//PNxxRVX4LOf/Sx2795t27a7uxv33XcfAODJJ59Ed3c3PvzhDwOY3DCiFtUNvlu2bEF3dze2bt0KAPjKV76C7u5ufPOb38T3v/99dHd34z3vec+k+7vvvvvQ3d2NL37xizM6d4ZhGObUIhAIAICII9/85jfR3d2NBx98EA8//DAuu+wyrF27Fu95z3vQ19c3pWFENpvFd77zHbz//e8X/VUf+MAH8PDDD8MwjAnbl0ol/Pu//zs+8IEPYMOGDVi/fj3e//7348EHH4SqqvN74gwzBSyeGOZNrr32WgDAL3/5y0m3oZI9CgzFYhG33norvvjFL+LVV19FU1MTVqxYgeHhYfznf/4n3vve9+LVV1+t6/MfffRRXH/99fjpT3+KVCqFpUuXoqOjA7FYDP/zP/+DG264QWS9gIrg6+joAABEIhFs3LgRK1eunNW5WwkGg9i4caMImgsXLhSf9e53vxtOpxMHDhxAT09PzfdbyxoZhmGY05ejR48CgIg1xFNPPYW7774biqKgq6sL+XweCxcunHQ/AwMD+OAHP4h7770XPT096OrqQnt7O9544w3cfffd+Md//Efb9slkEjfffDP+6Z/+CXv27EFLSwsWLVqEnp4efO1rX8OHPvQhJBKJuT9hhqkDFk8M8ybXXXcdJEnC66+/XtOe++DBg9i7dy88Hg/e9a53AQDuvvtu/P73v0dbWxt+/OMf46mnnsIjjzyCl156Ce9+97uRSqXwyU9+EiMjI1N+9ujoKO6++27ouo4777wTL7/8Mh555BE89dRTePLJJ3HuuedC0zR85zvfEe/54Q9/iA984AMAgLe+9a344Q9/iLvuuuu4r8Pq1avxwx/+EKtXrwYA3HLLLfjhD3+IP/3TP0VjYyMuvfRSAJUsXDWvv/46jhw5gs7OTmzatOm4j4VhGIY5OaTTafziF78AAFx++eW2v+3YsQMf+chH8Oyzz+LXv/41fvazn03ZA/wP//APOHToENavX4/f/OY3ePzxx/HEE0/g+9//Pnw+Hx577DHxWbT966+/jg0bNuDJJ5/EU089hV/84hd45pln8Ja3vAW7d+/m6gbmpMHiiWHepKurSzzw1yrdo4zK29/+dgQCARw7dgyPPfYYgEopw/r168W2oVAI//Iv/4I1a9YgkUjg3//936f8bCqRW79+Pf7yL/8SiqKIvy1cuBB/+Zd/CWDiPKeTAfWH1crQUfB73/veN61RBcMwDHNqYZom0uk0fvvb3+K2225DPB5HMBjEbbfdZtvO6XTi05/+tLjPNzY2TrrP1157DVu3boXP58O3vvUtm5nEhRdeiE9+8pMAgJ///OcAKv1Vzz33HCKRCL71rW9h0aJFYvuOjg5s3rwZfr8fzzzzDPbt2zdn584w9cLiiWEskDCodt0zTXNCyd6LL74IwzCwZs0arFu3bsK+FEXBTTfdBAB4/vnnp/zcd73rXdixYwf+8z//s+bfvV4vAKBQKNR9LvPFlVdeiXA4jIGBAbz22mvidV3Xheh83/ved7IOj2EYhqmTP/7jP54wJHfTpk34q7/6K7z++utoaGjAt771rQlleytXroTf76/rM2hW4JVXXommpqYJf7/hhhvwy1/+Et/+9rcBAM888wwA4G1ve1tNUdbU1CTGhPz2t7+t/2QZZo7gOU8MY+Gd73wnvvKVr2Dv3r04fPgwzjnnHACVlbOBgQG0t7fj4osvBlAZUAsAq1atmnR/a9assW07HU6nE9u2bcP+/ftx9OhRHD16FPv27cOxY8cAoGZT7YmGrN1/9KMf4Ze//CU2btwIoDJUd3R0FOeddx6WLVt2ko+SYRiGmY7zzjsPLpdL/C7LMnw+H9ra2rBhwwZcc801NWcFtrS01P0Z1De1YsWKmn8PBAJYvny5+J0qLF599VV86EMfqvkeiom9vb11HwfDzBUsnhjGQiAQwNvf/nY8/vjj+NWvfiXKCWrNdiL3oalW3yjolMtlqKoKt9s96bZPPfUU7r33XpvQkmUZK1aswNVXX42nnnrquM5tLrnuuuvwox/9CL/+9a/xhS98AYqiiJI9NopgGIY5PZhsSO50TBXLqkkmkwCmH9hOZLNZAEAsFkMsFpty20wmU/dxMMxcweKJYap4//vfbxNPmqbhiSeeEH8jKBDQjb4WdGN3OBxTBpuXXnoJn/rUp2CaJq688kq8853vRHd3N5YuXQqv14uXXnrplBJPZNV+5MgRbNmyBRdccAGeeeYZOBwOvPvd7z7Zh8cwDMOcIng8HgD1l51Tmfqdd94p+n0Z5lSCxRPDVHHxxRejvb0dhw4dwsGDBzE0NIREIoENGzZg6dKlYjv6ee/evZPui2YzTWXhCgDf+973YJomrr/+evzTP/3ThL8PDw/P5lTmlfe9733YvHkznn76aWiahnw+jyuuuKJmTTvDMAxzdkIzEQ8ePFjz72NjY/joRz+KRYsW4Wtf+xoWL14MYGqDpD179kCSJCxcuFCM1WCYEwUbRjBMFbIs473vfS8A4Nlnn8XTTz8NYNxMgrj00kshyzL27NmDHTt2TNiPruv40Y9+BAC45JJLpvzMgYEBALX7p0zTxKOPPir2aYWcjkzTnOasZs50+7722mshSRKee+45PPfcc+I1hmEYhiFovMXzzz8vSvis/OY3v8GuXbtw6NAhOBwOXHHFFQAqpezxeHzC9plMBrfccguuu+66KYfaM8x8weKJYWpA5XlPP/00nn32Wbjd7gnlaAsWLBCC6tOf/rRNQKXTafz93/899uzZg3A4jNtvv33Kz6OVuZ/+9KcYHR0Vr8diMXz2s58Vg3arp6pTv9Xg4OCMz3E6qCxxsn0vWLAAb3nLWzA4OIif//znCAQCuOqqq+b8OBiGYZjTl7e97W1Yt24dMpkMPvWpT9nmHm7duhX33nsvAODWW28FUJlbuGnTJqTTaXz0ox9FX1+f2D4ajeITn/gEUqkUWlpaxEInw5xIuGyPYWqwbNkynH/++di5cycA4N3vfjeCweCE7e666y4MDAxgy5YtuOGGG7BkyRL4/X4cPHgQqqoiEong//7f/4v29vYpP+/jH/84Xn75Zezfvx9XXXUVli5dinK5jCNHjqBcLuPCCy/Etm3bUCqVkEwmEYlEAADd3d0AgDfeeEP0SW3evHlOrkF3dzeee+45/Md//AdeeeUVXHPNNfjoRz9q2+baa6/FH/7wB+TzeXzgAx8Qte0MwzAMQ/zrv/4rbrnlFmzZsgVXXnklVqxYgUwmIwbS/+mf/qmtuuPee+/Fbbfdhtdffx3vfOc7sXz5csiyjMOHD0PTNAQCAXz3u9/lmMOcFDjzxDCTYDWHsP5sxefz4Xvf+x6+/OUvY+PGjRgdHcWhQ4fQ1dWFv/qrv8IvfvELYW0+FevXr8ejjz6Kq6++Gk1NTTh48CBGR0exYcMG3HPPPfiP//gPnH/++QAgSuSAyoreZz7zGbS3t2NgYAC7d++eMzvzO+64A+9///sRCARw+PBh7N+/f8I211xzjbC55ZI9hmEYphZdXV145JFH8Nd//ddYvHgxDh06hLGxMWzcuBFf//rXcc8999i2b2trw09/+lP8/d//PdasWYOBgQEcPnwYra2tuOGGG/Dzn/98yjEhDDOfSOYMmiV+97vf4f7770dPTw80TcOaNWtwxx1Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 1008x360 with 2 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(1, 2, figsize = (14, 5))\n",
|
||||
"\n",
|
||||
"ax[0].plot(train_x, pred_scale)\n",
|
||||
"ax[0].plot(test_x, vol_paths.t(), color = \"gray\", alpha = 0.5)\n",
|
||||
"\n",
|
||||
"ax[1].plot(test_x, px_paths.t(), color = \"gray\", alpha = 0.2)\n",
|
||||
"ax[1].plot(train_x, train_y)\n",
|
||||
"ax[1].plot(test_x, test_y)\n",
|
||||
"ax[1].plot(test_x, voltron.mean_module(test_x).detach().exp(), lw=3.)\n",
|
||||
"\n",
|
||||
"ax[0].set_title(\"Volatility\")\n",
|
||||
"ax[1].set_title(\"Price\")\n",
|
||||
"\n",
|
||||
"plt.show();"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 4
|
||||
}
|
||||
@@ -0,0 +1,270 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 15,
|
||||
"id": "ddc44d8f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import os\n",
|
||||
"import glob"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 29,
|
||||
"id": "79dbbbf6",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"start = r\"\"\"\n",
|
||||
"\\begin{figure}\n",
|
||||
" \\centering\n",
|
||||
" \\includegraphics[width=\\linewidth]{figs/\"\"\"\n",
|
||||
"\n",
|
||||
"mid = r\"\"\"}\n",
|
||||
" \\caption{Trading Strategy for \"\"\"\n",
|
||||
"end = \"\"\".}\n",
|
||||
"\\end{figure}\n",
|
||||
"\"\"\""
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 26,
|
||||
"id": "ef42d2c4",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"files = glob.glob(\"./trading*\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 27,
|
||||
"id": "40d9639b",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"'WFC'"
|
||||
]
|
||||
},
|
||||
"execution_count": 27,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"files[0].split(\"_\")[-1][:-4]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 30,
|
||||
"id": "c802c9e3",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"\\begin{figure}\n",
|
||||
" \\centering\n",
|
||||
" \\includegraphics[width=\\linewidth]{figs/./trading_strategy_WFC.pdf}\n",
|
||||
" \\caption{Trading Strategy for WFC.}\n",
|
||||
"\\end{figure}\n",
|
||||
"\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"print(start + files[0] + mid + files[0].split(\"_\")[-1][:-4] + end)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 37,
|
||||
"id": "b8a6a5fa",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"\\begin{figure}\n",
|
||||
" \\centering\n",
|
||||
" \\includegraphics[width=\\linewidth]{figs/trading_strategy_WFC.pdf}\n",
|
||||
" \\caption{Trading Strategy for WFC.}\n",
|
||||
"\\end{figure}\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\\begin{figure}\n",
|
||||
" \\centering\n",
|
||||
" \\includegraphics[width=\\linewidth]{figs/trading_strategy_PLD.pdf}\n",
|
||||
" \\caption{Trading Strategy for PLD.}\n",
|
||||
"\\end{figure}\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\\begin{figure}\n",
|
||||
" \\centering\n",
|
||||
" \\includegraphics[width=\\linewidth]{figs/trading_strategy_EOG.pdf}\n",
|
||||
" \\caption{Trading Strategy for EOG.}\n",
|
||||
"\\end{figure}\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\\begin{figure}\n",
|
||||
" \\centering\n",
|
||||
" \\includegraphics[width=\\linewidth]{figs/trading_strategy_COP.pdf}\n",
|
||||
" \\caption{Trading Strategy for COP.}\n",
|
||||
"\\end{figure}\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\\begin{figure}\n",
|
||||
" \\centering\n",
|
||||
" \\includegraphics[width=\\linewidth]{figs/trading_strategy_JPM.pdf}\n",
|
||||
" \\caption{Trading Strategy for JPM.}\n",
|
||||
"\\end{figure}\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\\begin{figure}\n",
|
||||
" \\centering\n",
|
||||
" \\includegraphics[width=\\linewidth]{figs/trading_strategy_BRK.B.pdf}\n",
|
||||
" \\caption{Trading Strategy for BRK.B.}\n",
|
||||
"\\end{figure}\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\\begin{figure}\n",
|
||||
" \\centering\n",
|
||||
" \\includegraphics[width=\\linewidth]{figs/trading_strategy_BLK.pdf}\n",
|
||||
" \\caption{Trading Strategy for BLK.}\n",
|
||||
"\\end{figure}\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\\begin{figure}\n",
|
||||
" \\centering\n",
|
||||
" \\includegraphics[width=\\linewidth]{figs/trading_strategy_C.pdf}\n",
|
||||
" \\caption{Trading Strategy for C.}\n",
|
||||
"\\end{figure}\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\\begin{figure}\n",
|
||||
" \\centering\n",
|
||||
" \\includegraphics[width=\\linewidth]{figs/trading_strategy_SLB.pdf}\n",
|
||||
" \\caption{Trading Strategy for SLB.}\n",
|
||||
"\\end{figure}\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\\begin{figure}\n",
|
||||
" \\centering\n",
|
||||
" \\includegraphics[width=\\linewidth]{figs/trading_strategy_BAC.pdf}\n",
|
||||
" \\caption{Trading Strategy for BAC.}\n",
|
||||
"\\end{figure}\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\\begin{figure}\n",
|
||||
" \\centering\n",
|
||||
" \\includegraphics[width=\\linewidth]{figs/trading_strategy_MS.pdf}\n",
|
||||
" \\caption{Trading Strategy for MS.}\n",
|
||||
"\\end{figure}\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\\begin{figure}\n",
|
||||
" \\centering\n",
|
||||
" \\includegraphics[width=\\linewidth]{figs/trading_strategy_SCHW.pdf}\n",
|
||||
" \\caption{Trading Strategy for SCHW.}\n",
|
||||
"\\end{figure}\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\\begin{figure}\n",
|
||||
" \\centering\n",
|
||||
" \\includegraphics[width=\\linewidth]{figs/trading_strategy_GS.pdf}\n",
|
||||
" \\caption{Trading Strategy for GS.}\n",
|
||||
"\\end{figure}\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\\begin{figure}\n",
|
||||
" \\centering\n",
|
||||
" \\includegraphics[width=\\linewidth]{figs/trading_strategy_CCI.pdf}\n",
|
||||
" \\caption{Trading Strategy for CCI.}\n",
|
||||
"\\end{figure}\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\\begin{figure}\n",
|
||||
" \\centering\n",
|
||||
" \\includegraphics[width=\\linewidth]{figs/trading_strategy_AMT.pdf}\n",
|
||||
" \\caption{Trading Strategy for AMT.}\n",
|
||||
"\\end{figure}\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\\begin{figure}\n",
|
||||
" \\centering\n",
|
||||
" \\includegraphics[width=\\linewidth]{figs/trading_strategy_XOM.pdf}\n",
|
||||
" \\caption{Trading Strategy for XOM.}\n",
|
||||
"\\end{figure}\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\\begin{figure}\n",
|
||||
" \\centering\n",
|
||||
" \\includegraphics[width=\\linewidth]{figs/trading_strategy_CVX.pdf}\n",
|
||||
" \\caption{Trading Strategy for CVX.}\n",
|
||||
"\\end{figure}\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\\begin{figure}\n",
|
||||
" \\centering\n",
|
||||
" \\includegraphics[width=\\linewidth]{figs/trading_strategy_EQIX.pdf}\n",
|
||||
" \\caption{Trading Strategy for EQIX.}\n",
|
||||
"\\end{figure}\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\\begin{figure}\n",
|
||||
" \\centering\n",
|
||||
" \\includegraphics[width=\\linewidth]{figs/trading_strategy_AXP.pdf}\n",
|
||||
" \\caption{Trading Strategy for AXP.}\n",
|
||||
"\\end{figure}\n",
|
||||
"\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"for stck in files:\n",
|
||||
" tckr = stck.split(\"_\")[-1][:-4]\n",
|
||||
" print(start + stck[2:] + mid + tckr + end)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "c9e72667",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,968 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 23,
|
||||
"id": "84cfc73b",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns\n",
|
||||
"import torch\n",
|
||||
"import gpytorch\n",
|
||||
"# from voltron.robinhood_utils import GetStockData\n",
|
||||
"import os\n",
|
||||
"# import robin_stocks.robinhood as r\n",
|
||||
"import pickle5 as pickle\n",
|
||||
"import pandas as pd\n",
|
||||
"from torch.distributions import Beta\n",
|
||||
"from scipy.special import betainc\n",
|
||||
"\n",
|
||||
"sns.set_style('white')\n",
|
||||
"# style.use('whitegrid')\n",
|
||||
"palette = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"sns.set(palette = palette, font_scale=2.0, style=\"white\", rc={\"lines.linewidth\": 4.0})"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "c4166d5f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"sns.palplot(palette)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 60,
|
||||
"id": "23f103e8",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"XLF_tckrs = list(pd.read_pickle(\"../../spdr-data/XLF.pkl\").symbol.unique())\n",
|
||||
"XLE_tckrs = list(pd.read_pickle(\"../../spdr-data/XLE.pkl\").symbol.unique())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 144,
|
||||
"id": "b3c79fa9",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"['BRK.B', 'JPM', 'BAC', 'WFC', 'C', 'MS']\n",
|
||||
"['XOM', 'CVX', 'EOG', 'COP', 'SLB']\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"print(XLF_tckrs[:6])\n",
|
||||
"print(XLE_tckrs)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 256,
|
||||
"id": "e0a102f3",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"spdrs = [\"XLF\", \"XLE\", \"XLRE\"]\n",
|
||||
"tckr_spdrs = []\n",
|
||||
"tckrs = []\n",
|
||||
"for spdr in spdrs:\n",
|
||||
" spdr_dat = pd.read_pickle(\"../../spdr-data/\" + spdr + \".pkl\")\n",
|
||||
" syms = list(spdr_dat.symbol.unique())\n",
|
||||
" tckrs += syms\n",
|
||||
" tckr_spdrs += [spdr for _ in range(len(syms))]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 124,
|
||||
"id": "30bb2631",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"SPDR = \"XLE\"\n",
|
||||
"tckr = \"CVX\"\n",
|
||||
"\n",
|
||||
"spdr_dat = pd.read_pickle(\"../../spdr-data/\" + SPDR + \".pkl\") \n",
|
||||
"data = spdr_dat[spdr_dat[\"symbol\"] == tckr]\n",
|
||||
"T = 5.\n",
|
||||
"ts = torch.linspace(0, T, data.shape[0])\n",
|
||||
"y = torch.FloatTensor(data['close_price'].to_numpy())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 223,
|
||||
"id": "0b67e09d",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"eval_times = [100, 200, 300, 400, 500, 600, 700, 800, 900, 1000, 1100, 1200]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 126,
|
||||
"id": "b43e33a5",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
" prices_at_time_y = y[torch.tensor(eval_times)]\n",
|
||||
"delta_y = prices_at_time_y[1:] - prices_at_time_y[:-1]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "c85e0f19",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Load Price Probabilities"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 127,
|
||||
"id": "a24b2245",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"voltron = torch.load(\"./outputs/voltron_\" + tckr + \".pt\")\n",
|
||||
"matern = torch.load(\"./outputs/matern_\" + tckr + \".pt\")\n",
|
||||
"specmix = torch.load(\"./outputs/sm_\" + tckr + \".pt\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "7da551d1",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Example Beta(17, 8) which is pretty right skewed."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 128,
|
||||
"id": "ad82c162",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"bought_func = lambda xs: betainc(17, 8, xs)\n",
|
||||
"held_voltron = 1000 * bought_func(voltron)\n",
|
||||
"held_matern = 1000 * bought_func(matern)\n",
|
||||
"held_specmix = 1000 * bought_func(specmix)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 129,
|
||||
"id": "af0e74e2",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def reward_risk(a, b, prob_incs):\n",
|
||||
" bought_func = lambda xs: betainc(a, b, xs)\n",
|
||||
" total_held = 1000 * bought_func(prob_incs)\n",
|
||||
" \n",
|
||||
" returns = total_held[1:] * delta_y\n",
|
||||
" #cum_returns = returns.cumsum(0)\n",
|
||||
" return returns.std(), returns.sum()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "3bbcf330",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Randomly sample to find pareto fronts. ideally, we'd do BO or something like that but it's 2d."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 130,
|
||||
"id": "8d3fa1b6",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"base = 10000"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 131,
|
||||
"id": "703cc564",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"prices_at_time_y = torch.tensor([y[0], *prices_at_time_y])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 132,
|
||||
"id": "5af01ccc",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def value_func(vec, base = 10000):\n",
|
||||
" portfolio_value = torch.zeros(12)\n",
|
||||
" portfolio_value[0] = base\n",
|
||||
" for i in range(11):\n",
|
||||
" price_of_stock = portfolio_value[i] * bought_func(vec)[i]\n",
|
||||
" amt_bought = price_of_stock / prices_at_time_y[i]\n",
|
||||
" cash_left = portfolio_value[i] - price_of_stock\n",
|
||||
" portfolio_value[i+1] = cash_left + amt_bought * prices_at_time_y[i+1]\n",
|
||||
" return portfolio_value\n",
|
||||
" "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 133,
|
||||
"id": "7c6a6113",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"plt_times = torch.tensor(eval_times) / 252\n",
|
||||
"hodl_strat = 10000 / y[0] * prices_at_time_y[1:]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 134,
|
||||
"id": "62e46650",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def running_sharpe_ratio(vec):\n",
|
||||
" returns = vec - 10000\n",
|
||||
" std_returns = torch.tensor([returns[:i].std(0) for i in range(len(vec))])\n",
|
||||
" # need avg return divided by sd of returns?\n",
|
||||
" return returns.cumsum(0) / std_returns / torch.arange(vec.shape[0])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 135,
|
||||
"id": "bfa6a9b3",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAACY0AAAIfCAYAAAA2ImFBAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdd3wUdfoH8M/M9pJk0yuQ0CK9KooKCipnxY6CoJ6eooeenp7n2c7CnXo/9U49wQMLCKKCCiJYUEBQKdI7CZAC6XWT7W3m98cmIZuZbclmS/K8Xy9e7H7nOzNPskl2duaZ52F4nudBCCGEEEIIIYQQQgghhBBCCCGEEEIIIaRXYCMdACGEEEIIIYQQQgghhBBCCCGEEEIIIYSQ8KGkMUIIIYQQQgghhBBCCCGEEEIIIYQQQgjpRShpjBBCCCGEEEIIIYQQQgghhBBCCCGEEEJ6EUoaI4QQQgghhBBCCCGEEEIIIYQQQgghhJBehJLGCCGEEEIIIYQQQgghhBBCCCGEEEIIIaQXoaQxQgghhBBCCCGEEEIIIYQQQgghhBBCCOlFKGmMEEIIIYQQQgghhBBCCCGEEEIIIYQQQnoRShqLUXfccQfuuOOOSIdBCCGEENJr0fEYIYQQQkhk0fEYIYQQQkjk0LEYIYQQEvukkQ6AdE5lZWWkQwjaydM1MFscHmNqlQwD+6ZFKCKKKRjRGBfFFLhojItiClw0xhWNMQHRGxfpmWLxeIwQQgghpCeh4zFCCCGEkMihYzFCCCEk9lGlMUIIIYQQQgghhBBCCCGEEEIIIYQQQgjpRajSGAmb7PRE8BzvMcawTISicaOYAheNcVFMgYvGuCimwEVjXNEYExC9cRFCCCGEEEIIIYQQQgghhBBCSDShpDESNiqFLNIhCFBMgYvGuCimwEVjXBRT4KIxrmiMCYjeuAghhBBCCCGEEEIIIYQQQgghJJpQe0pCCCGEEEIIIYQQQgghhBBCCCGEEEII6UUoaYwQQgghhBBCCCGEEEIIIYQQQgghhBBCehFKGiOEEEIIIYQQQgghhBBCCCGEEEIIIYSQXoSSxgghhBBCSES98MILyM/Px9tvvx2S7ZWXl2P+/PmYNm0aRowYgfPOOw+33HILlixZAqvVGpJ9EEIIIYQQQgghhBBCCCGEEBLLpJEOgBBCCCGE9F7bt2/Hp59+GrLtbdmyBY8++ihMJlPbmN1ux8GDB3Hw4EF88cUXWLhwIXJyckK2T0IIIYQQQgghhBBCCCGEEEJiDSWNEUIIIYSQiDh8+DDmzZsHjuNCsr2CggI89NBDsNlskEgkuPnmmzF+/HiYTCZ89dVX2LdvHwoLC/Hggw/is88+g0qlCsl+CSGEEEIIIYQQQgghhBBCCIk1lDRGCCGEEELCbsuWLXj88cdhNBpDts3nn3++LWFs4cKFmDx5ctuy2267DfPnz8fy5ctRUFCAjz76CPfff3/I9h0rOI5DeVUt9M0G6OLjkJ2RCpaljvWEEEIIIYQEgud5ODkeLg6QsICUZcAwTKTDIoQQQgjpFnTsQwghPR8ljZGwKS6vh8Vq9xhTKeXIy06OUEQUUzCiMS6KKXDRGBfFFLhojCsaYwKiNy5ylt1ux7vvvouFCxeGrMIYAOzbtw979+4FAFx//fUeCWMAwDAMnn76aWzfvh2nTp3CBx98gN///veQyWQhiyGaVdXW441Va7G8pgS1GknbeKrJhTvScvHnW65DRir9nhBCokO0JrjyPA9LowEOsxUytRKqxDg6WU4IIb2E08Wj1mhHZZMdNufZzzEKKYvMBDlStXJIJfSeQAghhJCegY59CCGk96CkMRI2LpcLTicnGIskiilw0RgXxRS4aIyLYgpcNMYVjTG1xhCNcRG3bdu24dlnn0VZWRkAQK1WY8aMGfjwww+7vO1vvvmm7fHMmTNF57Asi5kzZ+Kll16CXq/Hjh07cPHFF3d539Fu2doNmHt4G+xSBlB5Jl3Uqlj823ga7yx+G+8On4jZ110RoSgJISR6E1ytzSYc/WIL9i79Fk2nq9vGE/qmY+ydV2LoTZOhjNeEPS5CCCHhoTc7UFBtBscLl9mcHErqrTjdYEV+uho6de+4KYUQQgghPRcd+xBCSO8S+Vt1CSGEEEJIr7B27dq2hLHhw4dj1apVuPTSS0Oy7V27dgEAEhISMHToUK/zzj///LbHP//8c0j2Hc2Wrd2Ae45th0PCAAwDsB3uAGTd4w4Jg3uObceytRsiEyghpNdbtnYDBix+G/82nkatlwTXAYvfDvvfqZKt+7Fo4lxsnr8U+jM1MCklaIiTwaSUQH+mBpvnL8WiiXNRsnV/WOMihBASHnqzA8eqxC+atsfxwLEqM/RmR3gCI4QQQgjpBnTsQwghvQ9VGiOEEEIIIWGTlJSEefPm4bbbboNEIkF9fX2Xt+lyuXDq1CkAwIABA3y2MMvLy4NEIoHL5cLx48e7vO9oVlVbj7mHtwESBnzHZLEOeJYBw/GYe3gbLr9gHLWqJISEVWuCK1oTXDv+yWr5G+aQwD0PCEtlxJKt+/Hl71+BWcZgz/AkbBuRhIYERdvypCYbJh5qwLhCPb78/Su48YMnkTtpdLfHRQghJDycLh4F1eag1imoNmNc33hq10QIIYSQmEPHPoQQ0jtR0hghMaxeb8QTr6/Bjv1FGNA3Fc//8RqMHdon0mERQgghombNmoXnn38eSqUypNutq6uD3W4HAGRnZ/ucK5FIkJqaiqqqKlRUVIQ0jmjzxqq17paUTGAnbXiWgZ0B/r1qLV598O5ujo4QQtw6m+A6RpeCFF1Ct8VlN1rw1dzXUJCtxvJpfWCXskCHO60b4uRYd2EGNkxIwx3fn8HaB1/HfdvepVaVhBDSQ9Qa7X6rbHTE8e71MtslGRNCCCGExAI69iGEkN6JksYIiWH3/X0Fdh4sAQDsPXoGM//yAXZ8+gR0carIBkYIIYSIGDFiRLdst6Ghoe1xYmKi3/kJCQmoqqqCXq/vlniiAcdxWF5TAqiD7EbPA8tqSvAyx/ms2EYIIaHS2QTXK9evQt8aC9RWF9RWZ8v/LmisTqht7scqmwtskCe82yvso8WHV/VzP/FZAY3Fh1f1w93flOLol1sw9q6rOr9TQgghUYHneVQ22Tu1bmWTHRnxcjABvrcRQgghhEQaHfsQQkjvRUljhMSoihp9W8JYK6PZhq27TuC6KSMjExQhhBASARaLpe2xQuH/rrbWOVartdtiirTyqlrUaiTBr8gyqNVIUFFdh5zMtNAHRggh7XQ6wZVhUJWiQlWK75tlGJ6H0taaTNaSXNbyvH2yWcdlMhcPi5zF8mnuKs6BVkBbPq0Phi3/DmPuvJJOlhNCSIxzcjxsTq5T69qcHJwcDxm1aSKEEEJIjKBjH0II6b0oaYyQGLX78GnR8Xc/+5mSxgghhPQqLper7bFcLvc7v3WO0+nstpgiTd9s6NL6jU3NlDRGCOl2nU5wDRDPMLAopbAopagPYj2Zg4PExblbUgZVAY3FTxo7/qA3QpUY17mgCSGERAVX566Zeqwv6763OEIIIYSQkKJjH0II6b2o5wwhMcpmF7/QXVbdGOZICCGEkMhqX13M4XD4nW+3u0uty2Sybosp0nTxXUtWSEyID1EkhBDiXVcTXLuLQ8bCqpQGnDDWhgd+HZEMm9HcPYERQggJC57nUWPoXHumVhI6604IIYSQGNLVYxc69iGEkNhFf8IJiVE2u/hF8YYmMxav+iXM0RBCCCGRo9Fo2h7bbDa/81vnKJXKbosp0rIzUpFqcgEcH9yKHI9UkwtZ6SndExghhLTT1QTXqMMyaEhQwCillhyEEBKrTDYXDpUbUa73/7nCG4ZxJ56R6MPzPBwuDlYHB4eLo9eJEEIIaSFlGSiknUsbUEhZSFn6HEwIIbGKksYIiVF6g8Xrsn/87zufywkhhJCeJC7ubNJBU1OT3/mtc5KSkrotpkhjWRZ3pOUCwZ6vYYDZablgWfqYQAjpfpmpyUhutgPBXrDleTAcD2n3hNVlDlXPrWRJCCE9lYvjUVpvxcFyI0z2rvVn4nngSIUZFrsrRNGRrnK6eFQ22bDvjBG7Sw3Yd8bQ8r8RlU02OF2UPEYIIaR3YxgGmQnyTq2bmSAHE2ylbkIIIVEjWs+xEkL8aDZZvS5zujis33IYs645N4wREUIIIZGRlpYGjUYDk8mEiooKn3NdLhdqa2sBAFlZWeEIL2L+fMt1eGfx23BIAD6Au/0YjofMxePRW64LQ3SEkN6Oc7qw4YkFmFBcj28mZgS9/qNxffHKA3fB6LCj3mJBvcWMBosF9dazj+ssZjS0PK+3WNoeG+xdaznmT3y7tsmEEEKiX5PFiaJaC6zOriWLtWd1cjhcYUJ+uhrxKjoFH0l6swMF1WbRIsw2J4eSeitON1iRn66GTk2J34QQQnqvVK0cpxusQTUuYBn3eoQQQmIXfWIlJEY5HL7vViwqqwtTJIQQQkjkDRgwAAcPHsSpU6d8zisqKoLL5X4PHTx4cDhCi5iM1GS8O3wi7jm2HQzH+04ca6ny8w9lP2SkJocpQkJIb8U5Xfjmz2+jYN02NE3McPfxClD7BFeGYRAnVyBOrkBugi7gbdhdrnZJZeaWpDNL2+M6sxmfFxyB2eEIKjbwPPISdEhSqgJfhxBCSMQ4XTxKGyyoMTi6Z/scj6OVJgxIVSE1ji6mRoLe7MCxKrPfeRwPHKsyY0gGJY4RQgjpvaQSBvnp6oDeO1vlp6shlVCVMUIIiWWUNEZIjLI7fSeNSah/OCGEkF7k3HPPxcGDB1FfX48TJ05g0KBBovN27NjhsU5PN/u6KwAAcw9vg50BwMN9C2BHDIObNpUhRWEF/yhPJeUJId3G5XBi/Z/exInvdqIkQ41tIwJPVGVabnd+d8SFXUpwlUskyNBqkaHVep0zKj0Dj2/6HsE0q2IYBg+NP5/+hhJCSAyoNzlQXGeLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 3000x500 with 4 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(1, 4, figsize = (30, 5))\n",
|
||||
"ax[0].plot(ts, y)\n",
|
||||
"[ax[i].set_xlabel(\"Time\") for i in range(4)]\n",
|
||||
"ax[0].set_ylabel(\"Asset Price\")\n",
|
||||
"[ax[0].axvline(x=plt_times[i], alpha = 0.2, linestyle=\"--\") for i in range(len(eval_times))]\n",
|
||||
"\n",
|
||||
"ax[1].plot(plt_times, matern, marker = \".\", markersize = 20, label = \"Matern\", color=palette[4])\n",
|
||||
"ax[1].plot(plt_times, specmix, marker = \".\", markersize = 20, label = \"Spectral Mixture\", color=palette[2])\n",
|
||||
"ax[1].plot(plt_times, voltron, marker = \".\", markersize = 20, label = \"Voltron\", color = palette[-2])\n",
|
||||
"\n",
|
||||
"ax[2].plot(plt_times, value_func(matern), label = \"Matern\", marker = \".\", markersize = 20, color=palette[4])\n",
|
||||
"ax[2].plot(plt_times, value_func(specmix), label = \"SM\", marker = \".\", markersize = 20, color=palette[2])\n",
|
||||
"ax[2].plot(plt_times, hodl_strat, label = \"HOLD\", marker = \".\", markersize = 20, color=palette[1])\n",
|
||||
"ax[2].plot(plt_times, value_func(voltron), label = \"Voltron\", marker = \".\", markersize = 20, color = palette[-2])\n",
|
||||
"\n",
|
||||
"ax[3].plot(plt_times, running_sharpe_ratio(value_func(matern)), \n",
|
||||
" label = \"Matern\", marker = \".\", markersize = 20, color=palette[4])\n",
|
||||
"ax[3].plot(plt_times, running_sharpe_ratio(value_func(specmix)), \n",
|
||||
" label = \"SM\", marker = \".\", markersize = 20, color=palette[2])\n",
|
||||
"ax[3].plot(plt_times, running_sharpe_ratio(hodl_strat), label = \"HODL\", marker = \".\", markersize = 20,\n",
|
||||
" color=palette[1])\n",
|
||||
"ax[3].plot(plt_times, running_sharpe_ratio(value_func(voltron)), \n",
|
||||
" label = \"Voltron\", marker = \".\", markersize = 20, color = palette[-2])\n",
|
||||
"\n",
|
||||
"ax[1].set_ylabel(\"P(increase)\")\n",
|
||||
"ax[2].set_ylabel(\"Portfolio Value\")\n",
|
||||
"ax[3].set_ylabel(\"Sharpe Ratio\")\n",
|
||||
"\n",
|
||||
"ax[2].legend(ncol = 4, loc = \"lower center\", bbox_to_anchor = (-0.2, -0.4))\n",
|
||||
"plt.subplots_adjust(wspace=0.25)\n",
|
||||
"sns.despine()\n",
|
||||
"[ax[i].set_xlim((-0.1, 5.1)) for i in range(4)]\n",
|
||||
"plt.show()\n",
|
||||
"# plt.savefig(\"trading_strategy.pdf\", bbox_inches = \"tight\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 136,
|
||||
"id": "7a2735a3",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"tensor(5.)"
|
||||
]
|
||||
},
|
||||
"execution_count": 136,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"ts.max()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 166,
|
||||
"id": "e72f5845",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAACAoAAAHSCAYAAABv67sOAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdd3wU5dYH8N9sTXaT3Wx6QgJJaNLBiiCCgnKxgoqACFZ89Yody0X04lXseuVarldQFAtYELAjKEVp0gmdENJ7stnN9t2Zef9YErLZ2b6bbJLz/XyU7MwzMychJLPznOcchud5HoQQQgghhBBCCCGEEEIIIYQQQgjpFkQdHQAhhBBCCCGEEEIIIYQQQgghhBBC2g8lChBCCCGEEEIIIYQQQgghhBBCCCHdCCUKEEIIIYQQQgghhBBCCCGEEEIIId0IJQoQQgghhBBCCCGEEEIIIYQQQggh3QglChBCCCGEEEIIIYQQQgghhBBCCCHdCCUKEEIIIYQQQgghhBBCCCGEEEIIId0IJQoQQgghhBBCCCGEEEIIIYQQQggh3QglCnRSt956K2699daODoMQQgghhJB2R/fChBBCCCGku6J7YUIIIYSEi6SjAyDBqays7OgQCCGEEEII6RB0L0wIIYQQQroruhcmhBBCSLhQRQFCCCGEEEIIIYQQQgghhBBCCCGkG6FEAUIIIYQQQgghhBBCCCGEEEIIIaQboUQBQgghhBBCCCGEEEIIIYQQQgghpBuhRAFCCCGEEEIIIYQQQgghhBBCCCGkG6FEAUIIIYQQQgghhBBCCCGEEEIIIaQbkXR0AIQQQgghhBBCSFfRUFiBgys3oDq/ENYmE+TxCqQNycPQ6ROQmJfZ0eERQgghhJBOjO41CSGEhBMlChBCCCGEEEIIISGqOVqETYuWo3TbITBiEXiWa9lXvvsY9iz9AdmjBmPc07OROiCn4wIlhBBCCCGdDt1rEkIIiQRqPUAIIYQQQgghhISgeGs+Vty4AGU7jwCAy4Pb1q/Ldh7BihsXoHhrfrvHSAghhBBCOie61ySEEBIplChACCGEEEIIIYQEqeZoEdbMeQUOq93toW1bPMuBtdqxZs4rqDla1D4BEkIIIYSQTovuNQkhhEQSJQoQQgghhBBCCCFB2rRoOVi7A+B5v8bzPA/W7sDmRZ9GODJCCCGEENLZ0b0mIYSQSKJEAUK6geKKBnyyZge+33gQJrOto8MhhBBCCCGkS2gorEDptkM+V3e1xbMcSrblQ3u6MkKREUIIIYSQzo7uNQkhhEQaJQoQ0sVt21+IK+76D55e/B3u+9dKTH1kCfQGS0eHRQghhBBCSKd3cOUGMOLg3lYzYhEOrFgf5ogIIYQQQkhXQfeahBBCIo0SBQjp4l7/aD1MlrNVBA4cL8fa3w90YESEEEIIIYR0DdX5hQGv8GrGsxxqDp0Oc0SEEEIIIaSroHtNQgghkSbp6AAIIZHjYFn8lV/stv2/K7dg1nUXdUBEhBBCCCGEdA08x8PcoA/pHFa9MUzRkPbSUFiBgys3oDq/ENYmE+TxCqQNycPQ6ROQmJfZ0eERQgghpAuxNplCO57uNQkhhPhAiQKEdGG1DQbB7SWVWuw7WooRA7LbOSJCCCGEEEI6N57joKusR31RJfgQzyVXKcMSE4m8mqNF2LRoOUq3HQIjFrms7ivffQx7lv6A7FGDMe7p2UgdkNNxgRJCCCGky5DHK0I7nu41CSGE+ECtBwjpwiprPa9wuvnRpSgoqWnHaAghhBBCCOm8OJaFtqQap/7MR+Xh07AZLVD3SgcjYoI6HyNioO6VBp4PNd2ARFrx1nysuHEBynYeAQC3EsDNr8t2HsGKGxegeGt+u8dICCGEkK4nbUgeGHFwUziMWITUwblhjogQQkhXQ4kChHRhlbU6j/vMFjtWbzjQjtEQQgghhBDS+bAOFvVFlTj150FUHSuG3WJt2Zc9Zih4LriJfp7jkTQwFyV7jsNmtIQrXBJmNUeLsGbOK3BY7T57BPMsB9Zqx5o5r6DmaFH7BEgIIYSQLmvo9Ak+7z884VkOw2ZcEeaICCGEdDWUKEBIF2Y0W73uX/zpxnaKhBBCCCGEkM6FtTtQe6ocp/44gJoTpXBY7W5j4tISkdS/Z8BVBRgRg6RzekKZpoGpQY/TOw6h7nQFeC64B8EkcjYtWg7W7gD8rPzA8zxYuwObF30a4cgIIYQQ0tUl5mUie9TggKsKMGIReo4eAk1uRoQiI4QQ0lVQogAhXZjN5ujoEAghhBBCCOlUHFY7ak6UouCPA6g7Ve6cJPZiwNRxEEnEAONnsgDDgBGLMeCmcS2bOJZD7ckyFO08ArPOEEL0JJwaCitQuu1QwCv5eJZDybZ8aE9XRigyQgghhHQX456eDbFUAsbPe02GYSCWSjB2/qwIR0YIIaQroEQBQrowq531Oaa6Xt8OkRBCCCGEEBLd7GYrqo4W49SfB1BfVAnO4fteOiZegQFXXYwpS5+CRC71udqLETEQScQ4//4pUGWluu23NJlQ/NdRVB8v8ev6JLIOrtwQUl/gAyvWhzkiQgghhHQ3qQNyMHnJkxD7ea8plksxecmTSB2Q0z4BEkII6dQoUYCQLszqR0WB8256uR0iIYQQQgghJDrZTBZUHj6NU1vzoS2tBufH6vFYdRyyhvdFzshBUKUnodclQzFj1QvIvmgQALg9xG1+nX3xYNzw0T+QMbyfx3PzPI+G4ioUbj8EQ50uhM+MhKo6vzCkvsA1h06HOSJCCCGEdEe9Rg/xfq95pg1WYr9sXP/BE+g1eki7x0gIIaRzknR0AISQyLH5KJParKSyAT0zEiMcDaA3WMCDhzouNuLXIoQQQgghxBurwYS605VoqmoA72f/eUWiCsm5GVAkqtzKv6YOyMHUz56B9nQlDqxYj5pDp2HVGyFXKZE6OBfDZlzR0ieWc7CoLSiDtrTG47XtZitK9x6HOjMZqf2yIZFJQ/uEid84loW+qgFNVQ0hnceqN4YpIkIIIYR0d23vNSv3nYShRgtJrBzqnmnoOWYYlGkaSJT03JUQQoj/KFGAkC7Mn4oCAFBQUhvRRAG7g8Wzb3+Pz77fBZ7nMXn8MLz86GTEKeQRuyYhhBBCCCFCzDoD6k9XoqlG6/cxcckJSMrNgEIT73OsJjcD4+bP9jpGJBEj7ZxeUKUnofJIEawGk8exuoo6GOt0SO3fE6r0RL/705LA2c1WaMtqoCuvg8Nmh1gW2iMTmUoRpsgIIYQQQpya7zV5nkfh1oOwmawu+w11jXDY7JRkSgghxC+UKEBIF+ZvRQGxKLJdSL76ZS8+/e6vltdrfjuAfjmpePDWyyJ6XUIIIYQQQgBnOX+Ttgn1pythrPevnD/DMIhP1SApNwMxKmVE4opNiEPuyIGoL6pCXWEFeE64zL3DZkdF/inoK+uRNqAXZLGUcBsuzd8b2pJqGGobXSo8qHulQ3uqHDznX8WJ1hgRg9jkBFgNJsjjKGGAEEIIIeHFMAxUGcmoO1Xusp3neeirGpDYM62DIiOEENKZUKIAIV2Y1c9EAYvVHtE4vvplj9u2b37dR4kChBBCCCEkYA2FFTi4cgOq8wthbTJBHq9A2pA8DJ0+AYl5mS5jeZ6HsU6H+qJKmLRNfp3f+dA1CUk5GZC3Q8ssRiRCcl4m4tM0qDpS5DVOQ10jTNubkNInC5qs1JZ+tCRwnIOFrrIe2tIajxUdsscMxekNu4M6P8/xyLxgAE5vPwxNzzQk52VCLKVHMIQQQggJH3VGkluiAOCsSEWJAoQQQvxB71IJ6cL8bT1gjmCigINlsedwidv2wtI67DlSgvMG9ozYtQkhhBBCSNdRc7QImxYtR+m2Q2DEIvDs2dX35buPYc/SH5A9ajDGPT0bKef0QlONFvWnK2Hxs088IxIhITMZiTnpkCliIvVpeCRXxqLn+eegsawGtSfLwDpYwXGcg0X1sWLoK+uRMSiHVqsHyGayQFtaA11FHVgfidVxaYlI6t8TDSdLA6oqwIgYJPbLhjJNA57n0VBcBX1VPVL7ZkOVkUTtIwghhBASFjJFDBSaeLdEU4veSFWNCCGE+CWy9cYJIR3KZhN+uNiW2RK5RIGSCs+9X6+//30cOF4WsWsTQgghhJCuoXhrPlbcuABlO48AgEuSQOvXZTuO4IsbnsbOJd+j/ECBX0kCIrEIib3S0fuSoUgfmNMhSQLNGIaBJjsNuaOGID5F43WsWWfA6R1HUFtQDo4VbllAnHieh6FOh9J9J1C4NR8NxVU+kwSaDbvtbxBJJYC/k/sMA0YsxoCbxrlsdljtqDhUiOJdR/1OXiGEEEII8UWdkSy4XVdR186REEII6YwoUYCQLsxq8y8BwOLnuGBU1DR63f/eii0RuzYhhBBCCOn8ao4WYc2cV+Cw2t0SBNriOQ6szYHtb6yEvqzG61ixRIzkvEz0HjMMaf17QhojC2fYIZHGyNBjeB/0GNYHErnU4zie41BXWI6iHYf9bq3QnbAOFg0l1Sjclo/SvcdhqG0Ez/uuDMAwDFTpSeh1wQAMnzYeU5Y+BYlcCkbs/REKIxJBLBXj/PunQJWVKjjG3GhA0c4jqDpa5HeyAiGEEEKIJ/HpiRAJ3KPoKusDqohECCGke6JEAUK6qAadET9sPuTX2EhWFPjrULHX/T9uPgSWVkARQgghhBAPNi1a7pxQ9WOCFwDA8+BZFke/2SS4WyKTIqVvFnqPGYaUPlmQyDxPxHckhmGgSktE3qghSOiR4nWs1WhG8a6jzslnDy0LuhOr0Yyqo8Uo2LIf1ceKYTNa/DpOIpMiOa8Heo8Zhh5De0OhiQfDMOg1eghmrHoB2RcNAgC3hIHm19kXD8KMr19A/0kjwYg8P27heR7a0hoUbs1HY1kNPcQnhJAuTq/XY8yYMejfvz+eeuqpjg6HdDFiiRhxqe6VqBxWO4z1ug6IiBBCSGci6egACCHhV1RLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 2500x400 with 3 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"\n",
|
||||
"fig, ax = plt.subplots(1, 3, figsize = (25, 4))\n",
|
||||
"ax[0].plot(ts, y)\n",
|
||||
"[ax[i].set_xlabel(\"Time\") for i in range(3)]\n",
|
||||
"ax[0].set_ylabel(tckr)\n",
|
||||
"[ax[0].axvline(x=eval_times[i], alpha = 0.2, linestyle=\"--\") for i in range(len(eval_times))]\n",
|
||||
"\n",
|
||||
"# ax[1].plot(plt_times, matern, label = \"Matern\", color=palette[1], alpha = 0.5)\n",
|
||||
"# ax[1].plot(plt_times, specmix, label = \"Spectral Mixture\", color=palette[3], alpha = 0.5)\n",
|
||||
"# ax[1].plot(plt_times, voltron, label = \"Voltron\", color = palette[-1], alpha = 0.5)\n",
|
||||
"# ax[1].scatter(plt_times, matern, s = 120, label = \"Matern\", color=palette[0], zorder=4)\n",
|
||||
"# ax[1].scatter(plt_times, specmix, s = 120, label = \"Spectral Mixture\", color=palette[2], zorder=4)\n",
|
||||
"# ax[1].scatter(plt_times, voltron, s = 120, label = \"Voltron\", color = palette[-2], zorder=4)\n",
|
||||
"\n",
|
||||
"ax[1].plot(plt_times, value_func(matern), label = \"Matern\", color=palette[1], alpha = 0.5)\n",
|
||||
"ax[1].plot(plt_times, value_func(specmix), label = \"SM\", color=palette[3], alpha = 0.5)\n",
|
||||
"ax[1].plot(plt_times, hodl_strat, label = \"Hold\", color=palette[5], alpha = 0.5)\n",
|
||||
"ax[1].plot(plt_times, value_func(voltron), label = \"Voltron\", color = palette[-1], alpha = 0.5)\n",
|
||||
"ax[1].scatter(plt_times, value_func(matern), color=palette[0], zorder=4, s=120)\n",
|
||||
"ax[1].scatter(plt_times, value_func(specmix), color=palette[2], zorder=4, s=120)\n",
|
||||
"ax[1].scatter(plt_times, hodl_strat,color=palette[4], zorder=4, s=120)\n",
|
||||
"ax[1].scatter(plt_times, value_func(voltron),color = palette[-2], zorder=4, s=120)\n",
|
||||
"\n",
|
||||
"ax[2].plot(plt_times, running_sharpe_ratio(value_func(matern)), \n",
|
||||
" label = \"Matern\", color=palette[1], alpha = 0.5)\n",
|
||||
"ax[2].plot(plt_times, running_sharpe_ratio(value_func(specmix)), \n",
|
||||
" label = \"SM\", color=palette[3], alpha = 0.5)\n",
|
||||
"ax[2].plot(plt_times, running_sharpe_ratio(hodl_strat), label = \"HODL\", markersize = 20,\n",
|
||||
" color=palette[5], alpha = 0.5)\n",
|
||||
"ax[2].plot(plt_times, running_sharpe_ratio(value_func(voltron)), \n",
|
||||
" label = \"Voltron\", color = palette[-1], alpha = 0.5)\n",
|
||||
"ax[2].scatter(plt_times, running_sharpe_ratio(value_func(matern)), \n",
|
||||
" s = 120, color=palette[0], zorder=4)\n",
|
||||
"ax[2].scatter(plt_times, running_sharpe_ratio(value_func(specmix)), \n",
|
||||
" s = 120, color=palette[2], zorder=4)\n",
|
||||
"ax[2].scatter(plt_times, running_sharpe_ratio(hodl_strat),\n",
|
||||
" s = 120, color = palette[4], zorder=4)\n",
|
||||
"ax[2].scatter(plt_times, running_sharpe_ratio(value_func(voltron)), \n",
|
||||
" s = 120, color = palette[-2], zorder=4)\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"# ax[1].set_ylabel(\"P(increase)\")\n",
|
||||
"ax[1].set_ylabel(\"Portfolio Value\")\n",
|
||||
"ax[2].set_ylabel(\"Sharpe Ratio\")\n",
|
||||
"\n",
|
||||
"ax[2].legend(ncol = 4, loc = \"lower center\", bbox_to_anchor = (-0.85, -0.5))\n",
|
||||
"plt.subplots_adjust(wspace=0.35)\n",
|
||||
"sns.despine()\n",
|
||||
"[ax[i].set_xlim((-0.1, 5.1)) for i in range(3)]\n",
|
||||
"# plt.savefig(\"trading_strategy_\" + tckr + \".pdf\", bbox_inches = \"tight\")\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "b8c4b178",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "e33ac230",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Do everything"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 275,
|
||||
"id": "e2c15907",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"tckrs = [\"BAC\", \"BRK.B\", \"CVX\", \"EOG\", \"JPM\", \"XOM\", \"WFC\", \"COP\", \"C\", \"SLB\"]\n",
|
||||
"spdrs = [\"XLF\", \"XLF\", \"XLE\", \"XLE\", \"XLF\", \"XLE\", \"XLF\", \"XLE\", \"XLF\", \"XLE\"]\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"def running_sharpe_ratio(vec):\n",
|
||||
" returns = (vec - 10000)\n",
|
||||
" std_returns = torch.tensor([returns[:i].std(0) for i in range(len(vec))])\n",
|
||||
" # need avg return divided by sd of returns?\n",
|
||||
" return returns.cumsum(0) / std_returns / torch.arange(vec.shape[0])\n",
|
||||
"\n",
|
||||
"def reward_risk(a, b, prob_incs):\n",
|
||||
" bought_func = lambda xs: betainc(a, b, xs)\n",
|
||||
" total_held = 1000 * bought_func(prob_incs)\n",
|
||||
" \n",
|
||||
" returns = total_held[1:] * delta_y\n",
|
||||
" #cum_returns = returns.cumsum(0)\n",
|
||||
" return returns.std(), returns.sum()\n",
|
||||
"\n",
|
||||
"def value_func(vec, base = 10000):\n",
|
||||
" portfolio_value = torch.zeros(12)\n",
|
||||
" portfolio_value[0] = base\n",
|
||||
" for i in range(11):\n",
|
||||
" price_of_stock = portfolio_value[i] * bought_func(vec)[i]\n",
|
||||
" amt_bought = price_of_stock / prices_at_time_y[i]\n",
|
||||
" cash_left = portfolio_value[i] - price_of_stock\n",
|
||||
" portfolio_value[i+1] = cash_left + amt_bought * prices_at_time_y[i+1]\n",
|
||||
" return portfolio_value"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 276,
|
||||
"id": "8eba99f2",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAACAoAAAHeCAYAAAAYKXt1AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdd3hUZdoG8PvMTDKpk94bCb0LIiKIIrB2KUtRmnXxs6BYUBTBRQWVteyyYlllEUQBKYIgihSFVUAEpCTUQHovM5Oeqef7I2ZImJlkJpmScv+ui8uT85bzBIScOed5n1cQRVEEERERERERERERERERERERdQoSdwdARERERERERERERERERERErsNEASIiIiIiIiIiIiIiIiIiok6EiQJERERERERERERERERERESdCBMFiIiIiIiIiIiIiIiIiIiIOhEmChAREREREREREREREREREXUiTBQgIiIiIiIiIiIiIiIiIiLqRJgoQERERERERERERERERERE1IkwUYCIiIiIiIiIiIiIiIiIiKgTYaIAERERERERERERERERERFRJyJzdwBERERERERERERERB3B5cuXsX79evz6668oLCwEAMTFxeGWW27BAw88gODgYDdHSERERFRHEEVRdHcQRERERERERERERETt2erVq/Huu+9Cp9NZbA8JCcFHH32Ea665xrWBEREREVnARAEiIiIiIiIiIiIiolZYu3YtlixZAgDw9vbG5MmT0b9/f9TW1mLnzp04cuQIACAgIAA7d+5EWFiYO8MlIiIiYqIAEREREREREREREVFL5eTk4K677kJtbS2Cg4OxZs0a9OjRo1GfJUuWYO3atQCAmTNnYtGiRe4IlYiIiMhE4u4AiIiIiIiIiIiIiIjaqw8//BC1tbUAgH/9619mSQIA8OKLLyI4OBgA8MMPP7g0PiIiIiJLZO4OgIiIiIiIiIiIiIioPdJqtdi9ezcAYPTo0bj++ust9vP09MScOXOQkZGBoKAgaLVaeHp6ujJUIiIiokaYKEBERERERERERERE1AKHDx9GZWUlAGDixIlN9p0xY4YrQiIiIiKyCRMFiIiIiIiIiIiIiIha4Pz586bjgQMHmo6VSiXS0tKg0WiQkJCA2NhYd4RHREREZBUTBYiIiIiIiIiIiIiIWiA1NRVA3dYCERERyMrKwttvv40DBw5Ar9eb+vXv3x8LFizA4MGD3RUqERERUSMSdwdARERERERERERERNQeFRYWAgACAgJw9OhRjB8/Hvv27WuUJAAAycnJmDVrFnbu3OmOMImIiIjMMFGgnZo5cyZmzpzp7jCIiIiIiFyO98JERERE1FZUVVUBAGpqajBnzhxUV1dj8uTJ+O6775CcnIy9e/di9uzZkEgk0Ov1eOmll3Du3LkWX4/3wkREROQo3HqgncrPz3d3CEREREREbsF7YSIiIiJqK+oTBSorKwEATz/9NJ588klTe1xcHObNm4fY2Fj8/e9/h1arxTvvvINVq1a16Hq8FyYiIiJHYUUBIiIiIiIiIiIiIqJW6tGjB5544gmLbffddx8GDhwIADh48CBf+BMREZHbMVGAiIiIiIiIiIiIiKgFvL29Tcd33XUXBEGw2vf22283Hf/xxx9OjYuIiIioOUwUICIiIiIiIiIiIiJqAT8/P9NxUlJSk30TExNNx4WFhU6LiYiIiMgWTBQgIiIiIiIiIiIiImqB2NhYm/t6enqajo1GozPCISIiIrIZEwWIiIiIiIiIiIiIiFqgR48epuPc3Nwm+5aUlJiOIyIinBYTERERkS1k7g6AiIiIiIiIiIiI2i9lWh5Ob9iLwuQ0aCqqIff3QUT/JAy4byyCk6LdHR6RU1133XWm4//973946KGHrPY9efKk6bhhggERERGROzBRgIiIiIiIiIiIiOxWdC4D+5d+gexDKRCkEoiGK6XUc4+dx/GV3yFueD+MeuV+hPfu4r5AiZyof//+SEhIQGZmJg4fPozTp09jwIABZv1UKhV27twJAOjatSt69uzp6lCJiIiIGuHWA0RERERERERERGSXzIPJWD9pIXKOnAWARkkCDb/OOXIW6yctRObBZJfHSOQqTzzxBABAFEXMmzcPeXl5jdq1Wi1efPFFlJWVAQAefPBBV4dIREREZIYVBYiIiIiIiIiIiMhmRecysG32Mug1OkAUm+wrGowwGHXYNnsZpm1ZwsoC1CFNmDAB+/btw+7du5GZmYlx48ZhypQp6NOnD0pLS/H1118jLS0NADB06FBMmTLFzRETERERMVGAiIiIiIiIiIiI7LB/6Rcw6PTNJgnUE0URBp0eB5auxZQvFzk5OiL3eP/997Fw4UJs27YNFRUVWLVqlVmfG2+8EcuXL4cgCG6IkIiIiKgxJgoQERERAVCVVUPh5wWptPHOTEXKCvzwvzMICfTFmBt6wVvu4aYIiYiIiIjcT5mWh+xDKXaPEw1GZB1Khio9H0GJUU6IjMi9PDw8sGzZMkycOBGbNm3C8ePHUVpaisDAQHTv3h333nsv/vKXv0Ai4W7ARERE1DYwUYCIiIg6tQvphXj8tfW4mFmE0CBfLJk7Dnff3B8AcOJcNibN/RRanQEAcE2vWKx752Eo/LzcGTIRERERkduc3rAXglQC0WC0e6wgleDU+j0YteB+J0RG1DYMGzYMw4YNc3cYRERERM1i+iIRERF1WqIo4sklX+NiZhEAoERVhTlvfI284jKIoojnlm02JQkAwMnzOdixP7nF1zMYjFBX1KCmVtvq2ImIiIiI3KEwOa1FSQJAXVWBopR0B0dEREREREQtwUQBIiIi6rROX8zF+bSCRuf0BiMOnUjD6Yu5SM0sNhtzNDmjRdc6n16Avzzyb/Qb9wauv+8f2LLnRIvmISIiIiJyJ01FdevGl1c5KBIiIiIiImoNJgoQERFRp3X4pOXVTM+8tQmHTqRZbPvpyIUWXev5ZVtMlQuUZdV48Z2tyCtSt2guIiIiIiJ3kfv7tG68wtdBkRARERERUWswUYCIiIg6raw8pdW2pf/ZZfG8v6+X3dfJLlDh1IXcRuc0Oj2+3PG73XMREREREblTRP8kCNKWPVIUpBKE90t0cERERERERNQSTBQgIiKiTisz33qigDXSFjwUzS5QWTz/7y/32z0XEREREZE7DbhvLESDsUVjRYMRA6f9xcERERERERFRSzBRgIiIiDqtrBYkCqRll9g9RlVmfR/XC+mFds9HREREROQuwUnRiBveD4LEvseKglSC+BH9EZQY5aTIiIiIiIjIHkwUICIiok6rqLSiReNSUvPs6l+iqrTatmzl7hbFQERERETkLqMWzIJEJgEEwab+giBA6iHDzQtmOTkyIiIiIiKyFRMFiIiIqNOq1ehbNG7jruN29S9SWk9I2H3oHHR6Q4viICIiIiJyB6+QAFz7xERIZFIIkqaTBQSpBFK5ByZ8Nh/hvbu4JkAiIiIiImoWEwWIiIioU9IbDDAYW7a36vGz2Xb1zy1UN9m+4F/foqpG06JYiIiIiIhcyWgwojg1G6G9EjB8/nQE94gDALOEAUFa99gxblhfTNuyBAkj+rs8ViIiIiIisk7m7gCIiIiI3EGjbVk1AQA4dT7Hrv7ZBaom29fvPIZvdp/ExR8WQyplHicRERERtV2q7EJoq+uSXBWx4bj+mamoKlQh51AKNOoKaCuqIVf4IrxfIgZO+wuCEqPcHDEREREREVnCRAEiIiLqlLTa1pX7T80sQveEcJv6ZuU3nSgAABqdHrNf/QqrlnLfViIiIiJqm/RaHUrT8szO+0YE4ZZFDyAgOtQNURERERERUUtwyRoRERF1SrVaXavGb9t3yqZ+Gq0eBSXlNvXdfegcSlSVrQmLiIiIiMhpSi7lwqA3T7j1VvhCERXihoiIiIiIiKilmChAREREnVJrth4AgOVrf7apX26RGqIo2jzvbY+uaGlIREREREROU1tRDXVuscW28J7xEATBxREREREREVFrMFGAiIiIOqXWJgrYKtuGbQcaKiwpx+kLuU6KhoiIiIjIfqIoouhilsUEWEVkMHyC/N0QFRERERERtQYTBYiIiKhTsiVR4Jbre+Dum/u16jrKsiq7x6xYt9/i+V+PX8LsV7/CAy+vwc4DKa2Ki4iIiIjIVlUlZagqNd9OS5BIENYt1g0RERERERFRazFRgIiIiDql5hIFfL098cEr92L4oKRWXadGo7N7zPn0QrNzR06nY8aLq/HDL2ew77cL+L/F6/Dd/uRWxUZERERE1BzRaETRxWyLbcEJEfD08XJxRERERERE5AhMFCAiIqJOSaNt+gX+0rnjEOjvjdtH9rXa52hyZrPXqam1P1GgrKLG7Ny6747BYDQ2PrfzqN1zExERERHZQ5VTDE2V+f2pzNMDIYnRboiIiIiIiIgcgYkCRERE1Ck1V1HA+Of+q+HB1vdb3XXwbLPXsVZRoFt8GKLDAyy2qctrzPZ/3bLnhFm//x271Oz1iYiIiIhayqDTo+RyrsW2sG4xkMqkLo6IiIiIiIgchYkCRERE1CnVNpMooNUZTMcvP3qbxT7HUpqvKFBZVWvx/HcfP4Hfv56PubNuMWszGI2orNY0OzcRERERkTOVXM6FQWd+3+zl74OA6DA3RERERERERI7CRAEiIiLqlKyt9K83tH+C6fieUf0t9rG0RUBDWp0eK9YdsNjmJZcBAAL8vJudW6c3WOwDAAaD0WobEREREVFLaapqoMoustgW3jMegkRwcURERERERORITBQgIiKiTkmprrLa1j0hDN0Twk1fx0cFIyYi0KxfU6v+azQ6PPv2Zottnh5SyKR1ZVoDFZYTBdQNEgWKSiusXqe80nLFAiIiIiKi1ii6mG22HRYA+IcFwTdY4YaIiIiIiIjIkWTuDoCIiIjIHUrUlRbPD+gRg/+8Nh2C0HiF1PBrkrDpxz8anauu0eJyVjH2H03F2u1HoCyLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 2500x400 with 3 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAACAoAAAHeCAYAAAAYKXt1AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdd3hUZdoG8PvMTKal9xASUqghBGxUQVCwgVJEBERQV2F1BXtblF0s6LKufLKK7lpRVNhFBFEUBVRWAZEmCb2E9EYy6dNnzvdHzJAwJTPJTCbl/l2X186c9z3nPANZcuac530eQRRFEURERERERERERERERERERNQtSPwdABEREREREREREREREREREbUfJgoQERERERERERERERERERF1I0wUICIiIiIiIiIiIiIiIiIi6kaYKEBERERERERERERERERERNSNMFGAiIiIiIiIiIiIiIiIiIioG2GiABERERERERERERERERERUTfCRAEiIiIiIiIiIiIiIiIiIqJuhIkCRERERERERERERERERERE3QgTBYiIiIiIiIiIiIiIiIiIiLoRmb8DICIiIiIiIiIiIiLqCs6ePYu1a9fi559/RmlpKQAgMTERV199Ne68805ERET4OUIiIiKiBoIoiqK/g/Cl5557Dp9++ikWLlyIRYsW2Y1//vnn+POf/+zxcYcNG4Y1a9Y021ZYWIhrrrnGrf1TUlKwdetWj89LRERERERERERERB3P6tWr8Y9//AMmk8nheGRkJN58801ccskl7RsYERERkQNduqLAnj17sG7dOp8cWxAEu20nT570ybmIiIiIiIiIiIiIqONas2YNXn75ZQCASqXCrbfeioyMDOj1emzZsgV79+5FRUUFFixYgC1btiA6OtrPERMREVF312UTBY4cOYKFCxfCarW6nDdixAisWrWqxeMZDAb85S9/QV1dHWQyGf70pz/ZzWmaKLBixQooFAqnxwsMDGzxnERERERERERERETUsRUUFOAf//gHACAiIgIffvgh+vXrZxufOXMmXnzxRaxZswbV1dX417/+hSVLlvgrXCIiIiIAXTRRYOfOnXj88cdRV1fX4tz4+HjEx8e3OG/JkiW24z388MMYMWKE3ZzGRIHo6GhMmjTJw6iJiIiIiIiIiIiIqLNZtWoV9Ho9AOC1115rliTQ6Mknn8SWLVug0WjwzTffMFGAiIiI/K5LJQoYjUb861//wltvvdViJQFP7Ny5E//9738BAEOHDsU999zjcF5jooCjC0EiIiIiIiIiIiIi6lqMRiO+++47AMA111yD4cOHO5wnl8uxcOFC5OTkIDw8HEajEXK5vD1DJSIiImqmyyQK7N69G0uWLEFBQQEAQK1WY+bMmfjggw/adFydToelS5cCaLiYe+GFFyCRSOzmGQwG5ObmAmCiABEREREREREREVF3sGfPHlsl2mnTprmcO2fOnPYIiYiIiMgt9k+8O6nNmzfbkgQGDRqE9evX4+qrr27zcd9++20UFRUBAP7whz8gJSXF4bzTp0/DYrEAYKIAERERERERERERUXdw4sQJ2+shQ4bYXms0Guzfvx+7du2y3bcmIiIi6ki6TEUBAIiIiMDChQsxa9YsSKVSVFRUtOl45eXltooEUVFRWLBggdO5p06dsr3u27cvAKC0tBRnz56FKIpISEhAUlJSm+IhIiIiIiIiIiIioo7j9OnTABqq0cbGxiIvLw9/+9vfsHPnTpjNZtu8jIwMLF68GJdddpm/QiUiIiJqpsskCsyZMwdLly6FUqn02jHffvtt6HQ6AMCCBQsQGBjodO7JkycBABKJBPn5+Xj++eeRmZnZbE5ycjIeeOABTJ482WsxEhEREREREREREZF/lJaWAgBCQ0Oxb98+LFiwAFqt1m5eVlYW5s6di7///e+YNGlSe4dJREREZKfLtB7IyMjwapJAXV0dPvvsMwBAWFgYZsyY4XJ+Y6KA1WrFI488YpckAAA5OTl44okn8Mwzz8BqtbYpvjvuuAN33HFHm45BRERERNQZ8VqYiIiIiDqK+vp6AIBOp8PChQuh1Wpx66234quvvkJWVha2b9+O+fPnQyKRwGw24+mnn8bx48dbfT5eCxMREZG3dJmKAt62YcMG20XerFmzoFarXc5v2nogOjoa999/P8aPH4+IiAgUFhZi8+bNePfdd2E0GvHZZ58hLi4OixYtanV8xcXFrd6XiIiIiKgz47UwEREREXUUjfeQ6+rqAAAPPvggHnjgAdt4YmIiHn/8cSQkJOCvf/0rjEYjXnnlFbz//vutOh+vhamjOn+2EOVnC+2297v6MkgD+CiKiKgj6jIVBbxJFEV88sknABp6S7WUoVlVVQW9Xg8ASEpKwsaNGzFnzhzExcVBLpcjJSUFDz30EN5++23IZA2/EP/973+jqKjItx+EiIiIiIiIiIiIiNpFv3798Kc//cnh2KxZszBkyBAAwK5du/jAn7ocwd8BEBGRx5go4MAvv/yC3NxcAMANN9yA6Ohol/PDwsJw8OBB7N69G+vWrXM6f+TIkZg9ezYAwGQy4fPPP/du4ERERERERERERETUblQqle31pEmTIAjOH5fecMMNttcHDx70aVxE7c7Jz74oiu0cCBERuYuJAg5s27bN9rrpxVtLIiMjERER4XLOlClTbK8PHDjgeXBERERERERERERE1CEEBQXZXqemprqcm5KSYntdWlrqs5iI/IIlBYiIOh0mCjjwww8/AAACAwMxZswYrx676cViWVmZV49NRERERERERERERO0nISHB7blyudz22mq1+iIcIr8RnGUKsKAAEVGHxUSBixw/fhxFRUUAgKuvvrrZxZs3BAQEePV4REREREREREREROQf/fr1s70uLCx0Obe8vNz2OjY21mcxEfmFszwBth4gIuqwZP4OoKP56aefbK/Hjx/v1j4HDhzAoUOHUFFRgdtuu61ZCamLNS0pFR0d3fpAiYiIiIiIiIioW9BkFyFz3XaUZmXDUKuFIliN2IxUDJ41ARGp8f4Oj6hbGzp0qO31//73P9x9991O5/7222+2100TDIiIiIj8gYkCFzl06JDt9ZAhQ9za59dff8Vrr70GoCET1FWiwO7duz0+PhERERERERERdT9lx3Pw47KPkL/7CASpBKLlQqnywv0ncODdr5A4ahDGPTMPMWnJ/guUqBvLyMhAUlIScnNzsWfPHmRmZmLw4MF28yorK7FlyxYAQO/evdG/f//2DpXIpwTBaUmB9g2EiIjcxtYDFzly5AgAICwsDD179nRrnxEjRthef/bZZ7BYLA7nGQwGvP/++wAafmlOnjy5jdESEREREREREVFXlLsrC2unP4uCvccAoFmSQNP3BXuPYe30Z5G7K6vdYySiBn/6058ANJRYf/zxx22tbRsZjUY8+eSTqK6uBgDcdddd7R0ikd8wTYCIqONiokATNTU1KCsrAwCkpaW5vd+ll16K9PR0AMDp06fx8ssv2/XdMRqNeOKJJ5CTkwMAmDx5Mnr37u2dwImIiIiIiIiIqMsoO56DTfOXw2ww2SUIXEy0WGExmLBp/nKUHc9pnwCJqJmpU6fiuuuuAwDk5uZi8uTJWL58Ob788kusXr0aU6ZMwf/+9z8AwLBhwzBjxgx/hkvkG6woQETU6bD1QBN5eXm215GRkR7t+8ILL2DOnDnQ6XRYs2YNDh8+jClTpiAiIgJ5eXnYsGGD7fh9+vTB4sWLvRo7ERERERERERF1DT8u+wgWk9nthyuiKMJiMmPnsjWY8fESH0dHRI6sWLECzz77LDZt2oTa2lpbZdmmRo8ejZUrVzov0U7UifGnmoio82GiQBOlpaW210FBQR7tm56ejrfffhsPP/wwKioqkJmZiczMTLt5l112GVauXImwsLC2hktERETdjM5gwlc/ZEEiFXDV5X0QHRHs75CIiIiIyMs02UXI333E4/1EixV5u7NQea4Y4Sk9fBCZc5rsImSu247SrGwYarVQBKsRm5GKwbMmICI1vl1jIfKXgIAALF++HNOmTcP69etx4MABVFRUICwsDH379sXMmTNx7bXXQiJhkV/qopxlCrCgABFRh8VEgSbq6+ttr0NCQjzef9iwYdi6dSvWrVuH77//HmfPnoVOp0NERAQGDhyIm266CZMmTWLGKBEREXns2JlizH16NUoragEAcVEh+PSVu9EvOdbPkRERERGRN2Wu2w5BKmmx5YAjgkTAL29+jhEPzkCASoEAlRwBSjkkUqkPIm1okfDjso+Qv/uIXcyF+0/gwLtfIXHUIIx7Zh5i0pJ9EgNRRzNixAiMGDHC32EQ+YHj5x4iMwWIiDqsLp0oMHz4cJw8edLt+ZMnT8bkyZPbdM6QkBAsWLAACxYsaNNxiIiIiBqJoogHXlxnSxIAgJLyGrzz2S688vgtfoyMiIiIiLytNCu7VUkCACBaRZRkZaPkeE6z7TJFAAKUiobkAaX89wQChS2ZoDWJBLm7srBp/vKGFgmAXcyN7wv2HsPa6c9i6jtPIenKjFZ9LiIi6vicro9kngARUYfVpRMFiIiIiLqC/JJKnM49b7d9X1auH6IhIiIiIl8y1GrbtL9ZZ7DfZjDBbDBBV13ncB+ZPKBJBYLGZIIL7yWy5okEZcdzsGn+cpgNJkB0/QRItFhhsZqwaf5yzN7wIisLEBF1WcwUICLqbJgoQERERNTBnSuocLj9TJ598gARERERdW6KYHWb9pepFB7vYzaaYDaaoKt2ckx5QLPkge1L3oXFaG4xSaCRKIqwmMzYuWwNZny8xOP4iIioE3CSJ+DmrwoiIvIDib8DICIiIiLXzhWU+zsEIiIiImonsRmpEKStu2UnSASE9or1ckQNiQS6mnrUlGqQu/cYig+egmj1rD2CaLEib3cWKs8Vez0+IiLqAJz2HiAioo6KiQJEREREHdyOvSf9HQIRERERtZPBsyZAtHjLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 2500x400 with 3 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAB/kAAAHeCAYAAABwqrgpAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdeXxU9bk/8M85syUzyWRfyQQS1rCKiCKLcivWXVHAggu1tXhbq60/a11QK/aClnr1trdqF7yKK7gVFFArakEFRECWQMKafd8m20xmO+f8/ogZEmYmmUlmMlk+79fLlzPnfM85T0ggZ87zfZ6voCiKAiIiIiIiIiIiIiIiIiIiIhrwxHAHQERERERERERERERERERERP5hkp+IiIiIiIiIiIiIiIiIiGiQYJKfiIiIiIiIiIiIiIiIiIhokGCSn4iIiIiIiIiIiIiIiIiIaJBgkp+IiIiIiIiIiIiIiIiIiGiQYJKfiIiIiIiIiIiIiIiIiIhokGCSn4iIiIiIiIiIiIiIiIiIaJBgkp+IiIiIiIiIiIiIiIiIiGiQYJKfiIiIiIiIiIiIiIiIiIhokFCHOwAiIiIiIiIiIiIionA7c+YMNmzYgK+//hrV1dUAAJPJhP/4j//Aj3/8Y8THx4c5QiIiIqJ2gqIoSriDICIiIiIiIiIiIiIKl/Xr1+O///u/4XQ6ve5PSEjAiy++iPPOO69/AyMiIiLygkl+IiIiIiIiIiIiIhq2Xn/9daxevRoAEBkZicWLF2PKlCmw2WzYtm0b9u7dCwCIiYnBtm3bkJSUFM5wiYiIiJjkJyIiIiIiIiIiIqLhqaysDNdccw1sNhvi4+Px6quvYty4cV3GrF69Gq+//joA4LbbbsPjjz8ejlCJiIiI3MRwB0BEREREREREREREFA4vvPACbDYbAOBPf/qTR4IfAB588EHEx8cDAD7++ON+jY+IiIjIG3W4AyAiIiIiIiIiIiIi6m8OhwOffvopAOAHP/gBLrroIq/jtFot7rnnHhQVFSEuLg4OhwNarbY/QyUiIiLqgkl+IiIiIiIiIiIiIhp29uzZg9bWVgDAjTfe2O3YW2+9tT9CIiIiIvILk/xERERERERERERENOwcP37c/XratGnu1w0NDSgoKIDdbsfIkSORkZERjvCIiIiIfGKSn4iIiIiIiIiIiIiGnVOnTgFob8efkpKCkpIS/OEPf8DOnTvhcrnc46ZMmYKVK1fi/PPPD1eoRERERF2I4Q6AiIiIiIiIiIiIiKi/VVdXAwBiYmKwb98+3HDDDfj888+7JPgBIDc3F7fffju2bdsWjjCJiIiIPDDJHwa33XYbbrvttnCHQURERETU73gvTEREREQDhcViAQC0tbXhnnvugdVqxeLFi7F161bk5ubis88+w4oVKyCKIlwuFx5++GHk5+f3+nq8FyYiIqJgYbv+MKisrAx3CEREREREYcF7YSIiIiIaKDqS/K2trQCAX/3qV/jlL3/p3m8ymfDAAw8gIyMDTzzxBBwOB5555hm8/PLLvboe74VpoFEUBSc+2w9FUbpsjzOlIDVnZJiiIiIif7CSn4iIiIiIiIiIiIiGtXHjxuHuu+/2um/p0qWYNm0aAGDXrl1M1tOQIQgCRLXKY7ssSWGIhoiIAsEkPxERERERERERERENO5GRke7X11xzDQRB8Dn2yiuvdL/+7rvvQhoXUX/ymuR3MclPRDTQMclPRERERERERERERMNOVFSU+3V2dna3Y7Oystyvq6urQxYTUX9TeUnyS0zyExENeEzyExEREREREREREdGwk5GR4fdYrVbrfi3LcijCIQoLVvITEQ1OTPITERERERERERER0bAzbtw49+vy8vJux9bV1blfp6SkhCwmov4mqpjkJyIajNThDoCIiIiIiIiIiIjCo6GgAkc2fobq3ALYW6zQReuRMiUbU5cuQHx2erjDIwqpmTNnul9/+eWX+MlPfuJz7KFDh9yvO08OIBrsRA2T/EREgxGT/ERERERERERERMNMTX4Rdqx5DaW7j0JQiVCks+3Hy/cfx4GXtsI0ezLmP7ocyTmjwhcoUQhNmTIFI0eORHFxMfbs2YMjR45g6tSpHuPMZjO2bdsGABg9ejTGjx/f36EShYzKSyW/xCQ/EdGAx3b9REREREREREREw0jxrlxsWPQYyvbmAUCXBH/n92V787Bh0WMo3pXb7zES9Ze7774bAKAoCh544AFUVFR02e9wOPDggw+iqakJAHDHHXf0d4hEISWqPZP8iixDPud3AxERDSys5CciIiIiIiIiIhomavKLsHnFWrjsTkBRuh2rSDIk2YnNK9Zi2furWdFPQ9LChQvx+eef49NPP0VxcTGuv/56LFmyBBMnTkR9fT3efvttFBQUAAAuvPBCLFmyJMwREwWXt3b9ACBLEkQV60SJiAYqJvmJiIiIiIiIiIiGiR1rXoPkdPWY4O+gKAokpws717yOJW88HuLoiMLjueeew2OPPYbNmzejpaUFL7/8sseYuXPn4s9//jMEQQhDhESh461dPwDITgnQavo5GiIi8heT/ERERET9oKm1Ddt2HkWkToMr505EZIQ23CERERER0TDTUFCB0t1HAz5OkWSU7M6FubAScVlpIYiMKLw0Gg3Wrl2LG2+8Ee+++y4OHDiA+vp6xMbGYuzYsfjRj36Eyy+/HKLIqmYaery16wcAySX1cyRERBQIJvmJiIiIQuybw4VYfN869/uxI5Ow8dk7kZJgDGNURERERDTcHNn4GQSVCKUX6ywLKhGHN2zH/JXLQxAZ0cAwa9YszJo1K9xhEPUrX0l+WWKSn4hoIGOSn4iIiCiEjpwo75LgB4BTxbV4++MD+NVt/xGmqIiIiIhoOKrOLehVgh9or+avOVoY5IiIiCjcfCb5na5+joSGq4aCChzZ+Bmqcwtgb7FCF61HypRsTF26APHZ6eEOj2jAYpKfiIiIKIRW/ukDr9vf3LqPSX4iIiIi6lf2Fmvfjm+2BCkSIiIaKFS+2vX3clIYkb9q8ouwY81rKN191KPTUPn+4zjw0laYZk/G/EeXIzlnVPgCJRqguIgQERERUYhY2uw4dLzM677y6sb+DYaIiIiIhjVZknxWa/pLZzQEKRoiIhoofFbyu9iun0KneFcuNix6DGV78wDAo9NQx/uyvXnYsOgxFO/K7fcYiQY6JvmJiIiIQqTFYg93CEREREQ0zCmKgqbKOhTsyoU+OQ6CKPTqPIJKRPLkrCBHR0RE4cYkP/W3mvwibF6xFi67s8dlhBRJhmR3YvOKtajJL+qfAIkGCSb5iYiIiELE2uYIdwhERERENIxZG1tQ/G0+KnIL4LQ5YJo3FYqs9OpciiRj2rLLgxwhERGFm692/bLL1c+R0HCxY81rkJwuQPHvnkRRFEhOF3aueT3EkRENLkzyExEREYWIpa37Sn6Hkx+YiYiIiCj4nG12lOeeQfG3+WhranVvj0qJR8L4zICr+QWViMw5UxCXlRbsUImIKMwEUYSo8kwVSa7uK6yJeqOhoAKlu4/2WMF/LkWSUbI7F+bCyhBFRjT4MMlPREREFCKWHir5e9pPRERERBQI2SWh9nQ5Cnbnormy3uuYnCXzIahUgOBfol8QBKg0aly68vZghkpERAOIqPKs5me7fgqFIxs/g+BlUok/BJWIwxu2BzkiosGLSX4iIiKiEOkuiR9n1CM2OrIfoyEiIiKioUpRFDRW1OHMrlzUFZRD7qY6zpiRjDkP3Qq1VtPjQ3ZBJUKl02DhuoeQnDMqyFETEdFAIWq8JfnZfZCCrzq3IOAq/g6KJKPmaGGQIyIavNThDoCIiIhoqLJ2065/5V1XQvCzeoqIiIiIyBeruQU1J0rQ1mzpcawmUofkcSZEXz4Toy6ehJ1rXkfJ7lwIKrHLA/eO96ZZk3DpytuZ4CciGuJUXir5JVbyUwjYW6x9O96P+x2i4YJJfiIiIhqSyqsb8eSL23D4RDlSEqLx0J0/xJzzR/fb9c+U1OIXv9/odd8zD9yEZddc0G+xEBEREdHQ42izo/ZUKZqrGnocK6pVSMxKR1xminvd5eScUVjyxuMwF1bi8IbtqDlaCHuzBTqjAcmTszBt2eWIy0oL9ZdBREQDgKj2Vsnfu2prIl9kSfL6sxYIndEQpGiIBj8m+YmIiGjIabHYcN3df0VNQwuA9oT/z373Bna8+v+QkmAM+fUbW9pww71/87l/9vTskMdAREREREOT5JLQUFiJhpKqbtvyA4AgCIhJT0TSmAyodRqvY+Ky0jB/5fJQhEpERIOEqFahtboBpV8dQVNxFVw2BzT6CJTMm4qpSxcgPjs93CHSIKYoClqqG1Bzsgz65DgIogBFVgI+j6ASkTw5KwQREg1OTPITERHRkLNt51F3gr9Di8WO1z7Yi9/+9PKQX/+Lb06gsbnN5/6Y6MiQx0BEREREQ4siK2iqqEPtmTK47M4ex+vjjUgZn4mIaH0/REdERINVTX4Rdqx6GdVHzngkXxtOl+HAS1thmj0Z8x9dzuVbKGC2ZguqT5TAam5/TmeaNxWFn+3v1bkUSca0ZaF/rkc0WIjhDoCIiIgo2I6drvS6fee+U/1y/ZJK3y1T42P0iGWSn4iIiIgCYGloRtHeY6jMK+wxwa/VRyDjvLHInDGeCX4iIupW8a5cbFj0GGqOFgCAR3W18n3HmLK9ediw6DEU78rt9xhpcHI5nKjMK0LR3jx3gh8AolLikTA+E4IoBHQ+QRSROWcKlxIi6oRJfiIiIhpyXtm0x+v2E0XVkOXQrylX32jxuW+0KSnk1yciIiKiocFhtaHs0CmU7D8OW4u127EqtQrJ40zIungyopPjIAiBPTwnIqLhpSa/CJtXrIXL7uyxdboiyZDsTmxesRY1+UX9EyANSooso6G4CgW7ctFYVgNF8fzZylkyH4JKBfh7ryIIENUqXLry9iBHSzS4sV0/ERERDXqLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 2500x400 with 3 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAB/kAAAHeCAYAAABwqrgpAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdd3xUVfo/8M+90zPJpFfSgUCkiyACrqjYXYW1Yi9f/bn2XbEs6oq7uoptdde26roqrtgQGxZAF1SadEIv6b1nksn0e39/DAkJM5PMTGaSSfi8Xy9ezNx7zpknk0Du3Oec5wiyLMsgIiIiIiIiIiIiIiIiIiKisCcOdABERERERERERERERERERETkGyb5iYiIiIiIiIiIiIiIiIiIBgkm+YmIiIiIiIiIiIiIiIiIiAYJJvmJiIiIiIiIiIiIiIiIiIgGCSb5iYiIiIiIiIiIiIiIiIiIBgkm+YmIiIiIiIiIiIiIiIiIiAYJJvmJiIiIiIiIiIiIiIiIiIgGCSb5iYiIiIiIiIiIiIiIiIiIBgkm+YmIiIiIiIiIiIiIiIiIiAYJ5UAHQEREREREREREREQ00A4fPowlS5bgl19+QU1NDQAgIyMDp59+Oq6//nrExcUNcIRERERELoIsy/JAB0FERERERERERERENFDeeecdPPfcc7Db7R7Px8fH49VXX8XEiRP7NzAiIiIiD5jkJyIiIiIiIiIiIqLj1uLFi/HEE08AAHQ6HS699FKMGzcOFosFy5cvx8aNGwEA0dHRWL58ORITEwcyXCIiIiIm+YmIiIiIiIiIiIjo+FReXo4LLrgAFosFcXFxePfdd5GXl9etzRNPPIHFixcDAK655ho8+uijAxEqERERUSdxoAMgIiIiIiIiIiIiIhoIr7zyCiwWCwDgxRdfdEvwA8ADDzyAuLg4AMC3337br/EREREReaIc6ACIiIiIiIiIiIiIiPqbzWbDihUrAABnnHEGTj75ZI/t1Go17rzzThQXFyM2NhY2mw1qtbo/QyUiIiLqhkl+IiIiIiIiIiIiIjrurF+/Hm1tbQCAuXPn9tj26quv7o+QiIiIiHzCJD8RERERERERERERHXf27dvX+XjChAmdjxsbG1FYWAir1YqsrCykp6cPRHhEREREXjHJT0RERERERERERETHnYMHDwJwleNPTk5GaWkpnn76aaxZswYOh6Oz3bhx47BgwQKceOKJAxUqERERUTfiQAdARERERERERERERNTfampqAADR0dHYtGkTLr74Yvzwww/dEvwAUFBQgGuvvRbLly8fiDCJiIiI3DDJPwCuueYaXHPNNQMdBhERERFRv+O1MBERERGFC5PJBAAwm82488470d7ejksvvRRff/01CgoKsGrVKtxyyy0QRREOhwMPPfQQ9u7dG/Dr8VqYiIiIgoXl+gdAVVXVQIdARERERDQgeC1MREREROGiI8nf1tYGALj77rtxxx13dJ7PyMjA/PnzkZ6ejsceeww2mw3PPvss3n777YBej9fCx4+P5i1E+cY9AffPmDYGl3/wWBAjIiKioYYr+YmIiIiIiIiIiIjouJaXl4fbb7/d47krr7wSEyZMAACsXbuWyXrqVfK4XAiKwNIvgkJE0ticIEdERERDDZP8RERERERERERERHTc0el0nY8vuOACCILgte25557b+Xjr1q0hjYsGv/FXzobslALqKzslTJh3VpAjIiKioYZJfiIiIiIiIiIiIiI67kRGRnY+zs3N7bFtTs7RldU1NTUhi4mGhrjcNGRMH+v3an5BISJzxjjE5qSGKDIiIhoqmOQnIiIiIiIiIiIiouNOenq6z23VanXnY0kKbIU2HV9mPXwdFCpljxUiuhEEiEolTltwbWgDIyKiIYFJfiIiIiIiIiIiIiI67uTl5XU+rqio6LFtfX195+Pk5OSQxURDR1J+Nua8+SAUGlWvK/oFUYCoVGDaHy5DQl5GP0VIRESDmXKgAyAiIiIiIqLjQ2NhJXZ+uAo1BYWwtrZDExWB5HG5GH/lbMTlpg10eERERHScmTJlSufjn376CTfeeKPXttu3b+983HVyAFFPsmaMw7ylT2DNk4tRuq4AgkKE7DxaCUIQBciSjLi8DORfOguG9CTUH65AUl7mAEZNRESDAZP8REREREREFFK1e4ux+sn3ULZul9uNzYrN+7Dlra+RMX0sZj18HZLyswcuUCIiIjqujBs3DllZWSgpKcH69euxc+dOjB8/3q1dU1MTli9fDgAYPnw4Ro0a1d+h0iCWlJ+Ny95/FE1FVdixZCVqdxXBajRBVCuhS4hB5qkToE+O7WzfWFKDqKQ46GIiBzBqIiIKd0zyExERERERUciUrC3A57csgtPuAIBuCf6uz8s37sGSSx7BnDcfRNaMcf0eJxERER2fbr/9djz44IOQZRnz58/HO++8g7S0oxWGbDYbHnjgAbS0tAAAbrjhhgGKlAa72JxUzFpwXedzWZJRsmkvzC1t3drJsoyq3UXInjYGYi9l/omI6PjF3xBEREREREQUErV7i/H5LYvgsNrdkvvHkp0SnFY7Pr9lEWr3FvdPgERERHTcmzNnDs4++2wAQElJCS666CIsWrQIX331Fd555x1cfPHF+OmnnwAAU6dOxWWXXTaQ4dIQIogCUsfkQBDd0zRWkxn1hRUDEBUREQ0WTPITERERERFRSKx+8j3XCn5Z9qm9LMtw2h1Y8+TiEEdGREREdNQLL7yAOXPmAABaW1vx9ttvY/78+XjqqadQWFgIAJg5cyZee+01CIIwgJHSUKOJ1CFheJrHc43F1TC3mPo5IiIiGiyY5CcKY+u2F+KtT9fi14JiyD7eGCUiIiIiCgeNhZUoW7er1xX8x5KdEkrXFaCpqCpEkRERERF1p1KpsGjRIrz77ru48MILkZqaCrVajaSkJMyYMQP/+Mc/8OabbyIyknukU/DFZ6VCZ9C7He8o2y/5eT1NRETHB+VAB0BEnv3tje/w6pKfOp8nJxiw8cP7oVQoBjAqIiIiIiLf7PxwFQSF6HeSHwAEhYgdS1Z227OUiIiIKNSmTZuGadOmDXQYdJwRRAEpY3JQvHEPZKn7tbO1rR0NRVVIHDFsgKIjIqJwxZX8RGGotKoRr3/0c7djNfVGZM9+FJLEmZtEREREFP5qCgoDSvADrtX8tbuKghwREREREVF40kZFICEn1eO5hqJKWIws209ERN0xyU8Uhn7adBCS5Lk8/zc/7e7naIiIiIiI/Gdtbe9bf97IJCIiIqLjSHxOKrRREW7HO8r2H7vKn4iIjm9M8hOFoTWbD3k99/IHa/oxEiIiIiKiwGg83KD0q7+HfUmJiIiIiIYqQRSROjYXgiC4nbO0usr2ExERdWCSnyjM2OwO/Lhxv9fzuw5W9mM0RERERESBSR6XC0ER2EdOQSEiaWxOkCMiIiIiIgpv2qgIxOekeTxXX1QFSx+rZRER0dDBJD9RmKlvaoPV5hjoMIiIiIiI+iTv/FMgOwMrKSo7JYy/cnaQIyIiIiIiCn8JuanQRHoo2y9JqN5dBNnLNq9ERHR8YZKfKMw0t5oHOgQiIiIioj4xt5jQ1mhE/KhMCKJ7udGeCKKA+NGZsLRbuO8oERERER13BFFE6pgcj2X7zUYTGkpYtp+IiJjkJwo7jc2mHs8LggBZ5mxNIiIiIgpP1jYzyrcdgNPhRP5lsyAoFICHG5QeCQIEhQL5l85CS2U9yrYdgNPOKldEREREdHzRResRl53i8Vz94UpY21i2n4joeKcc6ACIqLvn/rOqx/OyLMNitUOnVfdTRERERETUobGwEjs/XIWagkJYW9uhiYpA8rhcjL9yNuJyPe+d2Z9+Xv0D9n70IlIbixHhsKJdqUFVXDbyr7gXp846M+SvbzNbUbZ1Pxw2OwDAkJ6Ek+6Yi82vLIPsdPZYWlQQXQn+k+6YC0N6EgDA1GBE6eZ9SJ+UBxWvf4mIiIjoOJKQOwxtdc2wtnWv/CpLEqp2FyNrSr7fVbOIiGjoYJKfKMxs3l3aaxuT2cYkPxEREVE/qt1bjNVPvoeydbsgKMRue81XbN6HLW99jYzpYzHr4euQlJ/d7/F9//UyWN5fiKktJciGACWOJtMdxnIoF/2ML17PgvaahTjnwrkhicFusaFsy37YLbZuxxNGZ2H6g1fh4JdrUbPzsNv71/E8Li8D+ZfO6kzwd7C0tqPk1z1In5QHbZT73qREREREREORqBCRekIOSjbtdavsam5pQ2NpNeKzUwcoOiIiGmhM8hOFEYfT6VO79mNunBIRERFR6JSsLcDntyzqLBvfNUHd9Xn5xj1YcskjmPPmg8iaMa7f4lvy75dx8mcLoZJc15JdE/xdn5/YUgr7a7dgSU0F5t18Z1BjcNodKN92ALZ2i8fzaRPzMPWG89FcUoMdS1aidlcRrEYTNAY9ksbmYMK8s6A2RKB8xyFIDvdrYrvFhtJNezFswgjo46ODGjsRERERUbjSxUQiLisZDcXVbufqD1cgMjEGGr1uACIjIqKBxiQ/URipqW/1qR2T/ERERET9o3ZvMT6/ZREcVjsgey81D7iS/U7Jjs9vWYR5S5/olxX933+9DCd/thBqyQFFL22VkCFIDpz82UJ8nzwsaCv6JYcTZVsPwNLqeV/QiNgoDJswHIIoIjYnFbMWXOd1rKwp+SjbegAOq/v1rvPI66SMyUFMWkJQYiciIiIiCncJw4ehta4ZNlP3CbWSU0LV7iJkncSy/URExyNxoAMgoqMOl9X51M5kZpKfiIiIqD+sfvI91wr+XhL8HWRZhtPuwJonF4c4MhfL+64V/L0l+DsoAKgkJ8zvLwzK60tOCeU7DsLc0ubxvNagR/rEkRAVvkWojYpA9tR8aCI9l+WXZRlVuwpRf7jCrWQpEREREdFQJCoUSD0hB4Lgnsg3N7ehqaxmAKIiIqKBxiQ/URh578uNPrUzM8lPREREFHKLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 2500x400 with 3 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAB/kAAAHeCAYAAABwqrgpAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdeXwTdfoH8M9MrjZt0/u+S7kvEVFEFLxXcREXURB1dV3UVXd1Xa9FXXUFXdbV1d+qu4rrASq4HiDixaGAAnIJtIVSWnrfR5KmSZprZn5/hJammZxNmrR93q8XL5KZ70yeQptO5vk+z5cRBEEAIYQQQgghhBBCCCGEEEIIIYQQQsIeG+oACCGEEEIIIYQQQgghhBBCCCGEEOIdSvITQgghhBBCCCGEEEIIIYQQQgghQwQl+QkhhBBCCCGEEEIIIYQQQgghhJAhgpL8hBBCCCGEEEIIIYQQQgghhBBCyBBBSX5CCCGEEEIIIYQQQgghhBBCCCFkiKAkPyGEEEIIIYQQQgghhBBCCCGEEDJEUJKfEEIIIYQQQgghhBBCCCGEEEIIGSIoyU8IIYQQQgghhBBCCCGEEEIIIYQMEZTkJ4QQQgghhBBCCCGEEEIIIYQQQoYIaagDIIQQQgghhBBCCCGEEEJC7dSpU1i3bh1+/PFHtLS0AACys7Nx8cUX49e//jUSEhJCHCEhhBBCiB0jCIIQ6iAIIYQQQgghhBBCCCGEkFB599138Y9//ANWq1V0f2JiIl5//XWcddZZgxsYIYQQQogISvITQgghhBBCCCGEEEIIGbHWrl2LFStWAAAiIyNx/fXXY/LkyTCZTPjyyy+xb98+AEBsbCy+/PJLJCcnhzJcQgghhBBK8hNCCCGEEEIIIYQQQggZmerr6zFv3jyYTCYkJCTgvffew5gxYxzGrFixAmvXrgUA3HzzzXjyySdDESohhBBCSC821AEQQgghhBBCCCGEEEIIIaHw2muvwWQyAQBefvllpwQ/ADzyyCNISEgAAHz99deDGh8hhBBCiBhpqAMghBBCCCGEEEIIIYQQQgabxWLBli1bAACXXHIJzjvvPNFxcrkc9913H6qrqxEfHw+LxQK5XD6YoRJCCCGEOKAkPyGEEEIIIYQQQgghhJARZ+/evdDr9QCA6667zu3YpUuXDkZIhBBCCCFeoSQ/IYQQQgghhBBCCCGEkBHnxIkTvY+nTp3a+1itVqOyshJmsxm5ubnIysoKRXiEEEIIIS5Rkp8QQgghhBBCCCGEEELIiFNeXg7A3o4/NTUVtbW1+Nvf/oadO3fCZrP1jps8eTKWL1+Os88+O1ShEkIIIYQ4YEMdACGEEEIIIYQQQgghhBAy2FpaWgAAsbGxOHDgAK699lps377dIcEPAMXFxbjlllvw5ZdfhiJMQgghhBAnlOQPgZtvvhk333xzqMMghBBCCCFk0NG1MCGEEEIICRcGgwEA0N3djfvuuw9GoxHXX389Nm/ejOLiYmzbtg3Lli0Dy7Kw2Wx47LHHUFpa6vfr0bUwIYQQQgKF2vWHQFNTU6hDIIQQQgghJCToWpgQQgghhISLniS/Xq8HAPzhD3/Avffe27s/OzsbDz30ELKysvDUU0/BYrHghRdewNtvv+3X69G1MCGEEEIChSr5CSGEEEIIIYQQQgghhIxoY8aMwT333CO6b/HixZg6dSoAYPfu3ZSsJ4QQQkjIUZKfEEIIIYQQQgghhBBCyIgTGRnZ+3jevHlgGMbl2F/84he9j3/++eegxkUIIYQQ4gkl+QkhhBBCCCGEEEIIIYSMONHR0b2PCwoK3I7Nz8/vfdzS0hK0mAghhBBCvEFJfkIIIYQQQgghhBBCCCEjTlZWltdj5XJ572Oe54MRDiGEEEKI1yjJTwghhBBCCCGEEEIIIWTEGTNmTO/jhoYGt2Pb29t7H6empgYtJkIIIYQQb0hDHQAhhBBCCCGEEEIIISQ0bPUVMH79Lqzlh8EbdGCjVJCNngblVbdBmlUY6vAICaoZM2b0Pt61axduv/12l2OPHDnS+7jv5ABCCCGEkFCgJD8hhBBCCCGEEEIIISOMtbIYujcfh+XoLoCVADzXu89y7CcYPnsV8qkXQXXnSsgKJocwUkKCZ/LkycjNzUVNTQ327t2LoqIiTJkyxWmcRqPBl19+CQAYNWoUxo4dO9ihEkIIIYQ4oHb9hBBCCCGEEEIIIYSMIObDO9H+4BWwFO+2b+iT4O/73FK8G+0PXgHz4Z2DHCEhg+eee+4BAAiCgIceegiNjY0O+y0WCx555BF0dnYCAG677bbBDpEQQgghxAlV8hNCCCGEEEIIIYQQMkJYK4uhfmYxYDEBguB+MM8BFjPUzyxG0ktbqKKfDEsLFizA9u3bsWXLFtTU1GD+/PlYtGgRJkyYgI6ODnz00UeorKwEAJx77rlYtGhRiCMmhBBCCKEkPyGEEEIIIYQQQgghI4buzccBq8Vzgr+HwANWC3Srn0Di858HNzhCQuSll17CE088gY0bN6Krqwtvv/2205jZs2fjlVdeAcMwIYiQEEIIIcQRJfkJIYQMKSXljfjpaBVG5SRjzjmFYFlaeYYQQgghhBBCvGGrr4Dl6C7fD+Q5WI7shK3hFKSZowIfGCEhJpPJsGrVKlx33XX4+OOPcejQIXR0dCAuLg6jR4/GjTfeiMsvv5zuQRBCCCEkbFCSnxBCyJCx5vOfsPzlTb3PF1w6Ff96/AaaRU8IIYQQQgghXjB+/S7ASuxt+H3FSmD86h2olq0IeFyEhIuZM2di5syZoQ6DEEIIIcQjSvITQggZEixWG/62eovDto3bj2LZogswdWxWiKIihBBCCCGEkKHDWn7YvwQ/APAcrBVHAhoPIYSMJOrKRhSt34aW4kqYu4xQxCiROrkAUxZfhoSCjFCHRwghZIgZVkl+i8WCTz75BF9//TXKyspgNBoRGxuLyZMnY8GCBbjyyivdVnsKgoDNmzfj008/RWlpKYxGI5KTkzFjxgwsXboUU6ZMGcSvhhBCSF9FZQ3QGUxO29/46Ae8/pclIYiIEEIIIYQQQoYW3qAL6fGEEDIStZZWY8fKNajbUwJGwkLg+N59DQdP4NBbm5E9axLmPn4rUsbnhS5QQgghQ8qwSfK3tLTgrrvuQmlpqcP29vZ2fP/99/j+++8xZ84cvPzyy1AqlU7Hm0wm3H///dixY4fD9oaGBjQ0NOCLL77AAw88gDvvvDOYXwYhhBARZVUtWPflQdF9h0vrBzkaQgghhBBCCBma2ChVSI8nhJCRpmZ3MTYuWwXOagMAhwR/3+f1+45j3cInsGD1o8i9YPKgx0kIIWToGRZJfqvV6pDgz83NxcKFC5Geno6qqiqsX78earUaO3fuxJ/+9Cf8+9//djrH448/3pvgHzVqFG644QYkJSXh2LFjWL9+PYxGI1588UWkpqbi2muvHcwvjxBCRrS//3cL/u/9HS731zVrBi8YQgghhBBCCBnCZKOnwXLsJ/9a9rMSyArPCnhMhBAyXLWWVmPjslWwma2AILgdK3A8ON6KjctWYcmnK6iinxBCiEdsqAMIhA0bNvQm+C+++GJs2rQJd911F+bPn4/7778fX375JcaNGwcA+O677/Djjz86HL97925s3rwZADBz5kxs2LABt912G6655ho8+uij+OSTTxAXFwcAeO6556DX6wfviyOEkBGstkntNsHfgxL9hBBCCCGEEOJZ5KwL/EvwAwDPQXn17YENiBBChrEdK9fYK/g9JPh7CIIAzmrDzpVrgxwZIYSQ4WBYJPm3bNkCAGBZFn/9618RERHhsD8hIQGPP/640/geb7/9NgBAKpVixYoVUCgUDvtHjRqFJ598EgCg1Wrx8ccfB/xrIIQQ4mzN5/u8GvfZ1sNBjoQQQgghhBBChja+9SBY437I8rMBhvHtYFYC+VlzIM0cFZzgCCFkmFFXNqJuT4lTe35PBI5H7Z5iaKqaghQZIYSQ4WJYJPnr6+3rMSckJCAlJUV0zNSpU3sfNzQ09D7WarXYs2cPAODCCy9Edna26PFXX301EhMTAQDffPNNQOImhBDiXkl5o1fjquo7ghwJIYQQQgghhAxNgiCAb/wBQvVmQBAQ/YuLAIkE8DbPz7CATA7VshVBjZMQQoaTovXbwEj8S78wEhZH120NcESEEEKGm2GR5I+JiQEAdHR0wGAwiI7pm9hPSEjofXzw4EHwvH023cyZM12+BsuymDFjBgDg6NGj6OzsHHDchBBC3OO9bGfW0KoNbiCEEEIIIYQQMgQJggChbguE+u2922RpyYhbOh+QSD1X9LMSQK5AwlPrISuYHORoCSFk+GgprvS5ir+HwPFoLakKcESEEEKGm2GR5J8yZQoA+weXntb7/b311lu9j2fPnt37uLy8vPfxmDFj3L5OYWFh7+ucPHnS73gJIYR4h+e9S/JL2GHx64wQQgghhBBCAkYQeAhVn0No3uu0T16Qg/hlN0I+dqJ9AytxHHD6uXzKbCS9tAWKaXOCHS4hhAwr5i7jwI7XiRczEkIIIT2koQ4gEH7961/js88+g9FoxOuvvw6dTofFixcjIyMDtbW1+O9//4vPP/8cAHDuuefimmuu6T22b4V/Zmam29dJS0tzOK6nsp+QYDlYUoPi8kZMHZuFaeOzwPi6Zh4hQ5zgZSW/tqs7yJEQQgghhBBCyNAhcFYIlZ9A0JS5HCM/ZzEUv3wVXGMljF+9A2vFEfAGHdgoFWSFZ0F59e2QZo4axKgJIWT4UMQoB3a8KipAkRBCCBmuhkWSPycnB6tXr8aDDz6IlpYWrFmzBmvWrHEYI5PJsHjxYvzpT3+CRHJmdrJare59HB8f7/Z14uLieh9rtdqAxE6IKy+8vRWvrP2+9/n9t1yMh39zeQgjImTweZvkLylvhCAINBGGODFZrLBaOcRERYQ6FEIIIYQQQgaFYDNBqFgPQVctPoBhwOZeAyZlOgBLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 2500x400 with 3 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAACAoAAAHeCAYAAAAYKXt1AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdeVxU9foH8M+ZGYZ92FdBQMUdzVwi06upbWqpueSSpnX1dtszU7Ms66eVdeu2d29WN9u0rDSXMpfScsnUXFBRUWTfhGEYGJj9/P5ARmiGgYGBgeHzfr163TPnu8yDv/p5OOc5zyOIoiiCiIiIiIiIiIiIiIiIiIiIOgSJqwMgIiIiIiIiIiIiIiIiIiKi1sNEASIiIiIiIiIiIiIiIiIiog6EiQJEREREREREREREREREREQdCBMFiIiIiIiIiIiIiIiIiIiIOhAmChAREREREREREREREREREXUgTBQgIiIiIiIiIiIiIiIiIiLqQJgoQERERERERERERERERERE1IEwUYCIiIiIiIiIiIiIiIiIiKgDYaIAERERERERERERERERERFRByJzdQBERERERERERERERO7g4sWLWLduHfbt24fCwkIAQGxsLG688Ubcc889CA4OdnGERERERNUEURRFVwdBRERERERERERERNSeffLJJ/jXv/4Fg8FgczwkJATvvfcerrnmmtYNjIiIiMgGJgoQERERERERERERETXDZ599hpUrVwIAvL29MWXKFCQlJUGr1WLbtm04dOgQACAgIADbtm1DWFiYK8MlIiIiYqIAEREREREREREREVFT5eTkYNy4cdBqtQgODsbatWvRvXv3OnNWrlyJzz77DABw9913Y/ny5a4IlYiIiMhC4uoAiIiIiIiIiIiIiIjaq3fffRdarRYA8MYbb1glCQDA4sWLERwcDAD48ccfWzU+IiIiIltkrg6AiIiIiIiIiIiIiKg90uv12LFjBwBg1KhRuO6662zOk8vleOihh5CRkYGgoCDo9XrI5fLWDJWIiIioDiYKEBERERERERERERE1wcGDB1FRUQEAmDRpkt25s2bNao2QiIiIiBqFiQJERERERERERERERE1w9uxZy3H//v0tx0qlEunp6dDpdIiLi0NMTIwrwiMiIiKqFxMFiIiIiIiIiIiIiIiaIC0tDUB1a4GIiAhkZWXh5Zdfxt69e2E0Gi3zkpKSsGzZMlx77bWuCpWIiIioDomrAyAiIiIiIiIiIiIiao8KCwsBAAEBATh8+DAmTJiA3bt310kSAICUlBTMnj0b27Ztc0WYRERERFaYKNBO3X333bj77rtdHQYRERERUavjtTARERERtRUajQYAUFVVhYceegiVlZWYMmUKtm7dipSUFOzatQvz58+HRCKB0WjE0qVLkZqa2uTv47UwEREROQtbD7RT+fn5rg6BiIiIiMgleC1MRERERG1FTaJARUUFAOCRRx7Bgw8+aBmPjY3FokWLEBMTg+eeew56vR6vvvoqPv744yZ9H6+FqaPSVVQi/cApq/Ne/j6IT+4DQRBcEBURUfvGigJERERERERERERERM3UvXt3PPDAAzbHpk+fjv79+wMA9u/fzwf+RA7y9POBd6Cf1XlteSW0ao0LIiIiav+YKEBERERERERERERE1ATe3t6W43Hjxtl9q/nWW2+1HP/5558tGheROwqMCbd5XpVzuZUjISJyD0wUICIiIiIiIiIiIiJqAj+/q284d+nSxe7chIQEy3FhYWGLxUTkrhQRQZB6WHfUVheUwGQ0uSAiIqL2jYkCRERERERERERERERNEBMT0+i5crnccmw2m1siHCK3JpFKoYgMsTpvNpmhzi9xQURERO0bEwWIiIiIiIiIiIiIiJqge/fuluPc3Fy7c4uLiy3HERERLRYTkTsLjAmzeV6VUwRRFFs5GiKi9s26RgsRERERERERERFRIynT83By/S4UpqRDV14JT38fRCR1Qb/pYxDcJdrV4RG1qMGDB1uOf/31V8ybN6/eucePH7cc104wIKLG8/L3gXeAH6rKKuqc15ZXQquuhHeAr4siIyJqf5goQERERERERERERA4rSs3AnlWfIvvAKQhSCUTT1VLquUfO4uiHWxE7tC9GPj0H4b3iXRcoUQtKSkpCXFwcMjMzcfDgQZw8eRL9+vWzmldaWopt27YBALp27YoePXq0dqhEbiMwJswqUQAAVLmXmShAROQAth4gIiIiIiIiIiIih2TuT8G6yc8g59AZAKiTJFD7c86hM1g3+Rlk7k9p9RiJWssDDzwAABBFEYsWLUJeXl6dcb1ej8WLF6OsrAwAMHfu3NYOkcitKCKCIZVJrc6rC0pgMppcEBERUfvERAEiIiIiIiIiIiJqtKLUDGyavxpGncEqQeCvRJMZJp0Bm+avRlFqRusESNTKJk6ciJtvvhkAkJmZiTvuuAOrV6/Gli1b8Mknn2DChAn49ddfAQBDhgzB1KlTXRkuUbsnkUmhiAqxOm82mlBeUOKCiIiI2icmChAREREREREREVGj7Vn1KUwGIyCKjZoviiJMBiP2rvqshSMjcp3XX38dEydOBACUl5fj448/xqJFi/DSSy8hPT0dADBs2DC8//77EATBhZESuYfAmHCb51U5l1s5EiKi9kvm6gCIiIiIOipRFHHgWDpS0wswoHcsBvbu7OqQiIiIiIjsUqbnIfvAKYfXiSYzsg6koPRSPoISologMiLX8vDwwOrVqzFp0iRs2LABR48eRUlJCQIDA5GYmIi77roLN910EyQSvrtH5Axe/j7wDvBDVVlFnfNVag2qyjTwDvB1UWRERO0HEwWIiIiIXOSZt7Zg7abfLZ+X/v1mPDRrZJ05Wr0Bl3JK0K1zGDxs9N8jIiIiImpNJ9fvgiCVNNhywBZBKsGJdTsxctmcFoiMqG1ITk5GcnKyq8Mg6hACO4VZJQoAQFnuZSYKEBE1AtMXiYiIiFwgI7ekTpIAALz84Q6oK7SWz1v3pKDvHStx031vof/EVdj/58XWDpOIiIiIqI7ClPQmJQkA1VUFik5dcnJERETUUSkigyG18VJFWUEJzEaTCyIiImpfWFGAiIiIqJVk5inxycaDMJnNKChW25zT+/YXkLLpGQgSAQ+v+hqGK7/YqjVaLFjxJU5sXAaZlJUFiIiIiMg1dOWVzVuv1jgpEiIi6ugkMikUkSEozSmqc95sNEFdoERgTJiLIiMiah+YKEBERETUCs5dKsT4B95DldbQ4NykiSvRJTbUkiRQo6y8CvuOXsTIId1bKkwiIiIiIrs8/X2at17BUtBEROQ8gTFhVokCAKDKvcxEASKiBrD1ABEREVEr+PDb/Y1KEqiRnl1s8/y5jEJnhURERERE5LCIpC4QpE27pShIJQjvm+DkiIiIqCPzUvjC20YSWlVZBbSsYkNEZBcTBYiIiIhawbptR5yyj0QQnLIPEREREVFT9Js+BqLJ3KS1osmM/jNucnJERETU0dVXOUCVe7mVIyEial+YKEBERETUwkRRdNpeEgkv34iIiIjIdYK7RCN2aF8IEscSWAWpBJ1vSEJQQlQLRUZERB2Vf2QIJDKp1Xl1fgnMf2nrSEREV/FOMxEREVELK1KWO20vqYM3ZImIiIiInO2GhdMhSKVAI6tdCYIAqYcMI5bNbuHIiIioI5LKpFBEhlidNxlNUBcqXRAREVH7wEQBIiIiohZWXFrhtL1MZudVJyAiIiIicpQoijCJZgx6cBIkMmmDlQUEqQRSTw9MXLME4b3iWydIIiLqcAI71dN+IIftB4iI6sNEASIiIqIWVlpW6bS9qnQGp+1FREREROQodYESmhI1QnvGYeiSmQjuHgsAVgkDgrT6tmNsch/M+HYl4m5IavVYiYio4/AO8IWXwtfqfFVZBbTlzrsvQ0TkTmSuDoCIiIjI3anKq5y2V5VW77S9iIiIiIgcYTIYUXQuy/JZEROO6x6bBk1hKYpOpKEirwQ6tQaeCl+E901A/xk3ISghyoURExFRRxLYKQwFao3VeVXOZUT2inNBREREbZvbJwo8//zz+PLLL/HQQw/h4YcfbnD+3r17sW7dOpw8eRJqtRohISHo168fZsyYgaFDhzbqO52xBxEREbmPUrXzMte1rChARERERC5yOS0HRr319WhA5wgMmDEGUpnUBVERERFVU0SFoOh8Fswmc53z6vxihHePgUTKv6eIiGpz60SBgwcPYv369Y2aazab8eyzz2LDhg11zhcUFKCgoAA7duzA3XffjWeeeQaCYLv3mjP2ICIiIvfTmESBmIhA5BSqGpxXyYoCREREROQClapylOYU2RyL6NmZSQJERORyUpkUisgQqHIv1zlvMpqgLixFYHSoiyIjImqb3DZR4NSpU3jooYdgNpsbngzgzTfftDzgj4qKwsyZM9GpUyekp6fjyy+/hFKpxOeff46QkBA88MADLbYHERERuR+VuuHWA107hzUqUaBKy4oCRERERNS6RLMZhamZNsf8wgLhHx7UyhERERHZFhgTZpUoAACqnCImChAR/YVbJgrs3bsXixYtQkVFRaPmX7x4EWvWrAEAJCYm4osvvkBAQIBl/K677sLdd9+NzMxMvPfee5gwYQI6derk9D2IiIjIPTWmokB0eECDcwCgiq0HiIiIiKiVKbMKoS23vqaVSCWI6BnHyplERNRmeCl84eXvY/X3VpWqArqKSnj6+bgoMiKitkfi6gCcSa/X46233sL9998PtVrd6HVr166FyWQCAKxYsaLOA34ACA8Px+rVqwEABoMB//vf/1pkDyIiInJPjUkUCA7wRY+EiAbnsaIAEREREbUmfZUOxRdzbY6FdukEubdnK0dERERUP0EQEBgTZnNMlWNdaYCIqCNzm0SBAwcO4LbbbsO7774Ls9kMHx8fzJs3r8F1ZrMZP/30EwCge/fuGDRokM15AwYMQJ8+fQAAP/30E0RLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 2500x400 with 3 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAACAoAAAHeCAYAAAAYKXt1AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdeXhU5fk38O+ZNclMJvtKAoQd2URAo4JQoVpRWQooiyLWQutuKaKi/qQWVGq12mrbV6wFUUFRQRYFd5C1IPtOCIHs+2Qyk8x+3j+GDAkzk8wksyX5fq6Li5nnPOeZOyHAzHPuc9+CKIoiiIiIiIiIiIiIiIiIiIiIqFOQhDoAIiIiIiIiIiIiIiIiIiIiCh4mChAREREREREREREREREREXUiTBQgIiIiIiIiIiIiIiIiIiLqRJgoQERERERERERERERERERE1IkwUYCIiIiIiIiIiIiIiIiIiKgTYaIAERERERERERERERERERFRJ8JEASIiIiIiIiIiIiIiIiIiok6EiQJERERERERERERERERERESdCBMFiIiIiIiIiIiIiIiIiIiIOhFZqAMgIiIiIiIiIiIiIuoIzp07h9WrV2PHjh0oLS0FAGRmZuIXv/gF7rvvPsTHx4c4QiIiIiIHQRRFMdRBEBERERERERERERG1ZytWrMBf//pXWCwWt8cTEhLwz3/+E1dffXVwAyMiIiJyg4kCRERERERERERERERtsGrVKixZsgQAEBkZialTp2LQoEEwGo3YvHkz9u7dCwCIiYnB5s2bkZSUFMpwiYiIiJgoQERERERERERERETUWgUFBbj99tthNBoRHx+PlStXok+fPk3mLFmyBKtWrQIA3HPPPXj++edDESoRERGRkyTUARARERERERERERERtVdvv/02jEYjAOCNN95wSRIAgIULFyI+Ph4A8NVXXwU1PiIiIiJ3ZKEOgIiIiIiIiIiIiIioPTKbzfj6668BADfffDOuu+46t/MUCgUeeeQR5OXlIS4uDmazGQqFIpihEhERETXBRAEiIiIiIiIiIiIiolbYvXs39Ho9AGDy5MnNzp01a1YwQiIiIiLyChMFiIiIiIiIiIiIiIha4dSpU87HQ4YMcT6uqqpCbm4uTCYTunXrhoyMjFCER0REROQREwWIiIiIiIiIiIiIiFrh7NmzABytBVJSUnDx4kW88sor2LZtG6xWq3PeoEGDsGjRIlxzzTWhCpWIiIioCUmoAyAiIiIiIiIiIiIiao9KS0sBADExMdi3bx8mTpyI7777rkmSAAAcPXoU9957LzZv3hyKMImIiIhcMFGgnbrnnntwzz33hDoMIiIiIqKg43thIiIiIgoXBoMBAFBfX49HHnkEdXV1mDp1KjZt2oSjR4/i22+/xdy5cyGRSGC1WvH000/j5MmTrX49vhcmIiIif2HrgXaquLg41CEQEREREYUE3wsTERERUbhoSBTQ6/UAgMceewwPP/yw83hmZiYWLFiAjIwMvPDCCzCbzXj11Vfx3nvvter1+F6YwlHZyTxsW7oKF3cdhSCVQLTZnccanne9cRBGL7oXyf27hy7QACjPKUBFbpHLuEQqQdb1A6GIighBVERE3mGiABERERERERERERFRG/Xp0wcPPfSQ22PTp0/H559/jsOHD2Pnzp0oLi5GWlpakCMkCozk/t0x7YPnUX2+GIdXf4OyY+dh0hmg1KiQPDALQ2b8EnFZHfPnPSErHbrSKpgNxibjdpsdJScvIPOaPhAEIUTRERE1j4kCREREREREREREREStEBkZ6Xx8++23N3tB8Fe/+hUOHz4MADhw4ABuv/32gMdHFExxWWkYs2h2qMMIKolUgrSrsnBhn2tLEUNlDXTFlYhJTwxBZERELZOEOgAiIiIiIiIiIiIiovZIrVY7H/fo0aPZuVlZWc7HpaWlAYuJiIIrKi4acRnJbo+VncmH1WwJckRERN5hogARERERERERERERUStkZGR4PVehUDgf2+32ZmYSUXuT1DsDMqXCZdxqtqDs9MUQRERE1DImChARERERERERERERtUKfPn2cjwsLC5udW1FR4XyckpISsJiIKPikchlS+3V1e6ymuBL6cm1wAyIi8oIs1AEQERERERERERFR+1WVW4Qja75F6dFcmGrroIyOQsqgHhg8fRzie6SHOjyigBoxYoTz8fbt23H//fd7nHvo0CHn48YJBkTUMUSnxCM6OQ61ZdUux0pOXkCPuGhIZNIQREZE5B4TBYiIiIiIiIiIiMhnZSfz8OPS95G/6xgEqQSi7XIp9cL9p/Dzu5uQecNAjHl2NpL7dw9doEQBNGjQIHTr1g0XLlzA7t27ceTIEQwePNhlXnV1NTZv3gwA6NmzJ/r27RvsUIkoCFL6dUNddS1sFmuTcYvRhPJzhUjp677qABFRKLD1ABEREREREREREfnkws6jWD3lORTsPQEATZIEGj8v2HsCq6c8hws7jwY9RqJgeeihhwAAoihiwYIFKCoqanLcbDZj4cKFqKmpAQDMmTMn2CESUZDIIxRI7p3h9lj1xVLUa/VBjoiIyDMmChAREREREREREZHXyk7mYf3cZbCaLC4JAlcSbXbYTBasn7sMZSfzghMgUZBNmjQJt9xyCwDgwoULmDBhApYtW4aNGzdixYoVmDhxIrZv3w4AuPbaazFt2rRQhktEARbTJQlRcdEu46IoouRkHkR78/93EhEFCxMFiIiIiIiIiIiIyGs/Ln3fUVJZFL2aL4oibBYrti1dFeDIiELn9ddfx6RJkwAAtbW1eO+997BgwQK8/PLLyM3NBQCMHDkS//rXvyAIQggjJaJAEwQBqVd1h0TqegnOWFuHyrySEERFRORKFuoAiIiIiDori9WG7fvOIq+oCjcM7YH+PVJDHRIRERERUbOqcouQv+uYz+eJNjsu7jqK6vPFiMtKC0BkRKEll8uxbNkyTJ48GWvXrsXPP/+MyspKxMbGonfv3rj77rvxy1/+EhIJ790j6gyUqkgk9EhH+dkCl2MVuUWITomDUhUZgsiIiC5jogARERFRCFhtNvzuhY/w9a6TAACJRMDfnp6KKb8cGuLIiIiIiIg8O7LmWwhSSYstB9wRpBIcXv0NxiyaHYDIiMJDdnY2srOzQx0GEYWBhG6pqC2pgrG2rsm4aLej5EQeug7vxwojRBRSTF8kIiIiCoG9h/OcSQIAYLeLeOmdrbCzTx0RERERhbHSo7mtShIAHFUFyo6d93NERERE4UmQSJB6VXe3yQB11bWoKSwPQVRERJcxUYCIiIgoBP7fJztcxkordMi5yA+JRERERBS+TFfcFenz+TqDnyIhIiIKf5ExasR1TXF7rOxMPixGc5AjIiK6jIkCRERERCFw+LRrjzoAqK0zBTkSIiIiIiLvKaOj2na+RuWnSIiIiNqHpJ5dII9UuozbrDaUnroQgoiIiBxkoQ6AiIiIqDMQRREfbvwfVn/1M9SRClRq3d9JVas3BjkyIiIiIiLvpQzqgcL9p1rVfkCQSpA8MCsAUREREYUviUyK1P7dkX/gtMux2rJq1JZWITolPgSREVFnx4oCREREREHw8Vc/4+m/fYHDpwqw82Cux3n3PLUC1bq2lXMlIiIiIgqUgVN/0aokAQAQbXYMmfFLP0dEREQU/tSJMYhJS3B7rOTURdgs1iBHRETERAEiIiKigLPZ7Hj1v996PX/15v0BjIaIiIiIqHVsFiv0VTok9O0KQSL4dK4glaDrjYMQl5UWoOiIiIjCW3LfrpAp5C7jVpMZZWfdt6gkIgokJgoQERERBdip8yUordB5PX/7/rMBjIaIiIiIyHc2ixX5P59Gvc6A/tPGQJBKAcG7ZAFBECCVyzB60b0BjpKIiCh8yRRyJPfNdHtMW1CGuuraIEdERJ2dLNQBEBEREXV0xeXeJwkAwI4D5wIUCRERERGR7xonCQCAJiMZwx+ejP1vr4Nos0G0ix7PFaQSSOUyTFr+FJL7dw9SxEREROFJk5oAXXEl9BU1Tcb1pVX45rnlMFbWwKSvhzI6CimDemDw9HGI75EeomiJqKNjogARERFRgOn0Rp/PMZotiHBTjo6IiIiIKJiuTBJokNivG254aiZOfb4dFSfyIEglEG125/GG55nZAzB60b1MEiAiIoKjyk5K/+6o23UUdpsduoIynFz7IypPX4QgEZok3xXuP4Wf392EzBsGYsyzs/l/KRH5XYdPFNDpdLj99ttRVlaGyZMn45VXXvE4VxRFbNq0CZ999hlOnjyJuro6JCUlYcSIEZg1axYGDx7c4uv5Yw0iIiLqWGoNvicK9Lr1BXy9/FFc1Ys9XImIiIgoNGwWK/IPnHFJEmiQ0DsTM9b+GYaSKhxe/Q3Kjp2HSWeAUqNC8sAsDJnxS8Rl8f0sERFRY4pIJZJ6ZeD4Fz85q/MAcKnQ05CAV7D3BFZPeQ6Tlj+FbjcOCnq8RNRxdfhEgZdeegllZWUtzjMajXj88cfx448/NhkvLCxEYWEhNm7ciCeeeALz5s0L6BpERETU8bQmUQAAHn95Lb75z2N+joaIiIiIqGXOJIEavdvj8kglug7vB0WkEoqsNIxZNDvIERIREbVfVn09fv7nOtitVsBzBx8AjoQBm92C9XOXYcZnS1hZgIj8pkMnCvz4449Yt26dV3OfffZZ5wX+nj174q677kJiYiKOHz+ONWvWoK6uDq+99hpSUlIwceLEgK1BREREHY+ulYkCJ3NLUFiqRZeUWP8GRERERETUjBaTBCIuJwkQERGR7358eRXsVnuLSQINRFGEzWLFtqWrMO2D5wMbHBF1Gh02UUCn0+H55737x3Lnzp3YtGkTACA7OxvvvPMOlErHB5077rgDU6dOxcyZM6HVavHSSy9h7NixUKvVfl+DiIiIOia9wdTqc0sqdUwUICIiIqKg8S5JoC+TBIiIiFqpKrcI+buO+XyeaLPj4q6jqD5Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 2500x400 with 3 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAACAoAAAHeCAYAAAAYKXt1AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdeVxU9foH8M+ZgWEfVgERBHEDc8nrkpmmqWlpuaSWS5nmtdWya6ZmWeZ1yRarX5b3ZnldSk1zzzSX3DXT3FBRUWQXEGbYYdbz+4MYGWdYBoaZAT7v18tXh+/3e77zUKaHc57zPIIoiiKIiIiIiIiIiIiIiIiIiIioUZDYOwAiIiIiIiIiIiIiIiIiIiKyHSYKEBERERERERERERERERERNSJMFCAiIiIiIiIiIiIiIiIiImpEmChARERERERERERERERERETUiDBRgIiIiIiIiIiIiIiIiIiIqBFhogAREREREREREREREREREVEjwkQBIiIiIiIiIiIiIiIiIiKiRoSJAkRERERERERERERERERERI0IEwWIiIiIiIiIiIiIiIiIiIgaESd7B0BERERERERERERE1BDcvHkT69evx7Fjx5CRkQEACAsLwyOPPILnn38efn5+do6QiIiIqJQgiqJo7yCIiIiIiIiIiIiIiOqzVatW4dNPP4VGozE77+/vj2+++Qb333+/bQMjIiIiMoOJAkREREREREREREREtbB27VosWLAAAODm5oZRo0ahQ4cOKCkpwa5du3Dq1CkAgLe3N3bt2oUmTZrYM1wiIiIiJgoQEREREREREREREdVUSkoKhgwZgpKSEvj5+WH16tVo06aN0ZoFCxZg7dq1AIBnn30Wc+fOtUeoRERERAYSewdARERERERERERERFRfff311ygpKQEAfPHFFyZJAgAwc+ZM+Pn5AQB2795t0/iIiIiIzHGydwBERERERERERERERPWRWq3G3r17AQD9+vXDAw88YHadTCbD1KlTkZCQAF9fX6jVashkMluGSkRERGSEiQJERERERERERERERDVw8uRJFBQUAABGjBhR6drx48fbIiQiIiKiamGiABERERERERERERFRDVy9etVw3KlTJ8OxQqFAfHw8VCoVwsPDERoaao/wiIiIiCrERAEiIiIiIiIiIiIiohqIi4sDUNpaICgoCElJSfjoo49w+PBhaLVaw7oOHTpgzpw5+Mc//mGvUImIiIiMSOwdABERERERERERERFRfZSRkQEA8Pb2xunTpzFs2DAcOHDAKEkAAGJiYvDcc89h165d9giTiIiIyAQTBeqpZ599Fs8++6y9wyAiIiIisjleCxMRERGRoygsLAQAFBcXY+rUqSgqKsKoUaPwyy+/ICYmBvv378eUKVMgkUig1Woxe/ZsxMbG1vjzeC1MRERE1sLWA/XU7du37R0CEREREZFd8FqYiIiIiBxFWaJAQUEBAOCNN97Aa6+9ZpgPCwvDjBkzEBoaig8++ABqtRqffPIJVq5cWaPP47UwERERWQsrChARERERERERERER1VKbNm3w6quvmp0bM2YMOnXqBAA4fvw4H/gTERGR3TFRgIiIiIiIiIiIiIioBtzc3AzHQ4YMgSAIFa597LHHDMdnz56t07iIiIiIqsJEASIiIiIiIiIiIiKiGvD09DQcR0ZGVrq2RYsWhuOMjIw6i4mIiIioOpzsHQARERERUX2Wm5uLDRs24ODBg7h16xYKCwvh5eWFtm3b4rHHHsNTTz0FmUxm9tw///wTzz33XLU+p1evXvj+++/NzomiiF9++QWbN29GbGwsioqK0KRJE3Tr1g3jx49Hx44dq9zfUfYgIiIiIqpPQkNDcfr06WqtLf9zgV6vr6uQiIiIiKqFiQJERERERDV08uRJTJ8+HQqFwmhcoVDg5MmTOHnyJH744QcsX74cYWFhJudfu3at1jGUlJRg2rRpOHTokNF4amoqUlNTsXPnTrz55pt48cUXHX4PIiIiIqL6pk2bNobj1NTUStdmZWUZjoOCguosJiIiIqLqYKIAEREREVENXL16Fa+88gqKi4sBlL7x379/f/j4+CAtLQ3btm1DXFwc4uLiMHnyZPz888+Qy+VGe5QlCnh4eODjjz+u9PP8/f3Njr/77ruGh/MtW7bE008/jYCAAFy+fBkbNmxAUVERPvvsMwQFBWHYsGEOvQcRERHVT4r4NFzcsB8ZMfFQ5RfBxcsdQR0i0XHMAPhFhtg7PKI61a1bN8PxkSNHMGnSpArXnj9/3nBcPsGAiIiIyB4EURRFewdBluvfvz8A4MCBA3aOhIiIiKhxevbZZw0lRufNm4exY8cazWu1WsyePRs7d+4EAEyaNAmzZ882WjN69GhcvHgRnTt3xoYNGyyO4fjx43jhhRcAAD169MC3334LFxcXw/zNmzcxbtw45OTkwMfHBwcOHDDqoepIe1iC18JERESOITM2AYcWrkHyiUsQpBKIurul1Mu+DuvZHn3fnYDA6Aj7BUpUxwYOHIjExEQIgoCNGzeabbmlVCoxaNAg5ObmomXLlvj1119r9Fm8FiYiIiJrkdg7ACIiIiKi+ubmzZuGJIEBAwaYJAkAgJOTExYuXIjAwEAAwJYtW6DT6Qzzer0eN27cAAC0bt26RnGsXLnS8FkLFiwwejgPlL7ZP3fuXABATk4ONm3a5LB7EBERUf2SeDwG60e+h5RTVwDAKEmg/Ncpp65g/cj3kHg8xuYxEtnKq6++CgAQRREzZsxAWlqa0bxarcbMmTORm5sLAJg4caKtQyQiIiIywUQBIiIiIiILnTx50nBcWRl9FxcXPPLIIwCA3NxcJCQkGOaSkpJQVFQEoGZlR3NycnDixAkAQO/evREWFmZ23eDBgw1tC/bs2eOQexAREVH9khmbgG1TlkCr0pgkCNxL1OmhU2mwbcoSZMYm2CZAIhsbPnw4Bg4cCABITEzE0KFDsWTJEuzcuROrVq3CsGHDcOTIEQBA9+7dMXr0aHuGS0RERASAiQJERERERBaTSCRo3bo1PD09ERERUelab29vw3FeXp7h+Nq1a4bjmiQKnDlzBnp96Y35Hj16VBprWd/UCxcuGN5icqQ9iIiIqH45tHANdBotUM2OpqIoQqfR4vDCtXUcGZH9LF26FMOHDwcA5OfnY+XKlZgxYwYWL16M+Ph4AECvXr2wfPlyCIJgx0iJiIiISjnZOwAiIiKixkqj1eHImRtITMvGg/dHIjoy2N4hUTWNGzcO48aNq9basvYCAODj42M4vn79uuG4rPVAcnIyEhMTIZVKER4ejpCQkAr3jYuLMxxXlWjQqlUrAKU36a9fv254YO8oexAREVH9oYhPQ/KJSxafJ+r0SDoRA+Wt2/Bt0bQOIiOyL2dnZyxZsgQjRozApk2b8NdffyE7Oxs+Pj5o3bo1nnnmGTz66KOQSPjuHhERETkGJgoQERER2YFWp8NLH6zD3hOxAACpRIKls0di5KOd7RwZWVNGRgaOHj0KAPD19UV4eLhhrqyiQEBAAI4cOYJvv/0WN2/eNDr/vvvuw5tvvomHH37YZO/U1FTDcbNmzSqNIzj4bhJKamqq4QG9o+xBRERE9cfFDfshSCVVthwwR5BKcGH9PvSdM6EOIiNyDD169Ki00hYRERGRo2D6IhEREZEdnLqQYEgSAACdXo9pizYhO6fAjlGRtS1ZsgQajQYAMGTIEKO3h8oSBbKysjBr1iyTJAEAuHz5MqZMmYJly5aZzCkUCsOxr69vpXGUr2SQk5PjcHsQERFR/ZERE1+jJAGgtKpA5qVbVo6IiIiIiIhqgokCRERERHbw895zZsc7jViE1Iwc2wZDdWLDhg3YtWsXAMDd3R0vvviiYa6oqAjJycmGr8PDw/Hxxx/j2LFjuHjxInbs2IFnn33W0Lv0q6++wpYtW4z2LykpMRy7uLhUGotMJjN7nqPsQURERPWHKr+odufnFVopEiIiIiIiqg0mChARERHZwdkrSRXOrf/1jA0jobqwf/9+zJ8/3/D1vHnzEBQUZPg6KSnJ8FC9c+fO2LJlC4YNG4YmTZrAxcUFbdu2xdy5c7F48WLDOUuWLEFh4d0b61qt1nBc/gG8OeXny5/nKHsQERFR/eHi5V678+UeVoqEiIiIiIhqg4kCRERERHYQ4OtZ4dwXa363YSRkbXv37sWbb74JnU4HAJgwYQKGDRtmtCYqKgrnz5/HkSNHsGLFCnh6mv/9MGLECPTr1w9Aaan+ffv2GeZcXV0Nx2XtDSqiVqsNx+Uf1jvKHkRERFR/BHWIhCCt2S1FQSpBYPsWVo6IiIiIiIhqgokCRERERHbg5803qRqiTZs24V//+pfhgfmIESMwZ84cs2sFQUBQUBC8vLwq3bN8ksGZM3erTbi7332bT6VSVbpH+Qf05dsDOMoeREREVH90HDMAok5fo3NFnR6dxj5q5YiIiIiIiKgmmChAREREZAdSiWDvEMjKli1bhvfee89QUn/kyJFYtGgRBKF2/60jIyMNx5mZmYZjuVxuOM7Jyal0j/Lzfn5+DrcHERER1R9+kSEI69ne4qoCglSC5g91gG+LpnUUGRERERERWYKJAkRERER28MvhS/YOgaxEr9fj/fffx1dffWUYmzBhAhYuXAiJpPaX205OTmbHIyIiDMe3b9+udI/09HTDcUhIiMPtQURERPVL33cnQOrsBFQzIVIQBEidndBnznN1HBkREREREVUXEwWIiIiIbCw9K6/KNboalnMl29Lr9Zg5cyZ++uknw9i0adPw7rvvVlpJ4MiRI/j222+xePFiKJXKSj8jIyPDcNykSRPDccuWLQ3HcXFxle5RNi8IAlq3bu1wexAREVH9EhgdgX4L/gmJkxRCFZWyBKkEUhdnDF8xC4HREbYJkIiIiIiIqsREASIiIiIbO3M5sco1eQUlNoiEamvevHnYuXMnAEAikWDevHl49dVXqzxLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 2500x400 with 3 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAB/kAAAHeCAYAAABwqrgpAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdd3hUZdoG8PtMSzLpvZAQEmqo0hQpgr1TBBREUHcXdlexrgVR17Kii66u7qpb8FMUFRSVbgMVVgEx9AQCCYT0nkx6pp/vj5ghYc4kM8lkWu7fdXk5Oe97zjwJkJk5z/s8ryCKoggiIiIiIiIiIiIiIiIiIiLyeDJ3B0BERERERERERERERERERET2YZKfiIiIiIiIiIiIiIiIiIjISzDJT0RERERERERERERERERE5CWY5CciIiIiIiIiIiIiIiIiIvISTPITERERERERERERERERERF5CSb5iYiIiIiIiIiIiIiIiIiIvAST/ERERERERERERERERERERF6CSX4iIiIiIiIiIiIiIiIiIiIvwSQ/ERERERERERERERERERGRl1C4OwAiIiIiIiIiIiIiInc7e/Ys1q9fj59++gnl5eUAgKSkJFx++eW48847ERER4eYIiYiIiFoJoiiK7g6CiIiIiIiIiIiIiMhd1q5di7/97W8wGAyS45GRkXj77bdx0UUXuTYwIiIiIglM8hMRERERERERERFRn7Vu3Tq88MILAICAgADMmzcPo0aNglarxY4dO3DgwAEAQGhoKHbs2IHo6Gh3hktERETEJD8RERERERERERER9U1FRUW48cYbodVqERERgffffx9DhgzpMOeFF17AunXrAAB33HEHnn76aXeESkRERGQhc3cARERERERERERERETu8NZbb0Gr1QIAXn/9dasEPwA89thjiIiIAAB89dVXLo2PiIiISIrC3QEQEREREREREREREbmaXq/Ht99+CwC44oorcMkll0jOU6lUWL58OfLy8hAeHg69Xg+VSuXKUImIiIg6YJKfiIiIiIiIiIiIiPqc/fv3o7GxEQAwZ86cTucuWrTIFSERERER2YVJfiIiIiIiIiIiIiLqc06dOmV5PGbMGMvjmpoa5ObmQqfTITk5GYmJie4Ij4iIiMgmJvmJiIiIiIiIiIiIqM/JyckB0NqOPzY2FgUFBfjrX/+KPXv2wGg0WuaNGjUKK1euxLhx49wVKhEREVEHMncHQERERERERERERETkauXl5QCA0NBQpKenY9asWfjuu+86JPgBICMjA4sXL8aOHTvcESYRERGRFSb53eCOO+7AHXfc4e4wiIiIiIhcju+FiYiIiMhTNDU1AQBaWlqwfPlyNDc3Y968edi+fTsyMjKwa9cuLF26FDKZDEajEStWrEBWVla3n4/vhYmIiMhZ2K7fDUpLS90dAhERERGRW/C9MBERERF5irYkf2NjIwDg/vvvx7333msZT0pKwiOPPILExEQ888wz0Ov1eOWVV/Duu+926/n4Xpg8SWNlLQqPZFsd7z9+KAIjQ90QEREROYKV/ERERERERERERETUpw0ZMgT33HOP5NiCBQswZswYAMDevXuZrCefIMgEyeOiWXRxJERE1B1M8hMRERERERERERFRnxMQEGB5fOONN0IQpJOeAHDddddZHh8+fLhX4yJyBUEmnR4ym80ujoSIiLqDSX4iIiIiIiIiIiIi6nOCgoIsj1NTUzudm5KSYnlcXl7eazERuYqtJL/IJD8RkVdgkp+IiIiIiIiIiIiI+pzExES756pUKstjVjqTL5DJbbTrN7FdPxGRN2CSn4iIiIiIiIiIiIj6nCFDhlgeFxcXdzq3qqrK8jg2NrbXYiJyFVbyExF5N4W7AyAiIiIiIiIiIiL3qMktwfENu1CekQtdQzP8gtWIHZWK0QuuQkRqgrvDI+pVEydOtDz+3//+h7vvvtvm3KNHj1oet18cQOStbCX52amCiMg7MMlPRERERERERETUx1Rk5WH3qg9QuC8TglwG0XQ+qVN88BQOvbMdSZNHYsaTSxCTNsB9gRL1olGjRiE5ORn5+fnYv38/jh8/jtGjR1vN02g02LFjBwBg4MCBGDp0qKtDJXI6mdxGJb+JSX4iIm/Adv1ERERERERERER9SP7eDKyf+xSKDpwEYJ3Qafu66MBJrJ/7FPL3Zrg8RiJXueeeewAAoijikUceQUlJSYdxvV6Pxx57DHV1dQCAu+66y9UhEvUKQSZIHhfNoosjISKi7mAlPxERERERERERUR9RkZWHzUtXw6gzAGLniRzRZIbJbMDmpaux8PMXWNFPPmn27Nn47rvv8O233yI/Px8zZ87E/PnzMXz4cFRXV+OTTz5Bbm4uAODiiy/G/Pnz3RwxkXOwXT8RkXdjkp+IiIiIiIiIiKiP2L3qA5gMxi4T/G1EUYTJYMSeVesw/8Onezk6Ivd47bXX8NRTT2Hz5s1oaGjAu+++azVn6tSpeOONNyAI0tXPRN5GEAQIMhnEC5L6F35NRESeiUl+IiIiol5gMJrwv4NnkFdcjcljUzEsJRY5+RXw91MiKS6cN4aIiIiIyOVqcktQuC/T4fNEkxkF+zKgOVeK8JT4XoiMyL2USiVWr16NOXPmYOPGjTh06BCqq6sRFhaGwYMH47bbbsPVV18NmY3KZyJvJZMJuGDHFrbrJyLyEkzyExERETmZ0WTCH59fj69/PCk5fu2UNLz59AIE+CldHBkRERER9WXHN+yCIJdBvDCjYwdBLsOx9TsxY+WSXoiMyDNMmjQJkyZNcncYRC7T2rLf1OGYuRuvEURE5HpcekhERETkZL9k5NtM8APAN3uzsH5HugsjIiIiIiICyjNyu5XgB1qr+Ssyzzk5IiIicidBbp0iYrt+IiLvwCQ/ERERkZP955Mfu5zzzw93934gRERERETt6Bqae3Z+fZOTIiEiIk8gtQUFk/xERN6BSX4iIiIiJztxprTLOZWaRhdEQkRERER0nl+wumfnhwQ6KRIiIvIEgkywOmY2i26IhIiIHMUkPxEREZGTKRVyd4dARERERGQlYmA/yYSOPQS5DDEjU5wcERERuZMgVcnfzW1diIjItZjkJyIiInIyhcK+t1hGk6mXIyEiIiIiam29XJFdgLChSRC7WaEpmswYs/BqJ0dGRETuJMjZrp+IyFsxyU9ERETkZEq5fZX8zS2GXo6EiIiIiPo6g1aPgkOnUZ1XhqDYCEQO7e9wNb8gl6H/lFEIT4nvpSiJiMgdZBKvB91dDEZERK7FJD8RERH5pOYWPd5avwdP/WMr9qTnQBRd9yFVYWe7/pAg/16OhIiIiIj6sqbqOuT9fALNmgbLsbT5MyDI5YBgX6JfEATIlQpMX7m4t8IkIiI3kWrXb2a7fiIir6BwdwBEREREzqY3GDFz+b9xKrcMALB208/48x9vwLJbpzrtOTR1zVi37QDyiqtx2YTBmHXFaAi/3ihVSLS7u9Cdsyc5LRYiIiIiovZEs4iqcyWozi2xWuwakhiDCffOwcG3NkE0mTqt2BTkMsiVCsxe8zhi0gb0ctRERORqbNdPROS9mOQnIiIin/PDgWxLgr/Na+9/h8UzL0aAv6rH129u0WP+w+9YnuPTrw/jbGEl/nTXVQAAlbLrSv67mOQnIiIiol5g1BtQknEWTdX1NufEjEjBzP88gsPvbEfBvkwIchnEdpWbbV8nTRqB6SsXM8FPROSjZBKV/EzyExF5Byb5iYiIyOd8f+C01bHGZh32HsnFVZcO69Y1TSYzTueVw1+lRHZ+hdUignc/34cHFl8OhVwOsx1bA8RFhXQrDiIiIiIiW5o1DSg+fhZGnd7mHJXaH/3GDIJ/sBoDZ4yD5lwpjq3fiYrMc9DVN8EvJBAxI1MwZuHVCE+Jd2H0RETkaoLMeusWcycdXoiIyHMwyU9EREQ+J/NMqeTx03nl3Uryl1fX464nPkBGTonNOXWNWmSfq8DwQfFo0Rq6vGaQ2s/hOIiIiIiIpIiiiJr8MlTmFFm1528vOCYc8SNSIFeevyUYnhKPGSuXuCJMIiLyMIKNSn7RLEouACAiIs/R9YaxRERERF4mMTZM8rimrrlb13v9g+87TfC3Mf3a0q5Za7tyqo0g8MMyEREREfWcyWBE0dEcVGQX2kzwC4KA2KH90W/MoA4JfiIi6tuk2vUDgCiyZT8Rkadjkp+IiIh8jk5vlDz+709+7Nb11m39xa55bYn7phZdp/OunDS0W3EQEREREbXXUteIcz+fQGNlrc05Sn8/JE9MQ0RyHBeaEhFRB7aq9c0mJvmJiDwdl+4SERGRz6lraLE5VlnTgOiIYLuvZTSZ7J7btm9dQ1PnSf7ubBlARERERNRGFEXUFlagPLsQotl2IiYoKgzxI1OgUCldGB0REXkLQW6jkt9se+sXIiLyDKzkJyIiIp/TWZJ/z8EzDl2rorrB7rl6gxF6g9FmJ4E2SoXcoRiIiIiIiNqYjCaUZJxF2al8mwl+QRAQPTgRiWMHM8FPREQ22WzX38kCMiIi8gys5CciIiKf0qIzIK+k2uZ4Y5PWoevll9TYPVenN6KpWd/lPBX3QSUiIiKibtA2NKP4+BnoO3lPq/BTImHUQARGhLgwMiIi8kYCk/xERF6Ld5iJiIjIp5w8Uwq9wXaL/eAgf4eul56Zb/dcncGIhuauFxGoVKzkJyIiIiLH1BZXovxUfqf7JKsjQtBv1EAo/Fi9T0REXRPkguRxs4nt+omIPB2T/ERERORTSipqOx3X620vAJBy+ly53XN1eiMam3RdzlMp+BaMiIiIiOxjNplQlpWPupKqTudFpSYgKrUfBJl0woaIiOhCrOQnIvJevMNMREREPqWqtqnT8aJyjUPXq2u0v72/3mBEY4sdSX4lK/k9RV1dHTZs2IAffvgB586dQ1NTE4KDgzF06FBcd911uOWWW6BSqWyeL4oitm/fjs8//xxZWVlLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 2500x400 with 3 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAACAoAAAHeCAYAAAAYKXt1AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdd3hUZdoG8PtMTe+dBJLQlSJgQQRBZVUsFAEFKeK62BUrIMLKIojY1sK6fuIqigqKAoLYKIIKiICU0AMhvZEy6VPP+f4YMyTMTDKZmnL/rotrZ973Pe88swYyc85znkeQJEkCERERERERERERERERERERdQgyXwdARERERERERERERERERERE3sNEASIiIiIiIiIiIiIiIiIiog6EiQJEREREREREREREREREREQdCBMFiIiIiIiIiIiIiIiIiIiIOhAmChAREREREREREREREREREXUgTBQgIiIiIiIiIiIiIiIiIiLqQJgoQERERERERERERERERERE1IEwUYCIiIiIiIiIiIiIiIiIiKgDYaIAERERERERERERERERERFRB6LwdQBERERERERERERERO3B2bNnsXr1avz2228oKioCACQlJeG6667DPffcg4iICB9HSERERGQmSJIk+ToIIiIiIiIiIiIiIqK2bOXKlXjttddgMBhszkdGRuLdd9/FZZdd5t3AiIiIiGxgogARERERERERERERkQtWrVqFxYsXAwD8/f0xYcIE9O3bF1qtFps3b8bevXsBAKGhodi8eTOio6N9GS4REREREwWIiIiIiIiIiIiIiJyVm5uLW2+9FVqtFhEREfj444/Ro0ePRmsWL16MVatWAQCmTp2KBQsW+CJUIiIiIguZrwMgIiIiIiIiIiIiImqr/vOf/0Cr1QIA3nzzTaskAQCYPXs2IiIiAADff/+9V+MjIiIiskXh6wCIiIiIiIiIiIiIiNoivV6Pn376CQBw/fXX46qrrrK5TqVS4dFHH0VmZibCw8Oh1+uhUqm8GSoRERFRI0wUICIiIiIiIiIiIiJywp49e1BdXQ0AGDduXJNrp0yZ4o2QiIiIiBzCRAEiIiIiIiIiIiIiIiecPHnS8rh///6Wx2VlZcjIyIBOp0OXLl2QmJjoi/CIiIiI7GKiABERERERERERERGRE9LT0wGYWwvExsYiOzsbL7/8Mnbu3Amj0WhZ17dvX8ybNw8DBw70VahEREREjch8HQARERERERERERERUVtUVFQEAAgNDcW+ffswZswYbNu2rVGSAACkpaVh2rRp2Lx5sy/CJCIiIrLCRIE2aurUqZg6daqvwyAiIiIi8jp+FiYiIiKi1qKmpgYAUFdXh0cffRS1tbWYMGECvv32W6SlpWHr1q2YOXMmZDIZjEYj5s6dixMnTjj9evwsTERERO7C1gNtVEFBga9DICIiIiLyCX4WJiIiIqLWoj5RoLq6GgDw+OOP45FHHrHMJyUl4ZlnnkFiYiJeeOEF6PV6vPrqq/jwww+dej1+FiZPqimrRPb+k1bjCX1SEZoQ5YOIiIjIk1hRgIiIiIiIiIiIiIjIRT169MDDDz9sc27SpEno378/AGDXrl284E+tkirAz+a4vlbr5UiIiMgbmChAREREREREREREROQEf39/y+Nbb70VgiDYXXvzzTdbHv/5558ejYvIGQq1EjKF3GpcX8NEASKi9oiJAkRERERERERERERETggKCrI8Tk1NbXJtSkqK5XFRUZHHYiJyliAINqsKsKIAEVH7xEQBIiIiIiIiIiIiIiInJCYmOrxWpVJZHoui6IlwiFxmL1FAkiQfRENERJ7ERAEiIiIiIiIiIiIiIif06NHD8jgvL6/JtSUlJZbHsbGxHouJyBWqQOtEAdEkwqjV+yAaIiLyJIWvAyAiIiIiIiIiIqK2qywjH0fWbEVRWgZ0VbVQBwcgtm8q+k0aiYjUBF+HR+RRV1xxheXxL7/8gnvvvdfu2kOHDlkeN0wwIGpNbFUUAMxVBZT+ai9HQ0REnsREASIiIiIiIiIiImqx4hOZ2LHkE+TsPgpBLoNkulBKPW//SRz44FskDemDEc9PR0zvZN8FSuRBffv2RZcuXZCVlYU9e/bgyJEj6Nevn9W68vJybN68GQDQtWtX9OzZ09uhEjnEbqJAjRaBkaFejoaIiDyJrQeIiIiIiIiIiIioRbJ2pWH1+PnI3XscABolCTR8nrv3OFaPn4+sXWlej5HIWx5++GEAgCRJeOaZZ5Cfn99oXq/XY/bs2aioqAAAzJgxw9shEjnMVusBwFxRgIiI2hcmChAREREREREREZHDik9kYsPMZTDqDFYJAheTTCJMOgM2zFyG4hOZ3gmQyMvGjh2LG2+8EQCQlZWF0aNHY9myZdi0aRNWrlyJMWPG4JdffgEAXHnllZg4caIvwyVqklwhh0KtshpnogARUfvD1gNERERERERERETksB1LPoHJYAQkyaH1kiTBZDBi55JVmPjpAg9HR+Qbb7zxBubPn48NGzagqqoKH374odWaoUOH4q233oIgCD6IkMhxqkA/GHX6RmO6GiYKEBG1N0wUICIiImrF9h45h4ycUvRMicXAS5J8HQ4RERERdXBlGfnI2X20xcdJJhHZu9NQfq4A4SnxHoiMyLeUSiWWLVuGcePGYe3atThw4ABKS0sRFhaG7t2746677sLf/vY3yGQs8kutnyrAD7VllY3GjFo9RJMImZw/w0RE7QUTBYiIiIhaqUXvfof31/4GABAEAc/NvAkPT77Wx1ERERERUUd2ZM1WCHJZsy0HbBHkMhxevQUj5k33QGRErcPgwYMxePBgX4dB5BJVgNpqTJIkGOq0UAcF+CAiIiLyBKZ+EREREbVC5/JK8cHXuyzPJUnCvz/Zhto6fRNHERERERF5VlFahlNJAoC5qkDx0XNujoiIiNxNHehvc1zP9gNERO0KEwWIiIiIWqEVX/4GUWzc87VOa8CewzyxSkRERES+o6uqde34yho3RUJERJ6iCvCzOa6rZaIAEVF7wkQBIiIiolZoz+EMm+OL3t2Mf3+8DZl5pV6OiIiIiIgIUAe7VnJaHRLopkiIiMhTlP4qCIJgNc6KAkRE7QsTBYiIiIhaoaKSKpvjZ3NK8PrKbbhp5jv483iOl6MiIiIioo4utm8qBLlzpxQFuQwxfVLcHBEREbmbIJNBGaC2GtezogARUbvCRAEiIiKiVkinNzY5X1Onx4q1v3kpGiIiIiIis36TRkIyiU4dK5lE9J/8NzdHREREnqAO8LcaM9TqfBAJERF5ChMFiIiIiFohvdHU7JpNO9K8EAkRERER0QURqQlIGtKnxVUFBLkMna/pi/CUeA9FRkRE7qQKtK4oYNQbYNQbfBANERF5AhMFiIiIiFwkiiLe//I33PbQu5jwxAr8fvicy3tKkuTWdURERERE7jLi+emQKRSAjf7VtgiCALlSgeHzpnk4MiIichdVgJ/NcbYfICJqPxS+DoCIiIiorXt39S94+YOfLM8nPLECfbsn4FxeKTrHR+CDF6egc3yEw/uZWlDKVVNZh/DQgBbFS0RERETkipjeyRg6Zwp+XboKkskESbSfvCrIZZArFRi7Yg5ieid7L0giInKJ0m6igA4BYcFejoaIiDyBFQWIiIiIXPT55n1WY2np+aiu1eH42QIMufs1HD9T4PB+lTWOZ+cXllY6vJaIiIiIyB3qKmoQmBCJIXPuRkSPJACAIGtcXaC+NUHS4Esx+evF6HJNX6/HSUREzlMH+tscN7TgnAUREbVurChARERE5AKd3ojsgvJm193+8H9x9qdFDu1ZWe34l+6gAOuegUREREREnlR6Lh8AEJIYg6ueuBM1ReXI/vUwas9rIBmMUIcEIqZPCvpP/hvCU+J9HC0RETlDrlJArpDDZDQ1GtfV1PkoIiIicjcmChARERG54HxZlUPrdAYjtHoD/FTK5tfqDQ7t2Ss1Dp1iQh1aS0RERETkDtqqWlQVN06UDYwNR59JN6DrsP6QK3m6kYioPRAEAaoAP9RV1jQa19fqfBQRERG5G1sPEBEREbmgqNSxRAEAuGbK6zhwLLvZdXqDqdk1oUF++GjJNMhk/DhHRERERN5Tes52S62wpBgmCRARtTPKQD+rMUOdFpIo+SAaIiJyN55ZJiIiInJBUWml42tLKjHm0ffw5ifbm1ynMxib3evYpn8iKS7c4dcmIiIiInKVrqYOVUVlVuMyuQwRXeJ8EBEREXmSOsA6UUA0iTDo9D6IhoiI3I2JAkREREQuWPze9y0+5rWPtmJfWpbdeb2+6USBK/t2afFrEhERERG5qvRcASTJ+i7SsMQYKBxosUVERG2LykZFAQDQ12i9HAkREXkCEwWIiIiInFRbp0d2QXnzC20Y9/j/Qas32JzTNZMo4K9WOfWaRERERETO0tdqUVlQajXOagJERO2XykZFAcD8O4GIiNo+JgoQEREROWnPoQyXjt+656TNcb3B1ORxJlF06XWJiIiIiFqqNLPQZjWB0IQoKP2YyEpE1B7ZTRRgRQEionaBiQJERERELWQwmrD4ve9xz7xPXNpn9bf7bY7rDU1XFDAYm04kICIiIiJyJ4NWj4r8EqtxQRAQkRzvg4iIiMgbZAq5zWQwVhQgImofmChARERE1EKvfbQV733xq8v71Gr1AACTqXGFAG0zrQcuXk9ERERE5EllmQWQbFS1Ck2Igspf7YOIiIjIW2xVFWBFASKi9oGJAkREREQt9MFXu9yyT26RBqPuX45uN7+AOx7/P2TllwFovqKASbQu+UpERERE5AlGnQGavPNW44IgIDKF1QSIiNo7W4kCBq0OoonVDomI2jomChARERG1gKaqDrpm7vh3VMH5CqSl58NgNOGPtCxMnfMRJEmCnhUFiIiIiKiVKMsqhGjj82dIXITd3tVERNR+qAJt/1uvr9V5ORIiInI3JgoQERERtUBOQZnLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 2500x400 with 3 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAB/kAAAHeCAYAAABwqrgpAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdd1xT1/sH8M/NZk8BQUTErbgHdVRb21pHq9Zatdra8bV7fbtra6cddvhrv921Q+1Qq1atte7VOmrdgltREFCQDYHs+/sDiUAGSUggwOf9evkyufecc58ERHKfc54jiKIogoiIiIiIiIiIiIiIiIiIiLyepKEDICIiIiIiIiIiIiIiIiIiIscwyU9ERERERERERERERERERNRIMMlPRERERERERERERERERETUSDDJT0RERERERERERERERERE1EgwyU9ERERERERERERERERERNRIMMlPRERERERERERERERERETUSDDJT0RERERERERERERERERE1EgwyU9ERERERERERERERERERNRIMMlPRERERERERERERERERETUSMgaOgAiIiIiIiIiIiIiooZ29uxZLFq0CDt27EB2djYAIDY2Ftdddx2mT5+O0NDQBo6QiIiIqIIgiqLY0EEQERERERERERERETWU+fPn48MPP4Rer7d6PiwsDF988QV69uxZv4ERERERWcEkPxERERERERERERE1Wz/++CNmz54NAPDx8cHtt9+OxMREaDQarFmzBnv27AEABAUFYc2aNWjRokVDhktERETEJD8RERERERERERERNU8ZGRkYPXo0NBoNQkNDsWDBAnTo0KFam9mzZ+PHH38EAEybNg2zZs1qiFCJiIiIzCQNHQARERERERERERERUUP4/PPPodFoAAAff/yxRYIfAJ5//nmEhoYCANauXVuv8RERERFZI2voAIiIiIiIiIiIiIiI6ptOp8OGDRsAANdffz0GDBhgtZ1CocBjjz2G8+fPIyQkBDqdDgqFoj5DJSIiIqqGSX4iIiIiIiIiIiIianZ2796N0tJSAMD48ePttp06dWp9hERERETkECb5iYiIiIiIiIiIiKjZOXHihPlxjx49zI/z8/ORmpoKrVaLuLg4tGrVqiHCIyIiIrKJSX4iIiIiIiIiIiIianZOnz4NoKIcf2RkJNLT0/Hee+9h+/btMBgM5naJiYmYOXMmevfu3VChEhEREVUjaegAiIiIiIiIiIiIiIjqW3Z2NgAgKCgIe/fuxdixY7F58+ZqCX4ASE5Oxl133YU1a9Y0RJhEREREFpjkbwDTpk3DtGnTGjoMIiIiIqJ6x9+FiYiIiMhbqNVqAEB5eTkee+wxlJWV4fbbb8cff/yB5ORkbNq0CTNmzIBEIoHBYMCLL76I48ePu3w9/i5MRERE7sJy/Q3g4sWLDR0CEREREVGD4O/CREREROQtKpP8paWlAIAnnngCjz76qPl8bGwsnn32WbRq1QqvvfYadDodPvjgA3z//fcuXY+/CxMREZG7cCU/ERERERERERERETVrHTp0wCOPPGL13OTJk9GjRw8AwM6dO5msJyIiogbHJD8RERERERERERERNTs+Pj7mx6NHj4YgCDbb3nzzzebHBw4c8GhcRERERLVhkp+IiIiIiIiIiIiImh1/f3/z47Zt29ptGx8fb36cnZ3tsZiIiIiIHMEkPxERERERERERERE1O61atXK4rUKhMD82mUyeCIeIiIjIYUzyExEREREREREREVGz06FDB/PjzMxMu21zc3PNjyMjIz0WExEREZEjZA0dABERERERERHVLj81C0cWb0J2ciq0JWVQBvgiMrEtuk++AaFtoxs6PCIiokanX79+5sd//fUX7r33XpttDx06ZH5cdXIAERERUUNgkp+IiIiIiIjIi+UcP49tby/EhV0pEKQSiMarJYIz953A/m//QOzAbhj28t2I6Nym4QIlIiJqZBITExEXF4e0tDTs3r0bR44cQffu3S3aFRQUYM2aNQCAhIQEdOzYsb5DJSIiIqqG5fqJiIiIiIiIvFTazmQsmvAKMvYcA4BqCf6qzzP2HMOiCa8gbWdyvcdIRETUmD3yyCMAAFEU8eyzzyIrK6vaeZ1Oh+effx5FRUUAgHvuuae+QyQiIiKywJX8RERERERERF4o5/h5rJwxBwatHhBFu21FowlGkx4rZ8zBlOWzuaKfiIjIQePGjcPmzZuxYcMGpKWl4dZbb8XEiRPRpUsX5OXlYcmSJUhNTQUA9O/fHxMnTmzgiImIiIiY5CciIiIiIiLyStveXgij3lBrgr+SKIow6g3Y/vaPmPjTLA9HR0RE1HTMnTsXr7zyClauXImSkhJ8//33Fm0GDx6MTz75BIIgNECERERERNUxyd/MXLhUgK17TiI02A/X9e8APx9lQ4dERERERERENeSnZuHCrhSn+4lGE9J3JaPg3EWExLf0QGRERERNj1wux5w5czB+/HgsXboU+/fvR15eHoKDg9G+fXtMmjQJN954IyQS7n5LRERE3oFJ/mZkz5FzuOuFBSjT6AAAPTrG4JcP70OQv08DR0ZERERERERVHVm8CYJUAtFocrqvIJXg8KKNGDbzbg9ERkRE1HQlJSUhKSmpocMgIiIiqhWT/M3IR/M3mxP8AHD4ZCZ+33IEd906oAGjIiIiIvIub7zxBn755Rc89thjePzxx2ttv337dixatAhHjhxBcXExwsLC0L17d0yZMgUDBw506JpNaQwie/JTs3Bk8SZkJ6dCW1IGZYAvIhPbovvkGxDaNrqhw/Mq2cmpLiX4gYrV/Dkp59wcERERERERERF5Cyb5m5FdB1Mtjr322R9M8hMRERFdsXv3bixevNihtiaTCa+++iqWLl1a7filS5dw6dIlbNiwAdOmTcMrr7xic9/OpjQGkT05x89j29sLcWFXisXq9Mx9J7D/2z8QO7Abhr18NyI6t2m4QL2ItqSsbv2L1W6KhIiIiIiIiIi8DZP8zYTBaLR6XKe3fpyIiIiouUlJScFjjz0Gk8mxlbOffPKJOSnesmVL3HnnnYiJiUFqaip++eUX5Ofn46effkJYWBgeeeSRJj8GkS1pO5OxcsYcGPUGALBYnV75PGPPMSya8ArGzXsBcYMS6z1Ob6MM8K1b/0A/N0VCRERERERERN6GSf5moqxc39AhEBEREXmt7du349lnn0VpaalD7c+ePYt58+YBANq3b4+ff/4ZQUFB5vOTJk3CtGnTkJaWhi+++AJjx45FTExMkx2jOfl722YcX/IxWuafh69BizKZEhdD26DzpKcwZNjwhg7P6+QcP4+VM+bAoNUDomi3rWg0wWjSY+WMOZiyfHazX9Ef2a0tMveegOjgxKOqBKkEEd3iPRAVEREREREREXkDSUMHQPWjtFzb0CEQEREReR2dTof//e9/eOihh1BcXOxwvwULFsB4pVLS66+/Xi0pDgARERGYM2cOAECv1+OHH35o0mM0B+v/WIFVk3ug3ZwJGHl+B/oWX0CXshz0Lb6Aked3oN2cCVg1uQfW/7GiQeM0ZJxB8bxXkPf8aFx+dAjynh+N4nmvwJBxpkHi2fb2wooV/FcS/AFyNXq1OInhsXtxc9xuDI/di14tTiJAXlFaXhRFGPUGbH/7xwaJ15vev9hBiRYJ/trev0qi0YQeU26sz3ABeNf7R0RERERERNSUMcnfTJTZSfLrrpTNJCIiImpOdu3ahZEjR+Lzzz+HyWSCr68v7r333lr7mUwmrF+/HgDQoUMH9O3b12q7Xr16oWvXrgCA9evXQ6yyirkpjdEcLPruM7T/cgZ6F6UDAGSo/vorn/cuSkf7L2dg0Xef1XuM+tRk5L14Ky7P6Av1yi+hS94JQ2oydMk7oV75JS7P6Iu8F2+FPjW53mLKT83ChV0pEI0mBCtLcH2rfbil7U50DElHpG8BQlUliPQtQMeQdNzSdieub7UPwcoSiEYT0nclo+DcxXqL1dvev9zULOh0eoR1bA1BIjj1/gkSAa0HJSIkvmW9xAp43/tHRERERERE1NQxyd9MlJbpbJ5re9OruP+VH/HUu0ux/2h6PUZFRERE1HB+//13ZGRkAAC6deuGpUuX4rrrrqu136lTp1BYWAgASEpKstu28nxOTg5OnjzZJMdo6tb/sQIDfnsdCpPBIrlfkwwiFCYDBvz2er2u6Nce3I7cp2+CLnlnxQGTsXqDK891yTuR+/RN0B7cXi9xHVm8CYJUgkjfPNzUeg8ifPMBABKh+vtY+TzCNx83td6DSN88CFIJDi/aWC9xetv7l3f+Ii6fqfjZ1HniMET5FTj1/kX5FaDPjFs9GmNV3vb+ERERERERETUHsoYOgOqHupZy/et3HgcA/LbpEH6acw+u7du+PsIiIiIialChoaF47LHHMHnyZEilUuTl5dXa59SpU+bHHTp0sNu2Xbt25scnTpxAp06dmtwYTZ3mp9chNxkhdbC9FIDcZET5T68DY8Z7MLIK+tRk5L8xGdBpat3zHiYjoNMi/43JCJ+7AfK2iR6NLTs5FUGyIgyNOQipYIIg2G8vEQDAhKExB7EhfQByUs55ND7A+96//LRLyDl1wfw8WFmCoa0OQjA68f61OoiyzBPQFLeDKtDP7TFW5W3vnz2GjDMoWzsf+tMHYVIXQ+IXCHn7XvAdeQ9krdrVPgDj8+r4iIiIiIiImhsm+ZuJsnLbK/mrMplELFj5D5P8RERE1ORNnToVr7/+OlQqlVP9MjMzzY9jYmLstm3Z8mq57Kr9mtIYTdnf2zajf1Ga0/1kEDGgKA1///kjBif1qXKmRpa2WtbWtXPFXzwN6LW1J1griSZAr0Xxl88gdNbXjvVxkkGvR2lOIUovXkLvFichcSDBX6kyUd2rxUmkFHSEWJTqkRgrufb+6VD8zUsIfetXQCIFBCkgSCAIdSuUl5+ejeyT1Sur+Wz/GoJodOr9E0UjlFuLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 2500x400 with 3 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAACAoAAAHeCAYAAAAYKXt1AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdd3hUZdoG8PtMTe+dhCRU6aCoiCIgqGsHlQVEsCCuruhasCziLq6iYl0+xXWVtYAKFgQLKk1ApUkndEJ6b5NMyvQ53x9DQsL0lkyS+3dd6Mw57znnSZ857/M+jyCKoggiIiIiIiIiIiIiIiIiIiLqFiQdHQARERERERERERERERERERG1HyYKEBERERERERERERERERERdSNMFCAiIiIiIiIiIiIiIiIiIupGmChARERERERERERERERERETUjTBRgIiIiIiIiIiIiIiIiIiIqBthogAREREREREREREREREREVE3wkQBIiIiIiIiIiIiIiIiIiKiboSJAkRERERERERERERERERERN0IEwWIiIiIiIiIiIiIiIiIiIi6EVlHB0BERERERERERERE1BWcOXMGK1euxO+//47y8nIAQFpaGsaPH4+77roLMTExHRwhERERkYUgiqLY0UEQEREREREREREREXVmH3/8MV5//XUYDAab+2NjY/Huu+9i+PDh7RsYERERkQ1MFCAiIiIiIiIiIiIi8sKKFSvw4osvAgCCg4Nx++23Y8iQIdBqtVi3bh12794NAIiMjMS6desQHx/fkeESERERMVGAiIiIiIiIiIiIiMhTRUVFuOGGG6DVahETE4NPPvkE/fr1azPmxRdfxIoVKwAAd955J5577rmOCJWIiIiohaSjAyAiIiIiIiIiIiIi6qyWLl0KrVYLAPj3v/9tlSQAAE899RRiYmIAAD/99FO7xkdERERki6yjAyAiIiIiIiIiIiIi6oz0ej02bNgAALjqqqtw6aWX2hynUCgwd+5c5OXlITo6Gnq9HgqFoj1DJSIiImqDiQJERERERERERERERB7YuXMnGhoaAACTJ092OHbGjBntERIRERGRS5goQERERERERERERETkgRMnTrQ8HjZsWMvjmpoa5OTkQKfTIT09HampqR0RHhEREZFdTBQgIiIiIiIiIiIiIvLA6dOnAVhaCyQmJqKgoACvvPIKtm3bBqPR2DJuyJAhmD9/Pi688MKOCpWIiIioDUlHB0BERERERERERERE1BmVl5cDACIjI7Fnzx7ccsst2Lx5c5skAQDIysrCzJkzsW7duo4Ik4iIiMgKEwU6qTvvvBN33nlnR4dBRERERNTu+FqYiIiIiAJFY2MjAECj0WDu3LloamrC7bffjh9++AFZWVnYtGkT5syZA4lEAqPRiGeeeQbHjx/3+Hp8LUxERES+wtYDnVRpaWlHh0BERERE1CH4WpiIiIiIAkVzokBDQwMA4JFHHsFDDz3Usj8tLQ3z5s1Damoq/vnPf0Kv1+O1117Dhx9+6NH1+FqYiIiIfIUVBYiIiIiIiIiIiIiIvNSvXz/89a9/tblv2rRpGDZsGABg+/btnPAnIiKiDsdEASIiIiIiIiIiIiIiDwQHB7c8vuGGGyAIgt2xf/rTn1oe79+/369xERERETnDRAEiIiIiIiIiIiIiIg+EhYW1PO7Vq5fDsZmZmS2Py8vL/RYTERERkSuYKEBERERERERERERE5IHU1FSXxyoUipbHZrPZH+EQERERuYyJAkREREREREREREREHujXr1/L4+LiYodjq6qqWh4nJib6LSYiIiIiV8g6OgAiIiIiIiKirqImpwSHV21CeVYOdPVNUIaHIHFILwydNhExvVI6OjwiIiLysYsvvrjl8a+//op77rnH7tiDBw+2PG6dYEBERETUEZgoQEREREREROSliuN52LpoOQp3HIEglUA0nSsnXLz3BPYt+wFpowdj3LOzkDAgo+MCJSIiIp8aMmQI0tPTkZ+fj507d+Lw4cMYOnSo1TiVSoV169YBAHr37o3+/fu3d6hEREREbbD1ABEREREREZEX8rdnYeVtC1C0+xgAtEkSaP28aPcxrLxtAfK3Z7V7jEREROQ/f/3rXwEAoihi3rx5KCkpabNfr9fjqaeeQl1dHQDg7rvvbu8QiYiIiKywogARERERERGRhyqO52HtnMUw6gyAKDocK5rMMJkNWDtnMaavfpGVBYiIiLqISZMmYfPmzdiwYQPy8/Nx8803Y8qUKRg4cCCqq6vxxRdfICcnBwBwySWXYMqUKR0cMRERERETBYiIiIiIiIg8tnXRcpgMRqdJAs1EUYTJYMS2RSsw5dPn/BwdERERtZc333wTCxYswNq1a1FfX48PP/zQaswVV1yBJUuWQBCEDoiQiIiIqC0mChB1A/klNdj6xynERIZgwqgLEBKs6OiQiIiIiIg6vZqcEhTuOOL2caLJjIIdWVDlliI6M9kPkREREVF7k8vlWLx4MSZPnoyvvvoK+/btQ3V1NaKiotC3b19MnToVV199NSQSdgMmIiKiwMBEAaIubsfBHNz99+Vo0uoBAMP698DK12cjIiyogyMjIiIiIurcDq/aBEEqgWgyu32sIJXg0MqNGDd/lh8iIyIioo4yatQojBo1qqPDICIiInKK6YtEXdzrH25sSRIAgEMni/HtL4c6MCIiIiIioq6hPCvHoyQBwFJVoOJIro8jIiIiIiIiIiJyDSsKEHVhRpMJf2TlW23/z6pfMfPmSzsgIiIiIiKirkE0i9DUqL06h07d6KNoqL3U5JTg8KpNKM/Kga6+CcrwECQO6YWh0yYipldKR4dHRERERERE5DImChB1YZU1DTa3F5SqcOB4IUYMSGvniIiIiIiIOjfRbEZdaTWq80ohenkuZUSoT2Ii/6s4noeti5ajcMcRq3YTxXtPYN+yH5A2ejDGPTsLCQMyOi5QIiIiIiIiIhex9QBRF1ZaaX+F058fX4bsgop2jIaIiIiIqPMym0xQFZTjzO9ZKD2aC32jFpHpSRAkgkfnEyQCItMTIYrephuQv+Vvz8LK2xagaPcxALBqN9H8vGj3May8bQHyt2e1e4xERERERERE7mKiAFEXVlpZZ3efRmvAmk2H2jEaIiIiIqLOx2Q0oTqvFGd+P4yyE/kwaHUt+9LGDIVo9myiXzSLiB2YiYJ9J6Fv1PoqXPKxiuN5WDtnMYw6g1WCwPlEkxkmnQFr5yxGxfG89gmQiIiIiIiIyENMFCDqwho1Oof7l6zY0k6REBERERF1LiaDEZVninHmt0OoOFUIo85gNSYsMQax/Xu6XVVAkAiIvaAnQhOj0VSjRu6uI6jKLYFodjwRTe1v66LlMBmMgIuVH0RRhMlgxLZFK/wcGREREREREZF3mChA1IXp9caODoGIiIiIqFMx6gyoOFWI7N8OoepMsWWS2IEBU8ZBIpMCgovJAoIAQSrFgNvHtWwym8yoPF2EvN3HoKlr8CJ68qWanBIU7jjitJLA+USTGQU7sqDKLfVTZERERERERETeY6IAURemM5icjimvVrdDJEREREREgc2g0aHseD7O/H4I1XmlMBudv5YOCg/BgOsvw+Rlz0CmlEOQOn6LLUgESGRSjHxoMiJSE6z2a+ubkP/HcZSfLHDp+uRfh1dtcvo1tUeQSnBo5UYfR0RERERERETkO7KODoCI/EfnQkWBi25/BUVbXmqHaIiIiIiIAo++SYvq3FLUlVa7XPo/ODIMsZnJCIuPgiAIiEiKxfTVL2LbohUo2JEFQSppswq9+XnaZYNxyQOTYDCZYdDabhMmiiJq8stQX6FC0oAMhMVF+uTjJPeVZ+W4XU2gmWgyo+JIro8jIiIiIiIiIvIdJgoQdWF6J2VSmxWU1qBncoyfowHUDVqIEBEZFuz3axEREREROaJraEJVbinqy2oguth/PiQmAnGZyQiJiYBwXquBhAEZmPLpc1DlluLQyo2oOJILnboRyohQJAzOxLDpVyM6MxkAYDaaUJldBFVhhd1rGzQ6FO4/iciUOCT0S4NMIffuAyaXmU0mqMtqUF9W49V5dOpGH0VERERERERE5HtMFCDqwlypKAAA2QWVfk0UMBhN+Mfb3+PT7/dAFEVMmjAMrzw+CWEhSr9dk4iIiIjIFk1dA6pzS1FfoXL5mLC4KMRmJiMkOtzp2OjMZIybP8vhGIlMisQL0hGRFIvSY3nQNTTZHVtXUoXGqjok9O+JiKQYqwQF8h2DRgdVUQXqiqtg1BsgVXh3y0QREeKjyIiIiIiIiIh8j4kCRF2YqxUFpBLP+m666suf92PFd3+0PF+7+RD6ZSTgkTvH+/W6RERERESApZx/k6oe1bmlaKyuc+kYQRAQnhCN2MxkBEWE+iWu4KgwZI4aiOq8MlTllNhtfWDUG1CSdQbq0mokDkiHIpgJt77S/L2hKihHQ2VtmwoPkelJUJ0phmh2reJEa4JEQHBcFHQNTVCGMWGAiIiIiIiIAg8TBYi6MJ2LiQJancGvcXz58z6rbV9vOMBEASIiIiJyW01OCQ6v2oTyrBzo6pugDA9B4pBeGDptImJ6pbQZK4oiGqvqUJ1XiiZVvUvnFwQBEcmxiM1IhrIdWmYJEgnieqUgPDEaZcfyHMbZUFWLpp31iO+TiujUBAgSVhfwlNloQl1pNVSFFXYrOqSNGYrcTXs9Or9oFpFy8QDk7jyK6J6JiOuVAqmct2CIiIiIiIgocPBdKlEX5mrrAY0fEwWMJhP2HS2w2p5TWIV9xwpw0cCefrs2EREREXUdFcfzsHXRchTuOAJBKoFoOrf6vnjvCexb9gPSRg/GuGdnIf6CdNRXqFCdWwqti33iBYkEUSlxiMlIgiIkyF8fhl3K0GD0HHkBaosqUHm6CCajyeY4s9GE8hP5UJdWI3lQBleru0nfpIWqsAJ1JVUwOUmsDkuMQWz/nqg5XehWVQFBIiCmXxpCE6MhiiJLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 2500x400 with 3 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAACAoAAAHeCAYAAAAYKXt1AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdd3hUZdoG8PtMS5n0npBAEkInIBZAREEFUVGBVQREsLC4qNixIbroisqqfMsq6oKrIgooIEVRqpJVQKSahB5Ceu996vn+GDJkMn0y6ffvupCZc97zzhMMYea8z/s8giiKIoiIiIiIiIiIiIiIiIiIiKhbkLR3AERERERERERERERERERERNR2mChARERERERERERERERERETUjTBRgIiIiIiIiIiIiIiIiIiIqBthogAREREREREREREREREREVE3wkQBIiIiIiIiIiIiIiIiIiKiboSJAkRERERERERERERERERERN0IEwWIiIiIiIiIiIiIiIiIiIi6ESYKEBERERERERERERERERERdSNMFCAiIiIiIiIiIiIiIiIiIupGZO0dABERERERERERERFRV3DhwgWsW7cOv/32GwoLCwEAMTExuPHGG/HAAw8gKCionSMkIiIiMhBEURTbOwgiIiIiIiIiIiIios7siy++wHvvvQeNRmPxfHBwMD766CNcccUVbRsYERERkQVMFCAiIiIiIiIiIiIiaoE1a9bgzTffBAB4eXnhnnvuQWJiIhoaGrB9+3YcOnQIAODv74/t27cjNDS0PcMlIiIiYqIAEREREREREREREZGrcnJyMHHiRDQ0NCAoKAirV69G3759Tca8+eabWLNmDQDg/vvvx6uvvtoeoRIREREZSdo7ACIiIiIiIiIiIiKizmrFihVoaGgAAPzrX/8ySxIAgBdeeAFBQUEAgJ9++qlN4yMiIiKyRNbeARARERERERERERERdUZqtRq7du0CANx0000YMWKExXEKhQLz589HRkYGAgMDoVaroVAo2jJUIiIiIhNMFCAiIiIiIiIiIiIicsHBgwdRU1MDAJgyZYrNsTNnzmyLkIiIiIgcwkQBIiIiIiIiIiIiIiIXnDlzxvh46NChxsdlZWVIT0+HSqVCr169EB0d3R7hEREREVnFRAEiIiIiIiIiIiIiIhecP38egKG1QHh4OLKysvDOO+8gKSkJWq3WOC4xMRELFy7ElVde2V6hEhEREZmQtHcARERERERERERERESdUWFhIQDA398fhw8fxqRJk7B3716TJAEASElJwaxZs7B9+/b2CJOIiIjIDBMFOqn7778f999/f3uHQURERETU5vhemIiIiIg6itraWgBAfX095s+fj7q6Otxzzz344YcfkJKSgj179mDu3LmQSCTQarV46aWXcPr0aZdfj++FiYiIyF3YeqCTys/Pb+8QiIiIiIjaBd8LExEREVFH0ZgoUFNTAwB48skn8fjjjxvPx8TEYMGCBYiOjsbf//53qNVqvPvuu/jss89cej2+FyYiIiJ3YUUBIiIiIiIiIiIiIqIW6tu3Lx577DGL56ZPn46hQ4cCAPbv388FfyIiImp3TBQgIiIiIiIiIiIiInKBl5eX8fHEiRMhCILVsbfeeqvx8bFjx1o1LiIiIiJ7mChAREREREREREREROQCHx8f4+P4+HibY+Pi4oyPCwsLWy0mIiIiIkcwUYCIiIiIiIiIiIiIyAXR0dEOj1UoFMbHer2+NcIhIiIichgTBYiIiIiIiIiIiIiIXNC3b1/j49zcXJtjS0pKjI/Dw8NbLSYiIiIiR8jaOwAiIiIiIiIiahtl6XlIXr8HhSnpUFXXwcPXG+GJ8RgyfRyC4qPaOzwiIqJO55prrjE+/t///oeHHnrI6tgTJ04YHzdNMCAiIiJqD0wUICIiIiIiIuriik5nYN+SL5F9IBWCVAJRd7ncce6RMzj66Q+IGTUYY1+ZjbABse0XKBERUSeTmJiIXr16ITMzEwcPHkRycjKGDBliNq68vBzbt28HAPTu3Rv9+vVr61CJiIiITLD1ABEREREREVEXlrk/BevuXoScQ6cAwCRJoOnznEOnsO7uRcjcn9LmMRIREXVmjz32GABAFEUsWLAAeXl5JufVajVeeOEFVFZWAgAefPDBtg6RiIiIyAwrChARERERERF1UUWnM7Bl7lJoVRpAFG2OFXV66PQabJm7FDM2vcnKAkRERA6aPHky9u7di127diEzMxN33XUXpk6dioEDB6K0tBTffPMN0tPTAQDDhw/H1KlT2zliIiIiIiYKEBEREREREblNWXoektfvQWFKOlTVdfDw9UZ4YjyGTB+HoPioNo9n35IvodNo7SYJNBJFETqNFklL1mDqV6+2cnRERERdx7Jly7Bo0SJs2bIF1dXV+Oyzz8zGjB49GsuXL4cgCO0QIREREZEpJgoQEV2iUmuRdPg8cgrLccPVCUjoGdbeIRERERFRJ1F0OgP7lnyJ7AOpEKQSk/L+uUfO4OinPyBm1GCMfWV2m+3UL0vPQ/aBVKevE3V6ZB1IQfnFfATGRbZCZERERF2PXC7H0qVLMWXKFGzYsAFHjx5FaWkpAgIC0KdPH0ybNg3jx4+HRMJuwERERNQxMFGAiAiGJIGHXvkS/zuSBgCQSSX4ZPF9uHX0wHaOjIiIiIg6usz9Kdgyd6lh5z5gkiTQ9HnOoVNYd/ciTF71Inpdl+jSazXu+NeptdCqNZcea6BVG47pNI2PNUhesxOCRICod6yaQFOCVII/1+3G2IWzXYqTiIiouxo5ciRGjhzZ3mEQERER2cVEASIiAEmHzxuTBABAq9PjnVU7mShARERERDYVnc7AlrlLoVVp7Jb3F3V66PQabJm7FDM2vYmwAbGXF/6bLv6rmyz+azSXjhuO6TRaiA62EShPz3MpSaAx1qLUiy5dS0REREREREQdHxMFiIgAvPmfn8yOpWUVo6K6HgG+Xu0QERERERF1BvuWfGmoJODg4r0oitCpNdjxwse4dsF06NSOL/w7S9ugbtH1qqpaN0VCRERERERERB0NGyIREQFIzy6xeFyn07VxJERERETUWZSl5yH7QKpZqwF7RL2I4pMXUZlV1GpJAgAg81S06HoPP6WbIiEiIiIiIiKijoaJAkTU7els3NhVaZgoQERERESWJa/fA0Hq2sdqQSIg69c/3RyRKf9eERAkgkvXClIJwgbHuTkiIiIiIiIiIuoomChARN3el9sOWT2nVmvbMBIiIiIi6kwKU9KdribQSNSLqMwqdGs8giBAppDDw8cL3kF+GDDlBoh61yoWiDo9hs4Y79b4iIiIiIiIiKjjkLV3AERErsrMK8P/rd6L8qo63DSiH2bdNRwSiXP5T6Io4tV/f2/1vFrDRAEiIiIiskxVXdei67X1KpvnBUGAVC6DVCGDVCGHzOyxHFKFDLJLv0tlMrMKAimjBiPn0CmnEhoEqQQxIwchMC7Spa+LiIiIiIiIiDo+JgoQUadUWVOP2//2ISprGgAAe38/i6Kyajz/sHO7nv48m2vzvJqtB4iIiIjICg9f7xZd7+mvRGBMGKRyOWSKpgv/lx5bWPh31thXZmPd3Yug02sgig5UF7iUnDBm4awWvS4RERERERERdWxsPUBEndLaHw4bkwQafbH5IDRa5xb2tyel2jyvYkUBIiIiIrIiPDEegtS1j9WNu/YjBsQiNKEHAnuGwy8iCMogP3j4eEOmkLc4SQAAwgbEYvKqFyH1kNuNVZAIkMikmPDe4wgbENvi1yYiIiIiIiKijouJAkTUKf3ry5/NjlXWNCA7v9ypeX4+dNbmebWaiQJEREREZNmQ6eOcKunflKjTY+gM56phuarXdYmYselNxIwYBABmCQONCQlBfWMw6sX7oAjwgah3oPoAEREREREREXVabD1ARJ1OTZ0KtfVqi+fSsooRHxPi8Fw5BbYTC9h6gIiIiIis8fBXIrhfT5Sdz3ZqYb2xmkBgXGQrRmcqbEAspn71Ksov5uPPdbtRlHoRDZU10Isi/KLD0PP6oVCGBwIAVDX1qMgtRmBMWJvFR0RERERERERtq8snCrz++utYu3Yt5s+fjyeeeMLm2Pr6emzatAm7d+/GuXPnUF1dDaVSifj4eNx8882YMWMGlEql1etzc3Nx0003ORRXXFwcduzY4dTXQkQGu/aftnou287Cf1MqtdZqwkEjNVsPEBEREZEFWpUGeanpGDB1LA4sXQtR1AGi/WQBQRAglcswZuGsNojSXGBcJMYunG18XpFThPxTGWbjSi7kwi8yGFKZtA2jIyIiIiIiIqK20qVbDxw8eBDr1693aOyZM2dw55134h//+Ad+//13lJWVQaPRoKKiAseOHcO7776L22+/HSdPnrQ6x9mztkuYE5GBWqNFnZ0Felv+d+S81XPOzFteVWd3jIqJAkRERETUjCiKyD+ZDq1KA7/oMFz9+BRIZFJjCX9rBKkEUg85Jq96EWEDYtsmWDv8o0Lh4eNtdlyr1qAsI78dIiIiIiIiIiKittBlKwqkpqZi/vz50Ovt94ssLCzEnDlzUFJSAgC44oorcPvttyMsLAzFxcX48ccfcfz4cRQUFGDOnDnYtGkTevToYTZP00SBZcuWwcPDw+pr2qpMQNRV6XR6vPHxj/jq+z8giiIm3TQUS5+bDIXcuR9FuUUVVs/V1qscnqesstbuGLWarQeIiIiIyFR5ViFqSiqNz0P698KoF+/D6Y37UHomC4JUAlF3+bNo4/OYkYMwZuGsDpMkAACCREBYn2hkHz9ndq4sswAB0WGQeyraITIiIiIiIiIiak1dMlEgKSkJCxYsQE1NjUPjly1bZkwSmDt3LhYsWGByfvbs2XjvvfewatUqlJeX45133sEHH3xgNk9jokBoaCgmTpzYwq+CqOtZ9+MR/HfTAePzDTuPISYyEM8+cLPN69QaLT7ffBC/Hr2AyFA/ZOSWWR1Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 2500x400 with 3 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAB/kAAAHeCAYAAABwqrgpAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdd3hUZfo38O85U9J7SCAkhITeRYqIKNhXQYqAgiiiLuz+FMsqiy7i2kAXu+9aFwuIAopKR6QJShHphJ5GKum9TT3vH2OGhDnTkpnMJPl+rsvLyXmec+YeCJPJuZ/7fgRJkiQQERERERERERERERERERGR1xM9HQARERERERERERERERERERE5hkl+IiIiIiIiIiIiIiIiIiKiVoJJfiIiIiIiIiIiIiIiIiIiolaCSX4iIiIiIiIiIiIiIiIiIqJWgkl+IiIiIiIiIiIiIiIiIiKiVoJJfiIiIiIiIiIiIiIiIiIiolaCSX4iIiIiIiIiIiIiIiIiIqJWgkl+IiIiIiIiIiIiIiIiIiKiVoJJfiIiIiIiIiIiIiIiIiIiolZC6ekAiIiIiIiIiIiIiIg8LTU1FatWrcLevXuRn58PAIiLi8ONN96IBx98EOHh4R6OkIiIiMhEkCRJ8nQQRERERERERERERESesmzZMrz11lvQ6XSy4xEREfjoo49w1VVXtWxgRERERDKY5CciIiIiIiIiIiKidmvFihVYtGgRAMDPzw9TpkzBgAEDUFdXh82bN+PgwYMAgJCQEGzevBkdOnTwZLhERERETPITERERERERERERUfuUnZ2NsWPHoq6uDuHh4Vi+fDl69uzZaM6iRYuwYsUKAMD999+PF154wROhEhEREZmJng6AiIiIiIiIiIiIiMgTPvzwQ9TV1QEA3nvvPYsEPwDMnz8f4eHhAICffvqpReMjIiIikqP0dABERERERERERERERC1Nq9Vi27ZtAICbbroJ11xzjew8tVqNuXPn4uLFiwgLC4NWq4VarW7JUImIiIgaYZKfiIiIiIiIiIiIiNqdAwcOoKqqCgAwadIkm3NnzJjREiEREREROYRJfiIiIiIiIiIiIiJqd86dO2d+PGjQIPPjkpISpKWlQaPRID4+HrGxsZ4Ij4iIiMgqJvmJiIiIiIiIiIiIqN1JTk4GYGrHHx0djczMTPznP//Bnj17oNfrzfMGDBiABQsW4Oqrr/ZUqERERESNiJ4OgIiIiIiIiIiIiIiopeXn5wMAQkJCcOjQIUyYMAE7d+5slOAHgKSkJDzwwAPYvHmzJ8IkIiIissAkvwfcf//9uP/++z0dBhERERFRi+NnYSIiIiLyFtXV1QCA2tpazJ07FzU1NZgyZQo2bdqEpKQk7NixA7Nnz4YoitDr9Xjuuedw9uzZJj8fPwsTERGRq7BdvwdcunTJ0yEQEREREXkEPwsTERERkbeoT/JXVVUBAJ544gk89thj5vG4uDjMmzcPsbGxePHFF6HVavHmm2/iiy++aNLz8bMwERERuQor+YmIiIiIiIiIiIioXevZsyceffRR2bFp06Zh0KBBAIB9+/YxWU9EREQe16Yq+bVaLb7//nv89NNPOH/+PGpqahASEoIBAwZg4sSJuP322yEIgtXzJUnCpk2b8MMPP+Ds2bOoqalBhw4dMGzYMMyYMQMDBw5swVdDRERERERERERERO7i5+dnfjx27Fib947/8pe/4MSJEwCAo0ePYuzYsW6Pj4iIiMiaNpPkz8/Px9/+9jeLPZGKiorwyy+/4JdffsHo0aPx3nvvwd/f3+L8uro6PPnkk9i9e3ej4zk5OcjJycHGjRvx1FNPYc6cOe58GURERERERERERETUAgIDA82PExMTbc5NSEgwP87Pz3dbTERERESOaBNJfp1O1yjBHx8fj8mTJ6NTp05IT0/H6tWrUVJSgj179uCZZ57Bxx9/bHGN559/3pzg79atG+655x5ERkbi9OnTWL16NWpqavD2228jOjoaEyZMaMmXR0REREREREREREQuFhsbi0OHDjk0V61Wmx8bjUZ3hURERETkkDaR5F+7dq05wX/jjTfivffeg6+vr3n8gQcewEMPPYRz585h165d2Lt3L0aNGmUe37dvHzZt2gQAGDFiBP73v//Bx8cHADBu3DhMmTIF9913H8rKyvDaa6/h5ptvbrTKk4iIiIiIiIiIiIhal549e5of5+Tk2JxbVFRkfhwdHe22mIiIiIgcIXo6AFfYtm0bAEAURbzyyiuNEvwAEB4ejueff95ifr0vvvgCAKBUKrFo0SJzgr9et27d8MILLwAAysrKsGbNGpe/BiIiIiIiIiJbStJysfu1r/Dt9Jfw1bj5+Hb6S9j92lcoScv1dGhERESt0rBhw8yPf/31V5tzjx8/bn7ccHEAERERkSe0iSR/dnY2AFMyPyoqSnbOoEGDzI8brsosKyvD/v37AQDXX3894uLiZM+/8847ERERAQDYunWrS+ImIiIiIiIisqfg7EV8d/8r+PKWp3D0yy3IPngGhWcuIvvgGRz9cgu+vOUpfHf/Kyg4e9HToRIREbUqAwYMQHx8PADgwIEDOHnypOy80tJSbN68GYCpIKxXr14tFiMRERGRnDaR5A8KCgIAFBcXo7q6WnZOw8R+eHi4+fHhw4fNeyiNGDHC6nOIomhe2XnixAmUl5c3O24iIiIiIiIiWzL2JWHV5IXIPngGACAZGu8BXP919sEzWDV5ITL2JbV4jERERK3Zo48+CgCQJAnz5s1Dbm7jDjlarRbz58833w+eNWtWS4dIREREZKFNJPkHDhwIwPRBrL71/pU+++wz8+NRo0aZHycnJ5sf22uz1L17d/PzXLhwocnxEhEREREREdlTcPYi1s1eAr1GZ5Hcv5JkMMKg0WHd7CWs6CciInLCxIkTcdtttwEAMjIyMH78eCxZsgQbN27EsmXLMGHCBHMr/+HDh2Pq1KmeDJeIiIgIQBtJ8j/44IPw9/cHAHz00UdYvHgxUlNTUVtbi/Pnz2P+/Pn44YcfAJg+iI0bN858bsMK/86dO9t8no4dO8qeR0RERERERORquxd/BYNOD0iSQ/MlSYJBp8eexSvcHBkREVHb8s4772DixIkAgMrKSnzxxReYN28eXn/9daSlpQEwFY59/PHHEATBg5ESERERmSg9HYArdOnSBUuXLsXTTz+N/Px8fPXVV/jqq68azVGpVJg2bRqeeeYZKBQK8/GSkhLz47CwMJvPExoaan5cVlbmktiJqHlSMwvx65EUdI4KxehhPeCjbhNva0RERETUzpWk5SJr/ymnz5MMRmTuT0Jp+iWEJXRyQ2RERERtj0qlwpIlSzBp0iSsWbMGR44cQXFxMUJDQ9GjRw/ce++9uPXWWyGKbaJmjoiIiNqANpMNGzp0KN555x089thjsgn44OBgJCQkNErwA0BdXZ35sY+Pj83nUKvVsucRkWds338Wc15cCZ3eAAAYdXU3LHt9JnzVKg9HRkRERETUPCdX74CgEO226ZcjKEScWLUdYxbMdENkREREbdeIESMwYsQIT4dBREREZFebSPLrdDo8++yz2Lx5MwDTh7HbbrsNYWFhyM7Oxvr165GSkoJXXnkFW7duxdKlS+Hr6wsA0Ov15us0TOLLaTje8Dwi8oz/LP3ZnOAHgL1HU/HTr6cx6ZarPBcUEREREZEL5CelNSnBD5iq+QtOpbs4IiIiIiIiIiLyFm2iv9AzzzxjTvC/8MILWL58OWbMmIE777wTc+bMwYYNG3DvvfcCAP744w8sWLDAfG59sh8wLRawRavVmh/bWxBARO5VVaPB+YsFFscfX/wd/vbSStTWaWXOIiIiIiJqHTSVNc07v6LaRZEQERERERERkbdp9Un+P/74Az///DMAYNKkSbj//vst5igUCrz44osYOHAgAGDz5s1ITk4GAPj7+5vnaTQam8/VMMlvr7U/EbmXVme9m8bmPaew6JOfWjAaIiIiIiLX8gnytz/J1vnBAS6KhIiIiIiIiIi8TatP8m/bts38+L777rM6T6FQNFoAsHv3bgBAcHCw+VhZWZnN52o4Hh4e7lygRORSdVrbW2Z8teGPFoqEiIiIiMj1ogckQlA07Vd2QSEiqn+CiyMiIiIiIiIiIm/R6pP8GRkZ5se9evWyObdv377mx9nZ2QCArl27mo9dunTJ5vl5eXnmxzExMc6ESUQuVqexvb2GJEktFAkRERERkesNnHYzJIOxSedKBiMGTb/VxRERERERERERkbdo9Un+hok8e+32RfHyy1UoFACAbt26mY/Vt/C3pn5cEAT06NHD6ViJyHU0dir5iYiIiIhaNYUCEb26QBAFp04TRAEd+naFf1Soe+IiIiIiIiIiIo9r9Un+jh07mh+fOnXK5twLFy6YH9dX4g8ePBgqlQoAcPDgQavnGgwGHDp0CADQu3fvRm3+iajl2avkJyIiIiJqraqLy1FwIQt9po6BoFAAgoOJfkGAoFCg1903IPPQOdRV1rg3UCIiIiIiIiLyiFaf5L/22mvNj7/66iur8yRJwsqVK81fjxo1CgAQHByMESNGAAB27dqF3Nxc2fM3b96MkpISAMAdd9zR7LiJqHkOn8qwO4ct+4mIiIiotdHW1CHnZCokSUJwbBSGPjYJolJht6JfEAWISgWGPjYJwbFR0Gt1yDx8DrVlVS0UORERERERERG1lFaf5L/11lvRuXNnAMAvv/yCjz/+2GKOJEl488038ccffwAwJfh79+5tHp81axYAQKfT4emnn0ZVVeObICkpKVi8eDEAICAgAFOnTnXHSyEiB124mI+XP9pid57B2LQ9TImIiIiIPMGgNyD7eDIMustbU0X2jsfIZ+9DeM84AICgaPxrfP3X4T3jMPLZ+xDZO/7y9XR6ZB49j+qSihaInoiIiIiIiIhaitLTATSXWq3GkiVL8NBDD0Gn0+G9997Djh07MG7cOERLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 2500x400 with 3 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAB/kAAAHeCAYAAABwqrgpAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdd3xTZdsH8N9J0nTvvaAtUIYMkSlDloshQ5ZQUBABXwVURFQEJ6iIAx+3IiIqIIhsRaYoQ2RDoYzSQXfp3s067x+1oSFJm7Rpk7a/7/vx9eTc41wBHzg917mvWxBFUQQRERERERERERERERERERHZPIm1AyAiIiIiIiIiIiIiIiIiIiLTMMlPRERERERERERERERERETUSDDJT0RERERERERERERERERE1EgwyU9ERERERERERERERERERNRIMMlPRERERERERERERERERETUSDDJT0RERERERERERERERERE1EgwyU9ERERERERERERERERERNRIMMlPRERERERERERERERERETUSDDJT0RERERERERERERERERE1EjIrB0AEREREREREREREZG1Xb9+HevXr8fhw4eRkZEBAAgNDcWgQYPw2GOPwcvLy8oREhEREVUQRFEUrR0EEREREREREREREZG1rFmzBu+//z6USqXBdm9vb3z++ee48847GzYwIiIiIgOY5CciIiIiIiIiIiKiZuuHH37A0qVLAQCOjo4YN24cOnXqhLKyMuzatQvHjx8HALi7u2PXrl3w9fW1ZrhERERETPITERERERERERERUfOUnJyM4cOHo6ysDF5eXvj+++8RGRmp02fp0qX44YcfAABTpkzBkiVLrBEqERERkZbE2gEQEREREREREREREVnDZ599hrKyMgDAypUr9RL8ALBw4UJ4eXkBAH7//fcGjY+IiIjIEJm1AyAiIiIiIiIiIiIiamgKhQJ79uwBAAwePBi9evUy2E8ul2POnDlISEiAp6cnFAoF5HJ5Q4ZKREREpINJfiIiIiIiIiIiIiJqdo4dO4aioiIAwJgxY6rtGxUV1RAhEREREZmESX4iIiIiIiIiIiIianYuX76sPe7SpYv2OCcnB3FxcSgvL0fLli0REhJijfCIiIiIjGKSn4iIiIiIiIiIiIianWvXrgGoKMfv7++PGzdu4N1338WhQ4egUqm0/Tp16oRFixbhrrvuslaoRERERDok1g6AiIiIiIiIiIiIiKihZWRkAADc3d1x4sQJjBo1Cvv379dJ8APAhQsXMHXqVOzatcsaYRIRERHpYZLfCqZMmYIpU6ZYOwwiIiIiogbHe2EiIiIishXFxcUAgNLSUsyZMwclJSUYN24cdu7ciQsXLmDfvn2YOXMmJBIJVCoVXnrpJcTExNT6erwXJiIiIkthuX4rSEtLs3YIRERERERWwXthIiIiIrIVlUn+oqIiAMC8efPw9NNPa9tDQ0OxYMEChISE4LXXXoNCocCKFSuwevXqWl2P98JERERkKVzJT0RERERERERERETNWmRkJJ566imDbY888gi6dOkCADhy5AiT9URERGR1TPITERERERERERERUbPj6OioPR4+fDgEQTDa98EHH9Qenz59ul7jIiIiIqoJk/xERERERERERERE1Oy4uLhojyMiIqrtGx4erj3OyMiot5iIiIiITMEkPxERERERERERERE1OyEhISb3lcvl2mONRlMf4RARERGZjEl+IiIiIiIiIiIiImp2IiMjtccpKSnV9s3KytIe+/v711tMRERERKaQWTsAIiIiIiKixiInLhXnN+xDxoU4lBeWwN7VCf6dItD5kXvhFRFk7fCw/2IsXt+7H9eKclEuqmEvSNHGxROv3zcEQ+5obe3wGB/jY2yMr17EJWXhp53/4tyVFBQWl8HV2QFd2gYjakRPRIT6WDs8xkdkw3r06KE9/uuvvzB9+nSjfc+ePas9rvpyABEREZE1CKIoitYOorkZMmQIAGD//v1WjoSIiIiIqGE11nvhzJgE/LlsLZKORkOQSiCqb5Vorfwc2qcjBr7yKPzahzV4fFtORePp33fipr0C0IiARLjV+N9n33J7fDZ0OMZ068j4GF+zic+WY2N8dXcpNg1vfL4Lh67Go6yFDAo3QJQJEFQi5AWAww0VBrSNwGv/NwwdWgcyvkYWH1FDuf/++5GYmAhBELBx40Z07txZr09ubi4eeOAB5Ofno1WrVvjtt99qda3Gei9MREREtofl+omIiIiIiKqReOQC1o9djOTjlwBAJ8Ff9XPy8UtYP3YxEo9caND4Pt57GBP3bsZNu/KKE1WTcFU+37Qrw8S9m/Hx3sOMj/E1i/hsOTbGV3eHT8XiwZe+xE67ZGQPsEdxSwmU3lKo3Cv+XdxSguwB9tgpS8KDL32Jw6diGV8jio+oIT311FMAAFEUsWDBAqSmpuq0KxQKLFy4EPn5+QCAadOmNXSIRERERHq4kt8K+MYmERERETVXje1eODMmAevHLoaqXAmY8KOTIAiQ2tth0ualDbKif8upaEzcuxmiBIAgILwkF5NTzqNzYQZcVeUolNnjvKs/1gV3RryTJyCKEDTAz/eNbZBVt4yP8VkrPluOjfHV3aXYNNy35EtkdpEAAhBelmc8PgcPQAT8zmmw960nG2RFOuMjanzmzp2LPXv2AABcXV0xfvx4dOjQAdnZ2fj5558RFxcHAOjZsyfWrl0LQRCqm86oxnYvTERERLaLSX4r4M0cERERETVXje1eeOOUN5F8/JLe6v3qCFIJQnvdgfE/LqnHyCoELX0XN+3K0b44C0ti/0S/3CSoBAGyKj/mVX4+7BmKt1oPRIyzD3yV9khd/BLjY3xNNj5bjo3x1d0DC7/AAa9MtC+5iSWxh0yIbwBinHwxJNcfu5c/yfhsPD4ia1AqlVi8eDG2bt1qtE+/fv3w8ccfw8XFpdbXaWz3wkRERGS7ZNYOgGxPXmEpDvxzBaXlSowY2BHuLo7WDomIiIiIqMHlxKUi6Wi02eNEtQY3jl5AbnwaPMPrb8Xj/ouxuGmvQN+cJKw+vxV2GjUA6CRpqn7unZuMrafW4/HOo3HEqwU2njiCuyOD6y2+I1eSGR/jMx7fyaPoYyg+ARBgeHWkUOX/G2sFgMNXkuoU2+ZT/6BvZGj1vwB18HcTj2/r6bO4p20YIAgQIPnv91OAIEgq/v3f/1WcM38lbHxyFv4W0tE3LxmrL2wzMb4NeLzTKPwNDaKvZyAs2Nvs6zK+ivjikrMQEeJTb/ERWYudnR2WL1+OMWPGYNOmTTh16hSys7Ph4eGBNm3aYOLEibjvvvsgkXD3WyIiIrINXMlvBbb8xmZKRh4mzF+FxNQcAIC7qyN2fz0HoQGeVo6MiIiIiJoCW74Xvt2fb6/F6e9+M2sVfyVBKsFd04dh4KJH6yGyCv1XfoX87AvYenoD5BoVpCaMUQNQSGQYfdcjsHMLwzdTe9ZbfDN/+BfKggTGx/gaPD5bjo3x3a4y4S+peE1DqPJSACRV2m+9FLBi7UmcKTiLrafWmx9ft0no5tYNLz7Wz8T4zPfu94dxuuBUk41vsP8grHz64XqLj6ipa0z3wkRERGTbuJKfdKzZ+o82wQ8A+YWlWPzxdnz/zmNWjIqIiIiIqOFlXIirVYIfqFjNnxkdb+GIdF0rysUn1w/BTqM2KUkDAFIAdho1Fl//C3PaBQCiut7iu1Fahs8ZH+OzQny2HBvj0ydW/bcJy1AO52Xg3et/1i6+a3/i5VaBmKtKMjk+c/2dl9qk43tNCKu32IiIiIiIyHRM8pOO4+f1H0Tu/+cKsnKL4ONZ+/2miIiIiIgam/LCkrqNLyi2UCSGBZRkol+u+YkgGUT0z72BwJJM+MWsrYfIKgSXuDK+OmB8tWfLsQGMr648FC7ol1fL+PKS4F6ezvjqEJ9TeVo9REVERERERObiJkLNWHpWAbbsO4tzV5KRmpmHBSt+xelLhn/Qi76W2sDRERERERFZl72rU93GuzlbKBJdGo0GGUVpeCT9PFS12M8aAFQQMCnjgoUj0/VIxgXGVweMr/ZsOTaA8dXV2KyLdYpvXHa0hSPS1eTjyzpn4YiIiIiIiKg2uJK/mcnMKcRfJ67hzxPXsHW/6T+YXU/KwsCekfUYGRERERGRbfHvFIGUk5drVbJfkErg1zHc4jHllmYjtSge5eoSdC3NhEw0oba1ATKIuKskA0CYReOrqltJOuOrA8ZXe7YcG8D46qpLHePrXJIBoKVlg6qi6cd308IRERERERFRbTDJ34ycuJCI6YvXIq+g1Oyx8SlZ9RAREREREZHt6vzIvTi1ametxopqDbpMus9isRSVFyKlKB7FyjxAFGFfmIhwSd22A2htp4C/c4hlAjRAIVPWaTzja/rx+TndHl/NiUexmk+VyusYWys7BXwcA+s0R3Uimnp8snJ4Onj/97sjQqz8fRJF7TngtmNR/7yxcf4or1N8ASiDXFa7leym8EdZncbbenyBUvNffCMiIiIiIstjkr8ZWfnDgVol+AHUehwRERERUWPlFRGE0D4dkXz8klmr+QWpBKG974BneN2TcOWqcqQUJiBPkQFoNLAvSIBz1jnIyvMgcRKAnNrP7ebqAIkgrXOMRud3s4eS8dV+/mYQn1RSP/G51zE2d1cH2Enllgvo9vmbenxujnCQuVguoNv4uzkA+XUY7+4AL4f6ewGmzM2xScfn6xdkuWCIiIiIiKjWmORvJkRRxKET12o9fuv+c3h3/mi4ONlbMCoiIiIiIts28JVHsX7sYqg1SoimlDcWBEhkUgxYNLVO11VpVEgrTEZWWTJEjQoOBfFwzjoHafmtzIwY4AExJRuSWpRdFgUJ7Np0hRDYt05xVseuzXUoEtMgiOav+mR8jK8ubDk2gPHVlVu76yhMToW0Fn/2qQUBbm3vgiSwfz1EVsG9XXyTjs+9Xfd6iIqIiIiIiMzFJH8zIIoLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 2500x400 with 3 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAACAoAAAHeCAYAAAAYKXt1AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdd3hUZdoG8PvMJJPeKyEhgdCriECkCALSlY4URdSFz4KiiIooK+tSlrWsrNgVQVSaCFKUXqREFpSSQIBASEgPKZM+mXa+P2KGhCmZSaak3L/r4nLmbedJAJmc87zPK4iiKIKIiIiIiIiIiIiIiIiIiIiaBYmjAyAiIiIiIiIiIiIiIiIiIiL7YaIAERERERERERERERERERFRM8JEASIiIiIiIiIiIiIiIiIiomaEiQJERERERERERERERERERETNCBMFiIiIiIiIiIiIiIiIiIiImhEmChARERERERERERERERERETUjTBQgIiIiIiIiIiIiIiIiIiJqRpgoQERERERERERERERERERE1IwwUYCIiIiIiIiIiIiIiIiIiKgZcXJ0AERERERERERERERETcGNGzewceNGnDhxAtnZ2QCAiIgIPPjgg3jiiSfg7+/v4AiJiIiIKgmiKIqODoKIiIiIiIiIiIiIqDFbt24d3nvvPahUKoP9AQEB+OSTT3DPPffYNzAiIiIiA5goQERERERERERERERUDxs2bMCyZcsAAG5ubpg8eTK6desGhUKBPXv24PTp0wAAHx8f7NmzB0FBQY4Ml4iIiIiJAkREREREREREREREdZWWloYxY8ZAoVDA398f69evR/v27WuMWbZsGTZs2AAAeOyxx7BkyRJHhEpERESkI3F0AEREREREREREREREjdXHH38MhUIBAPjwww/1kgQA4LXXXoO/vz8A4Ndff7VrfERERESGODk6ACIiIiIiIiIiIiKixkipVGL//v0AgCFDhqBv374Gx8lkMsybNw/Jycnw8/ODUqmETCazZ6hERERENTBRgIiIiIiIiIiIiIioDmJjY1FSUgIAmDBhgsmxM2fOtEdIRERERGZhogARERERERERERERUR1cuXJF97pHjx661/n5+UhKSkJFRQUiIyMRHh7uiPCIiIiIjGKiABERERERERERERFRHSQmJgKoPFogJCQEt27dwr/+9S8cO3YMarVaN65bt25YvHgx7r33XkeFSkRERFSDxNEBEBERERERERERERE1RtnZ2QAAHx8fnDlzBuPGjcOhQ4dqJAkAQFxcHB5//HHs2bPHEWESERER6WGiQCP12GOP4bHHHnN0GEREREREdsfPwkRERETUUJSWlgIAysvLMW/ePJSVlWHy5MnYvXs34uLicPDgQcyZMwcSiQRqtRqLFi1CQkJCna/Hz8JERERkLTx6oJHKzMx0dAhERERERA7Bz8JERERE1FBUJQqUlJQAAF588UU8//zzuv6IiAgsXLgQ4eHhePvtt6FUKvHuu+9i7dq1dboePwsT1U2FphSpJRf02oPdouEtC3FAREREjseKAkRERERERERERERE9dS+fXs899xzBvumTZuGHj16AABOnjzJB/5ERETkcEwUICIiIiIiIiIiIiKqAzc3N93rMWPGQBAEo2NHjhype/3nn3/aNC4iIiKi2jBRgIiIiIiIiIiIiIioDjw9PXWv27RpY3Js69atda+zs7NtFhMRERGROZgoQERERERERERERERUB+Hh4WaPlclkutdardYW4RARERGZjYkCRERERERERERERER10L59e93r9PR0k2Nzc3N1r0NCQmwWExEREZE5nBwdABERERERERERETVe+UkZuLjpILLjklBRXAYXL3eEdGuD7tOGwb9NmKPDI7Kp3r17617/9ttvePLJJ42OPX/+vO519QQDIiIiIkdgogARERERERERERFZLCchGUeXf4vUU/EQpBKImjul1NPPXsEfX+1GRL+uGPzmLAR3inJcoEQ21K1bN0RGRiIlJQWxsbG4ePEiunfvrjeuoKAAe/bsAQBER0ejQ4cO9g6ViIiIqIYmlSiQm5uLDRs24NixY0hJSQEAhIaGYsCAAXj00UfRtm1bk/NFUcTu3buxbds2JCQkoKysDEFBQejduzdmzpxp8AOeLdYgIiIiIiIiIiJqyFJOxmHHnFXQqNQAUCNJoPr7tNOXsXHSWxj/5euI7N/N7nES2cNzzz2H119/HaIoYuHChVi3bh3Cwu5U01AqlXjttddQWFgIAJg9e7aDIiUiIiK6QxBFUXR0ENZw8uRJLFiwAHK53GC/s7Mz5s2bh2eeecZgv0KhwPz583H06FGD/VKpFC+99BLmzp1rNAZrrGGuoUOHAgAOHTpU77WIiIiIiBoTfhYmIiJyrJyEZGyc9BbUFSrAjFuLgiBA6uKM6duWsbIANVkvvPAC9u/fDwDw8vLClClT0LlzZ+Tl5WHz5s1ISkoCAPTp0wfffvstBEGo03X4WZiobio0pUgtuaDXHuwWDW9ZiAMiIiJyvCZRUeDKlSt49tlnUVFRAQAYPHgwBgwYAE9PTyQkJGDz5s1QKBT4z3/+A09PTzz22GN6a7z55pu6B/zR0dGYOnUqAgMDcenSJWzatAllZWV4//33ERISgnHjxhmMwxprEBERERERERERNWRHl39bWUnAzP1HoihCo1Lj2PINmPLdEhtHR+QYH3zwAd566y3s2LEDxcXFWLt2rd6YAQMGYPXq1XVOEiAiIiKypiZRUWDmzJk4e/YsAGDp0qWYPn16jf7k5GRMmTIFRUVF8PT0xLFjx+Dp6anrP3nyJJ566ikAQExMDL744gu4uLjo+m/cuIEZM2ZALpfD19cXhw4dqjHfWmtYgpmjRERETV9BURkOn74KpVKDofd3QLC/l6NDImoQ+FmYiIjIcfKTMvDNsJfqPP+pQ6vh17qF9QIiamB+//13bN26FX/88Qfy8vLg6+uLdu3a4dFHH8VDDz0EiURSr/X5WZioblhRgIhIX/0+lTQA169f1yUJDBgwQC9JAACioqLwt7/9DQBQUlKC3377rUZ/VXank5MTli1bVuMBP1BZHWDJkspsZ7lcjq1bt+pdwxprEBEREVXJyJFj7LOfYP6KrXj1vZ8w/G//xbXkbEeHRURERETN3MVNByFI63ZLUZBKcGHjAStHRNSwxMTE4P3338fRo0cRFxeH48ePY+3atRgxYkS9kwSIiIiIrKnRfzKRy+WIiYlBYGAgRo4caXRcx44dda/T09NrzD916hQAYODAgYiIiDA4f/To0QgICAAA7N27Vy+G+q5BREREVN36n08jJSNf9z63oBSfbjruwIiIiIiIiIDsuCSIGm2d5ooaLXLib1o5IiIiIiIiqotGnyhw3333Yf369Th58iSmTJlidFxWVpbudVBQkO712bNnodVW/nATExNjdL5EIkHv3r0BABcuXEBhYaFV1yAiIiKq7uMfjum1bd33pwMiISIiIiK6o6K4rH7zi0qtFAkREREREdVHo08UMIdcLtcdDeDm5oYHHnhA15eYmKh73b59e5PrtG3bFgAgiiKuXbtm1TWIiIiIqpz884ajQyAiIiIiMsjFy71+8709rBQJERERERHVR5NNFKioqEBSUhK+/PJLPPzww0hOTgYALFq0CP7+/rpx1Y8haNmypck1Q0NDDc6zxhpEREREVT7Z9JujQyAiIiIiMiikWxsI0rrdUhSkEgR3bW3liIiIiIiIqC6cHB2ALcTHx2PSpEk12oKCgrB48WKMHj26Rnt+/p2zf/38/Eyu6+vrq3stl8utugYRERFRlWNnEmsfRERERETkAN2nDcMfX+2u01xRo0WP6Q9ZOSIiIiIiIqqLJpkokJWVpdcml8vx66+/onPnzoiKitK1KxQK3WsXFxeT68pkMoPzrLEGEREREQBotVqT/RqNFtI67uAiIiIiIsPykzJwcdNBZMcloaK4DC5e7gjp1gbdpw2Df5swR4fXoPi3CUNEv65I+/0SRK1o9jxBKkFETBf4tW5hw+iIiIiIiMhcTTJRwM/PD0uWLIG/vz+ysrLw888/48qVK9i/fz9iY2Oxfv16dOnSBQCgVqt186o/xDeken/1edZYg4iIiAgA8uSlJvvLK1TwdDedmEhERERE5slJSMbR5d8i9VQ8BKkEouZO0mb62Sv446vdiOjXFYPfnIXgTlGOC7SB6f/KNGyZthSiqAHE2pMFBEGA1NkJgxY/bofoiIiIiIjIHE1yO1qvXr3w2GOPYfTo0Xjqqaewfft2PProowCA4uJiLFy4EBqNBgDg6uqqm6dSqUyuq1Qqda+rP/C3xhpEREREAJCTX2yy/+fDF+0UCREREVHTlnIyDhsnvYW005cBoEaSQPX3aacvY+Okt5ByMs7uMTZEoihCC+C+5ydA4iSFIBFMjhekEkhdnDH+y9eZbEFERERE1IA0yUSBu0kkErz99tvo0KEDACApKQknTpwAALi7u+vGVVRUmFyn+kP+6kcMWGMNIiIiIgDIzjOdKPD6+9vx7toDdoqGiIiIqGnKSUjGjjmroK5Q6SUI3E3UaKGpUGHHnFXISUi2T4ANWElOAUpuyxHYMRL9Xp8B//YRAKCXMCD8dVxWREwXTN+2DJH9u9k9ViIiIiIiMq5ZJAoAgFQqxeTJk3Xv//zzTwCAt7e3rk0ul5tco3q/v7+/7rU11iAiIiICAHlRea1jPt30GwoKy+wQDREREVHTdHT5t9Co1GaVzQcqd9FrVGocW77BxpE1bBq1BtlXb+nee4cHo+9LUzHoH0+j47iBiIjpguDOUYiI6YJ7nxyNpw6txpQNS1hJgIiIiIioAXJydAD21Lp1a93r/Px8AEBUVJSuLTMzExEREUbnZ2Vl6V6HhYXpXltjDSIiIiIAKC03XZ0IAJQqDeIS0/HAfe3sEBERERFR05KflIHUU/EWzxM1Wtw6FYeCm5nwa93Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 2500x400 with 3 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAACAoAAAHeCAYAAAAYKXt1AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdd3gUZdcG8Ht203vvIYXeRREigiCoKChFQTq2F0UFRAVFBEUFEQufqKCCr4ooIIJ0pYooBJEOoYeQ3kjv2TbfH3mzJGzJ7mZbkvt3XeDsPGVPVMhm5sw5giiKIoiIiIiIiIiIiIiIiIiIiKhFkNg6ACIiIiIiIiIiIiIiIiIiIrIeJgoQERERERERERERERERERG1IEwUICIiIiIiIiIiIiIiIiIiakGYKEBERERERERERERERERERNSCMFGAiIiIiIiIiIiIiIiIiIioBWGiABERERERERERERERERERUQvCRAEiIiIiIiIiIiIiIiIiIqIWhIkCRERERERERERERERERERELQgTBYiIiIiIiIiIiIiIiIiIiFoQB1sHQERERERERERERETUHFy7dg3r1q3DoUOHkJOTAwCIjIzEvffeiyeeeAJ+fn42jpCIiIiohiCKomjrIIiIiIiIiIiIiIiImrLvv/8eH3/8MeRyudZxf39/rFixArfddpt1AyMiIiLSgokCRERERERERERERESNsGbNGixcuBAA4OrqilGjRqFr166oqqrCzp07cfToUQCAt7c3du7cicDAQFuGS0RERMREASIiIiIiIiIiIiIiU6Wnp2Po0KGoqqqCn58fVq9ejXbt2tWbs3DhQqxZswYAMHHiRMyfP98WoRIRERGpSWwdABERERERERERERFRU7V8+XJUVVUBAD799FONJAEAeO211+Dn5wcA+P33360aHxEREZE2DrYOgIiIiIiIiIiIiIioKZLJZNizZw8AYODAgejdu7fWeU5OTpg2bRqSk5Ph6+sLmUwGJycna4ZKREREVA8TBYiIiIiIiIiIiIiITHDkyBGUlZUBAEaOHKl37oQJE6wREhEREZFBmChARERERERERERERGSCS5cuqY+7d++uPi4oKEBSUhKqq6sRFRWFiIgIW4RHREREpBMTBYiIiIiIiIiIiIiITHD16lUANa0FgoODkZqaig8++AAHDx6EQqFQz+vatSvmzp2L22+/3VahEhEREdUjsXUARERERERERERERERNUU5ODgDA29sbx44dw/Dhw7F///56SQIAcO7cOUyaNAk7d+60RZhEREREGpgo0ERNnDgREydOtHUYRERERERWx8/CRERERGQvysvLAQCVlZWYNm0aKioqMGrUKOzYsQPnzp3Dvn37MGXKFEgkEigUCsyZMwcXL140+f34WZiIiIjMha0HmqisrCxbh0BEREREZBP8LExERERE9qI2UaCsrAwAMGPGDLz44ovq8cjISMyaNQsRERF4++23IZPJ8NFHH+Hbb7816f34WZjsRUmVAsUVNZUzJOWpcEv6AQDg4SKFRBBqzrd5HIJfJ5vFSERE+rGiABERERERERERERFRI7Vr1w4vvPCC1rGxY8eie/fuAIDDhw/zhj81eYKeVzeJlg+EiIhMxkQBIiIiIiIiIiIiIiITuLq6qo+HDh0KQdB1wxR48MEH1ccnT560aFxElqb7/3QiImoqmChARERERERERERERGQCDw8P9XFsbKzeuTExMerjnJwci8VEZBV1k2J0JciIrChARGTPmChARERERERERERERGSCiIgIg+c6OTmpj1UqlSXCIbIaXa0H6ucGMFGAiMieMVGAiIiIiIiIiIiIiMgE7dq1Ux9nZGTonZuXl6c+Dg4OtlhMRFbB3gNERE2eg60DICIiIiIiIiIioqarICkTZ9fvQ865JFSXVsDZ0w3BXWPRbex98IsNs3V4RBZ15513qo//+usvPPXUUzrnnj59Wn1cN8GAqCmqlyfA1gNERE0SEwWIiIiIiIiIiIjIaLkXk/Hnoh+QFp8AQSqBqLxZSj3j+CWc+GYHIvt0wYA3JyOoY7TtAiWyoK5duyIqKgopKSk4cuQIzp49i27dumnMKywsxM6dOwEArVu3Rvv27a0dKpFZ6Ww9UO88EwWIiOwZWw8QERERERERERGRUVIOn8O6x+Yh/egFAKiXJFD3dfrRC1j32DykHD5n9RiJrOWFF14AAIiiiFmzZiEzM7PeuEwmw2uvvYbi4mIAwJNPPmntEInMT2frASYHEBE1FawoQERERERERERERAbLvZiMLVOWQFEtb7CstKhUQamSY8uUJRi3aSErC1CzNGLECOzfvx979uxBSkoKhg0bhtGjR6NTp07Iz8/Hzz//jKSkJABAr169MHr0aBtHTNR4Qr1MgTrH9b4tMGmAiMieMVGAiIiIiIiIiIiIDPbnoh+glCsM7j0tiiKUcgUOLlqD0T/Ot3B0RLaxdOlSzJs3D1u2bEFpaSm+/fZbjTl9+/bFsmXLIOjq507UVAk6Wg8Y+H2CiIhsg4kCRERERHZIpVJh9+GLKC6tRK9u0YiNCLB1SEREREREKEjKRFp8gtHrRKUKqfHnUHg9C74xoRaIjMi2HB0dsWTJEowcORK//PILTpw4gfz8fPj4+KBt27YYM2YM7r//fkgk7AZMzQPzXYiImj4mChARERHZGZVKhTGv/hdHTl9Xn/v+/cm4764ONoyKiIiIiAg4u34fBKkEolJl9FpBKsGZdXsxYO5kC0RGZB/i4uIQFxdn6zCILE7Q8+omVhQgIrJnTF8kIiIisjPbDpyrlyQAAE/O/QHVMoWNIiIiIiIiqpFzLsmkJAGgpqpAbsL1hicSEVGTIupKFGDrASIiu8ZEASIiIiI789G3e7WeX/jV71aOhIiIiIiovurSisatLyk3UyRERGRLuloPMDWAiKjpYKIAERERkZ1JySzQev7n309YORIiIiIiovqcPd0at97L3UyREBGR3aibNSDWPWTaABGRPWOiABEREZEduVFQqnOsokpmxUiIiIiIiDQFd42FIDXtkqIglSCoS4yZIyIiIlsQ6rUb0FFegIiI7BoTBYiIiIjsSEJilq1DICIiIiLSqdvY+yAqVSatFZUqdB93v5kjIiIiWzCo9YDIigJERPaMiQJEREREdqSgmD1biYiIiMh++cWGIfSO9hAkxj09KkglaHV3V/jGhFooMiIish1d3xOYKEBEZM+YKEBERERkRyqr5LYOgYiIiIhIp6rSCrQdehcEqVT346S3EAQBUkcH9J87ycLRERGRtdT7FmDg9wMiIrIvTBQgIiIisiOVVTK94yLL9hERERGRjSjlCmScSYRHWAB6vjgSEgdpg5UFBKkEUmdHjFj1OoI6RlsnUCIispn6ly14DYOIyJ4xUYCIiIjIjlQ0kChQJVNYKRIiIiIioptEUUTW+euQVVQBAAI6RKHP6+Ph1y4SQE1CQF21ryPjOmPcpoWIururdQMmIiKLEnS+qpMcwIcdiIjsmoOtAyAiIiKimyoaaD1QVS2Hq7OjlaIhIiIiIqpRkJKN0tzCeue8IoLQe+bjcHCQIuOf88hNuI7qknI4e7kjqEsMuo+7H74xoTaKmIiILEkQhJr8ABG4NW2AiIiaBiYKEBEREdmRikr9FQWqWVGAiIiIiKysvKAEN66max1z9fFAVM8OaDvwDitHRUREtqbOE9CJFQWIiOwZWw8QERER2ZEGWw9U6684QERERERkTvIqGTLPXYOopXy0g5Mjwru1gSDhJUYiohZNuFlRoN53C7YeICKya/wUT0RERGQnqmRy5OSXNjiHiIiIiMgaRJUKmeeuQaElWVUQBIR1jYWji5MNIiMiInsgqBME2HqAiKgpYusBIiIiIhtLySzAA//5DOUNtB0AgKpqth4gIiIiIuu4kZiOikLtiawBbcLh7u9t5YiIiMieaE0PEHW+ICIiO8NEASIiIiIbSs7IR9+Jnxg8P6+wzILREBERERHVKM0pQH5yttYxj0Af+EeHWjkiIiKyXzpaDzBRgIjIrrH1ABEREZENLV6126j5yRn5FoqEiIiIiKiGrLwKWeevax1zcnNGWJfYOuWmiYiopeK3AiKipo2JAkREREQ2Ui1TYOfBBKPWLFi+E4XFFRaKiIiIiIhaOpVCifQziVAqlBpjEqkE4d3aQOrIIqVERFSHrowBkRUFiIjsGRMFiIiIiGwkLbvApHVf/vyXmSMhIiIiIgJEUUT2xWRUl2lPTA1uHwUXL3crR0VERPaqNj1AZOsBIqImiYkCRERERDZSWFxp0rptf5w1cyREREREREBR+g0UZ2lvdeUTHgifiEArR0RERPaMrQeIiJo2JgoQERER2UhhiWktBNJzilBcZlqSARERERGRNpXFZci5nKp1zMXTDcEdWlk5IiIisn/CLf9E/SICbD1ARGTXmChAREREZCOmJgoAQGLqDTNGQkREREQtmUImR8aZaxBVKo0xqaMDwru3gUQqtUFkRERkz9TpAfVKC4g6jomIyN4wUYCIiIjIRopKTU8UKC2rMmMkRERERNRSiaKIrIQkyKuqtY6Hdo6Bk5uLlaMiIqImQUvrAaYGEBE1HUwUICIiIrKRwmLTEwXKKrRfyCUiIiIiMkZ+UibK8oq1jvnHhMIzyNfKERERUVMhaDmqj2kDRET2jIkCRERERDbSmNYDpeVMFCAiIiKixinLK0ZeUqbWMXc/LwS2jrByRERE1JTc7DigI1FAZKIAEZE9Y6IAERERkQ0cO5eCn3YcM3l9eSUTBYiIiIjIdPLKamSeuwZRy00cB2cnhHVtDUGi6wlRIiIi7ZgbQETUdDBRgIiIiMjKdvx5Do/NXGnQXFcXR63nG1ONgIiIiIhaNpVShYyz16CUKzTGBEFAeLfWcHDW/jmUiIiollBbSUBg6wEioqaIiQJEREREVvb1hkNQqQz7YTmue4zW88vWHEDWDe2Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 2500x400 with 3 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAACAoAAAHeCAYAAAAYKXt1AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdd3hTZfsH8O9JuvegLRRKKXsVZBccIOACZINsUV94HQwHAiIgKojo+/oTRXHwosgUUJbIklEFKrL3hg5a2tI9s8/vj9rQmpM0SdMmab+f6+Iyfda5i0CT89znfgRRFEUQERERERERERERERERERFRrSCzdwBERERERERERERERERERERUfZgoQEREREREREREREREREREVIswUYCIiIiIiIiIiIiIiIiIiKgWYaIAERERERERERERERERERFRLcJEASIiIiIiIiIiIiIiIiIiolqEiQJERERERERERERERERERES1CBMFiIiIiIiIiIiIiIiIiIiIahEmChAREREREREREREREREREdUiTBQgIiIiIiIiIiIiIiIiIiKqRVzsHQARERERERERERERUU1w8+ZNrF+/HocPH0ZaWhoAICIiAo8++iieffZZBAUF2TlCIiIiohKCKIqivYMgIiIiIiIiIiIiInJm33//Pf7zn/9ArVZL9gcHB+PLL7/EAw88UL2BEREREUlgogARERERERERERERUSWsXr0aCxcuBAB4enpi+PDhiI6OhkKhwM6dO3Hs2DEAgL+/P3bu3ImQkBB7hktERETERAEiIiIiIiIiIiIiImvduXMH/fv3h0KhQFBQEFatWoXmzZuXG7Nw4UKsXr0aADBu3DjMmzfPHqESERER6cnsHQARERERERERERERkbP64osvoFAoAACffvqpQZIAAMycORNBQUEAgF27dlVrfERERERSXOwdABERERERERERERGRM1KpVNi7dy8AoHfv3ujWrZvkODc3N0yZMgXx8fEIDAyESqWCm5tbdYZKREREVA4TBYiIiIiIiIiIiIiIrBAXF4eCggIAwJAhQ0yOHTt2bHWERERERGQWJgoQEREREREREREREVnhypUr+tft27fXv87KysKtW7egVCoRGRmJBg0a2CM8IiIiIqOYKEBEREREREREREREZIXr168DKDlaICwsDImJifjwww8RGxsLjUajHxcdHY05c+agY8eO9gqViIiIqByZvQMgIiIiIiIiIiIiInJGaWlpAAB/f38cP34cgwYNwv79+8slCQDA+fPnMX78eOzcudMeYRIREREZYKKAkxo3bhzGjRtn7zCIiIiIiKod3wsTERERkaMoLCwEABQXF2PKlCkoKirC8OHD8csvv+D8+fP47bffMGnSJMhkMmg0GsyePRuXL1+2+np8L0xERES2wqMHnNTdu3ftHQIRERERkV3wvTAREREROYrSRIGCggIAwLRp0/DKK6/o+yMiIjBjxgw0aNAA77zzDlQqFT7++GOsXLnSquvxvTA5Gt3llRDzE8u1Cd7hkLWZbKeIiIjIXKwoQERERERERERERERUSc2bN8fLL78s2Tdq1Ci0b98eAHDkyBFu+FPNIRhuM4mizg6BEBGRpZgoQERERERERERERERkBU9PT/3r/v37QxAEo2OffPJJ/etTp05VaVxE1Ufqz7xY7VEQEZHlmChARERERERERERERGQFHx8f/evGjRubHBsVFaV/nZaWVmUxEVUriYoCYEUBIiKnwEQBIiIiIiIiIiIiIiIrNGjQwOyxbm5u+tc6HTdSqaZgRQEiImfFRAEiIiIiIiIiIiIiIis0b95c/zo5Odnk2IyMDP3rsLCwKouJqFqxogARkdNysXcARERERERERERE5LyybqXg3IbfkHb+FpT5RXD39UJYdGO0G9UXQY3D7R0eUZXq0qWL/vXvv/+O5557zujYM2fO6F+XTTAgcmqCREUBkRUFiIicARMFiIiIiIiIiIiIyGLpl+NxaNEPSDp6AYJcBlF7/wnS5BNXcHLFL4jo0Ra93p6A0FaN7BcoURWKjo5GZGQkEhISEBcXh3PnzqFdu3YG47Kzs7Fz504AQJMmTdCiRYvqDpWoikgVrmZFASIiZ8CjB4iIiIiIiIiIiMgiCUfOY/2wubhz7BIAlEsSKPv1nWOXsH7YXCQcOV/tMRJVl5dffhkAIIoiZsyYgZSUlHL9KpUKM2fORG5uLgBg4sSJ1R0iUdVhRQEiIqfFigJERERERERERERktvTL8dg6aQk0SnWFm0GiVgetTo2tk5Zg9E8LWVmAaqTBgwdj//792Lt3LxISEjBw4ECMGDECrVu3RmZmJn788UfcunULANC1a1eMGDHCzhET2ZAg8TyqyIoCRETOgIkCREREREREREREZLZDi36AVq0x+4lRURShVWsQu2g1RqyZV8XREdnHJ598grlz52Lr1q3Iz8/HypUrDcY89NBDWLp0KQSpJ7CJnJbUn2dWFCAicgY1LlEgNzcXGzZswMGDB3H79m0UFhbC19cXLVq0wJNPPomhQ4fCzc3N6HxRFPHLL7/gp59+wuXLl1FUVISQkBB06dIFY8eOlTxfqirWICIioppBp9NBqdLgl0MXAACPPdgKAb6edo6KiIiIiMg6WbdSkHT0gsXzRK0OiUfPI/v2XQRG1auCyIjsy9XVFUuWLMGQIUOwadMmnDx5EpmZmQgICECzZs3wzDPP4LHHHoNMxtOAqYZhRQEiIqdVoxIF4uLi8PrrryMrK6tce1ZWFuLi4hAXF4c1a9Zg+fLliIiIMJivUCgwffp0HDp0qFx7cnIykpOTsWPHDrz66quYPHmy0RhssQYRERE5v7TMPLz+4U+IPXG9XLuLXIbfVk5D04ahdoqMbK2mJKo6yhpERETk2M5t+A2CXAZRa/kmkCCX4ez6feg1Z0IVREbkGGJiYhATE2PvMIiqESsKEBE5K0EUzawR5uCuXLmCUaNGobi4GEBJGac+ffogICAAKSkp2Lp1K65fL7lRHxkZic2bN8PPz6/cGm+88QZ++eUXAECTJk0wcuRI1KlTBxcvXsSGDRtQVFQEAPjoo48waNAgyThssYY5+vTpAwDYv3+/1WsQERFR1Rky9Wscv5Ag2Tfu6a748PXB1RsQVQljiaplNWvWzOJE1VJyudzqRFVnXMNcfC9MRERkPz+OXoA7xy5ZPT8ipg1GrnvHhhER1S58L0yORnd7O8R7p8o3yt0g7zTHPgEREZHZakyiwLhx43D8+HEAwIIFCzB69Ohy/RqNBrNnz8aOHTsAAM899xxmz56t7z9y5Aief/55ACVZn9988w3c3d31/Tdv3sSYMWOQk5ODgIAA7N+/Hz4+PuWuYYs1zMU3hERERI4rKTUb3Ud/bHLMnYMfVFM0VFVqUqKqo6xhLr4XJiIisp8fBszEvUvxVs8Pbd0I43/5yHYBEdUyfC9MjkYXvwNi+snyjTJXyDu/bZ+AiIjIbDXiQKSbN2/qkwT69u1rkCQAAC4uLli0aBFCQ0vK/P7888/QarX6/pUrV+rHLVy4sNwGP1Bys3PevHkAgJycHGzatMngGrZYg4iIiJzf8fPSlQSoZlm4cKE+SWDBggX43//+hzFjxqBfv37417/+ha1bt+Lpp58GACQkJODLL78sN//IkSP6jfWYmBhs2bIFEydOxIABAzBr1ixs3rwZAQEBAIAPPvgABQUFBjHUpDWIiIjIObj7elVuvp+3jSIhIiLHILXNZPnxNEREVP1qRKJAXFyc/rWpp5Pc3d3x6KOPAig5SzY+Ph5Ayab90aNHAQAPP/ywZFlYAOjXrx+Cg4MBALt37y7XZ4s1iIiIqGbILSiucIymTMIiOZ+alKjqKGsQERGRcwiLbgxBbt0tRUEuQ2jbKBtHREREdiUIhm01o5A1EVGNVyMSBWQyGZo1awYfHx80atTI5Fh/f3/967y8PADAiRMnoNOVZLjFxMSYvE6XLl0AAGfPnkVubq6+zxZrEBERUc1w7mpyhWNy8ipOJiDHVVMSVR1lDSIiInIe7Ub1hai17klRUatD+9GP2TgiIiKyK0Fim0lkRQEiImdQIxIFxowZg19++QUnT55E8+bNTY69ceOG/nVp+dPSs2MBVDi/adOmAABRFHHt2jV9uy3WICIiIudRrFTj198v4H8/HcWtpAx9u0arxaY9pyqcn5VbWJXhURWrKYmqjrIGEREROY+gxuGI6NHW4qoCglyGhg9GIzCqXhVFRkRE9iFRUQAl+x9EROTYakSigLnS0tLwxx9/AAACAwMRGRkJAEhOvv/UX/369U2uUbduXf3rsvNssQYRERE5h2KFCuNmfofJ76zDO8t+Qd8XlmJ/3BUAwOGTN81aI69AUZUhUhWrKYmqjrIGEREROZdeb0+A3NVFuty0BEEQIHd1Qc8546s4MiIiqnZSFQUAVhUgInICtSpRYMmSJVCr1QCA/v37QyYr+fazsrL0YwIDA02uUXpzFygps1rKFmsQERGRc9h79DKOnYvXf61Sa7Hkf/sAAJdvpZq1hlqjrXgQOT1HT1R1lDWIiIjIuYS2aoS+H0yGzEUOQWY6WUCQyyB3d8Xgb2chtFWj6gmQiIiqkbGfA6woQETk6GpNosCGDRuwc+dOAICXlxcmT56s71Mo7j/R5+7ubnIdNzc3yXm2WIOIiIicw4JlOw3aLt28C5Vag9t3Ms1aQ2Plua7kXBw9UdVR1iAiIiLnIooiPOoEoMesMQhqHgEABgkDpUcTRMS0weifFiLywehqj5OIiKoBKwoQETktF3sHUB1+++03vPfee/qvFyxYgLCwMP3XGo1G/7rsJr6Usv1l59liDSIiInIO97ILJNtVai2SUrPNWqOiRAGtVoc9Ry4hJ78YXaMj0bRhqMVxkn05Q6Kqo6xBREREzqXgXg5URQr4NQhFt1dHojAtG4l/nEV+SgYEAB5+3ghtG4X2ox9DYFQ9e4dLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 2500x400 with 3 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAACAoAAAHeCAYAAAAYKXt1AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdd3hU1dYG8PfMZCaZ9B4SEkLoSFdAFBAEK72qSBFU1KtYLiIgAiIXUKzgtdyr6KUoIEUQRJQmKEUEFEjoEALpIb1Nn/P9kS9DhpkkM5NpSd7f89znO9ln730WfkpmzllnLUEURRFERERERERERERERERERETUKEjcHQARERERERERERERERERERG5DhMFiIiIiIiIiIiIiIiIiIiIGhEmChARERERERERERERERERETUiTBQgIiIiIiIiIiIiIiIiIiJqRJgoQERERERERERERERERERE1IgwUYCIiIiIiIiIiIiIiIiIiKgRYaIAERERERERERERERERERFRI8JEASIiIiIiIiIiIiIiIiIiokaEiQJERERERERERERERERERESNiJe7AyAiIiIiIiIiIiIiagiuXLmCdevW4eDBg8jOzgYAxMXF4d5778UTTzyB0NBQN0dIREREVEEQRVF0dxBERERERERERERERPXZypUr8f7770Or1Vo8HxYWhs8++wxdu3Z1bWBEREREFjBRgIiIiIiIiIiIiIioDtasWYNFixYBABQKBcaMGYNOnTpBpVJhx44dOHr0KAAgKCgIO3bsQEREhDvDJSIiImKiABERERERERERERGRvdLS0jB48GCoVCqEhoZi1apVaNOmjcmcRYsWYc2aNQCACRMmYN68ee4IlYiIiMhI4u4AiIiIiIiIiIiIiIjqq08//RQqlQoAsGzZMrMkAQCYOXMmQkNDAQA7d+50aXxERERElni5OwAiIiIiIiIiIiIiovpIo9Fg165dAIABAwbgzjvvtDhPLpdj2rRpSElJQUhICDQaDeRyuStDJSIiIjLBRAEiIiIiIiIiIiIiIjscOXIEpaWlAICRI0fWOHf8+PGuCImIiIjIKkwUICIiIiIiIiIiIiKyw/nz543HXbp0MR7n5+cjOTkZarUa8fHxiI2NdUd4RERERNViogARERERERERERERkR0uXboEoKK1QFRUFK5fv4533nkHBw4cgE6nM87r1KkT5syZg9tvv91doRIRERGZkLg7ACIiIiIiIiIiIiKi+ig7OxsAEBQUhGPHjmH48OHYu3evSZIAACQmJmLixInYsWOHO8IkIiIiMsNEgXpqwoQJmDBhgrvDICIiIiJyOX4WJiIiIiJPUVZWBgBQKpWYNm0aysvLMWbMGPz4449ITEzEnj17MHXqVEgkEuh0OsyePRvnzp2z+3r8LExERESOwtYD9VRmZqa7QyAiIiIicgt+FiYiIiIiT1GZKFBaWgoAeOmll/DCCy8Yz8fFxWHGjBmIjY3Fm2++CY1Gg/feew9ff/21XdfjZ+HGQ5d2GTemdrd7fcSKE/Bq2tKBERERUUPDigJERERERERERERERHXUpk0bPP/88xbPPfbYY+jSpQsA4NChQ3zgT7Uq37kSkEjtWyyRovyn/zk0HiIianiYKEBEREREREREREREZAeFQmE8Hjx4MARBqHbuQw89ZDz+66+/nBoX1X/aS38DBr19iw16aC+fdGg8RETU8DBRgIiIiIiIiIiIiIjIDv7+/sbjFi1a1Dg3ISHBeJydne20mKhhMJQVu3U9ERE1fEwUICIiIiIiIiIiIiKyQ2xsrNVz5XK58dhgMDgjHGpAJH6Bbl1PREQNHxMFiIiIiIiIiIiIiIjs0KZNG+Nxenp6jXNzc3ONx1FRUU6LiRoGWetugERq32KJFLJWXR0aDxERNTxe7g6AiIiIiIiIiIiI6q/85AycXr8H2YnJUJeUwzvAF1GdWqDzY/chtEWMu8MjcqoePXoYj3/77TdMmTKl2rknT540HldNMCCyxPfhySj7/hP7Fhv08B1U/b+LREREABMFiIiIiIiIiIiIyA4551Kwf/FqpB5OgiCVQNTfLKWefvw8Tqz4EXF3d0T/NyYhsn1z9wVK5ESdOnVCfHw8rl27hiNHjuD06dPo3Lmz2byCggLs2LEDANCyZUu0bdvW1aFSPeMV2wryLvdAk3gIMOitXyiRQt65D7yatnRecERE1CCw9QARERERERERERHZ5NqhRKwbPRdpR88CgEmSQNWf046exbrRc3HtUKLLYyRyleeffx4AIIoiZsyYgYyMDJPzGo0GM2fORFFREQBg8uTJrg6R6qnAZxYDMjkgWPkoR5AAMjkCpy5ybmBERNQgMFGAiIiIiIiIiIiIrJZzLgVbpy6FTq01SxC4lag3QK/WYuvUpcg5l+KaAIlcbMSIEXjggQcAANeuXcOwYcOwdOlSbN++HStXrsTw4cPx22+/AQB69uyJsWPHujNcqkdkLToh9M31gNwbkEhrniyRAnJvhL65HrIWnVwTIBER1WtMFCAiIiIiIiIiIiKr7V+8GnqtDhBFq+aLogi9VocDi9c4OTIi9/nwww8xYsQIAEBJSQm+/vprzJgxA2+//TaSk5MBAH369MHnn38OQRDcGCnVN97d+iH8w12Qd+5TMXDrvz///7O8cx+Ef7gL3t36uThCIiKqr7zcHQARERERmSspU2HPkfM4eyULQ/t3Que2Td0dEhERERER8pMzkHo4yeZ1ot6A64cTUXA1EyEJ0U6IjMi9ZDIZli5dipEjR2Ljxo04ceIE8vLyEBwcjNatW+PRRx/F/fffD4mE7+6R7WQtOiHs7R+gS7+C0tUvQ3c9BaJKDcHHG14xkVD0HQjv/m+4O0wiIqpnmChARERE5GHyi8ow9pUvcSElBwDw+frf0DQqGL+vmQ65jB/fiIiIiMh9Tq/fA0EqqbXlgCWCVIJT63aj/5xJToiMyDP06tULvXr1cncY1EB5NW2JwLFPQMw/Y3pCIkIUDRAEJqIQEZH1+FuDiIiIyMOs/+mEMUmgUnp2Idb/dNxNERERERERVchOTLYrSQCoqCqQk3TVwRERETUyikjzMYMOUOW7PhYiIqrXmChARERE5GFOnLlmcfz3E1dcHAkRERERkSl1SXnd1heXOSgSIqLGSbCUKAAAyhzL40RERNVgogARERGRh8nKLbY4XlqudnEkRERERESmvAN867Y+0M9BkRARNVK+lhMFRCYKEBGRjZgoQERERORhbhSUWhz//cRlXL6eA1EUXRwREREREREgiiKCmzeBIBHsWi9IJYjsmODgqIiIGhnvEEAiMx9nogAREdmIiQJEREREHkQUReRWkygAAP2fWIZZH2yFwWBfX1giIiIiInvodXpkJCYj9LbmEA32Ja6KegO6jLvfwZERETUugiCBoIgwG2dFASIishUTBYiIiIg8SF5hGTRafY1z1u44hiMnr7ooIiIiIiJq7FTFZUj54wyKs/LgHxWKsLbNbK4qIEglaNa7E0ISop0UJRFRI6Kw0H5AlQfRoHN9LEREVG8xUYCIiIjIg6Sk51k179N1B5wcCRERERE1dqIoouB6NlL+PAdNuco43n5sfwhSKSBYlywgCAKkMi/0mzPRWaESETUulhIFRBFQ5bo+FiIiqreYKEBERETkQa6mWZco8Nvxy06OhIiIiIgaM71Wh/RTl5F1/hrEW9peBcZGovsLIyHxktZaWUCQSiD1lmHEl7MQ2b65EyMmImo8BF8LiQIAUJ7t2kCIiKheY6IAERERkQfJKShxdwhERERE1MgpC0tx9Y8zKMkpqHZOVKeWGPn1bMT16gigIiGgqsqf43p1wLjNixDfu5PzAiYiamwsVRQAICpzXBwIERHVZ17uDoCIiIiIblJr2E+QiIiIiNxDFEXkX8vCjUtpEEWx2nmKYH807dQSMoU3mvfujIKrmTi1bjdykq5CXVwG70A/RHZMQJdx9yMkIdqFfwIiokZCFgB4+QA6lek4EwWIiMgGTBQgIiIi8iC2JAqIogjByr6wREREREQ10Wm0yDxzFaU3CmucF5YQjYiWTSFIblYQCEmIRv85k5wcIRERVRIEAYIiEmLJdZNxVhQgIiJbMFGAiIiIyIPYkiiwefdJjHmgmxOjISIiImoc8pMzcHr9HmQnJkNdUg7vAF9EdWqBzo/dh9AWMe4Oz+nKC0qQkXgFWpWm2jlechmiO7aAf3iQCyMjIqJqKaKAWxIFoC6CqFNB8PJxT0xERFSvMFGAiIiIyIOoNVqr5367/U8mChARERHVQc65FOxfvBqph5MgSCUQ9QbjufTj53FixY+Iu7sj+r8xCZHtm7svUCcRRRF5KZnIvZxeY6sB39BAxHRsAZmP3IXRERFRjRSRlseVN4CAONfGQkRE9ZKk9ilERERE5Cq2VBQ4lnTNiZEQERERNWzXDiVi3ei5SDt6FgBMkgSq/px29CzWjZ6La4cSXR6jM+k0WqT+dRE3LqVVmyQgCALCWzRFs9vbMkmAiMjDCL7VJQqw/QAREVmHiQJEREREHsSWRAEiIiIisk/OuRRsnboUOrXWLEHgVqLeAL1ai61TlyLnXIprAnSysvxiXD1yBmV5RdXO8fKWIe72Noho1RSCRHBhdEREZJVqKgqITBQgIiIrMVGAiIiIyIOomChARERE5HT7F6+GXqsDaii3X5UoitBrdTiweI2TI3Mu0SDixpV0pJ64AJ1aU+08v7BAJPTqCL+wIBdGR0REthC8FIA8wPyEMtv1wRARUb3k5e4AiIiIiOgmVhQgIiIicq785AykHk6yeZ2oN+D64UQUXM1ESEK0EyJzLq1Kg8ykZJTlF1c7RxAEhLdsirDm0awiQERUDwiKSIiaEpMxUZkDURQhCPx7nIiIasaKAkREREQeRK3R2jRfZeN8IiIiosbu9Po9EKT23RITpBKcWrfbwRE5X1leEVL+OFNjkoCXtxzLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 2500x400 with 3 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"spdrs = [\"XLF\", \"XLE\", \"XLRE\"]\n",
|
||||
"tckr_spdrs = []\n",
|
||||
"tckrs = []\n",
|
||||
"for spdr in spdrs:\n",
|
||||
" spdr_dat = pd.read_pickle(\"../../spdr-data/\" + spdr + \".pkl\")\n",
|
||||
" syms = list(spdr_dat.symbol.unique())\n",
|
||||
" tckrs += syms\n",
|
||||
" tckr_spdrs += [spdr for _ in range(len(syms))]\n",
|
||||
"\n",
|
||||
"for tckr, SPDR in zip(tckrs, tckr_spdrs):\n",
|
||||
" spdr_dat = pd.read_pickle(\"../../spdr-data/\" + SPDR + \".pkl\") \n",
|
||||
" data = spdr_dat[spdr_dat[\"symbol\"] == tckr]\n",
|
||||
" T = 5.\n",
|
||||
" ts = torch.linspace(0, T, data.shape[0])\n",
|
||||
" y = torch.FloatTensor(data['close_price'].to_numpy())\n",
|
||||
"\n",
|
||||
" eval_times = [100, 200, 300, 400, 500, 600, 700, 800, 900, 1000, 1100, 1200]\n",
|
||||
"\n",
|
||||
" prices_at_time_y = y[torch.tensor(eval_times)]\n",
|
||||
" delta_y = prices_at_time_y[1:] - prices_at_time_y[:-1]\n",
|
||||
"\n",
|
||||
" ## Load Price Probabilities\n",
|
||||
" if os.path.exists(\"./outputs/matern_\" + tckr + \".pt\"):\n",
|
||||
" voltron = torch.load(\"./outputs/voltron_\" + tckr + \".pt\")\n",
|
||||
" matern = torch.load(\"./outputs/matern_\" + tckr + \".pt\")\n",
|
||||
" specmix = torch.load(\"./outputs/sm_\" + tckr + \".pt\")\n",
|
||||
"\n",
|
||||
" bought_func = lambda xs: betainc(17, 8, xs)\n",
|
||||
" held_voltron = 1000 * bought_func(voltron)\n",
|
||||
" held_matern = 1000 * bought_func(matern)\n",
|
||||
" held_specmix = 1000 * bought_func(specmix)\n",
|
||||
"\n",
|
||||
" base = 10000\n",
|
||||
"\n",
|
||||
" prices_at_time_y = torch.tensor([y[0], *prices_at_time_y])\n",
|
||||
"\n",
|
||||
" plt_times = torch.tensor(eval_times) / 252\n",
|
||||
" hodl_strat = 10000 / y[0] * prices_at_time_y[1:]\n",
|
||||
"\n",
|
||||
" fig, ax = plt.subplots(1, 3, figsize = (25, 4))\n",
|
||||
" ax[0].plot(ts, y)\n",
|
||||
" [ax[i].set_xlabel(\"Time\") for i in range(3)]\n",
|
||||
" ax[0].set_ylabel(tckr)\n",
|
||||
" [ax[0].axvline(x=eval_times[i], alpha = 0.2, linestyle=\"--\") for i in range(len(eval_times))]\n",
|
||||
"\n",
|
||||
" # ax[1].plot(plt_times, matern, label = \"Matern\", color=palette[1], alpha = 0.5)\n",
|
||||
" # ax[1].plot(plt_times, specmix, label = \"Spectral Mixture\", color=palette[3], alpha = 0.5)\n",
|
||||
" # ax[1].plot(plt_times, voltron, label = \"Voltron\", color = palette[-1], alpha = 0.5)\n",
|
||||
" # ax[1].scatter(plt_times, matern, s = 120, label = \"Matern\", color=palette[0], zorder=4)\n",
|
||||
" # ax[1].scatter(plt_times, specmix, s = 120, label = \"Spectral Mixture\", color=palette[2], zorder=4)\n",
|
||||
" # ax[1].scatter(plt_times, voltron, s = 120, label = \"Voltron\", color = palette[-2], zorder=4)\n",
|
||||
"\n",
|
||||
" ax[1].plot(plt_times, value_func(matern), label = \"Matern\", color=palette[1], alpha = 0.5)\n",
|
||||
" ax[1].plot(plt_times, value_func(specmix), label = \"SM\", color=palette[3], alpha = 0.5)\n",
|
||||
" ax[1].plot(plt_times, hodl_strat, label = \"Hold\", color=palette[5], alpha = 0.5)\n",
|
||||
" ax[1].plot(plt_times, value_func(voltron), label = \"Voltron\", color = palette[-1], alpha = 0.5)\n",
|
||||
" ax[1].scatter(plt_times, value_func(matern), color=palette[0], zorder=4, s=120)\n",
|
||||
" ax[1].scatter(plt_times, value_func(specmix), color=palette[2], zorder=4, s=120)\n",
|
||||
" ax[1].scatter(plt_times, hodl_strat,color=palette[4], zorder=4, s=120)\n",
|
||||
" ax[1].scatter(plt_times, value_func(voltron),color = palette[-2], zorder=4, s=120)\n",
|
||||
"\n",
|
||||
" ax[2].plot(plt_times, running_sharpe_ratio(value_func(matern)), \n",
|
||||
" label = \"Matern\", color=palette[1], alpha = 0.5)\n",
|
||||
" ax[2].plot(plt_times, running_sharpe_ratio(value_func(specmix)), \n",
|
||||
" label = \"SM\", color=palette[3], alpha = 0.5)\n",
|
||||
" ax[2].plot(plt_times, running_sharpe_ratio(hodl_strat), label = \"Hold\", markersize = 20,\n",
|
||||
" color=palette[5], alpha = 0.5)\n",
|
||||
" ax[2].plot(plt_times, running_sharpe_ratio(value_func(voltron)), \n",
|
||||
" label = \"Volt\", color = palette[-1], alpha = 0.5)\n",
|
||||
" ax[2].scatter(plt_times, running_sharpe_ratio(value_func(matern)), \n",
|
||||
" s = 120, color=palette[0], zorder=4)\n",
|
||||
" ax[2].scatter(plt_times, running_sharpe_ratio(value_func(specmix)), \n",
|
||||
" s = 120, color=palette[2], zorder=4)\n",
|
||||
" ax[2].scatter(plt_times, running_sharpe_ratio(hodl_strat),\n",
|
||||
" s = 120, color = palette[4], zorder=4)\n",
|
||||
" ax[2].scatter(plt_times, running_sharpe_ratio(value_func(voltron)), \n",
|
||||
" s = 120, color = palette[-2], zorder=4)\n",
|
||||
"\n",
|
||||
"\n",
|
||||
" # ax[1].set_ylabel(\"P(increase)\")\n",
|
||||
" ax[1].set_ylabel(\"Portfolio Value\")\n",
|
||||
" ax[2].set_ylabel(\"Sharpe Ratio\")\n",
|
||||
"\n",
|
||||
" ax[2].legend(ncol = 4, loc = \"lower center\", bbox_to_anchor = (-0.85, -0.5))\n",
|
||||
" plt.subplots_adjust(wspace=0.35)\n",
|
||||
" sns.despine()\n",
|
||||
"# if tckr in [\"JPM\", \"BAC\"]:\n",
|
||||
" ax[2].set_ylim(0, 6)\n",
|
||||
" [ax[i].set_xlim((-0.1, 5.1)) for i in range(3)]\n",
|
||||
"# plt.savefig(\"trading_strategy_\" + tckr + \".pdf\", bbox_inches = \"tight\")\n",
|
||||
" plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "b4652987",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 246,
|
||||
"id": "a2f07158",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"tensor([10000.0000, 9132.0801, 9147.9814, 8791.2256, 8932.3867, 9046.4463,\n",
|
||||
" 10762.1602, 10762.1602, 10762.1602, 10762.1602, 10762.1602, 10762.1602])"
|
||||
]
|
||||
},
|
||||
"execution_count": 246,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"value_func(matern)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 239,
|
||||
"id": "c67195d4",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"'XLE'"
|
||||
]
|
||||
},
|
||||
"execution_count": 239,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"SPDR"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 278,
|
||||
"id": "98415c19",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"tckr = \"EOG\"\n",
|
||||
"SPDR = \"XLE\"\n",
|
||||
"spdr_dat = pd.read_pickle(\"../../spdr-data/\" + SPDR + \".pkl\") \n",
|
||||
"data = spdr_dat[spdr_dat[\"symbol\"] == tckr]\n",
|
||||
"T = 5.\n",
|
||||
"ts = torch.linspace(0, T, data.shape[0])\n",
|
||||
"y = torch.FloatTensor(data['close_price'].to_numpy())\n",
|
||||
"\n",
|
||||
"eval_times = [100, 200, 300, 400, 500, 600, 700, 800, 900, 1000, 1100, 1200]\n",
|
||||
"\n",
|
||||
"prices_at_time_y = y[torch.tensor(eval_times)]\n",
|
||||
"delta_y = prices_at_time_y[1:] - prices_at_time_y[:-1]\n",
|
||||
"\n",
|
||||
"## Load Price Probabilities\n",
|
||||
"\n",
|
||||
"voltron = torch.load(\"./outputs/voltron_\" + tckr + \".pt\")\n",
|
||||
"matern = torch.load(\"./outputs/matern_\" + tckr + \".pt\")\n",
|
||||
"specmix = torch.load(\"./outputs/sm_\" + tckr + \".pt\")\n",
|
||||
"\n",
|
||||
"bought_func = lambda xs: betainc(17, 8, xs)\n",
|
||||
"held_voltron = 1000 * bought_func(voltron)\n",
|
||||
"held_matern = 1000 * bought_func(matern)\n",
|
||||
"held_specmix = 1000 * bought_func(specmix)\n",
|
||||
"\n",
|
||||
"base = 10000\n",
|
||||
"\n",
|
||||
"prices_at_time_y = torch.tensor([y[0], *prices_at_time_y])\n",
|
||||
"\n",
|
||||
"plt_times = torch.tensor(eval_times) / 252\n",
|
||||
"hodl_strat = 10000 / y[0] * prices_at_time_y[1:]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 280,
|
||||
"id": "f0e2b903",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAACAoAAAGWCAYAAAD7IvN0AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdd3xT9foH8M85WW3TpG26F7Sl7C1eRUTAiXtcRUFciPhTL6hX8aqIigp6udd5FfSCewCKiHAdCDiqUkR2W3bp3iNJ0+xxzu+P0LRpRpM0bdP0eb9e3CZnfM9TrrQn5/t8n4fheZ4HIYQQQgghhBBCCCGEEEIIIYQQQgYEtq8DIIQQQgghhBBCCCGEEEIIIYQQQkjvoUQBQgghhBBCCCGEEEIIIYQQQgghZAChRAFCCCGEEEIIIYQQQgghhBBCCCFkAKFEAUIIIYQQQgghhBBCCCGEEEIIIWQAoUQBQgghhBBCCCGEEEIIIYQQQgghZAChRAFCCCGEEEIIIYQQQgghhBBCCCFkAKFEAUIIIYQQQgghhBBCCCGEEEIIIWQAoUSBfuq2227Dbbfd1tdhEEIIIYQQ0uvoXpgQQgghhAxUdC9MCCGEkGAR9nUAJDC1tbV9HQIhhBBCCCF9gu6FCSGEEELIQEX3woQQQggJFqooQAghhBBCCCGEEEIIIYQQQgghhAwglChACCGEEEIIIYQQQgghhBBCCCGEDCCUKEAIIYQQQgghhBBCCCGEEEIIIYQMIJQoQAghhBBCCCGEEEIIIYQQQgghhAwglChACCGEEEIIIYQQQgghhBBCCCGEDCDCvg6AEEIIIYQQQgghvUNZUoOCDTtRX1gCU6seElkUksfmYNzsS6DISevr8AghhBBCSD9G95qEENK/UKIAIYQQQgghhBAS5hqOleGXFR+jMr8IjIAFb+Mc+6r3Hcf+d79B5pQxmPHUHUgamdV3gRJCCCGEkH6H7jUJIaR/otYDhBBCCCGEEEJIGCvfVYj1Ny5F1Z6jAOD04Lbj+6o9R7H+xqUo31XY6zESQgghhJD+ie41CSGk/6JEAUIIIYQQQgghJEw1HCvD1wtWwmqyuDy07Yy3cbCZLPh6wUo0HCvrnQAJIYQQQki/RfeahBDSv1HrAUIIIYQQQgghJEhCrS/rLys+hs1iBXjep+N5nofNYkXeik8w69Onezg6QgghhBDSn9G9JiGE9G+UKEAIIWcYzRbszD+Oqno1pv9lKEbmpPR1SIQQQgghpJ8Ixb6sypIaVOYX+X0eb+NQkV8IVWkt4rJTeyAyQgghhBDS39G9JiGE9H+UKEAIIbAnCdz++IfYfagUAPDSmh+w6ulbcPWMsX0cGSGEEEIICXXluwrx9YKV9tVU6Lov6/VrH8fg8wO7z+Q5DlazFTaLFTazBVazBTazFTaz1fHaarF/Lfj0BzAsA57zbYVXR4yAxeH1OzBjyR0BxUkIIYQQQsJbwYadLgmyvqJ7TUIICQ2UKEAIIQB+3nPSkSQAADaOw7/e30GJAoQQQgghxKuOfVm7KrnK2zjYOHtf1jmbliNpZJabiX/7V/uEf6dkAIvVkYzgC3VpbUBJAm2xNhSVdn0gIYQQQggZkOoLSwJKEgDoXpMQQkIFJQoQQgiA5e9877KtpLIJ6lYDYmWRfRARIYQQQgjpDwLqy2q2YNs/VmPyI7P9mvj3l9Vo7tb5Jo0uSJEQQgghhJBwY2rVd+98utckhJA+x/Z1AIQQEgrKa5Rut9tstl6OhBBCCCGE9BdtfVn9XUnFczwaj5RBU9XYQ5HZCSPE3TpfIpcGKRJCCCGEEBJuJLKo7p1P95qEENLnKFGAEDLgWb0kA5gslChACCGEEELca+vLGgiGZVDx2+EgR+QsZnAKGJYJ6FxGwCJpTHaQIyKEEEIIIeEieWxO4PfCdK9JCCEhgRIFCCED3vubdnvcZzb3XClYQgghhBDSv3WrLyvHo6WiPqjxMCwLUYQYEXIpohNiMOqv08BzvrVEcInPxmH8nEuDGh8hhBBCCAkf42ZfEvi9MN1rEkJISBD2dQCEEBKo0xWNePmDnVC3GnDRucNw91+nQOBnFivP83j+7e887jf3YM9YQgghhBDSv3W3L6vVYPK6n2FZCMVCCERCCMQi+2uxCAKxEEKR/atALIRQLIJALAIrYMEwzhUECqaMQdWeo349xGUELDInj0ZcdmpA3xchhBBCCAl/ipw0ZE4Zg6o/jviVnEr3moQQEjooUYAQ0i+pNHpc/cBqtOrsD1d/21+MJrUOTy6Y6dc4B49Ved1vptYDhBBCCCHEg+72ZY2IjYZicMqZif8OSQBiEQQiIVihwGXi318znroD629cChtnAc/78ACXYSAQCTF9ye3dui4hhBBCCAl/05+4Hetuegq8xQb4cK/J0L0mIYSEFGo9QAjpl9Z9u9eRJNDm4y1/+F0B4Nu8Qq/7TVRRgBBCCCGEeNDdvqyZ545C8vBBSMhOQ2xGEmRJcYiKlUEcFQGBSNjtJAEASBqZhevXPg6BRNRlrAzLgBUKMPPfDyBpZFa3r00IIYQQQsKbOEaKsx+4wZ7gynq/d2VYBgKJCNevfZzuNQkhJERQogAhpF969YMfXba16kyoqFX6Nc4vf570ut9spkQBQgghhBDiXn/pyzr4/LGYs2k5Ms8dDQAuCQNtD3UVwzIx5fFbIYqNBs8F9n0RQgghhJCBged4NJfVImHEYEx5/FYohmUCgEvCQMd7zVs+fw6Dzx/b67ESQghxj1oPEEL6HY3W6HGl/+mKJuQOSvJ5rMo6tdf91HqAEEIIIYR4IpZFIX74IChPVYZ8X9akkVmY9enTUJXW4vD6HWgoKoVRowXH8ZBnJGHQBeMhTY4DAJh1RqirGhE3KLnX4iOEEEIIIf2Lpr4ZFoO94qs8IwnnPnwzdPUq1Px5FK1VjdArNRBGShAzKNlxrxkRJ+/jqAkhhHREiQKEkH5ne/4xj/sq61Q+j2M0W6A3mr0e428rA0IIIYQQMjBYjGbUHinFyFkzkL9yHXi+f/RljctOxYwldzjeq6sbUXuk1OW4ppIayFPjIRDRYwNCCCGEEOKM53k0l9a6bJcmx+HCpXdCEh2F07sKXPZr6psRm5HYGyESQgjxQdh/4n/uueewbt06LFy4EIsWLfJ6rMFgwKZNm7Bjxw6cPHkSra2tkEqlyMnJwcUXX4w5c+ZAKpV6PL+6uhoXXXSRT3FlZ2dj27Ztfn0vhIQLg9EMi5WDPDoioPN/21fscV9XE/8dqVr0XR7jqXIBIYQQQggZuHieR+2REljNFsgzknD2327AvlWbwdtsXisLMAIWApEwpPqyxqQmQFVRD2Or872x1WxBc1ktkoZm9lFkhBBCCCEkVGkb1TBpDS7bJdGRiE6KA8MwiJBLYdTonPbrla2wmi0QikW9FSohhBAv2K4P6b92796NDRs2+HTs8ePHcc011+CFF17AH3/8AaVSCYvFArVajQMHDuDf//43rrzyShw5csTjGCdOnAhW6ISEJavNhidf+xojr34eY69bjr+9sAEms/8T8TWNao/79AY/EgU0XScKmM3UeoAQQgghhDhTltVB16xxvHfpyypw/qjd9j5z8mjM2bQ8pPqyMiyDpGHukwFUFfWOcrKEEEIIIYQAnqsJAEB8VioYhgEAyJMVbs9trVf2aHyEEEJ8F7YVBYqKirBw4UJwHNflsfX19Zg/fz6ampoAABMmTMCVV16JpKQkNDY24rvvvsPBgwdRV1eH+fPnY9OmTUhPT3cZp2OiwKuvvgqJROLxmt4qExASrtZ9sxefbP3T8X7LTwXISo/HY3df6vU8o9mCtV/swm8HipGaEIOSyiaPx/pTUUCp1nV5DLUeIIQQQgghHRnUWjQWV7lsl2ck4ZIX/w9RcikKNuxEQ1EpTBodJHIpksZkY/ycSxGXndoHEXdNGh+D6IQYaJtanLZzNg6Np6uRNianjyIjhBBCCCGhRq/UwNCiddkuipRAntKeHCBLUaDhVKXLcZo6JeIyk3s0RkIIIb4Jy0SBvLw8LF68GFqt6y8rd1599VVHksCCBQuwePFip/133HEHXn75ZaxduxYqlQr//Oc/8eabb7qM05YokJiYiKuuuqqb3wUh4ef1j3922bbum71dJgo8/Z//Yf23+3y6ht5o8TmeomL3ma8dUaIAIYQQQghpY7NYUVN0Gjzv2l5AFCFB6uhsCERCzFhyRx9E1z2JQzOha9a4fG8tNU1QDEpGhJyS3QkhhBBCCNBc5qmaQAoYtr2yljhSgsiYaJekAoNaC4vRDFGEuEfjJIQQ0rWwShQwm81455138Pbbb/tUSQAAWltb8e233wIARo8ejUcffdTtcY8++ijy8/Nx5MgR7Ny5EyqVCnFxcU7HtCUKDBs2rBvfBSHhq0HZ6rKtUeU9oUfZosPGbQd8vobej9Ko+49UdHmM2UKtBwghhPjvueeew7p167Bw4UIsWrTI67EGgwGbNm3Cjh07cPLkSbS2tkIqlSInJwcXX3wx5syZ47UaVXV1NS666CKf4srOzsa2bds87s/Ly8P69etRUFAAjUaD+Ph4jBs3DnPmzMGUKVN8ukYwxiAkFPE8j7rj5TDrXe83GYZB2tgcCET99yN2hCwKMWkJUFc3uuyrP1mJQZOGO8rIEkIIIYSQgcmg1jq14GojlIgQk5bgsl2eonBJFGhrP6AYnNJjcRJCCPFN/32K0Ul+fj6efvppVFXZS0BGRUXhlltuwQcffOD1vH379sFisa9Avuaaazw++GAYBjNnzsSRI0fAcRwKCwsxbdo0x36TyYTy8nIAlChAiDtaNw9UffHrvmJYbb4l/gDAaS9tCTqrd3NT25mJKgoQQgjx0+7du7Fhwwafjj1+/DgWLlyIykrncoxqtRoHDhzAgQMH8Mknn2D16tUYPXq02zE6tr8KFMdxeOaZZ7Bx40an7XV1dairq8P27dtx222Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 2500x400 with 3 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(1, 3, figsize = (25, 4))\n",
|
||||
"ax[0].plot(ts, y)\n",
|
||||
"[ax[i].set_xlabel(\"Time\") for i in range(3)]\n",
|
||||
"ax[0].set_ylabel(\"Price\")\n",
|
||||
"[ax[0].axvline(x=eval_times[i], alpha = 0.2, linestyle=\"--\") for i in range(len(eval_times))]\n",
|
||||
"\n",
|
||||
"# ax[1].plot(plt_times, matern, label = \"Matern\", color=palette[1], alpha = 0.5)\n",
|
||||
"# ax[1].plot(plt_times, specmix, label = \"Spectral Mixture\", color=palette[3], alpha = 0.5)\n",
|
||||
"# ax[1].plot(plt_times, voltron, label = \"Voltron\", color = palette[-1], alpha = 0.5)\n",
|
||||
"# ax[1].scatter(plt_times, matern, s = 120, label = \"Matern\", color=palette[0], zorder=4)\n",
|
||||
"# ax[1].scatter(plt_times, specmix, s = 120, label = \"Spectral Mixture\", color=palette[2], zorder=4)\n",
|
||||
"# ax[1].scatter(plt_times, voltron, s = 120, label = \"Voltron\", color = palette[-2], zorder=4)\n",
|
||||
"\n",
|
||||
"ax[1].plot(plt_times, value_func(matern), label = \"Matern\", color=palette[1], alpha = 0.5)\n",
|
||||
"ax[1].plot(plt_times, value_func(specmix), label = \"SM\", color=palette[3], alpha = 0.5)\n",
|
||||
"ax[1].plot(plt_times, hodl_strat, label = \"Hold\", color=palette[5], alpha = 0.5)\n",
|
||||
"ax[1].plot(plt_times, value_func(voltron), label = \"Volt\", color = palette[-1], alpha = 0.5)\n",
|
||||
"ax[1].scatter(plt_times, value_func(matern), color=palette[0], zorder=4, s=120)\n",
|
||||
"ax[1].scatter(plt_times, value_func(specmix), color=palette[2], zorder=4, s=120)\n",
|
||||
"ax[1].scatter(plt_times, hodl_strat,color=palette[4], zorder=4, s=120)\n",
|
||||
"ax[1].scatter(plt_times, value_func(voltron),color = palette[-2], zorder=4, s=120)\n",
|
||||
"\n",
|
||||
"ax[2].plot(plt_times, running_sharpe_ratio(value_func(matern)), \n",
|
||||
" label = \"Matern\", color=palette[1], alpha = 0.5)\n",
|
||||
"ax[2].plot(plt_times, running_sharpe_ratio(value_func(specmix)), \n",
|
||||
" label = \"SM\", color=palette[3], alpha = 0.5)\n",
|
||||
"ax[2].plot(plt_times, running_sharpe_ratio(hodl_strat), label = \"Hold\", markersize = 20,\n",
|
||||
" color=palette[5], alpha = 0.5)\n",
|
||||
"ax[2].plot(plt_times, running_sharpe_ratio(value_func(voltron)), \n",
|
||||
" label = \"Volt\", color = palette[-1], alpha = 0.5)\n",
|
||||
"ax[2].scatter(plt_times, running_sharpe_ratio(value_func(matern)), \n",
|
||||
" s = 120, color=palette[0], zorder=4)\n",
|
||||
"ax[2].scatter(plt_times, running_sharpe_ratio(value_func(specmix)), \n",
|
||||
" s = 120, color=palette[2], zorder=4)\n",
|
||||
"ax[2].scatter(plt_times, running_sharpe_ratio(hodl_strat),\n",
|
||||
" s = 120, color = palette[4], zorder=4)\n",
|
||||
"ax[2].scatter(plt_times, running_sharpe_ratio(value_func(voltron)), \n",
|
||||
" s = 120, color = palette[-2], zorder=4)\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"# ax[1].set_ylabel(\"P(increase)\")\n",
|
||||
"ax[1].set_ylabel(\"Portfolio Value\")\n",
|
||||
"ax[2].set_ylabel(\"Sharpe Ratio\")\n",
|
||||
"\n",
|
||||
"ax[2].legend(ncol = 1, loc = \"lower center\", bbox_to_anchor = (-1.05, -0.02))\n",
|
||||
"plt.subplots_adjust(wspace=0.35)\n",
|
||||
"sns.despine()\n",
|
||||
"if tckr in [\"JPM\", \"BAC\"]:\n",
|
||||
" ax[2].set_ylim(0, 6)\n",
|
||||
"[ax[i].set_xlim((-0.1, 5.1)) for i in range(3)]\n",
|
||||
"plt.savefig(\"trade_strat.pdf\", bbox_inches = \"tight\")\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 185,
|
||||
"id": "7579d3d1",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"SPDR = \"XLRE\"\n",
|
||||
"spdr_dat = pd.read_pickle(\"../../spdr-data/\" + SPDR + \".pkl\") "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 187,
|
||||
"id": "bd42e44c",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAmgAAAHUCAYAAACOD9TaAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAACgJ0lEQVR4nOzdd3hT5dsH8G/SvUvpgFIoUEbZu2wQykb2BkGGAmIZAiJDkI38FBVEVAREUIZlb4EyC2Xv3bI76N4zTfP+0beHpDknOUlOmnV/rsvL9MyHQ0nuPOO+RTKZTAZCCCGEEGI0xIZuACGEEEIIUUQBGiGEEEKIkaEAjRBCCCHEyFCARgghhBBiZChAI4QQQggxMhSgEUIIIYQYGQrQCCGEEEKMDAVohBBCCCFGxtrQDSDaad68OQoKCuDl5WXophBCCCGEp8TERNja2uLGjRsqj6MAzUTl5+dDKpUauhmEEEII0UBhYSH4FHGiAM1EeXt7AwDCwsIM3BJCCCGE8BUcHMzrOJqDRgghhBBiZChAI4QQQggxMhSgEUIIIYQYGQrQCCGEEEKMDAVohBBCCCFGhgI0QgghhBAjQwEaIYQQQoiRoQCNEEIIIcTIUIBGCCGEEGJkKEAjhBBCCDEyFKARQgghhBgZCtAIIYQQYpEkhVI8iorD7cdveRUwL0tULJ0AAGQyGSQSCYqKigzdFGLBxGIxbGxsIBKJDN0UQoiZS03PwYSF23Ht/mtm25WdX8KvQjkDtuo9CtAsXE5ODtLT05GZmQmpVGro5hACKysruLi4wM3NDY6OjoZuDiHETG3Zf1khOAOAViO+w88LhmJAl8aGaZQcCtAsWGZmJqKjo2FjYwN3d3c4OTlBLBZT7wUxCJlMhqKiImRnZyMjIwNpaWnw8/ODi4uLoZtGCDFD565Fsm6fsWoPurWtAycHuzJukSIK0CxUTk4OoqOj4erqCl9fXwrKiNFwcnKCl5cXYmNjER0dDX9/f+pJI4QI7vbjt6zbpUVFOHr+AYb2aFbGLVJEiwQsVHp6OmxsbCg4I0ZJJBLB19cXNjY2SE9PN3RzCCFmRt2CgKg3iWXUEm4UoFkgmUyGzMxMuLq6UnBGjJZIJIKrqysyMzONbnUVIcS0FUhUz7m2Ehs+PDJ8C0iZk0gkkEqlcHJyMnRTCFHJ0dERUqkUEonE0E0hhJiRQjWL4sRiw3deUIBmgUpSaYiN4BsCIapYWVkBAKV/IYQISlKo+j3F2srwn4+GbwExGBreJMaOfkcJIfogkRSq3G9FARohhBBCSNmSSFX3oBnDCJPhW0AIIYQQUoYKC9UtEjB87z0FaIQQQgixKBI1qzhtrK3KqCXcKEAjhBBCiEWRqFnFWahmCLQsUIBGSBkbPXo0ateujdq1ayM6OtrQzQEAzJ07l2nT1atXDd0cQgjRi7wCCeZ8vx9dxq9TeZy6IdCyQAEaIYQQQizC4vVHsePodbXHSShAI4QQQgjRP5lMhiPn7/M6lgI0QgghhJAykFdQiLSMXF7HJqRk6bk16lGARgghhBCzl5HFLzgDgJfRSXpsCT8UoBFCCCHE7GVk5fE+9smLd5DJZHpsjXrWBr07If+vqKgIJ0+exLFjx3D//n0kJSXB2toaHh4eaNiwITp16oTevXsztRkBYMaMGTh+/DgAICQkBFOnTlV5j4SEBHzwwQeQSqWoWbMmjhw5AgDYt28f5s2bBwDYs2cPGjRogPDwcOzevRv3799HcnIyPDw8ULduXYwePRpt2rRhrhkbG4tt27bh7NmzePfuHWxsbBAYGIghQ4agX79+vP7sOTk5+Ouvv3D8+HG8ffsW1tbWqFy5Mjp37oyRI0fCw8ND4+eXnJwMa2treHt7IygoCAMGDECTJk14tYcQQsxRRjb/AC09Kw8x8Wnwq1BOjy1SjQI0YnApKSmYMmUKbt++rbC9oKAAOTk5iI6OxrFjx7BhwwZs3LgRVapUAQD079+fCdCOHDmiNkA7duwYpP+f+6Zv376sx0ilUsybNw/79u1T2P7u3Tu8e/cOZ86cwaxZszBx4kScPXsWX375JTIzM5nj8vLycP36dVy/fh3Xrl3DihUrVLbp7du3GDt2LN6+fauw/eHDh3j48CG2bduG5cuXo1u3bpzXiIqKwqxZs/DkyROF7fn5+Xj58iVevnyJ3bt3o2fPnlixYgWcnJxUtokQQsyRJj1oAPDo+TsK0IhlmzlzJhOclStXDp07d0blypUhkUjw6tUrnDx5EhKJBC9fvsT48eNx7Ngx2Nraol27dvD09ERSUhJevXqFhw8fol69epz3OXToEIDiGmtcAdq3336L27dvQywWo3379qhfvz4yMzNx8eJFvHz5EgDw448/wtXVFStWrEBBQQGaNWuGli1bAgDOnz+Phw8fAijujevcuTOCg4M52zR9+nSkp6fDwcEBXbp0QfXq1ZGSkoKTJ08iPj4e6enp+OKLL7BhwwZ07NhR6fyoqCiMHDkS6enpAAAHBwd88MEHqFmzJgoKCnD79m0mr9nx48fx5s0b/PPPP3BwcFD5d0IIIeZGkzloAPDsVTy6ta2jp9aoRwEaMahbt24hIiICABAQEIAdO3bA3d1d4Zg3b95g5MiRSExMxNu3b3Hs2DH0798f1tbW+PDDD7F161YAxb1oXAHa8+fPmcApKCgIFSpUYD3u9u3b8PT0xK+//oqGDRsy23NzczFhwgTcvHkTRUVF+Oabb2BtbY0ff/wRvXr1Yo6bNm0a5s+fz/TA7d27V2WAlp6ejlq1auHXX3+Fn58fs33WrFmYP38+jh07hsLCQixYsAAnTpyAs7Mzc4xEIsHnn3/OBGdNmjTBTz/9pPRnu3btGqZNm4bU1FQ8fPgQy5Ytw8qVKznbRAgh5kjTHjRNhkT1gRYJEIO6e/cu83ro0KFKwRkAVKlSBTNnzgQAiEQi3L//Po9N//79mdfHjh3jnNR5+PBh5rW6uWHLli1TCM6A4p6pSZMmKWwbO3asQnBW0r6StgLAgwcPVN7Lzc0NmzZtUgjOSu73/fffMwFnYmIi9u/fr3DMvn378OrVKwBApUqVsGnTJtbAMygoCL///jusrYu/j+3fvx8vXrxQ2S5CCDE3fFNslAg9cUtPLeGHAjRiUPKT/u/cucN5XM+ePXH06FHcvXsXCxcuZLbXqVMHtWvXBlA8T+zGjRus55cEaPb29irnc1WpUgWdO3dm3RcYGKjw89ChQ1mP8/LygpubG4Di+XWqfPzxx/Dx8WHdZ2VlhSlTpjA/lwzRlihZ5AAUL5KQ710rrVGjRkwwWVRUpBTsEUKIuXuXlK7R8Ymphs2FRgEaMaigoCDm9fHjxzF+/HgcP34cGRkZCsc5ODigRo0asLOzU7qGfC+afNBS4tatW0zNyy5duqgMZEr3nMnz8vJiXjs6OsLf35/zWEdHRwDFw5CqdO/eXeX+Dh06MD1fjx49Qk5ODoDiBQAl8/ZEIpHKoLNEz549mdfXr6svdUIIIeYkNlGzAA0A8gsK9dASfihAIwYVGBioMOR46dIlzJgxA61atcKwYcOwfv163Lt3T2U+mj59+jA9cf/99x8KCxX/QWkyvOnt7c25Tyx+/8/F1dVV5XXkj+ViY2OD6tWrqzzG1tYWlStXBgAUFhYygWZ8fDwT/Pn5+akMOkvUqfN+squxFGknhJCyEpeYof6gUjINOA+NAjRicMuXL8f48eOZniKgON3FnTt38PPPP2PIkCHo2LEjvv/+e9YhQy8vL7Rt2xYAkJqaikuXLjH7JBIJk4pD/jgufFc3yg/NasvV1ZVXIFcyXAqA6VlMS0tj3a+K/Pw++fMJIcQSxCakaXyOIRcKUIBGDM7W1hZfffUVzpw5g3nz5iEoKAg2NjYKx8THx+OPP/5Az549WSfecw1zhoeHIzU1FQCUEt2yEYlEOvxJ9EO+99DW1lZpG18lOeAA4/xzEkKIvuTmS5CSnqPxeZqu/BQSBWjEaPj4+GDs2LHYvn07rl27hk2bNmH8+PGoVq0ac0xaWhqmT5+uEGwAxXPLXFxcAABhYWEoKCgAABw9epQ5hm9m/7KSlcVvAmpJGg3g/dCq/BAr396wkkC19PmEEGLu4rSYfwYA2bn5AreEPwrQiFFydHRE+/bt8dVXX+HEiRNYt24d06sWHR2NW7cUlz/b2dmhR48eAIDs7GxERERAKpXiwoULAICaNWuibt26ZfuHUCM/Px/v3r1TeUxOTg7evHkDoPiZlKTj8PX1ZZ5HTEyMQjUDLvKVBkrmtRFCiCWIS9AuQCssLBK4JfxRgEYMatWqVRg2bBhatGiB+Ph4zuO6d++O1q1bMz+zBTbyw5xnzpzBrVu3mN4nrsoBhnb58mWV+8PCwlBUVPwG0ahRI2aenp2dHRo0aACgeLjz5MmTau914sQJ5nXjxo21bDEhhJieWK0DNKn6g/SEAjRiUG/fvsWdO3eQkZGhsNqSjfwCAbbcYc2aNWN6hs6ePYuzZ88CUF3aydA2btzIpM4oLS8vD+vXr2d+HjJkiML+AQMGMK9/+eUXlUOm9+7dUwjQSifYJYQQc/bT9jNanSeRUg8asVCDBg1iXq9duxbnz59nPW7r1q3M4gBvb2/WHiCRSMTMM4uPj8fu3bsBAC1btuQs7WRoL1++xPTp05WCq7S0NEyZMoWpFFCnTh2lnGn9+/dH1apVARQPc37yySesPYs3btzA5MmTmfQj/fr1U5nvjRBCTNWLt0lYuuEYvttyCgvXHcbny3Zhx9HreB3LnTTcq5wzvp7ck3WfIXvQqBYnMajg4GC0b98eFy9eREFBASZOnIgmTZqgfv368PLyQnp6Oq5fv4579+4BKA7C5s6dy6xmLK1///5Mr1NJ0GNsiwNKODg4wMXFBRcuXEBwcDC6d++OihUrIi4uDv/99x8z+d/d3R1r1qxRSEMCFK/oXLt2LUaNGoWsrCzcvn0bPXr0QKdOnVCLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 640x480 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"sns.lineplot(x=\"date\", y=\"close_price\", data=spdr_dat[spdr_dat.symbol == \"PSA\"], hue=\"symbol\")\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "ccee2327",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.8"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,408 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "4932e338",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from voltron.robinhood_utils import GetStockData\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns\n",
|
||||
"import numpy as np\n",
|
||||
"import torch\n",
|
||||
"import pandas as pd\n",
|
||||
"\n",
|
||||
"import gpytorch\n",
|
||||
"from torch.nn.functional import softplus\n",
|
||||
"from voltron.kernels import BMKernel, VolatilityKernel\n",
|
||||
"from voltron.models import BMGP, VoltronGP"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "792ecce6",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import os\n",
|
||||
"import robin_stocks.robinhood as r"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "e9425d23",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"# Header"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"id": "ad653471",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ntest = 100\n",
|
||||
"ntrain = 500\n",
|
||||
"tckrs = ['TSLA', \"F\", \"JPM\", \"SBUX\", 'AAPL', \"VIRT\"]\n",
|
||||
"tckr = \"SPY\"\n",
|
||||
"span = \"3month\"\n",
|
||||
"interval = 'hour'\n",
|
||||
"T = 5."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "3ad86623",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Data Setup"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"id": "e09fb2ab",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAicAAAGzCAYAAAD0T7cVAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAB9lklEQVR4nO3dd3xb9dU/8M/VsOQl7xEnjrOnE0ImYYSREAgphUJpC6FQnpZR0gKhtDTPQ9kl+RVKKSX0CbSEPoyG0UKBQokZSYBAMElMFtnLjuM95Kn9++PqXl1JV7Jky1r+vF8vv7ClK/laKNLROed7voLL5XKBiIiIKE5oYn0CREREREoMToiIiCiuMDghIiKiuMLghIiIiOIKgxMiIiKKKwxOiIiIKK4wOCEiIqK4wuCEiIiI4gqDEyIiIoorDE6IiIgorjA4IaJBtWvXLnz3u99FWVkZjEYjhg8fjgsvvBB/+tOf5GNGjRoFQRDkr8LCQpxzzjl44403AACrV6+GIAh4//33VX/HJZdcgqysLNTW1kblbyKiwSVwbx0iGixbtmzB+eefj5EjR+L6669HcXExqqur8cUXX+Dw4cM4dOgQADE4ycnJwS9+8QsAQG1tLdauXYsjR47gz3/+M3784x9j1qxZ6Orqwu7du5Gamir/jtdeew3f+973sGbNGtx6660x+TuJKLIYnBDRoFm6dCkqKytx4MABZGdne13X0NCAwsJCAGJwUl5ejnfeeUe+vq6uDuPGjcPw4cOxf/9+fPHFFzjrrLNw991345FHHgEAdHR0YNKkSRg5ciQ+++wzaDRMBhMlA/5LJqJBc/jwYUydOtUvMAEgByaBFBcXY/LkyTh69CgA4IwzzsAtt9yCxx57DHv37gUA3HPPPWhoaMAzzzzDwIQoifBfMxENmrKyMmzbtg27d+8O+7Y2mw3V1dXIy8uTL1u1ahUKCgpw8803Y9u2bVizZg3uuusuTJs2LZKnTUQxxuCEiAbNXXfdhe7ubsyYMQNnnnkm7r77bmzYsAE2m83vWJvNhqamJjQ1NWHnzp247rrrUF9fj6uuuko+xmQy4cknn8Snn36KxYsXo6ysDPfee280/yQiigL2nBDRoKqsrMSqVavw/vvvo7u7GwBQUFCAv/zlL/j2t78NQOw5OX78uNfttFotrrnmGqxdu9arARYQe1neffddvPfee7j44ouj84cQUdQwOCGiqLBarfj666/xxhtv4A9/+AMcDgeqqqowZcoUjBo1CsXFxXj44YchCALS0tIwefJk1V4VALj//vvxwAMPoLGxEfn5+dH9Q4ho0OlifQJENDSkpKRgzpw5mDNnDiZMmIAbbrgBr732Gu677z4AQH5+PhYtWhTjsySieMCeEyKKutmzZwMATp06FeMzIaJ4xOCEiAbNxx9/DLXK8bvvvgsAmDhxYrRPiYgSAMs6RDRofv7zn6O7uxvf+c53MGnSJFitVmzZsgWvvPIKRo0ahRtuuCHWp0hEcYjBCRENmsceewyvvfYa3n33XTzzzDOwWq0YOXIkbr31Vtxzzz0BG16JaGjjah0iIiKKK+w5ISIiorjC4ISIiIjiCoMTIiIiiisDCk5Wr14NQRBwxx13yJedd955EATB6+uWW27xup3v9YIgYP369QM5FSIiIkoS/V6tU1lZibVr12L69Ol+191444148MEH5Z/T0tL8jlm3bp3Xnhjs2iciIiKgn8FJZ2cnli1bhmeffRYPP/yw3/VpaWkoLi4Oeh/Z2dl9HkNERERDT7/KOsuXL8fSpUsD7oPx0ksvIT8/H+Xl5Vi5cqW8E6nvfeTn52Pu3Ll47rnnVKdISiwWC8xms/zV3t6OxsbGoLchIiKixBR25mT9+vXYvn07KisrVa+/5pprUFZWhpKSEuzcuRN333039u/fj3/+85/yMQ8++CAuuOACpKWlYcOGDbj11lvR2dmJ2267TfU+V61ahQceeMDv8vb2dphMpnD/BCIiIopjYQ1hq66uxuzZs1FRUSH3mpx33nmYMWMGnnjiCdXbfPTRR1i4cCEOHTqEsWPHqh5z7733Yt26daiurla93mKxwGKxyD+bzWaUlpYyOCEiIkpCYZV1tm3bhoaGBsycORM6nQ46nQ6bNm3Ck08+CZ1OB4fD4XebefPmAQAOHToU8H7nzZuHmpoarwBEyWAwwGQyeX0RERFRcgqrrLNw4ULs2rXL67IbbrgBkyZNwt133w2tVut3m6qqKgDAsGHDAt5vVVUVcnJyYDAYwjkdIiIiSkJhBSeZmZkoLy/3uiw9PR15eXkoLy/H4cOH8fLLL+OSSy5BXl4edu7ciRUrVmDBggVyGejtt99GfX09zjjjDBiNRlRUVOCRRx7BXXfdFbm/ioiIiBJWRHclTklJwQcffIAnnngCXV1dKC0txZVXXol77rlHPkav12PNmjVYsWIFXC4Xxo0bh8cffxw33nhjJE+FiIiIElRC7kpsNpuRlZXFhlgiIqIkxL11iIiIKK4wOCEiIqK4wuCEiIiI4gqDEyIiIoorDE6IiIgorjA4ISIiorjC4ISIiIjiCoMTIiIiiisRnRBLRBSPXC4X/lVVi/31HThtRDYuLi+O9SkRURAMTogo6VUea8Udr1TJP//47NH4zbemqB77r6qTePrjw5g/Ng93XTQRGQa+TBJFG8s6RJT0att65O8FAfjrp0dxrKlL9dj/+/w49td34Pktx/Dbf++N1ikSkQKDEyJKeu09NgDAJdOKMb4wAwBQ3dqNLosdj7z7Db45ZQYgln8ONXTKt3tn5yn02hzRP2GiIY7BCRElPSk4yUrVo8hkBADUtffigbf34JnNR3DF01sAAE2dVvnY3PQUdPTasXF/Q2xOmmgIY3BCREmnrduK2/6+Axc+vgktXZ6Aw5Sqx7AsT3Dy4Tdi4NHjzo4cbhSzJiNz03DVrBEAxOwJEUUXO72IKO5Y7A68ueMk5o3Owz+212BvrRlXzS7FRVOLIAhC0Nu6XC5c+9et2H1SLNV8dqjJK3Ni0IqfyU6Ze2Hutcm36+i1ySWdcYUZmDs6F2s3H8Hx5u7B+BOJKAgGJ0QUV2wOJ3728g5U7K1HVqpeDiw+3NeAl2+chzPH5ge9/cm2HjkwAcRsiDI4yU5NAQBUHm2BzeGSjzvY0CkHJ2ML0pGbLh7X0mWN3B9HRCFhWYeI4so/t9egYm89AE+viGT3yfY+b7+31uz186EG7+BEKuscVDS+AsDB+g65rDOuMIPBCVEMMTghorhyokUso+RniMGBViPgypli/8fB+s6AtwOAXpsDe90rb3LS9ACAw41dMKs0xPp67asabD3aAgCYUJQpByc9Ngd6rFyxQxRNLOsQUVzpsToBAN+dVYqSbCMKMgxwAfjH9hq/bIfSvjozvvXkp7A7xVLNt6aX4IUvjuNIYyey3YGKMnMiuWRaMd7dVYevjrcCABZNLsSM0mwAgF4rwOZwoaXbiuEpqRH+S4koEGZOiCiu9NjsAIC0FC2umz8KS6YNk2eTHGrohMvlUr3d+7vr5cAEABZPLUKKVgOL3Yl6swWAu+fEHahIblowVv5+zqgc/PEHp0MQBAiCIGdPWlnaIYoqZk6IKK50u0soaSla+bKyvHToNAI6LXacau9Fa7cVOWkpKMn2ZDO6rHav+5k2PAuj89Oxv75DviwrVe+12mdYlhEzSrPxt/+aC4NOg3mjc72uz0lLQb3ZgmZ3cHKooROvbavGT88di+y0lMj+4UQkY3BCRHFFCk5SFcFJik6Dsrw0HG7swt+/PIGnPj4EnUbAHYsmYPn54wAAp9p75eMnDzMhOy0F4wozvIKTTKOYNZk7KhdfHmvBikUTAADnTihQPRdl5qTH6sB3/3cL2rpt6Oy147ffmRbBv5qIlFjWIaK4Io2LV2ZOAGB8YSYA4E8fHYLLBdgcLjz6/n7Um8Wg5JR7/5y7Fk/ACz+eCwCYUmKSb59p1EGrEbMiT11zOl748Vx8b05p0HORgpPmLiv++OFBtHWLjbUvbT2B9m5bsJsS0QAwOCGiuCJnTvTewckl04d5/SztFvzpwSYAnszJ/LH5yM8wAABOG5EtH5+V6uk1KTQZcc549WyJkjJz8p/d3pNiX/nqRJ+3J6L+YXBCRHHFU9bxrjp/+7QSfH+2mOlYOKkQ180vAyBOgHU4XXIGpSTbsxpn2vCsAZ2LJ3NiQW2beP83njMaAPDpoeYB3TcRBcbghIjiSo/Vs1rH18PfKcfTy2bi0atOw9njxEmxnxxqwqn2HtidLmgEoMCdNQGALMXKnJrWnrDPRQpO9td1wOpwQiMAi6cWA/Af9ubry6Mt+OIIAxii/mBwQkRxRdqEz7esAwB6rQaXTBuG3PQUzCzLgUGnQWOHBWf/v48BAEUmI3TayL2s5bhX5OxyT6YdlpWK8pIsaASgqdOCho5e1dv12hy47rmtuO65L9HRy94UonAxOCGiuKK2WkeNUa/FLy+aKGc3APgNWAOAX140EQBwvbsMFI48931Le/AMz05FaooWYwrEuSuBsic1rT3otTlhtTuDDo4jInUMTogorvSozDkJ5CfnjMEbt54p/9zQYfE75pZzx+IfP52P/146Oexzyc3wnmUyPEecqzJlmLgKSBqV7+tkm6eEdKCuQ/UYIgqMwQkRxQ2r3SlPeU3ThzaGqSwvHfPH5AEArnDvwaOk1QiYVZYLg67vYMfX2IIMeVUQIGZOAM8S5UCZk1pFcKKcs0JEoeEQNiKKG1K/CQAYU0L/7LTuhjn4z+46LJxcGNHz0Ws1OHNsHja4d0ke4c6cTCwSZ64cbuxSvd1JRfNtX5sVEpE/Zk6IKG5IJR2tRkBKGI2tRr0Wl58+XJ4AG0nnTvTMQ5HKOtIqoEDNrieZOSEaEAYnRBQ3uqVlxHqt1x43sbRAMaxtWJYYnEilni6LXfU2ysxJY4eFGwcShYllHSKKG6Gu1Imm0tw0/OjMUTDLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 640x480 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"idx = -2\n",
|
||||
"data = GetStockData(tckr, span=span, interval=interval)\n",
|
||||
"\n",
|
||||
"ts = torch.linspace(0, T, data.shape[0])\n",
|
||||
"train_x = ts[:ntrain]\n",
|
||||
"test_x = ts[ntrain:(ntrain+ntest)]\n",
|
||||
"\n",
|
||||
"y = torch.FloatTensor(data['close_price'].to_numpy())\n",
|
||||
"train_y = y[:ntrain]\n",
|
||||
"test_y = y[ntrain:(ntrain+ntest)]\n",
|
||||
"\n",
|
||||
"dt = ts[1] - ts[0]\n",
|
||||
"\n",
|
||||
"plt.plot(train_x, train_y)\n",
|
||||
"plt.plot(test_x, test_y)\n",
|
||||
"plt.title(tckr);\n",
|
||||
"sns.despine()\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "384d8c44",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Model Setup"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "36931a8a",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Learn Train Vol Path"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 168,
|
||||
"id": "0fe7bc7b",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Iteration: 1, Func. Count: 6, Neg. LLF: 141119365.63998526\n",
|
||||
"Iteration: 2, Func. Count: 18, Neg. LLF: 67039164.31986434\n",
|
||||
"Iteration: 3, Func. Count: 30, Neg. LLF: 31901095.64321779\n",
|
||||
"Iteration: 4, Func. Count: 42, Neg. LLF: 1.0718228553175738e+17\n",
|
||||
"Iteration: 5, Func. Count: 55, Neg. LLF: 67236.99966236952\n",
|
||||
"Iteration: 6, Func. Count: 64, Neg. LLF: -1473.3917800663248\n",
|
||||
"Optimization terminated successfully (Exit mode 0)\n",
|
||||
" Current function value: -1473.3917726261616\n",
|
||||
" Iterations: 10\n",
|
||||
" Function evaluations: 64\n",
|
||||
" Gradient evaluations: 6\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/Users/gregorybenton/miniconda3/lib/python3.8/site-packages/arch/univariate/base.py:292: DataScaleWarning: y is poorly scaled, which may affect convergence of the optimizer when\n",
|
||||
"estimating the model parameters. The scale of y is 2.557e-05. Parameter\n",
|
||||
"estimation work better when this value is between 1 and 1000. The recommended\n",
|
||||
"rescaling is 100 * y.\n",
|
||||
"\n",
|
||||
"This warning can be disabled by either rescaling y before initializing the\n",
|
||||
"model or by setting rescale=False.\n",
|
||||
"\n",
|
||||
" warnings.warn(\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"garch = arch_model(np.log(train_y[1:]/train_y[:-1]), q=1, p=1).fit()\n",
|
||||
"v = (garch.conditional_volatility/np.sqrt(dt))\n",
|
||||
"vlog = v.log().float()\n",
|
||||
"# v = vlog.exp()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 169,
|
||||
"id": "6d9fbb28",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAmEAAAGCCAYAAACozRT6AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAA9hAAAPYQGoP6dpAAChl0lEQVR4nOyddZgb1frHP2ez7u3W3QUK1Ci0WJHi7u76w10uXPTicpF7gYsUd7eWUqBAobRUqbu3W9ntumfP748zk0yyyWp2k+y+n+eZZyYzZ86cbLKT77zve95Xaa0RBEEQBEEQWpaYcA9AEARBEAShLSIiTBAEQRAEIQyICBMEQRAEQQgDIsIEQRAEQRDCQGy4ByAIgiAIQuthzpw5aUBXxNBTDWwdNWpUYbAGSmZHCoIgCILQVObMmRMD3OVyuc5XSsUBKtxjCjNaa13pdrvfAh4eNWpUtX8DsYQJgiAIghAK7oqLi7uqS5cuFSkpKSVKqTZt5dFaq+Li4uTs7OyrKisrAR7ybyOWMEEQBEEQmsScOXPSXS7X7G7dusV16tQpJ9zjiSS2b9+etWXLlkq32z3K3zXZ1v21giAIgiA0nS5KqbiUlJSScA8k0rCsgnGYODkfRIQJgiAIgtBUYgDV1l2QgbD+JooAmktEmCAIgiAIQhgQESYIgiAIghAGRIQJgiAIgiCEARFhgiAIgiAIYUBEmCAIgiAIbZYxY8YMPv/883udf/75vdLS0oa3a9dur+uvv75bdXU18+bNS0xKShrx0ksvtbfbv/rqq+0SExNHzpkzJ7Gp15ZkrYIgCIIghJxqrSkpr2pxY09yQmx1jGpYsv5PP/0064wzztg5ffr0pX/88UfKTTfd1LtXr14VN9988857771306233trr0EMPLYqJidE333xz77vvvnvTqFGjypo6VknWKgiCIAhCk5gzZ86Q2NjYyQMHDixKTk4uAygqq4wZdt+UES09lkX3HT4vNTGuRomgYIwZM2ZwTk5O7MqVKxfHxBjN+H//93/dv//++8zVq1cvBjj44IMHFBYWuuLj43VMTIz+9ddfV9pt66KkpCRx5cqVqVVVVUeOGjVqmfOYuCMFQRAEQWjTjBw5stgpqsaNG1e8fv36hKqqKgDeeeeddcuXL09avHhx8rvvvruuvgKsLsQdKQiCIAhCyElOiK1edN/h88Jx3VD3OWvWrOTS0tKYmJgYNm7cGNe7d+/KUPQrIkwQBEEQhJAToxQNcQuGk3nz5qU4X8+YMSOld+/e5bGxsWzbts11xRVX9Lnuuuuys7Oz484///y+f//995LU1NQmx3OJO1IQBEEQhDbN1q1b4y+99NIeCxYsSHj55ZfbT5w4sdOVV165DeCiiy7q3bVr14rHHntsy0svvbSxurpaXXXVVT1DcV2xhAmCIAiC0KY5+eSTc0pLS2P233//oTExMVxyySXbb7755p0vvPBC1rRp0zJmzpy5JC4ujri4uOqJEyeuOfzww4ccd9xxeaeffnpBU64rIkwQBEEQhDZNXFycfv311zcCG5z7r7nmmpxrrrkmx7nv4IMPLqmsrJwbiuuKO1IQBEEQBCEMiAgTBEEQBEEIA+KOFARBEAShzTJr1qzl4bq2WMIEQRAEQRDCgIgwQRAEQRCEMCAiTBAEQRAEIQyICBMEIeQopfZQSn2ilFqvlCpTSm1WSv2glLrW0WadUko7ljKl1Eql1BNKqfZ+/b2hlCqq5XpFSqk3HK8fsvocH6Dtmdaxa0LyZgVBEBqJBOYLghBSlFLjgJ8x+XZeAbKBnsC+wPXA847m84GnrO1EYBRwA3AQMKYJw3gIOBN4SSm1p9a6whpbJvAM8Bfw3yb0LwiC0GREhAmCEGr+AeQDe2ut85wHlFKd/Npu1lq/43j9qmXxukUpNVBrvbIxA9BalymlrgKmAHcC91uHHgU6AkdpraOipp0gCK0XcUcKghBq+gOL/QUYgNZ6ez3Oz7bWVU0ZhNb6B+A94E6l1CCl1FjgcuBZrfX8pvQtCIJgo5Qa9fbbb2c25lyxhAmCEGrWA2OVUsO01ovqaBunlOpgbScCI4CbgF+11mtDMJabgKOAl4EsYBNwbwj6FQRBaDIiwgRBCDVPApOA+UqpWcBvwI/Az1rrSr+2hwM7/Pb9DpwcioForbcppe7AiDCAE7XWQQP8BUEQWhJxRwqCEFIsN+BY4CtgL+A24Htgs1LqeL/mM4EJ1nIsJp5sd+ArpVRSiIa001qXANND1KcgCK2AJ598skOnTp32dLvdPvsPPfTQ/qeddlofgMcee6xjz549h8XFxY3s06fPsP/85z/tA/XVGESECYIQcrTWf2mtTwbaYWY5PgKkAZ8opXZzNN2ptZ5qLd9qrR8GLgXGWesGXdZ/h1IqDXgOWA7EA481/N0IgtAodDWUFMS0+NKAOTcXXHDBrry8vNhvvvkmzd63bds212+//ZZx7rnn5rz11luZd999d8+rr75625w5cxZfeOGFO66//vq+X3/9dVpt/dYXcUcKgtBsWKkh/gL+UkqtACYCp+GdrRiIH631gXjTWZQBCUoppbX2EVtKKYWJJysL0Ne/gC4YIXgmZtblRK317418S4Ig1JfSohhOzBjR4tf9In8eyen1UmIdO3Z0H3jggfnvvvtu+xNOOKEQ4O23326XmZlZdeyxxxbuvffeQ0499dScO+64YwfAnnvuuW3WrFkpTz31VOfjjjuusKlDFUuYIAgtxWxr3bWOdvbDYapj33prf/8A7QcALquNB6XUaOBq4AWt9VyM8NuIyR0mD6CCIABw9tln506aNKldaWmpAvjwww+zTjjhhFyXy8Xq1asTx40b5xNHOnbs2KJVq1aFJFxCbkSCIIQUpdTBwDR/ixVwtLVeXkcXx1nrBY59k4CHgWswyVydXO1oY4/BhQnG3wrcA6C1LrYy9n8J3Ag8Udd7EQShCSSlVvNF/rywXLcBnHnmmXnXX399748++ihjv/32K54zZ07qM888s7G5hudERJggCKHmeSBZKfU5sAwTizUOOANYh3FJ2nRXSp1rbcdjAvmvwATTezLra63nK6VeBa5XSg0EfrAOTcCIu1e11k7Rdh0wEjhFa13o6OcrpdRXwL1KqQ+11htC9J4FQfBHxVBft2A4SU5O1kcccUTee++9l7Vy5cqEPn36lO2///4lAP379y/7448/Uq+99tocu/2MGTNSBw4cWBqKa4sIEwQh1NyCifs6GpMcNR5Twui/wEN+SVyHA29b29UY8fUZcI/WerNfv1cAC4GLMYH+YKxq1wH/sRsppXoADwDfaK0/CzC+a4ElGJF3QmPeoCAIrYtzzz0354wzzhi4YsWKxNNOOy3X3n/jjTdmX3zxxf2GDx9ecvTRRxd8+umnmVOmTGn3xRdfrAjFdUWECYIQUrTWk4HJ9WjXp4H9VmNmOj5XR7tNmJmYwY5vwDfeTBCENs5xxx1XmJGRUbVu3brECy+80GP1Ou+88/K2bNmy8YUXXuh899139+zevXvFs88+u/bYY49tclA+gKoZtiEIgiAIglB/5syZMyQ2NnbywIEDi5KTkwPNVG6zlJSUJK5cuTK1qqrqyFGjRi1zHpPZkYIgCIIgCGFARJggCIIgCEIYEBEmCIIgCIIQBkSECYIgCIIghAERYZiyJ0qp7lb5E0EQBEEQGkY1oLXW8jvqh/U30Zi/kQ+SosLQDdi0cWOLJMgVBEEQhNZGhta6sri4ODklJSUkiUxbC8XFxcla60pMBQ8fRIQJgiAIgtAkRo0aVTBnzpy3srOzrwKyUlJSSpRSbToHltZaFRcXJ2dnZ8e73e7XRo0aVSO3mIgwQRAEQRBCwcOVlZVs2bLlfKVUMtDWXZNaa13pdrtfw9S+rYEkawWUUt2x3JE9evQI93AEQRAEIdrwCK45c+akAV2RuPNqYGsgC5iNWMIEQRAEQQgZlugISVmf1k5bV6mCIAiCIAhhQUSYIAiCIAhCGBARJgiCIAiCEAZEhAmCIAiCIIQBEWGCIAiCIAhhQESYIAiCIAhCGBARJgiCIAiCEAbCKsKUUgcqpb5WSm1RSmml1ImOY3FKqceUUguVUsVWm7eUUt0C9HOMUmqmUqpUKbVLKfVFS74PQWgNaK2R5M2CIAgtR7gtYSnAAuDqAMeSgZHAg9b6ZGAw8JWzkVLqFOBtYCKwF7Af8F7zDVkQopf8kkqWbi2osb+s0s0R//6Vy96aHYZRCYIQlF3b4Px+8NxV4R6J0AyENWO+1noSMAlAKeV/LB+Y4NynlLoGmKWU6qW13qCUigWeBW7VWr/maLqkWQcuCFHK0c/9xua8UqbceCCDOqd59k9elM2KbUWs2FYE2esgZwvsPi58AxUEwTD9M8heC9+9Ahc+BOlZ4R6REELCbQlrKBmABvKs1yOB7kC1UmqeUmqrUmqSUmpYbZ0opRKUUun2AqQ266gFoQXZml/KrR8vYNX2Ip/9S7YUsDmvFIC/1uX6HnNax67bB27cD5bMaPaxCoJQB3N/MOtqN8z4qva2QtQRNSJMKZUIPAa8r7W2fzH6Wev7gIeAY4FdwDSlVPtaursTyHcsy5pjzIIQDk78z+98PGcTN38032f/p3M3ebYrq6p9jvm4KPO2m/U3LzbXEAVBqA/uKpj3o/f19M/CNxahWYgKEaaUigM+wlRpdzrG7fH/S2v9qdZ6DnARxlp2Wi1dPoKxqtnLkJAPWhDCgNaabQXlAKzLKfHZ/83fWzyvc4srWJ9TzKVvzmb2utyAcWL8/Uuzj1cQhFpYMRtKCiAu3ryeOwWKA/yvClFLWGPC6oNDgPUGDnFYwQC2WmtPDJjWulwptQboFaxPrXU5UO64RlqwtoIQTSzfVujZHtrV+7XeUVjuEWcAOcUVfDp3M1OXbqPSXc3OooqanW3fABVlEJ/YrGMWBCEIm1eZ9e77w85NsGkFzPoODj4zvOMSQkZEW8IcAmwLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 600x400 with 2 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {
|
||||
"needs_background": "light"
|
||||
},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(1,1,dpi=100)\n",
|
||||
"ax.plot(train_x, train_y, label='px')\n",
|
||||
"ax2 = ax.twinx()\n",
|
||||
"ax2.plot(train_x[1:], v, color='OrangeRed', label=\"vol\")\n",
|
||||
"sns.despine()\n",
|
||||
"fig.legend()\n",
|
||||
"ax.set_title(tckr)\n",
|
||||
"ax.set_ylabel(\"px\")\n",
|
||||
"ax2.set_ylabel(\"vol\")\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "e84c0ba3",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "cad3b907",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Train Latent GP"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 130,
|
||||
"id": "e0f00b97",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"tensor([1.5353], grad_fn=<SoftplusBackward>)\n",
|
||||
"tensor([0.3422])\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"vol_lh = gpytorch.likelihoods.GaussianLikelihood()\n",
|
||||
"vol_lh.noise.data = torch.tensor([1e-6])\n",
|
||||
"vol_model = BMGP(train_x[:-1], vlog, vol_lh)\n",
|
||||
"\n",
|
||||
"optimizer = torch.optim.Adam([\n",
|
||||
" {'params': vol_model.parameters()}, # Includes GaussianLikelihood parameters\n",
|
||||
"], lr=0.01)\n",
|
||||
"\n",
|
||||
"# \"Loss\" for GPs - the marginal log likelihood\n",
|
||||
"mll = gpytorch.mlls.ExactMarginalLogLikelihood(vol_lh, vol_model)\n",
|
||||
"\n",
|
||||
"for i in range(500):\n",
|
||||
" # Zero gradients from previous iteration\n",
|
||||
" optimizer.zero_grad()\n",
|
||||
" # Output from model\n",
|
||||
" output = vol_model(train_x[:-1])\n",
|
||||
" # Calc loss and backprop gradients\n",
|
||||
" loss = -mll(output, vlog)\n",
|
||||
" loss.backward()\n",
|
||||
"# print(loss.item(), model.covar_module.vol.item())\n",
|
||||
" optimizer.step()\n",
|
||||
" \n",
|
||||
"print(softplus(vol_model.covar_module.raw_vol))\n",
|
||||
"print(vol_model.mean_module.constant.data.exp())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "11f37719",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Train Data GP"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 106,
|
||||
"id": "24d67093",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"voltron_lh = gpytorch.likelihoods.GaussianLikelihood()\n",
|
||||
"voltron = VoltronGP(train_x[:-1], train_y[:-1].log(), voltron_lh, v)\n",
|
||||
"voltron.likelihood.raw_noise.data = torch.tensor([1e-6])\n",
|
||||
"voltron.vol_lh = vol_lh\n",
|
||||
"voltron.vol_model = vol_model"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 107,
|
||||
"id": "74164ae9",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"grad_flags = [False, True, True, True, False, False, False]\n",
|
||||
"\n",
|
||||
"for idx, p in enumerate(voltron.parameters()):\n",
|
||||
" p.requires_grad = grad_flags[idx]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 110,
|
||||
"id": "18f78b62",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"0\n",
|
||||
"1\n",
|
||||
"2\n",
|
||||
"3\n",
|
||||
"4\n",
|
||||
"5\n",
|
||||
"6\n",
|
||||
"7\n",
|
||||
"8\n",
|
||||
"9\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"nvol = 10\n",
|
||||
"npx = 10\n",
|
||||
"vol_paths = torch.zeros(nvol, ntest)\n",
|
||||
"px_paths = torch.zeros(npx*nvol, ntest)\n",
|
||||
"\n",
|
||||
"voltron.vol_model.eval();\n",
|
||||
"voltron.eval();\n",
|
||||
"\n",
|
||||
"for vidx in range(nvol):\n",
|
||||
" print(vidx)\n",
|
||||
" vol_pred = voltron.vol_model(test_x).sample().exp()\n",
|
||||
" vol_paths[vidx, :] = vol_pred.detach()\n",
|
||||
" \n",
|
||||
" px_pred = voltron.GeneratePrediction(test_x, vol_pred, npx).exp()\n",
|
||||
" px_paths[vidx*npx:(vidx*npx+npx), :] = px_pred.detach().T"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 111,
|
||||
"id": "404e736c",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAA2cAAAEICAYAAADbdozDAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Il7ecAAAACXBIWXMAAAsTAAALEwEAmpwYAAEAAElEQVR4nOy9d5hc+VXm/34r5+quTsphpNGMJo8neDwOjE3YcQB7WWMTDDZrMCzLEn4sGZa8mCWZBYzxgtcEY5u1sTHGAYNxtid6gjRBI41SS+pYOd/w/f1R/R7dKnVLLanV6pbO53n0qLvq1q1v3ZLq1nvPe95jrLVQFEVRFEVRFEVRLi+hy70ARVEURVEURVEURcWZoiiKoiiKoijKmkDFmaIoiqIoiqIoyhpAxZmiKIqiKIqiKMoaQMWZoiiKoiiKoijKGkDFmaIoiqIoiqIoyhpAxZlyxWGMqRtjrrnc61guxpj3GmN+c+Hnlxpjnr3A/bzLGPPLK7s6RVEU5WrEGHOfMWbyMj3354wxP7Dw8/cYY/7lAvfzSWPMm1d2dYpyaVFxpqx5jDFHjDGtBdE1vSBmMkttb63NWGufv5xruFCstV+01l63jPW8xRjzpYHH/rC19jdWek2KoijK+sQY8yljzK8vcvtrjTFTxpjIRezbGmMaC+fFE8aYPzDGhC9uxWdirX2ftfZblrGeXzXG/O3AY19prf2rlV6TolxKVJwp64VvtdZmALwAwJ0Afmlwg4s5yayjNSiKoijKcvkrAG8yxpiB278XwPuste5F7v/WhfPiNwL4bgA/OLiBnhcV5fxQcaasK6y1JwB8EsBNgFy5+6/GmOcAPBe4bffCz0ljzO8bY44aYyrGmC8ZY5IL991jjPmKMaZsjHncGHPfCq7hNcaYxxb2/RVjzC18vDHmdmPMo8aYmjHmgwASgfv6bCTGmK3GmH8wxswaY+aNMX9ijNkL4F0AXrRwxbK8sK3YIxd+/0FjzEFjTNEY8zFjzKbAfdYY88PGmOcW1vinPHkbY3YbYz6/cLzmFtaoKIqirD8+CmAEwEt5gzFmGMBrAPy1MSZujHmHMebkwp93GGPi5/sk1tpnAHwRwE3GmB0L55i3GmOOAfjswvP+Z2PM08aYkjHm08aY7YE1fbMx5pmF886fADCB+/qcIsaYG40xn1k4t00bY37BGHM/gF8A8MaF8+LjC9sG7ZEhY8wvLXwfmDHG/LUxJr9wH9f8ZmPMsYVz3y+e73FQlJVAxZmyrjDGbAXwKgBfD9z8OgAvBHDDIg/5PQB3ALgXQAHAzwDwjTGbAfwzgN9cuP2/A/iwMWbsYtdgjLkdwHsA/BB6J8U/B/CxhZNgDL2T5d8sPO//A/CflnieMICPAzgKYAeAzQA+YK19GsAPA/jqgoVzaJHHvgLAbwN4A4CNC/v4wMBmrwFwF4BbFrb7Dwu3/waAfwEwDGALgD8+1zFRFEVR1h7W2haAvwfwfYGb3wDgGWvt4wB+EcA9AG4DcCuAu7GIK+RcGGNuQE8ABs+L3wBgL4D/YIx5LXri6dsBjKEn5N6/8NhRAP+w8LyjAA4BePESz5MF8K8APgVgE4DdAP7NWvspAP8TwAcXzou3LvLwtyz8eTmAawBkAPzJwDYvAXAdepXA/7FwMVRRVhUVZ8p64aMLFaIvAfg8eh/C5LettcWFk5BgjAkB+M8Aftxae8Ja61lrv2Kt7QB4E4BPWGs/Ya31rbWfAfAweqLrYtfwNgB/bq19YOE5/wpAB70T4D0AogDeYa11rLUfAvDQEs93N3onn5+21jastW1r7ZeW2HaQ7wHwHmvtowuv9+fRq7TtCGzzdmtt2Vp7DMC/o3dyBgAHwHYAm87zORVFUZS1x18BeL0xhi6N71u4DeidK37dWjtjrZ0F8GvoWR6Xy6PGmBKAfwLwFwD+b+C+X104d7XQu6D429bapxeslP8TwG0L1bNXAdhvrf2QtdYB8A4AU0s832sATFlrf3/h/FSz1j6wzLV+D4A/sNY+b62to3de/E7Tb7v8NWtta0G4Po6eYFWUVUV9wMp64XXW2n9d4r7jS9w+ip5l8NAi920H8B3GmG8N3BZFT6Rc7Bq2A3izMea/BW6LoSe0LIAT1lobuO/oEvvcCuDoBfYEbALwKH+x1taNMfPoVd+OLNwcPPk10buKCPSqi78B4MGFk+7vW2vfcwFrUBRFUS4z1tovGWPmALzOGPMQehf+vn3h7k3oPwcdXbhtubzAWnsweIM53d42eF78I2PM7wc3Re+ctCm4rbXWGmOWOq9vxeLn9OWw2GuNAJgI3LbUeVFRVg2tnClXAnaJ2+cAtAHsWuS+4wD+xlo7FPiTtta+fQXWcBzAbw3sO2WtfT+AUwA2G9PXnL1tiX0eB7DNLN5MvdRrJifROxkCAIwxafQslifO+UKsnbLW/qC1dhN61sx3moUePkVRFGVd8tfoVczeBODT1trphdv7zhXonY9OrtBzDp4Xf2jgvJi01n4FvfPiVm64cH7cisU5jp4l8VzPtxiLvVYXwPTimyvK5UHFmXLFYq310ev9+gNjzCZjTNgY86KFZue/BfCtxpj/sHB7YiGMY8sKPPX/AfDDxpgXmh5pY8yrF7zyX0XvZPBjxpioMebb0buKuRgPonfSevvCPhLGGPrwpwFsWehhW4z3A/h+Y8xtC6/3fwJ4wFp75FyLN8Z8R+A4lNA74fnnftmKoijKGuWvAXwTemmKwWj59wP4JWPM2ELv1/9A7/y40rwLwM8bY24EAGNM3hjzHQv3/TOAG40x375wMfLHAGxYYj8fB7DRGPMTC33cWWPMCxfumwawY6GlYTHeD+AnjTE7TW8UDnvULjaxUlFWFBVnypXOfwfwJHp9XUUAvwMgZK09DoANyrPoXY37aazA/wlr7cPonQD/BD1xcxC9JmRYa7vo2UnesrCeN6LXCL3YfjwA34pew/MxAJML2wO99Kv9AKYW7CqDj/1XAL8M4MPoCbxdAL5zmS/hLgAPGGPqAD6GXs/eis6NUxRFUVaPhQtzXwGQRu9znfwmev3WT6B3rnx04baVfv6PoHf+/YAxpgpgH4BXLtw3B+A7ALwdwDyAawF8eYn91AB8M3rnxin0EpJfvnD3/1v4e94Y8+giD38PemFcXwBwGD1nzX9bZDtFuayY/tYXRVEURVEURVEU5XKglTNFURRFURRFUZQ1gIozRVEURVEURVGUNYCKM0VRFEVRFEVRlDWAijNFURRFURRFUZQ1wKoOob7//vvtpz71qdV8SkVRFOXyYM69iUL0/KgoinJVseQ5clUrZ3NzZyR+K4qiKMpVj54fFUVRFEBtjYqiKIqiKIqiKGsCFWeKoiiKoiiKoihrABVniqIoiqIoiqIoawAVZ4qiKIqiKIqiKGsAFWeKoiiKoiiKoihrABVniqIoiqIoiqIoawAVZ4qiKIqiKIqiKGsAFWeKoigKAMD3fbiue7mXoSiKoihrjunpaTiOc8mfR8WZoiiKAgBoNBpotVqXexmKoiiKsqbodDp4+umnsW/fvkv+XCrOFEVRlD48z7vcS1AURVGUNQPPi5VK5ZI/l4ozRVEUBQBgjAGg4kxRFEVRgqzmeVHFmaIoitKH9p0piqIoymmC4qzT6VzS51JxpiiKogAArLUIhUKIxWKXeymKoiiKctmx1sLzvD5xVqvVLulzqjhTFEVRYK0FAESjUUQikcu8mrWNMeY9xpgZY8xZO8ONMXcZY1xjzOtXa22KoijKynHgwAF88YtfxKFDhwAAe/bswfDw8CV9ThVniqIoiogz9p0pZ+W9AO4/2wbGmDCA3wHwL6uxIEVRFGXlaTQaAIBmswkAKBQKCIfDl/Q5VZwpiqIoKs7OA2vtFwAUz7HZfwPwYQAzl35FiqIoyqWg2+32/R4KXXrppOJMURRFge/7AHqNzu12+zKvZn1jjNkM4D8C+LPLvRZFURTlwrDWniHOLnXVDFBxpiiKoqAnzlzXhe/7UkVTLph3APhZa61/to2MMW8zxjxsjHl4dnZ2dVamKIqiLAvP8+TCJdHKmaIoirIqtFottNttWGtX5crgFc6dAD5gjDkC4PUA3mmMed3gRtbad1tr77TW3jk2NrbKS1QURVHOBqtmiUQCQK9qthrWf43kUhRFUWCthTEGxhgVZxeJtXYnfzbGvBfAx621H71sC1IURVHOi3q9jv379wMAUqnUqtr9VZwpiqJcxfi+j263K+IMWB3bxnrGGPN+APcBGDXGTAL4FQBRALDWvusyLk1RFEVZAY4cOYJWqwWgJ86KxeKqWf5VnCmKoqwA7XYb0Wh03VWdWq2WCLRQKIRIJKKJjefAWvtd57HtWy7hUhRFUZRLAK2MQE+cATij/+xSoZdHFUVRLhLHceA4zhmpTmudYLOz67qIRCJIJpOXeVWKoiiKcnkJCrHVPi9q5UxRFOUicRwHwPqbEeZ5HoDeurvdLsLhMHzfl94zRVEURbkacV0XAHDdddchEllduaSVM0VRlIvAcRwROavhR3ccZ8WsFa7ripWREfqNRkN89oqiKIpyNeI4DnK5HDZu3IhoNLqqz62VM0VRlIuAvVrGmEsmzhhxH4/H0W63EYvFEI/HL3q/gxU/BoGs9olIURRFUdYSruvKuVArZ4qiKOsE9mxFo9FLJs7Yz8YB0QBWZFC0tRatVguu64rwS6VSSCaTKs4URVGUqxr2YQNY9aAvrZwpiqJcIPSkR6PRFRFMi0HLJHC60uW6Lur1OkKhENLp9Hntz3VdeJ6HcDgs8fmu68J1XXQ6HeRyuRVdv6IoiqKsN4LizBiDQqGAiYmJVXluFWeKoigXiO/7Ymlk5Sw4L+x89uM4zqJWRe7PWitiMPg4rmG5z9NqteB5nuyz3W4jFArB8zyEQqELWr+iKIqiXCnwfBu0M95yyy2r9vxqa1QURblAgkKGf19I9azdbqPb7S4a9OH7PsLh8JILine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 1080x288 with 2 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {
|
||||
"needs_background": "light"
|
||||
},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(1,2,figsize=(15, 4))\n",
|
||||
"\n",
|
||||
"ax[0].set_title(\"Price Predictions\")\n",
|
||||
"ax[0].plot(train_x, train_y)\n",
|
||||
"ax[0].plot(test_x, test_y)\n",
|
||||
"ax[0].plot(test_x, px_paths.T, c='gray', alpha=0.1)\n",
|
||||
"# ax[0].set_ylim(50, 150)\n",
|
||||
"# ax[0].axvline(2., c='OrangeRed')\n",
|
||||
"sns.despine()\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"ax[1].set_title(\"Vol Prediction\")\n",
|
||||
"ax[1].plot(train_x[:-1], voltron.log_vol_path.exp().detach())\n",
|
||||
"ax[1].plot(test_x, vol_paths.T, c='gray', alpha=0.5)\n",
|
||||
"sns.despine()\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "7ced5dbc",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "a091b787",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.7.7"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,61 @@
|
||||
import matplotlib.pyplot as plt
|
||||
import seaborn as sns
|
||||
import numpy as np
|
||||
import torch
|
||||
import pandas as pd
|
||||
|
||||
import gpytorch
|
||||
from torch.nn.functional import softplus
|
||||
from voltron.kernels import BMKernel, VolatilityKernel
|
||||
from voltron.models import BMGP, VoltronGP
|
||||
import argparse
|
||||
from torch.distributions import Beta
|
||||
from scipy.special import betainc
|
||||
from Trainers import *
|
||||
|
||||
def main(args):
|
||||
full_data = pd.read_pickle("../../spdr-data/" + args.SPDR + ".pkl")
|
||||
tckrs = full_data.symbol.unique()
|
||||
|
||||
for tckr in tckrs:
|
||||
data = full_data[full_data["symbol"] == tckr]
|
||||
|
||||
ts = torch.linspace(0, data.shape[0]/252., data.shape[0])
|
||||
# train_x = ts[:ntrain]
|
||||
# test_x = ts[ntrain:(ntrain+ntest)]
|
||||
|
||||
y = torch.FloatTensor(data['close_price'].to_numpy())
|
||||
log_returns = torch.log(y[1:]) - torch.log(y[:-1])
|
||||
dt = ts[1] - ts[0]
|
||||
|
||||
eval_times = list(range(100, ts.shape[0], 100)) #+ [ts.shape[0]]
|
||||
prob_of_increases = []
|
||||
|
||||
for i, time in enumerate(eval_times):
|
||||
print("now running time: ", time)
|
||||
with gpytorch.settings.max_cholesky_size(2000):
|
||||
pred_vol = get_and_fit_gpcv(ts[:time], log_returns[:(time - 1)])
|
||||
vol_model = get_and_fit_vol_model(ts[:time], pred_vol)
|
||||
data_model = get_and_fit_data_model(ts[:time], y[:time],
|
||||
pred_vol, vol_model)
|
||||
end_ind = -1 if i + 1 >= len(eval_times) else eval_times[i+1]
|
||||
paths = predict_prices(ts[time:end_ind], data_model).detach()
|
||||
# now we predict the probability of increase at time i + 1
|
||||
prob_of_increase = (paths[..., -1] > y[time]).sum() / paths.shape[-2]
|
||||
# print("prob of stock increase: ", prob_of_increase.detach())
|
||||
|
||||
prob_of_increases.append(prob_of_increase.detach())
|
||||
|
||||
torch.save(obj=prob_of_increases, f="./outputs/voltron_" + tckr + ".pt")
|
||||
print(tckr, "Done")
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
parser.add_argument(
|
||||
"--SPDR",
|
||||
type=str,
|
||||
default="XLE",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
main(args)
|
||||
@@ -0,0 +1,625 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 55,
|
||||
"id": "8079c112",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns\n",
|
||||
"import torch\n",
|
||||
"# from voltron.robinhood_utils import GetStockData\n",
|
||||
"import os\n",
|
||||
"# import robin_stocks.robinhood as r\n",
|
||||
"import pickle5 as pickle\n",
|
||||
"\n",
|
||||
"sns.set_style('white')\n",
|
||||
"# style.use('whitegrid')\n",
|
||||
"palette = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"sns.set(palette = palette, font_scale=2.0, style=\"white\", rc={\"lines.linewidth\": 4.0})\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 56,
|
||||
"id": "e6613742",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAcwAAABECAYAAAAMTwWHAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAAClElEQVR4nO3arY4TURjH4XdpIc3CrthlwofgCioQFSgUigSNqMByLXgsqDo8N7ECMQJuAJY0ILZkKbDDIAgfTSh5BWcP2zyPnJMm/yaT/DLTbvV93wcA8Ffnag8AgLNAMAEgQTABIEEwASBhuO5guVxG27bRNE0MBoPT3AQAVXRdF/P5PMbjcYxGo5WztcFs2zam02nxcQDwv5nNZjGZTFaurQ1m0zQREfFx/3b0g+2yyyp5/Ohh7QlF3X/+tPaEop7c28z7MiLi1f0XtScUdePurdoTirr58lHtCUXt3VnUnlDM4WIYD57d+NnA360N5o/XsP1gO/rhxXLrKmquXK89oaiTnUu1JxS1f3Uz78uIiJ24UHtCUXvbu7UnFHXt/NfaE4q6vHtSe0Jxf/op0p9+ACBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIGG47qDruoiI2OqOT23MaZu/fV17QlHDxYfaE4p6d/i19oRiFvG59oSi3h8f1Z5Q1Jsvm/0s8ulobTrOvMPF9+/2o4G/2+r7vv/Thw4ODmI6nZZdBgD/odlsFpPJZOXa2mAul8to2zaaponBYHAqAwGgpq7rYj6fx3g8jtFotHK2NpgAwC+b/aIdAP4RwQSABMEEgATBBICEb0f+ZBgefApDAAAAAElFTkSuQmCC\n",
|
||||
"text/plain": [
|
||||
"<Figure size 576x72 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"sns.palplot(palette)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 57,
|
||||
"id": "ce6f398d",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"with open(\"../../spdr-data/XLF.pkl\", \"rb\") as handle:\n",
|
||||
" finance_data = pickle.load(handle)\n",
|
||||
"with open(\"../../spdr-data/XLE.pkl\", \"rb\") as handle:\n",
|
||||
" energy_data = pickle.load(handle)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 58,
|
||||
"id": "330416be",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"array(['AIG', 'AON', 'AXP', 'BAC', 'BK', 'BLK', 'BRK.B', 'C', 'CB', 'CME',\n",
|
||||
" 'COF', 'GS', 'ICE', 'JPM', 'MCO', 'MET', 'MMC', 'MS', 'MSCI',\n",
|
||||
" 'PGR', 'PNC', 'PRU', 'SCHW', 'SIVB', 'SPGI', 'TFC', 'TROW', 'TRV',\n",
|
||||
" 'USB', 'WFC'], dtype=object)"
|
||||
]
|
||||
},
|
||||
"execution_count": 58,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"np.unique(finance_data[\"symbol\"])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 59,
|
||||
"id": "9c80a0e6",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"array(['COP', 'CVX', 'EOG', 'SLB', 'XOM'], dtype=object)"
|
||||
]
|
||||
},
|
||||
"execution_count": 59,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"np.unique(energy_data[\"symbol\"])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 60,
|
||||
"id": "1f34d69f",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"data/matern_BAC.pt data/matern_CVX.pt data/matern_WFC.pt\r\n",
|
||||
"data/matern_BRK.B.pt data/matern_EOG.pt data/matern_XOM.pt\r\n",
|
||||
"data/matern_COP.pt data/matern_JPM.pt data/matern_v.pt\r\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"!ls data/matern_*"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 61,
|
||||
"id": "ac1df8ed",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"data/SM_BAC.pt data/SM_COP.pt data/SM_JPM.pt data/SM_XOM.pt\r\n",
|
||||
"data/SM_BLK.pt data/SM_CVX.pt data/SM_MS.pt\r\n",
|
||||
"data/SM_BRK.B.pt data/SM_EOG.pt data/SM_SLB.pt\r\n",
|
||||
"data/SM_C.pt data/SM_GS.pt data/SM_WFC.pt\r\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"!ls data/SM_*"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 62,
|
||||
"id": "1b02c605",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"data/voltron_BAC.pt data/voltron_CVX.pt data/voltron_SLB.pt\r\n",
|
||||
"data/voltron_BLK.pt data/voltron_EOG.pt data/voltron_WFC.pt\r\n",
|
||||
"data/voltron_BRK.B.pt data/voltron_GS.pt data/voltron_XOM.pt\r\n",
|
||||
"data/voltron_C.pt data/voltron_JPM.pt data/voltron_v.pt\r\n",
|
||||
"data/voltron_COP.pt data/voltron_MS.pt\r\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"!ls data/vol*"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 63,
|
||||
"id": "66bc45b4",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"tckrs = [\"BAC\", \"BRK.B\", \"CVX\", \"EOG\", \"JPM\", \"XOM\", \"WFC\", \"COP\"]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 64,
|
||||
"id": "2c30e3c6",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"y_list = []\n",
|
||||
"matern_list = []\n",
|
||||
"sm_list = []\n",
|
||||
"volt_list = []\n",
|
||||
"for tckr in tckrs:\n",
|
||||
" if tckr in np.unique(finance_data[\"symbol\"]):\n",
|
||||
" y_list.append(finance_data[finance_data[\"symbol\"] == tckr][\"close_price\"].values)\n",
|
||||
" elif tckr in np.unique(energy_data[\"symbol\"]):\n",
|
||||
" y_list.append(energy_data[energy_data[\"symbol\"] == tckr][\"close_price\"].values)\n",
|
||||
" \n",
|
||||
" matern_list.append(torch.tensor(torch.load(\"data/matern_\"+tckr+\".pt\")))\n",
|
||||
" sm_list.append(torch.tensor(torch.load(\"data/SM_\"+tckr+\".pt\")))\n",
|
||||
" volt_list.append(torch.tensor(torch.load(\"data/voltron_\"+tckr+\".pt\")))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 65,
|
||||
"id": "9d33383b",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"eval_times = list(range(100, y_list[0].shape[-1], 100))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 66,
|
||||
"id": "fcfcdeb0",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAbcAAAEgCAYAAAA39D0QAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAEAAElEQVR4nOz9eYxlaXrWi/6+Na+15yHmOcfKrLmr565uu93GvsYHHy7mGAyyLcuSsQQWQkIgEAILy0LIf1wwAhmQhbigwzX4YJnLNXYb2m73WFVdc1Zm5RzztOdhzdP9Y+3YEZGRWZU1ZE8nHimVmRF7r7322nt97/e+7/M+j0jTNOUUpzjFKU5xiu8jSN/pEzjFKU5xilOc4sPGaXA7xSlOcYpTfN/hNLid4hSnOMUpvu9wGtxOcYpTnOIU33c4DW6nOMUpTnGK7zucBreHQBRFbG5uEkXRd/pUTnGKU5ziFA+B0+D2ENjd3eULX/gCu7u73+lTOcUpTnGKUzwEToPbKU5xilOc4vsO3xPB7b/+1//KxYsX+da3vvWenre3t8c//If/kC984Qs89dRT/OiP/ij/8l/+S4IgeERneopTnOIUp/huwHd9cHv11Vf51V/91ff8vN3dXX7qp36K3/7t36ZYLPKDP/iD2LbNb/zGb/ALv/ALhGH4CM72FKc4xSlO8d2A7+rg9sUvfpFf+IVfwHGc9/zcX/mVX2F3d5e/+Tf/Jr/7u7/Lb/zGb/DFL36RT3/607z44ov8h//wHx7BGZ/iFKc4xSm+G/BdGdx2d3f5O3/n7/DLv/zLJElCvV5/T8+/c+cOf/Inf8Li4iK/9Eu/NP65ZVn82q/9GrIs8x//43/8sE/7FKc4xSlO8V2C78rg9s/+2T/j937v93jiiSf47d/+bc6cOfOenv/Vr36VNE35/Oc/jyQdf4uzs7NcvnyZra0tbt269WGe9ilOcYpTnOK7BMp3+gTuhzNnzvBP/+k/5Sd+4idOBKeHwUHQOn/+/AOP/+abb3Ljxg3OnTv3gc713bC13+I3v/oiH5md5zMXlzE1FVmSEJJAEgJJEsiSQIjszym+e+BHEX4ck1Ozz+zDRJIkRHFCnKTjfydximGomLr6ob7Wg5Cm6ffldy5NU/7bK28h4pTnluYp5k0MXUWRpe/L93uK++O7Mrj94i/+4gd6/v7+PgCTk5P3/f3ExAQAzWbzA73Ou+FPrt3kz/7OfyIWwNWXyP+BxKRssmgUWcwXmS0UKesmOU0lp2nkNQ1dV9FVGVmWTgRBSQiEJJAl6djPJPnoYyQ0VWaqXkRV5Ef6/r6dSJKEr91c4/peg6VqmR9+/PwjWajiJGF70OdGs8Vmt0cQxmiqTNW0qFkWRVUjr2jokkySpKMAlZAkCXGcEo/+Hydp9nec/Tn+2JR7naaGjs92o4cfRcxNlZifro4+a4GQJYTg8PsgjT5zSYz+n/3+4HdCCBJS0hSSNCVJU1LS8b+9KOLK3h7b3QErlTJ/9sIFqpb1oV/LbzcGtsfmfoc//3/+JzaSIakAPZEopSo5VPJCpaRolFSdoqpn952iYSoKpqZiqNkfXZexVAVd1VBlGUWWsk3o6J6UpJPfOyEJJqt5zi5MnAbQ7xJ8Vwa3DwrXdQEwDOO+vz/4+fshqrwX/L3f/59ZYBvBURK2Eocdz+UVp4G1J1OQteyG0w1yqkZR0bAkFVNRsFQVS9HIaSqWmgVAVZVRlezP/W6yA6xtt/n4U8vfNwHuD964ziub2wDs9geYmsbzF5Y/8HGDMGLo+Kw1OlzZ3uVWs0136OKHEWEYcTQGSZKEpsloqkJOU6mYJlXDpKjp5BUVWby/7C4II15a3WTD7ZOS8uagSXF3k3ol/z4XylEAHAXH7O/s+9LzPN5uNrC9ACEEb2m7XF3f45Pz83xscZ6pauF9VUu+U3DcgN1mn51mj412l//P1TdZGwU2AFdKCBOfHgEKEnIkkCKQHAkVCVVIaEJCEzKqkFARSEhICORUoErZY1Qho8qjvyUJTZHRFRlVUdAVBVNVWK5UMHWNuanyd/SanCLD92VwO7g5H7QwHOyaH7VP64VylVfcxvj/CdlOGsAXMb6IGaQRzcBF9QfoSOhCJqfo5DWNnKpgKiqmpCALCQHjm1FFxhgFwJyaBUBDVVEVGUPPPtYrN7d55rH57/md5NfevjsObAd4aXWDT51dQJYfLninaYrrhfSGLnutPo32kN12n7Vul13XxonCY49NSIk4yHogJYUE0ij7d/YzgBRZkVAVGUvVKGgaRcOgqGuokkycZllaCtm/SUfHP3ydRnfIy51t/DBCpIJSqpDra7R7DjOTJUxdfY+blJQkSUkSOBCMiyKP1W6X7eHgWMC2XR9NyLyxvcv13QZnCxXOTdaZrBWoV3IoD3l9v53wgpC95oDdZo/+0CNKEt7Y3+OFzU2+2dkmUQ8+IRAIIrJMNiBGSgVyChJJtgFIBVIKciKQECgIZCGjCIEqJJRURgakVECS/S0AmezxEtnzJQQvb21TKBj8P0+D23cFvi+DmzUqsXied9/f+74PgGmaj/Q8/l8/9b/xJ/98lf3UJxlthg92lIz+jtKESAI5TfAQKGlMLwpRAoGKjCZJGLKCqWbZnCmrWIqCLsnIiQRH5tElQBUyupBZyBe5FE5RLlqszNUe6ft8lHjzzg5fvn33xM/tIOCN9R2eXZk/8bs4Thi6Po32gP3WkEZ7SKs3xHZ8XD9kGAe0Ep9u4hOlCTHpib8P1n9ZQPxue6Bg9IfDSoAQoCsKOTUrNxd0jaKuo46ChWBUYkawPRgQxTECSEVKT4SISBD0Yvq2S6VgoaoKuqqg6wq6lv15mD5gkiS0Bg53u13c5KQ2appCa2jj5AugwNVek33XYalRRJNlquUcU7UCE9XCd7QKEEYx+60Bu80+7Z4NgB9E3Gm1eXFjg7t2n4Zw6YnoyKc32pRIkCQCCYhFSiJASlNEmn0OSKOgJUBOJSRi5IRx9iZEVu5XhIQqBLIQKCILegIxvpcHccQfvXWT//2zT35PZb/fr/i+DG4HvbYH9dQajcaxxz0qVPIWf/WJp3l5c4t9d0gjcEkUgUggkdPxTj4huwmjNCUVKcnohknTlJgYP02woxA9VtBlGS2Ss4VOUbAUBUvTyCkqmiTjuyH9MODtfot0Nev7lAsmleL3Xk/l1maDP3j7xjjbvRcvrW7x2OwkvYHHXqvPfmtAszOk3Xfoex5eFBGmCSEJbhLRHQW0IE14p3glA3lVp6oZ5DQNP4qww5AgjbM/JCAEB1XhAzKQAIQ0Clyj/mgKDAlxgpD9wCGn65RNg1reopqzCP2QRINCziRJEoaOT97QkYREDSNbYBOJmmWiKjJpkpL6Kb4XoioSlqahafI4axfpYZbZGzhstm32AhsDCUPSGC3H2GmEocg4SUwQxjQ8m6V8GYBm4NALPZasEkknpdkZIsQulZLFZLXARDWPoT160kscJzQ6Q3abfVqdIUEUM7A9BrZPe+jwdrvJmt2jJXnYUsQgDEgkEEiQgpSCmoAkBKkATcijfmQCkhjlXaMM7yBDTyHJkjSESEdZWrb5TNIsE5YRSCJBCIEiS0RBjJxKGJLMHbtLq+cwUck/8utzinfG92VwO2BJPojqf/v2bQAuXLjwSM9DCMHjE1Pc7ndIgLqRp2zpTEs5doIhrhzTdF36vk8QR4RJku0s0xSEQCgSEln5K0bgpjF+EiMikEOBISsYsoyuqMhyVhpTZRk/CKmqBk3fgdWMXPMjn7mMrn3vfNybux3+4K0b+EeyDUkI2j2Hnu3S93xsL+DK1U00RSYkyQLZKPOCUSmSCDuJ8NL4/i8kyPoqqkwtZzFfKjFXLVK0DAxdQ1WyHXgKBEGE64e4XkDH8ejY2WfnJiFhkjzU+3KjEHcQsjMYkKYpW+0ezaGNQJDKUMjp1DQTEYMlaSypBRQEaqxwZraOZWgPPLYkBIWcgR+GNNpDAj2iULfIJSZJmmQl1iRFVxRmzDzN2CVJU7aGfewkIpVAjN5GmCbcsju0ApflXAlNkml3bdpdm7fvQLlgMVkrMFnNY77DOb1XJElCu+ew2+yz1+zTH3oMHI++7eG6AWma0ghc3uo36cQeQykkkkCRJfwjn4GcZglVXdYhFXjEGKmCJmTyqoqQBbEMyII4iZESCdKEiMOy8wEBBwFIWdVFCJCRkFOBjISSQsMPSKMQZ5TdXV3d4Qcq92dqn+Lbh++d1e494LOf/SwAX/rSl/jbf/tvHysRbG9vc+3aNebm5h75GADA+XqVyrY5Ii4EJMATc9N8XjEYSBF3Bl2ankvDs+n7Pj3Po+8HhElEGGcLtSorWIoKpAgEipTdWKQpXprgRj5yJGHEMgoSURTTc130gkRV0nn7zi5CCH78B574nui/7Tb7/NGVm3TDw7Ky7frsdYesdjsMwsNarN0PWClXjj0/SGPsJMJNQ46GNCEEqiJhaAqWqZMzNco5k/lykdliVoY7gCrLmIqCLiuH7FQyVqp8wFIVgjRJ8fyI/tClMXRoDGzajsMw9DlYF4XISmLj7I6sBDa0ffbCHnoiMRQRbhQRKzCbL6Ehsed5WYaVyCiRYPVWn5XpKtW8hS4pGJKcnd/oM+0OHK7e2WXoeuzENsMkQpalcTnTMjSeOjfLjz9+kV7P5T9963WcOGSxWMZPY+bmyuQkje7AoTdwCcKYTujR6/ksmkUmdWv8/ekOHLoDhxurexRyBpO1AlO1AjlTf8+fd5qmdAcuO40e6zttWl2Hoe0xcHzSIwHLSyKuDFpseH0SKSXSQZYVFCFoey6jnBUlATnJSvQzRgFTVbFdnyBOsAyNME3QYxlLVkACPa8Ri4wpKyeQJNkYiBdFBEmcbQwgy5rhCFtVAgl8LyEkRkkSeqHPq6tb/MCzp8HtO43v+eC2vb2N67pUKhWq1SoACwsLfPazn+UrX/kK//yf/3P+1t/6W0DLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(eval_times, torch.stack(matern_list).t(), alpha = 0.3, color = palette[0], \n",
|
||||
" label = [\"Matern\", *[None]*(len(tckrs)-1)])\n",
|
||||
"plt.plot(eval_times, torch.stack(sm_list).t(), alpha = 0.3, color = palette[2], \n",
|
||||
" label = [\"SM\", *[None]*(len(tckrs)-1)])\n",
|
||||
"plt.plot(eval_times, torch.stack(volt_list).t(), alpha = 0.3, color = palette[6], \n",
|
||||
" label = [\"Volt\", *[None]*(len(tckrs)-1)])\n",
|
||||
"plt.xlabel(\"Time\")\n",
|
||||
"plt.ylabel(\"P(increase)\")\n",
|
||||
"plt.legend()\n",
|
||||
"sns.despine()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 67,
|
||||
"id": "fcf4e38a",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"eval_times = list(range(0, y_list[0].shape[-1], 100))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 68,
|
||||
"id": "c922b5b5",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"prices = [y_list[i][np.array(eval_times)] for i in range(len(y_list))]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 69,
|
||||
"id": "5b9b9ca3",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"(13,)"
|
||||
]
|
||||
},
|
||||
"execution_count": 69,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"prices[0].shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 70,
|
||||
"id": "1aff2590",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from torch.distributions import Beta\n",
|
||||
"from scipy.special import betainc"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 71,
|
||||
"id": "9b2b8f58",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"bought_func = lambda xs: betainc(17, 8, xs)\n",
|
||||
"\n",
|
||||
"def value_func(vec, prices_at_time_y, base = 10000):\n",
|
||||
" portfolio_value = torch.zeros(12)\n",
|
||||
" portfolio_value[0] = base\n",
|
||||
" for i in range(11):\n",
|
||||
" price_of_stock = portfolio_value[i] * bought_func(vec)[i]\n",
|
||||
" amt_bought = price_of_stock / prices_at_time_y[i]\n",
|
||||
" cash_left = portfolio_value[i] - price_of_stock\n",
|
||||
" portfolio_value[i+1] = cash_left + amt_bought * prices_at_time_y[i+1]\n",
|
||||
" return portfolio_value"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 72,
|
||||
"id": "aee0056a",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"matern_values = torch.stack([value_func(pinc, price) for pinc, price in zip(matern_list, prices)])\n",
|
||||
"sm_values = torch.stack([value_func(pinc, price) for pinc, price in zip(sm_list, prices)])\n",
|
||||
"volt_values = torch.stack([value_func(pinc, price) for pinc, price in zip(volt_list, prices)])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 73,
|
||||
"id": "e63fee88",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"torch.Size([8, 12])"
|
||||
]
|
||||
},
|
||||
"execution_count": 73,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"matern_values.shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 74,
|
||||
"id": "a6c6449b",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"hold_values = torch.stack([10000 / y[0] * torch.tensor(price[1:]) for y, price in zip(y_list, prices)])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 75,
|
||||
"id": "18d479ce",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAw0AAAFWCAYAAAArN3H9AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAEAAElEQVR4nOy9d5xcZ3m3f50+ve1sVS+2mrFwl6kmLlTHphNKAm9MIMY4BX7BzhteU0IKIXkd4E1IICSEONixY1oMtrHBuMvGYGzLsvpKu9q+0+upvz/OzOzO7mzVSlpJ5/p8VjM6c+bMMzNnZp7vc9/39xYcx3Hw8PDw8PDw8PDw8PCYAfFkD8DDw8PDw8PDw8PDY3njiQYPDw8PDw8PDw8Pj1nxRIOHh4eHh4eHh4eHx6x4osHDw8PDw8PDw8PDY1Y80eDh4eHh4eHh4eHhMSueaFjmmKZJf38/pmme7KF4eHh4eHh4eHicoXiiYZkzNDTE5ZdfztDQ0MkeioeHh4eHh4eHxxmKJxo8PDw8PDw8PDw8PGbFEw0eHh4eHh4eHh4eHrPiiQYPDw8PDw8PDw8Pj1nxRIOHh4eHh4eHh4eHx6x4osHDw8PDw8PDw8PDY1Y80eDh4eHh4eHh4eHhMSueaPDw8PDw8PDw8PDwmBVPNHh4eHh4eHh4eHh4zIonGjw8PDw8PDw8PDyWCbZl49j2yR7GNOSTPQAPDw8PDw8PDw+PMxVLN6lkC5TTecrpPHqhBEDbxpXE1nSd5NFN4IkGDw8PDw8PD49liG3bpNNpCoUClUoFexmuPnssHMdxcGwHx7bdS2fS+6oCCXd6XkoNM1hMgSAs+DFEUcTn8xEKhYjH44jisScXeaLBw8PDw8PDw2OZYZomfX19yLJMIpEgEAggiiLCIiaQHicXx7axLRvbtBaUeiQIAmrIv+D33HEcbNumVCqRyWTI5XKsWrUKWT62ab8nGjw8PDw8PDw8lhmpVApN0+ju7vaEwimGbds4po1tLUwkTEaURGRNXdR7LwgCkiQRDocJhUIMDg6SSqXo6OhY8LEm44kGDw8PDw8PD49lRjabZfXq1Z5gOAWwLRunJhBs08JxnAUfQ5REBElClERESUIQl+Z9FwSBtrY2jhw54okGDw8PDw8PD4/TDdM0UVX1ZA/DowVLJRJESUJYYpHQClVVMU3zmI/jiQYPDw8PDw8Pj2WIF2VYHtiWm2pUTzlaqEgQYEIcyG404US+t0v1WJ5o8PDw8PDw8PDw8KDubOQWLjv1wuVFiQQJUZ5IOTodBKAnGjw8PDw8PDw8PM5IGiLBnEg5OhaRUE85Oh1EwlQ80eDh4eHh4eHh4XFG4DgOjmVPpBwtUiQgiUiydFqLhKkce6cHDw8PDw8PDw8PjxPI3XffzaZNm9i0aROvetWrZmx85zgOtmlxzw//h02bNrF582Y+9Sefwqzq8y5iFgQBUZaQNZWB0SGu/6MbGU2PI2sqoiydEYIBvEiDh4eHh4eHh4fHKczo6CjPPPMMF1100aRIgtVIOXKAe++9d97HEwRhwgJVFhEmNdX7/euvp7e39/g8kWWOJxo8PDw8PDw8PDxOSSKRCLlcjh//6Eds3/qyhkiYTKlc5uFHH0FRFAzDmHaMhkioOxvN0nl7pojGmYCXnuTh4eHh4eHh4XHK4Ng2luH2HXjFJTvQNI2fPPAAlmlOEwwAP3/kYSqVCq+89BXuBgEkWUb2qahBP2rIjxLwIauKW6NwhqQbLRRPNHh4eHh4eHh4eJwSOLaDXqxgmxYA/kCAV176CkZGRvj1c8+1vM/9P7kfv9/PZZddBoCkyCgBDVlVsB2bO+64gw984ANccsklbNu2jUsuuYTf/d3f5ZFHHmkcY+fOnWzatIkjR44AcPnll7Np06amxxkaGuKWW27hda97Heeccw6vetWruOmmm+jr65s2pk2bNvG2t72NJ554gquuuoqXvexlvPGNb2RsbIyvfOUrbNq0iZ/97Gc88MADvOc97+G8887joosu4mMf+xh79uxZipdywXiiwcPDw8PDw8PD45TArOrTipevuvIqAH7y0wcAN91IUmQUn4ruWDz2xONcfvnlBELBpvs5jsPHPvYxbrnlFvbt28f27dt57WtfSygU4tFHH+XDH/4wDzzgHjOZTHL11VcTCAQAuOKKK7j66qsbx3rxxRe59tpruf3229E0jde97nW0t7fz3e9+l7e97W0810LQjIyMcP311+P3+3nlK19JJBIhmUw2br/zzjv52Mc+Rj6f51WvehXhcJgHHniA9773vQwPDy/Bq7kwvJoGDw8PDw8PD49TiFS2yO6DQ5TK+skeyoII+FW2rO8iEQ3OvXMLbMvCMsxaXwU30uBYNq++5BX4NI0Hf/pTbvqTT7kWqACCwIMPPki1WuUNr389xVLJvU9NdNx777089NBDnHfeefzbv/0bPp/PfRzb5q/+6q/41re+xW233cYVV1zBhg0b+NKXvsSVV17JkSNHuPnmm1m5ciUAuq5z4403kk6n+fSnP8373//+xpi/973vcdNNN/GHf/iH3Hvvvaiq2rhtdHSUq666ii9/+csIgjCtXuLBBx/kM5/5DL/1W7/VeJwPf/jDPPnkk9x111187GMfW9TruFi8SIOHh4eHh4eHxynE7gOnnmAAKJV1dh8YWtR9HcfBrOjgOFhVA8d2J/6CKBCKhnnlK17J4NAgzz77LGZFx6joGOUq99zzP4RDIS45/yL3/oBtWFRyRarFMpe95rXc+LGPI1pu2pNeqmBVDd569TUADBwdwKwamLqBpZsNwdHo9WDb3H/f/fT19XHllVc2CQaAa6+9lquuuoqjR49y//33T3teH/jABxo1FKLYPC0///zzG4IBQFVV3vWudwHw/PPPL+p1PBY80eDh4eHh4eHh4bGssQ0T23ILoJuyk2oT7quuuBKABx58oHFTLpfjyZ07ed1rX9e0wl/nDVe9nr//2//L+dtf3mj2VsgXePbZZ7mvNsHXDR2zqteESJX6g+vlKnqxjF4o8/hjjwFwwfbzqOZLVPMl9ELZvb1U4dKLdwDw5BNPYlbc49U5a8NGLMPENi3XJnZSs7nt27dPG3M9falUi5qcSLz0JA8PDw8PDw+PU4gtG7pO6fSkheLYDmbVqKUluSk8gtjscPSaV78Gn+bjgZ/+lD+68Q8BePChn2IYBm+46qoZj53L57nr7v/msSce59ChQ4ynxt3j1x2U5tH8rV5f8Fdf+iJ/9aUvzrjf0OAgpj5h+SqKIn5FwyhXm/azavuEAtPTuCRJqg1rYV2slwJPNHh4eHh4eHh4nEIkokFeed6Gkz2ME0a9+NnWXZtVQQBRdpNlRElC8alENJVXv/rV/OSBn7Bn/162btnKTx58gFgsxqWXXoooitOExr79+/nw9R8hnU7Tlmhj29atrF+/ns2bNrNm1Sp+67ffP20srbBst75ix8WXkEgkZtxvw/r1Tf+fy9rVqUUdlosFrCcaPDw8PDw8PDw8liW2WSt+tizsWh2DKMvgljojiAKSqgDwpje/iZ888BN++vOHWL1uLTufeoq3v/3tBKJhAGTNTVGSVBktHOCL//dLpNNprv/967nhho+5x6yt4O+t25oKtePXV/br9QeSgCCK4Di0J9sBuPrNb+Etb3rzkj7/5SIYwKtp8PDw8PDw8PDwWIY4jtPI/7cMdzVfEAVExe3cPJXXve51+P1+7rvvPh588EFM0+SNb3xjy2MLgtCwQf39638fSZaRZAlJkZEUmSee2tkYg+JTUfwail9rFCsrfh9ayI8WDnDJpW7NwuNPPYkWCqCG/KhBH2rAhxLw8Y/f+Cfe9f7f4rv/8wNkTUWuiRxwm8yJtU7UbjREpC6IRGV5re17osHDw8PDw8NjGpZhMrbnCId+/iz9T+2elnft4XG8sWrFz7ZhNdySpNpE2o02NOP3+3nNa17DoUOH+Jd/+ReSySSXXHLJjMfv6nLrKx588MGm7Q899BBf+cpXAKhWm897TdMAKBQKjW1vfvObaW9v55577uE/v/OfiKKIKEmIssTjTzzON//1X9m7dy8vP+/lyJqC7JsoylYCGmrA1+hMrYX8yJorKqamU51sPNHg4eHh4eHh0UR+cJy+J14g2z+CbZpU80WGnz94UoovPc5MHNvGqhrgONimW8sgyiKCJCKp8owT6npk4eDBg7zhDW+YZmM6mQ9+8IMA/NEf/RHve9/7uPHGG3nzm9/MRz7yEaLRKMFgkFwuh65PFJyvWbMGgBtvvJEbb7yRQqGA3+/n1ltvJRQK8bnPfY6rrrqKj33sY7z73e/muuuuQ9d1PvWpT7Fly5aleGlOGp5o8PDw8PDw8ABAL5Q5+ouXGHnxEJZhNt1WzRfJD4yfpJF5nGmYVQPHcbAMC8dxSwkkRUYQBOQW9ql1LrvsskbX5je96U2zPsZv/dZv8cUvfpGtW7eye/dunnzySWRZ5rrrruN73/sel1xyCaZp8vDDDzfu8yd/8idccMEFDA0N8eSTT9Lf3w/AhRdeyPe+9z3e+c53ous6P//5zxkYGOA1r3kN//Zv/8aHPvShJXhVTi6C4y0bLGv6+/u5/PLLefDBBxudBz08PDw8PJYS27RIHxokc2QYmHlaICkyqy49p5Ei4nH82L179ym/Mr1YbNNCL1XAdhu6ObjFy6IsuTUBmjLnMTyaWYrzyfvUe3h4eHh4nMEURtKM7enD0qd7/guCm5DgOK43vmWYpA8Nkjx71Qkdo8eZQ6PzM+755lArfpYlBNFNTfI4OXivvIeHh4eHxxmIUaowtreP0ni25e2BtijJTavJD46TPjTQ2J7tGyHSk0QN+U/UUD3OICzdxLZtHMvGtlyxWhcKsqYsKwvSMw1PNHh4eHh4eJxBOLZNuneITO9QI4IwGVlTSW5aTbA9BkBsTRf5wXHMSt1FxmFsbx8955994gbtcUbg2HajG3K9pkaLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 864x360 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.figure(figsize = (12, 5))\n",
|
||||
"plt.plot(eval_times[1:], matern_values.t(), alpha = 0.3, color = palette[0], \n",
|
||||
" label = [\"Matern\", *[None]*(len(tckrs)-1)])\n",
|
||||
"plt.plot(eval_times[1:], sm_values.t(), alpha = 0.3, color = palette[2], \n",
|
||||
" label = [\"SM\", *[None]*(len(tckrs)-1)])\n",
|
||||
"plt.plot(eval_times[1:], volt_values.t(), alpha = 0.3, color = palette[6], \n",
|
||||
" label = [\"Volt\", *[None]*(len(tckrs)-1)])\n",
|
||||
"plt.plot(eval_times[1:], hold_values.t(), alpha = 0.3, color = palette[4], \n",
|
||||
" label = [\"Hold\", *[None]*(len(tckrs)-1)])\n",
|
||||
"plt.xlabel(\"Time\")\n",
|
||||
"plt.ylabel(\"Portfolio Value\")\n",
|
||||
"plt.legend()\n",
|
||||
"sns.despine()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 76,
|
||||
"id": "3a02d7cb",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAw0AAAFWCAYAAAArN3H9AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAC9ZElEQVR4nOy9d3hc1Z3//7p3ep+RZtQly93YYDDVtEBCSQIhlGQJabuwy6YCSb6bJSS/sJRUNrvf9Tdhn5CETTaFBBZCIFmKjZ3QwRBjsHGRu9U1M5re5869vz9GGmnULNnqOq/n0TNzzz333jMj6d7zPp8maZqmIRAIBAKBQCAQCASjIM/0AAQCgUAgEAgEAsHsRogGgUAgEAgEAoFAMCZCNAgEAoFAIBAIBIIxEaJBIBAIBAKBQCAQjIkQDbMcRVFob29HUZSZHopAIBAIBAKBYIEiRMMsp7u7m0suuYTu7u6ZHopAIBAIBAKBYIEiRINAIBAIBAKBQCAYEyEaBAKBQCAQCAQCwZgI0SAQCAQCgUAgEAjGRIgGgUAgEAgEAoFAMCZCNAgEAoFAIBAIBIIxEaJBIBAIBAKBQCAQjIkQDQKBQCAQCAQCgWBMhGgQCAQCgUAgEAgEYyJEg0AgEAgEAoFAIBgTIRoEAoFAIBAIBIJZglpQ0VR1pocxDP1MD0AgEAgEAoFAIBBAwh8msOcoqqJQuawB96KamR5SCWFpEAgEAoFAIBAIZphUMErPzkOoigJA6GDnrLI4CNEgEAgEAoFAIBDMIJlYku6dBwGt1CYbZpdD0OwajUAgECww8uksyUCEVDBKJprAYDXjW9WE2WWf6aEJBAKBYBrIpzJ0v71/mFXBu7IRSZ496/tCNAgEAsE0k4klS0Ihl0iV7cslUnT8tYXK5Q24m6pnaIQCgUAgmA6UbJ6utw9QyCtl7d4VTdirPDM0qpERokEgEAimGE1VSYfiJAMRksEohVzuWEfQu7+NTCRB1epmZL1uWsYpEAgEgulDVQp0v7OffDpT1u5eVIOrsWqGRjU6QjQIBALBFFDIK6SCUZKBCOlQDLVQmPA5koEw7W+kqT5lCSaHdQpGKRAIBIKZQFNVunceJBsvtzbbayqpXNYwQ6MaGyEaBAKBYJLoj09IBiJkIgkGB7SNjoTZZcPmcyMb9PTuaysTGPl0ho439+Jd1YSzzjtlYxcIBALB9OHfc5R0KFbWZqlwUnXSohka0bERokEgEAiOE03TyMZTxfiEQIRcMj2u4yRZxlrhxOpzY/O60RkHbsUWt53unYfKYh00TSWw5wiZaALviiZk3ewJjBMIBALBxOg90E6iu7eszeSwUXPK0lkV+DwUIRoEAoFgAkw8PqGIzmjA5nVj9bqwVjpHfTAYrGbqz1xFcF8r8c5g2b54Z5BsLEXNKUswWM0n/FkEAoFAML1E2/xEjnaXtRksZmpOXTbr49eEaBAIBIJjUMgppHqL8Qmp3ui4i+0YbZY+a4ILk9OGJEnjOk7WyVSd1IzZZSfY0lp2vVwiRfsbe/Ctbp51mTUEAoFAMDqJnhDBfa1lbTqDntrTlqE3GWZoVONHiAaBQCAYgXwqQ7IvkHlC8QluO7Y+oXCi1gBnnReTw0rPzkNl2TXUQoGenQfJNFZTuax+VpuzBQKBQADpcBz/riNlbZIsU3Pa8jljORaiQSAQCOiLT4glSQajE49PqHRh87mxVrrK4hMmA5PDSsPZJxHYc5SEP1S2L9rWQzaWpPqUJehNxkm9rkAgEAgmh2w8RfeOg2jaYCu1RM3aZZidthkb10QRokEgECxY1IJKOhQjGSwWWivk8uM6Tmc0YvMWhYKlwjHlK/2yXkf1KUswt9kJ7mtjsNUjE03QvnU3VWuWYK10Tuk4BAKBQDAx8pkcXW8fQFXKi7dVrW6ec/dsIRoEAsGCopBTiiIhECEVio0/PsFuxdonFEwO67jjEyYTV2MVJqeV7h2HygKwC3mFrrf34Vlch2dx7YyMTSAQCATlFPIKXdv3DUuYUbG0AUdt5QyN6vgRokEgEMx78qnMQP2EaJLxxidYPI6SUDBYTFM9zHFhdtlpPGc1PbsODcvxHT7cSSaaoHrNkkl3kxIIBALB+FELKt3vHCCfKq/27GqowtNcM0OjOjHEU0UgEMw7SvEJfUJh6E17NGSdDmulq5gW1etCZ5idt0idUU/tacsJH+4ifLizbF86FKNt625q1i7B7LLP0AgFAoFg4aJpGv53D5GJJsra7VUVVK5onKFRnTiz84koEAgEE+SE4hP6sh1NR3zCZCFJEhVL6jC77Ph3HaKQH/CXLeRydPy1Be+KRlyNVTM4SoFAIFh4BPe2kgxGytrMbgdVa5rntPuoEA0CgWDOomTzpHqjxxWfYOurxmxyWqd4lFOLtdJJwzmr6dk5dFVLI7ivlUwkge+kRbO+aJBAIBDMB0KHOol1BsrajHYrNWtnd7Xn8SBEg0AgmFPkkhlSwf74hMSxDwD64xNsvmJF5tkSnzBZ6E1G6k5fQehgJ5HW8kqjCX+IbDxF9SlLMDnmtkASCASC2UysIzjMZVRvMlJ72rJZ6+46Eeb+JxAIBPMeTdNI9ISJtnaTjafGdUwpPsHnxlrpnBc37LGQZJnK5Q2YXDYCu4+gFgqlffl0ho6/7sW7sglnnXcGRykQCATzk2QgQmDv0bI2Wa+ndt2KeVNHZ07YSR5//HFWrlzJX//613H1v/nmm1m5ciVbt24dcX8sFuMHP/gB73//+1m7di3ve9/7+P73v08iMfKqZaFQ4OGHH+aaa65h3bp1nHvuuXzlK1/h8OHDo47h1Vdf5W//9m8555xzOP300/n0pz/NSy+9NK7xCwSCImpBJdrup/XVd/HvOnRMwaAzGnHWV1F72gqa33Mq1acswVFTMe8Fw2DsVR4azj4Jo73cqqCpKoE9R/DvOYJaGJ8bl0AgEAiOTSaaoOfdQwzOzCdJMjWnLsNomxvVnsfDrBcN27dv51vf+ta4+//2t78dc3KeSCT41Kc+xYMPPogkSVx88cVIksQvfvELPvaxjxGPx4cd881vfpO77rqL7u5uLrjgAurr63n66ae57rrr2L1797D+jz/+ODfddBPbt29n7dq1rFu3ju3bt3PzzTfzyCOPjPuzCAQLlUJeIXy4i9ZXdxJsaUXJZEfta7Rb8Syuo+Gs1TRfuBbfqiaslc457zt6IhisZurPXIVjBKtCvDNIx1/3jjujlGBhk8/kiHf1ko6M1xVQIFhY5JIZut85MCSmTqL6lCVY3PMrg92sXn7btGkTd9xxB6nU+NwRWltb+cEPfjBmnw0bNtDS0sL111/PPffcgyzLKIrCN77xDZ588kk2bNjAnXfeWTaGxx9/nDVr1vDLX/4Sh8MBwMMPP8xdd93FHXfcwZNPPlmKhvf7/dx11104HA5++9vfsmLFCgB27NjBTTfdxHe+8x0uvvhiqqurj+crEQjmNUo2R7TVT6wjUOZeMxSLx1mMT/C5MZjnh9l3spF1MlUnNWN22QnubUXTBh5ouUSK9jf2ULVmMTafe+YGKZiVaJpGujdGrCNAMhilf/XUd1KzcG8TCAahZHPF4m358mrPvlVN8/LeOiuX4rq7u7n99tu59dZbUVUVr/fYNylVVbn99tsxGAwsX758xD6xWIxHH30Uu93O1772NeS+lUi9Xs9dd92Fy+XiscceKxMpP//5zwG44447SoIB4IYbbuC8886jpaWlzA3qN7/5DblcjhtvvLEkGADWrl3LzTffTDabFdYGgWAIuWQG/54jtL7yLpHW7hEFgyTJOOq8NJ17MnWnr8DVWCUEwzhw1nmpP2sVBku5iVwtFOjecYDe/e3jzjolmN8o2Xyfhe9dut7Z35cycsDdInKke9RjBYKFRiGv0PX2AZRsebVnz+I6nPW+GRrV1DIrRcOGDRt48sknOfnkk3nkkUdYsmTJMY/52c9+xvbt27nzzjtHFRlvvvkmmUyG9evXY7eXm4xsNhvnnnsumUyGN998EyiKjLfffhu3282ZZ5457HyXXnopAC+++GKprd81qn/fYC677LJh/QWChUwmlqR7x0HaXn+XeGewbDW8H1mnw91UQ9P5J1N1UjMG6/zxD50uTA4rDWefhM3nGbYv0tpN51v7hj34BAuHVChGz85DHH15B6FDHaO6A+bTGXKJ9DSPTiCYfWiqSs/OQ+QS5Z4wjjovFUvqZmhUU8+sdE9asmQJ9913Hx/+8IdL1oCx2Lt3Lz/60Y94//vfz1VXXcXvf//7EfsdOHAAYFRLRL84aWlp4aKLLuLgwYNomsbSpUtHHEd//3379gFFk+6BAweQZXlEodPc3Iwsyxw4cABN0+Z0gQ+B4ERI9caIHO0mHY6N2kdnNOBqqMLZ4FtQgcxThazXUbN2KZHWHnr3tzN4BTkTTdC+dTdVa5ZgrXTO3CAF00YhrxDv6iXWHiCfHn98S8IfpsJumcKRCQSzG03T8O86Muz5ZfW68a1aNEOjmh5m5ZP4M5/5zLj75nI5br/9dpxOJ3ffffeYfQOBYrENn29ks1F/e29v77j6V1VVlfWPRqPkcjkqKiowGoe7Tej1ejweD729vSSTyWHWDoFgPqNpGkl/mMjRsdOm6s0m3IuqcdR6kXWz0hg6p3E3VWN22ejecYhCbsC6UDS178OzuA7P4lqxqDFPyUQTRNsDJHvCI1r2+pEkGVu1B73RUFb7I+kPz+uVVIHgWPTubyfhD5W1mZx2qk9eMu/vm7NSNEyE//f//h8tLS3853/+JxUVFWP27Y9VsFhGXiUxm81l/Y7V32QylfVLp9Nj9h98DSEaBAsFTVWJd/USOdoz5oqm0W7FvagGe7Vn3t94Zxqzy07jOavp2XWIdKh8tSx8uJNsLEnV6sXojHP+ESEAVKVAvLuXWEdwmDvFUAwWM84GH47aSnQGPYWcQqS1h37LVC6ZJpfMzKs0kgLBeIkc7SbLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 864x360 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.figure(figsize = (12, 5))\n",
|
||||
"plt.plot(eval_times[1:], matern_values.mean(0), alpha = 0.3, color = palette[0], \n",
|
||||
" label=\"Matern\")#, *[None]*(len(tckrs)-1)])\n",
|
||||
"plt.plot(eval_times[1:], sm_values.mean(0), alpha = 0.3, color = palette[2], \n",
|
||||
" label=\"SM\")#, *[None]*(len(tckrs)-1)])\n",
|
||||
"plt.plot(eval_times[1:], volt_values.mean(0), alpha = 0.3, color = palette[6], \n",
|
||||
" label = \"Volt\")#, *[None]*(len(tckrs)-1)])\n",
|
||||
"plt.plot(eval_times[1:], hold_values.mean(0), alpha = 0.3, color = palette[4], \n",
|
||||
" label = \"Hold\")#, *[None]*(len(tckrs)-1)])\n",
|
||||
"plt.xlabel(\"Time\")\n",
|
||||
"plt.ylabel(\"Portfolio Value\")\n",
|
||||
"plt.legend()\n",
|
||||
"sns.despine()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 77,
|
||||
"id": "e29486f7",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def running_sharpe_ratio(vec):\n",
|
||||
" returns = vec - 10000\n",
|
||||
" std_returns = torch.stack([returns[..., :i].std(-1) for i in range(vec.shape[-1])]).t()\n",
|
||||
" # need avg return divided by sd of returns?\n",
|
||||
" return returns.cumsum(-1) / std_returns / torch.arange(vec.shape[-1])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 78,
|
||||
"id": "4e00013e",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Automatic pdb calling has been turned OFF\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"%pdb"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 79,
|
||||
"id": "79ea67fd",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"matern_sharpe = running_sharpe_ratio(matern_values)\n",
|
||||
"sm_sharpe = running_sharpe_ratio(sm_values)\n",
|
||||
"volt_sharpe = running_sharpe_ratio(volt_values)\n",
|
||||
"hold_sharpe = running_sharpe_ratio(hold_values)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 80,
|
||||
"id": "288c389b",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# n = torch.arange(matern_sharpe.shape[-1])\n",
|
||||
"# stddev = (n - 1) / (n * (n - 3))\n",
|
||||
"\n",
|
||||
"stddev = matern_sharpe.shape[0]**(-0.5)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 83,
|
||||
"id": "d2b7f9d8",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"Text(0, 0.5, 'Sharpe Ratio')"
|
||||
]
|
||||
},
|
||||
"execution_count": 83,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAgEAAAE7CAYAAABE71I3AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAACUgElEQVR4nOydd4BcVdn/P+fcO3VbNtn0ShI2JIQQeu+IFJGuryI2UF+J+r4IIthFRQRULK/6ExsCiiBFBYRQBAJIFakJkkZIskk22c3uzk6995zfH/fO7MzO9jazu+cDk5m59dzZmft8z/M85zlCa60xGAwGg8Ew7pClboDBYDAYDIbSYESAwWAwGAzjFCMCDAaDwWAYpxgRYDAYDAbDOGVMiwDHcdi8eTOO45S6KQaDwWAwDDv9tXv2MLenpGzZsoWTTjqJW2+9lWnTppW6OQaDwWAwDCvbtm3j/PPPZ+XKlcydO7fX7ce0CGhsbATg/PPPL3FLDAaDwWAYORobG40ImDx5MoDxBBgMBoNhXJD1BGTtX2+MaRFgWRYA06ZNY9asWSVujcFgMBgMI0PW/vXGmE4MNBgMBoPB0D1GBBgMBoPBME4xIsBgMBgMhnGKEQEGg8FgMIxTjAgwGAwGg2GcYkSAwWAwGAzjFCMCDAaDwWAYpxgRYDAYDAbDOGVMFwsaCKlUiqamJtra2nBdt9TNMYwiLMuiqqqKiRMnEgqFSt0cg8Fg6BUjAvJIpVJs2rSJ2tpa5s2bRyAQQAhR6mYNGuUqQCOkHBPXU45orclkMrS2trJp0ybmzJljhIDBYCh7TDggj6amJmpra6mrqyMYDI4pg6mVRrkuWutSN2VMIoQgGAxSV1dHbW0tTU1NpW6SwWAw9IoRAXm0tbVRXV1d6mYMIwLluCilSt2QMU11dTVtbW2lbobBYDD0ihEBebiuSyAQKHUzhg0BIATaVcYrMIwEAgGTT2IwGEYFRgR0YiyFALpC4F2jFx5QRggMA2P9O2QwGMYORgSMU7KGyoQHDAaDYfxiRMA4xoQHDAaDYXxjRMA4ZzSFB8q5bQaDwTAaMSLAAPQtPHDXXXexaNEiFi1axJFHHtlrGOGBBx7IbX/FFVcMuG0bN27kwgsvZMuWLQM+hsFgMBiKMSLAkKM/4YHGxkZefPHFHo/397//fUja9alPfYonn3xySI5lMBgMhg6MCDAUkBMCPYQHsrUUHnzwwW6PE4/Hefzxx4dkyKVJXDQYDIbhwYgAQxHZPAHoOjxw5JFHEgqFWLlyZbfegn/84x8kEgmOOuqo4W6uwWAwGAaIEQGGbukuPBCNRjn66KPZvn07L730Upf73n///USjUY499tiidY7jcNttt3HBBRdwyCGHsPfee3PIIYdw4YUXsmrVqtx2zz77LIsWLWLTpk0AnHDCCSxatKjgWNu2bePrX/86xx13HEuXLuXII4/kiiuu4J133ik676JFizj77LP55z//yUknncQ+++zDKaecws6dO/nJT37CokWL+Mc//sHDDz/Mf/3Xf7Hffvtx0EEHsWLFCt58882BfYgGg8FQxhgRYOiR/PCAVh29/lNOOQXoOiQQi8VYtWoVxx9/POFwuGCd1poVK1bw9a9/nbfeeot9992XY445hsrKSp588kk+8YlP8PDDDwNQV1fH6aefTjQaBeDEE0/k9NNPzx3rjTfe4Mwzz+S2224jFApx3HHHMXnyZO6++27OPvtsXnnllaK27dixg4svvphIJMIRRxxBdXU1dXV1ufV33HEHK1asoK2tjSOPPJKqqioefvhhPvjBD7J9+/YBf44Gg8FQjphZBPtIS8Jhw84Eiczoik9HApK5E0NUBweu97JCAHwRoDXHHnss4XCYlStXcuWVVxZs//DDD5NKpTjllFNob28vWPfAAw/w2GOPsd9++/G73/0uJxKUUlxzzTXcdNNN3HrrrZx44oksWLCA66+/nne9611s2rSJK6+8klmzZgGQTqf53Oc+R3NzM1/96lf50Ic+lDvHPffcwxVXXMH//u//8sADDxAMBnPrGhsbOemkk/jxj3+MEKIo1PHII4/wjW98gw984AO583ziE5/gmWee4c9//jMrVqwY8OdoMBgM5YbxBPSR9aNQAAAkMoqNu5JDcizhyQG0hkg4wlFHHcXWrVuLetx///vfqaqq4uijjy46hlKK448/nssuu6zASyCl5LzzzgNg69atvbbloYce4p133uFd73pXgQAAOPPMMznppJPYsmULK1euLNr3ggsuyOU8SFn4E9h///1zAgAgGAzyvve9D4BXX32113YZDAbDaMKIAEP/EaCV4uR3nwx4vfssLS0tPPXUU5x44okFPfAsp512Gj//+c858MADc8vi8TivvPJKLrSQyWR6bcKzzz4LwCGHHNLl+mxC4nPPPVe0bq+99ur2uPvuu2/Rsmy4IB6P99oug8FgGE2YcEAfmV8XGdXhgKFGCMExRx+dCwlcfvnlgNdDz2QynHrqqd3u29raym233caqVatYv349O3fuzB2zrzQ0NADw7W9/m29/+9vdbrdt27aC91LKHqeLrqqqKlpmWRZgKhYaDIaxhxEBfaQmYrN8drGBGA0oV6GHYax9RUUFRx91NCsfWsmrr77K0qVL+fvf/86ECRM4/PDDu9znP//5Dx/5yEdoamqirq6OffbZhwULFrBkyRLmzp3LOeec06dzZ2P5hx9+OJMmTep2u4ULFxa8701omBkADQbDeMKIAMOgOPnkd7PyoZWsfHAlM6bP4JlnnuGcc87Btrv+an3rW9+iqamJFStW8NnPfrbA6PZnGN7kyZMBL/5/xhlnDO4iDAaDYZxicgIMg+LYY44lEonw0MMP8cgjj+A4DqecfHK322eTCP/7v/+7qNf91FNPAcUVArvqnWdzCp544okuz3PDDTdwxhlncPvtt/f9YgwGg2GcYUSAYVBEIhGOPuooNmzYwG9/91smTZrEgQcc2G3J4WnTpgHeULx8HnvsMX7yk58AkEqlCtaFQl5OQywWyy077bTTmDx5Mvfeey+33nprwfarVq3i17/+NW+++Sb77LPP4C/SYDAYxihGBBgGzcl+z3/9+vWc/O53Iy0LrRTaLc5D+OhHPwrAJZdcwvnnn8/nPvc5TjvtND71qU9RU1NDRUUFra2tpNPp3D5z584F4HOf+xyf+9zniMViRCIRbrjhBiorK7nqqqs46aSTWLFiBe9///u56KKLSKfTfPGLX2Tx4sXD/wEYDAbDKMWIAMOgOeboY4hGIgCcesqpubkHtPYmIcrnAx/4ANdeey1Llixh9erVPPPMM9i2zUUXXcQ999zDIYccguM4BW7+yy+/nAMOOIBt27bxzDPPsHnzZsALCdxzzz2cd955pNNpHn/8cbZu3crRRx/N7373Oz72sY+N2GdgMBgMoxGhx/C4p82bN3PCCSfwyCOP5CrN9cTq1avHZM8xOzqgFJnvGkBrhJQIKcZN9v1Y/S4ZDIbypr92z3gCDMNKx9wDXnhgDGtOg8FgGHUYEWAYdgrCA45bMBGRwWAwGEqHEQGGEUMIAUKgXLfb0QMGg8FgGDmMCDCMKCY8YDAYDOWDEQGGEadw9IAJDxgMBkOpMCLAUDK8kQJ+eEAZr4DBYDCMNEYEGEpKLjzgD2M0QsBgMBhGDiMCDCUnFx5QJjxgMBgMI4kRAYayoXN4wGAwGAzDi5lK2FBWCED74QFX6VyVwfFSadBgMBhGEiMCDGVHLk8AvFwB8MSAlEYMGAwGwxBiRIChbCkQA0qjlQtCIKVXdMgIAoPBYBgcRgQYyp6sGABvQqLszITjbVIig8FgGGrKPjFw9+7dHHnkkSxatKjUTTGUAdmRBNmqg8pxcR3XDC80GAyGAVD2noBvfvObNDY2lroZBp/m3bv5/e9/z+NPPM7mzZtJpVJMmjiRfZcv54z3nsGxxxxTsP2HP/oRnn/+eQAu/PiFXHbppT0e/78v/jSPP/44ADf99nccfPDBXW5nvAMGg8EweMraE3Dvvfdy//33l7oZBp/X33iDU045mV/8v1/Q0tLCvsv25eijj6Z24kQefPBBPn3xp7n8i5d3O7xv5UMrezx+a2srTz/9dL/bZbwDBoPBMDDK1hOwfft2vvWtb7Hffvvxyiuv4LpuqZs0rnEch/+95H9pbWvjqm9exTlnn42UHRpyzZo1fHrFxfzt3nvZe++lfOTDHy7Yv7q6mnfeeYfX33iDvZcs6fIcDz/yMJlMhkAgQCaT6XcbjXdg7KG19sSc0riui8q4KMdFSEG4KoqQZd2PMRjKnrL9BX35y18mlUrxve99r9RNMQD/+te/2Lx5M4cddhjnnXtugQAA2GuvvfjaV74KwB133F60//HHHQ/AypUPdnuOv//9Aaqrq1m2bNmg22u8A6MP5SrcjEMmmSYZSxBvbiO2s4X2Xa3Ed7eRbkvgpDJowM04tDe3oRzTOTAYBkNZioA//OEPrFq1issuu4y5c+eWujkGYFfTLgAE3femDz/8cE479TQOP/yIonXvOvFEbNtm5UMPdblv8+7dPPPsM5x4wgkEAoGhaTQdYiDrBVCu8ioSmmmMS4ZWnrF3UhlSsQTx3Vlj30J8d4xUWztOMo3WYAVs7FAQOxjECgawAjZSSqxAAIGgvakVJ91/r5HBYPAoOxGwadMmrrvuOg499FDOP//8UjfH4LOo3hud8dTTT/H/fvn/iMViRduEQiGuv+46vnTllUXrampqOPSQQ9m4cSNvvvlm0fqHVq7EcRxOPeXUoW+8T847QId3QDneXAVGEAw9Wilcx8VJZUi3J4jvjhHb2UJsVyvx5jaSrTEyyTRagbQt39gHsIJBz9hbPReHkraFDNgkdsdIx5Pmb2gwDICyyglwXZfLL78cIQTf/e53yyqG25TcxOrmR2h3mkrdlH5RYU9kUc3x1AZmDuo48+fP5+yzzuKuu+/mhh/9iJ/9/OcLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 576x360 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.figure(figsize = (8, 5))\n",
|
||||
"plt.plot(eval_times[1:], matern_sharpe.mean(0), label = \"Matern\", color = palette[1])\n",
|
||||
"plt.plot(eval_times[1:], sm_sharpe.mean(0), label = \"SM\", color = palette[3])\n",
|
||||
"plt.plot(eval_times[1:], volt_sharpe.mean(0), label = \"Volt\", color = palette[7])\n",
|
||||
"plt.plot(eval_times[1:], hold_sharpe.mean(0), label = \"Hold\", color = palette[5])\n",
|
||||
"\n",
|
||||
"plt.fill_between(eval_times[1:], \n",
|
||||
" matern_sharpe.mean(0) - 2 * stddev * matern_sharpe.std(0),\n",
|
||||
" matern_sharpe.mean(0) + 2 * stddev * matern_sharpe.std(0),\n",
|
||||
" color = palette[1], \n",
|
||||
" alpha = 0.1)\n",
|
||||
"plt.fill_between(\n",
|
||||
" eval_times[1:], \n",
|
||||
" sm_sharpe.mean(0) - 2 * stddev * sm_sharpe.std(0),\n",
|
||||
" sm_sharpe.mean(0) + 2 * stddev * sm_sharpe.std(0),\n",
|
||||
" color = palette[3],\n",
|
||||
" alpha = 0.1\n",
|
||||
")\n",
|
||||
"plt.fill_between(\n",
|
||||
" eval_times[1:], \n",
|
||||
" volt_sharpe.mean(0) - 2 * stddev * volt_sharpe.std(0),\n",
|
||||
" volt_sharpe.mean(0) + 2 * stddev * volt_sharpe.std(0),\n",
|
||||
" color = palette[7],\n",
|
||||
" alpha = 0.1\n",
|
||||
")\n",
|
||||
"plt.fill_between(\n",
|
||||
" eval_times[1:], \n",
|
||||
" hold_sharpe.mean(0) - 2 * stddev * hold_sharpe.std(0),\n",
|
||||
" hold_sharpe.mean(0) + 2 * stddev * hold_sharpe.std(0),\n",
|
||||
" color = palette[5],\n",
|
||||
" alpha = 0.1\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"plt.legend()\n",
|
||||
"plt.ylabel(\"Sharpe Ratio\")\n",
|
||||
"# plt.ylim((0.5, 3))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 91,
|
||||
"id": "7cca1b71",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import pandas as pd"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 92,
|
||||
"id": "b6734108",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df = pd.DataFrame({\"x\": np.repeat(np.arange(4),8), \n",
|
||||
" \"y\": torch.cat(\n",
|
||||
" [matern_sharpe[..., -1], sm_sharpe[..., -1], volt_sharpe[..., -1], hold_sharpe[..., -1]])\n",
|
||||
" })"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 99,
|
||||
"id": "9b3e838e",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAgoAAAFWCAYAAAAIZbVEAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAA4wUlEQVR4nO3deXxNd/7H8dcNEomIiFiC2ElprEMpWrGWolLbjNbaptWSaosapu3UoBtaisGghqKofU9I7EqipVVrLYktdrFESCS5vz/8ckckJ1xZbm7yfj4efTyac77nnM91xX3f7/d7vsdkNpvNiIiIiKTBwdYFiIiISM6loCAiIiKGFBRERETEkIKCiIiIGFJQEBEREUMKCo9ISEjg3LlzJCQk2LoUERERm1NQeMTFixdp2bIlFy9etHUpIiIiNqegICIiIoYUFERERMSQgoKIiIgYUlAQERERQwoKIiIiYkhBQURERAwpKIiIiIghBQURERExpKAgIiIihvLbugAREZGMOnDgAFOnTuXcuXO2LsVqZcuWZcCAAdSqVcvWpaRJPQoiImL3pkyZYpchAeDcuXNMmTLF1mUYUlAQERERQwoKIiJi9wIDA/H29rZ1GU/F29ubwMBAW5dhSHMURETE7tWqVYtp06Zl+nk7dOiQ4ue1a9dm+jVyOvUoiIiIiCEFBRERETGkoCAiIiKGFBRERETEkIKCiIiIGFJQEBEREUMKCiIiImJI6yiIiPw/e31eQE5/VoDYN/UoiIj8P3t9XkBOf1aA2DcFBRERETGkoQcRsSuPLqkrD0RFRWXZn01eXLZY/ifXBIXExEQWLFjA0qVLiYiIwNnZGV9fX3r37o2fn5+tyxMREbFLuSYojBgxglWrVuHq6srzzz/P/fv3CQ8PZ9euXQwaNIiBAwfaukQRERG7kyuCwvr161m1ahUVK1Zk/vz5eHp6AnD8+HF69OjBlClTaN++PRUqVLBtoSIiInYmVwSF1atXAzB06FBLSACoWrUqHTt25Mcff2TXrl0KCiK50MRWB21dQq7zQYivrUuQHCRXBIVJkyYRGRmZZhC4c+cOAPny5cvmqkREROxfrggKjo6OVKtWLdX2LVu2EBQUhIuLC61atbJBZSIiIvYtVwSFh927d49hw4Zx4sQJTp48SenSpRk7dmyKIQkRERF5MrluwaWoqCiCg4M5efKkZduxY8dsWJGIiIj9ynVBoVSpUuzZs4fw8HAmTpzI/fv3GT16NDNmzLB1aSIiInYn1wUFFxcXihYtSpEiRWjXrh1TpkzBZDLxn//8h7i4OFuXJyIiYldyXVB4VJ06dShXrhwxMTGcPXvW1uWIiIjYFbufzGg2mxk3bhwXLlxg3Lhx5M+f+iU5OjoCkJCQkN3liUgW0z3/IlnL7nsUTCYToaGhrF+/nl27dqXaf/bsWSIiInBxcaFixYo2qFBERMR+2X1QAOjevTsAY8aM4eLFi5btly5dYvDgwSQkJPDaa6/h5ORkqxJFRETskt0PPQD07t2bsLAwtm3bRrt27ahXrx6JiYn8/vvvxMbG0qxZM95//31blykiImJ3ckVQKFCgANOmTePHH39k+fLl7N27FwcHB6pVq0bnzp3p3r07Dg65ovNERB6hZz1kPs37kIfliqAAD57l0KtXL3r16mXrUkRERHINfc0WERERQwoKIiIiYkhBQURERAwpKIiIiIghBQURERExpKAgIiIihnLN7ZEiIpLzdejQwdYlZIg91r927doMHa+gICJ2TYsDiWQtDT2IiIiIIQUFERERMaShBxERsZleddvZuoRcZ97+DZl6PgUFEbErGZ2YZQuPToCzx9cgeZeGHkRERMSQgoKIiIgYUlAQERERQwoKIiIiYkhBQURERAwpKIiIiIghBQURERExpKAgIiIihhQURERExJCCgoiIiBhSUBARERFDCgoiIiJiKEMPhbp+/TqnT5/mzp07uLi4UL58eYoVK5ZZtYmIZKsDBw4wdepUzp07l6XXefQhURlVtmxZBgwYQK1atTL1vCLwlEFh9+7dTJw4kQMHDqTa5+Pjw4cffkizZs0yXJyISHaaMmUKUVFRti7DaufOnWPKlCnMmDHD1qVILmT10MOCBQt48803+f333zGbzbi6ulKiRAmcnZ0xm80cPXqUd955h7lz52ZFvSIiIpKNrAoKBw8e5PPPP8dsNtO7d282bdrE3r172bZtG/v27SMoKIiePXsCMG7cOA4fPpwlRYuIZIXAwEC8vb1tXYbVvL29CQwMtHUZkktZNfQwe/ZszGYzQ4cO5c0330y1v0KFCnzyySeUKlWK8ePHM2/ePL788stMK1ZEJCvVqlWLadOm2boMkRzFqh6FvXv34u7uzhtvvJFuuzfeeAN3d3fCw8MzVJyIiIjYllVBITo6Gm9vb0wmU/ondXDA29ubK1euZKg4ERERsS2rgkKRIkWeeEbwhQsXcHV1faqiREREJGewKijUqVOHa9eusWLFinTbLV++nKtXr1KnTp2M1CYiIiI2ZlVQ6NWrF2azmc8++4zvv/+emJiYFPtjYmKYNWsWI0eOxGQy0atXr0wtVkRERLKXVXc9NGrUiLfeeouZM2cyfvx4vv32W8qWLUuhQoWIiYnh/PnzJCUlYTabCQgI4Pnnn8+qukVERCQbWL0y45AhQ6hcuTJTpkzh3LlznD59OsX+cuXKMWDAAPz9/TOrxjwlu5aQzQpaRlZEJPd5qiWc/f398ff359SpU0RGRlqe9VCxYkUqVaqU2TXmKfa6hCxoGVkRkdwoQw+FqlSpkoKBiIhILmYYFLZt2wZAw4YNKViwYIpt1tDDoawTGBjItGnTOHv2rK1LsZq3tzfvvvuurcuwKQ0diUhuYxgU+vfvj4ODA+vWraNixYqWbY9bbOlhJpMp2573kJiYyMKFC1mxYgWnTp0iMTERb29vXn75ZQICAnBycsqWOjIqK5eQffTRtmvXrs2S6+RlGjoSkdzGMCiULl36QYP8+VNty2kSExMZMGAAW7duxcXFhdq1a5M/f35+//13Jk2axLZt25g7dy7Ozs62LlVERMSuGAaFzZs3P9G2nGDJkiVs3boVHx8fZs6cScmSJQG4fv06AwYMYP/+/UydOpUhQ4bYuFLJ7TR0JCK5TYYmM6YnJiaGM2fOUKNGjay6hEXySpH/+Mc/LCEBwMPDg5EjR9KpUyfWrVunoCBZLquGjjRsJCK2YtXKjNWrV6dnz55P1LZPnz68/fbbT1WUtYoWLUqlSpXSnIRVoUIFAC5fvpwttYiIiOQmVvUomM1mzGbzY9vdvn2bS5cucevWracuzBrTp0833PfHH38AUKpUqWypRUREJDcxDAonT56kT58+JCYmptj++++/p7s0s9lsJiYmhsTERKpUqZJ5lT4Fs9nMpEmTAGjTpo1NaxEREbFHhkGhcuXKtGrVikWLFlm2mUwmEhISiI6OfuyJCxYsyNChQzOnyqf07bffEh4ejqenJwEBATatRURExB6lO/Tw0Ucf8fLLLwMPvp336dOHatWq8cknnxge4+DggIuLC+XKlcPV1TVzq7XCd999x4wZM3B0dGTixIl4eHjYrBYRERF7lW5QKFSoEM8995zl5wYNGuDj45NiW06TkJDAqFGjWLx4MU5OTkyePJkGDRrYuiwRERG7ZNVkxnnz5mVVHZnizp07vP/+++zYsQM3NzemTp2qkCAiIpIBT7WOQmJiIpcuXeLu3bup7oJISEggPj6ey5cvs3nzZr744otMKfRxbt68Sb9+/Th06BBeXl7MmDGDatWqZcu1RUREciurg8LMmTOZOXMmt2/ffqL22REU4uPjefvttzl06BBVqlTh+++/1+2QIiIimcCqoLBx40a++eabJ2pbrlw52rZt+1RFWWvSpEn89ttveHl5MW/ePE1cFBERySRWBYWffvoJgI4dOzJ06FCcnJxo0qQJXbp04dNPP+XixYssXbqUWbNmkZSUlC0rM964ccMyd8LDwyPdHozx48dn2nUfXVLXHtnba9CyxSIi2c+qoHD48GGcnZ0ZOXIkhQoVAqBKlSrs2rWLAgUK4O3tzYcffkihQoWYMGECc+fOZeDAgVlSeLIDBw5w7949AA4dOsShQ4cM22ZmUBAREckLrAoKt27dolKlSpaQAFC1alXWrVvHzZs3KVKkCAC9e/dm2rRphIaGZnlQePHFFzl27FiWXkNERCSvsioouLi4YDKZUmzz9vYGHiz5XK9ePeDBqowVKlTg9OnTmVSmSOayt2GXR9lj/Ro6ErFPVgUFb29vIiIiiI2NxcXFBYDy5ctjNps5cuSIJSgA3Lt3j4SEhMytNgcbOWmOrUvIdUYO6mvrEkRE8jyrHjPdtGlTYmNj+fTTT4mJiQGgZs2aACxbtoz4+HjgwbyByMhIypQpk8nlioiISHayKij07t2bokWLsn79el544QXi4+OpXLkyDRs25MiRI3Tu3JlBgwbRr18/4EGwEBEREftl1dBDsWLFmD17Nh9//DHnz5/H0dERgI8//phevXpx4sQJTpw4AUCZMmUYMGBA5lcskgUGjbe/Mf+cbtJQzUkQyQ2sXpmxevXqLF++nAsXLli2VatWjXXr1rFs2TLOnz9PhQoV6NatG4ULF87UYkVERCR7PdWzHgC8vLxS/Ozp6Un//v1TbEtMTCRfvnxPewkRERGxMavmKFjj119/5dVXX82q04uIiEg2eGyPQkREBMuXL+fkyZOYzWZq1qzJX//6V4oVK5Zm+5s3bzJu3DiWL1+e6smSIiIiYl/SDQqzZ8/mm2++ISkpyfKhv3XrVn744QemTZtG3bp1U7Rfvnw548aN48aNG5jNZstKjSIiImKfDIcefv75Z8aOHUtiYiJFihTBz88PPz8/nJ2duXHjBu+99x537twBIDo6mv79+/Pxxx8THR2N2WzG39+fDRs2ZNsLERERkcxn2KMwf/58APz8/BgLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 576x360 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(figsize = (8, 5))\n",
|
||||
"sns.boxplot(x=\"x\", y=\"y\", data=df, palette=[palette[1], palette[3], palette[7], palette[5]])\n",
|
||||
"plt.ylabel(\"Sharpe Ratio\")\n",
|
||||
"plt.xlabel(\"Method\")\n",
|
||||
"ax.set_xticklabels([\"Matern\", \"SM\", \"Volt\", \"Hold\"])\n",
|
||||
"sns.despine()\n",
|
||||
"plt.savefig(\"sharpe_ratio.pdf\", bbox_inches = \"tight\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "4054cb36",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3.7.4 64-bit ('base': conda)",
|
||||
"language": "python",
|
||||
"name": "python37464bitbaseconda52eab690427c4f7ea56588deee120c46"
|
||||
},
|
||||
"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.7"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,891 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "c07fda77",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns\n",
|
||||
"import torch\n",
|
||||
"import gpytorch\n",
|
||||
"# from voltron.robinhood_utils import GetStockData\n",
|
||||
"import os\n",
|
||||
"# import robin_stocks.robinhood as r\n",
|
||||
"import pickle5 as pickle\n",
|
||||
"\n",
|
||||
"sns.set_style('white')\n",
|
||||
"# style.use('whitegrid')\n",
|
||||
"palette = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"sns.set(palette = palette, font_scale=2.0, style=\"white\", rc={\"lines.linewidth\": 4.0})\n",
|
||||
"\n",
|
||||
"import sys\n",
|
||||
"sys.path.append(\"../\")\n",
|
||||
"from voltron.likelihoods import VolatilityGaussianLikelihood\n",
|
||||
"from voltron.models import SingleTaskVariationalGP as SingleTaskCopulaProcessModel\n",
|
||||
"from voltron.kernels import BMKernel, VolatilityKernel\n",
|
||||
"from voltron.models import BMGP, VoltronGP\n",
|
||||
"from voltron.means import LogLinearMean\n",
|
||||
"from gpytorch.kernels import ScaleKernel, RBFKernel, MaternKernel"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"id": "d7661966",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAcwAAABECAYAAAAMTwWHAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAAClElEQVR4nO3arY4TURjH4XdpIc3CrthlwofgCioQFSgUigSNqMByLXgsqDo8N7ECMQJuAJY0ILZkKbDDIAgfTSh5BWcP2zyPnJMm/yaT/DLTbvV93wcA8Ffnag8AgLNAMAEgQTABIEEwASBhuO5guVxG27bRNE0MBoPT3AQAVXRdF/P5PMbjcYxGo5WztcFs2zam02nxcQDwv5nNZjGZTFaurQ1m0zQREfFx/3b0g+2yyyp5/Ohh7QlF3X/+tPaEop7c28z7MiLi1f0XtScUdePurdoTirr58lHtCUXt3VnUnlDM4WIYD57d+NnA360N5o/XsP1gO/rhxXLrKmquXK89oaiTnUu1JxS1f3Uz78uIiJ24UHtCUXvbu7UnFHXt/NfaE4q6vHtSe0Jxf/op0p9+ACBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIEEwASBBMAEgQTABIGG47qDruoiI2OqOT23MaZu/fV17QlHDxYfaE4p6d/i19oRiFvG59oSi3h8f1Z5Q1Jsvm/0s8ulobTrOvMPF9+/2o4G/2+r7vv/Thw4ODmI6nZZdBgD/odlsFpPJZOXa2mAul8to2zaaponBYHAqAwGgpq7rYj6fx3g8jtFotHK2NpgAwC+b/aIdAP4RwQSABMEEgATBBICEb0f+ZBgefApDAAAAAElFTkSuQmCC\n",
|
||||
"text/plain": [
|
||||
"<Figure size 576x72 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"sns.palplot(palette)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 5,
|
||||
"id": "4d2749ce",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"with open(\"./stock_data.pkl\", \"rb\") as handle:\n",
|
||||
" raw_data = pickle.load(handle)\n",
|
||||
" \n",
|
||||
"tckr = \"VIRT\"\n",
|
||||
"span = \"5year\"\n",
|
||||
"interval = 'day'\n",
|
||||
"T = 5.\n",
|
||||
"\n",
|
||||
"data = raw_data[raw_data[\"symbol\"] == tckr]\n",
|
||||
"\n",
|
||||
"ts = torch.linspace(0, T, data.shape[0])\n",
|
||||
"y = torch.FloatTensor(data['close_price'].to_numpy())\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 6,
|
||||
"id": "f6d81e7f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"eval_times = [100, 200, 300, 400, 500, 600, 700, 800, 900, 1000, 1100, 1200]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 7,
|
||||
"id": "dd9d7cd9",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"prices_at_time_y = y[torch.tensor(eval_times)]\n",
|
||||
"delta_y = prices_at_time_y[1:] - prices_at_time_y[:-1]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "eb26e423",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Load Price Probabilities"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 8,
|
||||
"id": "9cc4a7ac",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"voltron = torch.load(\"voltron_v.pt\")\n",
|
||||
"matern = torch.load(\"matern_v.pt\")\n",
|
||||
"specmix = torch.load(\"specmix_v.pt\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 9,
|
||||
"id": "45f1909a",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAbcAAAEgCAYAAAA39D0QAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAACSKklEQVR4nO2dd3wT9RvHP5fRtOneu5S2tOxZ9t5LtoAIyFARWU4EHD8VFFFREEHFBQiIDNl7b2SWPUoXbSndu02bdb8/StLcXZImbdIm6ff9evmS3F3uvne93Oee5/sMiqZpGgQCgUAg2BC8uh4AgUAgEAimhogbgUAgEGwOIm4EAoFAsDmIuBEIBALB5iDiRiAQCASbg4ibAcjlcqSmpkIul9f1UAgEAoFgAETcDCA9PR19+/ZFenp6XQ+FQCAQCAZAxI1AIBAINodViNvOnTsRFRWFa9euGfW9jIwM/O9//0Pfvn3RsmVLDBw4EGvWrIFUKjXTSAkEAoFgCVi8uMXExGDJkiVGfy89PR3jxo3D1q1b4eLigl69eqGkpASrVq3Cq6++CplMZobREggEAsESsGhxO3r0KF599VWUlpYa/d3PPvsM6enpeOutt7Br1y6sWrUKR48eRZcuXXDlyhVs3LjRDCMmEAgEgiVgkeKWnp6ODz74AHPnzoVSqYSXl5dR309ISMDp06cREhKCmTNnqpeLxWJ8+eWX4PP52LRpk6mHXS9YfvkCTicn6t3mdHIill++YHXHs+Vzq+3j1fa51Ta2fH62cm4WKW4rV67Enj170Lx5c2zduhVhYWFGff/8+fOgaRq9e/cGj8c8xYCAADRt2hRPnz5FXFycKYddL4j2D8CEvf/qvPlPJydiwt5/Ee0fYHXHs+Vzq+3j1fa51Ta2fH62cm6Cuh6ANsLCwvD1119j+PDhHHEyBJVoNWrUSOf+79y5g9jYWERERNRorPWNXiENsWX4GEzY+y/+HjYG3mJHZJWWAABuZqbjy0tn8VHnHqBp4NQT/W9/hvJBx654cfc2fNS5B1r7+KmXm+N4tXksWz5edrEMc9t2ZhzL19EJTTy9cCYlCRP2/ovfBo5EIxfLfkDqQvN3sHnYGPg6OiLS3RNCPl/98N8yfAx6hTSs66Eajea5bRk+BsHOrnCys4Ovo5NVnZtFituMGTNq9P3MzEwAgI+Pj9b13t7eAIDs7OwaHae+0iukITYPG4Nh//6NcoWCs/6D08fMclxd+zXH8WrzWLZ+PM199gxugHvZWfht4Ej4Cn3gJOKb/Hi1Ra+Qhlg3ZCSGbN8EBU0jyNkF3/UegNnHDlrFw18fKoEb/u8WSORy8CgKH3Toit9v37Cac7NIt2RNkUgkAAB7e3ut61XLqxOoQqjgbEqSVmEjEPRxJuUJ3mrbGb5CH0T6iuHqYJHv1wZzPzsTiuctMVOLCjH14G6refhXRSsfP0if/8aVNI3lVy5a1bnZpLipXJkURWldr+rPSvq0Vo/tD+/hy0vn6noYBCvlVkaOTQgbAOyMfcD4HOTsYjUP/6p4mJOtFm4AcLe3t6pzs/67SwtisRgAUFZWpnV9eXk5AMDBwaHWxmQr3Mh4hlcP7WEss+fzoaBpNPXyhru9+a5pXpkE97OzEODkjLTiIpMfr1ymBJ9HQcCndB5LrqChUNIQCU37XmjuczPX8WgaUCgrromCrvi/kvXOmC0pRkpxvvrzydRHcLEfVsMzqHvSiotw+dlTxrKE/DycTk60KhHQRXx+HuNzgbTcqs7NJsVNNdema04tKyuLsR3BMNKLizFm11ZINApIC3g8HBn/CqQKOSbs/RfLew8wy82vmsg+OHYieoU0VH825fEKJHLEZpQiQ5aJ14/s5hzrsy59zeJOq41zM8XxaJqGRKZEUZkCxeVyFJcrUCpV6j3WjcwUfHr5IIQ8HmTKim1zyyRYe/MaZrZpb9Lzqm32PH7IWcajKKsJuKiKk08SGJ9HRERZ1bnZpFtSFSWpK9Q/Pj4eABAZGVlrY7J2yuRyvLh7K1KLChnLfx7wAroEBjMirKrKkTEWbRFa5jieq4MAGbJMTDu4C78NHIlO/g1QVCZHj+BQ/DZwJKYd3IUMWaZZhM3c5wYAT/PLsf9xnN7j7X8ch6f5FZ4NqVyJ3BIZnuSW4V5aMa4kFeJWajESsiXILJIZKGyH8HnHIejkx3wYfnD6mMnvk9pm3Z0YzjKZUolV/Qab5XdQm5xOTsS2h/cYy3o3MN9v3BzYpLh1794dAHDy5EkolcwfYFpaGh48eIDAwECSBmAgNE1j1tEDHBfMW+06YmqL1urP5ngo6ws9NvXxTicn4vUju7FuyCj4CLwRk1KEu2kluPu0GL5CH6wbMgqvH9ltlecGADFZKWrhZh+ve1AoVvcZjmkHd+FoQjyuPynE9eQiPMooRVp+OQrLFBx3oy4oAA/yn+LzK4ewuu9wtPcLRs9A5m/NXeRgNQ9Jbex9/AgxGdq7hPiIHa1KBNio7suGru6M5RFuHmZ9iTU1Vi9uaWlpiI+PR25urnpZcHAwunfvjsTERPzwww/q5aWlpfj444+hUCgwbdq0uhiuVbLy2n/YeO8WY9mA0HAs69Wfs63q5r/2LM0kx772LE2vG8SUx1Md64VGERAJKn8aJVIlPJ0EeKFRhNWeGwA8zM3AuiGj4Cv0QU6JFJmFUsRnleJWahGuJBUiwN4Xn3UcjBsZzyBVGB5sZS/kwctJiFBPe7QIcESHhi7IleXhz8GjEO7kj0gfMXoENgSfqrymaSVF+LRLH5OdW23zz4M7OtfF5eea/G9Xm6juyyxJCWN5uLsHANPfl+aCoq0gZHDy5Mm4cuUKNm/ejOjoaK3r5syZg7lz56qXp6SkYMKECcjKykJkZCQaNmyIGzduICsrCz169MDPP/8MgcAw91Jqair69u2LEydOICgoyKTnZukcTniMETv/gVLjNony8MT5ia/CTUeqhS1wNakQcg1Thc8D2gQ7Q8i3+vdBZBdL8ThTUq3vCngUnER8ONnzK/4v4mu9Jqr5S9X85OPMUkw5tA1XM5PV27zRvAu+7NXLKqMmX9y1FXviHmldN79DFyzt2a+WR2Ra8svK4P3jN+rPIj4fhe98CJ6OCHRLxPruKgMJDg7G9u3bsWrVKpw9exZPnjxBcHAwXnnlFUyZMsVgYavPPMjJwsR9OxnC5ioSYeeol2xa2OQKmiFsAKBQAsm5ZQj3FtfRqExHdrFhHTEoAI7PBczJng9nER8iAU9nio0mxeUKRuCNi70APQLDGeL2X0YCisu7W524lUilOJIUr3M9O8rQGonPz2V8bujqblXCBliJuOmr4K9vnb+/P7766itzDMnmyZVIMHrnVhRKy9XLeBSFLcNeRKSHZx2OzPxkFWvv95dZJIOvi8Kqq2rklsiQVyrXus5eyFNbY84iPsQifrUfaIFuIsZnFwc+uvqH4/uYU1C9NtzKSoeCVwZAxPm+JXM0KR5lGhHDfIpi5IPZhrgxzyHc3V3HlpaL9ftYCCZHrlTi5X07EMd6e/umV3/0bxheR6OqHQokciTnas+PBIC4zFKrTf5XKGkk5TDdkUI+BT4FRPo6oE2wMxr5iOHvKoKTvcCkb+r2Ah4CnJzQzNOfsXy3lnB6S2dXLHPMYxs3Y3yOz8u12ntERXwe87cf7uZRRyOpPkTcCBzmnzqKE6zCuVNbtMa8dh3raES1g2qeyEMs1LmNRKZEcp5u8bNknuaXo1zOfOg28XNElJ8jErLKUCDRbtGZAoqi4OIgQPcA5suRtlwxS0aqUOBgQixj2eut2sLZzk79uVgmRWZpCfurVgX7xZaIG8Hq+eP2Day+cYWxrEtgMFb3G2LQXIs1o5onAus0eazPGQVSKAyNi7cQJFIF0vLLGcv8Xe3gKOLD1UGASF8xYjNKzSpwLvYC9GCJ2/nUZHVXCWvgdHIiCsorr6OXgxhdA0M4D/84luVjbbAttwjiliRYM+dTn2DusYOMZcHOLtg2YixE9SAAJ9BNBFcHASQyZm5kqKcDQ+8UNNSJztYATdNIzJFAU46FfApB7pVBQbUhcC4OfAQ6uSHctbL5sJKmsS8uVs+3LIvdj5kRksMiosDn8RDuxnz4W/u8G2fOjVhuBGslqSAfY3dvV5dIAgCxUIido16Cr6NTHY6sdqFpGmUyZrcDN7EA/q7MoIe0gnKUyfRX6LAUckpkKJAwzynU0x4ClkmqErjicvN0e7AX8GDHpzjWm7W4JhVKJfayxjqyURSAyhwwFexoQ2uiRCpFekmx+jOfohDi4lqHI6oeRNwIKJZKMXrXVmRLmC2A1g0eida+fjq+ZZvIlTQUGppFUYAdn0KQuwhCfqUY0DQ4wRmWiFxJIymHOUfo6iCAp6P2eUVXBwEn0tFUqObdegQyxe34kwQUllu+JfxfWioyNFyoznZ26NMgDADXskmwYsuNbbWFurpByLe+CGEibvUcJU1j2sHduJOVwVj+vy49MTqqSR2Nqu5gW2P2z/O6+DwKDTyYuX15pXLk6wirtxRScssg06g2QgFo6GlfZ/OnLvYChLl4IdCx0hKQKhQ4nKi9DqwlwbYwB4c1gv1zdz17Tsqa59zYVqc1uiQBIm71nsUXTnPCsUdHNsFHXXrU0YjqFra4OWi0tvFyEnJy3JJyJIwkd0uipFyB9EJmzl6AmwgOdnX3Fu7iwAdFUZyoyd2svmiWBk3TnN/JyEaN1f9mC4A1uyU5aQDuRNwIVoa2pqOtfPzw5+ARVleNwFSwg0nsNcSNoig09LTnbJ9RqD3puy6haRoJ2Uy3qUhAmc3laCjqeTdWIeVDCXGMxGhL43ZWBhIL8tWfRXw+BoVVnoO/kzMcNIKu8srKkCuxfLe1NuI4wSTWFykJEHGrt2hrOuojdsTOUePhqJGzU9/guCWFTCvHyV4Ab2fmfFVKXhlkCssKLsksknECQxp6OoDPzmuoZVTzbk09/OBp76hLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(eval_times, voltron, marker = \"x\", markersize = 10, label = \"Voltron\")\n",
|
||||
"plt.plot(eval_times, matern, marker = \"x\", markersize = 10, label = \"Matern\")\n",
|
||||
"plt.plot(eval_times, specmix, marker = \"x\", markersize = 10, label = \"Spectral Mixture\")\n",
|
||||
"\n",
|
||||
"plt.ylabel(\"Num Held\")\n",
|
||||
"plt.xlabel(\"Time\")\n",
|
||||
"plt.legend()\n",
|
||||
"sns.despine()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "90ba8290",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Example Beta(17, 8) which is pretty right skewed."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 10,
|
||||
"id": "3bc03425",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from torch.distributions import Beta\n",
|
||||
"from scipy.special import betainc\n",
|
||||
"\n",
|
||||
"xs = torch.linspace(0, 1, 100)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 11,
|
||||
"id": "611616e3",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"[<matplotlib.lines.Line2D at 0x7fe567fbad10>]"
|
||||
]
|
||||
},
|
||||
"execution_count": 11,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYsAAAEFCAYAAAASWssjAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAAr7UlEQVR4nO3de1wU9f4/8NcusFxEBASRVbygrgjeIC+ooZJ4yjJTs7xQKmmopZmVRy3PsUNHT57zrZ95jqV56eSlQlPLUvMS3jEpNIFECERFkYuKqNyW3Z3fH55dG1hggYXZZV/Px6MHj/nM7Ox7xmleO7fPyARBEEBERFQDudQFEBGR5WNYEBFRrRgWRERUK4YFERHVimFBRES1spe6gPoqKytDSkoKvL29YWdnJ3U5RERWQavVoqCgAD179oSTk5PJn7PasEhJSUFkZKTUZRARWaVt27ahX79+Jk9vtWHh7e0N4MECt23bVuJqiIisQ25uLiIjIw37UFPVOyx27dqFJUuW1Dmd8vLysGbNGpw6dQoFBQXw9fXFmDFj8PLLL0OhUJg8H/2pp7Zt26J9+/Z1rp+IyJbV9fR9vS5wnzt3Du+9916dP5ebm4vnn38esbGxcHNzw/Dhw1FcXIzVq1djxowZqKioqE85RETUyOocFgcPHsSMGTNQUlJS5y979913kZubi/nz52P37t1YvXo1Dh48iMGDByMhIQFbtmyp8zyJiKjxmRwWubm5+POf/4x58+ZBp9PBy8urTl906dIlHD16FB06dMDs2bMN7S4uLli+fDns7OywdevWOs2TiIiahslhsWrVKnz77bfo2bMnYmNj4e/vX6cvOnnyJARBQHh4OORy8dcqlUoEBgbi+vXryMjIqNN8iYio8Zl8gdvf3x8rV67EmDFjquzsTaEPgW7dulU7/+TkZKSnp6Nr1651nj8RkRQEQUBpWQWKS8tRrtY8+K9CA61WB41WB61WB50gQKcToBMEGO3oux59fwf4+6C1u2vDF8BEJodFdHR0g74oPz8fANCmTRuj4/W3cd28ebNB30NEZE6Fd0vw+5V8ZF69iczsAuTk30HerXsoKLyP20XFuHe/HFqdrsnrkstl+GjJcxgX0bdJvq/JnrMoLS0FgGqfGNS31+fCORGRueTfvocfT6fhl5Qr+OW3K8jMtswfsDqdgHXbTza/sNCfupLJZEbH6w/N+C4mImpqhUUl2HnoHPaf+A0JyVesZj/USenZZN/VZGHh4uIC4EGfTsaUl5cDAJydnZuqJCKycemX87BxZzy+PngO5WpNvefj6GCPFi4KODk6wFFhD4WDPRzs7WBnJ4edXAY7uRxy+YMfyvofzJV/N1f3Q7o6Pbsq8eqUofWuua6aLCz01yqquyZRUFAgmo6IqLFk5xbi72v3Y++xlFqntZPLoerUBl07eMPfzwudlK3h49USbVq3hJe7K1q2cIKjwmp7TjJZky2h/i6o6m6NzczMBACoVKqmKomIbExJqRprvjyGtbEnajyS6NVNiZFDemBgr07o26M9Wjg7NmGVlqnJwiIsLAwAEBcXh7feekt0+21OTg5SU1PRrl073jZLRI0iKe065sR8iSs5t42Ob+fjjqnPDMSY8N7wa+vRxNVZvkYJi5ycHJSWlsLDwwOeng8uwPj5+SEsLAwnTpzARx99hAULFgB4cPfT0qVLodVqERUV1RjlEJENEwQBG3fGY/m6H1Ch0VYZ36ubEq9OGYYnwgJhz3fjVKtRwmLRokVISEjA3LlzMW/ePEP7smXLMHnyZKxduxZxcXHo3Lkzzp49i4KCAgwdOhSTJ09ujHKIyEaVlqkxb8V2/HDiQpVx3h6uWPzy43ju8eB6PWhsa5r0qoyfnx927NiB1atX4/jx47hy5Qr8/PwwdepUTJs2Dfb2zf8iERE1jTv3ShH19mb8nHKlyrgXxwzA29FPoGUL098UZ+vqvXeuqYfYmsb5+vriH//4R32/loioVrk37+KFRf/FxUu5ona3Fk7458JxGD2sl0SVWS/+lCeiZuVGQREmvL6+yoXsAP+22PT3F9DBt+keZGtOGBZE1GwU3S/Fi4s/rxIUA3p1xGcrpqKVKx/6rS9e1SGiZqFMXYGX3tlS5dRTxKAAbPtnFIOigXhkQURWT6vV4bXl23Em6bKofeTgAHz6t0g42POW2IbikQURWb1VW+Kw7/hvorZ+QR3w8V8mMSjMhGFBRFbt2M+/Y9XmI6K2bh298dmKqXB2UkhUVfPDsCAiq3WjoAjzlseKuhT38miBrSuj4OHmImFlzQ/DgoisUoVGi1divsLtoocvTJPJZFizdBLa+bhLV1gzxbAgIqv0721Hqzyd/eb0ERgS0kWiipo3hgURWZ0LmTeweov4OsWw/t3w2gvDpSnIBjAsiMiqaLRavPXPXdBodYY2L48WWP32c+wQsBFxzRKRVVm3/SSS0q+L2la8/gxau7tKVJFtYFgQkdXIuJqPDz/7UdQ2elhPPDm0p0QV2Q6GBRFZBUEQ8JfV36O84uHrUD3cXPD3+WMkrMp2MCyIyCocir+IE4kZoraYeaPh5cHTT02BYUFEFq9crcHfPt4rans0pAvGjugjUUW2h2FBRBZv4854UbfjcrkM7859CjKZTMKqbAvDgogsWv7te/hoS5yo7cWnByCgc1uJKrJNDAsismj/t+kwikvVhuFWLZ3xVtRICSuyTQwLIrJYWddvIXZ/oqjtzekj4NGKnQQ2NYYFEVms//f5j9DqHj6p3alda7w4ZqCEFdkuhgURWaT0y3nYffi8qO3N6SP4MiOJMCyIyCL932eHRe+p6N7ZB8881lvCimwbw4KILE5y+vUqr0l9KyqCHQVKiGueiCzOh5+L+3/qrWqHJx4NlKgaAhgWRGRhLmbl4lD8RVHbwhkj+QCexBgWRGRRPv7yuGi4b0B7DO/fTaJqSI9hQUQW4+qN2/j2xyRR29zI4TyqsAAMCyKyGOtiT4ieq+jW0Rt/GhwgYUWkx7AgIotQcPsevqr0tPYrk4bxDigLwX8FIrIIG3fGo1z98MVGyjat8MwIPldhKRgWRCS5klI1tuw5I2qb9XwYFA72ElVElTEsiEhyOw+dQ9H9MsOwu5szJj/ZT8KKqDKGBRFJSqfTYePOeFHbC6MHwMVZIVFFZAzDgogkdfyXDGRcLTAM28nlmDo2VMKKyBiGBRFJakOlo4rRw3tC6d1KomqoOgwLIpJMxtV8HE1IF7W99OxgiaqhmjAsiEgym3aeFg0H9/DDI4EdJKqGasKwICJJ3Csuw9cHz4naZk7gUYWlYlgQkSR2HjyHkjK1YdjHyw1PDu0pYUVUE4YFETU5QRCw5bsEUduUp/rxlakWjGFBRE3u55QrSMvKMwzbyeWY/FR/CSui2jAsiKjJbdkjPqqIGBzA22UtHMOCiJrUrTv3sfdYsqht6piBElVDpmJYEFGT2v7DWagrtIbhjkpPhD3SRcKKyBQMCyJqMjqdDlsrXdh+4ekBfGeFFeC/EBE1mdO/ZuFKzm3DsMLBDs8/ESJhRWQqhgURNZkv9/0iGh4VFoTW7q4SVUN1wbAgoiZx514p9h//TdTGd1ZYD4YFETWJ3Yd/RXnFw9emdvD1wOBgfwkrorpgWBBRoxMEAV/uFZ+CmjjqEV7YtiL8lyKiRpecnoMLmTcMw3K5DM898YiEFVFdMSyIqNF9VenC9vD+Kj6xbWUYFkTUqErLK/DNj+dFbZOf4oVta8OwIKJGdeDkBdwtLjMMe3m0QMSgAAkrovpgWBBRo9rxw1nR8LMjg9kVuRViWBBRo8kpKMLxxAxR24TH+cS2NWJYEFGj2XXoVwiCYBju1U2JHv5tJayI6othQUSNQhAE7PghUdT2HPuBsloMCyJqFGdTs5GZfdMw7GBvh7Ej+khYETUEw4KIGkXlC9sRgwLg2aqFRNVQQzEsiMjsSssrsCcuSdTGrsitG8OCiMzucHxqlWcrhg9QSVgRNRTDgojM7uuD50TDY0f05bMVVo5hQURmdbPwPo4m/C5qm/CnYImqIXNhWBCRWX3z43lodTrDcPfOPgjq6ithRWQODAsiMqtdh34VDT87MhgymUyaYshsGBZEZDbpl/OQlH7dMCyTyTAugs9WNAcMCyIym52VjioeDekCX763ollgWBCRWeh0uqqnoHhhu9lgWBCRWZz+NQs3CooMw85ODhgVFihhRWRODAsiMovKRxWjwoLQwtlRmmLI7BgWRNRgpWVq7D2eImp7diRPQTUnDAsiarBD8Rdxv6TcMNzGsyWGhPhLWBGZG8OCiBps1+FfRcNjR/SBvR2792hOGBZE1CC37tzH0YR0Udv4kX2lKYYaDcOCiBpkz5FkaLQPu/dQdWzD7j2aIYYFETXIzko9zI4f2ZfdezRDDAsiqrdL2Tfx68VrorZxEX2lKYYaFcOCiOqt8oXt0D6d0c7HXZJaqHExLIioXgRBwO5KYcGjiuaLYUFE9ZL421VcybltGHZ0sMfo4T0lrIgak31dJo6Pj8fatWuRlpaGiooKBAUFITo6GmFhYSZ9XqPRIDg4GGq12uh4Hx8fHD9+vC4lEZFEKnfvETE4AK1cnaUphhqdyWGxa9cuLFmyBAqFAqGhodDpdDhz5gxmzpyJmJgYTJw4sdZ5ZGRkQK1Wo0OHDujTp2of9+7u7nUqnoikoa7QYM/RJFHbeJ6CatZMCov8/HwsW7YMLVu2xBdffAGVSgUASEpKQlRUFJYvX47hw4fDx8enxvmkpqYCAMaPH485c+Y0sHQiksrRhN9x526pYdjdzRnhA1USVkSNzaRrFlu3boVarcb06dMNQQEAvXv3xsyZM1FeXo7Y2Nha53PhwgUAQFBQUD3LJSJLsOuQ+NmKp4f3gsKhTme1ycqYFBYnTpwAAERERFQZN3LkSAAw6VqD/siCYUFkvYrul+JQ/EVRG3uYbf5q/SkgCAIyMjIgl8vh71+1F8lLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(xs, betainc(17, 8, xs))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 12,
|
||||
"id": "b117e956",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"bought_func = lambda xs: betainc(17, 8, xs)\n",
|
||||
"held_voltron = 1000 * bought_func(voltron)\n",
|
||||
"held_matern = 1000 * bought_func(matern)\n",
|
||||
"held_specmix = 1000 * bought_func(specmix)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 13,
|
||||
"id": "dcd3b158",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAckAAAEgCAYAAADfdIg7AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAACiEklEQVR4nOydd3wT9f/HX3fZ6V50QRld7DLK3nsJspH1FRSRIfwcyFAQAQcqKqKo+NUvKCBTtuwNBcres7RAC23pXkkz7/dHSJobaZM2aZNyz8eDh+Zzl1u93Ove7897EBRFUeDh4eHh4eFhQVb1AfDw8PDw8DgrvEjy8PDw8PBYgBdJHh4eHh4eC/AiycPDw8PDYwFeJHl4eHh4eCzAi2Qlo9VqkZKSAq1WW9WHwsPDw8NTBrxIVjJpaWno0aMH0tLSqvpQeHh4eHjKgBdJHh4eHh4eC7iESG7btg3R0dG4ePEi5/KkpCS8//776NKlC2JiYjBw4ECsW7cOer2ec/38/Hx888036NOnD5o2bYru3btj6dKlKCws5Fxfp9Nh48aNGDx4MJo3b4527drhvffeQ1JSkt3OkYeHh4fH+XB6kbxy5QqWLFlicfndu3cxfPhw/PvvvwgJCUGnTp2QlpaGJUuWYPbs2az1CwsLMW7cOPz+++8gCAJdu3YFQRBYvXo1Ro0ahYKCAtZ35s+fj4ULFyItLQ0dO3ZEaGgo9u7di6FDh+L27dt2PV8eHh4eHieCcmIOHDhANW/enIqKiqKioqKoCxcu0Jbr9Xpq4MCBVFRUFLVjxw7TeFZWlml8//79tO8sWbKEioqKoubPn0/pdDqKoihKo9FQH374IRUVFUUtXryYdQxRUVHUkCFDqPz8fNP4hg0bqKioKGrgwIGUXq+3+pySk5OpqKgoKjk52erv8PDw8PBUDU5pSaalpWH27NmYMWMG9Ho9/P39OdeLi4vDvXv30Lp1a7z66qumcV9fXyxcuBAAsHbtWtN4fn4+tmzZAnd3d8yZMwckaTh9oVCIhQsXwsvLC1u3boVCoTB953//+x8AYO7cufDw8DCNv/baa2jfvj3u3buH+Ph4+538S8Ky+Dgcf1K6u/r4kyQsi49zyf0tOHESex4klLrOngcJWHDiZIX3VZ3PrSr2V5lU9t+uMqku5+aUIrl8+XLs3LkTjRs3xqZNm1CvXj3O9U6dOgUA6NmzJ2tZy5Yt4efnh0uXLpnmGi9cuIDi4mK0bdsW7u7utPXd3NzQrl07FBcX48KFCwAMonr16lV4e3sjNjaWtQ/jfk+edL0fZ1UTGxyC0bv+sfgjOv4kCaN3/YPY4BCX3F+bkBBM3Lvd4sN9z4METNy7HW1CKr6/6nxuVbG/yqSy/3aVSXU5N2FVHwAX9erVw1dffYVBgwaZrD0uEhIMP5qoqCjO5XXr1kVWVhYePnyImJgY0/qRkZEW9wsA9+7dQ5cuXfDw4UNQFIXw8HDO4zCuf//+fetPjgcA0DWsLjYMGobRu/7BhkHDUM/bF4m52aAo4OrzNHx+9iQ+btcZFAUce2yfAKnZbTpg+I7N+LhdZzSrEWQad8T+3MQCvNeqPcb9uxWzW3VE86BgiAQEBASBs0+f4evzpzC7dSe4iQV22V+VnVvrTmgZFAS1ztBM6FJqKpZfjsO7LTqAArC7DAvQWqY3b4tx/27Fuy06oG1IKFqHhsBXJjUJ5Or+Q/BKZIRd9lWZMH8HtTy8IBYIUMvTyyQiGwYNQ9ewulV9qDbDPLeYGkHIUChQz9sHp1Meu8y5OaVITp482ar1nj9/DgAICAjgXG4cz8zMBABkZGRYtX5WVpZV69eoUYO2Po/1KDUaqHU6dKoVhr6b10HH0bFt9vFDDtm3pe06an8LzxzlHo87Yvd9Vfq5WTiHz+OPOWR/xu16i2VY2L4bFp055rICacQoJq/+sxEKrQYkQWBGy9ZYf+uGS4hIaRjPbcSOLdDodSjSaNAmOBQPc3Nc5tycUiStRalUAgCkUinncuO4cY7R+F+ZTGaX9SUSCW09HstQFIX7OVk4kPQQB5MScCL5MYr5qkM85SRXrcSckwewadAolxZII7FBoVDrdQAAPUXhx4vncWDUOJcQkbLoGlYXTQJq4FTKEwBAfOpT/LfvQJc5N5cWSaMLlCAIzuXUC+vE+F9Hr89Dp0CtwrHHj3AgKQEHHz3Eo7zcqj4knmqEr1ReLQQSAO5lZ0JrltetB4WmAUGlfMO1SC8qon1uWsN1zs2lRVIulwMAiouLOZerVCraetaub7Qcy1pfrVbT1n/ZoSgK1zPScTDpIQ4kJSDuaTLth18WIe7uyFAo0NA/AD7S8l1TH6EIbb390MDDCzIL89l6ioJGr4eAIKCjKIhIEqSFF6GKQlGG/WkpPUiCgJ6iICRICEjH7K8yz01PGf7mzHNz2LUEBbVORxu7fP06ZCKRQ/ZXmVAaNXZ36kUbS7h/Dx5iSRUdkf2gKArfNWkBoy2h0uuhzsxCllAMHx+fUuNOnAGXFskaNWrgzp07yMzMRHh4OGs5c07ROIdonKOs6PplzYm+DGQrlTj8ONFgLSY9RFoRd9UiLmp7eqFxQCBOJT/GXwOGYEBElClYYVm33ja7Y1QqFZ48eQIfHx94enpCJBKxvAAFKhUS83JQz8sHHhIJ67O9yShS4mlhHkLcvCAXiaDQaPCsKA91vLzhbWGaoLxU5rmptXrkFqvwrMhwbm4iMTSUBo/ych1ybgCQoyxGYl42bYwkSAQ7aH+VSVJuDohiJW2MJAjU9PZ1yH1ZmWQrFVAavUgUBRIAodEiOz8f+fn5qFWrFoRC55Ui5z0yK4iMjMSJEyeQkJCANm3a0JZRFIXExEQIBAKTgBqjWo1RrkwePnwIAIiOjgYAREREgCRJ0ziTxMREAJaja12NZfFxiA0OKVWcjjxKxM4HdxEgd8OBpARcSHsGvZXuZolAgC616qB33XD0qRuB1MICjNn9D/4ZMtK0T2ZEnC1CmZ2dDR8fH4t5tVyi4SGRoJ6Xj0PEJLe4mCaQACAXiRDi5mUQE9jv4V6Z50ZRFE0g5SIRRAICbkIp6sDb7ucGGK7l4/xcCEkBtPoSazLIzd0h+6tMClQq5HB4q3ykMoe+wFUGBSoVnuTnlQwQBGQiMUJ9/fEwNxueKg2ys7NNBokz4tx2bhl06tQJAHDkCDvC7vLly8jOzkbLli1NOZGtWrWCVCrF2bNnWcE2RUVFOHv2LORyOVq2bAkApv/PysrC5cuXWfs4fPgwAKBLly52Pa+qwlJeU2phAf66eQ29Nv6FflvW4ZerF7H4zAnEpz4tUyCjff0wo2Ub7Bk+Bs9nzMa/I8bi/2LbIq3IIJBcQmgulGUlI5tTUFAAT09P7mWlWFXmYlLwwuVeUXKLi/EoL5cmkEbkIhFC3A1CmWvBlW8LlX1u2cpimkASAEQCg8XuLZWijpe33c4NKLmWdby8IWNYHCRB2H1/lYnxbycScD+K7f23q0yM5+Ynk9PGJQIBPCQShHv7Il9IIiXjeRUdoXW4tEi2bt0akZGRiIuLw+bNm03j2dnZWLRoEQBg4sSJpnG5XI7BgwcjLy8PixYtMvV01Gq1WLx4MfLz8zFq1ChaoYExY8YAABYtWoTs7BJXz6ZNm3DmzBk0atSIZcW6KpGeIfhvn8EYvesfbLl7Cx+fPILYP39D2C/f4819O3E8+RHKshndRWIMiojGT7364/5bM3Dzzen4rnsf9KkbQROLi6nPSrUUjUJ5MfWZ1cev0+kgsjA/VaTRlPpGbhSTIo3G6v2VRpFGg9qe3iyBNCIXihDm4W2X/VXmuVEUhSKNhib+QgFJc2sbhdKe19LowpUI6CJZrNPafX+VSZFGg7pePtDo2HP3xVqt3e/LysR4XzJnqCUvXnQ8JBLU8/WHWuPcUe4E5QKhmePHj8f58+exfv16VuWb69ev4/XXX4dCoUBMTAxq1KiB8+fPIy8vDyNHjmQVR8/NzcVrr72GpKQk1KpVCw0bNsTt27eRnJyMhg0bYt26dXBzc6N9591338W+ffvg5eWF1q1bIz09HdevX4enpyf+/vtvi8UJuEhJSUGPHj1w5MgR1KxZs/wXxQHkKbW4n67A/cInmLhvW5mCaKRpQCD6vHChtgutBbFA4NDjtMSdO3fQoEGDKtk3FxqdHiqt5cAlkiAgE5EWo6edEa1Oj2LGOcnFAocF6zBJLypCSkGJ+85DLEGUr1+l7NtRFGu1uJXJtqaEJIkYF4oCtURCTjbyVCVWfh0vb5p16Wy/WyYuPScJAE2bNsWWLVuwYsUKxMfH48GDB6hduzbef/99jBgxgrW+t7c3Nm7ciJ9++gmHDx/GsWPHEBwcjEmTJmHKlCksgQSAZcuWoVmzZti6dStOnDgBHx8fDBgwADNnzkSdOnUq4SwrBy+ZEFGBcmxNyC5VIH2kUvSsXQ996kWgV51whLh7lLL2y4uecRGFJAGt2aCeoqDTUxAKXEckNTr6SQlJotIEEgBkDEtSxYh2dUVUFvKFtXo9dHo9BE4e/VkWKh39/JjeAGfHJSzJ6oQzW5JGpuz/F3/cuEQba+gbiP7hkRgUGYlWwaEQOuEP19neSIs1OpooSoQkdHqKNiYgCMjEVWN524pOT0GpoYuSTCRwWDoLF2qdDjcy0mljzQODK1Wo7U16USFSCvI5lzXwC7DosncFKIrCledptFzymBqBEJIl97yz/W6ZuJak81QKzDB7uVCEmc06YUj9SPjIhS7lHqxKmJYkQRAQCQhadKbuhTVZmUJTXpjzZgKCsPq4KYqyy30jIkkICBI6quRYVDotZELXFZLSKk+pdFqXFkmNXk8TSAFh+Pu5Eq51tDyVwo3n9Df1d5p2xpzT/2LDjXu4klyAp7kqzkADnhIoimJVYiIJQEAaipyb4+hr+eabbyI6OhrLli2zav3BgwcjOjqaFjWuZ1jAQElEq5Ho6Gg0bNiQNvbLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(eval_times, held_voltron, marker = \"x\", markersize = 10, label = \"Voltron\")\n",
|
||||
"plt.plot(eval_times, held_matern, marker = \"x\", markersize = 10, label = \"Matern\")\n",
|
||||
"plt.plot(eval_times, held_specmix, marker = \"x\", markersize = 10, label = \"Spectral Mixture\")\n",
|
||||
"plt.ylabel(\"Num Held\")\n",
|
||||
"plt.xlabel(\"Time\")\n",
|
||||
"plt.legend()\n",
|
||||
"sns.despine()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "79788082",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"reward is $\\sum \\Delta_y * Stocks$ "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 14,
|
||||
"id": "a445178c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def reward_risk(a, b, prob_incs):\n",
|
||||
" bought_func = lambda xs: betainc(a, b, xs)\n",
|
||||
" total_held = 1000 * bought_func(prob_incs)\n",
|
||||
" \n",
|
||||
" returns = total_held[1:] * delta_y\n",
|
||||
" #cum_returns = returns.cumsum(0)\n",
|
||||
" return returns.std(), returns.sum()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "89b647d1",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Randomly sample to find pareto fronts. ideally, we'd do BO or something like that but it's 2d."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 15,
|
||||
"id": "b9112d7e",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"vals = torch.rand(2500, 2) * 50\n",
|
||||
"mtps_voltron = torch.tensor([reward_risk(x[0], x[1], voltron) for x in vals])\n",
|
||||
"mtps_matern = torch.tensor([reward_risk(x[0], x[1], matern) for x in vals])\n",
|
||||
"mtps_specmix = torch.tensor([reward_risk(x[0], x[1], specmix) for x in vals])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "1a7cb011",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Voltron returns and SD of returns are normal."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 16,
|
||||
"id": "77c7a090",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"[Text(61.0, 0.5, 'b'), Text(395.7999999999999, 0.5, 'b')]"
|
||||
]
|
||||
},
|
||||
"execution_count": 16,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAwsAAAFdCAYAAABM7vF5AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAEAAElEQVR4nOydd5wU5d3Av8/MbLm9Xrg7OgiIYgNURKOxYoklRo29x5I3UaOxtxhssfckGo01GgtRiZ2IHSyIIlJEejmO43rbOjPP+8dsvZ3d2z04xGS++WzkZp55nmdmZ5/n+T2/JqSUEgcHBwcHBwcHBwcHhx4oP3QHHBwcHBwcHBwcHBy2ThxhwcHBwcHBwcHBwcHBFkdYcHBwcHBwcHBwcHCwxREWHBwcHBwcHBwcHBxscYQFBwcHBwcHBwcHBwdbHGHBwcHBwcHBwcHBwcEW7YfugMN/H+vWrePAAw/MeF4Igcfjoaqqip122olf/epX7LTTTput/eXLlzNq1KjNVp+Dg4PD/yIffvgh06dPZ968eTQ1NeF2u6murmaPPfbg2GOPZccdd0y7Jtv473K58Pl8DBkyhH322YcTTzyRgQMH5tWnH3p+AWeOcfjfQzh5Fhw2N8mD+Y477ojb7U45L6WktbWVNWvWYJomiqJw991387Of/WyT2m1sbOSWW25h7dq1/Otf/9qkuhwcHBz+V9F1ncsvv5w333wTgNraWqqrq+no6KC+vp5QKIQQgrPOOosrr7wy5dps438kEqGlpYX169cjpaSwsJBbb72VQw89NOe+/VDzCzhzjMP/Lo5mwaFfuf/++xkyZIjtubVr13LJJZfw7bffcu211/KTn/yE0tLSPrf18ccf89Zbb7HDDjv0uQ4HBweH/3Xuu+8+3nzzTUaOHMk999zDuHHj4ueCwSBPP/009957L48//jiDBg3itNNOs60n0/i/YcMG7rjjDt544w0uu+wySktL2XPPPfPu55acX8CZYxz+d3F8Fhx+MIYOHcq9996Lpmn4/X7eeOONH7pLDg4ODv/T+P1+nn32WcASGpIFBQCv18t5553H//3f/wHwyCOPYJpmXm3U1tZy9913c8ghhxCJRLj22muJRCKb5waiOPOLg8PmwxEWHH5Qhg4dysiRIwFYsWLFD9wbBwcHh/9tVq1ahd/vx+12s91222Usd/zxxwOWaU59fX3e7QghuOGGG/B4PNTV1fH666/3uc+ZcOYXB4fNgyMsOPzgCCEAy9Y0mXA4zJNPPsmxxx7LhAkTGD9+PL/4xS/4+9//TigUSil7wAEHcPXVVwOwcOFCxo4dywEHHABYNq5jx45l7NixdHd3p7X//fffx88nc9pppzF27Fi++uorpk6dysSJE5k4cSJnnnkmpmly1VVXMXbsWN5++22+++47LrzwQiZPnsxOO+3E4YcfzsMPP0w4HN5sz8nBwcGhv9E0yzo5HA7z6aefZixXW1vLq6++ynvvvUdtbW2f2qqsrGT//fcH4P333+9THb2RaX4BZ45xcMgVx2fB4QdlxYoVLF26FCAlYkVbWxvnnnsu8+fPR1EUhg4ditfrZcmSJSxatIg33niDv//975SXlwOWo5vL5WLVqlX4fD622247BgwYsFn6ePvttzNv3jy23XZb2traGDBgAIqSkLM/++wzLrvsMgBGjhxJQUEBy5Yt49577+Wbb77hr3/962bph4ODg0N/s80221BTU0NDQwO//e1vOeOMMzjyyCPZZptt0spuv/32m9zehAkTePvtt/nyyy83ua6eZJpfwJljHBzywdEsOPxgLF68mIsuuggpJUOGDOGwww6Ln7vqqquYP38+EyZM4J133mHGjBn8+9//ZubMmey2224sXLiQa6+9Nl7+gQce4PzzzweswfSf//wnDzzwwGbp57x583jooYd47bXX+PDDD1PaBfjnP//JT37yEz744ANee+013n///XiZ9957j/nz52+Wfjg4ODj0N5qmcf311yOEoLu7m7/85S8cdthh7L///lx55ZW8/PLLbNy4cbO1N3jwYACam5s3q99CtvkFnDnGwSEfHM2CQ7/yu9/9Li20XTgcZuPGjfEJZ9iwYTz88MN4PB4Avv32W95//33Kysr4y1/+QkVFRfzagQMH8sADDzBlyhRmzpzJd999l9WudnMwYcIEpkyZAoCiKJSVlaWcLysr4/7778fr9caPnX766TzzzDOsWbOGefPmsfPOO/drHx0cHBw2F1OmTOHRRx/lhhtuoK6uDoD169fz6quv8uqrr6IoCpMnT+ayyy7b5MhAhYWF8X+3t7dTVVWV87V9mV/AmWMcHPLFERYc+pUFCxbYHne5XBxyyCHsu+++HHnkkSkD/syZMwHYa6+9UgbxGJWVlUyePJmZM2fy0Ucf9ftAPn78+KznJ02alDKIxxg5ciRr1qyhq6urn3rm4ODg0D/ss88+zJgxg9mzZzNz5kxmz57NmjVrADBNk9mzZ3Pcccdxww03cOKJJ/a5nWRtQsy/IFf6Mr+AM8c4OOSLIyw49CszZ86Mx8EOh8PMmjWLW2+9lTVr1tDd3c0BBxyQNpAvX74cgC+//JKTTjrJtt5169YBsHLlyn7svUVvdqk1NTW2x2ODe75hBR0cHBy2BjRN46c//Sk//elPAaivr2f27Nm8/fbbfPTRR5imydSpU5kwYUKa826uJC90i4uL87q2L/MLOHOMg0O+OMKCwxbD7Xaz//77M27cOI499lg++eQTzj//fJ5++umUXZPY5JGsSs5EZ2dnv/YZSFFf2+FyubKed5KkOzg4/DcwcOBAjj32WI499lg+/fRTfvOb3+D3+5k2bVqanX2uxEKaDhkyxHZhnyu5zi/gzDEODvniCAsOW5yamhruuOMOzj77bL755hv+9Kc/MXXq1Pj5goICAK688krOPvvszdq23aAaDAY3axsODg4OP1YuvfRS5s2bx2WXXZbmFJzMnnvuyXHHHcfTTz/N6tWr+9ze119/DfRuipMrvc0v4MwxDg754kRDcvhB2GuvveJJfZ5//vmUeN7Dhw8HEqpiOxYtWsTixYtzstWMxQ0HbGNSb87IHg4ODg4/Zrq7u1m3bh0fffRRr2Vjzsg9HXJzpb6+Pj72ZxNM8iXb/ALOHOPgkC+OsODwg3H55ZfHbTWnTp0aH2T3228/AGbMmEFLS0vadZ2dnZx55pkcffTRvPXWW/HjyXGpkykpKYn/287+9L333uvzPTg4ODj8NxFbtL/++utZQ3IahsF//vMfAH7yk5/0qa2bbroJ0zQZPXp0PDnb5iLT/ALOHOPgkC+OsODwg1FcXMzll18OWAPso48+CsAee+zB7rvvTkdHB+eff36KiruhoYHf/OY3tLe3M2DAAI488sj4OZ/PB1i7OMkTg8/nizvf3XvvvXEbVF3Xefrpp3nllVf690YdHBwcfiQcfvjhTJgwgXA4zNlnn80zzzyTZre/fPlyfvOb3/Dtt98ybtw4fvazn+XVxqpVq7jooouYOXMmLpeLm266CVVVN+dtZJxfwJljHBzyxfFZcPhB+fnPf860adP44osveOSRRzjyyCMZNmwYd999N7/61a+YP38+hxxyCKNHj0ZRFFasWEEkEqGoqIhHH300xXFtzJgxCCFobGzkkEMOoba2ln/+858AXHTRRVx44YXMmTOHfffdlxEjRlBfX09LSwtnnHEGL7/88hZxZHNwcHDYmtE0jYcffphLLrmE2bNnc/PNN3P77bczdOhQioqKaGxspL6+HrCyIv/5z3/O6IDbMw9CKBRi48aNNDY2AlBUVMSdd97JxIkT++VeMs0vgDPHODjkgaNZcPjB+cMf/oDL5SIUCsUd0WpqanjppZe4/PLL2WGHHairq2PFihVUV1dzwgknMH36dLbffvuUekaOHMnNN9/MsGHDaGxsZO3atTQ1NQFw0EEH8dRTT7HPPvvEJ4QhQ4Zwxx13cM0112zxe3ZwcHDYWikrK+OJJ57gkUce4ZhjjmHIkCE0NzezePFiTNNk33335fbbb+eFF17IGNYTrDwIX331VfyzZMkSdF1n11135eKLL+Y///kPBxxwQL/ei938As4c4+CQD0I6MbccHBwcHBwcHBwcHGxwNAsODg4ODg4ODg4ODrY4woKDg4ODg4ODg4ODgy2OsODg4ODg4ODg4ODgYIsjLDg4ODg4ODg4ODg42OKETk0iGAyyYMECBgwYsNljPjs4OPQNwzBobGxkxx13TAljmAttbW05ZWBNpqioqM8ZaR3+u3DmBAeHrQ9nTtjyOMJCEgsWLOCUU075obvh4OBgw7PPPstuu+2Wc/m2tjamHHQgHZ35TQylpaXMmDHjf35ycHDmBAeHrRlnTthyOMJCErHU8M8++yy1tbU/cG8cHBwANmzYwCmnnBL/feZKV1cXHZ1dPP3n26kZUJXTNQ2NTZz+2yvp6ur6n54YHCycOcHBYevDmRO2PI6wkERMzVxbW8uQIUN+4N44ODgk01czkJqqCgbX5jipSLNPbTj8d+LMCQ4OWy9b85wwffp0nn32Wb7//ntM02TkyJEcc8wxnHrqqWn9XrlyJQ8++CBz586lra2NYcOGccIJJ3DyySejKOmuxR0dHTzyyCO8++671NfXU1VVxcEHH8wFF1xAUVFRWnnDMHjppZd4/vnnWb16NV6vl8mTJ3PRRRcxcuTInO7HcXB2cHD478Y08/s4ODg4OPz30s9zwh133MEVV1zB4sWLmThxInvssQdr1qzh1ltv5aKLLiI5F/J3333HcccdxxtvvMGgQYPYZ5992LBhAzfddBNXXHFFWt1dXV2ceuqpPPbYYwgh2G+//RBC8MQTT3DCCSfQ2dmZds11113HDTfcwIYNG9h7770ZPHgwb775JscccwyLFi3K6Z4czYKDg8N/NRKJzHF3SOIktHdwcHD4b6Y/54QlS5bw+OOPU1FRwXPPPRffuW9oaOCLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 792x360 with 4 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(1, 2, figsize = (11, 5))\n",
|
||||
"f = ax[0].scatter(vals[:,0], vals[:,1], c=mtps_voltron[:,1])\n",
|
||||
"ax[0].set_title(\"Return\")\n",
|
||||
"fig.colorbar(f, ax=ax[0])\n",
|
||||
"\n",
|
||||
"f = ax[1].scatter(vals[:,0], vals[:,1], c=mtps_voltron[:,0])\n",
|
||||
"ax[1].set_title(\"SD Return\")\n",
|
||||
"fig.colorbar(f, ax = ax[1])\n",
|
||||
"plt.tight_layout()\n",
|
||||
"\n",
|
||||
"[ax[i].set_xlabel(\"a\") for i in range(2)]\n",
|
||||
"[ax[i].set_ylabel(\"b\") for i in range(2)]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "97b967a1",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Matern is much more icky, mostly because we keep buying baskets of stocks."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 17,
|
||||
"id": "a5507bf3",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"[Text(61.0, 0.5, 'b'), Text(395.7999999999999, 0.5, 'b')]"
|
||||
]
|
||||
},
|
||||
"execution_count": 17,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAwwAAAFdCAYAAACuMuoAAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAEAAElEQVR4nOyddbwc1dnHv2dm1q5b3BNCIGiCS/CUFim00AKlaEtpobiVUqBIBVpKcShWvECLFi3BgwaIu9t1l5WZOef9Y9Z3dq8koeHt/PLZ3N2ZYzO7cx5/HqGUUnjw4MGDBw8ePHjw4MGDC7T/9gI8ePDgwYMHDx48ePCw9cITGDx48ODBgwcPHjx48JAXnsDgwYMHDx48ePDgwYOHvPAEBg8ePHjw4MGDBw8ePOSFJzB48ODBgwcPHjx48OAhLzyBwYMHDx48ePDgwYMHD3lh/LcX4OH/B9avX8+hhx6a97wQgkAgQE1NDTvttBM/+clP2GmnnTbb/CtWrGDChAmbbTwPHjx48JCJ9957jxdffJHZs2fT1NSE3+9n8ODB7LXXXhx33HHsuOOOOX0K0Qafz0dRUREjR45k2rRpnHjiiQwbNqxfa/pv0x7w6I+H/w0Irw6Dh82B9E17xx13xO/3Z5xXStHa2sratWuRUqJpGrfccgtHHHHEJs3b2NjI7373O9atW8e//vWvTRrLgwcPHjzkwrIsLrvsMl599VUAhg4dyuDBg+no6KC2tpZoNIoQgjPOOIMrrrgio28h2mCaJi0tLWzcuBGlFMXFxfz+97/n29/+dp/X9t+iPeDRHw//W/AsDB42O2677TZGjhzpem7dunVcdNFFzJs3j6uuuor99tuP8vLyAc/1wQcf8Nprr7HDDjsMeAwPHjx48JAff/3rX3n11VcZN24cf/nLX5g8eXLyXCQS4dFHH+XWW2/loYceYvjw4Zxyyimu4+SjDXV1ddx888288sorXHrppZSXl7PPPvv0e51fJ+0Bj/54+N+CF8Pg4WvFqFGjuPXWWzEMg56eHl555ZX/9pI8ePDgwUMe9PT08MQTTwCO4JAuLAAEg0F+9rOf8Ytf/AKA++67Dyllv+YYOnQot9xyC4cffjimaXLVVVdhmubmuYA4PNrjwcOmwRMYPHztGDVqFOPGjQNg5cqV/+XVePDgwYOHfFi9ejU9PT34/X622267vO1++MMfAo6bTm1tbb/nEUJw7bXXEggE2LBhA//+978HvOZ88GiPBw8DhycwePivQAgBOP6l6YjFYvz973/nuOOOY8qUKey6665873vf48EHHyQajWa0PeSQQ7jyyisBWLBgAZMmTeKQQw4BHL/WSZMmMWnSJLq7u3PmX7p0afJ8Ok455RQmTZrEl19+yXXXXcfUqVOZOnUqp59+OlJKfvWrXzFp0iRef/11Fi9ezHnnncfee+/NTjvtxJFHHsm9995LLBbbbPfJgwcPHv6bMAzHczkWi/Hxxx/nbTd06FBeeOEF3n77bYYOHTqguaqrqzn44IMBeOeddwY0Rm/IR3vAoz8ePBSCF8Pg4WvHypUrWbZsGUBGtoq2tjbOOuss5s6di6ZpjBo1imAwyJIlS1i4cCGvvPIKDz74IJWVlYAT4Obz+Vi9ejVFRUVst912DBo0aLOs8aabbmL27Nlsu+22tLW1MWjQIDQtJV9/8sknXHrppQCMGzeOUCjE8uXLufXWW5kzZw733HPPZlmHBw8ePPw3MX78eIYMGUJ9fT3nnnsup512GkcffTTjx4/Pabv99ttv8nxTpkzh9ddfZ9asWZs8Vjby0R7w6I8HD73BszB4+FqxaNEizj//fJRSjBw5ku985zvJc7/61a+YO3cuU6ZM4Y033uDNN9/kpZdeYsaMGey+++4sWLCAq666Ktn+9ttv5+yzzwacTfOpp57i9ttv3yzrnD17NnfeeScvv/wy7733Xsa8AE899RT77bcf7777Li+//DLvvPNOss3bb7/N3LlzN8s6PHjw4OG/CcMwuPrqqxFC0N3dzd133813vvMdDj74YK644gqee+45GhoaNtt8I0aMAKC5uXmzxjEUoj3g0R8PHnqDZ2HwsNlxwQUX5KS2i8ViNDQ0JAnL6NGjuffeewkEAgDMmzePd955h4qKCu6++26qqqqSfYcNG8btt9/O9OnTmTFjBosXLy7oS7s5MGXKFKZPnw6ApmlUVFRknK+oqOC2224jGAwmj5166qk89thjrF27ltmzZ7Pzzjtv0TV68ODBw9eB6dOnc//993PttdeyYcMGADZu3MgLL7zACy+8gKZp7L333lx66aWbnDGouLg4+b69vZ2ampo+9x0I7QGP/njw0Bd4AoOHzY758+e7Hvf5fBx++OEceOCBHH300Rkb+4wZMwDYd999MzbrBKqrq9l7772ZMWMG77///hbfsHfdddeC5/fcc8+MzTqBcePGsXbtWrq6urbQyjx48ODh68e0adN48803+eijj5gxYwYfffQRa9euBUBKyUcffcTxxx/Ptddey4knnjjgedKtCol4g75iILQHPPrjwUNf4AkMHjY7ZsyYkcyFHYvFmDlzJr///e9Zu3Yt3d3dHHLIITkb9ooVKwCYNWsWJ510kuu469evB2DVqlVbcPUOevNFHTJkiOvxxCbe37SCHjx48LC1wzAMDjjgAA444AAAamtr+eijj3j99dd5//33kVJy3XXXMWXKlJyA3r4indktLS3tV9+B0B7w6I8HD32BJzB42KLw+/0cfPDBTJ48meOOO44PP/yQs88+m0cffTRDQ5IgEumm43zo7OzcomsGMszVbvD5fAXPewXUPXjw8P8dw4YN47jjjuO4447j448/5pxzzqGnp4d//vOfOX73fUUi3enIkSNdmfu+oq+0Bzz648FDX+AJDB6+FgwZMoSbb76ZM888kzlz5vCHP/yB6667Lnk+FAoBcMUVV3DmmWdu1rndNs9IJLJZ5/DgwYOH/4+45JJLmD17NpdeemlOoHA69tlnH44//ngeffRR1qxZM+D5vvrqK6B3t5y+ojfaAx798eChL/CyJHn42rDvvvsmi/v84x//yMjpPWbMGCBlGnbDwoULWbRoUZ/8MxO5wwHXvNSbM6uHBw8ePPx/RXd3N+vXr+f999/vtW0iQDk7SLevqK2tTdKFQsJJf1GI9oBHfzx46As8gcHD14rLLrss6Z953XXXJTfTgw46CIA333yTlpaWnH6dnZ2cfvrpHHvssbz22mvJ4+m5qdNRVlaWfO/mc/r2228P+Bo8ePDg4X8FCcb93//+d8F0nbZt85///AeA/fbbb0Bz3XDDDUgp2WabbZIF3DYX8tEe8OiPBw99gScwePhaUVpaymWXXQY4G+n9998PwF577cUee+xBR0cHZ599doZJu76+nnPOOYf29nYGDRrE0UcfnTxXVFQEOBqbdAJQVFSUDLq79dZbk36nlmXx6KOP8vzzz2/ZC/XgwYOH/wc48sgjmTJlCrFYjDPPPJPHHnssx49/xYoVnHPOOcybN4/JkydzxBFH9GuO1atXc/755zNjxgx8Ph833HADuq5vzsvIS3vAoz8ePPQFXgyDh68dxxxzDP/85z/57LPPuO+++zj66KMZPXo0t9xyCz/5yU+YO3cuhx9+ONtssw2aprFy5UpM06SkpIT7778/I2Bt4sSJCCFobGzk8MMPZ+jQoTz11FMAnH/++Zx33nl8/vnnHHjggYwdO5ba2lpaWlo47bTTeO65576WADYPHjx4+KbCMAzuvfdeLrroIj766CNuvPFGbrrpJkaNGkVJSQmNjY3U1tYCTvXku+66K29QbnadhGg0SkNDA42NjQCUlJTwpz/9ialTp26Ra8lHewCP/njw0As8C4OH/wquueYafD4f0Wg0GYA2ZMgQnn32WS677DJ22GEHNmzYwMqVKxk8eDAnnHACL774Ittvv33GOOPGjePGG29k9OjRNDY2sm7dOpqamgA47LDDeOSRR5g2bVpy4x85ciQ333wzv/71r7/2a/bgwYOHbyIqKip4+OGHue+++/j+97/PyJEjaW5uZtGiRUgpOfDAA7npppt4+umn86b8BKdOwpdffpl8LVmyBMuy2G233bjwwgv5z3/+wyGHHLJFr8WN9oBHfzx46A1Cefm3PHjw4MGDBw8ePHjwkAeehcGDBw8ePHjw4MGDBw954QkMHjx48ODBgwcPHjx4yAtPYPDgwYMHDx48ePDgwUNeeAKDBw8ePHjw4MGDBw8e8sJLq5qGSCTC/PnzGTRo0GbPAe3Bg4fCsG2bxsZGdtxxx4zUhX1BW1tbnyqwpqOkpGTAFWk9ePDohQcP/z149OLrhycwpGH+/PmcfPLJ/+1lePDwP40nnniC3Xffvc/t29ramH7YoXR09o8AlJeX8+abb/7PEwEPA4NHLzx4+O/DoxdfHzyBIQ2JsvFPPPEEQ4cO/S+vxoOH/y3U1dVx8sknJ5/DvqKrq4uOzi4evesmhgyq6VOf+sYmTj33Crq6uv6nCYCHgcOjFx48/Pfg0YuvH57AkIaEWXno0KGMHDnyv7waDx7+NzFQ944hNVWMGNpH4qHkgObw4CEBj1548PDfh0cvvj54AoMHDx7+f0BK59XXth48ePDg4X8THr3oNzyBwYMHD/8voFCoPmqCFF6Bew8ePHj4X4VHL/oPT2Dw4MHD/w94GiMPHjx48NAXePSi3/AEBg8ePPz/gJJ99zX1fFI9ePDg4X8XHr3oNzyBYSuEUgrV1Y0IBRHG1vkVKdNC9oTRykoQQvy3l7PVQClF87KNKCWp2XZkr/emvb6F+iXrqR4zhOoxQ/o8j7RtbNPGF/Rv6pL//0BKkHbf23rw8P8ASimHoRHaVrsXKylBgBBerdh0KNtERboR/iDCV7iWgMMXNKMsE62sBqH7+jGP5fw+NO/+J+HRi35j6+RG/0cQ/vATOh9/BruljeB+e1F2yg+Jzp5H25/Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 792x360 with 4 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(1, 2, figsize = (11, 5))\n",
|
||||
"f = ax[0].scatter(vals[:,0], vals[:,1], c=mtps_matern[:,1])\n",
|
||||
"ax[0].set_title(\"Return\")\n",
|
||||
"fig.colorbar(f, ax=ax[0])\n",
|
||||
"\n",
|
||||
"f = ax[1].scatter(vals[:,0], vals[:,1], c=mtps_matern[:,0])\n",
|
||||
"ax[1].set_title(\"SD Return\")\n",
|
||||
"fig.colorbar(f, ax = ax[1])\n",
|
||||
"plt.tight_layout()\n",
|
||||
"\n",
|
||||
"[ax[i].set_xlabel(\"a\") for i in range(2)]\n",
|
||||
"[ax[i].set_ylabel(\"b\") for i in range(2)]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 18,
|
||||
"id": "ca963564",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"[Text(61.0, 0.5, 'b'), Text(395.7999999999999, 0.5, 'b')]"
|
||||
]
|
||||
},
|
||||
"execution_count": 18,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAwoAAAFdCAYAAACjLJpHAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAEAAElEQVR4nOydeZwlVXn+v+dU3aX3dfYFhkU2QUFEUBE3gkuMJrhEiUFNfjExxsTELS5BomIkGo2aROOWYFwhRsQFUVRAQEFW2QaYgdmnp/e++6065/z+OHX3qtv39nQPg9bz+TRM9z1Vdaruvc857/a8whhjiBEjRowYMWLEiBEjRow6yMd6AjFixIgRI0aMGDFixDj8EBsKMWLEiBEjRowYMWLEaEFsKMSIESNGjBgxYsSIEaMFsaEQI0aMGDFixIgRI0aMFsSGQowYMWLEiBEjRowYMVoQGwoxYsSIESNGjBgxYsRogftYTyDGby52797N8573vMjXhRCkUinGx8c5+eST+ZM/+RNOPvnkZbv+tm3bOProo5ftfDFixIjxm4zrrruOK6+8kjvvvJOpqSmSySSrV6/maU97Gueffz5PfOITW45px/OJRILe3l42btzI2WefzR/+4R+ybt26rub0WK8jEK8lMX67IeI+CjFWCvUE/8QnPpFkMtnwujGG2dlZdu7cidYaKSUf+9jHeNGLXnRQ152cnORDH/oQu3bt4n//938P6lwxYsSI8ZsO3/d5+9vfzve//30A1q5dy+rVq1lYWGDfvn2USiWEELz+9a/nne98Z8Ox7Xje8zxmZmbYu3cvxhj6+vq45JJLeMELXtDx3B6rdQTitSRGDIgjCjEOEf71X/+VjRs3hr62a9cu3vrWt/LrX/+a97znPTzjGc9gaGhoyde64YYb+MEPfsBJJ5205HPEiBEjxm8LPvGJT/D973+fLVu28C//8i+ceOKJ1deKxSKXXXYZH//4x/niF7/I+vXree1rXxt6niie379/P5deeinf+973eNvb3sbQ0BBnnXVW1/M8lOsIxGtJjBgQ1yjEOAywadMmPv7xj+O6Lvl8nu9973uP9ZRixIgR47cC+Xyer3zlK4A1GOqNBIB0Os2f/dmf8Rd/8RcAfPazn0Vr3dU11q5dy8c+9jHOO+88PM/jPe95D57nLc8NBIjXkRgxVgaxoRDjsMCmTZvYsmULANu3b3+MZxMjRowYvx149NFHyefzJJNJjj/++Mhxr3zlKwGbjrNv376uryOE4KKLLiKVSrFnzx6++93vLnnOUYjXkRgxlh+xoRDjsIEQArA5p/Uol8v813/9F+effz6nnnoqT37yk/n93/99vvCFL1AqlRrGPve5z+Xv//7vAbj33ns57rjjeO5znwvYXNfjjjuO4447jlwu13L9Bx98sPp6PV772tdy3HHHcfvtt3PxxRdz2mmncdppp/G6170OrTXvete7OO6447j66qt54IEH+Ku/+ivOPPNMTj75ZF784hfzmc98hnK5vGzPKUaMGDGWC65rM5DL5TI333xz5Li1a9fy7W9/m5/85CesXbt2SdcaGxvjOc95DgA//elPl3SOxRC1jkC8lsSIsRTENQoxDgts376dhx56CKBBsWJubo7/9//+H3fffTdSSjZt2kQ6nWbr1q3cd999fO973+MLX/gCIyMjgC12SyQSPProo/T29nL88cezatWqZZnjRz7yEe68806e8IQnMDc3x6pVq5CyZmv/4he/4G1vexsAW7Zsoaenh4cffpiPf/zj3HXXXfzHf/zHsswjRowYMZYLRx11FGvWrGFiYoK//Mu/5MILL+QlL3kJRx11VMvYE0444aCvd+qpp3L11Vfzq1/96qDP1YyodQTitSRGjKUijijEeMxx//3385a3vAVjDBs3buSFL3xh9bV3vetd3H333Zx66qn88Ic/5JprruE73/kO1157Laeffjr33nsv73nPe6rjP/nJT/LGN74RsAT7ta99jU9+8pPLMs8777yTT3/601x11VVcd911DdcF+NrXvsYznvEMfvazn3HVVVfx05/+tDrmJz/5CXffffeyzCNGjBgxlguu6/K+970PIQS5XI5///d/54UvfCHPec5zeOc738m3vvUtDhw4sGzX27BhAwDT09PLWqfQbh2BeC2JEWOpiCMKMQ4J/vqv/7pF1q5cLnPgwIHqIrR582Y+85nPkEqlAPj1r3/NT3/6U4aHh/n3f/93RkdHq8euW7eOT37yk5x77rlce+21PPDAA23za5cDp556Kueeey4AUkqGh4cbXh8eHuZf//VfSafT1b/98R//MV/+8pfZuXMnd955J6eccsqKzjFGjBgxusW5557L5z73OS666CL27NkDwN69e/n2t7/Nt7/9baSUnHnmmbztbW87aAWgvr6+6r/n5+cZHx/v+NilrCMQryUxYhwMYkMhxiHBPffcE/r3RCLBeeedxznnnMNLXvKShkXg2muvBeDpT396A7FXMDY2xplnnsm1117L9ddfv+Lk/uQnP7nt62eccUYDsVewZcsWdu7cSTabXaGZxYgRI8bB4eyzz+aaa67hpptu4tprr+Wmm25i586dAGituemmm3j5y1/ORRddxB/+4R8u+Tr1UYRKPUGnWMo6AvFaEiPGwSA2FGIcElx77bVV/etyucyNN97IJZdcws6dO8nlcjz3uc9tIfdt27YB8Ktf/YpXv/rVoefdvXs3AI888sgKzt5isfzUNWvWhP69QvjdSgrGiBEjxqGE67o861nP4lnPehYA+/bt46abbuLqq6/m+uuvR2vNxRdfzKmnntpSqNsp6je5AwMDXR27lHUE4rUkRoyDQWwoxDjkSCaTPOc5z+HEE0/k/PPP5+c//zlvfOMbueyyyxq8KJUFpT6sHIVMJrOicwYaQtlhSCQSbV+Pm6DHiBHj8YR169Zx/vnnc/7553PzzTfzpje9iXw+zxVXXNGSV98pKrKlGzduDN3Ud4pO1xGI15IYMQ4GsaEQ4zHDmjVruPTSS3nDG97AXXfdxYc//GEuvvji6us9PT0AvPOd7+QNb3jDsl47jGiLxeKyXiNGjBgxDnf83d/9HXfeeSdve9vbWgqA63HWWWfx8pe/nMsuu4wdO3Ys+Xp33HEHsHj6TadYbB2BeC2JEeNgEKsexXhM8fSnP73ayOfrX/96g473EUccAdTCxmG47777uP/++zvK2azohQOhWtTLqewRI0aMGI8H5HI5du/ezfXXX7/o2ErhcXPxbafYt29flePbGSXdot06AvFaEiPGwSA2FGI85nj7299ezdm8+OKLq8T77Gc/G4BrrrmGmZmZluMymQyve93reNnLXsYPfvCD6t/r9ajrMTg4WP13WB7qT37ykyXfQ4wYMWI8HlHZsH/3u99tK7uplOJHP/oRAM94xjOWdK0PfOADaK055phjqo3XlgtR6wjEa0mMGAeD2FCI8ZhjYGCAt7/97YAl3c997nMAPO1pT+OpT30qCwsLvPGNb2wId09MTPCmN72J+fl5Vq1axUte8pLqa729vYD16tQvFr29vdUCvI9//OPVXFTf97nsssv4v//7v5W90RgxYsQ4zPDiF7+YU089lXK5zBve8Aa+/OUvt+Tpb9u2jTe96U38+te/5sQTT+RFL3pRV9d49NFHectb3sK1115LIpHgAx/4AI7jLOdtRK4jEK8lMWIcDOIahRiHBV760pdyxRVXcMstt/DZz36Wl7zkJWzevJmPfexj/Mmf/Al333035513HscccwxSSrZv347nefT39/O5z32uoXjt2GOPRQjB5OQk5513HmvXruVrX/saAG95y1v4q7/6K2699VbOOeccjjzySPbt28fMzAwXXngh3/rWtw5JMVuMGDFiHA5wXZfPfOYzvPWtb+Wmm27igx/8IB/5yEfYtGkT/f39TE5Osm/fPsB2O/63f/u3yGLb5j4HpVKJAwcOMDk5CUB/fz///M//zGmnnbYi9xK1jgDxWhIjxhIRRxRiHDb4h3/4BxKJBKVSqVqMtmbNGi6//HLe/va3c9JJJ7Fnzx62b9/O6tWredWrXsWVV17JCSec0HCeLVu28MEPfpDNmzczOTnJrl27mJqaAuD5z38+//3f/83ZZ59dXSQ2btzIpZdeyrvf/e5Dfs8xYsSI8VhjeHiYL33pS3z2s5/lD/7gD9i4cSPT09Pcf//9aK0555xz+MhHPsI3vvGNSOlOsH0Obr/99urP1q1b8X2fpzzlKfzN3/wNP/rRj3juc5+7ovcSto5AvJbEiLFUCBPrbMWIESNGjBgxYsSIEaMJcUQhRowYMWLEiBEjRowYLYgNhRgxYsSIESNGjBgxYrQgNhRixIgRI0aMGDFixIjRgthQiBEjRowYMWLEiBEjRgtiedQ6FItF7rnnHlatWrXsGs8xYsRYPiilmJyc5IlPfGKDnGEnmJub66j7aj36+/uX3I02xuGPmPtjxHh8IOb+Q4/YUKjDPffcwwUXXPBYTyNGjBgd4itf+Qqnn356x+Pn5uY49/nPYyHT3WIxNDTENddc81u/YPymIub+GDEeXzicuf/KK6/kK1/5Cg8++CBaa7Zs2cIf/MEf8Ed/9EctjojrrruOL3zhC9xzzz1orTnqqKN42ctexgUXXBDqtFhYWOCzn/0sP/7xj9m3bx/j4+P8zu/8Dm9+85vp7+9vGa+U4vLLL+frX/86O3bsIJ1Oc+aZZ/KWt7yFLVu2dHQ/saFQh0r796985SusXbv2MZ5NjBgxorB//34uuOCC6ne2U2SzWRYyWS77t4+wZtV4R8dMTE7xx3/5TrLZbGwo/IYi5v4YMR4fONy5/9JLL+ULX/gCyWSSpz71qTiOw69+9SsuueQSbrnlFj796U8jhADgiiuu4D3veQ9SSk4//XT6+vq44447+NCHPsR1113HZz/7WVy3tk3PZrP80R/9EVu3bmXLli08+9nP5t5Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 792x360 with 4 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(1, 2, figsize = (11, 5))\n",
|
||||
"f = ax[0].scatter(vals[:,0], vals[:,1], c=mtps_specmix[:,1])\n",
|
||||
"ax[0].set_title(\"Return\")\n",
|
||||
"fig.colorbar(f, ax=ax[0])\n",
|
||||
"\n",
|
||||
"f = ax[1].scatter(vals[:,0], vals[:,1], c=mtps_specmix[:,0])\n",
|
||||
"ax[1].set_title(\"SD Return\")\n",
|
||||
"fig.colorbar(f, ax = ax[1])\n",
|
||||
"plt.tight_layout()\n",
|
||||
"\n",
|
||||
"[ax[i].set_xlabel(\"a\") for i in range(2)]\n",
|
||||
"[ax[i].set_ylabel(\"b\") for i in range(2)]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "147e7ddc",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Plot pareto fronts."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 19,
|
||||
"id": "c2939e6d",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from botorch.utils.multi_objective import is_non_dominated"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 20,
|
||||
"id": "cf7a2321",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAxgAAAHbCAYAAABInjxXAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOydeZwkRZn3vxGRmXV0zwUDc3Aqp6ByLHIJ4gCKAl6ICgsouquir+fr7a63rq6vircCIgi6uyqKiLfigQeC7CCCHIqIDDPDAAMzPd1dVZkZEe8fT2R19TVn91zE9/MZ6COrKiurujueeJ7f76e8955IJBKJRCKRSCQSmQL0lj6BSCQSiUQikUgksv0QC4xIJBKJRCKRSCQyZcQCIxKJRCKRSCQSiUwZscCIRCKRSCQSiUQiU0YsMCKRSCQSiUQikciUEQuMSCQSiUQikUgkMmXEAiMSiUQikUgkEolMGbHAiEQikUgkEolEIlNGLDAiWwVnn302Z5999pY+jUgkEolEIpHIJpJs6ROIRACWL1++pU8hEolEIpFIJDIFxA5GJBKJRCKRSCQSmTJigRGJRCKRSCQSiUSmjFhgRCKRSCQSiUQikSkjFhiRSCQSiUQikUhkyogFRiQSiUQikUgkEpkyYoERiUQikUgkEolEpoxYYEQikUgkEolEIpEpIxYYkUgkEolEIpFIZMqIBUYkEolEIpFIJBKZMmKBEYlEIpFIJBKJRKaMWGBEIpFIJBLZ7hgYbHPVz2/u/h8Y9fnAYHsLn2Eksv2SbOkTiEQikUgkEllfBgbb/Pz6O3nWoifiAa3g6l/8iUWH78fM/nr3mLPfdgmLb1vCLvNmsXTFaq685k/87He3s8u82SxdsYpDD9iNyz9ybvc2Sqkt+Kwike2L2MGIRCKRSCSyVdDbbajo7TYMDLY5662X8JoPfp0vfOM3rGmXfOHrv+X/fODrnPXWS1i1psXqNS3OeqsUFwBLV6wG4Ge/uz18vgqAxbct4Z/fegn3PTTImralUzq895vpmUYi2zfKx5+myFbACSecAMA111yzhc8kEolEIpuDgcE2v7jhTp5z/EEArF7T4pTzPs89y1byjleezDnPOZL/vvp6PvCF73PoAbvxhXefyXnv+29uun1J9z52W7ADS5Y/3P38gL0XAnDbXcvW+zwO2n9XPv++s2k2amRG01czJEZjncN5UIDRKnY4IpENIBYYka2CWGBEIpHIo4feEaZ3veoUnvu0g3jWeZ9j2QOru8eMLR6ajYzhVj4t5/Oe1z+P4496HF5BLdEopXAOlAKFopFJ4VFPzbQ8fiSyvRELjMhWQSwwIpFIZPuj6lI8O+glFPD1H97IZd+9gT/dubR7XJYa8sJukXN81TlP49kn/RN4cM7jwnmmiaaZatJEU1rpYsxpJjRrBh27GZHIWoki70gkEolEIlNOpZe46fYlLHlggLOffSQXf/PXnH/pz8YdO5XFxXlnn8hVP7mR5Q+sWuexC+fN5vRnPgnnHRaF857CQn8NtFJ0So9SHmsdw7mnXZTs0EzIUulmJHF0KhKZkFhgRCKRSCQSmRIGBtv84NpbAMXXvveHrl7iIxf8kMuuvI5l67HoXxsz+xsMDLYm/f6CnWejlFqv4gJg2YpVfOuHf+A5zzgMHeoEY8ArhdbQLiyDHYtR4DwMe2gVlkQrmplhZiOlL9NorXFhHiTRCqNj0RF5dBMLjEgkEolEIpvMwGCbM998MTf3jD71Mt3FBcDyB1bxhct/ukH3+/nLf0r/jAbHH30g1ovtrXXgnKVTIKNdocAoPagSvPIMuZLSwaphSI1Ca40CskTRSA3NzMTuRuRRSywwIpFIJBKJbBIDg23OeNPF/OkvExcXU/MYay8uNpb99lrIQQc8VmxqHVjABXmqCwWHc+C9ePsrAAOgaBcW6zyp1jQyjzEGmzs6hcd7T389nZZzjkS2dmKBEYlEIpFIZKOpOhfTWVxsCPvsOR8U/PXv96/z2P32Wsg73/BCao0a1oXuhQccXVE6SHEB8n0fCo629XgFiQbrfbidR2twzrGm48lSQ2Zi5Fjk0Ud810cikUgkEllnyN1k/OKGOycdi5pq6rXRHYGF82aP+nzfvRbynjefwfvfcgb7h0yMeXNnAnDkofsCsHP4fN/HLuTf3/BC+po1KSa8jEJ5Rv/Dy/9deAzvoSwh9+AdaK2wToTgzjnahSUvHQMty/2r2jww0GHVcE67sN3OSCSyvRM7GJFIJBKJPMoYG3I3MNjm5PM+xz1LV7Li4UFe+ryj+Mp3ruO9n5WQu6/+50uZ2V+f8L6ec/xBLH9wDR/84g+m9BzTxFCUI+5S+z52Ie94zXP46Be+y+1/XcprXnIiJ5/wT3z7xzfypa9dw4H77MIH3vICkqwGwAff8kJu/vPf+acnPJbfLr6LYw4/gN/94TYOecJe/O+f/sZBB+5FvSHFBUAJaAepEq1FVWDY8H8I3Q3PqNtQyGdeOXThcXgUUqwYrUB5OoWnlXv6aoa+moki8Mh2T8zBiGwVxByMSCQS2Tz0hty9+9Wn8PynH8LJr/wsS1es6h4zNuRuXUXGYKfkhHPPZ+n9j2zUOc3feTb394jAd547k09/4GX87De38KWvXcO+j13Ie970QnaaXaeTF/x+8d849vD9yK2MYvzvH//CoiP2I00T1nREF1E4TxLcoHAei4IwyqTCmFPuRs5BAZmBVEPhILdSSGhGCgpFT2cjYACl5XbeywiVtTJONatpqCcJ1jm88tSMFBgz6kkUgEe2a2KBEdkqiAVGJBKJTD8TibHXN+SuKjJm9NXGLY6/+I3f8MEvbFwH49wXHMdppxzJVT++kYv/+xoWzJvDJ97zEvr7angPv/r9bRz2hL3YcXadzEjoncPTsQ5bQj2FvlpCPU1AQScIrzulx1qLCu0E5yTjQqkg3AbyckRrkYRKopbK/9tF6FAEquJiIuoh4Lt0Igyv6pZEQZZALTXgoJYqaqlhx/6UWhJTwSPbL3FEKhKJRCKR7Zwqn+Lyq24YJ8Ze35C7xbct4bvX/plTn/pEtFLUUk1qFFf/4k8bVVzsved8nvO0Q3nacQeRW8epTz+MHWY1OeKQvchqNRkxUnDikw9E4aknJhQL4FEoIEfaEe3SU7qSNNF45ymcw3uP0hKY532wnvVgdI+eIvxfI90H68EXUDNyXNnT4Vjbbqx1cls35uvOy32o0qEA7TQud6xplWT9OnYxItstscCIRCKRSGQ7ZmCwzdlvvYTFIfRuY3nVOSdy1GH70cottVRTtB1ZojnusH05aL9dNkjove9j5vOf7ziDuXP68UDdwlDheOpRB4BSaCXZErVEY5Q4NrXDIj1LNForvFH40qGVBNtJJ8JThGKi6kxIceGp1vLWSpgeyHiTY0TMXS33LdLpqLoWhuAgFe6vKjaqgsJNUFxopEgBcM5LEeIsqXFB7K2YUTfS3YhEtjNigRGJRCKRyHbKwGCbs956STdRe2NZsPNsTnvm4XgUndKSJZos0eSlp1FLueTD53L2Wy/htruWrfV+9nvsAp7ztEM49sjH0ayLGLuZGbwzeF9gnVi/1hJNPdUUFjqlFVG0l4W8B1Kl0FpRJka6FHgc0qmoqgKjFI3MkJeAcsjkk6MoxQWqGpVShM4FwVozVBlpAoWVwkEr0VkoRLtRhfEx8nDjRqh8OBZGRq2ki+Lx3rKqBa3SMqeZ0l+Ly7HI9kW0qY1EIpFIZBtmMnvZpStWcfbbNr24AEnIvurHN6KVwnkpLJz3aAXt0pJkKV/60Es4IFjDArzhpU/ndec+vfv5/nst4FPvPouTjz+EOf0N6qmmdJ41rZLhwqKMollL6Ksl1FKD8yo8hiIziv56wqxGQiM1oEAbTX9No1EUTlGWjtJ6CutBeRKjSY0mSxT11FBLFPVES6EQRq2MBq1lFCtRkCQi1k61JlGaJHwPJZa0RisSE74W8EjBUX1N9Xzdhn8u/At3hfOeTuEYalseWJ2zplVgrYx1RSLbA7FkjkQikUhkK6eylV10+H78/Po7edaiJ/Ldn/+JJz1+d85485e5Z+lKlj+0hhc/50gu+87v+eAXf8Ceu+zIPUtXTtk5fPYrP2HHOX0sOuoAhnOHdTKqlFtHohWzZtS58EMv5ge/upUkTTnhyQfgnKfRV6fTKXjakw9kp9l9aA3ee6xXaOVwzoeFu6aeJmjl8SElG69IU01hfTdVOzEKV4jIe0Y9wWgZn+oUiFuTlUIh0QrnoZ5KQWSdJkscdQelpdvt0H6kO2GU2ONKV8OTpprEuq6jlELOo+pMaEQcXrqR2/eOUPWiCZa34b6sd9SAlvM8sMbTSA2NmqEeukPRyjayLRNdpCJbBdFFKhKJRCam11Z24c6zWPbAahYduT+/+P0d47Iidp0/h/t6rGKr46eCx+29kA+//QzSWopG08wMqRFNAwoambQG8lKKBu+h9J52YcmUYscZGbXUdIXN3ksXpAxzT4kZWVQ77xnqWFm0a0W7lN39vHSit/CeTgkz64Ys0bRyS2E9iVFY68iddD4S5emrSzhfUTo61uOsp3CWTulJtJLxKi/ZF1opUqO7AnNQdKzFWdFwuKDT6E38rmWwph2O11K8KKSYmAwd7sOHoqlRU/RlCY1UkxhDohV9NUNiYpER2TaJHYxIJBKJRLYiKsenei3juCfty4vffml3zKkqFn7x+zsARhUXwKjiojo+TQ3FejpFVYwtXPbbayHvfdMLMUlKaaGZynhRXobug4ahjmglmpmIlj0irvbOdRfnvagQm61DhLbzHlMNGPnqGOlK5IUNjyOp2QDOOYY6ntIZrPXUEhmHsonCdhzKO0qn6BSiGVFKk2hHI9NoZVg5VGA9JFrjvKNhjGg8vDynzChahROxtlEYpSi9EzG3gzRR1LQmTRXOl3jk3Fxwq9KMF35XeKSggSoZ3EMGDhnBss4znFtm1E10mopLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 800x500 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.figure(figsize = (8, 5), dpi = 100)\n",
|
||||
"\n",
|
||||
"scaling = torch.ones(2)\n",
|
||||
"scaling[0] = -1.\n",
|
||||
"pf_bool_vol = is_non_dominated(mtps_voltron * scaling)\n",
|
||||
"plt.scatter(*mtps_voltron.t(), alpha = 0.2, c = [palette[1]] * 2500)\n",
|
||||
"plt.scatter(*mtps_voltron[pf_bool_vol].t(), label = \"Voltron Pareto Front\", color = palette[0], marker = \"x\")\n",
|
||||
"\n",
|
||||
"pf_bool = is_non_dominated(mtps_matern * scaling)\n",
|
||||
"plt.scatter(*mtps_matern.t(), alpha = 0.2, c = [palette[3]] * 2500)\n",
|
||||
"plt.scatter(*mtps_matern[pf_bool].t(), label = \"Matern Pareto Front\", color = palette[2], marker = \"x\")\n",
|
||||
"\n",
|
||||
"pf_bool_sm = is_non_dominated(mtps_specmix * scaling)\n",
|
||||
"plt.scatter(*mtps_specmix.t(), alpha = 0.2, c = [palette[5]] * 2500)\n",
|
||||
"plt.scatter(*mtps_specmix[pf_bool_sm].t(), label = \"Matern Pareto Front\", color = palette[4], marker = \"x\")\n",
|
||||
"\n",
|
||||
"plt.ylabel(\"Return\")\n",
|
||||
"plt.xlabel(\"SD Return\")\n",
|
||||
"plt.legend()\n",
|
||||
"\n",
|
||||
"sns.despine()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "522440db",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Show the buy functions corresponding to the pareto optimal sets."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 21,
|
||||
"id": "5e840dd0",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"Text(0, 0.5, '# of Stocks Held')"
|
||||
]
|
||||
},
|
||||
"execution_count": 21,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAiIAAAFWCAYAAABU7/ZgAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAEAAElEQVR4nOz9acx0e1rXj35+4xqq6h6eZ0/djTQI2Al4UEEbMCL+AcWIAiEIMkQ0GuNJ1BN40Rqj6dOiRNEX/RcTY8KRY3BA6RAJAQUhBgkQaBuNL8RmPEoPe3iG+65prfUbz4trVT3j7r27ew/dvetKnux9V62qtapqDd91Xd9B1VorpzrVqU51qlOd6lSvQ+nXewNOdapTnepUpzrVG7dOQORUpzrVqU51qlO9bnUCIqc61alOdapTnep1qxMQOdWpTnWqU53qVK9bnYDIqU51qlOd6lSnet3qBERe40op8YEPfICU0uu9Kac61alOdapTve51AiKvcT377LN8xVd8Bc8+++zrvSmnOtWpTnWqU73udQIipzrVqU51qlOd6nWrExA51alOdapTnepUr1udgMipTnWqU53qVKd63eoERE51qlOd6lSnOtXrVp8UQOSHf/iHedvb3sZ//a//9bHP/9Zv/Rbf+Z3fyZd92Zfxe37P7+FP/ak/xb/8l/+SUspjl1+v1/zDf/gP+aqv+io+//M/ny//8i/n7//9v892u33s8jlnfvAHf5Cv+7qv4/f9vt/Hl3zJl/Ad3/Ed/NZv/dYr9hlPdapTnepUp3oj1ic8EPlv/+2/8V3f9V0v+vz/+l//i2/4hm/gx37sx3jzm9/Ml37pl/Lss8/yXd/1XbzjHe94ZPntdsu3fdu38X3f930opfgjf+SPoJTi+7//+/mmb/omNpvNI6/5W3/rb/HOd76TZ599lj/0h/4Qb3nLW/jxH/9xvv7rv57/+T//5yv6eU91qlOd6lSneiPVJzQQ+cmf/En+wl/4C+z3+8c+X2vlHe94B9vtlu/5nu/h3/ybf8M/+Sf/hJ/4iZ/gbW97Gz/6oz/KT/zETzzwmne/+928//3v5xu/8Rv58R//cf7xP/7H/MRP/ARf+7Vfy6//+q/z7ne/+5Ft+OEf/mE+7/M+j//0n/4T3/u938t73vMe3vWud7Hf7/kbf+NvcAowPtWpTnWqU53qY6tPSCDy7LPP8o53vIO/+lf/KqUUnnjiiccu93M/93O8//3v5+1vfztf+7Vfe3z8xo0bvPOd7wTgB37gB46Pr9drfuiHfojlcslf/+t/Ha3l41treec738n5+Tnvec97HgA+//yf/3MA/sbf+BusVqvj43/mz/wZ/uAf/IO8//3v5xd/8RdfuQ9/qlOd6lSnOtUbqOzrvQGPq3e/+938yI/8CL/7d/9uvvu7v5u/+3f/Lrdu3XpkuZ/92Z8F4Cu/8isfee4Lv/ALuXnzJu973/vYbrcsl0ve+973Mo4jX/mVX8lyuXxg+cViwZd8yZfwH//jf+S9730vX/ZlX8Z6vea///f/zsXFBb//9//+R9bxlV/5lfz8z/88/+W//Be++Iu/+BX69Kf6RKkHOl0Pdb0+qibYp0DH7GV1/V5kmce+dn7s/udqhVLqvfd5zOsepn3V+96nliJ/lwK1Ukoh5wKlUEul5CzLlEKOiZwSJWdKjJSUKCkzjZFpOxH3e2KMpBBJY2CaJqYxELYTKWXG3UBJmRQnck6kqy3kDClBCDBGyMAIhJf+6k71qV0ViEBCdonhMX8XZJeJr9M2Plz/r//P/5Ov+L/+r9dkXZ+QQOR3/s7fyT/4B/+Ar/marzl2LR5Xv/7rvw7A7/pdv+uxz3/mZ34mt2/f5jd+4zf4Pb/n9xyX/5zP+ZwXXS/A+9//fr7sy76M3/iN36DWymd91mc9djsOy//qr/7qy/9wp3rVq9ZKSZlaqlygcrl3kar1+DgV+fvhi+JjLoD3P1cPr7t/2fseBx7/3L1n7v39mHU+vOwjm/SRANKjG/7AZzh89uP3kws1Z0oplHzvQl7y/L3kQjle3KHWQs0VqJRcKCVT0vzaXEkpQirknKklURMzAJB15pSpuVBSouZMTomcCiVnUszUUChhJIdMyYWYEmUMM6jI5ClRY4BcSDmTY4WcIBc5i5/qVC+zhtdhnQro5n+f6PXzf+Gf8gP8U/6/v/nvXvV1fUICkb/0l/7Sy1ru+eefB+DJJ5987POHxw/dlBdeeOFlLX/79u2XtfxTTz31wPKnem3rcEEtMZFTlovrDDoOzx9BRykIDrkHRmAGCPNzx78P4OQhcPGJWLUUuUinRA6FEpNc5GMmxyQAI2VqyZRYoAqwqIVHQUqtcARvmVKg5ESOhVoyNUZZR5bugQAR6TbUUqEW6gxKSpm3J0VSSJQo25Ti/NrD75YzNVdKSOSSZnBUqTVDlu15FF2d6lQvrw5dh1N99KWAt75G6/qEBCIvt4ZBMG3bto99/vD4gfNx+G/XPR6PfrTLN03zwHKnenWr1kqeInGYiMNEjulet+MBwHEAGp+cdRwzHC7w82fKMZGmQJ4CKWTyFGdAkF/y8x5BWZHug3Ql8jyaKHNHZH48Z2qaOyU5U8t9QOUI7golFtIUBHRMmTQGSkqkKVNKoqZMjJGaKjVlak2UVOb3ExDF4Z/gwlOd6uOukdOu9EqVeo3W80kNRA7jEqUe/3U93HJ/tZc/1cdXtRRp85ciF655rp9jIg4TaZAL3Wta9/32L7YffHRvJ+9RsnQZ5LPO45BSBFhV+dxpjORw71/Jj/fFUUqhjLzvYRRV8qE7kqXrkDI5F2osAiJmAHDYd0uWLoh0Qu6BDeaRSslZRigxkqZAGiNxipQpSDcqJko5jHwSJSNg5t4M6jjqOXU5TvVq1OsxavlUr/cA/+/XYD2f1ECk73sAxnF87PPTND2w3Mtd/tABeanlQwgPLH+ql1f37sTn/84cgYdHICUX4n4kjR89209pjTIarRUojdIKpRVojVKg0KBmnKHUPbwx/89jQYdSx8eVml+sePBv7r0nIKCgZGo88CkqtcyfVR1fTElZPusQiONETgVlNLZvsH0zL/fgthxARw6RPCVyCJSYqQqqAoxGK0ArqtZoQPUzcNGaeuB5pHLsTtRSyTGSYiKPkTQlShEgOK0HAR8pypglZQErM7+05gxUUAZtZL3lSBKt9z6zOvBiPuqf9VRv5NLAkat3bx8aFPOxIcc3RpMr5FLJSpHm/S6jpPM3jydTgVwreR7LJmVkn0UBMq4t2lDV4ViV56qSc0oFigKFOjC/yFqDUqR52aIURUFVh9cKKbWiyMjjVavjY7IcVKXJ8/Fua8XWiqoFB+hS0fMada3ysZnBvdIPHFYVWZ9sh2ynzplm/swRRVLgULiS6VwiPWMZvWbXO7bd46cNr3R9UgORp556il/5lV/h1q1bfNZnfdYjzz/M8ThwOh6nwPlYln8pjsqpEJ5ClJZ8npUJL8W5EAAykabwkssqrdHOoK1FW4PWCmX0ERioA3jQcvGVh2cQMv+NmkEK87KK+17PAwDkJT9rEoD1yGdVoM18gcZSciGFSNwNxN1ImuID4xVjzSPvX2ulxEwcp3lEk6ilyDbO26+NRlsr2+s02liMUyg0SsvnLrkQh5GwHUnjRA6JNEzSgYnS7UghEDZ74j6QY5STchLOSQWUs1ir5bfNFQVo4+91dXKmZNCqoo2mKE0x94GVU1fkU7csoCwYhbUaZcw9EH3gaSk5To3WKGvQ2qCsRhuZ1MWcUUrLxblUKoVSISsNRoExoA3ZW7qmAa0Y0GRVyboSYmUfE6oqASMxE0Ii5EzKhTzfFKhaKMiFXI4RhVGyjepw8daQigCHWAqp1CPwKSiKRv47N/yU1nLxR24I6gxU6gwEKpA5AJDDc/N7UNEoXCk0tWJnQGQNaKWpBeL8jqbONxwIT/se4BHgEZUmKUXSmqAVFUWTI73ShPvOZ7YUIpWiFKE16F4RjWZsLcW8Ng4fn9RA5HM+53P4mZ/5GX7913+dL/qiL3rguVorv/mbv4kx5ghSDmqZg3rm4fqN3/gNAN72trcB8Nmf/dlorY+PP1y/+Zu/Cby4aueNVrVW8kxMzPO/j5boGcdA2A6Pf90MGGzjcZ3Hdi3Gmvs6HWrufMwdkFdglPJi9bF81pILeQpMu5G4HUghvujFWM2AqtZCGmQckkNEaYPzFr86lxN1rVTmM2DlHmeGelTFlAQlB9KUCbuBNIzEYaSETM5JuhkpE4ZA2OyIY6BOcb5ggNaGFBNaG7QxMzE4UVLlAH4KQC1oYyhGQVIYC3LGzhATNch2Cap5RX+OU72WpQA7tyG0QVuNMQrtHGiF0RptBFgopalaUWtBKY2xFqUFgGBAaYMx92TZMRdKVTRaUZ2WK7BryK0hu4aua9FOs4+ZCjithVA9ZVTOmFIJUyTbjGsy01TIpTCWyqQ0xVSMRhisqkLRWK1wWmGsphgD1qGsYsqVWCohZWIWQJBSoiCfsVKPY1GFFll4LaQI+dARVRCZwZOGajQFRayQi3Q1FdLZsLXS1UJrwDeK1lhKzJAEvmRhj+MP41QgKemiZBTBKCKaqDVJK7LSAsIAXws3SLRZQOA9SlYlac2gFK2utDqRlSI6zeT0a3av8EkNRL70S7+U7/u+7+Onf/qn+dZv/dYHnvvlX/5l7ty5w9vf/vajZ8gf+AN/gLZt+YVf+AX2+/1x9AKw2+34hV/4Bfq+5wu/8AsBjv//3ve+l1/+5V/mC77gCx5Yx0/91E8B8GVf9mWv5sf8hK6SMmmaeQwfA/A4VK2VsBspMWGcRVs9AwoZsVhn8csO23q5s3+N637gkUKkxJfHVamlkEIi7kfCdhCCaS4ore51cLSRLo7WKKOOnJg8TFSlMM7hlg1aGxnhjBNpmGQM82Ik1UM3yGgBGfuBaT1SUqCELOCDSp4S03pL2IkfRk0zoTUlSpQ7UVALine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 576x360 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.subplots(figsize = (8, 5))\n",
|
||||
"[plt.plot(xs, 1000. * betainc(*v, xs), color = palette[0], alpha = 0.1) for v in vals[pf_bool_vol]]\n",
|
||||
"[plt.plot(xs, 1000. * betainc(*v, xs), color = palette[2], alpha = 1.0) for v in vals[pf_bool]]\n",
|
||||
"[plt.plot(xs, 1000. * betainc(*v, xs), color = palette[4], alpha = 0.1) for v in vals[pf_bool_sm]]\n",
|
||||
"sns.despine()\n",
|
||||
"plt.xlabel(\"P(increase)\")\n",
|
||||
"plt.ylabel(\"# of Stocks Held\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "df048755",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Finally, show the pareto optimal points for the voltron ones."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 22,
|
||||
"id": "29a955cc",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"[Text(61.0, 0.5, 'b'), Text(395.7999999999999, 0.5, 'b')]"
|
||||
]
|
||||
},
|
||||
"execution_count": 22,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAwsAAAFdCAYAAABM7vF5AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAEAAElEQVR4nOydd5wU5d3Av8/MbLm9Xrg7OgiIYgNURKOxYoklRo29x5I3UaOxtxhssfckGo01GgtRiZ2IHSyIIlJEejmO43rbOjPP+8dsvZ3d2z04xGS++WzkZp55nmdmZ5/n+T2/JqSUEgcHBwcHBwcHBwcHhx4oP3QHHBwcHBwcHBwcHBy2ThxhwcHBwcHBwcHBwcHBFkdYcHBwcHBwcHBwcHCwxREWHBwcHBwcHBwcHBxscYQFBwcHBwcHBwcHBwdbHGHBwcHBwcHBwcHBwcEW7YfugMN/H+vWrePAAw/MeF4Igcfjoaqqip122olf/epX7LTTTput/eXLlzNq1KjNVp+Dg4PD/yIffvgh06dPZ968eTQ1NeF2u6murmaPPfbg2GOPZccdd0y7Jtv473K58Pl8DBkyhH322YcTTzyRgQMH5tWnH3p+AWeOcfjfQzh5Fhw2N8mD+Y477ojb7U45L6WktbWVNWvWYJomiqJw991387Of/WyT2m1sbOSWW25h7dq1/Otf/9qkuhwcHBz+V9F1ncsvv5w333wTgNraWqqrq+no6KC+vp5QKIQQgrPOOosrr7wy5dps438kEqGlpYX169cjpaSwsJBbb72VQw89NOe+/VDzCzhzjMP/Lo5mwaFfuf/++xkyZIjtubVr13LJJZfw7bffcu211/KTn/yE0tLSPrf18ccf89Zbb7HDDjv0uQ4HBweH/3Xuu+8+3nzzTUaOHMk999zDuHHj4ueCwSBPP/009957L48//jiDBg3itNNOs60n0/i/YcMG7rjjDt544w0uu+wySktL2XPPPfPu55acX8CZYxz+d3F8Fhx+MIYOHcq9996Lpmn4/X7eeOONH7pLDg4ODv/T+P1+nn32WcASGpIFBQCv18t5553H//3f/wHwyCOPYJpmXm3U1tZy9913c8ghhxCJRLj22muJRCKb5waiOPOLg8PmwxEWHH5Qhg4dysiRIwFYsWLFD9wbBwcHh/9tVq1ahd/vx+12s91222Usd/zxxwOWaU59fX3e7QghuOGGG/B4PNTV1fH666/3uc+ZcOYXB4fNgyMsOPzgCCEAy9Y0mXA4zJNPPsmxxx7LhAkTGD9+PL/4xS/4+9//TigUSil7wAEHcPXVVwOwcOFCxo4dywEHHABYNq5jx45l7NixdHd3p7X//fffx88nc9pppzF27Fi++uorpk6dysSJE5k4cSJnnnkmpmly1VVXMXbsWN5++22+++47LrzwQiZPnsxOO+3E4YcfzsMPP0w4HN5sz8nBwcGhv9E0yzo5HA7z6aefZixXW1vLq6++ynvvvUdtbW2f2qqsrGT//fcH4P333+9THb2RaX4BZ45xcMgVx2fB4QdlxYoVLF26FCAlYkVbWxvnnnsu8+fPR1EUhg4ditfrZcmSJSxatIg33niDv//975SXlwOWo5vL5WLVqlX4fD622247BgwYsFn6ePvttzNv3jy23XZb2traGDBgAIqSkLM/++wzLrvsMgBGjhxJQUEBy5Yt49577+Wbb77hr3/962bph4ODg0N/s80221BTU0NDQwO//e1vOeOMMzjyyCPZZptt0spuv/32m9zehAkTePvtt/nyyy83ua6eZJpfwJljHBzywdEsOPxgLF68mIsuuggpJUOGDOGwww6Ln7vqqquYP38+EyZM4J133mHGjBn8+9//ZubMmey2224sXLiQa6+9Nl7+gQce4PzzzweswfSf//wnDzzwwGbp57x583jooYd47bXX+PDDD1PaBfjnP//JT37yEz744ANee+013n///XiZ9957j/nz52+Wfjg4ODj0N5qmcf311yOEoLu7m7/85S8cdthh7L///lx55ZW8/PLLbNy4cbO1N3jwYACam5s3q99CtvkFnDnGwSEfHM2CQ7/yu9/9Li20XTgcZuPGjfEJZ9iwYTz88MN4PB4Avv32W95//33Kysr4y1/+QkVFRfzagQMH8sADDzBlyhRmzpzJd999l9WudnMwYcIEpkyZAoCiKJSVlaWcLysr4/7778fr9caPnX766TzzzDOsWbOGefPmsfPOO/drHx0cHBw2F1OmTOHRRx/lhhtuoK6uDoD169fz6quv8uqrr6IoCpMnT+ayyy7b5MhAhYWF8X+3t7dTVVWV87V9mV/AmWMcHPLFERYc+pUFCxbYHne5XBxyyCHsu+++HHnkkSkD/syZMwHYa6+9UgbxGJWVlUyePJmZM2fy0Ucf9ftAPn78+KznJ02alDKIxxg5ciRr1qyhq6urn3rm4ODg0D/ss88+zJgxg9mzZzNz5kxmz57NmjVrADBNk9mzZ3Pcccdxww03cOKJJ/a5nWRtQsy/IFf6Mr+AM8c4OOSLIyw49CszZ86Mx8EOh8PMmjWLW2+9lTVr1tDd3c0BBxyQNpAvX74cgC+//JKTTjrJtt5169YBsHLlyn7svUVvdqk1NTW2x2ODe75hBR0cHBy2BjRN46c//Sk//elPAaivr2f27Nm8/fbbfPTRR5imydSpU5kwYUKa826uJC90i4uL87q2L/MLOHOMg0O+OMKCwxbD7Xaz//77M27cOI499lg++eQTzj//fJ5++umUXZPY5JGsSs5EZ2dnv/YZSFFf2+FyubKed5KkOzg4/DcwcOBAjj32WI499lg+/fRTfvOb3+D3+5k2bVqanX2uxEKaDhkyxHZhnyu5zi/gzDEODvniCAsOW5yamhruuOMOzj77bL755hv+9Kc/MXXq1Pj5goICAK688krOPvvszdq23aAaDAY3axsODg4OP1YuvfRS5s2bx2WXXZbmFJzMnnvuyXHHHcfTTz/N6tWr+9ze119/DfRuipMrvc0v4MwxDg754kRDcvhB2GuvveJJfZ5//vmUeN7Dhw8HEqpiOxYtWsTixYtzstWMxQ0HbGNSb87IHg4ODg4/Zrq7u1m3bh0fffRRr2Vjzsg9HXJzpb6+Pj72ZxNM8iXb/ALOHOPgkC+OsODwg3H55ZfHbTWnTp0aH2T3228/AGbMmEFLS0vadZ2dnZx55pkcffTRvPXWW/HjyXGpkykpKYn/287+9L333uvzPTg4ODj8NxFbtL/++utZQ3IahsF//vMfAH7yk5/0qa2bbroJ0zQZPXp0PDnb5iLT/ALOHOPgkC+OsODwg1FcXMzll18OWAPso48+CsAee+zB7rvvTkdHB+eff36KiruhoYHf/OY3tLe3M2DAAI488sj4OZ/PB1i7OMkTg8/nizvf3XvvvXEbVF3Xefrpp3nllVf690YdHBwcfiQcfvjhTJgwgXA4zNlnn80zzzyTZre/fPlyfvOb3/Dtt98ybtw4fvazn+XVxqpVq7jooouYOXMmLpeLm266CVVVN+dtZJxfwJljHBzyxfFZcPhB+fnPf860adP44osveOSRRzjyyCMZNmwYd999N7/61a+YP38+hxxyCKNHj0ZRFFasWEEkEqGoqIhHH300xXFtzJgxCCFobGzkkEMOoba2ln/+858AXHTRRVx44YXMmTOHfffdlxEjRlBfX09LSwtnnHEGL7/88hZxZHNwcHDYmtE0jYcffphLLrmE2bNnc/PNN3P77bczdOhQioqKaGxspL6+HrCyIv/5z3/O6IDbMw9CKBRi48aNNDY2AlBUVMSdd97JxIkT++VeMs0vgDPHODjkgaNZcPjB+cMf/oDL5SIUCsUd0WpqanjppZe4/PLL2WGHHairq2PFihVUV1dzwgknMH36dLbffvuUekaOHMnNN9/MsGHDaGxsZO3atTQ1NQFw0EEH8dRTT7HPPvvEJ4QhQ4Zwxx13cM0112zxe3ZwcHDYWikrK+OJJ57gkUce4ZhjjmHIkCE0NzezePFiTNNk33335fbbb+eFF17IGNYTrDwIX331VfyzZMkSdF1n11135eKLL+Y///kPBxxwQL/ei938As4c4+CQD0I6MbccHBwcHBwcHBwcHGxwNAsODg4ODg4ODg4ODrY4woKDg4ODg4ODg4ODgy2OsODg4ODg4ODg4ODgYIsjLDg4ODg4ODg4ODg42OKETk0iGAyyYMECBgwYsNljPjs4OPQNwzBobGxkxx13TAljmAttbW05ZWBNpqioqM8ZaR3+u3DmBAeHrQ9nTtjyOMJCEgsWLOCUU075obvh4OBgw7PPPstuu+2Wc/m2tjamHHQgHZ35TQylpaXMmDHjf35ycHDmBAeHrRlnTthyOMJCErHU8M8++yy1tbU/cG8cHBwANmzYwCmnnBL/feZKV1cXHZ1dPP3n26kZUJXTNQ2NTZz+2yvp6ur6n54YHCycOcHBYevDmRO2PI6wkERMzVxbW8uQIUN+4N44ODgk01czkJqqCgbX5jipSLNPbTj8d+LMCQ4OWy9b85wwffp0nn32Wb7//ntM02TkyJEcc8wxnHrqqWn9XrlyJQ8++CBz586lra2NYcOGccIJJ3DyySejKOmuxR0dHTzyyCO8++671NfXU1VVxcEHH8wFF1xAUVFRWnnDMHjppZd4/vnnWb16NV6vl8mTJ3PRRRcxcuTInO7HcXB2cHD478Y08/s4ODg4OPz30s9zwh133MEVV1zB4sWLmThxInvssQdr1qzh1ltv5aKLLiI5F/J3333HcccdxxtvvMGgQYPYZ5992LBhAzfddBNXXHFFWt1dXV2ceuqpPPbYYwgh2G+//RBC8MQTT3DCCSfQ2dmZds11113HDTfcwIYNG9h7770ZPHgwb775JscccwyLFi3K6Z4czYKDg8N/NRKJzHF3SOIktHdwcHD4b6Y/54QlS5bw+OOPU1FRwXPPPRffuW9oaOCLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 792x360 with 4 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(1, 2, figsize = (11, 5))\n",
|
||||
"f = ax[0].scatter(vals[:,0], vals[:,1], c=mtps_voltron[:,1])\n",
|
||||
"ax[0].scatter(vals[pf_bool_vol,0], vals[pf_bool_vol,1], color = \"blue\")\n",
|
||||
"ax[0].set_title(\"Return\")\n",
|
||||
"fig.colorbar(f, ax=ax[0])\n",
|
||||
"\n",
|
||||
"f = ax[1].scatter(vals[:,0], vals[:,1], c=mtps_voltron[:,0])\n",
|
||||
"ax[1].scatter(vals[pf_bool_vol,0], vals[pf_bool_vol,1], color = \"blue\", label = \"Pareto Pts\")\n",
|
||||
"ax[1].set_title(\"SD Return\")\n",
|
||||
"ax[1].legend(loc = \"lower right\")\n",
|
||||
"fig.colorbar(f, ax = ax[1])\n",
|
||||
"plt.tight_layout()\n",
|
||||
"\n",
|
||||
"[ax[i].set_xlabel(\"a\") for i in range(2)]\n",
|
||||
"[ax[i].set_ylabel(\"b\") for i in range(2)]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 23,
|
||||
"id": "1ec30148",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"[tensor(0.8000),\n",
|
||||
" tensor(0.7400),\n",
|
||||
" tensor(0.6100),\n",
|
||||
" tensor(0.7400),\n",
|
||||
" tensor(0.5900),\n",
|
||||
" tensor(0.6500),\n",
|
||||
" tensor(0.6200),\n",
|
||||
" tensor(0.5500),\n",
|
||||
" tensor(0.4900),\n",
|
||||
" tensor(0.5100),\n",
|
||||
" tensor(0.5700),\n",
|
||||
" tensor(0.5400)]"
|
||||
]
|
||||
},
|
||||
"execution_count": 23,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"voltron"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 24,
|
||||
"id": "0e0411cf",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"tensor([18.9000, 17.5000, 15.9000, 35.2000, 21.8000, 26.8800, 21.7000, 16.6700,\n",
|
||||
" 23.6700, 27.3250, 27.5900, 27.6600])"
|
||||
]
|
||||
},
|
||||
"execution_count": 24,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"prices_at_time_y"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 25,
|
||||
"id": "31cbf69f",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"array([0.91082877, 0.7290717 , 0.22038326, 0.7290717 , 0.16599621,\n",
|
||||
" 0.35751542, 0.2512742 , 0.08629335, 0.02554779, 0.03962997,\n",
|
||||
" 0.121481 , 0.07189631], dtype=float32)"
|
||||
]
|
||||
},
|
||||
"execution_count": 25,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"bought_func(voltron)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 26,
|
||||
"id": "56047e29",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"base = 10000"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 27,
|
||||
"id": "4c124db1",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"prices_at_time_y = torch.tensor([y[0], *prices_at_time_y])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 28,
|
||||
"id": "89f9a042",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def value_func(vec, base = 10000):\n",
|
||||
" portfolio_value = torch.zeros(12)\n",
|
||||
" portfolio_value[0] = base\n",
|
||||
" for i in range(11):\n",
|
||||
" price_of_stock = portfolio_value[i] * bought_func(vec)[i]\n",
|
||||
" amt_bought = price_of_stock / prices_at_time_y[i]\n",
|
||||
" cash_left = portfolio_value[i] - price_of_stock\n",
|
||||
" portfolio_value[i+1] = cash_left + amt_bought * prices_at_time_y[i+1]\n",
|
||||
" return portfolio_value\n",
|
||||
" "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 29,
|
||||
"id": "069bae75",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"hodl_strat = 10000 / y[0] * prices_at_time_y[1:]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 30,
|
||||
"id": "ca127889",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"Text(0, 0.5, 'Portfolio Value')"
|
||||
]
|
||||
},
|
||||
"execution_count": 30,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAkUAAAFWCAYAAABn1OlpAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAADMfElEQVR4nOydd3wUZf7H3zPbkk3vhBI6oQmEroIgxYLS7cqpd+pZuaJnufLDeqeed+ehnt7ZFQQFEURpgqKAUlR6DwRCTy+72T7z+yPsZmdLkk02ySbM+/XiFfaZZ2aebTOf/VZBlmUZFRUVFRUVFZULHLGlF6CioqKioqKiEgmookhFRUVFRUVFBVUUqaioqKioqKgAqihSUVFRUVFRUQFUURTxOJ1OTp48idPpbOmlqKioqKiotGlUURThnD17lvHjx3P27NmWXoqKioqKikqbRhVFKioqKioqKiqookhFRUVFRUVFBVBFkYpKRJBXsZUSa35YjlVizSevYmtYjqWioqJyIaGKIhWVCCBB346dxV80WhiVWPPZWfwFCfp2YVqZioqKyoWDKopUVCKA5KgsBqZc2yhh5BZEA1OuJTkqK8wrVFFRUWn7qKJIRSVCaIwwUgWRioqKSuNRRZGKSgTREGGkCiIVFRWV8KCKIhWVCCMUYaQKIhUVFZXwoYoiFZUIpD7CSBVEKioqKuFFFUUqKhGKrzBySnYk2QWogkhFRUWlKdC29AJUVFSC4xZG24uWIskuNIKOrnHDOWb6URVEKioqKmFGFUVtEJvNRklJCZWVlbhcrpZejkojkWWZJNcIz+PSUpk0zWjOVZg5x/6wnksURaKiooiNjSUpKQlRVI3JKioqFw6qKGpj2Gw28vPzSUpKokuXLuh0OgRBaOllqTQCh2SjylmqGIvRJqMV9WE9jyzLSJJEVVUVZWVlVFRU0KlTJ7Ra9TKhoqJyYaD+DGxjlJSUkJSURGpqKnq9XhVEbQCX5PAbs7uqwn4eQRDQaDTExcXRsWNHDAYDJSUlYT+PioqKSqSiiqI2RmVlJfHx8S29DJUw4pRtfmMO2YpT8h8PF4IgkJKSQnl5eZOdQ0VFRSXSUEVRG8PlcqHT6Vp6GSphwinZcMn+liKAKmd5kwojvV6P0+lssuOrqKioRBqqKGqDqC6ztoFTslHlDG6p0YmGJhVG6udIRUXlQkMVRSoqEYhbEOnF6KBzXLITozahyS1GKioqKhcKqihSqZP/LPiOTduPhOVYm7Yf4T8LvgvLsdoqbkFk1CYgIwed55IdiIJGFUYqKioqYUIVRSp1MrB3B+57akGjhdGm7Ue476kFDOzdIUwra33IcnCRA0pBpBUNuGS7YruA0qXlkGxoRYMqjFRUVFTCgCqKVOrk0pzuvD7n5kYJI7cgen3OzVya0z3MK6w/S5YsITs7m+zsbEaNGoUkSbXOX7VqlWf+448/3uDzHjt2jF/96lecOnUq6BxfQSTLEi5ZGeisF2P89gFUYdSCvLRlE+vz88JyrPX5eby0ZVNYjqWiohI6qihSqReNEUaRIoh8KSws5Keffqp1zsqVK8Nyrl//+tds3Lgx6HZfQQT4ZZ1p0KLTRCn3k2v6oanCqGUYmtmemz//tNHCaH1+Hjd//ilDM9uHaWUqKiqhoooilXrTEGEUqYLIXctp9erVQedUVVXx7bffhqXEQV0WKXfQtFsQVY/5iCJRh0bQohGUFaa9BZBbGPlamFSajrFZXVkwZWajhJFbEC2YMpOxWV3DvEIVFZX6EpGiyOVyMW/ePGbOnElOTg4DBgzgmmuu4bXXXsNm8/8FnJeXx+9//3vGjBnDwIEDmTx5MvPmzQt6I6qoqODvf/87V155JQMGDGDcuHE8//zzmEymoOtZuHAh06ZNIycnh4svvpjf/e535OUFvwB+//33/OIXv2DEiBEMHjyYWbNmsWHDhoa9IBFEKMIoUgURwKhRozAYDKxZsyZonM8333yDxWJh9OjRTb4egyZGIYjAv5K1RqgWZ1pBaS1y+FiFtKIBg0bpZlNpWhojjFRBpKISOUScKHK5XNx///0888wzHD16lIEDBzJ8+HAKCgqYO3cus2bNwmKxeOYfOHCA6667ji+//JL27dszevRozp49yzPPPMOjjz7qd3yTycRtt93GW2+9hSAIjB07FkEQePfdd7nxxhuprKz02+fPf/4zc+bM4ezZs4waNYoOHTqwYsUKZsyYwb59+/zmL1myhDvvvJPt27czYMAAcnJy2L59O3fddRcff/xxeF+wFqA+wiiSBRGA0Wjksssu49y5c2zfvj3gnBUrVmA0Ghk7dqzfNqfTycKFC5k1axYjRoygX79+jBgxgl/96lcK8btlyxays7PJz88HYPz48WRnZyuOdfbsWebMmcPll19O//79GTVqFI899hjHTxxXzNMIOrKzs7nlhlls3fwTM6+5hVFDJjJj8g0UFhbwyiuvkJ2dzTfffMPatWu56aabyMnJYdiwYTzwwAMcPHgwpNforMmENUzFGy+EWJmGCCNVEKmoRBYRJ4oWLVrE+vXryc7OZtWqVbz33nu89dZbrF69mpycHHbu3Ml//vMfoDqT59FHH8VkMvHiiy+yYMECXn31VVavXk12djbLly/3c4+8/PLLHDx4kBtuuIEVK1Ywd+5cVq9ezdSpU8nNzeXll19WzF+zZg1LliyhX79+fPXVV7zyyissXryYp556iqqqKh5//HGFpaGgoIA5c+YQFxfHp59+yptvvsnbb7/NRx99RGxsLM899xznzp1r8texqalNGEW6IHJz9dVXA4FdaCaTiQ0bNjBu3DiiopSWGVmWeeCBB5gzZw6HDx9m4MCBjBkzhtjYWDZu3Mjdd9/N2rVrAUhNTWXy5MkYjUYAJkyYwOTJkz3H2rdvH9OmTWPhwoUYDAYuv/xy0tLSWLp0KbNu+BX79uwHqrPOxPNus4KCQh79zZ+Iio5m+MVDiI2LJTE5znPMRYsW8cADD1BZWcmoUaOIi4tj7dq13HLLLSF99mJ0OoosVWqsTAiEIoxUQaSiEnlEXPvrzz77DIA//vGPZGRkeMaTk5N58sknmTp1Kl9++SUPP/wwmzZt4uDBgwwfPpypU6cq5s6ZM4dbbrmFDz/8kCuvvBKodpstWrSI2NhYHnvsMUSxWhNqtVrmzJnD+vXrWbx4MQ8//LDnJvbOO+8A8PjjjxMXV3Pjuemmm1i9ejXff/89W7ZsYeTIkQDMmzcPu93Or3/9a3r16uWZP2DAAO666y5efvllPv74Y2bPnt0UL1+tbNp+hD+9/Dm5+YVhPe6Nv387pPFQ6ZGVxnO/nRJ2gTV27FiioqJYs2YNTzzxhGLb2rVrsdlsXH311ZjNZsW2VatWsX79enJycnjvvfc8okmSJJ5//nnef/995s+fz4QJE+jevTsvvfQSEydOJD8/nyeeeIKOHTsCYLfbmT17NqWlpfzlL3/htttu85xj8ZJP+PMf/48nHnmSxcvnEa2P8VSYLiwsZPyEy/nrP+cgCAKSJClcaOvWrePJJ5/k5ptv9pzn7rvvZvPmzSxevJgHHnigXq9PnMFAarSRSY24cV+IN35vYRTseV+Ir4uKSmsg4ixFSUlJdOvWjQEDBvht69KlC1BtjQE8booJEyb4zR0yZAgpKSn89NNPnlihbdu2YbVaGTlyJLGxsYr5MTExXHzxxVitVrZt2wZUi6gdO3aQmJjI0KFD/c7hPu9339UUI6xtTRMnTvSb35w8/s+lYRdEzUFufiGP/3Np2I8bExPDZZddxunTp9m1a5di28qVK4mLi+Oyyy7z20+SJMaNG8cjjzyisCKJosj1118PwOnTp+s8/1dffcWJEyeYOHGiQhABXDPlKi6fcBlnTp3l66++9cQTuZk1a5ZHJImiiFO2eSyWgwcP9ggiqO5hdsMNNwCwe/fuOtflTZRWq8bKNABvYbTu2FF2F56j4nw85IX8uqioRDoRJ4reeOMNVq5c6bHUeOO+oLdr1w6A3NxcAIVFxpuuXbsiSRJHjhxRzO/Zs2fA+d26dQPwxF4cOXIEWZbp3r27x6oUaP6hQ4eAardKbm4uoih6tnnTpUsXRFEkNze3ziJ+Ks2D24W2atUqz1h5eTmbNm1iwoQJ6PV6v32uueYaXn/9dYVQrqqqYteuXR5XnMMRuImrN1u2bAFgxIgRfttcsoORlw4H4Ocfd6IRlaKob5/+iF5fXxkZierU/IEDB/odLzU11bPOUFFjZRrGmE5deHDwcCYtns/g9/5L+9deYuLCD5j52SfMu3bGBfu6qKhEMhHnPguGLMvMnTsXgCuuuAKosRilpaUF3Mc9XlRUBFS7Heozv7i4uF7z09PTFfPLy8ux2+0kJycHvJlqtVqSkpIoLi7GbDb7Wauamud/P40///tzDh9vXdainp3TePY3U5rk2N4uNHdg/ldffYXD4WDSpElB96uoqGDhwoVs2LCBo0ePej5joTRRPXPmDADPPvsszz77bNB5BWcLFJYiURRJSEjA4qzALtWIHOl8Gr63m9eNRqMB6q6oHYz6uITcqIIIcktL+P3Xq1h5NNczZnO5WH/iGAB3rfqc2/oN4Bf9B9IzKaWFVqmiouJLqxFF//znP9m6dSupqancddddAJ4sNN9AWDfucfevY/ff6OjATTZDnW8wGBTz3OsJNt/7HC0hii7N6c437/0urMf0DapuLUHWboxGI2PGjGH16tXs3buXfv36sXLlShITE7nkkksC7nPo0CFuv/12SkpKSE1N5aKLLqJ79+707duXzp07M3PmzHqd210y4pJLLiElpebGKCPhlGrae3Tr3hVR0Hgeu4WXVjT4iCKXYnu4UWNl6qbK4eDFLRv5+9bvsbtcQeedrKzg+c0beX7zRi7t0InbLxrEddl9idMbgu6joqLS9LQKUfTvf/+b//3vf+j1el5++WWSk5MBPC6tYDcB969i99/mml8Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 576x360 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.figure(figsize = (8, 5))\n",
|
||||
"plt.plot(eval_times, value_func(matern), label = \"Matern\", marker = \"x\", markersize = 20)\n",
|
||||
"plt.plot(eval_times, value_func(specmix), label = \"SM\", marker = \"x\", markersize = 20)\n",
|
||||
"plt.plot(eval_times, value_func(voltron), label = \"Voltron\", marker = \"x\", markersize = 20)\n",
|
||||
"plt.plot(eval_times, hodl_strat, label = \"HODL\", marker = \"x\", markersize = 20)\n",
|
||||
"plt.legend()\n",
|
||||
"sns.despine()\n",
|
||||
"plt.xlabel(\"time\")\n",
|
||||
"plt.ylabel(\"Portfolio Value\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 31,
|
||||
"id": "0eecd76f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"eval_times = torch.tensor(eval_times) / 252"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 32,
|
||||
"id": "5fc079ab",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def running_sharpe_ratio(vec):\n",
|
||||
" returns = vec - 10000\n",
|
||||
" std_returns = torch.tensor([returns[:i].std(0) for i in range(len(vec))])\n",
|
||||
" # need avg return divided by sd of returns?\n",
|
||||
" return returns.cumsum(0) / std_returns / torch.arange(vec.shape[0])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 57,
|
||||
"id": "6dfd94ca",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAABtUAAAGHCAYAAADPxN2IAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAAEAAElEQVR4nOzdd5xU1fn48c+5d/r2pVpAOmoSLFgTFSu2WGOsscYYv0YTzZdY8vUXW0yi8s3XIEZjiUnssSRqYkFR0VhQUERRURCQuizbp8+99/z+GLbMnS2zu7O7s8vzfr14wdy5M+fMsDvlPOd5HqW11gghhBBCCCGEEEIIIYQQQgghOmQM9ASEEEIIIYQQQgghhBBCCCGEKHQSVBNCCCGEEEIIIYQQQgghhBCiCxJUE0IIIYQQQgghhBBCCCGEEKILElQTQgghhBBCCCGEEEIIIYQQogsSVBNCCCGEEEIIIYQQQgghhBCiCxJUE0IIIYQQQgghhBBCCCGEEKILnoGegBBCCCGEEEIIIYQQQgwV69ev58477+Q///kPtbW1VFRUcPDBB/PTn/6UESNGDPT0hBBC9ILSWuuBnoQQQgghhBBCCCGEEEIMdh9//DHnn38+TU1NTJkyhbFjx/LJJ5+wadMmxo4dy5NPPklZWdlAT1MIIUQPSflHIYQQQgghhBBCCCGE6KVkMsmsWbNoamri2muv5bnnnuPOO+9k3rx5HHnkkXz99dfccccdAz1NIYQQvSBBNSGEEEIIIYQQQgghhOil559/ntWrV3Pcccdx9tlntxz3+/1cc801DB8+nFWrVg3gDIUQQvSW9FQTQgghhBBCCCGEEEKIXpo3bx4A559/ftZ12223HW+99VZ/T0kIIUSeSVBNCCGEEEIIIYQQQggheunTTz/F6/Wy8847s3HjRp577jm+/vprysvLmTlzJtOmTRvoKQohhOglpbXWAz2JQmFZFps2bWL06NF4PBJvFEIIMXTIe5wQQoihSN7fhBBCFIpkMsm3vvUtRo8ezZVXXsn//M//EIvFMs754Q9/yJVXXpnT/cl7nBBCFCbpqdbGpk2bOOyww9i0adNAT0UIIYTIK3mPE0IIMRTJ+5sQQohCEQ6HAWhoaOCqq67i8MMP58UXX+T999/n//7v/ygvL+f+++/n8ccfz+n+5D1OCCEKkwTVhBBCCCGEEEIIIYQQohcSiQQAsViMfffdl9mzZzN+/HhKS0s55phjuOWWWwC48847kcJhQggxeEnucIGJxZMZl4MBn4xVoGMNxcckY8lYhTKWEEIIIYQQQggxmASDwZZ/n3HGGVnXH3zwwYwaNYqqqirWrFnDuHHj+nF2Qggh8kWCagXmyzXVGZenTd1BxirQsYbiY5KxZKxCGUsIIYQQQgghhBhMSkpK8Hq9pFIpdtxxx3bP2X777amqqqKurk6CakIIMUhJ+UchhBBCCCGEEEIIIYToBdM0mThxIgBVVVXtnrNlyxYAKisr+21eQggh8kuCakIIIUQfe/rpp5k6dSqLFi3q1u2qqqr41a9+xWGHHca0adM48sgjufPOO0kmk13fWAghhBBCCCFEvzrooIMAePHFF7Ou++qrr1i/fj0jR45kzJgx/T01IYQQeSJBNSGEEKIPffjhh9x0003dvt2mTZs49dRTefzxxyktLeXggw8mEokwZ84cfvjDH5JKpfpgtkIIIYQQQggheur0008nFArxz3/+k+eee67leENDA9deey2O43DWWWdhGLIkK4QQg5X0VBNCCCH6yLx587j66quJRqPdvu3111/Ppk2b+NnPfsYll1wCQDQa5Sc/+Qlvv/02Dz74IBdccEG+pyy6aVM4zHMrl1Mfj1MRCPLdiVMYXVwsY8lY/TpWfxmKjwmG7uMSYjCJVNez8pVFxBvCBMqLmXjYXhSNKB/oaQkhRLftsMMO3HzzzfziF79g1qxZPPDAA4wcOZIlS5ZQV1fHfvvtxw9/+MOBnqYYguzaKhLvvoATrscoqcC/71GYlaMGelpCDEkSVBNCCCHybNOmTfz+97/nmWeeIRgMMnz48Jba+bn46quveP311xk7diwXX3xxy/FQKMTNN9/M4YcfzkMPPSRBtQG0JRrlopeeZd6qlZiGQdKy8Hk8XD7/BWaOn8g9Rx7P8FBIxpKx+nSs/jIUHxMM3cclxGASrW1k3tV3s/qNj1CmgZ1MYfq8vHr9A4w7aDdm/u5iQpWlAz1NIYTolmOOOYbx48dz11138d5777FixQrGjBnDBRdcwPnnn4/X6x3oKYohxGmoof72S0ksfhVlmGgrifL40HddiX/6oZRfPhejbNhAT1OIIUVyjYUQQog8u/3223nmmWf45je/yeOPP86ECRO6dfv//Oc/aK055JBDssqCbL/99uy6666sX7+eFStW5HPaIkdbolH2/ts9vPjVChK2TTSVwtKaaCpFwrZ58asV7P23e9jSgwxFGUvGKjRD8THB0H1cQgwm0dpGHjruKla9/iF2MoUVS6BtByuWwE6mWLVgCQ8ddxXR2saBnqoQQnTbLrvswpw5c3j33XdZunQp//73v7noooskoCbyymmoofqyg0gsegVSCXQiCraV/juVILHoFaovOwinoWagpyrEkCKZamJI0Vrz8ZfrCfq9TBo7cqCnI4TYRk2YMIFbbrmF448/vke18puDZZMnT+7w/j/++GO++OILJk2a1Ku5iu676KVnqYqESTlOu9enHIeqSJgLXniSh487qldjXfDCCzLWEBzropee5emTTu/VWP0l15/3wfSYYOg+LiEGk3lX301kSwOOZbd7vZOyiGxpYN7Vd3PiPVf28+yEEEKIwld/+6U4dZvB6qDnupXCqdtM/e2XUXndI/07OSGGMAmqiSHlF7P/wWPPLwLgwlO+w8mH7z6wExJCbJMuuuiiXt1+8+bNAIwc2f7mgBEjRgB0q6SkyI9N4TDzVq3scCG+WcpxmL96DZ/ULGNEyN+jsaqjCeavWStjDaaxVn9NytFdjvXSVyvZFA4XfN+u7vy8z1s1OB4TDN3HJcRgEqmuZ/UbH+GkrE7Pc1IWq99YQqS6XnqsCSGEEG3YtVUkFr/acUCtmZUisXg+dm2V9FgTIk+k/KMYMrbUhVsCagD3PfnWAM5GCCF6LhaLARAIBNq9vvl4VMqS9bvnVi7HzDH70DAUr63peeDztTXVGCq3c2WswhjLtjoPqDWzbYd/rfyix2P1l+78vJvKGBSPCYbu4xJiMFn5yiKUmdvvoTJNVs5f1PWJQgghxDYk8e4LKMPM6VxlGCQWvtjHMxJi2yFBNTFkfL2xbqCnIIQQedFcMlKp9qMBWuuMv0X/qY/HSVqd76pvZtkOjYkudg12oiFhkbI7z6SRsQpnrI31cWxyDKpph3V19T0eq7905+c96djUx2N9PKP8GKqPS4jBJN4Qxk7m9pprpyziDZE+npEQQggxuDjherSVzOlcbaVwwvV9OyEhtiFS/lEMGbFE9huJznFxSwghCkkoFAIgHo+3e30ikQAgGAz225xEWnkggM/jwUp1vRDoNU1Ghioo8Y7o0VgjQw34TJNYB71mZKzCGEujSdma6vWfoxzQOWxZUxrWr63v1jgDoTs/7z7DpDwwOF6ThurjEmIwCZQVY/q8WLFEl+eaXg+BsqJ+mJUQQggxeBjF5SiPD213vVlMebwYxeV9PykhthESVBNDRlMk+wtZysptd7oQQhSS5l5qHfVMq66uzjhP9J/jJk7livm5lc1wNJy5y4GMCvWsF9NZu2zHdW9+IGMV2Fhap4NoCcshnnJIWA5oGO98hVZrcxpLAxNVSU7nDqTu/Lzb2uG7E6f08YzyY6g+LiEGk4mH78WrNzyQ07natpl42F59PCMhhBBicPHvdzT67qtyOlc7Dv59j+rjGQmx7ZDyj2LIaApnZ3QkkrmV9hFCiEIyefJkAFasWNHu9StXrgRgyhRZ6O1vo4uLmTFmFJ4uGnV5DYOZ4ycyurhnwaDmsWaOn4i3i95P+Rrr4B3HQVd7URw4ZMdxg+px5WMs29FEkja1kRQbG5JUNSapj1rEUw7NVVi3LyklWKPB7iJL3tYEazVjyst78Ij6V3/+X/Wnofq4hBhMikaUM+6g3TA8ne/zNbwexh20O0UjyvtnYkIIIcQgYVaOwj/9UPB4Oz/R48U//TDMylH9MzEhtgESVCswlWWhjD8yVu4+Xbkx61jQ33fJmEPt+ZOxZKxCGmtbd+CBBwLw6quv4jiZUY4NGzbw2WefscMOOzBp0qSBmN42LZKq4fqDJjE86OswsOY1DEYVFXPPkcf3erx7jjyeUUXFHS7+53Os0qUpzEQnQSFbYyY0JR/3vO9Ys/58XD0ZS2tNPOXQELOoakyyoSFBbThFJGFjO+0/PwfuPZWKT22MJJ0+h0YSKpbZHPHtnXv70PpF6/PX9z/v/ak/fwaFEO2b+buLCQ0rRZnt/x4aXg9Fw8uY+buL+3lmQgghxOBQfvlcjPKR0MF7KR4vRsVIyi+/o38nJsQQJ+UfC8yOoytkrB6698m3so6Vl/bdwv9Qe/5kLBmrkMbalmzYsIFYLEZFRQWVlZUAjBkzhgMPPJA333yTP/zhD1xxxRUARKNRrr32Wmzb5vzzzx/IaW+TbG1RHVtFRcDH09/bl2sXfMZb62owlUnSsfEZJrZ2mDl+IvcceTzDQ71/DxoeCvH+ORdx0UvPMm/VSkxl9MlYm2ubWPjeV1Ri0fgtL4kRRrpGoUE6e02Bv9qh9OMU77KSzbVNjKzsefnC/npczWPNO+ksZt73V9aZERSgVbqvmQZGJQO8dNZZlPkDNMWtlrKOuhttWU1DMXZ0GQfvNpnX3v2S2l2MDp/Dys8cDttzSq+ev/40PBTi3R+czznPP8J/1tVgKLBsB49p4GiYOW4i9x51Ql7+r/pT88/gBS88wfw1a1sel9c0048rjz+DQoj2hSpLOfH+q5nLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 2160x360 with 4 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(1, 4, figsize = (30, 5))\n",
|
||||
"ax[0].plot(ts, y)\n",
|
||||
"[ax[i].set_xlabel(\"Time\") for i in range(4)]\n",
|
||||
"ax[0].set_ylabel(\"Asset Price\")\n",
|
||||
"[ax[0].axvline(x=eval_times[i], alpha = 0.2, linestyle=\"--\") for i in range(len(eval_times))]\n",
|
||||
"\n",
|
||||
"ax[1].plot(eval_times, matern, label = \"Matern\", color=palette[1], alpha = 0.5)\n",
|
||||
"ax[1].plot(eval_times, specmix, label = \"Spectral Mixture\", color=palette[3], alpha = 0.5)\n",
|
||||
"ax[1].plot(eval_times, voltron, label = \"Voltron\", color = palette[-1], alpha = 0.5)\n",
|
||||
"ax[1].scatter(eval_times, matern, s = 120, label = \"Matern\", color=palette[0], zorder=4)\n",
|
||||
"ax[1].scatter(eval_times, specmix, s = 120, label = \"Spectral Mixture\", color=palette[2], zorder=4)\n",
|
||||
"ax[1].scatter(eval_times, voltron, s = 120, label = \"Voltron\", color = palette[-2], zorder=4)\n",
|
||||
"\n",
|
||||
"ax[2].plot(eval_times, value_func(matern), label = \"Matern\", color=palette[1], alpha = 0.5)\n",
|
||||
"ax[2].plot(eval_times, value_func(specmix), label = \"SM\", color=palette[3], alpha = 0.5)\n",
|
||||
"ax[2].plot(eval_times, hodl_strat, label = \"Hold\", color=palette[5], alpha = 0.5)\n",
|
||||
"ax[2].plot(eval_times, value_func(voltron), label = \"Voltron\", color = palette[-1], alpha = 0.5)\n",
|
||||
"ax[2].scatter(eval_times, value_func(matern), color=palette[0], zorder=4, s=120)\n",
|
||||
"ax[2].scatter(eval_times, value_func(specmix), color=palette[2], zorder=4, s=120)\n",
|
||||
"ax[2].scatter(eval_times, hodl_strat,color=palette[4], zorder=4, s=120)\n",
|
||||
"ax[2].scatter(eval_times, value_func(voltron),color = palette[-2], zorder=4, s=120)\n",
|
||||
"\n",
|
||||
"ax[3].plot(eval_times, running_sharpe_ratio(value_func(matern)), \n",
|
||||
" label = \"Matern\", color=palette[1], alpha = 0.5)\n",
|
||||
"ax[3].plot(eval_times, running_sharpe_ratio(value_func(specmix)), \n",
|
||||
" label = \"SM\", color=palette[3], alpha = 0.5)\n",
|
||||
"ax[3].plot(eval_times, running_sharpe_ratio(hodl_strat), label = \"HODL\", markersize = 20,\n",
|
||||
" color=palette[5], alpha = 0.5)\n",
|
||||
"ax[3].plot(eval_times, running_sharpe_ratio(value_func(voltron)), \n",
|
||||
" label = \"Voltron\", color = palette[-1], alpha = 0.5)\n",
|
||||
"ax[3].scatter(eval_times, running_sharpe_ratio(value_func(matern)), \n",
|
||||
" s = 120, color=palette[0], zorder=4)\n",
|
||||
"ax[3].scatter(eval_times, running_sharpe_ratio(value_func(specmix)), \n",
|
||||
" s = 120, color=palette[2], zorder=4)\n",
|
||||
"ax[3].scatter(eval_times, running_sharpe_ratio(hodl_strat),\n",
|
||||
" s = 120, color = palette[4], zorder=4)\n",
|
||||
"ax[3].scatter(eval_times, running_sharpe_ratio(value_func(voltron)), \n",
|
||||
" s = 120, color = palette[-2], zorder=4)\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"ax[1].set_ylabel(\"P(increase)\")\n",
|
||||
"ax[2].set_ylabel(\"Portfolio Value\")\n",
|
||||
"ax[3].set_ylabel(\"Sharpe Ratio\")\n",
|
||||
"\n",
|
||||
"ax[2].legend(ncol = 4, loc = \"lower center\", bbox_to_anchor = (-0.2, -0.4))\n",
|
||||
"plt.subplots_adjust(wspace=0.25)\n",
|
||||
"sns.despine()\n",
|
||||
"[ax[i].set_xlim((-0.1, 5.1)) for i in range(4)]\n",
|
||||
"plt.savefig(\"trading_strategy.pdf\", bbox_inches = \"tight\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 99,
|
||||
"id": "42799762",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"tensor(5.)"
|
||||
]
|
||||
},
|
||||
"execution_count": 99,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"ts.max()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "26cd9568",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
File diff suppressed because it is too large.
Load diff
@@ -0,0 +1,601 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns\n",
|
||||
"import torch\n",
|
||||
"import gpytorch\n",
|
||||
"# from voltron.robinhood_utils import GetStockData\n",
|
||||
"import os\n",
|
||||
"import pandas as pd\n",
|
||||
"# import robin_stocks.robinhood as r\n",
|
||||
"import pickle5 as pickle\n",
|
||||
"\n",
|
||||
"sns.set_style(\"whitegrid\")\n",
|
||||
"sns.set_palette(\"bright\")\n",
|
||||
"\n",
|
||||
"sns.set(font_scale=2.0)\n",
|
||||
"sns.set_style('whitegrid')\n",
|
||||
"\n",
|
||||
"import sys\n",
|
||||
"sys.path.append(\"../\")\n",
|
||||
"from botorch.models import SingleTaskGP\n",
|
||||
"from voltron.means import LogLinearMean\n",
|
||||
"from gpytorch.kernels import ScaleKernel, RBFKernel, MaternKernel"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# with open(\"../../stock_data.pkl\", \"rb\") as handle:\n",
|
||||
"# raw_data = pickle.load(handle)\n",
|
||||
" \n",
|
||||
"with open(\"./stock_data.pkl\", \"rb\") as handle:\n",
|
||||
" raw_data = pickle.load(handle)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"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>date</th>\n",
|
||||
" <th>symbol</th>\n",
|
||||
" <th>open_price</th>\n",
|
||||
" <th>close_price</th>\n",
|
||||
" <th>high_price</th>\n",
|
||||
" <th>low_price</th>\n",
|
||||
" </tr>\n",
|
||||
" </thead>\n",
|
||||
" <tbody>\n",
|
||||
" <tr>\n",
|
||||
" <th>0</th>\n",
|
||||
" <td>2016-09-26</td>\n",
|
||||
" <td>XOM</td>\n",
|
||||
" <td>83.52</td>\n",
|
||||
" <td>83.06</td>\n",
|
||||
" <td>84.5000</td>\n",
|
||||
" <td>82.9300</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>1</th>\n",
|
||||
" <td>2016-09-27</td>\n",
|
||||
" <td>XOM</td>\n",
|
||||
" <td>82.59</td>\n",
|
||||
" <td>83.24</td>\n",
|
||||
" <td>83.3400</td>\n",
|
||||
" <td>82.2900</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>2</th>\n",
|
||||
" <td>2016-09-28</td>\n",
|
||||
" <td>XOM</td>\n",
|
||||
" <td>83.46</td>\n",
|
||||
" <td>86.90</td>\n",
|
||||
" <td>87.2300</td>\n",
|
||||
" <td>83.3400</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>3</th>\n",
|
||||
" <td>2016-09-29</td>\n",
|
||||
" <td>XOM</td>\n",
|
||||
" <td>86.97</td>\n",
|
||||
" <td>86.46</td>\n",
|
||||
" <td>87.2000</td>\n",
|
||||
" <td>85.6750</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>4</th>\n",
|
||||
" <td>2016-09-30</td>\n",
|
||||
" <td>XOM</td>\n",
|
||||
" <td>86.84</td>\n",
|
||||
" <td>87.28</td>\n",
|
||||
" <td>87.8100</td>\n",
|
||||
" <td>86.6500</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>...</th>\n",
|
||||
" <td>...</td>\n",
|
||||
" <td>...</td>\n",
|
||||
" <td>...</td>\n",
|
||||
" <td>...</td>\n",
|
||||
" <td>...</td>\n",
|
||||
" <td>...</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>6290</th>\n",
|
||||
" <td>2021-09-17</td>\n",
|
||||
" <td>SLB</td>\n",
|
||||
" <td>28.70</td>\n",
|
||||
" <td>28.31</td>\n",
|
||||
" <td>29.2600</td>\n",
|
||||
" <td>28.0000</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>6291</th>\n",
|
||||
" <td>2021-09-20</td>\n",
|
||||
" <td>SLB</td>\n",
|
||||
" <td>27.37</td>\n",
|
||||
" <td>27.25</td>\n",
|
||||
" <td>27.7077</td>\n",
|
||||
" <td>26.7050</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>6292</th>\n",
|
||||
" <td>2021-09-21</td>\n",
|
||||
" <td>SLB</td>\n",
|
||||
" <td>27.61</td>\n",
|
||||
" <td>26.93</td>\n",
|
||||
" <td>27.7800</td>\n",
|
||||
" <td>26.6400</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>6293</th>\n",
|
||||
" <td>2021-09-22</td>\n",
|
||||
" <td>SLB</td>\n",
|
||||
" <td>27.51</td>\n",
|
||||
" <td>27.15</td>\n",
|
||||
" <td>27.8100</td>\n",
|
||||
" <td>27.1200</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>6294</th>\n",
|
||||
" <td>2021-09-23</td>\n",
|
||||
" <td>SLB</td>\n",
|
||||
" <td>27.32</td>\n",
|
||||
" <td>28.88</td>\n",
|
||||
" <td>29.1250</td>\n",
|
||||
" <td>27.2265</td>\n",
|
||||
" </tr>\n",
|
||||
" </tbody>\n",
|
||||
"</table>\n",
|
||||
"<p>6295 rows × 6 columns</p>\n",
|
||||
"</div>"
|
||||
],
|
||||
"text/plain": [
|
||||
" date symbol open_price close_price high_price low_price\n",
|
||||
"0 2016-09-26 XOM 83.52 83.06 84.5000 82.9300\n",
|
||||
"1 2016-09-27 XOM 82.59 83.24 83.3400 82.2900\n",
|
||||
"2 2016-09-28 XOM 83.46 86.90 87.2300 83.3400\n",
|
||||
"3 2016-09-29 XOM 86.97 86.46 87.2000 85.6750\n",
|
||||
"4 2016-09-30 XOM 86.84 87.28 87.8100 86.6500\n",
|
||||
"... ... ... ... ... ... ...\n",
|
||||
"6290 2021-09-17 SLB 28.70 28.31 29.2600 28.0000\n",
|
||||
"6291 2021-09-20 SLB 27.37 27.25 27.7077 26.7050\n",
|
||||
"6292 2021-09-21 SLB 27.61 26.93 27.7800 26.6400\n",
|
||||
"6293 2021-09-22 SLB 27.51 27.15 27.8100 27.1200\n",
|
||||
"6294 2021-09-23 SLB 27.32 28.88 29.1250 27.2265\n",
|
||||
"\n",
|
||||
"[6295 rows x 6 columns]"
|
||||
]
|
||||
},
|
||||
"execution_count": 4,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"pd.read_pickle(\"../../spdr-data/XLE.pkl\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 25,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"array(['AAPL', 'F', 'JPM', 'SBUX', 'TSLA', 'VIRT'], dtype=object)"
|
||||
]
|
||||
},
|
||||
"execution_count": 25,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"np.unique(raw_data[\"symbol\"])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"# Header"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 26,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ntest = 200\n",
|
||||
"ntrain = 200\n",
|
||||
"tckrs = ['TSLA', \"F\", \"JPM\", \"SBUX\", 'AAPL', \"VIRT\"]\n",
|
||||
"tckr = \"VIRT\"\n",
|
||||
"span = \"5year\"\n",
|
||||
"interval = 'day'\n",
|
||||
"T = 5."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Data Setup"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 27,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAX8AAAEdCAYAAADkeGc2AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAABGfUlEQVR4nO2deWATdfr/37nb9KD3BaUt0HKfFQQVUERUkENUXHUB8Yu4y0/XAxdXXEVB3ZXDo3iirquoiCioHIKggCxIq+W+CuW+Su8jTXPn98dkJpNkkibN2eR5/dNkZjLzSTPznmeez3OIzGazGQRBEEREIQ72AAiCIIjAQ+JPEAQRgZD4EwRBRCAk/gRBEBEIiT9BEEQEQuJPEAQRgUiDPQCC8AezZs3Cjh07cN111+GTTz5x6zOPPvootmzZggkTJmDx4sWYOnUqSkpK8Oijj+Kxxx7jtlu2bBnefvttp/uRSCRQKpXo1KkTRowYgYcffhhxcXHc+lGjRuHSpUsefZ8hQ4ZgxYoVHn2GIFxB4k+EJZMmTcKOHTtQXFyMmpoaJCcnu9y+qakJO3bsAADceeedbh0jNjYWBQUFDssNBgMuXbqEY8eO4dixY1i/fj2++eYbJCUlAQD69OmD9PR0m8+oVCqcOHECADBo0CCHfQodhyC8gcSfCEtGjx6N+Ph4NDY2YtOmTXjggQdcbr9p0ybodDpkZmZi6NChbh2jV69eLq3xn376CU8//TQuXbqEV155BUuXLgUAFBUVOWxbXFyMadOmAQBWrlzp1vEJwhvI50+EJXK5HLfffjsAYOPGja1uv27dOgDAxIkTIRb75rIYM2YMHnnkEQDA5s2boVKpfLJfgvAFJP5E2DJp0iQAQGlpKa5evep0u6tXr+L3338H4L7Lx11GjBgBANDr9Th37pxP900Q3kDiT4QtgwYNQk5ODsxmM3788Uen261fvx4mkwkDBw5Ebm6uT8cgEom411RGiwglSPyJsGbixIkAgA0bNjjdhnX5TJ482efH/+mnnwAA0dHRyM/P9/n+CaKtkPgTYc2kSZMgEolw8OBBXLhwwWF9eXk5jh07hqioKIwdO9Znx9Xr9Vi9ejU+/vhjAMC0adOgUCh8tn+C8BaK9iHCmo4dO2Lw4MEoKSnBjz/+iFmzZtmsZ63+0aNHIzY21qN9Hz16FPfdd5/D8ubmZly4cAFqtRoAcwPi5wkQRChA4k+EPZMmTUJJSQk2btxoI/5ms9krl49KpcLevXsF12VlZWHSpEkYP368YNw+QQQbcvsQYc+tt96K6OhoHDt2DKdPn+aW7927F5cuXUJGRgaGDRvm8X6HDBmCsrIylJWV4fjx4/jtt9/w2GOPQSqVorq6Gp06dSLhJ0IWEn8i7ImNjcXo0aMB2Mb8+zK2XyQSISkpCY8++igWLFgAnU6HRYsWuV1agiACDYk/ERGw8fus+Ov1emzatMlmna+46667MGHCBADA4sWLnbqGCCKYkPgTEcGwYcOQkZGBU6dOoby8HHv27EFdXR0GDhyIvLw8nx/vn//8J1JTU2E0GvHss89Cp9P5/BgE4Q0k/kREIBaLMX78eADAL7/8gq1btwKwZgH7mg4dOmDevHkAgLNnz+KDDz7wy3EIoq2Q+BMRA+ve2bp1K3755RcoFAqMGzfOb8cbO3YsN5G8fPlynD171m/HIghPIfEnIoauXbuib9++OHDgACorKzF69GibOvv+4IUXXoBMJoNOp8OCBQv8eiyC8AQSfyKi4E/u+nqiV4guXbrgoYceAgDs2rUL69ev9/sxCcIdRGaqNkUQBBFxkOVPEAQRgZD4EwRBRCAk/gRBEBEIiT9BEEQEEvLibzabodVqqQsSQRCEDwl58dfpdDh8+HCb0+OPHDni4xERoQj9zpEB/c6+I+TF31s0Gk2wh0AEAPqdIwP6nX1H2Is/QRAE4QiJP0EQRARC4k8QBBGBkPgTBEFEICT+BEEQEQiJP0EQRARC4k8QBBGBkPgTfqPsXC3Gz/kel6pUwR4KQRB2kPgTfuOn4vMAgIMnq4I8EoIg7CHxJ/yGRmcAACjk0iCPhCAIe0j8Cb+h0xsBAFKJKMgjIQjCHhJ/wm+whVi1OmNwB0IQhAMk/oRfMJvNKD5SAQDQkPgTRMhB4k/4Bb61r9WT+BNEqEHiT/iFFq2Be200moI4EoIghCDxJ/wCX/wNRurCRhChBok/4RfUGqv4n7ncgOfe2wW1Rh/EEREEwYcCsAmfM37O95BLrXYFO/F79EwtrumZHqxhEQTBgyx/wi/oDI5+fh1N/BJEyEDiTwSMxmZdsIdAEIQFEn8iYBhNNPFLEKECiT/hU8xm5wJvNFHIJ0GECiT+hE8xubDuXa0jCCKwkPgTPoXv2nlwXC/bdRTvTxAhA4k/4VNY8Z9xR2/cNSpfcB1BEMGHxJ/wKazASwTKOFOZB4IIHUj8CZ/CCrxEzIj/O3+/CcufHQ2xiCx/ggglKMOX8CnspC4r/p0z4gEAYrGYxJ8gQgiy/Amfwgq8WGx7akkkIhJ/ggghSPwJn8L5/O3OLIlYRHH+BBFCkPgTPoUVeAfLXyyCiUI9CSJkIPEnfIq9z59FIhHDQG4fgggZSPwJn+I01NMMtPBq/BMEEVxI/Amfwmbx2lv+9Sotduy7GIwhEQQhAIk/4VM0Osa6V8gpipggQhkSf8KnaLRMw5ZoO/G/4/o8yKR0uhFEqEBXI+EzGlRafLvtJAAgSiGxWRcTLYPRaHJZ8pkgiMBBz+aET6hr0uCpN39FdX0LACBaYXtqyWUSmMyAwWiCTCoR2gVBEAGExJ/wmnMVjXh08TabZfbir5Azgq/VGUn8CSIEILcP4TX2wg8AHWIVNu8VMov4UxN3gggJSPwJr9BoHWP3k+IVDsvYyV69gUo8EEQoQOJPeMVXW8oclslljm4dNu6fWjkSRGhA4k94RWOzzmGZkE9fYqn1Q5U9CSI0IPEnvEJIzBUyx9NKbCn3QOJPEKEBiT/hFSLHbo1OLH+L+FMrR4IICUj8Ca+ob9I6LFO48PmT5U8QoQGJP+EVDc06JMbZRvf0y09x2I71+dOEL0GEBiT+hFdotAbkZyciJSGaW3b3qHyH7VjL30BuH4IICSjDl/AKjc6IWKUMnzw/Bt//egp5WfEQCUwE0IQv0d7RG4zYsOssxt+QB4l9n9J2CIk/0WaMJjOq61u40g0TR3R1ui35/In2zprt5fj8x+OIkktw27DcYA/Ha9r/7YsIGtv+OA8A+OWPC61uK5WQz59o37Cd6FQt+iCPxDeQ5U+0GdZ9r9W1Xq9HTKGeRDuHPYc/3XAU+09UokduEv58W88gj6rtkOVPtJmEWDkA4OkHClvdltw+RHtHzJvLOnCyGqu2nBDczmQy47/rj6CipjlQQ2sTJP5EmzFZGrN0zohrdVsSf6K94+4kb2WdGt9uK8dTb+7w84i8g8SfaDOskIvFAmm+drAXDvuZBpUWDSrHBDGCCFXYMOU4pczldqwbtEmtD+nQZhJ/os0YjYyQS9wRf66qJ3Mx/Hn+Jvx5/ib/DY4gfEyzZaK3Se16wpffs6K2QePXMXkDiT/RZlgrns3edYWYS/Iitw/RPlEJiL6QX58fAGEK4Z7VJP5Em2GteE8sf3ufv1oTHmFzRHhjNJmxY99Fh+XNAmGfGp21wVEohzaT+BNthrP8Je6IvyXO32iyuQHc+9xGbCk+558BEoSP2FoifI7+Z90Rh2V8t08oBzi4Lf5GoxGfffYZJkyYgL59+2LIkCF46KGHsH37dsHtz5w5g6eeegojR45E//79MX78eHz++eectUi0fzyb8LVa/vaTYGXn63w/OILwIUIWPgAcLK926Esddm6fZ599Fq+88gouXbqEYcOGoXfv3igpKcEjjzyCd955x2bb48eP4+6778aGDRuQlZWF4cOHo6KiAgsXLsTcuXN9/iWI4GCd8G39NGLdPicv1kNvd7FEKyjXkAhtpFLrOX7Pzfm466Zu3Hv7SV3+zSCU3T5uXXUbN27E999/j7y8PHz++edISWFK9p48eRL33Xcf3n77bYwbNw65ubkwm82YO3cuVCoVFi1ahIkTJwIAamtr8eCDD2LdunW45ZZbcOutt/rvWxEBwTrh636o5/bSi4iPkdusEyoERxChhIwX4z9tbC8AgM5gwrqdp6Fq0QGI4dbzLf927/b54YcfAABPP/00J/wAkJ+fj/Hjx8NkMmHXrl0AgF27dqGsrAxDhgzhhB8AkpKSMH/+fADAihUrfPYFiODRlglfANh/ogoA8H8T+gCAw5MAQYQarIYvfmw4t+yG/lkAHEM/124vt36uvYt/UVER1q1bhxEjRjisa25mQp0kEqay486dOwEAo0ePdti2sLAQycnJKC0thUqlavOgieDy26ErmPriJmgsFo47E7781PgmS9P3DrFyJMVHYePuM4KN4AkiVGCteX42e5ySeYJVqW3P3Tped7t27/OXy+UoKCiAXG77uL5t2zZs2rQJSqWSE/vycuauV1BQILivvLw8mEwmnDp1yptxE0Hko+8Pob5Ji6q6FgCA2IM4f8B6cUglYtQ2amAyA6999rt/BksQXnDkdA3u+sd6fLKeiepRyK2e8thoJtPXVdKXMYTzWjyeadNoNJg7dy7Ky8tx6tQpZGVlYdGiRZw7qLKyEgCQmpoq+Hl2eXV1dVvHTAQZ1n+vs7hr3PD6CCLjTaIdLKfzgQg9Pt1wlDvPAVv3ZaylzAPj82cw21n67d7y53P58mVs3rzZxnIvKyvjXre0MNZgVFSU4OfZ5Wq12tNDEyGC1OLmUWsNkEvFbZ6wlUnF6Jad4MOREYRvaXaRhCiTShAll6Cp2bpNi5ZJ8LqmZzqA0Pb5e2z5Z2RkYM+ePRCLxdi9ezdeeeUVLFy4EGq1GrNmzeJcAM4Egb0z2t8hW+Pw4cOeDpWjtLS0zZ8lHNFpGbdNRVUdpJK2/38Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"idx = -2\n",
|
||||
"# data = GetStockData(tckr, span=span, interval=interval)\n",
|
||||
"data = raw_data[raw_data[\"symbol\"] == tckr]\n",
|
||||
"\n",
|
||||
"ts = torch.linspace(0, T, data.shape[0])\n",
|
||||
"# train_x = ts[:ntrain]\n",
|
||||
"# test_x = ts[ntrain:(ntrain+ntest)]\n",
|
||||
"\n",
|
||||
"y = torch.FloatTensor(data['close_price'].to_numpy())\n",
|
||||
"log_returns = torch.log(y[1:]) - torch.log(y[:-1])\n",
|
||||
"# train_y = y[:ntrain]\n",
|
||||
"# test_y = y[ntrain:(ntrain+ntest)]\n",
|
||||
"\n",
|
||||
"dt = ts[1] - ts[0]\n",
|
||||
"\n",
|
||||
"plt.plot(ts, y)\n",
|
||||
"plt.title(tckr);\n",
|
||||
"sns.despine()\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 28,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"torch.Size([1259])"
|
||||
]
|
||||
},
|
||||
"execution_count": 28,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"ts.shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Now apply GCPV"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 33,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def get_and_fit_data_model(train_x, train_y, pred_vol=None, vol_model=None):\n",
|
||||
" voltron_lh = gpytorch.likelihoods.GaussianLikelihood()\n",
|
||||
"# voltron = VoltronGP(train_x, train_y.log(), voltron_lh, pred_vol)\n",
|
||||
" model = SingleTaskGP(train_x.view(-1,1), train_y.log().view(-1,1), likelihood=voltron_lh)\n",
|
||||
" model.mean_module = gpytorch.means.LinearMean(1)\n",
|
||||
" model.mean_module = LogLinearMean(1)\n",
|
||||
" model.mean_module.initialize_from_data(train_x, train_y.log())\n",
|
||||
" model.likelihood.raw_noise.data = torch.tensor([1e-6])\n",
|
||||
"\n",
|
||||
" # Use the adam optimizer\n",
|
||||
" optimizer = torch.optim.Adam([\n",
|
||||
" {'params': model.parameters()}, # Includes GaussianLikelihood parameters\n",
|
||||
" ], lr=0.1)\n",
|
||||
"\n",
|
||||
" # \"Loss\" for GPs - the marginal log likelihood\n",
|
||||
" mll = gpytorch.mlls.ExactMarginalLogLikelihood(voltron_lh, model)\n",
|
||||
"\n",
|
||||
" for i in range(500):\n",
|
||||
" # Zero gradients from previous iteration\n",
|
||||
" optimizer.zero_grad()\n",
|
||||
" # Output from model\n",
|
||||
" output = model(train_x)\n",
|
||||
" # Calc loss and backprop gradients\n",
|
||||
" loss = -mll(output, train_y.log())\n",
|
||||
" loss.backward()\n",
|
||||
" # print(loss.item())\n",
|
||||
" optimizer.step()\n",
|
||||
" return model"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 43,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def predict_prices(test_x, voltron, nvol=10, npx=10):\n",
|
||||
" ntest = test_x.shape[0]\n",
|
||||
" # vol_paths = torch.zeros(nvol, ntest)\n",
|
||||
" px_paths = torch.zeros(npx*nvol, ntest)\n",
|
||||
"\n",
|
||||
" # voltron.vol_model.eval();\n",
|
||||
" # voltron.eval();\n",
|
||||
"\n",
|
||||
" px_paths = voltron.posterior(test_x).sample(torch.Size(((nvol * npx),))).exp().squeeze(-1)\n",
|
||||
"# for vidx in range(nvol * npx):\n",
|
||||
"# vol_pred = voltron.vol_model(test_x).sample().exp()\n",
|
||||
"# vol_paths[vidx, :] = vol_pred.detach()\n",
|
||||
"\n",
|
||||
"# px_pred = voltron.GeneratePrediction(test_x, vol_pred, npx).exp()\n",
|
||||
"# px_paths[vidx*npx:(vidx*npx+npx), :] = px_pred.detach().T\n",
|
||||
" return px_paths"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 44,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"eval_times = list(range(100, ts.shape[0], 100)) #+ [ts.shape[0]]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 45,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"now running time: 100\n",
|
||||
"prob of stock increase: tensor(0.6900)\n",
|
||||
"now running time: 200\n",
|
||||
"prob of stock increase: tensor(0.8200)\n",
|
||||
"now running time: 300\n",
|
||||
"prob of stock increase: tensor(0.7900)\n",
|
||||
"now running time: 400\n",
|
||||
"prob of stock increase: tensor(0.)\n",
|
||||
"now running time: 500\n",
|
||||
"prob of stock increase: tensor(0.9700)\n",
|
||||
"now running time: 600\n",
|
||||
"prob of stock increase: tensor(0.7300)\n",
|
||||
"now running time: 700\n",
|
||||
"prob of stock increase: tensor(0.9600)\n",
|
||||
"now running time: 800\n",
|
||||
"prob of stock increase: tensor(0.9900)\n",
|
||||
"now running time: 900\n",
|
||||
"prob of stock increase: tensor(0.4300)\n",
|
||||
"now running time: 1000\n",
|
||||
"prob of stock increase: tensor(0.2900)\n",
|
||||
"now running time: 1100\n",
|
||||
"prob of stock increase: tensor(0.3400)\n",
|
||||
"now running time: 1200\n",
|
||||
"prob of stock increase: tensor(0.4200)\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"prob_of_increases = []\n",
|
||||
"\n",
|
||||
"for i, time in enumerate(eval_times):\n",
|
||||
" print(\"now running time: \", time)\n",
|
||||
" with gpytorch.settings.max_cholesky_size(2000):\n",
|
||||
" # pred_vol = get_and_fit_gpcv(ts[:time], log_returns[:(time - 1)])\n",
|
||||
" # vol_model = get_and_fit_vol_model(ts[:time], pred_vol)\n",
|
||||
" data_model = get_and_fit_data_model(ts[:time], y[:time], None, None)\n",
|
||||
" end_ind = -1 if i + 1 >= len(eval_times) else eval_times[i+1]\n",
|
||||
" paths = predict_prices(ts[time:end_ind], data_model).detach()\n",
|
||||
" # now we predict the probability of increase at time i + 1\n",
|
||||
" prob_of_increase = (paths[..., -1] > y[time]).sum() / paths.shape[-2]\n",
|
||||
" print(\"prob of stock increase: \", prob_of_increase.detach())\n",
|
||||
"\n",
|
||||
" prob_of_increases.append(prob_of_increase.detach())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 68,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"True"
|
||||
]
|
||||
},
|
||||
"execution_count": 68,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"i+1 >= len(eval_times)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 69,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"Text(0, 0.5, 'Prob of Stock Increasing')"
|
||||
]
|
||||
},
|
||||
"execution_count": 69,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAbcAAAE2CAYAAADie9yWAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAABbrklEQVR4nO3dd1xT1/sH8M8NO0BkQ5QNRmU4EFfd1g61jtpW695WbbX112Ft7beu1u5aqlWrVuu2Wq1V615oVVQEBUWRDbI3YSe5vz8wqUgCXEjI4Hm/Xr5eeu/NzZPrJQ/3nPOcw7Asy4IQQggxIDxtB0AIIYSoGyU3QgghBoeSGyGEEINDyY0QQojBoeRGCCHE4FByI4QQYnAouRFCCDE4lNwIIYQYHEpuhBBCDI4xl4OnTp3a+BMbG8PMzAyOjo7o1KkTXn75Zdja2nIOkBBCCOGK4TL9VseOHWtexDAAAGUvVbaPYRjY2Njg22+/Rb9+/ZoVMCGEENIQTsntxo0b+P3333Hu3Dm0bdsWY8aMgZ+fHywtLVFaWorY2Fj8/fffSEpKgr+/P1566SUUFxfjypUrePDgASwtLXH48GG4u7tr8jMRQghp5Tglt/Pnz2PBggUYNGgQ1q5dC3Nz8zrHSCQSfPTRRzh58iS2bNmC5557DgCwevVq7Nq1C5MnT8ayZcvU9wkIIYSQZ3BKbhMnTsT9+/dx6dIltGnTRuVxYrEY/fv3R0BAAHbu3AkAKC8vR9++feHk5ISTJ082P3JCCCFEBU6jJWNiYuDr61tvYgMAKysr+Pj4IDo6WrHNwsIC7u7uyMrKalqkhBBCSCNxSm58Ph85OTmNOjYnJwfGxrUHY0qlUpiZmXF5S0IIIYQzTsmtU6dOyM7Oxh9//FHvcX/++SeysrLg5+en2Jafn4+kpCS4uro2LVJCCCGkkTjVuc2aNQtXrlzBihUrkJiYiHHjxsHLy0uxPyEhAYcOHcL27dvBMAxmzJgBAIiKisK3334LiUSCF198Ub2fgBBCCHkGpwElAPDbb7/h22+/VfzbxMQEFhYWKCsrg0QiUWxfvHgx5s6dCwAYN24c7t69i7Zt2+Lvv/+GlZWVmsInhBBC6uKc3ADg7t272LBhA65du4aKigrFdmNjYzz33HOYP38+unXrptg+Y8YMeHt746233oKTk5N6IieEEEJUaFJyk6uqqsLjx49RWFgICwsLeHl50YARQgghWtes5EYIIYToIk4DSuTy8vJw9+5diMViSKXSeo8dM2ZMU96CEEIIaTJOT24sy2LNmjXYs2dPg0lNLiYmpsnBEUIIIU3B6clt79692LFjB4CaGUfatWtHfWyEEEJ0DqfkdvDgQTAMg5kzZ+K9996DiYmJpuIihBBCmoxTs2TXrl1hbW2N0NBQxbpthBBCiK7h9ORmZmYGBweHVpfYZDIZSktLYWJi0uo+OyGENBXLsqiuroalpSV4PE6zPTYbp+QWEBCA27dvo7S0FJaWlpqKqY5Dhw5h6dKl2L17N4KDgxv9uqysLKxfvx7//vsvcnJyIBQKMWrUKMyZMwempqaNPo98IVZCCCHciUQiWFtbt+h7ckpuc+fOxfTp0/HVV19h1apVmoqploiIiCa9V2ZmJsaPH4/MzEz4+fnB398ft2/fRkhICK5fv47ffvut0X2G8uNEIhGnpKhvoqOjERAQoO0wdApdk7romihH16Uu+bJn2hifwSm5WVlZYdKkSdi9ezciIyPRr18/ODs71xv4pEmTmhzc6dOn8fHHH6OsrIzza5cvX47MzEy8++67WLBgAQCgrKwMb7/9Nq5evYqdO3di5syZjTqXvCnS1NTU4EeHGvrnawq6JnXRNVGOroty2ujO4ZTcXnvtNTAMA5Zl8ejRI8TFxTX4mqYkt8zMTPzwww84cuQILCws4ODggNzc3Ea/PiEhARcvXoS7uzvmzZun2M7n8/HFF19g6NCh2LVrV6OTGyGEEP3CKbn16NFDU3HUsnbtWhw5cgQBAQH48ssvsXr1ak7J7cqVK2BZFoMHD67Tidm2bVv4+fkhKioKcXFx8PX1VXf4hBBCtIxTctu5c6em4qjF29sbX3/9NUaNGtWkETbyJ8r27durPH9UVBRiY2MpuRFCiAFq0tySmiZfB66psrOzAUDl8jqOjo4AwOlpkBBCiP5QmdzKy8sB1Eyz9ew2Lp5+fUuRx2lubq50v3x7UwaqEEII0X0qk1u3bt3A4/Fw/PhxeHl5AQCCgoI4nZxhGNy/f795ETaBvClT1Qgd+aQsXFf7kQ9rNWTh4eHaDkHn0DWpi66JcnRddEe9zZIymazWv7kmA20tFcfn8wGg1irhT6usrATA/akyICDAoIf6hoeHo3v37toOQ6fQNamrMdekWiJF/OMixCYX4GFyAR6mFMDM1Agh7w+GEc8wZ/mhe6UubSZ7lcnt3LlzAABnZ+c623SdvK9NVZ9aTk5OreMIIU3Hsiwy88rwMKUAD5PzEZtSgITHRZBIa365dbCxgJ3ADLEphUjOKIZ3uzZajpi0BiqTW7t27Rq1TRfJR0mqqsOLj48HUDPjCCGEm/IqGSIeZj9JZgWITSlAcWkVAMDc1Ai+bjYYPcAHHTxsIXK3hX0bC+QUlGPm6tOIis+l5EZahNpGS1ZWVuL69euQyWTo3r07BAKBuk7NWf/+/QEA58+fxwcffFCrnCA9PR0xMTFo164dlQEQ0gCpVIbkzBI8TM5XJLO0bDGAdDAM4OpkjV7+LhC526KDhy3cna1hZFS3fMfR1gIu9nxEx+di9ACflv8gpNXhnNzS0tKwceNGtGvXDvPnzwdQ8yQ0c+ZMxRB8Pp+PlStXYsSIEeqNVon09HSUl5fD1tYWdnZ2AAA3Nzf0798fly9fxk8//YTFixcDqBkduWzZMkilUsyYMUPjsRH9dup6EvadysR6/2rwzVvH2oV5ReV4kFxQ01eWUoC4tEJUVkkBAG2sTCFyt4XIhYfBffzR3s0WlhaNvy4B3g4Iu5cBmYwFz0D73Yju4JTcsrKyMG7cOBQUFGDgwIGK7cuWLUNWVhbMzc3h4OCAtLQ0fPTRR/D29kanTp3UHvTTlixZghs3buCdd97BwoULFds///xzTJgwARs3bsT58+fh5eWF27dvIycnBwMGDMCECRM0GhfRf3fjcpFbLMGBc48wbYSftsNRu4pKCeLSChGbUlCT0FIKkFdUMwjL2IgHn3Zt8FIvD8VTmbMdHwzDIDw8HF1F3PurA3zscfZmCpIzi+HVlpomiWZxSm5bt25Ffn4+unbtqpizMT4+HhERETA2Nsaff/4JHx8f/P7771izZg22b9+Or7/+WiOBN8TNzQ0HDhxASEgIQkNDkZycDDc3N0ydOhXTpk2DsbFO1q8THZKeWwoA+OtSPF7q7QEX+5Zb5knTEh4X4eP1V1BeKQEAuNjz4e9tjw4etujgbgvvdm1gYmyk1vcM8HEAAETH51FyIxrH6Rv+ypUrMDc3xy+//KJoApSPoOzduzd8fGra0qdOnYpNmzbhxo0bagmyvmm/6tsnFAqxZs0atcRAWheWZZGRI0ZHV3MkZVfjt6P38Mn0ntoOSy1kMhYb/rwDUxMePpjUCx08bNHGSvMlLs52fDjZWiA6IRcj+3tr/P1I68Zp4saMjAx4eXkpEhtQk/AYhlEM4gBqiqeFQiFNb0X0VnFpFUorJPBwMsMbz4twLSoDd2JztB2WWpy7mYIHyQWY8Yo/evq7tEhikwvwcUB0fJ7WamBJ68EpuZmYmEAqlSr+XVZWhtu3bwOoeXJ7WlFRkVYWqCNEHTKeNEnaWRtjzEAfONnxsflIFKRSWQOv1G0lZVXYfvw+/LzsMCTYrcXfP9DHHsWlVUjJKmnx9yatC6fk5u7ujpSUFJSU1NyYly5dgkQigbOzMzp06KA4Ljo6GmlpafD09FRrsIS0FHl/m721MUxNjDBrpD+SM0tw8nqyliNrnp0nYiAur8a8sZ21soDk0/1uhGgSp+Q2YMAAVFRUYMGCBdixYwe+/PJLMAyjGPJfXl6Of/75BwsWLADDMHjhhRc0EjQhmpaeKwaPAWwsa7ql+wQK0dnXAbtPPkBJWZWWo2uauNRCnLyWhFf6emltQIezHR8ONhaIiqcuC6JZnJLbrFmz0KlTJ9y8eRNr1qxBTk5OrdWuo6Ki8H//93/Izs5G165dMX36dE3ETIjGZeSWwtGWD2OjmqcbhmEwe3QASsursPf0Qy1Hx51MxmLDoTuwsTLDxJc6ai0OhmEQ4GOPe9TvRjSM02hJS0tL7Nu3DwcPHkRsbCzc3d3xxhtvwNraGkDNIqCenp4YOXIk5syZA1NTU40ETYimpeeWQuhQe+i/V9s2eKm3J47/m4iXe3vA3UV7s/BwdeZGMmJTCvH+xCBOhdeaEODtgIvhaUjLFsPN2VqrsRDDxbnYy8zMDJMmTVK6z8HBASdPnmx2UIRok7wMYGCQKwBJrX2TXu6I0MjH2HwkGivn9tFKvxVXxaVV+P34ffh72z/5TNoV6GsPAIiOz6XkRjSGU7MkV7QYKNFH8jIAoYNVnX1trMww8cUOiIzNwc37WVqIjrsd/9xHaYUE87U0iORZQntL2AnMaVAJ0SjOT24SiQRnz55FXFwcKioq6qz5JpVKUVlZiezsbNy6dUtthdyEtBR5GUBbR0ugvLDO/uF9vXDiWhK2/B2Nbh0c1T6Thzo9TM7H6bBkjB7gAw+hbjSjyvvdouNzwbKsTiRcYng4JTexWIzJkyfj4cOGO9TppiX6Sl4G0NbBElmpdfcbG/EwZ3QgPt98DUcvJ2LsYN1cXUIqY7Hx0F3YWptjwosdGn5BCwr0cUBoxGOk55ainWPdJ2RCmotTs+S2bdvw4MED8Hg89O7dG88//zxYlkXHjh0xfPhwdO/eHUZGNb/F9ujRA//8849GgiZEk+RlAM52queSDOrohB5+zth35iEKSpSv+K5Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(eval_times, prob_of_increases)\n",
|
||||
"plt.xlabel(\"Time\")\n",
|
||||
"plt.ylabel(\"Prob of Stock Increasing\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 70,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"torch.save(obj=prob_of_increases, f=\"matern_v.pt\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 48,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from torch.distributions import Beta\n",
|
||||
"from scipy.special import betainc"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 49,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"xs = torch.linspace(0, 1, 100)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 50,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"prob_incs = torch.tensor(prob_of_increases)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 51,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"bought_func = lambda xs: betainc(8, 1, xs)\n",
|
||||
"total_held = 1000 * bought_func(prob_incs)[:-1] "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 52,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"Text(0.5, 0, 'Time')"
|
||||
]
|
||||
},
|
||||
"execution_count": 52,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAaYAAAEgCAYAAAD/mNfGAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8rg+JYAAAACXBIWXMAAAsTAAALEwEAmpwYAABIwklEQVR4nO3deXhTZdo/8G+SNm3TtE23lLZpaWka9oKlIKCMgCKCw4yoAyoo4DAM4vJ7HXVAB4ZBXMDXUQZhXp3RUVEUQQEHFRdgBAawQGmpLN2gSbd0oUvSpFuW8/sjTehKmzbJOSe5P9flJZycnjwJp7lzznM/9y1gGIYBIYQQwhFCtgdACCGEdESBiRBCCKdQYCKEEMIpFJgIIYRwCgUmQgghnOLH9gD4zGq1wmg0wt/fHwKBgO3hEEIILzAMA5PJhODgYAiF3a+PKDANgtFoREFBAdvDIIQQXlKpVAgJCem2nQLTIPj7+wOwvblisZjl0XjOhQsXMGbMGLaHwRv0fjmP3jPn8O39amtrQ0FBgeMztCsKTINgv30nFosREBDA8mg8y9de72DR++U8es+cw8f3q7cpEEp+IIQQwikUmAghhHAKBSZCCCGcQoGJEEIIp1BgIoTw3hdHCpFbVNNpW25RDb44UsjSiMhgUGAihPBeaqIMm3ecdQSn3KIabN5xFqmJMnYHRgaE0sUJIbyXpozGY/emYe3bJzF5TCwuXq3F6kcykKaMZntoZADoiokQwntWK4PvMjVgGODUz1rMmZpEQYnHKDARQnjvwH+vIqfAdhsvWhaEgyfV3eacCH9QYCKE8Jpaq8f7By7CTyTAhBFymCxWrH4ko9OcE+EXCkyEEN5qM1nw151Z8PcT4rnFGRibEoWGxlYMi5dh9SMZKCxpYHuIZAAo+YEQwls7vrkMtVaP9csnI2NkDDIvaAEA5dWNSFNG0zwTT9EVEyGEl7Lzq/HlsSu4+5ZkZIyMAQDEy6UAgPIaA5tDI4NEgYkQwjt6Yxu27MqGQi7F0l+OcmwfEhkMkVCAsmoKTHxGgYkQwisMw2D75znQG1vx7KIJCBRfn5HwEwkxJDKYAhPPUWAihPDK4TMlOJmrxeK7RiJFIev2uEIupcDEcxSYCCG8ob1mxD/2/4wxKZG4Z7qyx30Ucim014ywWKweHh1xFQpMhBBesFiseOOTLAgFAjz9YDpEwp67nyrkUpgtVlTVN3l4hMRVKDARQnhh96EC5Gnq8dh94yAPl/S6X3x0CACgnG7n8RYFJkII5+Vp6rDrUAGmpytwW7rihvvaU8Zpnom/KDARQjitqcWEN3aeQ2RYIFbem9bn/qHBYoQGi2ktE49RYCKEcNq7X15AZZ0Rf3gwHcFB/v36GcrM4zcKTIQQzjr1cwV+OF2C+2emYkxKVL9/Lj5aSnNMPEaBiRDCSbW6Zry1OwdKRRgevHOEUz+rkIegwdAKQ1Obm0ZH3IkCEyGEc6wMgy27stFqsuIPD02Av59zH1UKewIEzTPxEgUmQgjnnM43IKegBst/NRoJMSFO/7wjM6+KAhMfUWAihHCKWqvHDzk6TBwVg7umJA3oGDEREviJBJSZx1MUmAghnGFv/BcoFuKpBTdBIOi5ukNfrhdzbXTxCIknUGAihHDGRwdtjf9+fXM4ZCEBgzqWQi6lKyaeosBECOGEnIJq7D96BXOnJkEVHzTo48VHUzFXvqLARAhhXWNTG9781Nb4b9m80S45pkIeArOFQVUdFXPlGwpMhBBWMQyD7XvOQ29sxTNdGv8NBqWM8xcFJkIIqw6fKcWJ3AosumsklD00/hsoShnnLwpMhBDWVNYa8Y/9uRg9LBLze2n8N1AhEjHCpFTMlY8oMBHiBl8cKURuUU2nbblFNfjiSCFLI+Iei8WKv+60Nf77w0O9N/4bDIU8hFLGeYgCEyFukJoow+YdZ3G+sAYNRjNyi2qwecdZpCbK2B4aZ+w+XNivxn+DER9NKeN85JpZRkJIJ2nKaKx+JAMv/+s0mlrNkAbV4vmlE5GmjGZ7aJyQr6nDrh/ycdtNfTf+GwyFXIrvM9vQ2NSGEInYbc9DXIuumAhxkzRlNIbGhgIARiVHUlBq19xqxl8/aW/8d1/fjf8Gw54AQS0w+IUCEyFukltUg8LSegC2xaNd55x81T/3/4zKWlvjP2k/G/8NlCNlnOaZeIUCEyFuYJ9TkgTa7pbHREqwecdZnw9O9sZ/981wrvHfQMWE24q5UjdbfqHARIgbFJY04InfjIfeaIK/nwBVtU14dtEEFJY0sD001tga/51HiiIMD812rvHfQIlEQsRGBVNg4hkKTIS4wX0zUxEcZLtaGp0YhDazFbKQANw3M5XlkbHDamXwt13ZaDVZ8MwAGv8NhkIeQpl5PEOBiRA3UVfoAQA3DQsGABSWNrA4GnZ9deIqsgtq8NsBNv4bDHsxVzMVc+UNCkyEuElxhR4yaQASosWQBPqhqKyB7SGxQqPV44OvLmHiqBjMGWDjv8FQyKWwWKmYK59QYCLETdRaHZLiQiEUCKBUyFDkg1dMJrMFr+/MQnCg/6Aa/w0GpYzzDwUmQtzAYrFCU9mIpPZ1TEqFDMUVepjMvnU7acc3tsZ/Ty4cP+jGfwOliKaUcb6hwESIG1RcM8JktiI5LgwAoEyQwWyxQlOpZ3lknnO+oAb7j17BnKlJmDRqCGvjkErEkEkDKDOPRygwEeIG9sSH5LjrV0wAcMWL55k6Fq5tbGrDm7vOISosEJGhgSyPzHY7jwITf1BgIsQNirU6iIQCR+WBIZESBAf5e3VmXsfCtds/P496fQua2ywYmRzB9tCgkFMxVz6hIq6EuEFxhR4JMSHw9xMBAAQCAVIVMq/OzLMXrn3pX6fR3GpGgFiEFzhSuFYhl0JvbIPe2IbQYCrmynV0xUSIG6grdI7EB7sURRg0Wj1MZgtLo3K/NGU0YiJtLSx+/YsUTgQlwLaWCaDMPL6gwESIizU2teGarsUxv2SXmhAOs4WBWuu9CRC5RTUo0eoRHx2Mb0+pOVMbUCG3LeqlzDx+oMBEiIvZEx+SYsM6bVcmyADAa9cz5RbVYNOHZ2BlgNsnJmL1IxmcKVwrj5DATySkeSaeoMBEiIsVa3UAgKQuV0zy8CCESMRemwBRWNKA+dOVAIARQyMcc05cKFwrEgqomCuP9Jr8UFFR4ZIniIuLc8lxCOELdYUeYVIxwrssKBUIBFAqwnClTMfSyNzrvpmp+PjgZQiFAqS2Xx2mKaM5M8+kkEtRUkm38vig18B0++23D/rgAoEAly5dGvRxCOGTYq0eSbGhPZbfUSbIsPc/RWg1WRDgL2JhdO6Vr6lHUmwoAgO4l/CrkEtx+mIlzBYr/ER0s4jLev3XYRimX/9JpVLExsYiIiKi0/awsDBERLC/foEQT7JYGZRo9Y6KD12lJshgsTJQV3jfVZPFyiC/pB7Dh4azPZQe2Yu5VtYa2R4K6UOvX2vOnTvXbZvJZMLjjz+O8+fPY8WKFbj//vsRGxvreFyn02Hv3r3YunUrkpOT8d5777ln1IRwVEWNAW1ma7dUcTulwvahXVTagOFDveuLW2lVI5pbzRjB0dcV76iZZ3Bk6RFu6vWKSSKRdPvvo48+QlZWFjZt2oQnn3yyU1ACgLCwMCxbtgxvvPEGsrOz8be//c3tL4AQLrGngvcWmKJkgQiTilHkhfNMeeo6AMCIJG5eMcW3ByNay8R9Tt1o3b9/P2JjY3H33XffcL8ZM2YgMTERBw8eHNTgCOEbtVYPoVCAxCE9fyMX2FtgeGEFiHxNPUKDxYiNDGZ7KD2SBvlDFhJAKeM84FRgqqmpgUwm69e+EokEjY2UAUN8S3GFDgq51FGKqCfKBBlKKvVoaTN7cGTul6epw/Ch4az0XOovBRVz5QWnAlNcXBwKCwtRVVV1w/2uXLmCgoICJCYmDmpwhPCNWqtHcmzPiQ92SoUMVgYoLveeChCNTW0oqzZwdn7JLj6aAhMfOBWY5s6dC5PJhFWrVqGsrKzHffLy8rBq1SowDIP58+e7ZJCE8IGhqQ019c3dFtZ2ZV/j40238/I19QC4O79kp5CHoLGpDTpDK9tDITfg1GKDZcuW4fvvv8fFixdx1113Ydy4cUhJSYFEIkFTUxMuX76MCxcugGEY3HzzzVi0aJG7xk0I5/SV+GAXERqI8JAArwtMQoGtHiCX2duQlNcYECZlp6Mu6ZtTgUkqleKDDz7ASy+9hIMHDyIrKwtZWVkQCARgGAYAIBKJ8OCDD+IPf/gD/P393TJoQrjIHpi6Fm/tSiAQQJkg86rSRHmaOiTFhiGIgwtrO7IHprJqA0YlR7I8GtIbp8+iyMhIvPnmm3juuefw3//+F2q1GgaDAaGhoUhOTsaMGTNoYS3xScUVeoRIxIjoR8dWpUKGrMtVaG41c/7DvC9WK4OCknrcdpOC7aH0KTpcAn8/IaWMc9yAfyPi4uKwYMECV46FEF5Ta3VIjuu5FFFXygRbAsTVch1GD+P3N/fSqkY0tZg5P78E2Iq5xlExV86jglGEuIDFykCtbewz8cFOqZAB8I4EiDxN+8Jajmfk2cXLpSivoaUsXNbrFdP9998/6IMLBALs2bNn0MchhOsqa41oM1n6TBW3iwgNRGRYoFf0ZsrX1CNEIkZsFDcX1nYVHy3FTxcqYTJb4e9H3825qNfAdOHChUEfnMsL7QhxJUdzwH5eMQHwmgoQfFhY25FCHgJrezHXhBiqmcdFvQamV1991ZPjIITXiit0tlJETnzQKRNkOH2pEk0tJkgC+ZnBamhqQ2mVAbelcz/xwa5jZh4FJm7qNTDR4lhC+k+t1SM+WgqxEz2WlAoZGAa4Uq7D2JQoN47OffJL2hfW8mR+CbheZZxq5nGXS26wGo3U34T4tuIKHZL7WFjblSMBgsfzTNcX1srYHkq/BQf5IzwkAGXVlADBVQMKTJWVldi0aRPmzp2L0aNLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(eval_times[:-1], total_held, marker = \"x\")\n",
|
||||
"plt.ylabel(\"Num Held\")\n",
|
||||
"plt.xlabel(\"Time\")"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 4
|
||||
}
|
||||
Vendored
BIN
Binary file not shown.
@@ -0,0 +1,152 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
import pandas as pd
|
||||
import gpytorch
|
||||
import argparse
|
||||
import datetime
|
||||
import warnings
|
||||
|
||||
from voltron.data import make_ticker_list, GetStockHistory
|
||||
import sys
|
||||
sys.path.append("../trading/")
|
||||
from GenerateMultiMeanPreds import GenerateStockPredictions, GenerateBasicPredictions
|
||||
from gpytorch.utils.warnings import NumericalWarning
|
||||
warnings.simplefilter("ignore", NumericalWarning)
|
||||
|
||||
def main(args):
|
||||
|
||||
if args.end_date is None:
|
||||
end_date = datetime.date.today()
|
||||
else:
|
||||
end_date = datetime.datetime.strptime(args.end_date, "%Y-%m-%d")
|
||||
|
||||
data = GetStockHistory(args.ticker, history=args.ntrain + args.lookback)
|
||||
|
||||
ntest = args.forecast_horizon
|
||||
ntrain = args.ntrain
|
||||
n_test_times = args.n_test_times
|
||||
ntime = data.shape[0]
|
||||
|
||||
test_idxs = torch.arange(ntrain, ntime-ntest,
|
||||
int((ntime-ntest-ntrain)/n_test_times))
|
||||
|
||||
train_x = torch.arange(ntrain) * dt
|
||||
test_x = torch.arange(ntest) * dt + train_x[-1] + dt
|
||||
|
||||
if torch.cuda.is_available():
|
||||
train_x = train_x.cuda()
|
||||
test_x = test_x.cuda()
|
||||
|
||||
####################
|
||||
## setup filename ##
|
||||
####################
|
||||
savepath = "./saved-outputs/" + args.ticker + "/"
|
||||
if not os.path.exists(savepath):
|
||||
os.mkdir(savepath)
|
||||
|
||||
if args.model.lower() == 'lstm':
|
||||
modelname = "lstm"
|
||||
else:
|
||||
if args.model.lower() == 'gp':
|
||||
modelname = "gp_" + args.kernel + "_"
|
||||
elif args.model.lower() == 'volt':
|
||||
modelname = "volt_"
|
||||
|
||||
if args.mean.lower() == 'constant':
|
||||
modelname += 'constant' + "_"
|
||||
elif args.mean.lower() in ['ewma', 'dewma', 'tewma']:
|
||||
modelname += args.mean + args.k + "_"
|
||||
|
||||
|
||||
###############
|
||||
## Main Loop ##
|
||||
###############
|
||||
|
||||
for last_day in test_idxs:
|
||||
date = str(data.index[last_day.item()].date())
|
||||
train_y = data.Close[last_day.item()-ntrain:last_day.item()].to_numpy()
|
||||
train_y = torch.FloatTensor(train_y).to(train_x.device)
|
||||
|
||||
if args.model.lower() == 'lstm':
|
||||
model = LSTM(train_x, train_y, 10, 128, 1)
|
||||
model.Train(args.train_iters)
|
||||
elif args.model.lower() == 'gp':
|
||||
model = BasicGP(train_x, train_y, kernel=args.kernel,
|
||||
mean=args.mean, k=args.k)
|
||||
model.Train(args.train_iters)
|
||||
elif args.model.lower() == 'volt':
|
||||
model = Volt(train_x, train_y, mean=args.mean, k=args.k)
|
||||
model.Train(gpcv_iters=args.train_iters,
|
||||
vol_mod_iters=args.train_iters,
|
||||
data_mod_iters=args.train_iters)
|
||||
else:
|
||||
print("ERROR: Model not found")
|
||||
|
||||
|
||||
samples = model.Forecast(test_x).squeeze()
|
||||
torch.save(samples, savepath + modelname + date + ".pt")
|
||||
torch.cuda.empty_cache()
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--ticker",
|
||||
type=str,
|
||||
default='F',
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ntrain",
|
||||
type=int,
|
||||
default=400,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--n_test_times",
|
||||
type=int,
|
||||
default=25,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--forecast_horizon",
|
||||
type=int,
|
||||
default=100,
|
||||
)
|
||||
parser.add_argument(
|
||||
'--kernel',
|
||||
type=str,
|
||||
default="matern",
|
||||
)
|
||||
parser.add_argument(
|
||||
'--model',
|
||||
type=str,
|
||||
default="volt",
|
||||
)
|
||||
parser.add_argument(
|
||||
'--mean',
|
||||
type=str,
|
||||
default="ewma",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--nsample",
|
||||
type=int,
|
||||
default=1000,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--train_iters",
|
||||
type=int,
|
||||
default=500,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--end_date",
|
||||
default=None,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--lookback",
|
||||
type=int,
|
||||
default=500,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--k",
|
||||
type=int,
|
||||
default=100,
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
main(args)
|
||||
@@ -0,0 +1,6 @@
|
||||
{
|
||||
"cells": [],
|
||||
"metadata": {},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,408 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "9754a692-edb8-41bf-8159-fd82edcb9195",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Warning no robinhood utils.\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import pandas as pd\n",
|
||||
"import torch\n",
|
||||
"from torch import nn\n",
|
||||
"import seaborn as sns\n",
|
||||
"import time\n",
|
||||
"import copy\n",
|
||||
"import sys\n",
|
||||
"from torch.utils.data import DataLoader\n",
|
||||
"from voltron.data import make_ticker_list, GetStockHistory\n",
|
||||
"sns.set_style('white')\n",
|
||||
"# style.use('whitegrid')\n",
|
||||
"palette = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"sns.set(palette = palette, font_scale=2.0, style=\"white\", rc={\"lines.linewidth\": 2.0})"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "dbc24c1a-f499-4b8f-8544-6d4a66f48488",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Dataset Setup"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "62f91be1-38aa-4781-8b65-c5caf9752045",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from torch.utils.data import Dataset\n",
|
||||
"\n",
|
||||
"class SequenceDataset(Dataset):\n",
|
||||
" def __init__(self, data, sequence_length=5):\n",
|
||||
" self.sequence_length = sequence_length\n",
|
||||
" self.X = data.float()\n",
|
||||
"\n",
|
||||
" def __len__(self):\n",
|
||||
" return self.X.shape[0]-1\n",
|
||||
"\n",
|
||||
" def __getitem__(self, i): \n",
|
||||
" if i >= self.sequence_length - 1:\n",
|
||||
" i_start = i - self.sequence_length + 1\n",
|
||||
" x = self.X[i_start:(i + 1)]\n",
|
||||
" else:\n",
|
||||
" padding = self.X[0].repeat(self.sequence_length - i - 1, 1).squeeze(-1)\n",
|
||||
" x = self.X[0:(i + 1)]\n",
|
||||
" x = torch.cat((padding, x), 0)\n",
|
||||
" \n",
|
||||
" return x.unsqueeze(0), self.X[i+1]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 239,
|
||||
"id": "566c9386-c297-430f-b4ef-a0f3ba7aec37",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"tckr = \"JPM\"\n",
|
||||
"ntrain = 400\n",
|
||||
"lookback = 1\n",
|
||||
"data = GetStockHistory(tckr, end_date=\"2021-12-07\", history=ntrain + lookback).Close.to_numpy()\n",
|
||||
"data = torch.FloatTensor(data).log()\n",
|
||||
"\n",
|
||||
"data = (data - data.mean())/data.std()\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"# xin = torch.linspace(0, 5*np.pi, 250)\n",
|
||||
"# data = torch.sin(xin) + 0.2 * torch.randn(xin.shape)\n",
|
||||
"\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 269,
|
||||
"id": "5dece53d-785f-4ca9-9cbf-03e16de71e9b",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"seq_len = 25\n",
|
||||
"dset = SequenceDataset(data, seq_len)\n",
|
||||
"\n",
|
||||
"trgts = []\n",
|
||||
"for i in range(len(dset)):\n",
|
||||
" trgts.append(dset[i][1].item())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 270,
|
||||
"id": "2317ba11-4f62-443f-9ad8-63f9a0996ee9",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"train_loader = DataLoader(dset, batch_size=20, shuffle=True)\n",
|
||||
"X, y = next(iter(train_loader))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 271,
|
||||
"id": "ffecfde0-6938-4dee-b1b0-6627012e3329",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"class ShallowRegressionLSTM(nn.Module):\n",
|
||||
" def __init__(self, input_size, hidden_units=128):\n",
|
||||
" super().__init__()\n",
|
||||
" self.input_size = input_size # this is the number of features\n",
|
||||
" self.hidden_units = hidden_units\n",
|
||||
" self.num_layers = 5\n",
|
||||
"\n",
|
||||
" self.lstm = nn.LSTM(\n",
|
||||
" input_size=input_size,\n",
|
||||
" hidden_size=hidden_units,\n",
|
||||
" batch_first=True,\n",
|
||||
" num_layers=self.num_layers\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
" self.linear = nn.Linear(in_features=self.hidden_units, out_features=2)\n",
|
||||
" \n",
|
||||
"\n",
|
||||
" def forward(self, x):\n",
|
||||
" batch_size = x.shape[0]\n",
|
||||
" h0 = torch.zeros(self.num_layers, batch_size, self.hidden_units).requires_grad_()\n",
|
||||
" c0 = torch.zeros(self.num_layers, batch_size, self.hidden_units).requires_grad_()\n",
|
||||
"\n",
|
||||
" _, (hn, _) = self.lstm(x, (h0, c0))\n",
|
||||
" out = self.linear(hn[0]) # First dim of Hn is num_layers, which is set to 1 above.\n",
|
||||
" output = torch.zeros_like(out)\n",
|
||||
" output[:, 0] = out[:, 0]\n",
|
||||
" output[:, 1] = torch.exp(out[:, 1])\n",
|
||||
" return output\n",
|
||||
" \n",
|
||||
"from torch.autograd import Variable \n",
|
||||
"class LSTM1(nn.Module):\n",
|
||||
" def __init__(self, num_classes, seq_len, hidden_size, num_layers):\n",
|
||||
" super(LSTM1, self).__init__()\n",
|
||||
" self.num_classes = num_classes #number of classes\n",
|
||||
" self.num_layers = num_layers #number of layers\n",
|
||||
" self.input_size = seq_len #input size\n",
|
||||
" self.hidden_size = hidden_size #hidden state\n",
|
||||
"\n",
|
||||
" self.lstm = nn.LSTM(input_size=seq_len, hidden_size=hidden_size,\n",
|
||||
" num_layers=num_layers, batch_first=True) #lstm\n",
|
||||
" self.fc_1 = nn.Linear(hidden_size, 128) #fully connected 1\n",
|
||||
" self.fc = nn.Linear(128, num_classes) #fully connected last layer\n",
|
||||
"\n",
|
||||
" self.relu = nn.ReLU()\n",
|
||||
" self.softplus = nn.Softplus()\n",
|
||||
" \n",
|
||||
" def forward(self,x):\n",
|
||||
" h_0 = Variable(torch.zeros(self.num_layers, x.size(0), self.hidden_size)) #hidden state\n",
|
||||
" c_0 = Variable(torch.zeros(self.num_layers, x.size(0), self.hidden_size)) #internal state\n",
|
||||
" # Propagate input through LSTM\n",
|
||||
" output, (hn, cn) = self.lstm(x, (h_0, c_0)) #lstm with input, hidden, and internal state\n",
|
||||
"\n",
|
||||
" hn = hn[self.num_layers-1]\n",
|
||||
" hn = hn.view(-1, self.hidden_size) #reshaping the data for Dense layer next\n",
|
||||
" out = self.relu(hn)\n",
|
||||
" out = self.fc_1(out) #first Dense\n",
|
||||
" out = self.relu(out) #relu\n",
|
||||
" out = self.fc(out) #Final Output\n",
|
||||
" \n",
|
||||
" output = torch.zeros_like(out)\n",
|
||||
" output[:, 0] = out[:, 0]\n",
|
||||
" output[:, 1] = self.softplus(out[:, 1])\n",
|
||||
" return output"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 272,
|
||||
"id": "3bc83b15-3b0f-4f21-803b-29e354cb558f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# model = ShallowRegressionLSTM(seq_len)\n",
|
||||
"model = LSTM1(2, seq_len, 128, 1)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 273,
|
||||
"id": "41547213-7851-4ad1-ac04-03cab5e0a91c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def NLL(targets, outputs):\n",
|
||||
" dist = torch.distributions.Normal(outputs[:, 0], outputs[:, 1])\n",
|
||||
" return -dist.log_prob(targets).sum()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 274,
|
||||
"id": "e6958f51-5cd0-414a-b419-3ce55f12b486",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def train_model(data_loader, model, loss_function, optimizer, epochs=200):\n",
|
||||
" num_batches = len(data_loader)\n",
|
||||
" total_loss = 0\n",
|
||||
" model.train()\n",
|
||||
" for epoch in range(epochs):\n",
|
||||
" for X, y in data_loader:\n",
|
||||
" output = model(X)\n",
|
||||
" loss = loss_function(y, output)\n",
|
||||
"\n",
|
||||
" optimizer.zero_grad()\n",
|
||||
" loss.backward()\n",
|
||||
" optimizer.step()\n",
|
||||
"\n",
|
||||
" total_loss += loss.item()\n",
|
||||
"\n",
|
||||
" if epoch%10 == 0:\n",
|
||||
" avg_loss = total_loss / num_batches\n",
|
||||
" print(f\"Train loss: {avg_loss}, Epoch: {epoch}\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 275,
|
||||
"id": "92484e45-11c9-4fe1-bc5b-87b0f27d5a2c",
|
||||
"metadata": {
|
||||
"tags": []
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Train loss: 14.94985544681549, Epoch: 0\n",
|
||||
"Train loss: 2.123787060379982, Epoch: 10\n",
|
||||
"Train loss: -118.20948788821697, Epoch: 20\n",
|
||||
"Train loss: -225.282495072484, Epoch: 30\n",
|
||||
"Train loss: -393.8007117182016, Epoch: 40\n",
|
||||
"Train loss: -546.9331771463155, Epoch: 50\n",
|
||||
"Train loss: -724.232698109746, Epoch: 60\n",
|
||||
"Train loss: -916.3166337996721, Epoch: 70\n",
|
||||
"Train loss: -1073.7028237611055, Epoch: 80\n",
|
||||
"Train loss: -1267.8471668988466, Epoch: 90\n",
|
||||
"Train loss: -1434.5194520920516, Epoch: 100\n",
|
||||
"Train loss: -1630.7973297566175, Epoch: 110\n",
|
||||
"Train loss: -1831.4267205685378, Epoch: 120\n",
|
||||
"Train loss: -2038.7122355431318, Epoch: 130\n",
|
||||
"Train loss: -2246.2204612165688, Epoch: 140\n",
|
||||
"Train loss: -2457.420981016755, Epoch: 150\n",
|
||||
"Train loss: -2669.7478996187447, Epoch: 160\n",
|
||||
"Train loss: -2884.8198515325785, Epoch: 170\n",
|
||||
"Train loss: -3108.3130769401787, Epoch: 180\n",
|
||||
"Train loss: -3337.9945754677055, Epoch: 190\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"optimizer = torch.optim.Adam(model.parameters(), lr=0.01)\n",
|
||||
"train_model(train_loader, model, NLL, optimizer)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 276,
|
||||
"id": "07eab392-37e2-4dae-bfa4-dd177e3d552f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"means = []\n",
|
||||
"vrs = []\n",
|
||||
"for X, y in dset:\n",
|
||||
" output = model(X.unsqueeze(0))\n",
|
||||
" means = means + list(output[:, 0].detach().numpy())\n",
|
||||
" vrs = vrs + list(output[:, 1].detach().numpy())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 277,
|
||||
"id": "b10866df-b554-44a6-bf73-b98abb2e6400",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"<matplotlib.collections.PathCollection at 0x7fa59d357370>"
|
||||
]
|
||||
},
|
||||
"execution_count": 277,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYIAAAEFCAYAAADuT+DpAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/MnkTPAAAACXBIWXMAAAsTAAALEwEAmpwYAABsMElEQVR4nO29eXxdV3nv/Vt7PJOOZsmTPMay4zF2BpyQYAIJIW1DKVMCobS0lFvgbfsylHuhvL2XhLZMbSGhgUJoS5sAKSUlDcRJEy44TuIMTpzYluMhii3Lg2zpSEdn3tNa7x9r731m6WiWddb388kn9jn77LP3tvQ8az3D7yGMMQaBQCAQ1C3SXF+AQCAQCOYW4QgEAoGgzhGOQCAQCOoc4QgEAoGgzhGOQCAQCOocZa4voBZyuRwOHTqE9vZ2yLI815cjEAgEFwWO42BwcBCbNm1CIBCoetxF4QgOHTqE22+/fa4vQyAQCC5K7r//flxxxRVV378oHEF7ezsAfjOLFi2a46sRCASCi4OBgQHcfvvtvg2txkXhCLxw0KJFi7Bs2bI5vhqBQCC4uBgvpC6SxQKBQFDnCEcgEAgEdY5wBAKBQFDnCEcgEAgEdY5wBAKBQFDnCEcgEAgEdY5wBAKBQDAN2JTBoRfneBfhCAQCgWCKMMZwLm6gL5ZFzqJzfTkT5qJoKBMIBIK5xqEMhAASIUWvmzbFuVEDpsOgEIKBhIFFUR2xtImwJsNhQFNQgSyRKmeee4QjEAgEgnEwLIrT8RyCqoQlTcXibYmcDcth0GUCQggMm+JMPAfGgIzhAAAsh6KzQQMh89MZCEcgEAiqwhgDZZjXq9mZZiRtIZa2IBEgY1I4lPnPgzKGRNaG6joBAFAkAocx6AqPvDPGkMw6oNRERJegKTIypoOILhd9zjtf6Y5jNhCOQCAQVIQxhtGsjZGMjWXNOlS5/lKKqZyNWNqCKhNI7mo/Z1GEda7dkzEdOAxQC4y3LBHIyP+dEAJdAXKWg7TpALAAAMNpC5GA7O8UGGPoH86hOaQgGlRn9T7r719WIBCMSzJn42Qsi8GkBZsyDCYtMHZxVsRMBMYYchaFTRkMm2IgYUKRiL9KlwgwmrUBAImshXOj/P3xIIRAlSUEFAm6TBBQJGgyQTLrIOMml23KYNoMQ6nZf9ZiRyAQ1CHVQj6eIRxMmpAIQUDlK9W06cCwGQLqwg4R5WyK08M5gBDoCr/XwmekSAQZ04FpU4zmHCgSqckRFOKFggghUGRgKGki2BKAafNkNGWAQwFlFmdwiR2BQFCHxLM2To/k/JWn7TBQxmA5DGfiBoC8ASSEuCtha86ud7ZI5RwQiUCTCbIWLTPynhEfzdowLAp5in5RkQgsh2E4bSFnOf7rpjO7JahiRyAQ1BkO5YaHMsB0GCTCcGo4B5kQhHS+NizNB6gSQTLnoC3C5n3iOGdR5CwHTaGJxdkdypDI2VDdUFBIrbwkVyTih4emowpIkwniWRuSe24eIqIIabO3JRA7AoGgDqCM+av/nEXBGEAAZE0HGcMBYwBjwGjWqRjq8AzecNqCYc/vhqlY2sRwmsfZHcqQNZ3xPwSeGGasvE+gFIkAINyATweEEEgAGPguTJG4Y5jNPIFwBALBAocyHu5JuKvYtOGAgK8+h1IWhtIWN0BuErPail+VCUYyNoZScxciylkUqZxd9f20YSNjUlDG4/2xtIUzcQNWDaGWjEVRy2aHEAJdlqa1J0CVJWjuLkyWCGyHIZmzkTUdjKRn/nkLRyAQLHDShoOs6WAobYG61TCyRCBLBKpEIJPaEp4SIQgoBFnTgT1HmjqjWQtnR01krfJVPmO8ukmVCCQJOJ8wMZq1QQgwMGrCsOmYWkCWMzc1/JVQJYILKd6/MJKZ+Soi4QgEggVOzqKQCQFjgGFTmHZ+5Su5DqFWvFWwMQd6OowxpA0HEgFG0naZUbcpg00ZJAKokgSHMmgygSZLMB2K/uEczsRzuJAwK+oB2U5tO4LZQJJ4uChtOHx3M8PPWzgCgWCBk7Oob+xThgNGpp7kzFVYkc80psNAwWPzGdPB6XiuaKVs2vzP3r1psuSv8DWZ1+2bNsNIxkLKsGE5FBnT4aW0lJfTzidUWUJQ5SbameEdgagaEggWMIzxChRV5iv/pOHwrOQUkCVeWjmbMMaQM/m1E7fGP2dTZK18dU3adDCWeyOEQJYA4jaFjWZtUAosbdahyvlS2fkEv56Z91BiRyAQzBEOZRMKsUwmTmxTBgbXCBJ+DmmKv/Uy4SEmOs715CwK28kfQxmrGNuvhVPDOQxnbMiFUg4EGM3Y/rmTORvKOJU8isRDRQSAQghAuHroXOU85gvCEQgEc0TasHEhadZ0rE15rX9igk1dhaWThHAjqE7RExA333Bu1Kh6DGUM5+I59I/kfIeRMSnOxs0JG13mNrrxJHf+dUUiSJsOLIci61YK1ZrsVWWJ50dcp+Y4whEIBII5IGNSZC2naNVcjVjKhOlQDE5QhyaZc2YkAarJxDW+5dfCXEVOh/Fdjxe7T+Z4XD6entg9eH4jrMlFoRvvzyMZG/GMNakuX09ILmPRMcNKCx3hCASCWcS0Kc7FDeQsBxnTAUCKEq+WU25cLYcimXOgy5Jb+TO2EXUor66xHFpRJmE68IywVeLETJvi9IiBwRQv4wQAw3L8ip+gKmEky3dCtTqD0u8oRJN5l2/anNx9SgSwbIaM4cz7jumZRCSLBYJZwnIohtMWEjkbhkPhMB7nzlgUkQBfSZ+NGwjrsp8AHU5bUGQCgnziMGM6UGWCrOWUrZIBLo18dtRAUJX8/MDM3ROD7loRyhjOjhpwaH5Ii8QYEjkHulv9IhECXeY7lcZgbSJ2Nq2eRyFubwNjbFL3SQgBA4PpUASUya2LYz29OLt7H8xEGlo0jCU7r0DrxjWTOtdcIRyBQDAL2A7XmrcpQ1CVkLOZL1HgxfENm4dRDJvPAAAAMJ7s9YyUIhHecWo5SBsOljUHyjRpUoYDVeIhj6nmA8bDsPiAFQBIZvmkrkKDqkgEOdvBYCr/Gc+h5SwHAVUCZcx3dN4MBMNm6GhQQQgPQY1n4qfi7DSZgIFM+Byxnl70P74XTs7EsGHjQDyH7ZYD8+HdSJ0+jxU3XTPpa5pthCMQCKYJL6RTmrB0KMOFpAHKuEEnhCDoroQZYzAdromTMmwQCQjIUkHYhICyvKGTCK+nNx3GNWkydpEjcKi7Y5AIVDKzTkAmxG90silDLG1Bq6DWqckSciUxeJkQZAyKxiBD/7ABTSFYFNWQMXkeBAB0hQuw1VINNBUIIVUdTaGxBwA5qKPrhh0AgL5dT4HZ3In/4PVh9KZM/OLMKD6/qRPYfwRD+4/k79f93HzdKQhHIBBMEyNpCwxAW0Tjf89YCKkyUoaNtOFAV8r1abzVcdqw+cjDAulnj0IbSAiBJDFIICCEN3Z5TiOWtmDajAvKzUI9vCxxPZ9E1kIyxztg1QoGWyIEhLCSKV5A1nZgOQyWQ2E6QN9wDpbD/GfgD2ghpOizs0XfY88UGXMAcLIG+h7ZA0lTwWwHhkPx1cMXcN7VP8o6DD/pi2NjUwAxw8HSoIrVEQ3NAE4+vBv9Tzw7Lx2CcAQCwTTAGPObtVrDXOpgKGUhrDnIWRRaBSfgIROCwRSXha7F4BWGeyx3aAzABeEYmz5VzPHwQjnnk3ye71jfq5fIWnufTRncgGoyAaXwcwtAfh7CXExGi/X0+k6AMob7TozgVNrEDYsbsKMtDCfLS2cPxHO+E7iyNYT9wxkciOdwIJ7zzxWQCP5sfTu6whp3JLueAgDfGcyHHINwBALBNGA5zC8DNV3lSAIgbXL9mrHq22UJACWoIn8/LufiBjRF4sZ4kgnPqcCdwOS+N+GWt0qEQKpy/3PR7dv/xLMAeKjtp/1xPB/LAACeOJfEjrawf9yBkSwA7gQ+uLIZo6aDY8ni/oocZfjZ6VH8ybp2AACzHZx8eLf/fmGIyUykyxzFbCAcgUAwRRhjGEqZIMTT9LeRyNnQ5NoSkN7IwskgSzyHkLUoQur0OgFKGR564iVsuGQJ1q1eXPGYWu+xGpbDK4zmmsJVuRzQ/JzAA6fieGYw7R83kLORsSlCigSLMhwe5Sv/31oahSwRvHVRBMeSBq7rCOMtnQ1wGMPXDl/A0YSBr/Scx462ME6kDEiE4C0PPI6uaNB3AknLgUUZWgD0P75XOAKB4GLCtBnSJoUuc6PsSR/PxkpWcaekqNLkyidLKYyL77mQwgN9cYQVCd//83dhxRUbyo6fyneqUnEifK6I9fQWrco9J3BgJIu9rhO4fWUz9g6l8XrKxMm0iQ2NARxL5GBQhmUhFa1uDe3GpiDu2LIITZrs7wJvW9mMH50cQX/GQv+puP+9z8cy+O1ljbhxcQMA4K6jgxjI2vjkpe1Y7V7XbDkD0VAmEEyRZC5v+CXCf6lKq2dmmul0AmcyJr50cAAP9MUBAGmb4t5/fhR9jz0z5e8oRHKH4cw1Z3fv853AiZSB8zkLZzMW/rk3BgbgbYsbcHV7GKvcIoAjozmcyZj49vEYAGBLU7DofC26UhQKvLI1hDu3Lsa17Tyk1BVS0e46jkfPJpC0HGRsinNZGwzAT9znfnb3vhm862LEjkAgmCIpw/GTvIQQaMrcG7dqVEtMeslRxhh+fDKOgZwNXSJYGlLxesrEM4Np/Ma+wwBwUdXH14KZ4Kv+YcPG3x8ZLJKjXhnWcMvSKABga1MQvxxI4cXhDKyCg65sDY37HWFFwm0rm/HOrkYE3HzKPceGcHg0h6cupLGuUfePPZe1QBnzr2s2EI5AIJgCjjsMZbYqdaaCt+LP2BT/8noMy0Oj+K3kk+h79CkwywFlDH996DwGcjY0ieDOrYsRUiR87fB59KUtHIjnoO0/gsiyznlX/jgVtGgYZiKNVxO5spkElzUH/d3WqoiGNl3GkOFgjxsy+pN1bWgP1G5Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"xx = torch.arange(data.shape[0])\n",
|
||||
"\n",
|
||||
"plt.plot(xx[1:], means)\n",
|
||||
"plt.fill_between(xx[1:], means - 2*np.sqrt(vrs), means + 2*np.sqrt(vrs), color=palette[1], alpha=0.5)\n",
|
||||
"plt.scatter(xx, data, color=palette[5])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "6be11181-0698-429c-9dc7-a0f4b6c99584",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Rollouts"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 267,
|
||||
"id": "81bb90d7-e790-4e31-ba87-fe61449eb9ce",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"nroll = 50\n",
|
||||
"roll_len = 100\n",
|
||||
"xin, xout = dset[len(dset)-1]\n",
|
||||
"xx = torch.cat((xin[0, 1:], xout.unsqueeze(0)))\n",
|
||||
"xx = xx.repeat(nroll, 1).unsqueeze(1)\n",
|
||||
"roll_pxs = torch.zeros(nroll, roll_len)\n",
|
||||
"with torch.no_grad():\n",
|
||||
" for idx in range(roll_len):\n",
|
||||
" out = model(xx)\n",
|
||||
" smpl = torch.normal(out[:, 0], out[:, 1])\n",
|
||||
" roll_pxs[:, idx] = smpl\n",
|
||||
" xx = torch.cat((xx[..., 1:], smpl.unsqueeze(-1).unsqueeze(-1)), -1)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 268,
|
||||
"id": "1aec9c78-083b-4161-af02-4ef9e6590858",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAX8AAAEFCAYAAAAL/efAAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/MnkTPAAAACXBIWXMAAAsTAAALEwEAmpwYAABLQElEQVR4nO3de3xcdZ34/9eZM/dMMrk0lyZN06ZtStMrKWBQsCK3roAXEIpU2WVl/e667rq7qLuy7E+BXffL7uoq+FUUdFcEEYEqolCkUEqFtkDT0qalCb0kTdrmfp3J3Of8/pg5JzOZSZq0SdNk3s/Hw4fNzJmZk055n895f96f90fRNE1DCCFERjFN9wkIIYQ49yT4CyFEBpLgL4QQGUiCvxBCZCAJ/kIIkYHM030C4+H3+6mvr6ewsBBVVaf7dIQQYkaIRCJ0dnayYsUK7HZ70nMzIvjX19ezcePG6T4NIYSYkZ544gkuuuiipMdmRPAvLCwEYr9ASUnJNJ+NEELMDG1tbWzcuNGIoYlmRPDXUz0lJSXMmzdvms9GCCFmlnTpcpnwFUKIDCTBXwghMpAEfyGEyEAS/IUQIgNJ8BdCiAwkwV8IITKQBH8hhIjr6uoiEAic9fuEQiE6Ojom4YymjgR/IUTGOHbsGF1dXaM+f/ToUbq7u8/ovTVNIxKJ0NzcTFNTE+++++6ZnuY5IcFfCJEx3nnnHXp7e9M+19LSQiAQwOv1ntF7nzp1ir1799LV1UVjYyPt7e1A7C6gubmZUChEMBgc13t1dXXR399/RucxXhL8hRAzXk9PT8pjoVCIgYEB42ev14vD4cDn8/Haa6/h8/mSjt+zZw99fX1JwX+8wRqgvb2d/v5+hoaG8Hq9hMNhAAYGBnjvvfc4cuQIR44cSXndwMAAx48fByAQCNDS0kJdXR3Nzc3j/uwzIcFfCDHjvfjii3i9Xurr643HmpqaOHz4MBAL7O+88w4XXnghQ0NDNDY20tfXZxzb0dFBR0cHzc3N9Pb20tbWBsAzzzxj/FnTNLxeLy+//LLxmE7TNLZu3UokEqG/v59oNIrb7cbn83H06FF6e3vZu3cvp06d4sSJEwQCATRNY+fOndTX1/P222+jaRoDAwM0NDQwODhopJ9aWlqm5O9sRvT2EUKIsTgcDhobG2loaODkyZMUFxfT39+P2Wzm6NGj7Ny5E6fTic/nIxQKMX/+fLq7uyksLMRsNvPSSy/R3t6OyWTiwIEDdHd3U1paSk5ODg0NDZSUlDAwMEBdXR379u1j//793HDDDSxZsoTt27fT3NyMx+Phvffew2azkZWVRUFBAb29vdTV1REMBunp6aG7u5tIJEJlZSXZ2dns2rULRVGw2+2cPHmSV199lVOnTmG1Wjl58iQXXngh77//PuXl5ZP+dybBXwgxo0UiERRFoaOjg7a2Ntra2mhubmbBggVEo1H27t3LsmXL2LdvHwAej4fCwkK2b99OOBxm9erVtLa2Eg6HMZvN9PT0YDKZOHToEAUFBbjdbt555x1KS0vp7u7GZrOhaRrbt29nyZIlHDp0iI6ODjRNo6+vj9zcXLxeL263m56eHrq6unA4HKiqysDAgHE34na7CQQCWK1W+vv7efPNN2lqaiIajZKVlUVfXx8vvvgiOTk5U/L3JsFfCDEjHTlyhEWLFnHy5EkjXRIOh7HZbPT19dHe3s68efMYGBhg0aJFDA4OAjA4OIjf7wdg9+7dzJ8/38jth8NhTCYT4XCYSCRCT08PHo+HU6dOcd111xk5fZPJxMDAANFolM7OTiO/DxiPNzU10drais1mw+v1YjabCQaDdHZ24nK5UFWVaDSKx+MhEolw/Phx4zGfz0d2djZtbW14PJ4p+fuT4C+EmJG2b9+O3+9n+/btRCIRAoEAiqKgaRqhUIihoSGi0SihUIgdO3ZgMpnw+/1Eo1EikQiqqqIoCi+88AJms9kI4KqqGpU2kUgEn8+H1Wpl7969BINBwuGwEaR37NhBJBIxzkkv9wSwWq10d3djsViIRqOYTMNTrK2trXR0dBAKhYhGo0DsjkRRFCA2OW2327HZbMaFarJJ8BdCzEitra04nU66urqMwKoHcT2Fc+LECbq7uzGZTGiaZlwgAoEAOTk5tLW1UVlZmRTAw+EwmqYBsWDudDoZGhriyJEj2Gw2wuEwoVCI/Px8tm7dahw7ks/nIxqNMjQ0BMSqjwAURSEYDKIoihH4dYnv1dvbi8lkwuVyTd5fWgKp9hFCzDh6sG5qakoKojabzXhOr+AJBoOoqkooFELTNGMk7vF4CIVCNDQ0JAXdkcFcD96KojA0NGQ87/V6iUQixs/6qF3X39+fdFHR/2w2x8bcp1tJHI1GCYfDOJ3OCfzNjJ8EfyHEjLNnzx4cDoeR5tEDamJ+PBwO43A40DTNSJ3oKSH9eSApHTOWxDsCRVGS8vz6e4/HeNYO2Gw248/6YrHJJmkfIcSUaalr5ODmnfj6PDhyXVSvr6W8puqs37e1tRWr1UpPTw92uz1p1K0HYX3CFkgK1ImjcSAl9QIYF5VEiT+nS9mMh8lkGtfrEu8KxntRmfC5TMm7CiEy3t5N29j9yy34+jwMhiJ0dw2w+5db+OOPnjvr9+7u7jZy4ZqmJQXUxJH8mQTodIF/JE3Tzigon8lr9LuaySYjfyHEWRs5wi++oIKmnQcA8IQi3L+/jZAGGypyqeUEv/naD4BYoK34QDVrblw3oc/z+/1GINXLMnUjq2/GS1EUFEUhNzc3bbuIRGc6Gtc0DZvNNma+f+TFJ93m65NBgr8Q4qy01DVS9/SrDAXDDAQjFGmDRuDXNI3nWvsZisSC2ePHenn51CBWk8LqPAdr851o8WMncgEIhUJGHn9kGudsRKNRysvLjeA/njSNnnIaTyooMZef+FpN04zS05FzCfocxWSTtI8Q4qzs3bSNaDjCDxq6uL++naeP9xnP7e/zs6NrKOn4dn+YlqEQvzsxwM+OxoJs084DtNQ1jvszvV5vSg5+NOMdOVssFkwmE2azGavVCsTaRqSbEB45x6CnnkZ+VnZ2dtLPJpMpKZgnvreiKFit1nFPQJ8tGfkLIc7Y3k3biARD7O7xccwbq2J5t9fHLRV5ABzoH16gtDrPwbu9sU6aORYTA6EoTd4gf/N2KzfNdxN9+lUAymuqjDTSUO8gzrzspIniUChEOBzGYrEY761X/IwcNcPwnUFiOsVisSQFYf21iqJw4MAB4/30VJBOVVWjnUTixUe/Q7BYLEmjd5fLRTQaxev1GoE/8e7A6XQaK4/10k6d/hllZWXj/DYmRoK/EOKMtNQ10rTzAJ5QhCeahnvk94eiBCJRrCaFxoFYbvvPKvO5wG0jFI1ySUEWFxU4+dnRHt7uHkIDnjnez54eH18Ov0x30yladjfwu6YeXu/w8LcXhAk++xoQuzC88847QGo6JDFIJ14I9EVaFouFSCSSkibSF1L5/f6k+QO9lFS/I9A/T78AJH6ufiGw2WxG8FdVFZfLhdvt5tChQ0nnZzKZyM3NJT8/n9LSUhobG3E6ncYdjcViMS5GU3UnIGkfIabBG3VH+M+fvow/ODX53KnUUtfIS996jN2/3ALAjq4hQlGNqmwbJfbYeLLDH+aYJ0hnIIzLbOLCfAcus8oXqwq5qCC2aOnm+bncuaiAuY7Ya454gjxwoINXt9QRCYX5/ckBBsNRftDYRSQU5uDmnQBJbZsTA2riqHnJkiVJxyiKQiQSwW63A7ELh8lkwuFwUFZWht/vNwJ1NBo1WkAEAgFsNhuFhYVALLjPmTMHGE4nqaqK0+nEZDJhMpmwWq2YzWZyc3Nxu91cdtllZGdno2kaeXl5WCwWFixYgNPpZN68eRQUFGC1Wlm5ciWqqnLBBRcwf/58nE4niqIk7UkwmST4C3EOtdQ18vz/9yif++pP+d7Pt/KFP//OhHLd062lrpG9z76Gr8+Dpmm0+0P8/kSsD85Hil0UO2KpmCZvkP+J5/MvyneipsnJO80m1uQ7+KflxXyxag4Wk8IJX4jftibvYNUbjI2yfX2xBVyjBcPENMzChQuNEbPX6zVG0CUlJcYFw2Kx4Ha7ufbaa1FVlTvvvJPs7GyjFYTJZKKwsJCCggIKCwux2+2sWbMGi8WCxWIx8vOKopCVlUVubi6hUIi5c+dy2WWXsXz5cjRNw+fzsXr1agAWLFiA1WolOzub7OxsAoEAbW1trFy5krKyMvLz87nmmmsIBAKUlpYaF62pIMFfiHPkjz96jv99+Hn+5a1mgtFYoPrjqQHe/MUfZswF4ODmnURCYbacGuRv3znB/fvbCWtQ7baxMtdOpSs2Ufp86wC9wQgWBa4tzR7zPVVFodpt595VJQAcHgzS4U++IwpFNSNoJzY60wO+1WpFVVVj9L1q1SojMANGPr6srIzy8nJsNhsulwu73U5JSQl5eXnYbDacTic2m41oNEplZSWBQICKigrC4TD5+fnk5ubidDqZP3++cRHQP6OiooKioiLmz5+PzWajsrISk8nE888/j8ViMd4/NzeXwcFB5s2bx7Fjx1AUheLiYsrKyoy7A6/XS21tLXl5eRQVFU3CN5fqnAX/TZs2sXTpUiNfJ8RsFo1GGfQOB6m9m7bRcbiVX7f00xUYHskFoxr7e/3seXbrdJzmhOmj7z29sVy9bllObJVtTZ4DgKFIbFLz4oIssi3jq7bJsahcnB9LCT3fmjy6HwhFjNYMenoncQSvV+hkZ2eTl5eHyWQyKm304/Sfr732WqLRKEuWLMFisaCqqpEmcrvd2Gw2LBYLH/vYx/D5fNTU1BAOh1m4cCENDQ2YzWbcbrcx4Wwymejv7zc2cTGbzdhsNvLy8jh8+DCqqnLw4EFWrFjBkiVLqKioIDs7mzVr1hAKhbj11ltxuVz4fD6WLVsGwBVXXMG8efNYvHgxNTU14/x2JuacBP89e/Zw//33n4uPEmLatdQLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"full_x = torch.arange(data.shape[0])\n",
|
||||
"test_x = torch.arange(data.shape[0], data.shape[0] + roll_len)\n",
|
||||
"plt.plot(full_x[1:], means)\n",
|
||||
"plt.scatter(full_x, data, color=palette[5])\n",
|
||||
"plt.plot(test_x, roll_pxs[:20, :].T.detach(), color='gray', alpha=1., lw=0.5)\n",
|
||||
"plt.show()"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.12"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,454 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 7,
|
||||
"id": "92e5b691-e01a-438a-b2b9-543c5fa3f8b2",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import pickle as pkl\n",
|
||||
"import pandas as pd\n",
|
||||
"import numpy as np\n",
|
||||
"import torch\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns\n",
|
||||
"import os\n",
|
||||
"from voltron.data import make_ticker_list, GetStockHistory\n",
|
||||
"\n",
|
||||
"sns.set_style('white')\n",
|
||||
"palette = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"sns.set(palette = palette, font_scale=2.0, style=\"white\", rc={\"lines.linewidth\": 4.0})"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 8,
|
||||
"id": "a519062e-3992-4d40-9bf1-990c80990fc0",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"matern_df = pd.read_pickle(\"./matern_calib.pkl\")\n",
|
||||
"volt_df = pd.read_pickle(\"./volt_calib.pkl\")\n",
|
||||
"matern_df.columns = volt_df.columns"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 9,
|
||||
"id": "37d94b48-6e05-4a59-91a3-11291bfc7f4d",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"array(['tewma', 'linear', 'constant', 'ewma'], dtype=object)"
|
||||
]
|
||||
},
|
||||
"execution_count": 9,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"matern_df.Mean.unique()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 10,
|
||||
"id": "1f59e48e-eb80-48b2-8ce5-91596ad136f2",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def ECDF(sample_pxs, true_px): \n",
|
||||
" return (torch.sum(sample_pxs < true_px, 0)/sample_pxs.shape[0])\n",
|
||||
" \n",
|
||||
"def Calibration(pcts, percentile=0.95):\n",
|
||||
" in_band = np.where((pcts < percentile))[0].shape[0]\n",
|
||||
" return in_band/pcts.shape[0]\n",
|
||||
"\n",
|
||||
"def GetCalibration(model, horizon=np.arange(75,100), \n",
|
||||
" logger=[], exp=True):\n",
|
||||
" \n",
|
||||
"\n",
|
||||
" ntrain = 400\n",
|
||||
" n_test_times = 20\n",
|
||||
" ntest = 100\n",
|
||||
" pcts = torch.tensor([])\n",
|
||||
" for tckr in ticker_list:\n",
|
||||
" data = GetStockHistory(tckr, history=1000, end_date=end_date)\n",
|
||||
" for idx, date in enumerate(data.index):\n",
|
||||
"\n",
|
||||
" fpath = \"./saved-outputs/\"+ tckr + \"/\"\n",
|
||||
" fname = model + \"_\"\n",
|
||||
" if model == 'volt':\n",
|
||||
" fname += \"constant\"\n",
|
||||
"\n",
|
||||
" fname += str(date.date()) + \".pt\"\n",
|
||||
"\n",
|
||||
"\n",
|
||||
" if os.path.exists(fpath + fname): \n",
|
||||
" preds = torch.load(fpath + fname)\n",
|
||||
" if isinstance(preds, tuple):\n",
|
||||
" preds = preds[0]\n",
|
||||
" \n",
|
||||
" if preds.shape[-1] == 100:\n",
|
||||
" preds = preds[:, horizon]\n",
|
||||
"\n",
|
||||
" test_y = torch.tensor(data.iloc[idx:idx+100].Close.to_numpy())\n",
|
||||
" if test_y.shape[0] == 100:\n",
|
||||
" if exp:\n",
|
||||
" preds = preds.exp()\n",
|
||||
" pcts = torch.cat((pcts, ECDF(preds, test_y[horizon])))\n",
|
||||
" \n",
|
||||
" if pcts.numel() == 0:\n",
|
||||
" return logger\n",
|
||||
" \n",
|
||||
" pcts = pcts.flatten().numpy()\n",
|
||||
" percentiles = np.linspace(0.05, 0.95, 19)\n",
|
||||
" for pct in percentiles:\n",
|
||||
" clb = Calibration(pcts, pct)\n",
|
||||
" logger.append([clb, np.round(pct, 2), model, \"Constant\", 100])\n",
|
||||
" \n",
|
||||
" return logger"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 11,
|
||||
"id": "5fe6b4d6-b49e-4e06-8874-9be43b4ee4bd",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"data_path = \"../../voltron/data/\"\n",
|
||||
"ticker_list = make_ticker_list(data_path + \"test_tickers.txt\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 12,
|
||||
"id": "da419f67-e0d9-4008-b038-96fd341e10f3",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"end_date = \"2022-01-13\""
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 13,
|
||||
"id": "d9ed9129-be2d-47e8-9866-5a708c27d225",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"log = []\n",
|
||||
"for model in ['volt', 'lstm']:\n",
|
||||
" log = GetCalibration(model, horizon=np.arange(75,100), \n",
|
||||
" logger=log, exp=True)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 14,
|
||||
"id": "89054cf9-fb1c-45a1-aa8b-8884713f4c07",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df = pd.DataFrame(log)\n",
|
||||
"df.columns = ['Calibration', 'Percentile', \"Model\", \"Mean\", \"k\"]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 26,
|
||||
"id": "9097eb62-95d8-443d-bbd4-aae81a3445fa",
|
||||
"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>Calibration</th>\n",
|
||||
" <th>Percentile</th>\n",
|
||||
" <th>Model</th>\n",
|
||||
" <th>Mean</th>\n",
|
||||
" <th>k</th>\n",
|
||||
" </tr>\n",
|
||||
" </thead>\n",
|
||||
" <tbody>\n",
|
||||
" <tr>\n",
|
||||
" <th>0</th>\n",
|
||||
" <td>0.024371</td>\n",
|
||||
" <td>0.05</td>\n",
|
||||
" <td>volt</td>\n",
|
||||
" <td>Constant</td>\n",
|
||||
" <td>100</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>1</th>\n",
|
||||
" <td>0.031473</td>\n",
|
||||
" <td>0.10</td>\n",
|
||||
" <td>volt</td>\n",
|
||||
" <td>Constant</td>\n",
|
||||
" <td>100</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>2</th>\n",
|
||||
" <td>0.042644</td>\n",
|
||||
" <td>0.15</td>\n",
|
||||
" <td>volt</td>\n",
|
||||
" <td>Constant</td>\n",
|
||||
" <td>100</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>3</th>\n",
|
||||
" <td>0.058020</td>\n",
|
||||
" <td>0.20</td>\n",
|
||||
" <td>volt</td>\n",
|
||||
" <td>Constant</td>\n",
|
||||
" <td>100</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>4</th>\n",
|
||||
" <td>0.073815</td>\n",
|
||||
" <td>0.25</td>\n",
|
||||
" <td>volt</td>\n",
|
||||
" <td>Constant</td>\n",
|
||||
" <td>100</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>...</th>\n",
|
||||
" <td>...</td>\n",
|
||||
" <td>...</td>\n",
|
||||
" <td>...</td>\n",
|
||||
" <td>...</td>\n",
|
||||
" <td>...</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>337</th>\n",
|
||||
" <td>0.793388</td>\n",
|
||||
" <td>0.75</td>\n",
|
||||
" <td>volt</td>\n",
|
||||
" <td>tewma</td>\n",
|
||||
" <td>400</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>338</th>\n",
|
||||
" <td>0.812314</td>\n",
|
||||
" <td>0.80</td>\n",
|
||||
" <td>volt</td>\n",
|
||||
" <td>tewma</td>\n",
|
||||
" <td>400</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>339</th>\n",
|
||||
" <td>0.840826</td>\n",
|
||||
" <td>0.85</td>\n",
|
||||
" <td>volt</td>\n",
|
||||
" <td>tewma</td>\n",
|
||||
" <td>400</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>340</th>\n",
|
||||
" <td>0.879091</td>\n",
|
||||
" <td>0.90</td>\n",
|
||||
" <td>volt</td>\n",
|
||||
" <td>tewma</td>\n",
|
||||
" <td>400</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>341</th>\n",
|
||||
" <td>0.944463</td>\n",
|
||||
" <td>0.95</td>\n",
|
||||
" <td>volt</td>\n",
|
||||
" <td>tewma</td>\n",
|
||||
" <td>400</td>\n",
|
||||
" </tr>\n",
|
||||
" </tbody>\n",
|
||||
"</table>\n",
|
||||
"<p>494 rows × 5 columns</p>\n",
|
||||
"</div>"
|
||||
],
|
||||
"text/plain": [
|
||||
" Calibration Percentile Model Mean k\n",
|
||||
"0 0.024371 0.05 volt Constant 100\n",
|
||||
"1 0.031473 0.10 volt Constant 100\n",
|
||||
"2 0.042644 0.15 volt Constant 100\n",
|
||||
"3 0.058020 0.20 volt Constant 100\n",
|
||||
"4 0.073815 0.25 volt Constant 100\n",
|
||||
".. ... ... ... ... ...\n",
|
||||
"337 0.793388 0.75 volt tewma 400\n",
|
||||
"338 0.812314 0.80 volt tewma 400\n",
|
||||
"339 0.840826 0.85 volt tewma 400\n",
|
||||
"340 0.879091 0.90 volt tewma 400\n",
|
||||
"341 0.944463 0.95 volt tewma 400\n",
|
||||
"\n",
|
||||
"[494 rows x 5 columns]"
|
||||
]
|
||||
},
|
||||
"execution_count": 26,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"df"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 15,
|
||||
"id": "185272ab-5724-4dc8-add6-42f054832940",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df = pd.concat([df, matern_df, volt_df])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 16,
|
||||
"id": "81e78fa5-dfcb-4c71-9b09-094be0e58b36",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"mat_df = df[(df['Model'] == 'matern') & (df['Mean'] == 'tewma') & (df['k'] == 400)]\n",
|
||||
"lstm_df = df[df['Model']=='lstm']\n",
|
||||
"volt_df = df[(df['Model'] == 'volt') & (df['Mean'] == 'ewma') & (df['k']==100)]\n",
|
||||
"plt_df = pd.concat([lstm_df, mat_df, volt_df])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 27,
|
||||
"id": "3b8bf1a9-42f3-4a5a-aafb-3385cd32ea9e",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"mat_df = df[(df['Model'] == 'matern') & (df['Mean'] == 'constant')]\n",
|
||||
"volt_df = df[(df['Model'] == 'volt') & (df['Mean'] == 'Constant')]\n",
|
||||
"const_df = pd.concat([mat_df, volt_df])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 41,
|
||||
"id": "bf42d9c5-0bc3-46a8-a1ab-14b87fd8851b",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAABPoAAAIRCAYAAADTKdPXAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAABcSAAAXEgFnn9JSAAEAAElEQVR4nOzdd3hTZfsH8O/Jbrp3gQKFliGj7A0qU+EHshRxgDhwgeBCxdetr7gnr4ATxYlQQFBBEGTvVUaBDgp00b2SZp/fH6WxadI2SdOWlu/nurhIznlWmrQ5ufM8zy2IoiiCiIiIiIiIiIiImjRJYw+AiIiIiIiIiIiI6o6BPiIiIiIiIiIiomaAgT4iIiIiIiIiIqJmgIE+IiIiIiIiIiKiZoCBPiIiIiIiIiIiomaAgT4iIiIiIiIiIqJmgIE+IiIiIiIiIiKiZoCBPiIiIiIiIiIiomaAgT4iIiIiIiIiIqJmgIE+IiIiIiIiIiKiZoCBPiIiIiIiIiIiomaAgT4iIiIiIiIiIqJmgIE+IiIiIiIiIiKiZoCBPiIiIiIiIiIiomaAgT4iIiIiIiIiIqJmQNbYAyAiau46depkc3/u3Ll47LHHGmk01460tDSMHDnS5tiiRYswZcqURhoRNUeffvopFi9ebHPs7Nmz9VZvxowZOHDggPV+//79sWLFCidHS/WJzw0RERFdDRjoI2omTCYTkpKSkJKSguLiYhQXF8NiscDLywtqtRoRERFo1aoVIiMjoVAoGnu4RHSNuXz5MlJSUpCeno7i4mLodDqoVCr4+vrC398fbdq0QYcOHSCVSht7qERERERETRYDfURNmMFgwObNm7F69WocPnwYOp2u1jpyuRwdOnRA9+7d0a9fPwwZMgRBQUENMFpqSkaMGIH09HSnyspkMvj4+MDX1xctWrRA165dERsbi+HDh8PLy6ueR0pXK4vFgp07d2LTpk3YtWsXLl++XGsdLy8vdOnSBTfeeCMmTJiAFi1aNMBIiYiIiIiaDwb6iJqov//+G6+99hqysrJcqmc0GnH69GmcPn0av/zyCyQSCe666y688MILtdblsiRyxGQyobCwEIWFhbh06ZL1NeLr64uJEydi7ty5CAwMbORRUkMRRRFr1qzB0qVLceHCBZfqlpWV4fDhwzh8+DA++OADDBgwAHPmzEH//v3rabR0LXvuueewZs0a6/1WrVph69atjTgiIiIiorpjMg6iJkYURbzyyit49NFHXQ7yOWKxWJCRkeGBkRHZKikpwffff4/x48dj+/btjT0cagAXL17E9OnTsXDhQpeDfFWJooh9+/ZhxowZePDBB5GWluahURIRERERNV+c0UfUxLz88sv45ZdfHJ5r2bIlBg4ciJiYGAQFBcHLywtarRbFxcVITU3FqVOncObMGRgMhgYeNTUHnTt3dnjcaDSiuLgYOTk5Ds/n5uZizpw5+PzzzzF48OD6HCI1ot27d2PevHkoLS11eF6hUKB3796IjY1FUFAQAgMDoVQqodFokJGRgcTERBw8eBAFBQV2dbdv344DBw4gMjKyvh8GEREREVGTxkAfUROyZcsWh0G+rl27YsGCBRg4cCAEQaixjbKyMuzcuRObN2/Gli1boNVq62u41MysW7euxvP5+fn4559/8PXXXyMxMdHmnNFoxGOPPYa//voLwcHB9TlMq8jISKeymFLd/fPPP5g7dy6MRqPduejoaMydO9epPRstFgsOHDiAX3/9FRs3boTJZKqvITe6xx57jNm3mxluZUFERERXAy7dJWoiRFHEm2++aXd8zJgx+PnnnzFo0KBag3xA+Wb3Y8aMwbvvvosdO3Zg4cKFaNu2bX0Mma4xQUFBmDJlCuLi4nDbbbfZnS8tLcXixYsbYWRUnxISEvDEE0/YBfnkcjleeuklrF+/HuPGjXMqMYtEIsHAgQPx/vvv448//sDw4cPra9hERERERM0SA31ETcSRI0fssqCGh4dj0aJFUCgUbrXp6+uLWbNm4dlnn/XEEIkAlC/RfO211zBkyBC7c2vWrOHS8WZEr9fjySeftJsZrFarsWzZMtx1112QSqVutd22bVssXboUb7/9NtRqtSeGS0RERETU7DHQR9RE7Nixw+7Y5MmT4ePj0wijIaqZRCLB008/bXe8IqsqNQ9Lly5FSkqK3fGPPvrIYaDXHZMmTcJPP/2EiIgIj7RHRERERNSccY8+oibCUWbcbt26NcJIGsb58+eRnJyMvLw8FBYWwsvLC8HBwYiIiEBsbCzkcnm99V1WVoYTJ04gJycHBQUFKCkpgUqlgq+vL6KiohAdHY3AwMB667+56NKlC1q1amU3E/XUqVMYNGiQ2+2mpqbizJkzyMrKglarhVwuR2hoKCZNmlTHETvHYDDg5MmTyMrKQmFhIYqLi6FQKODj44PWrVsjJiYGoaGhHuvvwoULSEpKQn5+PgoKCqBQKBAQEICIiAj07NkTKpXKY325Ii8vD8uXL7c7fscdd+CGG27waF/VJYKpzeXLl5GSkoK0tDSUlpZCp9PBx8cH/v7+aNmyJbp37w6lUunRsV4tLl26ZH2d6nQ6BAUFITw8HD179kRAQEC991/X39Nr+bmrLCMjAwkJCTa//0FBQQgLC2uw33+j0Yj4+HgkJyejoKAAMpkMQUFBiIqKQmxsrNuzdomIiKh+MNBH1ETk5+fbHXNmz6u66tSpU7XnDhw4UOP5Cn///bdT2TIvX76ML7/8Elu3bkVaWlq15by9vTFo0CDMnDkTAwYMqLVdZ+j1emsCgGPHjjlMKlBBEAR06tQJN9xwA6ZMmYKoqCiPjKE6BoMBzz//PNavX29zPCIiAp9//rlTz0Fj6dChg12gz9FrGbB/rc2dO9earECr1eL777/HypUrcenSJYf1qwYQ0tLSMHLkSJtjixYtwpQpU1x5CAAAs9mM9evXY/369Th8+DDKyspqLB8VFYXrr78ekydPRpcuXVzu79KlS1i+fDm2b99e7eMFAKVSib59++Kee+7xeHCtNitXrrRbsuvr64sFCxY06Dgqy8/Px5YtW7Bnzx4cPHgQubm5NZaXy+Xo2bMn7rrrLtx0002QSBpuocOnn35qt2dlXZPHiKKIuLg4LF++HOfOnXNYRi6XY+DAgXjwwQfRv39/l/vw9O9phYZ67kaMGGH3N6lCenq6U39Pv/vuO4fvPTNmzMCBAwes9/v37+9Wgo7i4mJ8/fXX2Lx5M5KSkqotp1Qq0a9fP0yfPh2jR492uZ+4uDgsXLjQ5ljl9+ucnBx8/vnnWLNmDUpKShy24efnh8mTJ+ORRx7hF2BERERXCQb6iJoIR/vwOZrl1xSZzWYsXrwY33zzTa0BFADQaDTYsmULtmzZghtuuAGvvPIKWrZs6Xb/P/30E/73v/8hJyfHqfKiKOLMmTM4c+YMli1bhk8++QQ33XST2/3XpKioCHPnzrX58AgA1113HZYtW4bw8PB66ddTHC0tr+4DY3WOHz+Oxx9/vNFe75s2bcL777+PCxcuOF0nNTUVqamp+O677/Dcc8/h3nvvdapeaWkpPvjgA6xcubLGYHMFvV6P3bt3Y/fu3ejTpw/ee++9Ov0uuCIuLs7u2KRJk+Dt7d0g/Vf11FNPuZyp12g04uDBgzh48CCio6Px8ccfo0OHDvU4yvqTn5+Pxx57DIcOHaqxnNFoxM6dO7Fr1y5MnToVL774okdmhdXl9/Raf+4qW7FiBT799FMUFRXVWlav12PXrl3YtWsXevXqhVdffdVjX/xs2rQJL7zwAoqLi2ssV1xcjG+//Rbr1q3DsmXL0LNnT4/0T0RERO7jHn1ETYSjpYB//vlnI4zEs8rKyjBnzhx89tlnTgX5qtq+fTtuv/12nDlzxuW6er0eTz/9NF555RWng3yOaDQat+vW5NKlS5g+fbpdkG/YsGH4/vvvr/ogH1AeuKrK19fX6foHDx7EjBkzGiXIZ7FY8M4772DevHkuBfmqcvQzcCQ9PR133HEHfvjhB6eCfFUdPnwYt912G44dO+ZyXVclJSXh4sWLdsdvv/32eu+7OkePHnUpUFRVcnIypk2bhj179nhwVA2jqKgId911V61BvspEUcSqVavw8MMPQ6fT1an/uv6eXsvPXQWz2YyXXnoJb7zxhlNBvqqOHj2KO++8E3v37q3zWH766SfMnz+/1iBfZYWFhbj33nuRkJBQ5/6JiIiobjijj6iJ6NWrF3755RebY3v27MGKFSswY8aMeuu38t5YFy9etFmqp1ar0aZNm1rbqG4/PYvFgkcffdThhzNvb28MHz4csbGxCA0NRWlpKVJTU7Flyxa7oEt2djbuvvturF69Gm3btnXqcRmNRtx///04ePCg3TmJRIKuXbti0KBBaNGiBQICAmAwGFBYWIizZ88iPj6+xuVUnhAfH4+HH34YeXl5NsenTZuGl19+GTJZ0/jznZiYaHcsKCjIqbo5OTmYO3cu9Hq99VhsbCyGDBmCVq1awdvbG9nZ2UhOTsbGjRs9NuYKCxYswIYNGxye69ixIwYPHow2bdogMDAQRqMRRUVFSEpKwsmTJ3H69GmIouh0X+np6Zg2bZrD5YqxsbHo3bs32rVrBz8/PxiNRuTk5ODo0aPYsWOHTRbj3NxcPPTQQ4iLi0OrVq1cf9BO2r9/v92xkJCQq2ZGlVQqRZcuXdChQwe0a9cOgYGB1pmGFX9Ljh8/jiNHjsBisVjrabVaPPHEE1i7di1atGjRWMN32TPPPGOTFKVFixYYPXo0oqOj4efnh9zcXJw8eRJ///23XeB57969eOKJJ7BkyRK3+vb072l9P3fR0dHWLxsyMzNtgmpyuRzR0dG1jrE+skC/+OKLWL16td1xpVKJoUOHol+/fggNDYVOp0N6ejr+/vtvu6XepaWlmD17Nr799lv06dPHrXHs2LEDr7/+uvXvl6+vL4YMGYJevXohODgYFosF6enp+OeLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 900x450 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"from matplotlib.lines import Line2D\n",
|
||||
"fig, ax = plt.subplots(1,1,dpi=150, figsize=(6, 3))\n",
|
||||
"\n",
|
||||
"percentiles = np.linspace(0.05, 0.95, 19)\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"pal = [palette[5], palette[7]]\n",
|
||||
"# pal = [palette[5]]\n",
|
||||
"sns.lineplot(x='Percentile', y=\"Calibration\", hue='Model', data=const_df, ax=ax, alpha=0.2,\n",
|
||||
" palette=pal, legend=True)\n",
|
||||
"sns.scatterplot(x='Percentile', y=\"Calibration\", hue='Model', data=const_df, ax=ax, s=120, legend=False, zorder=4,\n",
|
||||
" palette=pal, alpha=0.35)\n",
|
||||
"\n",
|
||||
"pal = [ palette[0], palette[4], palette[6]]\n",
|
||||
"sns.lineplot(x='Percentile', y=\"Calibration\", hue='Model', data=plt_df, ax=ax, alpha=0.5,\n",
|
||||
" palette=pal, legend=True)\n",
|
||||
"sns.scatterplot(x='Percentile', y=\"Calibration\", hue='Model', data=plt_df, ax=ax, s=120, legend=False, zorder=4,\n",
|
||||
" palette=pal)\n",
|
||||
"x = np.linspace(0.05,0.95)\n",
|
||||
"y = np.linspace(0, len(percentiles))\n",
|
||||
"ax.plot(x, x, color=\"gray\", lw=1., ls=\"--\")\n",
|
||||
"ax.set_title(\"Stock Price Calibration\")\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"plt.tick_params(labelsize=16)\n",
|
||||
"sns.despine()\n",
|
||||
"\n",
|
||||
"custom_lines = [Line2D([0], [0], color=palette[0], lw=2),\n",
|
||||
" Line2D([0], [0], color=palette[4], lw=2),\n",
|
||||
" Line2D([0], [0], color=palette[6], lw=2),\n",
|
||||
" Line2D([0], [0], color=palette[5], lw=2),\n",
|
||||
" Line2D([0], [0], color=palette[7], lw=2)]\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"plt.legend(custom_lines, ['LSTM', r\"Matérn + Magpie\", \"Volt + Magpie\", \"Matérn + Constant\",\n",
|
||||
" \"Volt + Constant\"],\n",
|
||||
" fontsize=14, frameon=False, bbox_to_anchor=(1., 0.85))\n",
|
||||
"# ax.legend(fontsize=14, bbox_to_anchor=(1., 0.75))\n",
|
||||
"# plt.label(\"Percentile\")\n",
|
||||
"plt.savefig(\"./stock_calibration.pdf\", bbox_inches=\"tight\")\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "957ddda1-e78b-4bb7-9003-886f18e424f7",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "d9fa07e1-f9e5-4309-9311-160c53b43d0d",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,152 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
import pandas as pd
|
||||
import gpytorch
|
||||
import argparse
|
||||
import datetime
|
||||
import warnings
|
||||
|
||||
from voltron.data import make_ticker_list, GetStockHistory
|
||||
import sys
|
||||
sys.path.append("../trading/")
|
||||
from GenerateMultiMeanPreds import GenerateStockPredictions, GenerateBasicPredictions
|
||||
from gpytorch.utils.warnings import NumericalWarning
|
||||
warnings.simplefilter("ignore", NumericalWarning)
|
||||
|
||||
def main(args):
|
||||
|
||||
if args.end_date is None:
|
||||
end_date = datetime.date.today()
|
||||
else:
|
||||
end_date = datetime.datetime.strptime(args.end_date, "%Y-%m-%d")
|
||||
|
||||
data = GetStockHistory(args.ticker, history=args.ntrain + args.lookback)
|
||||
|
||||
ntest = args.forecast_horizon
|
||||
ntrain = args.ntrain
|
||||
n_test_times = args.n_test_times
|
||||
ntime = data.shape[0]
|
||||
|
||||
test_idxs = torch.arange(ntrain, ntime-ntest,
|
||||
int((ntime-ntest-ntrain)/n_test_times))
|
||||
|
||||
train_x = torch.arange(ntrain) * dt
|
||||
test_x = torch.arange(ntest) * dt + train_x[-1] + dt
|
||||
|
||||
if torch.cuda.is_available():
|
||||
train_x = train_x.cuda()
|
||||
test_x = test_x.cuda()
|
||||
|
||||
####################
|
||||
## setup filename ##
|
||||
####################
|
||||
savepath = "./saved-outputs/" + args.ticker + "/"
|
||||
if not os.path.exists(savepath):
|
||||
os.mkdir(savepath)
|
||||
|
||||
if args.model.lower() == 'lstm':
|
||||
modelname = "lstm"
|
||||
else:
|
||||
if args.model.lower() == 'gp':
|
||||
modelname = "gp_" + args.kernel + "_"
|
||||
elif args.model.lower() == 'volt':
|
||||
modelname = "volt_"
|
||||
|
||||
if args.mean.lower() == 'constant':
|
||||
modelname += 'constant' + "_"
|
||||
elif args.mean.lower() in ['ewma', 'dewma', 'tewma']:
|
||||
modelname += args.mean + args.k + "_"
|
||||
|
||||
|
||||
###############
|
||||
## Main Loop ##
|
||||
###############
|
||||
|
||||
for last_day in test_idxs:
|
||||
date = str(data.index[last_day.item()].date())
|
||||
train_y = data.Close[last_day.item()-ntrain:last_day.item()].to_numpy()
|
||||
train_y = torch.FloatTensor(train_y).to(train_x.device)
|
||||
|
||||
if args.model.lower() == 'lstm':
|
||||
model = LSTM(train_x, train_y, 10, 128, 1)
|
||||
model.Train(args.train_iters)
|
||||
elif args.model.lower() == 'gp':
|
||||
model = BasicGP(train_x, train_y, kernel=args.kernel,
|
||||
mean=args.mean, k=args.k)
|
||||
model.Train(args.train_iters)
|
||||
elif args.model.lower() == 'volt':
|
||||
model = Volt(train_x, train_y, mean=args.mean, k=args.k)
|
||||
model.Train(gpcv_iters=args.train_iters,
|
||||
vol_mod_iters=args.train_iters,
|
||||
data_mod_iters=args.train_iters)
|
||||
else:
|
||||
print("ERROR: Model not found")
|
||||
|
||||
|
||||
samples = model.Forecast(test_x).squeeze()
|
||||
torch.save(samples, savepath + modelname + date + ".pt")
|
||||
torch.cuda.empty_cache()
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--ticker",
|
||||
type=str,
|
||||
default='F',
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ntrain",
|
||||
type=int,
|
||||
default=400,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--n_test_times",
|
||||
type=int,
|
||||
default=25,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--forecast_horizon",
|
||||
type=int,
|
||||
default=100,
|
||||
)
|
||||
parser.add_argument(
|
||||
'--kernel',
|
||||
type=str,
|
||||
default="matern",
|
||||
)
|
||||
parser.add_argument(
|
||||
'--model',
|
||||
type=str,
|
||||
default="volt",
|
||||
)
|
||||
parser.add_argument(
|
||||
'--mean',
|
||||
type=str,
|
||||
default="ewma",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--nsample",
|
||||
type=int,
|
||||
default=1000,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--train_iters",
|
||||
type=int,
|
||||
default=500,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--end_date",
|
||||
default=None,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--lookback",
|
||||
type=int,
|
||||
default=500,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--k",
|
||||
type=int,
|
||||
default=100,
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
main(args)
|
||||
@@ -0,0 +1,131 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "0f0afec5-bfbe-437e-8c42-5702d8c6be91",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Warning no robinhood utils.\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import torch\n",
|
||||
"import pandas as pd\n",
|
||||
"import gpytorch\n",
|
||||
"import argparse\n",
|
||||
"import datetime\n",
|
||||
"import warnings\n",
|
||||
"import os\n",
|
||||
"from voltron.data import make_ticker_list, GetStockHistory\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"id": "4fb8850d-b17f-41a8-8285-d4e74e6d7b16",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ntrain = 400\n",
|
||||
"lookback = 2\n",
|
||||
"tckr = \"TSLA\"\n",
|
||||
"data = GetStockHistory(tckr, history= ntrain + lookback)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 21,
|
||||
"id": "ef536909-31a6-46ef-99b2-c9c956e3575a",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"pxs = data.Close[:-1].to_numpy()\n",
|
||||
"preds = torch.load(\"./saved-outputs/TSLA/lstm_2021-12-28.pt\")\n",
|
||||
"\n",
|
||||
"trx = np.arange(pxs.shape[0])\n",
|
||||
"tex = np.arange(pxs.shape[0], pxs.shape[0] + preds.shape[-1])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 25,
|
||||
"id": "c779bb74-3b07-4c78-9924-911efae14177",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"[<matplotlib.lines.Line2D at 0x7fcb5e8153d0>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7fcb5e8154f0>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7fcb5e815610>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7fcb5e815730>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7fcb5e815850>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7fcb5e815970>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7fcb5e815a90>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7fcb5e815bb0>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7fcb5e815cd0>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7fcb5e815df0>]"
|
||||
]
|
||||
},
|
||||
"execution_count": 25,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAX0AAAD4CAYAAAAAczaOAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/MnkTPAAAACXBIWXMAAAsTAAALEwEAmpwYAAA6KElEQVR4nO3deXzcdZ348dd7JpM5cidN2jTp3dITKKWU+0ZBcC0uIgURdkWrLL9V193lUFdXkfVcUFZltyKCiLBdEUGQo1SOFiilpdD7PtPmvo+55/P7Y74zmSSTHkkmM03ez8ejj8x8vt/MfPKFvOeT9+fzfX/EGINSSqnRwZbuDiillBo+GvSVUmoU0aCvlFKjiAZ9pZQaRTToK6XUKJKV7g4cy5gxY8zkyZPT3Q2llDqprF+/vsEYU9q7PeOD/uTJk1m3bl26u6GUUicVETmQrF3TO0opNYpo0FdKqVFEg75SSo0iGvSVUmoU0aCvlFKjiAZ9pZQaRTToK6XUKKJBXymlRhEN+kopNQyamprIhP1LNOgrpdQw6Orqor6+Pt3d0KCvlFLDwe/34/f7aW9vT2s/NOgrpdQwCIVC+Hw+qqqq0prm0aCvlFLDpL29nY6ODrq6utLWBw36Sik1DGw2Gw0NDUQiEbxeb9r6kfGllZVSaiRwOBy0t7fjcDgQkbT1Q0f6Sik1DLKysvD7/RhjNKevlFIjnd1uJxKJYIyhpaUlbf3QoK+UUsOspaWFmpqatLy3Bn2llBoGoVAIiI74g8EgTU1NaVmzr0FfKaWGQSAQwGazYbPZCAQCBINBgsHgsPdDg75SSg2DQCAAQDgcpq2tjezsbCKRyLD3Q4O+UkoNg0AggIjEJ3QdDkda+qFBXymlhkE4HCYYDOLz+TDGpCW1A8cR9EXkERGpE5HNSY79i4gYERmT0HaPiOwWkR0icmVC+5kissk69qCk8+4EpZQaRrG1+cYYwuEw4XA4nu4Zbscz0n8UuKp3o4hMAD4CHExomwMsAeZa3/NLEbFbhx8ClgIzrH99XlMppUaiWJG12Ai/vLwcv9+fmTl9Y8ybQFOSQw8AdwKJt5YtBp4yxviNMfuA3cAiESkH8o0x75jorWi/Ba4dbOeVUirThUIhAoEAxpj46p2cnBy6urrSUoNnQDl9EfkEcNgY82GvQxXAoYTnVVZbhfW4d3t/r79URNaJyLpM2HRAKaUGqrOzE4/HQyQSQUTweDw4HA58Pl9mjvR7ExEP8A3gW8kOJ2kzR2lPyhizzBiz0BizsLS09ES7qJRSGSMYDGK32wkEAvFia1lZWWRlZaUl6A+kyuY0YArwoTUXWwm8LyKLiI7gJyScWwkcsdork7QrpdSIJiI0NDQQCATiyzVjSzfTUXjthEf6xphNxpgyY8xkY8xkogF9gTGmBngOWCIiThGZQnTCdq0xphpoF5FzrFU7twDPDt2PoZRSmcvr9RIOh3E4HITDYUQko5dsPgm8A8wUkSoRua2/c40xW4DlwFbgJeAOY0zYOnw78DDRyd09wIuD7LtSSmU8t9uNMYZIJILT6SQUCtHa2pq2oH/M9I4x5sZjHJ/c6/l9wH1JzlsHzDvB/iml1EnN4/EQDAbjZZUjkUi8+Fpvxhj8fj8ulytl/dE7cpVSKsUCgQChUIhQKITb7cbr9eJ0OvucFwqFaGtrS2lfNOgrpVQKxSZrs7OzERGKiopobW0lLy+vz7l+v59wONynfShp0FdKqRQKh8P4fD46Ozux2WxkZWURCoWw2fqG39iyzlTSoK+UUikUW6Lp8/nIysrCbrdjt9uTjuhj5zY0NKSsPxr0lVIqhWI3YAWDQbKzs7Hb7WRlZcVLMySTyj10NegrpVQKhcNhjDGICDabDY/Hg81mi6/UaWrqW9oslXl9DfpKKZVCsVLKsTtw29rasNlshMNhjhw50mO1znBUnB9IGQallFLHKZbe8Xg8hEIhurq6sNvt8bt0/X5//NzhKMugQV8ppVIoVnYhOzubUCjUY9Qf09raGl+1EwwGk67hHyoa9JVSKoVi+XyIfgBEIpH46h0RiW+sEivL4PV6ycnJSVl/NKevlFIpFAv4IhJfkgnRD4POzs74pK3X68UYQ3NzM1lZWSlL9ehIXymlhkEs6MduyiooKCAQCBAMBhER6urqcLlcPf4SyMoa+hCtI32llBoGsQ1UIDrK93g8uN1u/H4/xhhKSkqoqanBGIPX6+23KNtgadBXSqkUiqVp3G43NpstPuJ3uVy4XK748dzcXDo7O+NBP1WllzXoK6XUMLDZbPFVO4FAIL5ePzbRG5vUNcZQV1enI32llDoZJa7cyc7OJhKJEAgEcLlciEj85i2IFlyz2Wx4vV46OjpS0h8N+kopNQxiKZ1IJEJOTg42m41AIEAgEODIkeiW4X6/Pz7iT7xpayhp0FdKqWFgjInX0C8oKEBE8Pv98RU8sRG/zWYjFAqlbMmmBn2l1KhT2+bj35/bQrtvePepje2RW1ZWFl+nH6uh39XVhc/ni9fe14lcpZQaIo+8tY9H397P957fNmzvabfbsdls8fROLKcfu2GrpaUFt9tNc3MzXq83ZZupHDPoi8gjIlInIpsT2n4sIttFZKOIPCMihQnH7hGR3SKyQ0SuTGg/U0Q2WccelOEoJ6eUUkms2hndpOTlrTUpfy9jDKFQiOzsbGw2G263G4fDgcfjITc3F7fbTU5ODsFgELfbTSAQICcnB5/Pl5L+HM9I/1Hgql5tK4B5xpjTgJ3APQAiMgdYAsy1vueXImK3vuchYCkww/rX+zWVUiqlDjZ28fCqvWytbiPbbqOlK0gwHEnpe7a3txOJROL5/NidtrFlmrERvdfrxeVyEQqFyMvLS1lN/WMGfWPMm0BTr7ZXjDGxRaRrgErr8WLgKWOM3xizD9gNLBKRciDfGPOOic5O/Ba4doh+BqWUOi5/95u1fO+FaErnrClFADR3BVL6ns3NzUQikXjlzJKSkniqx+l0kp2dHS+6lp2dTWtrK7m5uSkrujYUOf3PAS9ajyuAQwnHqqy2Cutx7/akRGSpiKwTkXX19fVD0EWllIJ2f/cNT6dWFALQ2JHaoO92uwkGg/ERvdvtpr29HbvdHv8giEQihMNhHA4HbrebSCSSkro7MMigLyLfAELAE7GmJKeZo7QnZYxZZoxZaIxZWFpaOpguKqVUXElOdvzxaZUFADR1pjboFxQU9BjpJ34IxLZODAaD2Gy2eKDPysqKb74y1Ab8USIitwIfBy433QtKq4AJCadVAkes9sok7UopNWxKcruD/oyyXAAaOlJzE1RMKBTCZrPFR/qxssmxfH44HMblcpGd3d23xBLMQ21AI30RuQq4C/iEMaYr4dBzwBIRcYrIFKITtmuNMdVAu4icY63auQV4dpB9V0qpARuTGx15pzq9U19f32PCFqJbJ/a+49YYg8/ni4/y8/PzU9Kf41my+STwDjBTRKpE5Dbg50AesEJEPhCR/7Y6vQVYDmwFXgLuMMbEpqBvBx4mOrm7h+55AKWUGhYd/mg4KnA7KHA7sNsk5emd2G5ZLpcrvk7f4/Hg9Xqx2+3Y7XZaWloIBoP4fL741oqpGukfM71jjLkxSfOvj3L+fcB9SdrXAfNOqHdKKTWEOnxBLpwxhoduPhObTSjyZNPYmTy9E4kYRBh08E3cECUW9GP1dUQEp9OJz+ejsbGRkpISOjs7yc7OpqSkZFDv2x+9I1cpNWp0+sOUF7jIdUbHu7lOO53+5Ovhlz6+ngX3rhj0e4pIvI5+LOjHOJ1OAoFAvOxCrNRybHetVNCgr5QaNTr8IXKd3bl1l8OOL5g86L+6rZbmrmC88NmWI610+E+8xr3D4YgH8VjZhRin0xlfow/R0sqxlE8678hVSqmTXiRirKBvj7e5HHZ8ob5LI6uau9entHlDvH+wmWseXM0DK3YO6L1jk7ixVE/ssd0e7UtTUxMOh4NQKBT/qpuoKKXUIHRZI/pcV/dUpsthwxfoO9Jff6A5/rimzccv/ro7/vhEJd6YlTjSj6V8srOzcbvd8aWcEK2rn5ube8LvdTw06CulRoUOX3TknONMDPp2fKG+Qf9QU/dIv7bNR1WzF4CuAaR3jDHxAJ6Y04+N+mOllWPbKaa6FqUGfaXUqBDLx+cmBv2s5Dn9WJAH2FHTHr+Bq7btxG/kipVQBuITtRAN+rHNUmJpndgduanaQAU06CulRomkQd9hwxfsm9M/3OJl1rhoVcz7/rKNRmstf137iad3YuUXej93OBwEg0FEhK6uLlwuFy6XK57nTxUN+kqpUaGmNRqwy/Jc8bb+Vu9UNXuZVtYzp16Sk01jZ+CESzHHRveRSASbzUZBQbTmT6y0clFRUTytM27cOLKyslI2iQsa9JVSo0QsTz+x2BNvcznseHsF/UjEcLjFS2Whm5e/elG8fc74fIw5sVo9/lA4PkEbW5aZWF7BGEN+fj6hUAiXy8W0adMoKSnRoK+UUoN1qLmLPFcWBZ6e6/T9vdI7r++sIxCKMLU0h5lWigfg1IroCH1nbcdxvZ8xhnO//1d+t+YANputz6YosZr67e3tOJ1OnE4ndrs9/pdA2jZRUUqpkeBgU1ePUT5Ec/qBcIRwpHvi9IEVu5helsu1Z0S3/LDbopOwn5g/njG52fx69b7jer/qVh9NnQGKPbb4aD6R0+mMl1VOnMQNBoPk5ubqzVlKKTUYyYN+dNLUby3bjEQMu+s6uGhGKc6s6LG/tYL/hCIPi+dXsGZP43GtrtlW3QZAucfG2LFj+4zcY2vzbbboh0Ji6WWPx6NBXymlBioSMVQ1e5nQO+hnRUPg/oZovr+23Yc3GGZKafdWhfd98lRW3XkpOc4sSvOcBMIRvMEwdW0+9tb3n+qJBf2p5cU4HI4+m6KICOPHj8fhcMRH/RDdTjGVa/U16CulRry6dj+BUKRv0LdG+lc/uIon1x5kX0MnAFNKuoNLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {
|
||||
"needs_background": "light"
|
||||
},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(trx, pxs)\n",
|
||||
"plt.plot(tex, preds[:10, :].T.exp(), color='gray', alpha=0.75, lw=0.2)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "eb8f6699-b9ff-45cc-8901-e2becfcc01b0",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.12"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
Binary file not shown.
@@ -0,0 +1,408 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "9754a692-edb8-41bf-8159-fd82edcb9195",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Warning no robinhood utils.\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import pandas as pd\n",
|
||||
"import torch\n",
|
||||
"from torch import nn\n",
|
||||
"import seaborn as sns\n",
|
||||
"import time\n",
|
||||
"import copy\n",
|
||||
"import sys\n",
|
||||
"from torch.utils.data import DataLoader\n",
|
||||
"from voltron.data import make_ticker_list, GetStockHistory\n",
|
||||
"sns.set_style('white')\n",
|
||||
"# style.use('whitegrid')\n",
|
||||
"palette = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"sns.set(palette = palette, font_scale=2.0, style=\"white\", rc={\"lines.linewidth\": 2.0})"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "dbc24c1a-f499-4b8f-8544-6d4a66f48488",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Dataset Setup"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "62f91be1-38aa-4781-8b65-c5caf9752045",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from torch.utils.data import Dataset\n",
|
||||
"\n",
|
||||
"class SequenceDataset(Dataset):\n",
|
||||
" def __init__(self, data, sequence_length=5):\n",
|
||||
" self.sequence_length = sequence_length\n",
|
||||
" self.X = data.float()\n",
|
||||
"\n",
|
||||
" def __len__(self):\n",
|
||||
" return self.X.shape[0]-1\n",
|
||||
"\n",
|
||||
" def __getitem__(self, i): \n",
|
||||
" if i >= self.sequence_length - 1:\n",
|
||||
" i_start = i - self.sequence_length + 1\n",
|
||||
" x = self.X[i_start:(i + 1)]\n",
|
||||
" else:\n",
|
||||
" padding = self.X[0].repeat(self.sequence_length - i - 1, 1).squeeze(-1)\n",
|
||||
" x = self.X[0:(i + 1)]\n",
|
||||
" x = torch.cat((padding, x), 0)\n",
|
||||
" \n",
|
||||
" return x.unsqueeze(0), self.X[i+1]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 239,
|
||||
"id": "566c9386-c297-430f-b4ef-a0f3ba7aec37",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"tckr = \"JPM\"\n",
|
||||
"ntrain = 400\n",
|
||||
"lookback = 1\n",
|
||||
"data = GetStockHistory(tckr, end_date=\"2021-12-07\", history=ntrain + lookback).Close.to_numpy()\n",
|
||||
"data = torch.FloatTensor(data).log()\n",
|
||||
"\n",
|
||||
"data = (data - data.mean())/data.std()\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"# xin = torch.linspace(0, 5*np.pi, 250)\n",
|
||||
"# data = torch.sin(xin) + 0.2 * torch.randn(xin.shape)\n",
|
||||
"\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 269,
|
||||
"id": "5dece53d-785f-4ca9-9cbf-03e16de71e9b",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"seq_len = 25\n",
|
||||
"dset = SequenceDataset(data, seq_len)\n",
|
||||
"\n",
|
||||
"trgts = []\n",
|
||||
"for i in range(len(dset)):\n",
|
||||
" trgts.append(dset[i][1].item())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 270,
|
||||
"id": "2317ba11-4f62-443f-9ad8-63f9a0996ee9",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"train_loader = DataLoader(dset, batch_size=20, shuffle=True)\n",
|
||||
"X, y = next(iter(train_loader))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 271,
|
||||
"id": "ffecfde0-6938-4dee-b1b0-6627012e3329",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"class ShallowRegressionLSTM(nn.Module):\n",
|
||||
" def __init__(self, input_size, hidden_units=128):\n",
|
||||
" super().__init__()\n",
|
||||
" self.input_size = input_size # this is the number of features\n",
|
||||
" self.hidden_units = hidden_units\n",
|
||||
" self.num_layers = 5\n",
|
||||
"\n",
|
||||
" self.lstm = nn.LSTM(\n",
|
||||
" input_size=input_size,\n",
|
||||
" hidden_size=hidden_units,\n",
|
||||
" batch_first=True,\n",
|
||||
" num_layers=self.num_layers\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
" self.linear = nn.Linear(in_features=self.hidden_units, out_features=2)\n",
|
||||
" \n",
|
||||
"\n",
|
||||
" def forward(self, x):\n",
|
||||
" batch_size = x.shape[0]\n",
|
||||
" h0 = torch.zeros(self.num_layers, batch_size, self.hidden_units).requires_grad_()\n",
|
||||
" c0 = torch.zeros(self.num_layers, batch_size, self.hidden_units).requires_grad_()\n",
|
||||
"\n",
|
||||
" _, (hn, _) = self.lstm(x, (h0, c0))\n",
|
||||
" out = self.linear(hn[0]) # First dim of Hn is num_layers, which is set to 1 above.\n",
|
||||
" output = torch.zeros_like(out)\n",
|
||||
" output[:, 0] = out[:, 0]\n",
|
||||
" output[:, 1] = torch.exp(out[:, 1])\n",
|
||||
" return output\n",
|
||||
" \n",
|
||||
"from torch.autograd import Variable \n",
|
||||
"class LSTM1(nn.Module):\n",
|
||||
" def __init__(self, num_classes, seq_len, hidden_size, num_layers):\n",
|
||||
" super(LSTM1, self).__init__()\n",
|
||||
" self.num_classes = num_classes #number of classes\n",
|
||||
" self.num_layers = num_layers #number of layers\n",
|
||||
" self.input_size = seq_len #input size\n",
|
||||
" self.hidden_size = hidden_size #hidden state\n",
|
||||
"\n",
|
||||
" self.lstm = nn.LSTM(input_size=seq_len, hidden_size=hidden_size,\n",
|
||||
" num_layers=num_layers, batch_first=True) #lstm\n",
|
||||
" self.fc_1 = nn.Linear(hidden_size, 128) #fully connected 1\n",
|
||||
" self.fc = nn.Linear(128, num_classes) #fully connected last layer\n",
|
||||
"\n",
|
||||
" self.relu = nn.ReLU()\n",
|
||||
" self.softplus = nn.Softplus()\n",
|
||||
" \n",
|
||||
" def forward(self,x):\n",
|
||||
" h_0 = Variable(torch.zeros(self.num_layers, x.size(0), self.hidden_size)) #hidden state\n",
|
||||
" c_0 = Variable(torch.zeros(self.num_layers, x.size(0), self.hidden_size)) #internal state\n",
|
||||
" # Propagate input through LSTM\n",
|
||||
" output, (hn, cn) = self.lstm(x, (h_0, c_0)) #lstm with input, hidden, and internal state\n",
|
||||
"\n",
|
||||
" hn = hn[self.num_layers-1]\n",
|
||||
" hn = hn.view(-1, self.hidden_size) #reshaping the data for Dense layer next\n",
|
||||
" out = self.relu(hn)\n",
|
||||
" out = self.fc_1(out) #first Dense\n",
|
||||
" out = self.relu(out) #relu\n",
|
||||
" out = self.fc(out) #Final Output\n",
|
||||
" \n",
|
||||
" output = torch.zeros_like(out)\n",
|
||||
" output[:, 0] = out[:, 0]\n",
|
||||
" output[:, 1] = self.softplus(out[:, 1])\n",
|
||||
" return output"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 272,
|
||||
"id": "3bc83b15-3b0f-4f21-803b-29e354cb558f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# model = ShallowRegressionLSTM(seq_len)\n",
|
||||
"model = LSTM1(2, seq_len, 128, 1)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 273,
|
||||
"id": "41547213-7851-4ad1-ac04-03cab5e0a91c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def NLL(targets, outputs):\n",
|
||||
" dist = torch.distributions.Normal(outputs[:, 0], outputs[:, 1])\n",
|
||||
" return -dist.log_prob(targets).sum()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 274,
|
||||
"id": "e6958f51-5cd0-414a-b419-3ce55f12b486",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def train_model(data_loader, model, loss_function, optimizer, epochs=200):\n",
|
||||
" num_batches = len(data_loader)\n",
|
||||
" total_loss = 0\n",
|
||||
" model.train()\n",
|
||||
" for epoch in range(epochs):\n",
|
||||
" for X, y in data_loader:\n",
|
||||
" output = model(X)\n",
|
||||
" loss = loss_function(y, output)\n",
|
||||
"\n",
|
||||
" optimizer.zero_grad()\n",
|
||||
" loss.backward()\n",
|
||||
" optimizer.step()\n",
|
||||
"\n",
|
||||
" total_loss += loss.item()\n",
|
||||
"\n",
|
||||
" if epoch%10 == 0:\n",
|
||||
" avg_loss = total_loss / num_batches\n",
|
||||
" print(f\"Train loss: {avg_loss}, Epoch: {epoch}\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 275,
|
||||
"id": "92484e45-11c9-4fe1-bc5b-87b0f27d5a2c",
|
||||
"metadata": {
|
||||
"tags": []
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Train loss: 14.94985544681549, Epoch: 0\n",
|
||||
"Train loss: 2.123787060379982, Epoch: 10\n",
|
||||
"Train loss: -118.20948788821697, Epoch: 20\n",
|
||||
"Train loss: -225.282495072484, Epoch: 30\n",
|
||||
"Train loss: -393.8007117182016, Epoch: 40\n",
|
||||
"Train loss: -546.9331771463155, Epoch: 50\n",
|
||||
"Train loss: -724.232698109746, Epoch: 60\n",
|
||||
"Train loss: -916.3166337996721, Epoch: 70\n",
|
||||
"Train loss: -1073.7028237611055, Epoch: 80\n",
|
||||
"Train loss: -1267.8471668988466, Epoch: 90\n",
|
||||
"Train loss: -1434.5194520920516, Epoch: 100\n",
|
||||
"Train loss: -1630.7973297566175, Epoch: 110\n",
|
||||
"Train loss: -1831.4267205685378, Epoch: 120\n",
|
||||
"Train loss: -2038.7122355431318, Epoch: 130\n",
|
||||
"Train loss: -2246.2204612165688, Epoch: 140\n",
|
||||
"Train loss: -2457.420981016755, Epoch: 150\n",
|
||||
"Train loss: -2669.7478996187447, Epoch: 160\n",
|
||||
"Train loss: -2884.8198515325785, Epoch: 170\n",
|
||||
"Train loss: -3108.3130769401787, Epoch: 180\n",
|
||||
"Train loss: -3337.9945754677055, Epoch: 190\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"optimizer = torch.optim.Adam(model.parameters(), lr=0.01)\n",
|
||||
"train_model(train_loader, model, NLL, optimizer)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 276,
|
||||
"id": "07eab392-37e2-4dae-bfa4-dd177e3d552f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"means = []\n",
|
||||
"vrs = []\n",
|
||||
"for X, y in dset:\n",
|
||||
" output = model(X.unsqueeze(0))\n",
|
||||
" means = means + list(output[:, 0].detach().numpy())\n",
|
||||
" vrs = vrs + list(output[:, 1].detach().numpy())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 277,
|
||||
"id": "b10866df-b554-44a6-bf73-b98abb2e6400",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"<matplotlib.collections.PathCollection at 0x7fa59d357370>"
|
||||
]
|
||||
},
|
||||
"execution_count": 277,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYIAAAEFCAYAAADuT+DpAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/MnkTPAAAACXBIWXMAAAsTAAALEwEAmpwYAABsMElEQVR4nO29eXxdV3nv/Vt7PJOOZsmTPMay4zF2BpyQYAIJIW1DKVMCobS0lFvgbfsylHuhvL2XhLZMbSGhgUJoS5sAKSUlDcRJEy44TuIMTpzYluMhii3Lg2zpSEdn3tNa7x9r731m6WiWddb388kn9jn77LP3tvQ8az3D7yGMMQaBQCAQ1C3SXF+AQCAQCOYW4QgEAoGgzhGOQCAQCOoc4QgEAoGgzhGOQCAQCOocZa4voBZyuRwOHTqE9vZ2yLI815cjEAgEFwWO42BwcBCbNm1CIBCoetxF4QgOHTqE22+/fa4vQyAQCC5K7r//flxxxRVV378oHEF7ezsAfjOLFi2a46sRCASCi4OBgQHcfvvtvg2txkXhCLxw0KJFi7Bs2bI5vhqBQCC4uBgvpC6SxQKBQFDnCEcgEAgEdY5wBAKBQFDnCEcgEAgEdY5wBAKBQFDnCEcgEAgEdY5wBAKBQDAN2JTBoRfneBfhCAQCgWCKMMZwLm6gL5ZFzqJzfTkT5qJoKBMIBIK5xqEMhAASIUWvmzbFuVEDpsOgEIKBhIFFUR2xtImwJsNhQFNQgSyRKmeee4QjEAgEgnEwLIrT8RyCqoQlTcXibYmcDcth0GUCQggMm+JMPAfGgIzhAAAsh6KzQQMh89MZCEcgEAiqwhgDZZjXq9mZZiRtIZa2IBEgY1I4lPnPgzKGRNaG6joBAFAkAocx6AqPvDPGkMw6oNRERJegKTIypoOILhd9zjtf6Y5jNhCOQCAQVIQxhtGsjZGMjWXNOlS5/lKKqZyNWNqCKhNI7mo/Z1GEda7dkzEdOAxQC4y3LBHIyP+dEAJdAXKWg7TpALAAAMNpC5GA7O8UGGPoH86hOaQgGlRn9T7r719WIBCMSzJn42Qsi8GkBZsyDCYtMHZxVsRMBMYYchaFTRkMm2IgYUKRiL9KlwgwmrUBAImshXOj/P3xIIRAlSUEFAm6TBBQJGgyQTLrIOMml23KYNoMQ6nZf9ZiRyAQ1CHVQj6eIRxMmpAIQUDlK9W06cCwGQLqwg4R5WyK08M5gBDoCr/XwmekSAQZ04FpU4zmHCgSqckRFOKFggghUGRgKGki2BKAafNkNGWAQwFlFmdwiR2BQFCHxLM2To/k/JWn7TBQxmA5DGfiBoC8ASSEuCtha86ud7ZI5RwQiUCTCbIWLTPynhEfzdowLAp5in5RkQgsh2E4bSFnOf7rpjO7JahiRyAQ1BkO5YaHMsB0GCTCcGo4B5kQhHS+NizNB6gSQTLnoC3C5n3iOGdR5CwHTaGJxdkdypDI2VDdUFBIrbwkVyTih4emowpIkwniWRuSe24eIqIIabO3JRA7AoGgDqCM+av/nEXBGEAAZE0HGcMBYwBjwGjWqRjq8AzecNqCYc/vhqlY2sRwmsfZHcqQNZ3xPwSeGGasvE+gFIkAINyATweEEEgAGPguTJG4Y5jNPIFwBALBAocyHu5JuKvYtOGAgK8+h1IWhtIWN0BuErPail+VCUYyNoZScxciylkUqZxd9f20YSNjUlDG4/2xtIUzcQNWDaGWjEVRy2aHEAJdlqa1J0CVJWjuLkyWCGyHIZmzkTUdjKRn/nkLRyAQLHDShoOs6WAobYG61TCyRCBLBKpEIJPaEp4SIQgoBFnTgT1HmjqjWQtnR01krfJVPmO8ukmVCCQJOJ8wMZq1QQgwMGrCsOmYWkCWMzc1/JVQJYILKd6/MJKZ+Soi4QgEggVOzqKQCQFjgGFTmHZ+5Su5DqFWvFWwMQd6OowxpA0HEgFG0naZUbcpg00ZJAKokgSHMmgygSZLMB2K/uEczsRzuJAwK+oB2U5tO4LZQJJ4uChtOHx3M8PPWzgCgWCBk7Oob+xThgNGpp7kzFVYkc80psNAwWPzGdPB6XiuaKVs2vzP3r1psuSv8DWZ1+2bNsNIxkLKsGE5FBnT4aW0lJfTzidUWUJQ5SbameEdgagaEggWMIzxChRV5iv/pOHwrOQUkCVeWjmbMMaQM/m1E7fGP2dTZK18dU3adDCWeyOEQJYA4jaFjWZtUAosbdahyvlS2fkEv56Z91BiRyAQzBEOZRMKsUwmTmxTBgbXCBJ+DmmKv/Uy4SEmOs715CwK28kfQxmrGNuvhVPDOQxnbMiFUg4EGM3Y/rmTORvKOJU8isRDRQSAQghAuHroXOU85gvCEQgEc0TasHEhadZ0rE15rX9igk1dhaWThHAjqE7RExA333Bu1Kh6DGUM5+I59I/kfIeRMSnOxs0JG13mNrrxJHf+dUUiSJsOLIci61YK1ZrsVWWJ50dcp+Y4whEIBII5IGNSZC2naNVcjVjKhOlQDE5QhyaZc2YkAarJxDW+5dfCXEVOh/Fdjxe7T+Z4XD6entg9eH4jrMlFoRvvzyMZG/GMNakuX09ILmPRMcNKCx3hCASCWcS0Kc7FDeQsBxnTAUCKEq+WU25cLYcimXOgy5Jb+TO2EXUor66xHFpRJmE68IywVeLETJvi9IiBwRQv4wQAw3L8ip+gKmEky3dCtTqD0u8oRJN5l2/anNx9SgSwbIaM4cz7jumZRCSLBYJZwnIohtMWEjkbhkPhMB7nzlgUkQBfSZ+NGwjrsp8AHU5bUGQCgnziMGM6UGWCrOWUrZIBLo18dtRAUJX8/MDM3ROD7loRyhjOjhpwaH5Ii8QYEjkHulv9IhECXeY7lcZgbSJ2Nq2eRyFubwNjbFL3SQgBA4PpUASUya2LYz29OLt7H8xEGlo0jCU7r0DrxjWTOtdcIRyBQDAL2A7XmrcpQ1CVkLOZL1HgxfENm4dRDJvPAAAAMJ7s9YyUIhHecWo5SBsOljUHyjRpUoYDVeIhj6nmA8bDsPiAFQBIZvmkrkKDqkgEOdvBYCr/Gc+h5SwHAVUCZcx3dN4MBMNm6GhQQQgPQY1n4qfi7DSZgIFM+Byxnl70P74XTs7EsGHjQDyH7ZYD8+HdSJ0+jxU3XTPpa5pthCMQCKYJL6RTmrB0KMOFpAHKuEEnhCDoroQZYzAdromTMmwQCQjIUkHYhICyvKGTCK+nNx3GNWkydpEjcKi7Y5AIVDKzTkAmxG90silDLG1Bq6DWqckSciUxeJkQZAyKxiBD/7ABTSFYFNWQMXkeBAB0hQuw1VINNBUIIVUdTaGxBwA5qKPrhh0AgL5dT4HZ3In/4PVh9KZM/OLMKD6/qRPYfwRD+4/k79f93HzdKQhHIBBMEyNpCwxAW0Tjf89YCKkyUoaNtOFAV8r1abzVcdqw+cjDAulnj0IbSAiBJDFIICCEN3Z5TiOWtmDajAvKzUI9vCxxPZ9E1kIyxztg1QoGWyIEhLCSKV5A1nZgOQyWQ2E6QN9wDpbD/GfgD2ghpOizs0XfY88UGXMAcLIG+h7ZA0lTwWwHhkPx1cMXcN7VP8o6DD/pi2NjUwAxw8HSoIrVEQ3NAE4+vBv9Tzw7Lx2CcAQCwTTAGPObtVrDXOpgKGUhrDnIWRRaBSfgIROCwRSXha7F4BWGeyx3aAzABeEYmz5VzPHwQjnnk3ye71jfq5fIWnufTRncgGoyAaXwcwtAfh7CXExGi/X0+k6AMob7TozgVNrEDYsbsKMtDCfLS2cPxHO+E7iyNYT9wxkciOdwIJ7zzxWQCP5sfTu6whp3JLueAgDfGcyHHINwBALBNGA5zC8DNV3lSAIgbXL9mrHq22UJACWoIn8/LufiBjRF4sZ4kgnPqcCdwOS+N+GWt0qEQKpy/3PR7dv/xLMAeKjtp/1xPB/LAACeOJfEjrawf9yBkSwA7gQ+uLIZo6aDY8ni/oocZfjZ6VH8ybp2AACzHZx8eLf/fmGIyUykyxzFbCAcgUAwRRhjGEqZIMTT9LeRyNnQ5NoSkN7IwskgSzyHkLUoQur0OgFKGR564iVsuGQJ1q1eXPGYWu+xGpbDK4zmmsJVuRzQ/JzAA6fieGYw7R83kLORsSlCigSLMhwe5Sv/31oahSwRvHVRBMeSBq7rCOMtnQ1wGMPXDl/A0YSBr/Scx462ME6kDEiE4C0PPI6uaNB3AknLgUUZWgD0P75XOAKB4GLCtBnSJoUuc6PsSR/PxkpWcaekqNLkyidLKYyL77mQwgN9cYQVCd//83dhxRUbyo6fyneqUnEifK6I9fQWrco9J3BgJIu9rhO4fWUz9g6l8XrKxMm0iQ2NARxL5GBQhmUhFa1uDe3GpiDu2LIITZrs7wJvW9mMH50cQX/GQv+puP+9z8cy+O1ljbhxcQMA4K6jgxjI2vjkpe1Y7V7XbDkD0VAmEEyRZC5v+CXCf6lKq2dmmul0AmcyJr50cAAP9MUBAGmb4t5/fhR9jz0z5e8oRHKH4cw1Z3fv853AiZSB8zkLZzMW/rk3BgbgbYsbcHV7GKvcIoAjozmcyZj49vEYAGBLU7DofC26UhQKvLI1hDu3Lsa17Tyk1BVS0e46jkfPJpC0HGRsinNZGwzAT9znfnb3vhm862LEjkAgmCIpw/GTvIQQaMrcG7dqVEtMeslRxhh+fDKOgZwNXSJYGlLxesrEM4Np/Ma+wwBwUdXH14KZ4Kv+YcPG3x8ZLJKjXhnWcMvSKABga1MQvxxI4cXhDKyCg65sDY37HWFFwm0rm/HOrkYE3HzKPceGcHg0h6cupLGuUfePPZe1QBnzr2s2EI5AIJgCjjsMZbYqdaaCt+LP2BT/8noMy0Oj+K3kk+h79CkwywFlDH996DwGcjY0ieDOrYsRUiR87fB59KUtHIjnoO0/gsiyznlX/jgVtGgYZiKNVxO5spkElzUH/d3WqoiGNl3GkOFgjxsy+pN1bWgP1G5Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"xx = torch.arange(data.shape[0])\n",
|
||||
"\n",
|
||||
"plt.plot(xx[1:], means)\n",
|
||||
"plt.fill_between(xx[1:], means - 2*np.sqrt(vrs), means + 2*np.sqrt(vrs), color=palette[1], alpha=0.5)\n",
|
||||
"plt.scatter(xx, data, color=palette[5])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "6be11181-0698-429c-9dc7-a0f4b6c99584",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Rollouts"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 267,
|
||||
"id": "81bb90d7-e790-4e31-ba87-fe61449eb9ce",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"nroll = 50\n",
|
||||
"roll_len = 100\n",
|
||||
"xin, xout = dset[len(dset)-1]\n",
|
||||
"xx = torch.cat((xin[0, 1:], xout.unsqueeze(0)))\n",
|
||||
"xx = xx.repeat(nroll, 1).unsqueeze(1)\n",
|
||||
"roll_pxs = torch.zeros(nroll, roll_len)\n",
|
||||
"with torch.no_grad():\n",
|
||||
" for idx in range(roll_len):\n",
|
||||
" out = model(xx)\n",
|
||||
" smpl = torch.normal(out[:, 0], out[:, 1])\n",
|
||||
" roll_pxs[:, idx] = smpl\n",
|
||||
" xx = torch.cat((xx[..., 1:], smpl.unsqueeze(-1).unsqueeze(-1)), -1)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 268,
|
||||
"id": "1aec9c78-083b-4161-af02-4ef9e6590858",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAX8AAAEFCAYAAAAL/efAAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/MnkTPAAAACXBIWXMAAAsTAAALEwEAmpwYAABLQElEQVR4nO3de3xcdZ34/9eZM/dMMrk0lyZN06ZtStMrKWBQsCK3roAXEIpU2WVl/e667rq7qLuy7E+BXffL7uoq+FUUdFcEEYEqolCkUEqFtkDT0qalCb0kTdrmfp3J3Of8/pg5JzOZSZq0SdNk3s/Hw4fNzJmZk055n895f96f90fRNE1DCCFERjFN9wkIIYQ49yT4CyFEBpLgL4QQGUiCvxBCZCAJ/kIIkYHM030C4+H3+6mvr6ewsBBVVaf7dIQQYkaIRCJ0dnayYsUK7HZ70nMzIvjX19ezcePG6T4NIYSYkZ544gkuuuiipMdmRPAvLCwEYr9ASUnJNJ+NEELMDG1tbWzcuNGIoYlmRPDXUz0lJSXMmzdvms9GCCFmlnTpcpnwFUKIDCTBXwghMpAEfyGEyEAS/IUQIgNJ8BdCiAwkwV8IITKQBH8hhIjr6uoiEAic9fuEQiE6Ojom4YymjgR/IUTGOHbsGF1dXaM+f/ToUbq7u8/ovTVNIxKJ0NzcTFNTE+++++6ZnuY5IcFfCJEx3nnnHXp7e9M+19LSQiAQwOv1ntF7nzp1ir1799LV1UVjYyPt7e1A7C6gubmZUChEMBgc13t1dXXR399/RucxXhL8hRAzXk9PT8pjoVCIgYEB42ev14vD4cDn8/Haa6/h8/mSjt+zZw99fX1JwX+8wRqgvb2d/v5+hoaG8Hq9hMNhAAYGBnjvvfc4cuQIR44cSXndwMAAx48fByAQCNDS0kJdXR3Nzc3j/uwzIcFfCDHjvfjii3i9Xurr643HmpqaOHz4MBAL7O+88w4XXnghQ0NDNDY20tfXZxzb0dFBR0cHzc3N9Pb20tbWBsAzzzxj/FnTNLxeLy+//LLxmE7TNLZu3UokEqG/v59oNIrb7cbn83H06FF6e3vZu3cvp06d4sSJEwQCATRNY+fOndTX1/P222+jaRoDAwM0NDQwODhopJ9aWlqm5O9sRvT2EUKIsTgcDhobG2loaODkyZMUFxfT39+P2Wzm6NGj7Ny5E6fTic/nIxQKMX/+fLq7uyksLMRsNvPSSy/R3t6OyWTiwIEDdHd3U1paSk5ODg0NDZSUlDAwMEBdXR379u1j//793HDDDSxZsoTt27fT3NyMx+Phvffew2azkZWVRUFBAb29vdTV1REMBunp6aG7u5tIJEJlZSXZ2dns2rULRVGw2+2cPHmSV199lVOnTmG1Wjl58iQXXngh77//PuXl5ZP+dybBXwgxo0UiERRFoaOjg7a2Ntra2mhubmbBggVEo1H27t3LsmXL2LdvHwAej4fCwkK2b99OOBxm9erVtLa2Eg6HMZvN9PT0YDKZOHToEAUFBbjdbt555x1KS0vp7u7GZrOhaRrbt29nyZIlHDp0iI6ODjRNo6+vj9zcXLxeL263m56eHrq6unA4HKiqysDAgHE34na7CQQCWK1W+vv7efPNN2lqaiIajZKVlUVfXx8vvvgiOTk5U/L3JsFfCDEjHTlyhEWLFnHy5EkjXRIOh7HZbPT19dHe3s68efMYGBhg0aJFDA4OAjA4OIjf7wdg9+7dzJ8/38jth8NhTCYT4XCYSCRCT08PHo+HU6dOcd111xk5fZPJxMDAANFolM7OTiO/DxiPNzU10drais1mw+v1YjabCQaDdHZ24nK5UFWVaDSKx+MhEolw/Phx4zGfz0d2djZtbW14PJ4p+fuT4C+EmJG2b9+O3+9n+/btRCIRAoEAiqKgaRqhUIihoSGi0SihUIgdO3ZgMpnw+/1Eo1EikQiqqqIoCi+88AJms9kI4KqqGpU2kUgEn8+H1Wpl7969BINBwuGwEaR37NhBJBIxzkkv9wSwWq10d3djsViIRqOYTMNTrK2trXR0dBAKhYhGo0DsjkRRFCA2OW2327HZbMaFarJJ8BdCzEitra04nU66urqMwKoHcT2Fc+LECbq7uzGZTGiaZlwgAoEAOTk5tLW1UVlZmRTAw+EwmqYBsWDudDoZGhriyJEj2Gw2wuEwoVCI/Px8tm7dahw7ks/nIxqNMjQ0BMSqjwAURSEYDKIoihH4dYnv1dvbi8lkwuVyTd5fWgKp9hFCzDh6sG5qakoKojabzXhOr+AJBoOoqkooFELTNGMk7vF4CIVCNDQ0JAXdkcFcD96KojA0NGQ87/V6iUQixs/6qF3X39+fdFHR/2w2x8bcp1tJHI1GCYfDOJ3OCfzNjJ8EfyHEjLNnzx4cDoeR5tEDamJ+PBwO43A40DTNSJ3oKSH9eSApHTOWxDsCRVGS8vz6e4/HeNYO2Gw248/6YrHJJmkfIcSUaalr5ODmnfj6PDhyXVSvr6W8puqs37e1tRWr1UpPTw92uz1p1K0HYX3CFkgK1ImjcSAl9QIYF5VEiT+nS9mMh8lkGtfrEu8KxntRmfC5TMm7CiEy3t5N29j9yy34+jwMhiJ0dw2w+5db+OOPnjvr9+7u7jZy4ZqmJQXUxJH8mQTodIF/JE3Tzigon8lr9LuaySYjfyHEWRs5wi++oIKmnQcA8IQi3L+/jZAGGypyqeUEv/naD4BYoK34QDVrblw3oc/z+/1GINXLMnUjq2/GS1EUFEUhNzc3bbuIRGc6Gtc0DZvNNma+f+TFJ93m65NBgr8Q4qy01DVS9/SrDAXDDAQjFGmDRuDXNI3nWvsZisSC2ePHenn51CBWk8LqPAdr851o8WMncgEIhUJGHn9kGudsRKNRysvLjeA/njSNnnIaTyooMZef+FpN04zS05FzCfocxWSTtI8Q4qzs3bSNaDjCDxq6uL++naeP9xnP7e/zs6NrKOn4dn+YlqEQvzsxwM+OxoJs084DtNQ1jvszvV5vSg5+NOMdOVssFkwmE2azGavVCsTaRqSbEB45x6CnnkZ+VnZ2dtLPJpMpKZgnvreiKFit1nFPQJ8tGfkLIc7Y3k3biARD7O7xccwbq2J5t9fHLRV5ABzoH16gtDrPwbu9sU6aORYTA6EoTd4gf/N2KzfNdxN9+lUAymuqjDTSUO8gzrzspIniUChEOBzGYrEY761X/IwcNcPwnUFiOsVisSQFYf21iqJw4MAB4/30VJBOVVWjnUTixUe/Q7BYLEmjd5fLRTQaxev1GoE/8e7A6XQaK4/10k6d/hllZWXj/DYmRoK/EOKMtNQ10rTzAJ5QhCeahnvk94eiBCJRrCaFxoFYbvvPKvO5wG0jFI1ySUEWFxU4+dnRHt7uHkIDnjnez54eH18Ov0x30yladjfwu6YeXu/w8LcXhAk++xoQuzC88847QGo6JDFIJ14I9EVaFouFSCSSkibSF1L5/f6k+QO9lFS/I9A/T78AJH6ufiGw2WxG8FdVFZfLhdvt5tChQ0nnZzKZyM3NJT8/n9LSUhobG3E6ncYdjcViMS5GU3UnIGkfIabBG3VH+M+fvow/ODX53KnUUtfIS996jN2/3ALAjq4hQlGNqmwbJfbYeLLDH+aYJ0hnIIzLbOLCfAcus8oXqwq5qCC2aOnm+bncuaiAuY7Ya454gjxwoINXt9QRCYX5/ckBBsNRftDYRSQU5uDmnQBJbZsTA2riqHnJkiVJxyiKQiQSwW63A7ELh8lkwuFwUFZWht/vNwJ1NBo1WkAEAgFsNhuFhYVALLjPmTMHGE4nqaqK0+nEZDJhMpmwWq2YzWZyc3Nxu91cdtllZGdno2kaeXl5WCwWFixYgNPpZN68eRQUFGC1Wlm5ciWqqnLBBRcwf/58nE4niqIk7UkwmST4C3EOtdQ18vz/9yif++pP+d7Pt/KFP//OhHLd062lrpG9z76Gr8+Dpmm0+0P8/kSsD85Hil0UO2KpmCZvkP+J5/MvyneipsnJO80m1uQ7+KflxXyxag4Wk8IJX4jftibvYNUbjI2yfX2xBVyjBcPENMzChQuNEbPX6zVG0CUlJcYFw2Kx4Ha7ufbaa1FVlTvvvJPs7GyjFYTJZKKwsJCCggIKCwux2+2sWbMGi8WCxWIx8vOKopCVlUVubi6hUIi5c+dy2WWXsXz5cjRNw+fzsXr1agAWLFiA1WolOzub7OxsAoEAbW1trFy5krKyMvLz87nmmmsIBAKUlpYaF62pIMFfiHPkjz96jv99+Hn+5a1mgtFYoPrjqQHe/MUfZswF4ODmnURCYbacGuRv3znB/fvbCWtQ7baxMtdOpSs2Ufp86wC9wQgWBa4tzR7zPVVFodpt595VJQAcHgzS4U++IwpFNSNoJzY60wO+1WpFVVVj9L1q1SojMANGPr6srIzy8nJsNhsulwu73U5JSQl5eXnYbDacTic2m41oNEplZSWBQICKigrC4TD5+fnk5ubidDqZP3++cRHQP6OiooKioiLmz5+PzWajsrISk8nE888/j8ViMd4/NzeXwcFB5s2bx7Fjx1AUheLiYsrKyoy7A6/XS21tLXl5eRQVFU3CN5fqnAX/TZs2sXTpUiNfJ8RsFo1GGfQOB6m9m7bRcbiVX7f00xUYHskFoxr7e/3seXbrdJzmhOmj7z29sVy9bllObJVtTZ4DgKFIbFLz4oIssi3jq7bJsahcnB9LCT3fmjy6HwhFjNYMenoncQSvV+hkZ2eTl5eHyWQyKm304/Sfr732WqLRKEuWLMFisaCqqpEmcrvd2Gw2LBYLH/vYx/D5fNTU1BAOh1m4cCENDQ2YzWbcbrcx4Wwymejv7zc2cTGbzdhsNvLy8jh8+DCqqnLw4EFWrFjBkiVLqKioIDs7mzVr1hAKhbj11ltxuVz4fD6WLVsGwBVXXMG8efNYvHgxNTU14/x2JuacBP89e/Zw//33n4uPEmLatdQLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"full_x = torch.arange(data.shape[0])\n",
|
||||
"test_x = torch.arange(data.shape[0], data.shape[0] + roll_len)\n",
|
||||
"plt.plot(full_x[1:], means)\n",
|
||||
"plt.scatter(full_x, data, color=palette[5])\n",
|
||||
"plt.plot(test_x, roll_pxs[:20, :].T.detach(), color='gray', alpha=1., lw=0.5)\n",
|
||||
"plt.show()"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.12"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,454 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 7,
|
||||
"id": "92e5b691-e01a-438a-b2b9-543c5fa3f8b2",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import pickle as pkl\n",
|
||||
"import pandas as pd\n",
|
||||
"import numpy as np\n",
|
||||
"import torch\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns\n",
|
||||
"import os\n",
|
||||
"from voltron.data import make_ticker_list, GetStockHistory\n",
|
||||
"\n",
|
||||
"sns.set_style('white')\n",
|
||||
"palette = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"sns.set(palette = palette, font_scale=2.0, style=\"white\", rc={\"lines.linewidth\": 4.0})"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 8,
|
||||
"id": "a519062e-3992-4d40-9bf1-990c80990fc0",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"matern_df = pd.read_pickle(\"./matern_calib.pkl\")\n",
|
||||
"volt_df = pd.read_pickle(\"./volt_calib.pkl\")\n",
|
||||
"matern_df.columns = volt_df.columns"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 9,
|
||||
"id": "37d94b48-6e05-4a59-91a3-11291bfc7f4d",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"array(['tewma', 'linear', 'constant', 'ewma'], dtype=object)"
|
||||
]
|
||||
},
|
||||
"execution_count": 9,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"matern_df.Mean.unique()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 10,
|
||||
"id": "1f59e48e-eb80-48b2-8ce5-91596ad136f2",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def ECDF(sample_pxs, true_px): \n",
|
||||
" return (torch.sum(sample_pxs < true_px, 0)/sample_pxs.shape[0])\n",
|
||||
" \n",
|
||||
"def Calibration(pcts, percentile=0.95):\n",
|
||||
" in_band = np.where((pcts < percentile))[0].shape[0]\n",
|
||||
" return in_band/pcts.shape[0]\n",
|
||||
"\n",
|
||||
"def GetCalibration(model, horizon=np.arange(75,100), \n",
|
||||
" logger=[], exp=True):\n",
|
||||
" \n",
|
||||
"\n",
|
||||
" ntrain = 400\n",
|
||||
" n_test_times = 20\n",
|
||||
" ntest = 100\n",
|
||||
" pcts = torch.tensor([])\n",
|
||||
" for tckr in ticker_list:\n",
|
||||
" data = GetStockHistory(tckr, history=1000, end_date=end_date)\n",
|
||||
" for idx, date in enumerate(data.index):\n",
|
||||
"\n",
|
||||
" fpath = \"./saved-outputs/\"+ tckr + \"/\"\n",
|
||||
" fname = model + \"_\"\n",
|
||||
" if model == 'volt':\n",
|
||||
" fname += \"constant\"\n",
|
||||
"\n",
|
||||
" fname += str(date.date()) + \".pt\"\n",
|
||||
"\n",
|
||||
"\n",
|
||||
" if os.path.exists(fpath + fname): \n",
|
||||
" preds = torch.load(fpath + fname)\n",
|
||||
" if isinstance(preds, tuple):\n",
|
||||
" preds = preds[0]\n",
|
||||
" \n",
|
||||
" if preds.shape[-1] == 100:\n",
|
||||
" preds = preds[:, horizon]\n",
|
||||
"\n",
|
||||
" test_y = torch.tensor(data.iloc[idx:idx+100].Close.to_numpy())\n",
|
||||
" if test_y.shape[0] == 100:\n",
|
||||
" if exp:\n",
|
||||
" preds = preds.exp()\n",
|
||||
" pcts = torch.cat((pcts, ECDF(preds, test_y[horizon])))\n",
|
||||
" \n",
|
||||
" if pcts.numel() == 0:\n",
|
||||
" return logger\n",
|
||||
" \n",
|
||||
" pcts = pcts.flatten().numpy()\n",
|
||||
" percentiles = np.linspace(0.05, 0.95, 19)\n",
|
||||
" for pct in percentiles:\n",
|
||||
" clb = Calibration(pcts, pct)\n",
|
||||
" logger.append([clb, np.round(pct, 2), model, \"Constant\", 100])\n",
|
||||
" \n",
|
||||
" return logger"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 11,
|
||||
"id": "5fe6b4d6-b49e-4e06-8874-9be43b4ee4bd",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"data_path = \"../../voltron/data/\"\n",
|
||||
"ticker_list = make_ticker_list(data_path + \"test_tickers.txt\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 12,
|
||||
"id": "da419f67-e0d9-4008-b038-96fd341e10f3",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"end_date = \"2022-01-13\""
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 13,
|
||||
"id": "d9ed9129-be2d-47e8-9866-5a708c27d225",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"log = []\n",
|
||||
"for model in ['volt', 'lstm']:\n",
|
||||
" log = GetCalibration(model, horizon=np.arange(75,100), \n",
|
||||
" logger=log, exp=True)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 14,
|
||||
"id": "89054cf9-fb1c-45a1-aa8b-8884713f4c07",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df = pd.DataFrame(log)\n",
|
||||
"df.columns = ['Calibration', 'Percentile', \"Model\", \"Mean\", \"k\"]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 26,
|
||||
"id": "9097eb62-95d8-443d-bbd4-aae81a3445fa",
|
||||
"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>Calibration</th>\n",
|
||||
" <th>Percentile</th>\n",
|
||||
" <th>Model</th>\n",
|
||||
" <th>Mean</th>\n",
|
||||
" <th>k</th>\n",
|
||||
" </tr>\n",
|
||||
" </thead>\n",
|
||||
" <tbody>\n",
|
||||
" <tr>\n",
|
||||
" <th>0</th>\n",
|
||||
" <td>0.024371</td>\n",
|
||||
" <td>0.05</td>\n",
|
||||
" <td>volt</td>\n",
|
||||
" <td>Constant</td>\n",
|
||||
" <td>100</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>1</th>\n",
|
||||
" <td>0.031473</td>\n",
|
||||
" <td>0.10</td>\n",
|
||||
" <td>volt</td>\n",
|
||||
" <td>Constant</td>\n",
|
||||
" <td>100</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>2</th>\n",
|
||||
" <td>0.042644</td>\n",
|
||||
" <td>0.15</td>\n",
|
||||
" <td>volt</td>\n",
|
||||
" <td>Constant</td>\n",
|
||||
" <td>100</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>3</th>\n",
|
||||
" <td>0.058020</td>\n",
|
||||
" <td>0.20</td>\n",
|
||||
" <td>volt</td>\n",
|
||||
" <td>Constant</td>\n",
|
||||
" <td>100</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>4</th>\n",
|
||||
" <td>0.073815</td>\n",
|
||||
" <td>0.25</td>\n",
|
||||
" <td>volt</td>\n",
|
||||
" <td>Constant</td>\n",
|
||||
" <td>100</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>...</th>\n",
|
||||
" <td>...</td>\n",
|
||||
" <td>...</td>\n",
|
||||
" <td>...</td>\n",
|
||||
" <td>...</td>\n",
|
||||
" <td>...</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>337</th>\n",
|
||||
" <td>0.793388</td>\n",
|
||||
" <td>0.75</td>\n",
|
||||
" <td>volt</td>\n",
|
||||
" <td>tewma</td>\n",
|
||||
" <td>400</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>338</th>\n",
|
||||
" <td>0.812314</td>\n",
|
||||
" <td>0.80</td>\n",
|
||||
" <td>volt</td>\n",
|
||||
" <td>tewma</td>\n",
|
||||
" <td>400</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>339</th>\n",
|
||||
" <td>0.840826</td>\n",
|
||||
" <td>0.85</td>\n",
|
||||
" <td>volt</td>\n",
|
||||
" <td>tewma</td>\n",
|
||||
" <td>400</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>340</th>\n",
|
||||
" <td>0.879091</td>\n",
|
||||
" <td>0.90</td>\n",
|
||||
" <td>volt</td>\n",
|
||||
" <td>tewma</td>\n",
|
||||
" <td>400</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>341</th>\n",
|
||||
" <td>0.944463</td>\n",
|
||||
" <td>0.95</td>\n",
|
||||
" <td>volt</td>\n",
|
||||
" <td>tewma</td>\n",
|
||||
" <td>400</td>\n",
|
||||
" </tr>\n",
|
||||
" </tbody>\n",
|
||||
"</table>\n",
|
||||
"<p>494 rows × 5 columns</p>\n",
|
||||
"</div>"
|
||||
],
|
||||
"text/plain": [
|
||||
" Calibration Percentile Model Mean k\n",
|
||||
"0 0.024371 0.05 volt Constant 100\n",
|
||||
"1 0.031473 0.10 volt Constant 100\n",
|
||||
"2 0.042644 0.15 volt Constant 100\n",
|
||||
"3 0.058020 0.20 volt Constant 100\n",
|
||||
"4 0.073815 0.25 volt Constant 100\n",
|
||||
".. ... ... ... ... ...\n",
|
||||
"337 0.793388 0.75 volt tewma 400\n",
|
||||
"338 0.812314 0.80 volt tewma 400\n",
|
||||
"339 0.840826 0.85 volt tewma 400\n",
|
||||
"340 0.879091 0.90 volt tewma 400\n",
|
||||
"341 0.944463 0.95 volt tewma 400\n",
|
||||
"\n",
|
||||
"[494 rows x 5 columns]"
|
||||
]
|
||||
},
|
||||
"execution_count": 26,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"df"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 15,
|
||||
"id": "185272ab-5724-4dc8-add6-42f054832940",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df = pd.concat([df, matern_df, volt_df])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 16,
|
||||
"id": "81e78fa5-dfcb-4c71-9b09-094be0e58b36",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"mat_df = df[(df['Model'] == 'matern') & (df['Mean'] == 'tewma') & (df['k'] == 400)]\n",
|
||||
"lstm_df = df[df['Model']=='lstm']\n",
|
||||
"volt_df = df[(df['Model'] == 'volt') & (df['Mean'] == 'ewma') & (df['k']==100)]\n",
|
||||
"plt_df = pd.concat([lstm_df, mat_df, volt_df])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 27,
|
||||
"id": "3b8bf1a9-42f3-4a5a-aafb-3385cd32ea9e",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"mat_df = df[(df['Model'] == 'matern') & (df['Mean'] == 'constant')]\n",
|
||||
"volt_df = df[(df['Model'] == 'volt') & (df['Mean'] == 'Constant')]\n",
|
||||
"const_df = pd.concat([mat_df, volt_df])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 41,
|
||||
"id": "bf42d9c5-0bc3-46a8-a1ab-14b87fd8851b",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAABPoAAAIRCAYAAADTKdPXAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAABcSAAAXEgFnn9JSAAEAAElEQVR4nOzdd3hTZfsH8O/Jbrp3gQKFliGj7A0qU+EHshRxgDhwgeBCxdetr7gnr4ATxYlQQFBBEGTvVUaBDgp00b2SZp/fH6WxadI2SdOWlu/nurhIznlWmrQ5ufM8zy2IoiiCiIiIiIiIiIiImjRJYw+AiIiIiIiIiIiI6o6BPiIiIiIiIiIiomaAgT4iIiIiIiIiIqJmgIE+IiIiIiIiIiKiZoCBPiIiIiIiIiIiomaAgT4iIiIiIiIiIqJmgIE+IiIiIiIiIiKiZoCBPiIiIiIiIiIiomaAgT4iIiIiIiIiIqJmgIE+IiIiIiIiIiKiZoCBPiIiIiIiIiIiomaAgT4iIiIiIiIiIqJmgIE+IiIiIiIiIiKiZoCBPiIiIiIiIiIiomaAgT4iIiIiIiIiIqJmQNbYAyAiau46depkc3/u3Ll47LHHGmk01460tDSMHDnS5tiiRYswZcqURhoRNUeffvopFi9ebHPs7Nmz9VZvxowZOHDggPV+//79sWLFCidHS/WJzw0RERFdDRjoI2omTCYTkpKSkJKSguLiYhQXF8NiscDLywtqtRoRERFo1aoVIiMjoVAoGnu4RHSNuXz5MlJSUpCeno7i4mLodDqoVCr4+vrC398fbdq0QYcOHSCVSht7qERERERETRYDfURNmMFgwObNm7F69WocPnwYOp2u1jpyuRwdOnRA9+7d0a9fPwwZMgRBQUENMFpqSkaMGIH09HSnyspkMvj4+MDX1xctWrRA165dERsbi+HDh8PLy6ueR0pXK4vFgp07d2LTpk3YtWsXLl++XGsdLy8vdOnSBTfeeCMmTJiAFi1aNMBIiYiIiIiaDwb6iJqov//+G6+99hqysrJcqmc0GnH69GmcPn0av/zyCyQSCe666y688MILtdblsiRyxGQyobCwEIWFhbh06ZL1NeLr64uJEydi7ty5CAwMbORRUkMRRRFr1qzB0qVLceHCBZfqlpWV4fDhwzh8+DA++OADDBgwAHPmzEH//v3rabR0LXvuueewZs0a6/1WrVph69atjTgiIiIiorpjMg6iJkYURbzyyit49NFHXQ7yOWKxWJCRkeGBkRHZKikpwffff4/x48dj+/btjT0cagAXL17E9OnTsXDhQpeDfFWJooh9+/ZhxowZePDBB5GWluahURIRERERNV+c0UfUxLz88sv45ZdfHJ5r2bIlBg4ciJiYGAQFBcHLywtarRbFxcVITU3FqVOncObMGRgMhgYeNTUHnTt3dnjcaDSiuLgYOTk5Ds/n5uZizpw5+PzzzzF48OD6HCI1ot27d2PevHkoLS11eF6hUKB3796IjY1FUFAQAgMDoVQqodFokJGRgcTERBw8eBAFBQV2dbdv344DBw4gMjKyvh8GEREREVGTxkAfUROyZcsWh0G+rl27YsGCBRg4cCAEQaixjbKyMuzcuRObN2/Gli1boNVq62u41MysW7euxvP5+fn4559/8PXXXyMxMdHmnNFoxGOPPYa//voLwcHB9TlMq8jISKeymFLd/fPPP5g7dy6MRqPduejoaMydO9epPRstFgsOHDiAX3/9FRs3boTJZKqvITe6xx57jNm3mxluZUFERERXAy7dJWoiRFHEm2++aXd8zJgx+PnnnzFo0KBag3xA+Wb3Y8aMwbvvvosdO3Zg4cKFaNu2bX0Mma4xQUFBmDJlCuLi4nDbbbfZnS8tLcXixYsbYWRUnxISEvDEE0/YBfnkcjleeuklrF+/HuPGjXMqMYtEIsHAgQPx/vvv448//sDw4cPra9hERERERM0SA31ETcSRI0fssqCGh4dj0aJFUCgUbrXp6+uLWbNm4dlnn/XEEIkAlC/RfO211zBkyBC7c2vWrOHS8WZEr9fjySeftJsZrFarsWzZMtx1112QSqVutd22bVssXboUb7/9NtRqtSeGS0RERETU7DHQR9RE7Nixw+7Y5MmT4ePj0wijIaqZRCLB008/bXe8IqsqNQ9Lly5FSkqK3fGPPvrIYaDXHZMmTcJPP/2EiIgIj7RHRERERNSccY8+oibCUWbcbt26NcJIGsb58+eRnJyMvLw8FBYWwsvLC8HBwYiIiEBsbCzkcnm99V1WVoYTJ04gJycHBQUFKCkpgUqlgq+vL6KiohAdHY3AwMB667+56NKlC1q1amU3E/XUqVMYNGiQ2+2mpqbizJkzyMrKglarhVwuR2hoKCZNmlTHETvHYDDg5MmTyMrKQmFhIYqLi6FQKODj44PWrVsjJiYGoaGhHuvvwoULSEpKQn5+PgoKCqBQKBAQEICIiAj07NkTKpXKY325Ii8vD8uXL7c7fscdd+CGG27waF/VJYKpzeXLl5GSkoK0tDSUlpZCp9PBx8cH/v7+aNmyJbp37w6lUunRsV4tLl26ZH2d6nQ6BAUFITw8HD179kRAQEC991/X39Nr+bmrLCMjAwkJCTa//0FBQQgLC2uw33+j0Yj4+HgkJyejoKAAMpkMQUFBiIqKQmxsrNuzdomIiKh+MNBH1ETk5+fbHXNmz6u66tSpU7XnDhw4UOP5Cn///bdT2TIvX76ML7/8Elu3bkVaWlq15by9vTFo0CDMnDkTAwYMqLVdZ+j1emsCgGPHjjlMKlBBEAR06tQJN9xwA6ZMmYKoqCiPjKE6BoMBzz//PNavX29zPCIiAp9//rlTz0Fj6dChg12gz9FrGbB/rc2dO9earECr1eL777/HypUrcenSJYf1qwYQ0tLSMHLkSJtjixYtwpQpU1x5CAAAs9mM9evXY/369Th8+DDKyspqLB8VFYXrr78ekydPRpcuXVzu79KlS1i+fDm2b99e7eMFAKVSib59++Kee+7xeHCtNitXrrRbsuvr64sFCxY06Dgqy8/Px5YtW7Bnzx4cPHgQubm5NZaXy+Xo2bMn7rrrLtx0002QSBpuocOnn35qt2dlXZPHiKKIuLg4LF++HOfOnXNYRi6XY+DAgXjwwQfRv39/l/vw9O9phYZ67kaMGGH3N6lCenq6U39Pv/vuO4fvPTNmzMCBAwes9/v37+9Wgo7i4mJ8/fXX2Lx5M5KSkqotp1Qq0a9fP0yfPh2jR492uZ+4uDgsXLjQ5ljl9+ucnBx8/vnnWLNmDUpKShy24efnh8mTJ+ORRx7hF2BERERXCQb6iJoIR/vwOZrl1xSZzWYsXrwY33zzTa0BFADQaDTYsmULtmzZghtuuAGvvPIKWrZs6Xb/P/30E/73v/8hJyfHqfKiKOLMmTM4c+YMli1bhk8++QQ33XST2/3XpKioCHPnzrX58AgA1113HZYtW4bw8PB66ddTHC0tr+4DY3WOHz+Oxx9/vNFe75s2bcL777+PCxcuOF0nNTUVqamp+O677/Dcc8/h3nvvdapeaWkpPvjgA6xcubLGYHMFvV6P3bt3Y/fu3ejTpw/ee++9Ov0uuCIuLs7u2KRJk+Dt7d0g/Vf11FNPuZyp12g04uDBgzh48CCio6Px8ccfo0OHDvU4yvqTn5+Pxx57DIcOHaqxnNFoxM6dO7Fr1y5MnToVL774okdmhdXl9/Raf+4qW7FiBT799FMUFRXVWlav12PXrl3YtWsXevXqhVdffdVjX/xs2rQJL7zwAoqLi2ssV1xcjG+//Rbr1q3DsmXL0LNnT4/0T0RERO7jHn1ETYSjpYB//vlnI4zEs8rKyjBnzhx89tlnTgX5qtq+fTtuv/12nDlzxuW6er0eTz/9NF555RWng3yOaDQat+vW5NKlS5g+fbpdkG/YsGH4/vvvr/ogH1AeuKrK19fX6foHDx7EjBkzGiXIZ7FY8M4772DevHkuBfmqcvQzcCQ9PR133HEHfvjhB6eCfFUdPnwYt912G44dO+ZyXVclJSXh4sWLdsdvv/32eu+7OkePHnUpUFRVcnIypk2bhj179nhwVA2jqKgId911V61BvspEUcSqVavw8MMPQ6fT1an/uv6eXsvPXQWz2YyXXnoJb7zxhlNBvqqOHj2KO++8E3v37q3zWH766SfMnz+/1iBfZYWFhbj33nuRkJBQ5/6JiIiobjijj6iJ6NWrF3755RebY3v27MGKFSswY8aMeuu38t5YFy9etFmqp1ar0aZNm1rbqG4/PYvFgkcffdThhzNvb28MHz4csbGxCA0NRWlpKVJTU7Flyxa7oEt2djbuvvturF69Gm3btnXqcRmNRtx///04ePCg3TmJRIKuXbti0KBBaNGiBQICAmAwGFBYWIizZ88iPj6+xuVUnhAfH4+HH34YeXl5NsenTZuGl19+GTJZ0/jznZiYaHcsKCjIqbo5OTmYO3cu9Hq99VhsbCyGDBmCVq1awdvbG9nZ2UhOTsbGjRs9NuYKCxYswIYNGxye69ixIwYPHow2bdogMDAQRqMRRUVFSEpKwsmTJ3H69GmIouh0X+np6Zg2bZrD5YqxsbHo3bs32rVrBz8/PxiNRuTk5ODo0aPYsWOHTRbj3NxcPPTQQ4iLi0OrVq1cf9BO2r9/v92xkJCQq2ZGlVQqRZcuXdChQwe0a9cOgYGB1pmGFX9Ljh8/jiNHjsBisVjrabVaPPHEE1i7di1atGjRWMN32TPPPGOTFKVFixYYPXo0oqOj4efnh9zcXJw8eRJ///23XeB57969eOKJJ7BkyRK3+vb072l9P3fR0dHWLxsyMzNtgmpyuRzR0dG1jrE+skC/+OKLWL16td1xpVKJoUOHol+/fggNDYVOp0N6ejr+/vtvu6XepaWlmD17Nr799lv06dPHrXHs2LEDr7/+uvXvl6+vL4YMGYJevXohODgYFosF6enp+OeLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 900x450 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"from matplotlib.lines import Line2D\n",
|
||||
"fig, ax = plt.subplots(1,1,dpi=150, figsize=(6, 3))\n",
|
||||
"\n",
|
||||
"percentiles = np.linspace(0.05, 0.95, 19)\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"pal = [palette[5], palette[7]]\n",
|
||||
"# pal = [palette[5]]\n",
|
||||
"sns.lineplot(x='Percentile', y=\"Calibration\", hue='Model', data=const_df, ax=ax, alpha=0.2,\n",
|
||||
" palette=pal, legend=True)\n",
|
||||
"sns.scatterplot(x='Percentile', y=\"Calibration\", hue='Model', data=const_df, ax=ax, s=120, legend=False, zorder=4,\n",
|
||||
" palette=pal, alpha=0.35)\n",
|
||||
"\n",
|
||||
"pal = [ palette[0], palette[4], palette[6]]\n",
|
||||
"sns.lineplot(x='Percentile', y=\"Calibration\", hue='Model', data=plt_df, ax=ax, alpha=0.5,\n",
|
||||
" palette=pal, legend=True)\n",
|
||||
"sns.scatterplot(x='Percentile', y=\"Calibration\", hue='Model', data=plt_df, ax=ax, s=120, legend=False, zorder=4,\n",
|
||||
" palette=pal)\n",
|
||||
"x = np.linspace(0.05,0.95)\n",
|
||||
"y = np.linspace(0, len(percentiles))\n",
|
||||
"ax.plot(x, x, color=\"gray\", lw=1., ls=\"--\")\n",
|
||||
"ax.set_title(\"Stock Price Calibration\")\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"plt.tick_params(labelsize=16)\n",
|
||||
"sns.despine()\n",
|
||||
"\n",
|
||||
"custom_lines = [Line2D([0], [0], color=palette[0], lw=2),\n",
|
||||
" Line2D([0], [0], color=palette[4], lw=2),\n",
|
||||
" Line2D([0], [0], color=palette[6], lw=2),\n",
|
||||
" Line2D([0], [0], color=palette[5], lw=2),\n",
|
||||
" Line2D([0], [0], color=palette[7], lw=2)]\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"plt.legend(custom_lines, ['LSTM', r\"Matérn + Magpie\", \"Volt + Magpie\", \"Matérn + Constant\",\n",
|
||||
" \"Volt + Constant\"],\n",
|
||||
" fontsize=14, frameon=False, bbox_to_anchor=(1., 0.85))\n",
|
||||
"# ax.legend(fontsize=14, bbox_to_anchor=(1., 0.75))\n",
|
||||
"# plt.label(\"Percentile\")\n",
|
||||
"plt.savefig(\"./stock_calibration.pdf\", bbox_inches=\"tight\")\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "957ddda1-e78b-4bb7-9003-886f18e424f7",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "d9fa07e1-f9e5-4309-9311-160c53b43d0d",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.12"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,722 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "9754a692-edb8-41bf-8159-fd82edcb9195",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Warning no robinhood utils.\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import pandas as pd\n",
|
||||
"import torch\n",
|
||||
"from torch import nn\n",
|
||||
"import seaborn as sns\n",
|
||||
"import time\n",
|
||||
"import copy\n",
|
||||
"import sys\n",
|
||||
"from torch.utils.data import DataLoader\n",
|
||||
"from voltron.data import make_ticker_list, GetStockHistory\n",
|
||||
"sns.set_style('white')\n",
|
||||
"# style.use('whitegrid')\n",
|
||||
"palette = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"sns.set(palette = palette, font_scale=2.0, style=\"white\", rc={\"lines.linewidth\": 2.0})"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "dbc24c1a-f499-4b8f-8544-6d4a66f48488",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Dataset Setup"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 29,
|
||||
"id": "62f91be1-38aa-4781-8b65-c5caf9752045",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from torch.utils.data import Dataset\n",
|
||||
"\n",
|
||||
"class SequenceDataset(Dataset):\n",
|
||||
" def __init__(self, data, sequence_length=5):\n",
|
||||
" self.sequence_length = sequence_length\n",
|
||||
" self.X = data.float()\n",
|
||||
"\n",
|
||||
" def __len__(self):\n",
|
||||
" return self.X.shape[0]-1\n",
|
||||
"\n",
|
||||
" def __getitem__(self, i): \n",
|
||||
" if i >= self.sequence_length - 1:\n",
|
||||
" i_start = i - self.sequence_length + 1\n",
|
||||
" x = self.X[i_start:(i + 1)]\n",
|
||||
" else:\n",
|
||||
" padding = self.X[0].repeat(self.sequence_length - i - 1, 1).squeeze(-1)\n",
|
||||
" x = self.X[0:(i + 1)]\n",
|
||||
" x = torch.cat((padding, x), 0)\n",
|
||||
" \n",
|
||||
" return x.unsqueeze(0), self.X[i+1]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 32,
|
||||
"id": "566c9386-c297-430f-b4ef-a0f3ba7aec37",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"tckr = \"JPM\"\n",
|
||||
"ntrain = 400\n",
|
||||
"lookback = 1\n",
|
||||
"pxs = GetStockHistory(tckr, end_date=\"2021-12-07\", history=ntrain + lookback).Close.to_numpy()\n",
|
||||
"data = torch.FloatTensor(pxs).log()\n",
|
||||
"\n",
|
||||
"data = (data - data.mean())/data.std()\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"# xin = torch.linspace(0, 5*np.pi, 250)\n",
|
||||
"# data = torch.sin(xin) + 0.2 * torch.randn(xin.shape)\n",
|
||||
"\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 33,
|
||||
"id": "5dece53d-785f-4ca9-9cbf-03e16de71e9b",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"seq_len = 25\n",
|
||||
"dset = SequenceDataset(data, seq_len)\n",
|
||||
"\n",
|
||||
"trgts = []\n",
|
||||
"for i in range(len(dset)):\n",
|
||||
" trgts.append(dset[i][1].item())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 34,
|
||||
"id": "2317ba11-4f62-443f-9ad8-63f9a0996ee9",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"train_loader = DataLoader(dset, batch_size=20, shuffle=True)\n",
|
||||
"X, y = next(iter(train_loader))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 35,
|
||||
"id": "ffecfde0-6938-4dee-b1b0-6627012e3329",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"class ShallowRegressionLSTM(nn.Module):\n",
|
||||
" def __init__(self, input_size, hidden_units=128):\n",
|
||||
" super().__init__()\n",
|
||||
" self.input_size = input_size # this is the number of features\n",
|
||||
" self.hidden_units = hidden_units\n",
|
||||
" self.num_layers = 5\n",
|
||||
"\n",
|
||||
" self.lstm = nn.LSTM(\n",
|
||||
" input_size=input_size,\n",
|
||||
" hidden_size=hidden_units,\n",
|
||||
" batch_first=True,\n",
|
||||
" num_layers=self.num_layers\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
" self.linear = nn.Linear(in_features=self.hidden_units, out_features=2)\n",
|
||||
" \n",
|
||||
"\n",
|
||||
" def forward(self, x):\n",
|
||||
" batch_size = x.shape[0]\n",
|
||||
" h0 = torch.zeros(self.num_layers, batch_size, self.hidden_units).requires_grad_()\n",
|
||||
" c0 = torch.zeros(self.num_layers, batch_size, self.hidden_units).requires_grad_()\n",
|
||||
"\n",
|
||||
" _, (hn, _) = self.lstm(x, (h0, c0))\n",
|
||||
" out = self.linear(hn[0]) # First dim of Hn is num_layers, which is set to 1 above.\n",
|
||||
" output = torch.zeros_like(out)\n",
|
||||
" output[:, 0] = out[:, 0]\n",
|
||||
" output[:, 1] = torch.exp(out[:, 1])\n",
|
||||
" return output\n",
|
||||
" \n",
|
||||
"from torch.autograd import Variable \n",
|
||||
"class LSTM1(nn.Module):\n",
|
||||
" def __init__(self, num_classes, seq_len, hidden_size, num_layers):\n",
|
||||
" super(LSTM1, self).__init__()\n",
|
||||
" self.num_classes = num_classes #number of classes\n",
|
||||
" self.num_layers = num_layers #number of layers\n",
|
||||
" self.input_size = seq_len #input size\n",
|
||||
" self.hidden_size = hidden_size #hidden state\n",
|
||||
"\n",
|
||||
" self.lstm = nn.LSTM(input_size=seq_len, hidden_size=hidden_size,\n",
|
||||
" num_layers=num_layers, batch_first=True) #lstm\n",
|
||||
" self.fc_1 = nn.Linear(hidden_size, 128) #fully connected 1\n",
|
||||
" self.fc = nn.Linear(128, num_classes) #fully connected last layer\n",
|
||||
"\n",
|
||||
" self.relu = nn.ReLU()\n",
|
||||
" self.softplus = nn.Softplus()\n",
|
||||
" \n",
|
||||
" def forward(self,x):\n",
|
||||
" h_0 = Variable(torch.zeros(self.num_layers, x.size(0), self.hidden_size)) #hidden state\n",
|
||||
" c_0 = Variable(torch.zeros(self.num_layers, x.size(0), self.hidden_size)) #internal state\n",
|
||||
" # Propagate input through LSTM\n",
|
||||
" output, (hn, cn) = self.lstm(x, (h_0, c_0)) #lstm with input, hidden, and internal state\n",
|
||||
"\n",
|
||||
" hn = hn[self.num_layers-1]\n",
|
||||
" hn = hn.view(-1, self.hidden_size) #reshaping the data for Dense layer next\n",
|
||||
" out = self.relu(hn)\n",
|
||||
" out = self.fc_1(out) #first Dense\n",
|
||||
" out = self.relu(out) #relu\n",
|
||||
" out = self.fc(out) #Final Output\n",
|
||||
" \n",
|
||||
" output = torch.zeros_like(out)\n",
|
||||
" output[:, 0] = out[:, 0]\n",
|
||||
" output[:, 1] = self.softplus(out[:, 1])\n",
|
||||
" return output"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 36,
|
||||
"id": "3bc83b15-3b0f-4f21-803b-29e354cb558f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# model = ShallowRegressionLSTM(seq_len)\n",
|
||||
"model = LSTM1(2, seq_len, 128, 1)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 37,
|
||||
"id": "41547213-7851-4ad1-ac04-03cab5e0a91c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def NLL(targets, outputs):\n",
|
||||
" dist = torch.distributions.Normal(outputs[:, 0], outputs[:, 1])\n",
|
||||
" return -dist.log_prob(targets).sum()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 38,
|
||||
"id": "e6958f51-5cd0-414a-b419-3ce55f12b486",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def train_model(data_loader, model, loss_function, optimizer, epochs=200):\n",
|
||||
" num_batches = len(data_loader)\n",
|
||||
" total_loss = 0\n",
|
||||
" model.train()\n",
|
||||
" for epoch in range(epochs):\n",
|
||||
" for X, y in data_loader:\n",
|
||||
" output = model(X)\n",
|
||||
" loss = loss_function(y, output)\n",
|
||||
"\n",
|
||||
" optimizer.zero_grad()\n",
|
||||
" loss.backward()\n",
|
||||
" optimizer.step()\n",
|
||||
"\n",
|
||||
" total_loss += loss.item()\n",
|
||||
"\n",
|
||||
" if epoch%10 == 0:\n",
|
||||
" avg_loss = total_loss / num_batches\n",
|
||||
" print(f\"Train loss: {avg_loss}, Epoch: {epoch}\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 39,
|
||||
"id": "92484e45-11c9-4fe1-bc5b-87b0f27d5a2c",
|
||||
"metadata": {
|
||||
"tags": []
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Train loss: 9.58264362514019, Epoch: 0\n",
|
||||
"Train loss: -7.665388387441635, Epoch: 10\n",
|
||||
"Train loss: -96.57648058235645, Epoch: 20\n",
|
||||
"Train loss: -222.45588295161724, Epoch: 30\n",
|
||||
"Train loss: -393.4549728780985, Epoch: 40\n",
|
||||
"Train loss: -557.4210761398077, Epoch: 50\n",
|
||||
"Train loss: -733.938936367631, Epoch: 60\n",
|
||||
"Train loss: -916.8391911536455, Epoch: 70\n",
|
||||
"Train loss: -1094.614848962426, Epoch: 80\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"ERROR:root:Internal Python error in the inspect module.\n",
|
||||
"Below is the traceback from this internal error.\n",
|
||||
"\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Traceback (most recent call last):\n",
|
||||
" File \"/home/greg_b/miniconda3/lib/python3.8/site-packages/IPython/core/interactiveshell.py\", line 3444, in run_code\n",
|
||||
" exec(code_obj, self.user_global_ns, self.user_ns)\n",
|
||||
" File \"<ipython-input-39-1bf22962a03c>\", line 2, in <module>\n",
|
||||
" train_model(train_loader, model, NLL, optimizer)\n",
|
||||
" File \"<ipython-input-38-1b396ffa6e98>\", line 11, in train_model\n",
|
||||
" loss.backward()\n",
|
||||
" File \"/home/greg_b/miniconda3/lib/python3.8/site-packages/torch/_tensor.py\", line 307, in backward\n",
|
||||
" torch.autograd.backward(self, gradient, retain_graph, create_graph, inputs=inputs)\n",
|
||||
" File \"/home/greg_b/miniconda3/lib/python3.8/site-packages/torch/autograd/__init__.py\", line 154, in backward\n",
|
||||
" Variable._execution_engine.run_backward(\n",
|
||||
"KeyboardInterrupt\n",
|
||||
"\n",
|
||||
"During handling of the above exception, another exception occurred:\n",
|
||||
"\n",
|
||||
"Traceback (most recent call last):\n",
|
||||
" File \"/home/greg_b/miniconda3/lib/python3.8/site-packages/IPython/core/interactiveshell.py\", line 2064, in showtraceback\n",
|
||||
" stb = value._render_traceback_()\n",
|
||||
"AttributeError: 'KeyboardInterrupt' object has no attribute '_render_traceback_'\n",
|
||||
"\n",
|
||||
"During handling of the above exception, another exception occurred:\n",
|
||||
"\n",
|
||||
"Traceback (most recent call last):\n",
|
||||
" File \"/home/greg_b/miniconda3/lib/python3.8/site-packages/IPython/core/ultratb.py\", line 1101, in get_records\n",
|
||||
" return _fixed_getinnerframes(etb, number_of_lines_of_context, tb_offset)\n",
|
||||
" File \"/home/greg_b/miniconda3/lib/python3.8/site-packages/IPython/core/ultratb.py\", line 248, in wrapped\n",
|
||||
" return f(*args, **kwargs)\n",
|
||||
" File \"/home/greg_b/miniconda3/lib/python3.8/site-packages/IPython/core/ultratb.py\", line 281, in _fixed_getinnerframes\n",
|
||||
" records = fix_frame_records_filenames(inspect.getinnerframes(etb, context))\n",
|
||||
" File \"/home/greg_b/miniconda3/lib/python3.8/inspect.py\", line 1503, in getinnerframes\n",
|
||||
" frameinfo = (tb.tb_frame,) + getframeinfo(tb, context)\n",
|
||||
" File \"/home/greg_b/miniconda3/lib/python3.8/inspect.py\", line 1461, in getframeinfo\n",
|
||||
" filename = getsourcefile(frame) or getfile(frame)\n",
|
||||
" File \"/home/greg_b/miniconda3/lib/python3.8/inspect.py\", line 708, in getsourcefile\n",
|
||||
" if getattr(getmodule(object, filename), '__loader__', None) is not None:\n",
|
||||
" File \"/home/greg_b/miniconda3/lib/python3.8/inspect.py\", line 744, in getmodule\n",
|
||||
" for modname, module in sys.modules.copy().items():\n",
|
||||
"KeyboardInterrupt\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"ename": "TypeError",
|
||||
"evalue": "object of type 'NoneType' has no len()",
|
||||
"output_type": "error",
|
||||
"traceback": [
|
||||
"\u001b[0;31m---------------------------------------------------------------------------\u001b[0m",
|
||||
"\u001b[0;31mKeyboardInterrupt\u001b[0m Traceback (most recent call last)",
|
||||
" \u001b[0;31m[... skipping hidden 1 frame]\u001b[0m\n",
|
||||
"\u001b[0;32m<ipython-input-39-1bf22962a03c>\u001b[0m in \u001b[0;36m<module>\u001b[0;34m\u001b[0m\n\u001b[1;32m 1\u001b[0m \u001b[0moptimizer\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mtorch\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0moptim\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mAdam\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mmodel\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mparameters\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mlr\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0;36m0.01\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m----> 2\u001b[0;31m \u001b[0mtrain_model\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mtrain_loader\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mmodel\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mNLL\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0moptimizer\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m",
|
||||
"\u001b[0;32m<ipython-input-38-1b396ffa6e98>\u001b[0m in \u001b[0;36mtrain_model\u001b[0;34m(data_loader, model, loss_function, optimizer, epochs)\u001b[0m\n\u001b[1;32m 10\u001b[0m \u001b[0moptimizer\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mzero_grad\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 11\u001b[0;31m \u001b[0mloss\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mbackward\u001b[0m\u001b[0;34m(\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 12\u001b[0m \u001b[0moptimizer\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mstep\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/lib/python3.8/site-packages/torch/_tensor.py\u001b[0m in \u001b[0;36mbackward\u001b[0;34m(self, gradient, retain_graph, create_graph, inputs)\u001b[0m\n\u001b[1;32m 306\u001b[0m inputs=inputs)\n\u001b[0;32m--> 307\u001b[0;31m \u001b[0mtorch\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mautograd\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mbackward\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mgradient\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mretain_graph\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mcreate_graph\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0minputs\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0minputs\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 308\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/lib/python3.8/site-packages/torch/autograd/__init__.py\u001b[0m in \u001b[0;36mbackward\u001b[0;34m(tensors, grad_tensors, retain_graph, create_graph, grad_variables, inputs)\u001b[0m\n\u001b[1;32m 153\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 154\u001b[0;31m Variable._execution_engine.run_backward(\n\u001b[0m\u001b[1;32m 155\u001b[0m \u001b[0mtensors\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mgrad_tensors_\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mretain_graph\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mcreate_graph\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0minputs\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;31mKeyboardInterrupt\u001b[0m: ",
|
||||
"\nDuring handling of the above exception, another exception occurred:\n",
|
||||
"\u001b[0;31mAttributeError\u001b[0m Traceback (most recent call last)",
|
||||
"\u001b[0;32m~/miniconda3/lib/python3.8/site-packages/IPython/core/interactiveshell.py\u001b[0m in \u001b[0;36mshowtraceback\u001b[0;34m(self, exc_tuple, filename, tb_offset, exception_only, running_compiled_code)\u001b[0m\n\u001b[1;32m 2063\u001b[0m \u001b[0;31m# in the engines. This should return a list of strings.\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 2064\u001b[0;31m \u001b[0mstb\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mvalue\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_render_traceback_\u001b[0m\u001b[0;34m(\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 2065\u001b[0m \u001b[0;32mexcept\u001b[0m \u001b[0mException\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;31mAttributeError\u001b[0m: 'KeyboardInterrupt' object has no attribute '_render_traceback_'",
|
||||
"\nDuring handling of the above exception, another exception occurred:\n",
|
||||
"\u001b[0;31mTypeError\u001b[0m Traceback (most recent call last)",
|
||||
" \u001b[0;31m[... skipping hidden 1 frame]\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/lib/python3.8/site-packages/IPython/core/interactiveshell.py\u001b[0m in \u001b[0;36mshowtraceback\u001b[0;34m(self, exc_tuple, filename, tb_offset, exception_only, running_compiled_code)\u001b[0m\n\u001b[1;32m 2064\u001b[0m \u001b[0mstb\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mvalue\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_render_traceback_\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 2065\u001b[0m \u001b[0;32mexcept\u001b[0m \u001b[0mException\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 2066\u001b[0;31m stb = self.InteractiveTB.structured_traceback(etype,\n\u001b[0m\u001b[1;32m 2067\u001b[0m value, tb, tb_offset=tb_offset)\n\u001b[1;32m 2068\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/lib/python3.8/site-packages/IPython/core/ultratb.py\u001b[0m in \u001b[0;36mstructured_traceback\u001b[0;34m(self, etype, value, tb, tb_offset, number_of_lines_of_context)\u001b[0m\n\u001b[1;32m 1365\u001b[0m \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 1366\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mtb\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mtb\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1367\u001b[0;31m return FormattedTB.structured_traceback(\n\u001b[0m\u001b[1;32m 1368\u001b[0m self, etype, value, tb, tb_offset, number_of_lines_of_context)\n\u001b[1;32m 1369\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/lib/python3.8/site-packages/IPython/core/ultratb.py\u001b[0m in \u001b[0;36mstructured_traceback\u001b[0;34m(self, etype, value, tb, tb_offset, number_of_lines_of_context)\u001b[0m\n\u001b[1;32m 1265\u001b[0m \u001b[0;32mif\u001b[0m \u001b[0mmode\u001b[0m \u001b[0;32min\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mverbose_modes\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 1266\u001b[0m \u001b[0;31m# Verbose modes need a full traceback\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1267\u001b[0;31m return VerboseTB.structured_traceback(\n\u001b[0m\u001b[1;32m 1268\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0metype\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mvalue\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mtb\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mtb_offset\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mnumber_of_lines_of_context\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 1269\u001b[0m )\n",
|
||||
"\u001b[0;32m~/miniconda3/lib/python3.8/site-packages/IPython/core/ultratb.py\u001b[0m in \u001b[0;36mstructured_traceback\u001b[0;34m(self, etype, evalue, etb, tb_offset, number_of_lines_of_context)\u001b[0m\n\u001b[1;32m 1122\u001b[0m \u001b[0;34m\"\"\"Return a nice text document describing the traceback.\"\"\"\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 1123\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1124\u001b[0;31m formatted_exception = self.format_exception_as_a_whole(etype, evalue, etb, number_of_lines_of_context,\n\u001b[0m\u001b[1;32m 1125\u001b[0m tb_offset)\n\u001b[1;32m 1126\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/lib/python3.8/site-packages/IPython/core/ultratb.py\u001b[0m in \u001b[0;36mformat_exception_as_a_whole\u001b[0;34m(self, etype, evalue, etb, number_of_lines_of_context, tb_offset)\u001b[0m\n\u001b[1;32m 1080\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 1081\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1082\u001b[0;31m \u001b[0mlast_unique\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mrecursion_repeat\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mfind_recursion\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0morig_etype\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mevalue\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mrecords\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 1083\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 1084\u001b[0m \u001b[0mframes\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mformat_records\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mrecords\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mlast_unique\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mrecursion_repeat\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/lib/python3.8/site-packages/IPython/core/ultratb.py\u001b[0m in \u001b[0;36mfind_recursion\u001b[0;34m(etype, value, records)\u001b[0m\n\u001b[1;32m 380\u001b[0m \u001b[0;31m# first frame (from in to out) that looks different.\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 381\u001b[0m \u001b[0;32mif\u001b[0m \u001b[0;32mnot\u001b[0m \u001b[0mis_recursion_error\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0metype\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mvalue\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mrecords\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 382\u001b[0;31m \u001b[0;32mreturn\u001b[0m \u001b[0mlen\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mrecords\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;36m0\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m 383\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 384\u001b[0m \u001b[0;31m# Select filename, lineno, func_name to track frames with\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;31mTypeError\u001b[0m: object of type 'NoneType' has no len()"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"optimizer = torch.optim.Adam(model.parameters(), lr=0.01)\n",
|
||||
"train_model(train_loader, model, NLL, optimizer)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 40,
|
||||
"id": "07eab392-37e2-4dae-bfa4-dd177e3d552f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"means = []\n",
|
||||
"vrs = []\n",
|
||||
"for X, y in dset:\n",
|
||||
" output = model(X.unsqueeze(0))\n",
|
||||
" means = means + list(output[:, 0].detach().numpy())\n",
|
||||
" vrs = vrs + list(output[:, 1].detach().numpy())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 41,
|
||||
"id": "b10866df-b554-44a6-bf73-b98abb2e6400",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"<matplotlib.collections.PathCollection at 0x7f0b7c18d910>"
|
||||
]
|
||||
},
|
||||
"execution_count": 41,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYwAAAEFCAYAAADwhtBaAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAABpqklEQVR4nO39d5xdV33vjb/XbqdNL+rNkiXbkgyWMWAbgxNicAmQELghYAIJuUDim4QbQgKEPAnce/klEG4S8vBwUzGEnlBjwPY1TeDIYMu2bEtGki1ZdSxp+plTd1u/P3aZU2fO9JFnvV8vXmhO2/ssz6zP+nYhpZQoFAqFQjEN2lLfgEKhUCguDpRgKBQKhaIllGAoFAqFoiWUYCgUCoWiJZRgKBQKhaIljKW+gblSKpU4ePAg/f396Lq+1LejUCgUFwWe5zE4OMju3btJJpMtveeiF4yDBw9y++23L/VtKBQKxUXJ5z//ea655pqWXnvRC0Z/fz8QfOk1a9Ys8d0oFArFxcG5c+e4/fbb4z20FS56wYjcUGvWrGHDhg1LfDcKhUJxcTETV74KeisUCoWiJZRgKBQKhaIllGAoFAqFoiWUYCgUCoWiJZRgKBQKhaIllGAoFAqFoiWUYCgUCsUC40uJ6138o4cu+joMhUKhWO4M5RzGiy49aYOejIkQYqlvaVYowVAoFIo5UnZ9fClJmdVFcFJKBsbLFMo+liEYybtYhobrSzQBng8ZSydhXhzOHiUYCoVCMQeKtsfAeBkpYVNPEsuY3PyLjk/B9kkYAiEEuiY5l7URSCSBlTFWcNjUm8LQlr/VcXHImkKhWHZIKcmVPXx58fvmZ0vB9jg7VkYLXUzZkhs/J6VktOCiC2IXlKlrWLogYegkDY2koSGB0yMlBidsHM8nV/Yo2h6ev/zWVVkYCoViRkgpKdg+2ZLLRMkjbWms70pctH752WK7PuezNromMDSBJmC86NKdNtE1EVoXHgm9el20mnWydA3Pl4wXXcaLLhIQgK4J1nclYovF8XykpMqCWWyUYCgUipbxfMlI3mGs4KAJQdIQFG2fsXCjfK5TsD3Kro8mYCTvIpGYWrCBa0LgS0mu7JI0dQbGyxiaaElIdU2g17ikbM/n/ITNhlCML0zYFG2fzb1JTH1pREMJhkKhqEJKyUTJpS1hoFVsYlJKzoyWcDxJwtDijdDSYTTv0Jky6k7PzyU8X/LseJnIUySARM1p39AE2aKHE6bQziUuYWqCkhNYcu0Jg6LjIwliJmZKCYZCoVgG5Msez2ZtetOSvnaL4ZzNWNElZWrYniRZs0lqmsBxJfmyR3tyeW8pUkocT87KrVNyfHwJSUPDl1HIuhpdQMn1Kbs+pj438RRCYOlBSq4mBMjg83Nln47UnD561qigt0KhiJFSMlJwsDSN0aLLcM5mpOBihKdds8mJWdcEo3kHdxkGaivJlT1OjZQoOT5Syvj/W2Gi5BJ9fU00djUJEcQzgpjG3K0tTQikhAsTNpoWWCwF28Px/Dl/9qzuZ0muqlAolg3ZohO6mnwcT2K7El0DSxeMFFx0Ldi4TF2r87NH6AJsX3J2tNTyBjzfRALQrKLa9SWDORuAgfESgzmH0yMlzmftaTdg1w8ywpoJZiVGg3jEXLB0gQBMLXADCmC04E73tgVheduPCoViQQk2UQdPSkbzLqYRbHTRxpQ0Wtv4hBAkdEHZ9Sk4Phmr9Slu80Wu7HEua2NogjUdFqmaexjJO/h+EHdwfUm26GIZgrztURjx6EwFVdjR96mkaHsNH18MhBBV7i1TF2SLLm0JnaLtLWrluLIwFIoVTNH2kBKSusZ4KUjrnEugVhOQLS7+6VfKIHvL0AJf/5mxMvny5H34YSA/2ngNTZAwNDQhsHQNXQhG8g5nx8qM5B2klHi+jK2lkus3jFksBUIEKbznsjZDOYd82Vu0ayvBUCiWIW6UkbPAMYFc2UMLC8sMIbBdyVy8KZGPfaHvu5aCE7jTDE1g6AJTE5yfsOPit0LZw5f1NRARmha43EqOz2jBZWC8zPGhIiP5QHRKtl+VMbbUGFqQwmtogYW0WCjBUCiWEYWyy3jRIV9yGS+4jJemP637UjKUszk3Xp5RdbDnSwq2F1sUhi5Imdqc3BsiDNIOh6f0qa5tu9VxA7fiRD8Tio7H+fEyesV965rA94NgseP5XMjZ08Yf9NDqMDRBOQzwT5RdpAzudY5JT/NK4ALUWGwPmYphKBTLhKLjMTBux5W+liEYL7h0pYwpN/HhnMNY0QUJaculI9VaAV2u5CLl/PvlLV0wVnTpSBkkGsRApJQMjJUpuT7ruxKkLR3Pl5waKdGR1OlMGS0XvEHgApOSOP5SeR/5shdbUXqLO76uCXREIBRekC4sWZr4xXJDCYZCsQwYyTuM5B10IeJTo64FQeSgbkBQsF18n7CJXdCaQtc0xosuCV3gySB7pj3ZXGACX39QoZydY7yiGUHAXFK0vbrCtlzJZSTvYHsSXQQFfylTI1d28fyg99JowWVNh9VSTUfUpqTR9xBCkDCCjX82m330PS7knEU/yS9XlGAoFEuI70vGig7DeRdLb5y7X7A9TF0wlHMoOX6wkQlC942Im9vpBCfiXNmj6PikTY22mk235PqMFBwEIATzmv5ZiS6C031XRbuQsaLD4ISDIQLXj5SSguMzOOFQcr24fsH1gr5KjQTDD1NnU2E7cNcPgtNTWQ9zsQxMXeB4ctZFeMOHjjGwdz92No/VkWHdjdfQu2vbrO9nqVGCoVAsAr6UlB0/TvXMFh3yZR9dC9w3zcTC0ARjYeFc2QmqrCs3QM+fDFIHQhK1z4aJEqQsPRYFKSXjBReNhW9gp2tQDGdEaCKwlIYmnKrvGfjhYbzkBi64cFPWNeJ6Ctv1GCm4ZCwdISBve+TLHp0pk3KYuSRYOHeREAKrxdTiSoYPHeP0fQ/glez4MTub58Rdexl6/CiXvfHW+bzNRUMJhkKxCBRtn4HxMhu6EyQMjeG8Gweok0Zzf72uCUquz7NZG1Ovf12thWBVNKUruz5FxydjBdk/IwUnmM2wCNFbEZhAlJ0gu+jceBlN1GcpCSHQhZx8T/z/MhCZnIPrBVZF8FzQmiNXciG4RCw0i02tKOipBBtvuhaAk3ffj3QbZy/lTj7Lw3/5qfjn6H0Xg+WhBEOhWATKbnBizhbduHCstidTM5Kh+2amp2hNwMBYie60STY8xScaiM5C8mzWxvclQlSLWSXNOq9mS24Qv2lwz7M59c8nw4eOceJbPwoUK8Qrljn5nR+jWWaVWJQ8nx+ez7GjPcHW9kTdZ0XvA5a9aCjBUCgWgaLtY+lBK/Bc2cOaYexgNpt80MpDMlpw0DTRdMNeKExdxJXVM8XQBNmSizmDbKnF5PR3f1IlFhHS8/GK5arHvnM2y/fP5wB48yXdXNuXafi+E9/6EbC8RUMJhkKxwEgpKblB4NoO+xxpi+RG0cIA81JsuZoQaLPsEKJrgpSYW03IQjF86FidKOQcj5ShVdWCAPzoQi4WC4AvnRjlkjaLCcen09QxBHQnwm1YSk7ctZfcmfNsvvn6qustl8C5EgyFYoEpuX5FvYNc9NnNF+uMiuUiFpUbtp60qgLZAF88Mcp/DuZZmzJ4/67V8XqXPJ+vnhoDYE93CtuXHBov8T+fOF/1/tdv6uTnVrfHPw89epjSyDiXvfHWOteXnc0vqSWiBEOhWGDyZY/oiD8b98xyYzmdeBea4UPHqgLYtWLx0HCB/xzMA/Bs0eV03mFzmwXA0WwZT0J/wuA3t/XwyEiRQ+Olumt862yWF/dlSFW4DHMnn+WRv/p0EIiqdX1Jyen7HlCCoVA815BSkit5GMvktDwXak+7Zwo2Xt7GbuBGea4wsHd/02yn03mbzz4zAgTnAQkcHC/GghGJw4v60mhC8LyuJFsyFr0Jnddu7KLo+Xz55ChPT9j889PD3Laug41pk1HHY3XSRHo+NGkTVStci4USDIViAXF9ievLJUv9nC+GDx3jxF1745+P58p8/PAgvoT/emkvz3/0MMBzTjTsbL7pc3cPZPElXN+f4fldSf7PU8McHCvxi+s78XzJY6NFAK7sSgJBlth7dq6K39+Fzms3dPF3RwY5ki1zJDuIIcCVsC5lcsu6dq7uSS/sF5whSjAUigWkaF88fYhO3ruPoQNHCAMu9F11GZtvvr5KLPKuz788PczRicmg75dOjHJFR4KhRw/TtmH1c8o9ZXVk6kQj63g8OFTg8bEShoBXre8gpWtYmuB0wWHM9jieK5NzfVYnDdZP0dtrc5vFH+1cxXfOZnlktIgbep8Gig6fOT7C6qTJ+rTJsYkyP7qQ49L2BC9d1Yaeqk/PXQyUYCgUC0jR8efULnyxOHnvPoZCKwEAKRl69DBDjx2BsMDQk5J/OznK0YkyCU1wTW+aYxNlzpVcHhgqcOPqNgb27n9OCca6G6+JxfJbZ8fZN5gn60x22X3jlm46zCAV7LKOBE+MlfjPwRzfPRdkRl3fnwmKE5MWumVWxX1yZ84zdOAIa1Imv7mth5dO2ORcj23tCb58cozHRovcO5DlbZf2cvdAlsPZMg+PFNnVleLqV1+7+IuBEgyFYkEpOn5dquVyZOjAkfjfR7MlRmyPF/emERUdyL9wYpSHR4oI4N1X9LM+bbF/uMCnj4/w6GggGFO5cJaS2Qbqe3dt49Q9/4lTdvjeuRxORfv4XkvnxRU1FVd0JnlirMTdAxNAICA/v7oNoWtsfMV1ddfr3bUtduEd+eLdbD/5bPzc6zd18vhokcfGimQdj7MFJ37ujO3zCpVWq1A8t/B8ietdJPGLMJCdczz+7sgQAD2Wzra2BA8OFxAiyAgCeNOWbtang8Durs4kmoBj4em4zdAZPnRs0a2MqQSh1nqys3lO3n0/MH1q6vChY/iux9mig+NLLE2wpzvFT4cL3LKuo+q1l2Ssqp+v7EqhCcHLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"xx = torch.arange(data.shape[0])\n",
|
||||
"\n",
|
||||
"plt.plot(xx[1:], means)\n",
|
||||
"plt.fill_between(xx[1:], means - 2*np.sqrt(vrs), means + 2*np.sqrt(vrs), color=palette[1], alpha=0.5)\n",
|
||||
"plt.scatter(xx, data, color=palette[5])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "6be11181-0698-429c-9dc7-a0f4b6c99584",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Rollouts"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 42,
|
||||
"id": "81bb90d7-e790-4e31-ba87-fe61449eb9ce",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"nroll = 50\n",
|
||||
"roll_len = 100\n",
|
||||
"xin, xout = dset[len(dset)-1]\n",
|
||||
"xx = torch.cat((xin[0, 1:], xout.unsqueeze(0)))\n",
|
||||
"xx = xx.repeat(nroll, 1).unsqueeze(1)\n",
|
||||
"roll_pxs = torch.zeros(nroll, roll_len)\n",
|
||||
"with torch.no_grad():\n",
|
||||
" for idx in range(roll_len):\n",
|
||||
" out = model(xx)\n",
|
||||
" smpl = torch.normal(out[:, 0], out[:, 1])\n",
|
||||
" roll_pxs[:, idx] = smpl\n",
|
||||
" xx = torch.cat((xx[..., 1:], smpl.unsqueeze(-1).unsqueeze(-1)), -1)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 43,
|
||||
"id": "1aec9c78-083b-4161-af02-4ef9e6590858",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYcAAAEFCAYAAAAIZiutAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAABZ9UlEQVR4nO29d3wc9Z3//5ztVVr1bss2lhum2ARsDgIJJDHpCZeeu8uFSy4hR658c5dy973Aldwdxy+58r1CkguBhJAGoRy9xIADxmBhXIQt27LktWW1VdveZn5/zM5oV7uyLVvFlt7Px4MH1uzszme1q89r3l3RNE1DEARBEPKwzPcCBEEQhHMPEQdBEAShCBEHQRAEoQgRB0EQBKEIEQdBEAShCNt8L+BsSSQS7N27l5qaGqxW63wvRxAE4bwgm80yODjIhRdeiMvlKnr8vBeHvXv38qlPfWq+lyEIgnBecu+993LZZZcVHT/vxaGmpgbQ32B9ff08r0YQBOH8oK+vj0996lPmHjqZ814cDFdSfX09zc3N87waQRCE84up3PESkBYEQRCKEHEQBEEQihBxEARBEIoQcRAEQRCKEHEQBEEQijjvs5UEQTh/CbZ3svuhF0nHkwA4PC7Wv/8qWja0zfPKBBEHQRDmnMmiYJCKJWj/xXMA55VAaJqGoiinPO/YsWPnTcq9iIMgCLPKrgeep+eVDnMDrVreyMjRfrLpTMnztazKzp8+Q/vPnmXpFWu55MPXzPGKp+bNN9+krKyMbDZLNptl2bJlRCIR2tvbeetb32qed+LECRoaGoqe/8tf/pJbbrnlvGj1IzEHQRBmjV0PPE/39n0YAyc1TWPo8PECYfh5zwj/tK+f/kS64LmaptG9fR+7Hnh+Ttc8GU3TOHjwIAAHDx7k4MGDPPLII3R0dBCLxTh27BhjY2MAqKrKjh07ePLJJ4teR1VV4vE4R44cmdP1nykiDoIgzBrdr+w76eOxjMoLA1GCsTTfOxhCLTG1uHv7PoLtnbO1xFMSjUZpb29nfHyc48ePE41Gueaaa/B4PLz44ovs2bPHdClFo1H27dtHKBQiHo8XvM6bb74JQFdX15y/hzNB3EqCIMwKwfZOOMWE+gePjZn/7ktk2DuaYHWZE5tFwZLnw991/1ZAj0ME2zvpeGI78dEI7oCPtVs2zVp8QlVVTpw4QTwe59ChQ4RCIfx+P+FwGIfDQW9vL4qisGbNGlRVZWRkhOHhYZYsWcJ9993H7//+76MoCpqmsW3bNhwOBxaLhVQqhcPhOO1YhUE2m0VVVUZHR6fsiTRTiDgIgjAr7Hl4m/nvw+Ekvzw6SoPbzu8sq0BRFEaSGV4ajAJQ5bASSmX57qGQ+Zz3NJVxQ2MZANl0hp0/fYZQ9wmCOw+Ybqn4aKRAOGaa3t5eHnroIVatWsXg4KC5odfU1LBv3z4sFgvpdJrOzk5OnDiB0+nkggsuwGq1kk6nicVieL1e0uk0AwMDLF26FLfbzYEDBwgEAmzdupVrrrmGJUuWnHItmqbx6KOPUlZWRiwW493vfveMv998xK0kCMKMEWzv5Mlv3cODf/GfpGIJ8/hDx8YIxtLsCMUYSmb1c2N6jKHJbedjrRVFr/Xo8XFGU9mCY93b9xUFsrPpDLsfenGm3woAo6OjXHbZZTQ0NNDb24vH4yEUCmG1WvnMZz5DOp0mmUwyPDxMd3c3R48eZf369TgcDjweD0eOHGF4eJhoNIrT6aSurg5FUejo6OCpp54iEonw+OOPA7Bnzx5GRkZKrmNgYIDx8XH6+vpQVZXq6mp2795txnJmAxEHQTiHyN9cn/zWPfPqa58uwfZOdt2/lfhopOD4oXCSrkjK/PlYTP/38bguDqvKnLR6HRjOlSuqPdhzP+wbK/TbT0U6npyV39Xo6Chr1qyhoaGBwcFBc6hYNpvF7/cD+h29qqpomkY0GuXnP/85dXV1OBwOfvOb3/DEE0+wdetWAoEAFRUVHDlyhOPHj3P11VdTVlZGKKRbS8ePH2f//v1mcNvghRde4J577qG/vx+Px8PBgwex2+38+te/5sCBAzP+ng1EHAThHCDY3smDX/svdv70GQaHxni2L0xoaJz2nz973ghExxPbS6anPnyscLM7lrMYjhuWg8eOx2bhT1bX8Odra/mdZZX89tIAAA8FxzgcTnIsliKVVU95/dPh6aefPq3z3njjDY4fP86DDz5IdXU1qqrS2NhIIBDA7XaTSCSwWq04nU4sFguqqpoisWPHDtLpNIODg/T29tLf34+iKFRUVHDixAn8fj8rVqwgmUxitVoZGBggGo3y4osv0tHRQX9/PyMjI6TTaXbt2kU6nWb//v3U19cTCoXo6OhgbGzMzKKaDeY85tDV1WVG+Pfu3Ut3dzeapvGv//qvbNmyZa6XIwjzTrC9k50/fcb8+ec9o7QPxzkYTvKFldXsfujF86IgbLLFADCeznIkksKmwMdbK/jxkRGOxdJomsaRiF4At9TrAGCF32k+7/IqL1v7I5yIZ/jO/kEAyu0W3tNUzpU13tO+vkE4HGZsbAyPx8P+/ft529vehtVqnTIYnM1m6e/vJx6PMzIywvHjxykvL6empoZkMklzczPxeBy73U44HEZVVdxuN6qqUlZWhqIo7N+/H03T8Hq9+P1+MxDt9/tpbm5mcHCQ48ePA/DKK68QDodJJBK8/vrrZDIZBgcHGR0dxWq1oqoqwWCQtrY2MpkMoVAITdM4ceLEaXwyZ8aci8N9993HPffcM9eXFYRzglKZNsYdb1bT+N9j47QP666UvaMJBhMZZjcnZeZwB3xFG/QLAxE0YFWZixU+ffMPxlKMpLKMpVU8VoVaV/E2ZLco/NnqWv79wCBHcxbGWFrl5z0jrA+4cFktdIwlaHTbqck93x3wTbm2rq4uDhw4wIkTJ1i2bBnPPPMMiqLwzne+s0ggEokE//u//0sgEOAtb3kLzz77LE899RSVlZV4vV5SqRSVlZWMjo6STqexWCyUl5czPj6Ooiik02lGRkZMSyIUClFbW0tNTQ3t7e0oisKRI0ewWq3U1tYyMDDAnj17qK2txWazMTo6SiKRoK2tjZdeeolEIoGmaYyPjzM2Nobb7SYcDgO622u2mHO3UltbGzfddBPf+c53ePrpp7n88svnegmCMC/seuB5dv70GXMDNTJtjJ9fGozydF+44Dm7R3WhOB9cS2u3bAIgq2q0D8f44eEQT/Tq7+e6eh9VTisui8J4WuXR4+MALPc5sSiKuUG7Az42fvx6Nn78erxuO3+yppaPLQ3wByuqaPHYyWiwIxTjpcEo3zsU4rY9fUQyWax2m3n9UoRCIS6//HIuuugiMpkMkUiE5uZment7i87dt28f1113HW63m66uLlwuF6FQiPLycjweD+FwGK/XyxtvvIGmaVRWVnLBBReQzWax2WxkMhnS6bT5nqxWKx0dHXR0dDA4OIimaQwPD5uv6XA4yGQypkvKYrGwZ88edu7cicfjIZPJkM3qgfnu7m5sNht2u33mPrgpmHPL4SMf+chcX1IQ5p1geyfd24sLwsyUzKzKMyf0jfSDzeVUOKzc1TXMm2MJrqv30/HE9nPetdSyoY2dP3uGHaEY93ZPZN00e+y0lbkAPb5wOJLilVAMgHc1+lGsFjZ85O1F769lQxu7HngeR+73ltE0fpj7nVQ7J7auYc3CNTdeS8uGNrq6uli+fDmgu5Jee+01rr32WmKxGEeOHKG2tpaxsTG8Xi9Llixh//799Pb2ctFFF+F06pbN6OgoFRUVjI+Ps3//fiorKwF9FHEqlTI3/q6uLjKZDPF4HLfbjd1u5z3veQ9bt24lHA6zdOlSuru7SaX0APzQ0BA2m414PE42m6W3t5clS5ZQUVFBf38/qVTKFIeWlhaOHDmC3W7HYrHgcrnwer2EQiEzEyqRSJRs0TFTSJ2DIMwB+cHS0VSWR46NcX2Dnwa3fgf4+PFxQqksHquFq2u9RDN68DWY88+fzJ8+m0yn4CzY3onFZjUDzgZe24SDotZl43Auc8lrs7DM52TpW9ZM+ZqXfPgaU1TbyvTN+3AkVVBbN5JRzed3dHSY4tDe3s7hw4dJpVLE43F6e3tZsWIFXq+XNWvW4PP5iEQijI2NUVFRQUVFBVVVVVgs+nrtdjvxeByLxUJzczPV1dXs2LGD8fFxMpkM4+PjZuZSIpFg8+bN7N6923T59PT0YLPZzLt+ox+T0VcpmUxy+PBhACwWC/F4HKvVitfrZWxsjGQySTKZ5NJLLyUYDBKLxcwAuCE4VVVVp/oIzxjJVhKEOSB/c7/3yDCvhGJ892CItKrxWijGrhHdffSRpQGcVgsVDitem4VoRmUkl+s/166lyamphhus1DqMc9V0lsGkbg2t9DtR0IvZDGry4gt1uX/37+856TocHt3qKLNbaXDbSKsaB8YnurmeGIma/x4cHCSTyXD06FF27NiB1Wpl7969hMNhYrEYiqJgt9tpbGw03T6RSISuri527txZUDcQj8fxer3YbDYuu+wybDYbbrcbl8vF2NiYufFns1mGh4cZHh7m+PHjpNO6OKqqisPhKHo/RlW0USORSqXIZrMkk0nT5ZTvZhoZGWHdunUArFq1yrQaFEUxhWg2EHEQhDnACJZqmsabuY1tMJnhvw8O8cOuYYZzArCuXN8IFUWhxaNbFcad+Ommak6XqWor9jy8jWw6w3g6y7aBCFlVI5vOlFxHfhrrYEL//0eWBPiXjU0s901kIdXkuYPqXfr7O5VVtP79V5n/XuV3FT0+nMoWCFZXVxcvvPACmqaZPv5YLEZdXR3t7e0MDQ0xPq7HPBRFwWKxMDQ0REdHB+Pj42b6aSQSYc2aNdjtdpYsWcLRo0ex2+0sW7aMaDSKx+NBVVVcLhfd3d309PQQi8VMgXE4HDQ2Nhat19j0DQKBABaLxRQawKx1cDqdhEIhqqqqiMfjrF271kyLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"full_x = torch.arange(data.shape[0])\n",
|
||||
"test_x = torch.arange(data.shape[0], data.shape[0] + roll_len)\n",
|
||||
"plt.plot(full_x[1:], means)\n",
|
||||
"plt.scatter(full_x, data, color=palette[5])\n",
|
||||
"plt.plot(test_x, roll_pxs[:20, :].T.detach(), color='gray', alpha=1., lw=0.5)\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "0b7e9075-855e-4c81-8334-8507718152f1",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Just playing"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 44,
|
||||
"id": "83db2178-1011-4e7b-9070-8bb5955f943f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import torch\n",
|
||||
"import gpytorch\n",
|
||||
"import sys\n",
|
||||
"\n",
|
||||
"sys.path.append(\"../\")\n",
|
||||
"from voltron.likelihoods import VolatilityGaussianLikelihood\n",
|
||||
"from voltron.models import SingleTaskVariationalGP\n",
|
||||
"from voltron.kernels import BMKernel, VolatilityKernel, FBMKernel\n",
|
||||
"from voltron.models import BMGP, VoltronGP, MaternGP, SMGP, VoltMagpie\n",
|
||||
"from voltron.means import LogLinearMean, EWMAMean, DEWMAMean, TEWMAMean\n",
|
||||
"from gpytorch.kernels import ScaleKernel, RBFKernel, MaternKernel\n",
|
||||
"from gpytorch.means import ConstantMean\n",
|
||||
"from voltron.train_utils import LearnGPCV, TrainVolModel, TrainVoltMagpieModel\n",
|
||||
"from voltron.train_utils import TrainBasicModel, TrainVoltModel\n",
|
||||
"from voltron.rollout_utils import GeneratePrediction\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 45,
|
||||
"id": "c9c02c1e-c7ec-4719-9178-cacbb7c628a4",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"train_y = torch.FloatTensor(pxs).log()\n",
|
||||
"train_x = (torch.arange(train_y.numel())/252.)[1:]\n",
|
||||
"test_x = torch.arange(100)/252. + train_x[-1]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 46,
|
||||
"id": "dbe354ea-dc9d-479f-8025-48a1ca9cae22",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"vol = LearnGPCV(train_x, train_y, train_iters=200,\n",
|
||||
" printing=False)\n",
|
||||
"vmod, vlh = TrainVolModel(train_x, vol, \n",
|
||||
" train_iters=200, printing=False)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 47,
|
||||
"id": "b89d9a80-cf7f-4bf7-9cbf-59ee9a0563f6",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"voltron_lh = gpytorch.likelihoods.GaussianLikelihood()\n",
|
||||
"voltron = VoltronGP(train_x, train_y[1:], voltron_lh, vol)\n",
|
||||
"voltron.mean_module = ConstantMean()\n",
|
||||
"\n",
|
||||
"voltron.likelihood.raw_noise.data = torch.tensor([1e-5])\n",
|
||||
"voltron.vol_lh = vlh\n",
|
||||
"voltron.vol_model = vmod"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 48,
|
||||
"id": "50591bec-1240-4a8e-81b9-46b8295e642b",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Iter 1/200 - Loss: 7.323\n",
|
||||
"Iter 51/200 - Loss: 1.576\n",
|
||||
"Iter 101/200 - Loss: 0.906\n",
|
||||
"Iter 151/200 - Loss: -0.836\n",
|
||||
"Iter 201/200 - Loss: -1.960\n",
|
||||
"Iter 251/200 - Loss: -1.965\n",
|
||||
"Iter 301/200 - Loss: -1.965\n",
|
||||
"Iter 351/200 - Loss: -1.965\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"voltron.train();\n",
|
||||
"voltron_lh.train();\n",
|
||||
"voltron.vol_lh.train();\n",
|
||||
"voltron.vol_model.train();\n",
|
||||
"\n",
|
||||
"# Use the adam optimizer\n",
|
||||
"optimizer = torch.optim.Adam([\n",
|
||||
" {'params': voltron.parameters()}, # Includes GaussianLikelihood parameters\n",
|
||||
"], lr=0.1)\n",
|
||||
"\n",
|
||||
"# \"Loss\" for GPs - the marginal log likelihood\n",
|
||||
"mll = gpytorch.mlls.ExactMarginalLogLikelihood(voltron_lh, voltron)\n",
|
||||
"\n",
|
||||
"print_every = 50\n",
|
||||
"for i in range(400):\n",
|
||||
" # Zero gradients from previous iteration\n",
|
||||
" optimizer.zero_grad()\n",
|
||||
" # Output from model\n",
|
||||
" output = voltron(train_x)\n",
|
||||
" # Calc loss and backprop gradients\n",
|
||||
" loss = -mll(output, train_y[1:])\n",
|
||||
" loss.backward()\n",
|
||||
" if True:\n",
|
||||
" if i % print_every == 0:\n",
|
||||
" print('Iter %d/%d - Loss: %.3f' % (i + 1, 200, loss.item()))\n",
|
||||
" optimizer.step()\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "2bbbe667-d72d-4ece-898e-727ca1069d86",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 49,
|
||||
"id": "de189faa-0fea-4f7e-b7b5-2d8a795ffceb",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"voltron.eval();\n",
|
||||
"voltron_lh.eval();"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 50,
|
||||
"id": "91d4bf03-5636-45f1-8fdb-3cd6b30bf61f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"pred_vols = vmod(test_x).sample(torch.Size((10,)))\n",
|
||||
"samples = GeneratePrediction(train_x, train_y[1:], test_x, pred_vols.exp(), voltron)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 51,
|
||||
"id": "79430340-faf5-4693-91cb-53556ad9276b",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAY8AAAEFCAYAAAAbsWtZAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAABJj0lEQVR4nO3dd3hUZdrA4d9M2qT3ShJIgCT0jhRBBUR0AQuWZRHs6Opadj+liCtgWyyA2FhURBYB6wK6IEgRRDpSQ0IJJCSk996mfH9MZjKTQjLp5bmvy8uTc868553hzHnm7QqdTqdDCCGEsICytTMghBCi/ZHgIYQQwmISPIQQQlhMgocQQgiLSfAQQghhMevWzkBTKikpITIyEm9vb6ysrFo7O0II0S5oNBrS09Pp27cvKpWqXq/pUMEjMjKSGTNmtHY2hBCiXVq/fj1Dhw6t17kdKnh4e3sD+g/Az8+vlXMjhBDtQ0pKCjNmzDA+Q+ujQwUPQ1WVn58fgYGBrZwbIYRoXyyp7pcGcyGEEBaT4CGEEMJiEjyEEEJYTIKHEEIIi0nwEEIIYTEJHkKINkGt0XAlIaO1syHqSYKHEKJNWP3DQcbOWkbEnxYTFZPc2tkRdZDgIYRodTsPRvP6yp8BKCgq5e3VvwCQkV1AbGJma2ZN1KJDDRIUQrRtOp2OlV/v58q1DF7722Qc7G0BeGTBOrPzdh++wIg/v0NSei5arY4tHz3FkD7BrZFlUQspeQghWsyWPWd469PtfL3tONt/jwKgpKy8xnOvpeag1epXyT58JrbF8ijqR4KHEKJFaLVavvrpqPHvlIw8AC7Hpxv32atsanztvz7dwdOvbaS0TN28mRT1JsFDCNEiNmw9zuHTlSWIlIxcAKKvpALg4erAzs+fMx6/59aBbPrgSePfP/56lt2Hz1dLd8ue0xw9G9dMuRa1keAhhGgRe49dBGBw7yCgsuRxIVYfPB6+eyTdungaz3dQ2RIe4muWxtuf/8LBU1eMfyem5vDM699wz3OfotZomjX/wpwEDyFEszl06goJKdkcPHWF7fv1bRx/vkO/XkRqZj4A52NTAIioCBSG4w/fPQIXJxVfvjWLYH93AC4nZDDjpTUkpuYAkJyea7zW6fOJzf+GhJEEDyHaqLSsfNZsOkR2XlFrZ6VB4pOzuO/vnzNy+ru8/fkvxv0jB4YC+pLHoVNX+PWIvkQSEaJfg+ed/7uLs5tfMf49YWQEBze8xP9WPo2nmyPlag1Pv/41pWVqUjLzjOkeMimRdGQbo84ydO2nnMtIa9V8SPAQoo3JzCngT3/9hOH3v80/P/iJF9/9b2tnqUFir1WOz/jjXDwAi/82mQBvV+xsrUlMzWHpl7uN53QN8ABAqVTi7upQLb2BEYH89MlfcXdx4I9z8ew8GE1aRekFID27oLneSqtTa7UkF+Tz/YUoZm3dxOm0FF769Ze6X9iMJHgI0cb8cuA8p89fQ63RArDj9yjOXmx/VTKZOYXV9v35jiHY2VobSx+GBvSP//kAVlZ1P46C/T14evpYAHYcMA8e7amEduqHvez76AfUtXRTruqV33YTvHI503/83rhvZ9wVen76Afdu+oZr+XnXeXXzkOAhRCspLC7lk42/GevvDc5cvAbA3/5yE0/efyMAH2/Y19LZa7T07Pxq+xzt7QAYPSjUbP+Q3vUfAHjbjb0BfVC9bDIXVk5ecUOy2SJ0Oh0ZFSUjjVpD3JEosuNTeXvFl8xZu4GycjX7/rWOTY8vIS0qrtpr1507Y/z7vVsmMiJAv1JqXG4OW2Iu8NqBlr8/ZIS5EK1k9sIN7Dt2idMXrrFq0V8A2HPkAut+1I+FGDO0B/5erqz69neOVVT7tCdpWbVXI/l6uhi3FQoFvl4utZ5bVWigFyMGhHD4dCw/7z9n3N8SJY+sqylc/v0M/abeiMq5etVabT7//gCLP9nGvxdN58Yefsb9fdLLIT2HdSu/JueznwC4sucEYX8ZT++/3UU2Gs6mp5JWVIivgyOXn3weO2trXOzsOJx0zZjOmrMn6evtw3NDbmi6N1sHCR5CtILSMjX7jl0CYN+xSxQUlbL8P3s4YjIOYmBEIPZ2NtirbEjNyCM7rwh3l/o/sFpbxnWCh6ebo3Hb19MZG+v6r50NcOe4/mZjRgBy8hsWPHKTMijMysO/dwgKpeK65/72sb79SVOm5oaHb0ehqPl8tUbDe1/sori0nKenj2XxJ9sAeGrRRh4K9WCYp/m/oza6slrSxsmemO/28cemvaS5q4gNcIBhvjzSfxB21vpH9kN9B+Jmp2J811D+sWcHayNPMX/fLu4P74Ofk1ODPgdLSbWVEK0gPjnLuG2lVLDww/+x6pv9nDqv/zX5+esP4mhvh1KpJKybvgvr+SsprZJXg6S0HB6c+2W1h3ZtDNVWL8+ehMrOhpdnTzIeMw0eXXzdLM7LqIGh1fY1pNoq5rdT/Pr+txz9z3bi/ziPVqNFp9PVeG5ucmUHgJToOHa+sRZNmZrC9Bzifjtl9rof95zlow37WP3DQYbcu8QsnY1x2eSVm49JcS7TYm1vx8TlzxI86Qa6TR1NaO8e9CpUMOF4On/972VmFDmi0+koLylDqVBwd1gvXOzs+Pz2qTwzeDi2SivKtC031kVKHkK0AtOeSLkFJXyz/Q+z4xEmg+Miuvly+vw1Ll1NMzY0t4ZnXv+GY5FXOXY2jgvbFtV5flzFbLjjR4TzxH2jzUoXpsGjW4BntdfWJTTIC1cnFbkFJcZ9uQXFaLValErz38QleYVY2dpgo7I17suMS+birydIjb5q3Hfyu185+d2v2DmocHJU0fXG/pTkFKByc8K3Tyh7l34NSgVKTSleOafIKQtj+yufkHxsG3s1nmRb+eHUtydKO1uOn7tKVcN6+HEsJoUyrY6XTyVze14mI3xd8OwbgsrTlW4zxnP5wFlAX5XnEuKPS4g/1/acoJ+TE7/MWcmWZ5ah9HXlqe//hZOvO4nHL+Ds78mDcWrm3foAfi6uFn+WDSXBQ4hWEJt4/UWPAv3cjNshgfqHa1xiVi1nNx2NRsv76/bQM9iHqeP6mx07Fql/IBYWl9WZTm5BMfHJ2djZWtM92AtrK/NqKQ9Xk+DRxcPifCoUCjzcHI3Bw8VRRV5hCTn5xWZp5yZlsO+jH1AoFPS/awxOPYP4aP0+rp64SGlBMV52VjwwtjdpFxP071+n42RyDvlXkvBauRnQEeycwtfuPfG01/CUyzUOZ5SQaqVmstUWjmgCeYdJFFlVBKbIyrYpD62G4TmZbHf3xqu8jNvzMpnY15c3I/Uj6n928eT3/HIezi0h3FWFLqey2i3dSou3Rh8EQ4f5Y6csIWvQMOboEkjxUKFZuZ6gsymknI4BwMHfk6gdR7lz2bO4Bnpb/Hk2hAQPIVqIVqtl0cdbSUjJJq/ioVf11zPA0D7BZg9bw5QdcS2wrsXGrcdYvnYPAFNu6Wes04+JrxyQZmdb92PjXMViThGhftUCR9U0TEshlnB1sjduR4T6cvTsVT78ai8FxaU8fNsgyi9eJe5oNFRUJ5387le2F8P/Iq+ZpVN2LpmBAZ7sP3qJgxlF6EOjLQP9A0iwVqCy8iFR5wolcKLEhRScQQs/qntxyD8Qja0ChzgNXRXZXHapGKtSBnMiAinT+dM/PgNnVPj5WaGxKucVq2I+LLEi18aWfGsbPryYwRsRDgRYF1Nk34WfVcU89ecpxK/ZQWjCd3jl6ntaLR74N6556NtK3tem8n+nL5HmaUdcTy/u9g/Fwcqa4rxCXJHgIUSHsv5/x/jiv4fM9s26cwQfrt8LwOP3juau8QPoEWz+5TcEjxPR8eh0ulobaZvCDztPGbfjk7ONA/c2bqusVistU3PkTCwqOxsGhAfWmI5hptyIbr41HjcVEujVoLxGhPoZ24gev3c0R89e5bPvD+Bmo2BG0QZy1CGgCyFHpcTe2hrnzGQOni8ElPiprEFpRUpRKT+diuOnU9XTP2Vb0aitqwxSKTgbt4+6dCF3oL7E4RVuT45CSZZS/7dbWTGjDy5CpdWQqHKhS0ke6AfS00uhoo+NG5utevCVwwDy+9rwUU4Sbo4lRDupCbKzYfjWrxnhksP6vEgAyhRKtjiWA/r0M13tePVJfZdltdKKzeTwpEMAfwoPatBn2RASPIRoIYa5nUyNuyEcR3tbNm47zt0TBtT4MDYEj4zsQj777gCzK8Z+NDWtVkvU5crlX89eSjQGj5ir5lNhTHv+M2ysrTj89Uv4eDjz+4nL9O7uh6ebvqePofHa4zqlii/fmsXpC9e4cXD3BuX35dm3kZlTwEN3jiCsm49xf6BjCos8R/NbeRnhyljGlZ1j8JVMUuI1FAXejb1tEf8YNZC0nt1xT0nmdFwppy+dIzkznYkOEDyyC4d+vERysZICnS3dVWW4uwfT29uZ3WpbrsSnEjowkAPOKYB+kF+clRrDgx0gx9ae+wY/wNcnv9MHDhMO1iX0sE7hL3YFfHbDIMqtlZx2r+y+a/ikD6vcGDvyUW7Iz6BbWREZdo6EFWTwl6SzLAq7BbXSvES3qiiJxzPSGOjr36DP01ISPIRoIaaT+Bl4ujvytxk387cZN9f6OicHOyaMjGDXofMcOn2l2YJHfHK2WXvGoVOxTL6pHwDJ6foH4EN33sDaLUcAKFdr+HLTYfy9XXj5/R8JD/Fl9xfPA5BToA8ers721GbCyAgmjIxocH49XB1Z8+YsQB/4VHY2lJSWo/HqztYifU+v4zhznBEQgv4/NIAdD6VFQloks/uOwqq3Lb9p1BDqziWAvCKe6l/KishtFPYbT+ygKVy7qsF3RBfu6TWc5IJi7ti0msLycnp5ejGj0ImFxXHY6RQ8l+/MTlUxf9iVc9rFj5tGzsZfDTPTUxnW6wbsQwLw+fyvXPII4v4+kylWm4ww1+pAqSAsPYOL3vrSWKLKhf+qKsfA3JV2iUevnUSLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(train_x, train_y[1:])\n",
|
||||
"plt.plot(test_x, samples.T.detach());"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "7faf7644-ea71-442c-a5d1-15f197d0517c",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Check Volt + Constant Preds"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 52,
|
||||
"id": "088cecfb-8a26-4d08-bc87-906bab10004c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"preds = torch.load(\"./saved-outputs/AMZN/volt_constant2020-01-09.pt\").cpu()\n",
|
||||
"# preds = torch.load(\"../trading/saved-outputs/AMZN/volt_dewma100_2019-12-18.pt\").cpu()\n",
|
||||
"\n",
|
||||
"tckr = \"AMZN\"\n",
|
||||
"ntrain = 400\n",
|
||||
"lookback = 1\n",
|
||||
"data = GetStockHistory(tckr, end_date=\"2019-12-18\", \n",
|
||||
" history=ntrain + lookback).Close.to_numpy()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 53,
|
||||
"id": "df65a48d-ad1a-41a1-8122-d1c2d6087103",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"train_x = torch.arange(data.shape[0])\n",
|
||||
"test_x = torch.arange(preds.shape[-1]) + train_x[-1]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 54,
|
||||
"id": "74fc106c-b377-4724-8391-6534499de38c",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAZ4AAAEFCAYAAADT3YGPAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAABijUlEQVR4nO2deZgcZbX/v9V79/Qy+5pZk5msE7KRkASEiwghkusliBCDkUVQUS9eBQHxSuL1Khe5bNcFjJBgQvJD0bCIEBEQEEImELJNJjPZZt969ume3rt+f9S8b1f1MntPzyTn8zw89FRXVddMd/pb57znfI8giqIIgiAIgpgkVIm+AIIgCOL8goSHIAiCmFRIeAiCIIhJhYSHIAiCmFRIeAiCIIhJRZPoC5jquN1uHDt2DBkZGVCr1Ym+HIIgiGlBIBCA3W7HggULYDAYFM+R8AzDsWPHsHHjxkRfBkEQxLTk+eefx7JlyxTbSHiGISMjA4D0x8vOzk7w1RAEQUwPWltbsXHjRv4dKmdEwuPz+fDxxx/j3XffxcGDB9Hc3Iyenh6kpKRg8eLF2LhxI1asWBFx3H333Yc9e/bEPG9xcTHeeOONmM+/+uqr2L17N6qrqxEMBlFcXIzrrrsOGzZsgEoVe3lqrMdFg6XXsrOzMWPGjFEdSxAEcb4TbYliRMJz4MAB3HLLLQCkCGD+/PkwGo04ffo09u7di7179+LOO+/EXXfdFfX4JUuWoLCwMGJ7NCVkbNmyBbt27YJer8fKlSuh0Wiwb98+/OQnP8G+ffvwxBNPRP2FxnocQRAEMTmMSHgEQcBVV12FTZs2ReTq/vrXv+Luu+/Gr3/9a6xYsQIXXXRRxPHXX3891q9fP+KL2rt3L3bt2oWMjAzs3LkTRUVFAICOjg5s2rQJb775Jnbu3ImvfvWrE3IcQRAEMXmMKO+0cuVKPPnkkxGiAwBr167FtddeCwB45ZVXJuSinn76aQDA3XffzcUDANLT07F582YAwNatWxEMBifkOIIgCGLymJA+nnnz5gEA2traxn2u1tZWVFZWQqvVYs2aNRHPL1++HFlZWbDb7Th06NC4jyMIgiAmlwmpaqutrQUQe81m//79qK6uxsDAANLS0rB06VKsXr066kL/8ePHAQClpaURtd+M8vJytLW1oaqqCkuWLBnXcQRBEMTkMm7hsdvtvHLtyiuvjLrPSy+9FLFt1qxZePTRRzF79mzF9sbGRgBAbm5uzNfMyclR7Due4wiCIM4nRFGEy+WC0WiEIAgJuYZxpdr8fj/uuece9Pf3Y+XKlbj88ssVz8+ZMwc/+tGP8Nprr+HTTz/F+++/j6effhpz5szBqVOncMstt0Sk5wYGBgAARqMx5usmJSUBAJxO57iPIwiCOJ9ob29HRUUF6uvrE3YN4xKeBx98EPv27UNOTg5+8YtfRDx/88034ytf+QpmzZoFk8mEzMxMXHbZZfjjH/+IRYsWobOzkxcEMNhcutEq8ViPIwiCOJ84ffo0AODs2bMJu4YxC89Pf/pTvPjii8jIyMD27duH7MkJR6fT4Y477gAAvPvuu4rnWFTCIphosIiF7Tue4wiCIKYboiiiuroaTU1Nozqmo6ND0cfo9/vh9XrR398fj8uMyZjWeB566CHs2LEDqamp2L59u6J0eaSUlJQAiKyEy8vLAwA0NzfHPLa1tVWx73iOIwiCmG60t7ejpaUFwMi/z+x2Oy/CYjgcDhw+fBiiKGL58uUwmUwTfq3RGHXE8/DDD2Pbtm1ITk7Gtm3bMGvWrDG9cE9PD4DI6IOVZp88eRJutzvqsUePHgUAzJ07d9zHEQRBTDc6OjpGfQz7zpUzMDDAlylcLtd4L2vEjEp4HnnkETzzzDOw2WzYtm0b5syZM+YXfv311wEACxYsUGzPycnB/Pnz4fP5ovq4VVRUoLW1FRkZGVi8ePG4jyMIgphueDwe/nikDfHRrMJqamr4Y/n6uCiK6OzshM/nG8dVxmbEwvP4449j69atsFqtePbZZ3mEEYuqqiq88847CAQCiu1+vx/btm3Djh07AEgFCOGw9Z9HHnkEdXV1fHtnZye2bNkCALj99tsj+oDGehxBEMR0wu/3R3081P4NDQ0jPmd7ezuOHj2Kw4cPj/0ih2BEazxvvfUWfvOb3wAACgoKsHPnzqj7lZSU8C//pqYmfOtb30JycjKKioqQlZUFp9OJmpoatLe3Q6VS4e6778Yll1wScZ41a9Zgw4YN2L17N9atW4dVq1Zxs0+Hw4ErrrgCN91004QdRxAEMZ2QRyLhN/fROHnypOLn/Pz8CCGSn4el8hwOx3guMyYjEp7e3l7++NixYzh27FjU/ZYvX86FZ/bs2di0aROOHj2KpqYmHD9+HIIgIDs7G+vXr8fGjRsj0mxyNm/ejKVLl+L5559HRUUFgsEgSkpKhh1vMNbjCIIgpgOiKI464uns7OSPrVYriouLoVKpFJkhufDE+3tyRMKzfv36UblLA5KiPvDAA2O6KMa6deuwbt26STuOIAhiqhMIBHhBADAy4TGZTOjr6wMgNfarVCpotVrFPvLzyM8fCAQmfJQM3f4TBEFME3p7e3HgwAHFtpEIj1w4mJelRqOMO+QRj7x4IVaV8Hgg4SEIgpgmnDx5UiEKAFBZWTlsI6nP54MoisjIyOCRT3jEIxcer9fLH8djjAwJD0EQxDQhVsorvHggHL/fD5fLhebmZhw6dAhOpzMi4pFHTkzc5s6dC7PZPM6rjoSEhyAIYpoQLhbBYHDE5dSBQIAXDfT19cWMeAKBAILBIARBQGZmZlz8L0l4CIIgpglMZJKTk1FcXIzGxkY0NjYOWVLNquCCwSAXHo/HE3ONh5Vqa7XauJkuT8ggOIIgCCL+MOGZOXMmDAYDX3+J5jAgiiIEQeDHCILAhcTr9UKr1cJqtcLtdsPr9fJzyIUnXlDEQxAEMU1gUYlGo4FareYVaqIoKkqgT506hX379sHn83EhkUcvHo8HgiBg8eLFWLp0KQCQ8BAEQRCRsOhFo9HA7/cjKysLgiDwdRlGY2MjvF4v2tvbufuAXEhY8YAgCHw7q3ybDOGhVBtBEMQ0QO5YoFar4XK5IAgCVCoVLzIIr3oTRZHP2tHpdHy7vCRbpVJBrVYjEAjg5MmTfLQMRTwEQRDnOSyiUalUUKlUXIRUKhUCgUDM6rZoEY/P51MUJDBRks8zCy8+mEhIeAiCIKYooiiirq4OnZ2dqK6uhsPh4IIRLjxsu3ytJxgMRl3jAZSOBNGiG6vVOrG/jAxKtREEQUxRuru7cfbsWXR0dMBisaCjo4ObNssFhUU8TqcTjY2NvHTa7/dzgWIRk9FohMvlwoEDB3DJJZdArVYr0nAAsGzZsrg0jjJIeAiCIBIMK30Oh4nGwMAAr2DT6/U4cuQIj2zUajV3Jjh16hT6+vrQ1dWFzMxMLjxsfYil6Rgulwtmszki4jEajfH6VQFQqo0gCCKhtLe34/3334fdbo/6fDAY5P8BgNlsRldXF7q7uwGA9/N0dXVhYGAAwWAQAwMDAEJrOX6/H4IgQK/XIy0tjZ9bXp7d3d2Nrq4uiKIY97EIJDwEQRAJhKXGKisrIww55amyQCCAnJwcHvkwjEYjAoEAent7I5pKPR4PRFFEIBCAIAgwmUwoLCxUnJ/1APX29qKvr4/vG09IeAiCIBKIvAQ6fOKnXHhEUYRer496PCun9vl8XHjcbjd3opav76jVamRkZPDzHzlyBMePH+fni8cYhHBIeAiCIBKIvKdGPu0ZgKJazeVyRZRMWywWaDQaCIKAYDCoaCSVn5c9Zms3rFS6paUF3d3digiHhIcgCOIcZyjh8Xg8fGy1KIo8ggGA0tJSLFy4EBqNBiqVCj09PQBC0Y28rJpVwLGIiUVZ7Bj5mo7X68WRI0fiMoeHQcJDEASRINi4AkZ/fz/q6+tRVVWF/v5+HDt2DEAoRSaPeFg1msfjgUqlQl9fH0RR5IJhs9l4hMNegwlPeHOoXHhUKhW6urrw8ccfo6KiQjEUbqKgcmqCIIgEwSIYk8kEt9sNj8eDM2fOAJAiDyYixcXFqK+vVwiPvPeGCYfP5+PbRVGEzWaDy+UaVnhYqs1ms/HiBVYZN9TIhbFCEQ9BEESCaGlpAQBkZ2dHNGx2d3cjGAzCZDIhJyeHV6cBklAwgcnOzlYIDxMXURQxc+ZMZGVl8WICdky4pxs73mQyKcqtLRZLXHp6SHgIgiAShNPpBACkpqYiKSkp4vlgMIikpCReJh0MBjFjxgyUl5dz8SgtLYXNZgMgRUnyiEej0fBqOK1WywVGHvEwo1FAEia5VY5chCYSEh6CIIgEwQoL9Hp9RGQRCATg8/mg0WhgsVi48KSlpSE1NZXvp1arkZycDADcoZqlzgKBAFwuFwBwcQKUabqysjIUFBTwMddysYmXXxut8RAEQSQAVligUqmg0WhgMpn4c2q1Gv39/QgGg0hJSYHBYOCNnuFpMiBk8ul0OhXWOKy3BwBKSkr4/larFWVlZRBFEVlZWcjOzubO1DabDXq9Hh6Ph4SHIAjiXEIe7QiCoIh4SkpKUFFRAYfDgebmZj6kDVCWSTPkqbNgMAi1Wg21Wg2fz8cLEuR+bIIgIDc3V3GOwsJCiKIInU6HZcuWIRgMxm00AqXaCIIgEoBceABpYT8vLw8zZ85EdnY230+lUqGjo4P/XF1dHXGu0tJSGAwG5OTkAJAiJpVKBa/Xq5haOhTFxcU8KtJqtVFdEiYKingIgiASQLjwCIKA0tJS/nx6ejp6enqgUqlQXV0NnU4Xc+CbyWRCfn4+T6upVCoIgsDXd+TrPlMBingIgiASABOJ8JEEwWAQn376Kex2OwRB4Gs6qampCpdq+Xk6OzuxZMkSvk2r1UIQBN6LE89pomNhal0NQRDEeUIs4Tl16hQOHz7Mf2alzuz/gUBAMb/nk08+gdvtRnl5OT+GRVELine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(train_x, data)\n",
|
||||
"plt.plot(test_x, preds[:10, :].T.exp(), color='gray', alpha=0.5);\n",
|
||||
"# plt.ylim(1000, 5000)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "b94a72c1-92ab-4d5f-8f3d-b9498fd9ed15",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,393 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 18,
|
||||
"id": "92e5b691-e01a-438a-b2b9-543c5fa3f8b2",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import pickle as pkl\n",
|
||||
"import pandas as pd\n",
|
||||
"import numpy as np\n",
|
||||
"import torch\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns\n",
|
||||
"import os\n",
|
||||
"from voltron.data import make_ticker_list, GetStockHistory\n",
|
||||
"\n",
|
||||
"sns.set_style('white')\n",
|
||||
"palette = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"sns.set(palette = palette, font_scale=2.0, style=\"white\", rc={\"lines.linewidth\": 4.0})"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 19,
|
||||
"id": "1f59e48e-eb80-48b2-8ce5-91596ad136f2",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def ECDF(sample_pxs, true_px): \n",
|
||||
" return (torch.sum(sample_pxs < true_px, 0)/sample_pxs.shape[0])\n",
|
||||
" \n",
|
||||
"def Calibration(pcts, percentile=0.95):\n",
|
||||
" in_band = np.where((pcts < percentile))[0].shape[0]\n",
|
||||
" return in_band/pcts.shape[0]\n",
|
||||
"\n",
|
||||
"def GetCalibration(model, mean='ewma', k=100, horizon=np.arange(75,100), \n",
|
||||
" logger=[], exp=True):\n",
|
||||
" \n",
|
||||
"\n",
|
||||
" ntrain = 400\n",
|
||||
" n_test_times = 20\n",
|
||||
" ntest = 100\n",
|
||||
" pcts = torch.tensor([])\n",
|
||||
" for tckr in ticker_list[:5]:\n",
|
||||
" data = GetStockHistory(tckr, history=1000, end_date=end_date)\n",
|
||||
" for idx, date in enumerate(data.index):\n",
|
||||
"\n",
|
||||
" fpath = \"./saved-outputs/\"+ tckr + \"/\"\n",
|
||||
" fname = model + \"_\"\n",
|
||||
" if model == 'volt':\n",
|
||||
" fname += mean + str(k) + \"_\"\n",
|
||||
"\n",
|
||||
" fname += str(date.date()) + \".pt\"\n",
|
||||
"\n",
|
||||
"\n",
|
||||
" if os.path.exists(fpath + fname): \n",
|
||||
" preds = torch.load(fpath + fname).cpu()\n",
|
||||
" if isinstance(preds, tuple):\n",
|
||||
" preds = preds[0]\n",
|
||||
"\n",
|
||||
" preds = preds[:, horizon]\n",
|
||||
"\n",
|
||||
" test_y = torch.tensor(data.iloc[idx:idx+100].Close.to_numpy())\n",
|
||||
" if test_y.shape[0] == 100:\n",
|
||||
" if exp:\n",
|
||||
" preds = preds.exp()\n",
|
||||
" pcts = torch.cat((pcts, ECDF(preds, test_y[horizon])))\n",
|
||||
" \n",
|
||||
" if pcts.numel() == 0:\n",
|
||||
" return logger\n",
|
||||
" \n",
|
||||
" pcts = pcts.flatten().numpy()\n",
|
||||
" percentiles = np.linspace(0.05, 0.95, 19)\n",
|
||||
" for pct in percentiles:\n",
|
||||
" clb = Calibration(pcts, pct)\n",
|
||||
" logger.append([clb, np.round(pct, 2), model, mean, k])\n",
|
||||
" \n",
|
||||
" return logger"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 20,
|
||||
"id": "5fe6b4d6-b49e-4e06-8874-9be43b4ee4bd",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"data_path = \"../../voltron/data/\"\n",
|
||||
"ticker_list = make_ticker_list(data_path + \"test_tickers.txt\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 21,
|
||||
"id": "da419f67-e0d9-4008-b038-96fd341e10f3",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"logger = []\n",
|
||||
"end_date = \"2022-01-13\"\n",
|
||||
"tckr = ticker_list[0]\n",
|
||||
"data = GetStockHistory(tckr, history=1000, end_date=end_date)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 22,
|
||||
"id": "b47fa8fa-293a-4f22-8bfa-38c3719d0231",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"log = GetCalibration('gpcv', horizon=np.arange(75,100), \n",
|
||||
" logger=[], exp=True)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 23,
|
||||
"id": "7ce0ce5d-11a3-4728-b242-c5c40b190acc",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"[[0.0032876712328767125, 0.05, 'gpcv', 'ewma', 100],\n",
|
||||
" [0.006575342465753425, 0.1, 'gpcv', 'ewma', 100],\n",
|
||||
" [0.009863013698630137, 0.15, 'gpcv', 'ewma', 100],\n",
|
||||
" [0.009863013698630137, 0.2, 'gpcv', 'ewma', 100],\n",
|
||||
" [0.010958904109589041, 0.25, 'gpcv', 'ewma', 100],\n",
|
||||
" [0.01315068493150685, 0.3, 'gpcv', 'ewma', 100],\n",
|
||||
" [0.020821917808219178, 0.35, 'gpcv', 'ewma', 100],\n",
|
||||
" [0.04547945205479452, 0.4, 'gpcv', 'ewma', 100],\n",
|
||||
" [0.09917808219178083, 0.45, 'gpcv', 'ewma', 100],\n",
|
||||
" [0.16383561643835617, 0.5, 'gpcv', 'ewma', 100],\n",
|
||||
" [0.24493150684931506, 0.55, 'gpcv', 'ewma', 100],\n",
|
||||
" [0.3150684931506849, 0.6, 'gpcv', 'ewma', 100],\n",
|
||||
" [0.38904109589041097, 0.65, 'gpcv', 'ewma', 100],\n",
|
||||
" [0.44767123287671234, 0.7, 'gpcv', 'ewma', 100],\n",
|
||||
" [0.5090410958904109, 0.75, 'gpcv', 'ewma', 100],\n",
|
||||
" [0.581917808219178, 0.8, 'gpcv', 'ewma', 100],\n",
|
||||
" [0.6547945205479452, 0.85, 'gpcv', 'ewma', 100],\n",
|
||||
" [0.7682191780821918, 0.9, 'gpcv', 'ewma', 100],\n",
|
||||
" [0.8936986301369862, 0.95, 'gpcv', 'ewma', 100]]"
|
||||
]
|
||||
},
|
||||
"execution_count": 23,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"log"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 5,
|
||||
"id": "d9ed9129-be2d-47e8-9866-5a708c27d225",
|
||||
"metadata": {
|
||||
"collapsed": true,
|
||||
"jupyter": {
|
||||
"outputs_hidden": true
|
||||
},
|
||||
"tags": []
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"ERROR:root:Internal Python error in the inspect module.\n",
|
||||
"Below is the traceback from this internal error.\n",
|
||||
"\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Traceback (most recent call last):\n",
|
||||
" File \"/home/greg_b/miniconda3/envs/rpp/lib/python3.8/site-packages/IPython/core/interactiveshell.py\", line 3437, in run_code\n",
|
||||
" exec(code_obj, self.user_global_ns, self.user_ns)\n",
|
||||
" File \"<ipython-input-5-43a8e5b361f3>\", line 5, in <module>\n",
|
||||
" log = GetCalibration(tckr, 'volt', mean=mean, k=k, horizon=np.arange(75,100),\n",
|
||||
" File \"<ipython-input-2-2db4e6563d03>\", line 17, in GetCalibration\n",
|
||||
" data = GetStockHistory(tckr, history=1000, end_date=end_date)\n",
|
||||
" File \"/home/greg_b/voltron/voltron/data/MakeData.py\", line 39, in GetStockHistory\n",
|
||||
" data = yf.download(tickers=ticker, period='10y', progress=False)\n",
|
||||
" File \"/home/greg_b/miniconda3/envs/rpp/lib/python3.8/site-packages/yfinance/multi.py\", line 112, in download\n",
|
||||
" _time.sleep(0.01)\n",
|
||||
"KeyboardInterrupt\n",
|
||||
"\n",
|
||||
"During handling of the above exception, another exception occurred:\n",
|
||||
"\n",
|
||||
"Traceback (most recent call last):\n",
|
||||
" File \"/home/greg_b/miniconda3/envs/rpp/lib/python3.8/site-packages/IPython/core/interactiveshell.py\", line 2061, in showtraceback\n",
|
||||
" stb = value._render_traceback_()\n",
|
||||
"AttributeError: 'KeyboardInterrupt' object has no attribute '_render_traceback_'\n",
|
||||
"\n",
|
||||
"During handling of the above exception, another exception occurred:\n",
|
||||
"\n",
|
||||
"Traceback (most recent call last):\n",
|
||||
" File \"/home/greg_b/miniconda3/envs/rpp/lib/python3.8/site-packages/IPython/core/ultratb.py\", line 1101, in get_records\n",
|
||||
" return _fixed_getinnerframes(etb, number_of_lines_of_context, tb_offset)\n",
|
||||
" File \"/home/greg_b/miniconda3/envs/rpp/lib/python3.8/site-packages/IPython/core/ultratb.py\", line 248, in wrapped\n",
|
||||
" return f(*args, **kwargs)\n",
|
||||
" File \"/home/greg_b/miniconda3/envs/rpp/lib/python3.8/site-packages/IPython/core/ultratb.py\", line 281, in _fixed_getinnerframes\n",
|
||||
" records = fix_frame_records_filenames(inspect.getinnerframes(etb, context))\n",
|
||||
" File \"/home/greg_b/miniconda3/envs/rpp/lib/python3.8/inspect.py\", line 1515, in getinnerframes\n",
|
||||
" frameinfo = (tb.tb_frame,) + getframeinfo(tb, context)\n",
|
||||
" File \"/home/greg_b/miniconda3/envs/rpp/lib/python3.8/inspect.py\", line 1473, in getframeinfo\n",
|
||||
" filename = getsourcefile(frame) or getfile(frame)\n",
|
||||
" File \"/home/greg_b/miniconda3/envs/rpp/lib/python3.8/inspect.py\", line 708, in getsourcefile\n",
|
||||
" if getattr(getmodule(object, filename), '__loader__', None) is not None:\n",
|
||||
" File \"/home/greg_b/miniconda3/envs/rpp/lib/python3.8/inspect.py\", line 754, in getmodule\n",
|
||||
" os.path.realpath(f)] = module.__name__\n",
|
||||
" File \"/home/greg_b/miniconda3/envs/rpp/lib/python3.8/posixpath.py\", line 391, in realpath\n",
|
||||
" path, ok = _joinrealpath(filename[:0], filename, {})\n",
|
||||
" File \"/home/greg_b/miniconda3/envs/rpp/lib/python3.8/posixpath.py\", line 425, in _joinrealpath\n",
|
||||
" if not islink(newpath):\n",
|
||||
" File \"/home/greg_b/miniconda3/envs/rpp/lib/python3.8/posixpath.py\", line 167, in islink\n",
|
||||
" st = os.lstat(path)\n",
|
||||
"KeyboardInterrupt\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"ename": "TypeError",
|
||||
"evalue": "object of type 'NoneType' has no len()",
|
||||
"output_type": "error",
|
||||
"traceback": [
|
||||
"\u001b[0;31m---------------------------------------------------------------------------\u001b[0m",
|
||||
"\u001b[0;31mKeyboardInterrupt\u001b[0m Traceback (most recent call last)",
|
||||
" \u001b[0;31m[... skipping hidden 1 frame]\u001b[0m\n",
|
||||
"\u001b[0;32m<ipython-input-5-43a8e5b361f3>\u001b[0m in \u001b[0;36m<module>\u001b[0;34m\u001b[0m\n\u001b[1;32m 4\u001b[0m \u001b[0;32mfor\u001b[0m \u001b[0mmean\u001b[0m \u001b[0;32min\u001b[0m \u001b[0;34m[\u001b[0m\u001b[0;34m'ewma'\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m'dewma'\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m'tewma'\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m----> 5\u001b[0;31m log = GetCalibration(tckr, 'volt', mean=mean, k=k, horizon=np.arange(75,100), \n\u001b[0m\u001b[1;32m 6\u001b[0m logger=log, exp=True)\n",
|
||||
"\u001b[0;32m<ipython-input-2-2db4e6563d03>\u001b[0m in \u001b[0;36mGetCalibration\u001b[0;34m(tckr, model, mean, k, horizon, logger, exp)\u001b[0m\n\u001b[1;32m 16\u001b[0m \u001b[0;32mfor\u001b[0m \u001b[0mtckr\u001b[0m \u001b[0;32min\u001b[0m \u001b[0mticker_list\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 17\u001b[0;31m \u001b[0mdata\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mGetStockHistory\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mtckr\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mhistory\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0;36m1000\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mend_date\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mend_date\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 18\u001b[0m \u001b[0;32mfor\u001b[0m \u001b[0midx\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mdate\u001b[0m \u001b[0;32min\u001b[0m \u001b[0menumerate\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mdata\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mindex\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/voltron/voltron/data/MakeData.py\u001b[0m in \u001b[0;36mGetStockHistory\u001b[0;34m(ticker, end_date, history)\u001b[0m\n\u001b[1;32m 38\u001b[0m \u001b[0mend_date\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mdatetime\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mdatetime\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mstrptime\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mend_date\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m\"%Y-%m-%d\"\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mdate\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 39\u001b[0;31m \u001b[0mdata\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0myf\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mdownload\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mtickers\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mticker\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mperiod\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0;34m'10y'\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mprogress\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0;32mFalse\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 40\u001b[0m \u001b[0mend_idx\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mnp\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mwhere\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mdata\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mindex\u001b[0m \u001b[0;34m==\u001b[0m \u001b[0mpd\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mto_datetime\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mend_date\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0;36m0\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0;36m0\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/yfinance/multi.py\u001b[0m in \u001b[0;36mdownload\u001b[0;34m(tickers, start, end, actions, threads, group_by, auto_adjust, back_adjust, progress, period, show_errors, interval, prepost, proxy, rounding, timeout, **kwargs)\u001b[0m\n\u001b[1;32m 111\u001b[0m \u001b[0;32mwhile\u001b[0m \u001b[0mlen\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mshared\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_DFS\u001b[0m\u001b[0;34m)\u001b[0m \u001b[0;34m<\u001b[0m \u001b[0mlen\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mtickers\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 112\u001b[0;31m \u001b[0m_time\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0msleep\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;36m0.01\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 113\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;31mKeyboardInterrupt\u001b[0m: ",
|
||||
"\nDuring handling of the above exception, another exception occurred:\n",
|
||||
"\u001b[0;31mAttributeError\u001b[0m Traceback (most recent call last)",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/IPython/core/interactiveshell.py\u001b[0m in \u001b[0;36mshowtraceback\u001b[0;34m(self, exc_tuple, filename, tb_offset, exception_only, running_compiled_code)\u001b[0m\n\u001b[1;32m 2060\u001b[0m \u001b[0;31m# in the engines. This should return a list of strings.\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 2061\u001b[0;31m \u001b[0mstb\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mvalue\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_render_traceback_\u001b[0m\u001b[0;34m(\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 2062\u001b[0m \u001b[0;32mexcept\u001b[0m \u001b[0mException\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;31mAttributeError\u001b[0m: 'KeyboardInterrupt' object has no attribute '_render_traceback_'",
|
||||
"\nDuring handling of the above exception, another exception occurred:\n",
|
||||
"\u001b[0;31mTypeError\u001b[0m Traceback (most recent call last)",
|
||||
" \u001b[0;31m[... skipping hidden 1 frame]\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/IPython/core/interactiveshell.py\u001b[0m in \u001b[0;36mshowtraceback\u001b[0;34m(self, exc_tuple, filename, tb_offset, exception_only, running_compiled_code)\u001b[0m\n\u001b[1;32m 2061\u001b[0m \u001b[0mstb\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mvalue\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_render_traceback_\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 2062\u001b[0m \u001b[0;32mexcept\u001b[0m \u001b[0mException\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 2063\u001b[0;31m stb = self.InteractiveTB.structured_traceback(etype,\n\u001b[0m\u001b[1;32m 2064\u001b[0m value, tb, tb_offset=tb_offset)\n\u001b[1;32m 2065\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/IPython/core/ultratb.py\u001b[0m in \u001b[0;36mstructured_traceback\u001b[0;34m(self, etype, value, tb, tb_offset, number_of_lines_of_context)\u001b[0m\n\u001b[1;32m 1365\u001b[0m \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 1366\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mtb\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mtb\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1367\u001b[0;31m return FormattedTB.structured_traceback(\n\u001b[0m\u001b[1;32m 1368\u001b[0m self, etype, value, tb, tb_offset, number_of_lines_of_context)\n\u001b[1;32m 1369\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/IPython/core/ultratb.py\u001b[0m in \u001b[0;36mstructured_traceback\u001b[0;34m(self, etype, value, tb, tb_offset, number_of_lines_of_context)\u001b[0m\n\u001b[1;32m 1265\u001b[0m \u001b[0;32mif\u001b[0m \u001b[0mmode\u001b[0m \u001b[0;32min\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mverbose_modes\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 1266\u001b[0m \u001b[0;31m# Verbose modes need a full traceback\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1267\u001b[0;31m return VerboseTB.structured_traceback(\n\u001b[0m\u001b[1;32m 1268\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0metype\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mvalue\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mtb\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mtb_offset\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mnumber_of_lines_of_context\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 1269\u001b[0m )\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/IPython/core/ultratb.py\u001b[0m in \u001b[0;36mstructured_traceback\u001b[0;34m(self, etype, evalue, etb, tb_offset, number_of_lines_of_context)\u001b[0m\n\u001b[1;32m 1122\u001b[0m \u001b[0;34m\"\"\"Return a nice text document describing the traceback.\"\"\"\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 1123\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1124\u001b[0;31m formatted_exception = self.format_exception_as_a_whole(etype, evalue, etb, number_of_lines_of_context,\n\u001b[0m\u001b[1;32m 1125\u001b[0m tb_offset)\n\u001b[1;32m 1126\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/IPython/core/ultratb.py\u001b[0m in \u001b[0;36mformat_exception_as_a_whole\u001b[0;34m(self, etype, evalue, etb, number_of_lines_of_context, tb_offset)\u001b[0m\n\u001b[1;32m 1080\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 1081\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1082\u001b[0;31m \u001b[0mlast_unique\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mrecursion_repeat\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mfind_recursion\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0morig_etype\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mevalue\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mrecords\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 1083\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 1084\u001b[0m \u001b[0mframes\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mformat_records\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mrecords\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mlast_unique\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mrecursion_repeat\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;32m~/miniconda3/envs/rpp/lib/python3.8/site-packages/IPython/core/ultratb.py\u001b[0m in \u001b[0;36mfind_recursion\u001b[0;34m(etype, value, records)\u001b[0m\n\u001b[1;32m 380\u001b[0m \u001b[0;31m# first frame (from in to out) that looks different.\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 381\u001b[0m \u001b[0;32mif\u001b[0m \u001b[0;32mnot\u001b[0m \u001b[0mis_recursion_error\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0metype\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mvalue\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mrecords\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 382\u001b[0;31m \u001b[0;32mreturn\u001b[0m \u001b[0mlen\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mrecords\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;36m0\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m 383\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 384\u001b[0m \u001b[0;31m# Select filename, lineno, func_name to track frames with\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;31mTypeError\u001b[0m: object of type 'NoneType' has no len()"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"log = []\n",
|
||||
"\n",
|
||||
"for k in [25, 50, 100, 200, 300, 400]:\n",
|
||||
" for mean in ['ewma', 'dewma', 'tewma']:\n",
|
||||
" log = GetCalibration('volt', mean=mean, k=k, horizon=np.arange(75,100), \n",
|
||||
" logger=log, exp=True)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 24,
|
||||
"id": "89054cf9-fb1c-45a1-aa8b-8884713f4c07",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df = pd.DataFrame(log)\n",
|
||||
"df.columns = ['Calibration', 'Percentile', \"Model\", \"Mean\", \"k\"]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 25,
|
||||
"id": "95915269-4bbf-4949-a944-37e076db8889",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# pd.to_pickle(df, \"./volt_calib.pkl\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 28,
|
||||
"id": "bf42d9c5-0bc3-46a8-a1ab-14b87fd8851b",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAA80AAAJQCAYAAAC0OYbcAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMSwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Z1A+gAAAACXBIWXMAABcSAAAXEgFnn9JSAADavElEQVR4nOzdd1xV9/0/8Ne5l72HiMhQQZEhyBAQZQmOGE2iNjHbNMNmT5s2afttun61bZo20yRNbJqYNLFJNXHGwRIRRXEwBUEQZe8Nl3vv+f1BuQEFucDhXsbr+Xjk8bj38jmfz5tW4L7u+QxBFEURRERERERERHQDmb4LICIiIiIiIhqvGJqJiIiIiIiIBsHQTERERERERDQIhmYiIiIiIiKiQTA0ExEREREREQ2CoZmIiIiIiIhoEAzNRERERERERINgaCYiIiIiIiIaBEMzERERERER0SAYmomIiIiIiIgGwdBMRERERERENAiGZiIiIiIiIqJBMDQTERERERERDcJA3wXQxLZlyxYUFxdjzpw5eOONN/RdDhERERERkaQYmmlUiouLkZOTo+8yiIiIiIiIxgSnZxMRERERERENgqGZiIiIiIiIaBAMzURERERERESDYGgmIiIiIiIiGgRDMxEREREREdEgGJqJiIiIiIiIBsHQTERERERERDQIhmYiIiIiIiKiQTA0ExEREREREQ2CoZmIiIiIiIhoEAzNRERERERERINgaCYiIiIiIiIaBEMzERERERER0SAYmomIiIiIiIgGwdBMRERERERENAiGZiIiIiIiIqJBMDQTERERERERDYKhmYiIiIiIiGgQDM1EREREREREg2BoJiIiIiIiolEpLi7GF198gc7OTn2XIjkDfRdAREREREREE1N5eTni4+Nx+fJlAEBaWhqWLVum56qkxdBMREREREREw1JbW4vExETk5uYCAGQyGYKDgxESEqLnyqTH0ExERERERERaEUUR+/fvx9mzZyGKIgDA398fMTExsLW11XN1Y4OhmYiIiIiIiLQiCAKUSiVEUcT8+fMRGxuL6dOn67usMcXQTERERERERAPq6urCyZMn4evri2nTpgEAli1bhuDgYLi6uuq5Ot1gaCYiIiIiIqJ+lEolMjIycOzYMbS3t6O6uhp33XUXAMDa2hrW1tZ6rlB3GJqJiIiIiIgIAKBWq5GZmYmkpCQ0NTUBAOzs7ODj46PnyvSHoZmIiIiIiIhQUFCAI0eOoLa2FgBgaWmJ6OhoBAQEQC6X67k6/WFoJiIiIiIiIlRXV6O2thYmJiaIjIxESEgIDA0N9V2W3jE0ExERERERTUFlZWUQRREuLi4AgLCwMKjVaoSGhsLExETP1Y0fDM1ERERERERTSG1tLRISEpCXlwdHR0c8/vjjEAQBhoaGiIqK0nd54w5DMxERERER0RTQ1NSEpKQkXLhwAaIoAgAcHR2hUChgbGys5+rGL4ZmIiIiIiKiSaytrQ3Hjx/H6dOnoVKpAADz589HbGwspk+fLskYoihCEAStX59IGJqJiIiIiIgmsdLSUpw8eRIAMGvWLMTFxcHV1XXU/fYNxEVXa5GZfw2Xr9WhW6mCraUpfOfNRKC3KyzMJvZdbIZmIiIiIiKiSUSpVKKmpgZOTk4AAC8vLwQEBMDX1xceHh6S3fkVRWDn92fwya405BRWDNjG1MQQd8T646l7o+HuMm1C3nlmaCYiIiIiIpoE1Go1Lly4gOTkZCgUCjz//PMwNjaGIAi44447JB2rpKwOL/7pG5zOvnLTdh2d3fjqQAa+PXoBP314OZ64J2rCBWeGZiIiIiIioglMFEVcvHgRCQkJqK2tBQBYWlqirq4OM2fOlHy8nMJy3LPln2hobtf6mk6FEn/48HsUXKnG335+54QKzgzNREREREREE9Tly5cRHx+P8vJyAICpqSkiIiIQEhICQ0NDSccSRRH1TW144Gf/GlZg7us/35/FjGnW+NmjKyStbSwxNBMREREREU1ADQ0N2LFjBwDA0NAQixcvxpIlS2BiYjIm4wmCgF/8fQ9qGlpH1c97/07GyqXeCPBykaiyscXQTERERERENEG0tbXB3NwcAGBra4vAwEAYGhoiMjISFhYWYzr2ubyr2H8se9T9qNRqbP3oEHa+8agEVY09hmYiIiIiIqJxrqmpCUlJScjKysKTTz4Je3t7AMBtt92ms7XBn353UrK+Us8WobC0Bh6u08b92maGZiIiIiIionGqra0NKSkpOHPmDFQqFQCgoKAA4eHhAKDTwJlwskDS/pLSCzDXzUHSPscCQzMREREREdE409XVhbS0NKSlpUGhUAAAZs+ejbi4OLi46HYtsCiKKK9uQn1Tm6T9ZhWUSdrfWGFoJiIiIiIiGkfUajX+8Y9/oL6+HgDg5OSEuLg4uLu762UqsyAIqKprlrzfilrp+xwLDM1ERERERER6plarIQgCBEGATCbDwoULkZmZiWXLlsHHx0fv637HYnx9f0/aYmgmIiIiIiLSE1EUkZeXh8TERKxatQpz584FACxZsgQRERGQyWR6rrCnxpkO1pL36zxd+j7HAkMzERERERGRHly+fBnx8fEoLy8HAKSlpWlCs4HB+IlqgiDAcZoVHO0tUVXXIlm//p7OkvU1lsbP/xNERERERERTQFlZGeLj41FcXAwAMDQ0RHh4OJYsWaLnym5uxRJvfL43XZK+ZDIBceFeEEVx3E/TZmgmIiIiIiLSkcOHDyMtLQ0AIJfLsWjRIkRGRsLc3FzPlQ3toTvCJAvNsWHz4TrDVpK+xhpDMxERERERkY64uLhAEAQsXLgQ0dHRsLGx0XdJWvP2cMLdq4Ox82DGqPoxNjTAL35yi0RVjT2GZiIiIiIiojHQ1taGlJQU2NvbIyQkBADg7e2Np59+Gvb29nqubni6u5WQyWT4zdNrkHb+MkorGkbc188eWwHP2dMlrG5sMTQTERERERFJqKurCydOnMDJkyehUChgZmaGgIAAGBoaQhCECReYRVHEkbSLsDQ3QWTwXHz510ew8aXtKKtqHHZfT94Ticc3Rk6Itcy9GJqJiIiIiIgkoFQqcfr0aaSkpKCjowMA4OTkhLi4uHG1G/ZwZV8qR97lSgCAtaUp/D2dsXfbk3j1b9/iUGqeVn1YW5jgt8+uxZ0rgyZUYAYYmomIiIiIiEatqKgIe/bsQXNzMwDA3t4esbGx8Pb2nlAB8Xo19S1IOJWveX7kRB46OhUI8ZuN7X94EPEnL+Kfu9Jw7EwhRFG84fpptha459ZgPLphCRzsLCdcYAYYmomIiIiIiEbN3Nwczc3NsLKyQnR0NAICAiCTyfRd1qh0KZTYm5QFpUrd7/XjZ4tw+VotVi7xRtxiL8Qt9kJDczsy88tQfK0W3UoVbKzM4DdvJubNng65TKYJ1BMtMAOAIA70cQCRljZs2ICcnBz4+vpi165d+i6HiIiIiGjMiaKIy5cvo7q6GuHh4ZrXCwoK4O7uPqGnYvcSRREHjuXgYnHlTdvNdrZHmN9sODva3BCIJ+Jd5YFM/P83iYiIiIiIdKSsrAzx8fEoLi6GTCbD/PnzYWdnBwDw9PTUc3XSySwoGzIwA0BVbQusLEwHDMeTITADDM1ERERERERDqqmpQUJCAi5evAgAkMvlWLRoEUxMTPRcmfSq6pqRlF6gVdtbIn1gZTH5/jfoi6GZiIiIiIhoEK2trYiPj8eFCxc0040XLlyI6Oho2NjY6Ls8yXUplNiXlH3DOuaBhPrNhrvLNB1UpV8MzURERERERIOQyWTIzc2FKIrw8vJCbGwsHBwc9F3WmBBFEUdO5KGxpX3Iti6Otlga6K6DqvSPoZmIiIiIiOh/urq6kJOTg8DAQAiCADMzM6xduxa2trZwcXHRd3lj6kJ+GfJLqoZsZ2pshFujfCf87uDaYmgmIiIiIqIpT6lUIj09HcePH0dHRwesra3h4eEBAPDz89NzdWOvsla7dcyCIGB1lC8szSf3Oua+GJqJiIiIiGjKUqvVOH/+PJKTk9Hc3AwAsLe3nzQ7P2ujU9GNfUlZUKmHXscc5jcbc5ztdVDV+MHQTEREREREU44oisjNzUViYiLq6uoAAFZWVoiJicHChQunzNRjURRxJDUPTa0dQ7Z1nWGL8IA5OqhqfGFoJiIiIiKiKUcURSQkJKC+vh6mpqaIjIxESEgIDAymVkQ6l3cNBVeqh2xnZmKEW6MWTJkPE/qaWv8iiIiIiIhoyiorK8OMGTMgl8shk8mwfPlyVFVVITw8HMbGxvouT+cqappw7MylIdsJgoBboxbAwmzq/W8EMDQTEREREdEkV11djcTERFy8eBGrV69GaGgoAMDb2xve3t56rk4/Orq6sT85W7t1zP6zMWumnQ6qGp8YmomIiIiIaFJqbGxEUlISLly4AKDnjmlTU5Oeq9I/URRxODVXq3XMbjPsEL5w6q1j7ouhmYiIiIiIJpXW1lakpKTgzJkzUP/vTqq3tzeWLVsGBwcHPVenf2fzrqKwtGbIduamxrg1euqcxzwYhmYiIiIiIppU9u/fj4sXLwIA5syZg7i4ODg7O+u5qvGhvKYJx05ru47ZF+amU3Mdc18MzURERERENKF1d3dDpVLBxMQEABAVFYWWlhbExsbC3d1dz9WNHx1d3diflAW1KA7ZNnzhHLg5Td11zH0xNBMRERER0YSkVqtx/vx5JCUlwcvLC7feeisAwMnJCY8++igEQdBzheOHKIo4dDwXzW2dQ7ad5WSHMP/ZY1/UBMHQTEREREREE4ooisjNzUViYiLq6uoAAIWFhVAqlZpzlhmY+8vIKUXR1aHXMVuYGU/Z85gHw9A8QgqFAgcOHMD+/ftRWFiI2tpaWFtbw8XFBStWrMD69ethZzc20xnOnTuH7777DhcuXEBZWRna2tpgbGyMadOmwdvbG8uXL8eqVatgZGQ0JuMTEREREemDKIq4fPky4uPjUVFLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 1050x600 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"def PlotCalib(df, ax, title):\n",
|
||||
" sub_df = df[df['k'] == 100]\n",
|
||||
" pal = [ palette[0]]#, palette[4], palette[6]]\n",
|
||||
" sns.lineplot(x='Percentile', y=\"Calibration\", hue='Mean', data=sub_df, ax=ax, alpha=0.5,\n",
|
||||
" palette=pal, legend=False)\n",
|
||||
" sns.scatterplot(x='Percentile', y=\"Calibration\", hue='Mean', data=sub_df, ax=ax, s=120, legend=False, zorder=4,\n",
|
||||
" palette=pal)\n",
|
||||
" x = np.linspace(0.05,0.95)\n",
|
||||
" y = np.linspace(0, len(percentiles))\n",
|
||||
" ax.plot(x, x, color=\"gray\", lw=1., ls=\"--\")\n",
|
||||
"# ax.set_title(title)\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"from matplotlib.lines import Line2D\n",
|
||||
"fig, ax = plt.subplots(1,1,dpi=150, figsize=(7, 4))\n",
|
||||
"\n",
|
||||
"percentiles = np.linspace(0.05, 0.95, 19)\n",
|
||||
"PlotCalib(df, ax, \"Wind Speed Calibration\")\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"plt.tick_params(labelsize=16)\n",
|
||||
"sns.despine()\n",
|
||||
"\n",
|
||||
"# custom_lines = [Line2D([0], [0], color=palette[0], lw=2),\n",
|
||||
"# Line2D([0], [0], color=palette[4], lw=2),\n",
|
||||
"# Line2D([0], [0], color=palette[6], lw=2)]\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"# plt.legend(custom_lines, ['LSTM', r\"GP-Matérn\", \"Volt + Magpie\"],\n",
|
||||
"# fontsize=14, frameon=False, bbox_to_anchor=(0.45, 0.6))\n",
|
||||
"# ax.legend(fontsize=14, bbox_to_anchor=(1., 0.75))\n",
|
||||
"# plt.label(\"Percentile\")\n",
|
||||
"# plt.savefig(\"./wind_calibration.pdf\", bbox_inches=\"tight\")\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 18,
|
||||
"id": "957ddda1-e78b-4bb7-9003-886f18e424f7",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"data = GetStockHistory('ATVI', history=1000, end_date='2020-05-14')\n",
|
||||
"preds = torch.load(\"./saved-outputs/ATVI/gpcv_2020-04-16.pt\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 19,
|
||||
"id": "e75ab9ec-9a86-412b-954c-06c5e9cde875",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"px = data.Close.to_numpy()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "fb01ab2e-c488-4dd1-99ce-8b040c034ae1",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.8"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,125 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
import pandas as pd
|
||||
import gpytorch
|
||||
import argparse
|
||||
import datetime
|
||||
import warnings
|
||||
|
||||
from voltron.data import make_ticker_list, GetStockHistory
|
||||
from GenerateMultiMeanPreds import GenerateStockPredictions, GenerateBasicPredictions, GenerateGPCVPredictions
|
||||
from gpytorch.utils.warnings import NumericalWarning
|
||||
warnings.simplefilter("ignore", NumericalWarning)
|
||||
|
||||
def main(args):
|
||||
|
||||
|
||||
data_path = "../../voltron/data/"
|
||||
ticker_file = args.ticker_fname + ".txt"
|
||||
tckr_list = make_ticker_list(data_path + ticker_file)
|
||||
|
||||
if args.end_date.lower() == "none":
|
||||
end_date = datetime.date.today()
|
||||
else:
|
||||
end_date = datetime.datetime.strptime(args.end_date, "%Y-%m-%d").date()
|
||||
|
||||
for tckr in tckr_list:
|
||||
try:
|
||||
data = GetStockHistory(tckr, history=args.ntrain + args.lookback,
|
||||
end_date=str(end_date))
|
||||
except:
|
||||
print(tckr, "FAILED")
|
||||
data = None
|
||||
if data is not None:
|
||||
if args.kernel.lower() == 'volt':
|
||||
GenerateStockPredictions(tckr, data, forecast_horizon=args.forecast_horizon,
|
||||
train_iters=args.train_iters,
|
||||
nsample=args.nsample,
|
||||
ntrain=400, save=args.save, ntimes=args.ntimes,
|
||||
vol_kernel=args.vol_kernel.lower())
|
||||
elif args.kernel.lower() == 'gpcv':
|
||||
GenerateGPCVPredictions(tckr, data, forecast_horizon=args.forecast_horizon,
|
||||
train_iters=args.train_iters,
|
||||
nsample=args.nsample,
|
||||
ntrain=400, ntimes=args.ntimes)
|
||||
|
||||
else:
|
||||
GenerateBasicPredictions(tckr, data, forecast_horizon=args.forecast_horizon,
|
||||
kernel_name=args.kernel, mean_name=args.mean, k=args.k,
|
||||
train_iters=args.train_iters,
|
||||
nsample=args.nsample,
|
||||
ntrain=args.ntrain, save=args.save)
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--ntimes",
|
||||
type=int,
|
||||
default=25,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--forecast_horizon",
|
||||
type=int,
|
||||
default=100,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ticker_fname",
|
||||
type=str,
|
||||
default='nasdaq100',
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ntrain",
|
||||
type=int,
|
||||
default=400,
|
||||
)
|
||||
parser.add_argument(
|
||||
'--kernel',
|
||||
type=str,
|
||||
default="volt",
|
||||
)
|
||||
parser.add_argument(
|
||||
'--vol_kernel',
|
||||
type=str,
|
||||
default="bm",
|
||||
)
|
||||
parser.add_argument(
|
||||
'--mean',
|
||||
type=str,
|
||||
default="ewma",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--nsample",
|
||||
type=int,
|
||||
default=1000,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--printing",
|
||||
type=bool,
|
||||
default=False
|
||||
)
|
||||
parser.add_argument(
|
||||
"--train_iters",
|
||||
type=int,
|
||||
default=300,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--end_date",
|
||||
default='2022-04-08',
|
||||
)
|
||||
parser.add_argument(
|
||||
"--lookback",
|
||||
type=int,
|
||||
default=500,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--save",
|
||||
type=bool,
|
||||
default=True,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--k",
|
||||
type=int,
|
||||
default=100,
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
main(args)
|
||||
@@ -0,0 +1,182 @@
|
||||
import matplotlib.pyplot as plt
|
||||
import seaborn as sns
|
||||
import numpy as np
|
||||
import torch
|
||||
import pandas as pd
|
||||
import os
|
||||
import gpytorch
|
||||
import argparse
|
||||
import datetime
|
||||
|
||||
from botorch.models import SingleTaskGP
|
||||
from botorch.optim.fit import fit_gpytorch_torch
|
||||
from gpytorch.likelihoods import GaussianLikelihood
|
||||
from gpytorch.mlls import ExactMarginalLogLikelihood
|
||||
from gpytorch.means import ConstantMean, LinearMean
|
||||
from gpytorch.kernels import SpectralMixtureKernel, MaternKernel, RBFKernel, ScaleKernel
|
||||
from voltron.means import EWMAMean, DEWMAMean, TEWMAMean
|
||||
from voltron.train_utils import LearnGPCV, TrainVolModel, TrainVoltMagpieModel, TrainBasicModel
|
||||
from voltron.models import VoltMagpie
|
||||
from voltron.means import LogLinearMean
|
||||
|
||||
from voltron.rollout_utils import GeneratePrediction, Rollouts
|
||||
from voltron.data import make_ticker_list, DataGetter
|
||||
|
||||
|
||||
def main(args):
|
||||
|
||||
|
||||
data_path = "../../voltron/data/"
|
||||
ticker_file = "test_tickers.txt"
|
||||
tckr_list = make_ticker_list(data_path + ticker_file)
|
||||
print("Downloading Data.....")
|
||||
|
||||
if args.end_date.lower() == "none":
|
||||
end_date = str(datetime.date.today())
|
||||
else:
|
||||
end_date = args.end_date
|
||||
|
||||
DataGetter(fpath = data_path, ticker_file=ticker_file, end_date=end_date)
|
||||
print("Data Downloaded.")
|
||||
use_cuda = torch.cuda.is_available()
|
||||
|
||||
ntest = 20
|
||||
dt = 1./252
|
||||
|
||||
print("Producing Forecasts.....")
|
||||
for tckr in tckr_list:
|
||||
dat = pd.read_csv(data_path + tckr + ".csv")
|
||||
train_x = torch.arange(dat.shape[0]-1) * dt
|
||||
test_x = torch.arange(ntest) * dt + train_x[-1] + train_x[1]
|
||||
train_y = torch.FloatTensor(dat.Close.to_numpy())
|
||||
|
||||
if use_cuda:
|
||||
train_x = train_x.cuda()
|
||||
test_x = test_x.cuda()
|
||||
train_y = train_y.cuda()
|
||||
|
||||
train_iters=150
|
||||
mean = 'ewma'
|
||||
nsample = 1000
|
||||
if args.kernel == "volt":
|
||||
vol = LearnGPCV(train_x, train_y, train_iters=args.train_iters,
|
||||
printing=False)
|
||||
vmod, vlh = TrainVolModel(train_x, vol,
|
||||
train_iters=args.train_iters, printing=False)
|
||||
voltron, lh = TrainVoltMagpieModel(train_x, train_y[1:],
|
||||
vmod, vlh, vol,
|
||||
printing=False,
|
||||
train_iters=args.train_iters,
|
||||
k=300, mean_func=args.mean)
|
||||
vmod.eval();
|
||||
if args.mean in ['ewma', 'dewma', 'tewma']:
|
||||
save_samples = Rollouts(train_x, train_y, test_x, voltron,
|
||||
nsample=nsample)
|
||||
|
||||
else: ## VOLT + STANDARD MEAN
|
||||
voltron.vol_model.eval()
|
||||
predvol = voltron.vol_model(test_x).sample(torch.Size((nsample, ))).exp()
|
||||
save_samples[idx, ::] = GeneratePrediction(train_x, train_y, test_x,
|
||||
predvol, voltron).detach()
|
||||
del predvol
|
||||
|
||||
del voltron, lh, vmod, vlh, vol
|
||||
torch.cuda.empty_cache()
|
||||
else:
|
||||
kernel_possibilities = {"sm": SpectralMixtureKernel, "matern": MaternKernel, "rbf": RBFKernel}
|
||||
kernel = kernel_possibilities[args.kernel.lower()]
|
||||
if type(kernel) is not SpectralMixtureKernel:
|
||||
kernel = ScaleKernel(kernel())
|
||||
else:
|
||||
kernel = kernel()
|
||||
kernel.initialize_from_data_empspect(train_x, train_y.log())
|
||||
|
||||
train_y = train_y[1:]
|
||||
|
||||
model = SingleTaskGP(
|
||||
train_x.view(-1,1),
|
||||
train_y.log().reshape(-1, 1),
|
||||
covar_module=kernel,
|
||||
likelihood=GaussianLikelihood()
|
||||
)
|
||||
mean_name = args.mean.lower()
|
||||
if mean_name == "loglinear":
|
||||
model.mean_module = LogLinearMean(1)
|
||||
model.mean_module.initialize_from_data(train_x, train_y.log())
|
||||
elif mean_name == 'linear':
|
||||
model.mean_module = LinearMean(1)
|
||||
elif mean_name == "constant":
|
||||
model.mean_module = ConstantMean()
|
||||
elif mean_name == "ewma":
|
||||
model.mean_module = EWMAMean(train_x, train_y.log(), k=args.k).to(train_x.device)
|
||||
elif mean_name == "dewma":
|
||||
model.mean_module = DEWMAMean(train_x, train_y.log(), k=args.k).to(train_x.device)
|
||||
elif mean_name == "tewma":
|
||||
model.mean_module = TEWMAMean(train_x, train_y.log(), k=args.k).to(train_x.device)
|
||||
|
||||
if use_cuda:
|
||||
model = model.to(train_x.device)
|
||||
|
||||
mll = ExactMarginalLogLikelihood(model.likelihood, model)
|
||||
fit_gpytorch_torch(mll, options={'maxiter':train_iters, 'disp':False})
|
||||
|
||||
if mean_name in ["loglinear", "constant", 'linear']:
|
||||
save_samples[idx] = model.posterior(test_x).sample(torch.Size((nsample, ))).squeeze(-1).cpu().detach()
|
||||
else:
|
||||
save_samples[idx] = Rollouts(
|
||||
train_x, train_y, test_x, model, nsample=nsample, method = "nonvol"
|
||||
).cpu().detach()
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
del model, kernel
|
||||
|
||||
model_name = args.kernel + "_" + args.mean
|
||||
savepath = "./saved-outputs/" + tckr + "/"
|
||||
|
||||
if not os.path.exists(savepath):
|
||||
os.mkdir(savepath)
|
||||
|
||||
torch.save(save_samples, savepath + str(datetime.date.today()) + ".pt")
|
||||
if args.printing:
|
||||
print("\t" + tckr + " done.")
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
parser.add_argument(
|
||||
"--kernel",
|
||||
type=str,
|
||||
default="volt",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--mean",
|
||||
type=str,
|
||||
default="ewma",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--k",
|
||||
type=int,
|
||||
default=100,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--nsample",
|
||||
type=int,
|
||||
default=1000,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--printing",
|
||||
type=bool,
|
||||
default=False
|
||||
)
|
||||
parser.add_argument(
|
||||
"--train_iters",
|
||||
type=int,
|
||||
default=400,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--end_date",
|
||||
default="none",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
main(args)
|
||||
@@ -0,0 +1,138 @@
|
||||
import matplotlib.pyplot as plt
|
||||
import seaborn as sns
|
||||
import numpy as np
|
||||
import torch
|
||||
import pandas as pd
|
||||
import os
|
||||
import gpytorch
|
||||
import argparse
|
||||
import datetime
|
||||
|
||||
from botorch.models import SingleTaskGP
|
||||
from botorch.optim.fit import fit_gpytorch_torch
|
||||
from gpytorch.likelihoods import GaussianLikelihood
|
||||
from gpytorch.mlls import ExactMarginalLogLikelihood
|
||||
from gpytorch.means import ConstantMean, LinearMean
|
||||
from gpytorch.kernels import SpectralMixtureKernel, MaternKernel, RBFKernel, ScaleKernel
|
||||
from voltron.means import EWMAMean, DEWMAMean, TEWMAMean
|
||||
from voltron.train_utils import LearnGPCV, TrainVolModel, TrainVoltMagpieModel, TrainBasicModel
|
||||
from voltron.models import VoltMagpie
|
||||
from voltron.means import LogLinearMean
|
||||
|
||||
from voltron.rollout_utils import GeneratePrediction, Rollouts
|
||||
from voltron.data import make_ticker_list, DataGetter, GetStockHistory
|
||||
|
||||
|
||||
def GenerateStockPredictions(ticker, dat,
|
||||
forecast_horizon=20,
|
||||
train_iters=400, nsample=1000,
|
||||
ntrain=400, mean='ewma', kernel='volt',
|
||||
save=False, k=300):
|
||||
|
||||
end_idxs = torch.arange(ntrain, dat.shape[0])
|
||||
ntest = forecast_horizon
|
||||
dt = 1./252
|
||||
|
||||
model_name = kernel + "_" + mean + str(k) + "_"
|
||||
savepath = "./saved-outputs/" + ticker + "/"
|
||||
|
||||
for last_day in end_idxs:
|
||||
date = str(dat.index[last_day.item()].date())
|
||||
try:
|
||||
train_y = torch.FloatTensor(dat.Close[last_day.item()-ntrain:last_day.item()].to_numpy())
|
||||
train_x = torch.arange(train_y.shape[0]-1) * dt
|
||||
test_x = torch.arange(ntest) * dt + train_x[-1] + train_x[1]
|
||||
# try:
|
||||
use_cuda = torch.cuda.is_available()
|
||||
if use_cuda:
|
||||
train_x = trin_x.cuda()
|
||||
test_x = test_x.cuda()
|
||||
train_y = train_y.cuda()
|
||||
|
||||
# print("Producing " + ticker + " Forecasts.....")
|
||||
if kernel == "volt":
|
||||
vol = LearnGPCV(train_x, train_y, train_iters=train_iters,
|
||||
printing=False)
|
||||
vmod, vlh = TrainVolModel(train_x, vol,
|
||||
train_iters=train_iters, printing=False)
|
||||
voltron, lh = TrainVoltMagpieModel(train_x, train_y[1:],
|
||||
vmod, vlh, vol,
|
||||
printing=False,
|
||||
train_iters=train_iters,
|
||||
k=k, mean_func=mean)
|
||||
vmod.eval();
|
||||
if mean in ['ewma', 'dewma', 'tewma']:
|
||||
save_samples = Rollouts(train_x, train_y, test_x, voltron,
|
||||
nsample=nsample)
|
||||
|
||||
else: ## VOLT + STANDARD MEAN
|
||||
voltron.vol_model.eval()
|
||||
predvol = voltron.vol_model(test_x).sample(torch.Size((nsample, ))).exp()
|
||||
save_samples[idx, ::] = GeneratePrediction(train_x, train_y, test_x,
|
||||
predvol, voltron).detach()
|
||||
del predvol
|
||||
|
||||
del voltron, lh, vmod, vlh, vol
|
||||
torch.cuda.empty_cache()
|
||||
else:
|
||||
kernel_possibilities = {"sm": SpectralMixtureKernel, "matern": MaternKernel, "rbf": RBFKernel}
|
||||
kernel = kernel_possibilities[kernel.lower()]
|
||||
if type(kernel) is not SpectralMixtureKernel:
|
||||
kernel = ScaleKernel(kernel())
|
||||
else:
|
||||
kernel = kernel()
|
||||
kernel.initialize_from_data_empspect(train_x, train_y.log())
|
||||
|
||||
train_y = train_y[1:]
|
||||
|
||||
model = SingleTaskGP(
|
||||
train_x.view(-1,1),
|
||||
train_y.log().reshape(-1, 1),
|
||||
covar_module=kernel,
|
||||
likelihood=GaussianLikelihood()
|
||||
)
|
||||
mean_name = mean.lower()
|
||||
if mean_name == "loglinear":
|
||||
model.mean_module = LogLinearMean(1)
|
||||
model.mean_module.initialize_from_data(train_x, train_y.log())
|
||||
elif mean_name == 'linear':
|
||||
model.mean_module = LinearMean(1)
|
||||
elif mean_name == "constant":
|
||||
model.mean_module = ConstantMean()
|
||||
elif mean_name == "ewma":
|
||||
model.mean_module = EWMAMean(train_x, train_y.log(), k=k).to(train_x.device)
|
||||
elif mean_name == "dewma":
|
||||
model.mean_module = DEWMAMean(train_x, train_y.log(), k=k).to(train_x.device)
|
||||
elif mean_name == "tewma":
|
||||
model.mean_module = TEWMAMean(train_x, train_y.log(), k=k).to(train_x.device)
|
||||
|
||||
if use_cuda:
|
||||
model = model.to(train_x.device)
|
||||
|
||||
mll = ExactMarginalLogLikelihood(model.likelihood, model)
|
||||
fit_gpytorch_torch(mll, options={'maxiter':train_iters, 'disp':False})
|
||||
|
||||
if mean_name in ["loglinear", "constant", 'linear']:
|
||||
save_samples[idx] = model.posterior(test_x).sample(torch.Size((nsample, ))).squeeze(-1).cpu().detach()
|
||||
else:
|
||||
save_samples[idx] = Rollouts(
|
||||
train_x, train_y, test_x, model, nsample=nsample, method = "nonvol"
|
||||
).cpu().detach()
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
del model, kernel
|
||||
|
||||
if save:
|
||||
if not os.path.exists(savepath):
|
||||
os.mkdir(savepath)
|
||||
|
||||
torch.save(save_samples, savepath + model_name + date + ".pt")
|
||||
except:
|
||||
nans = torch.ones(nsample, ntest) * torch.nan
|
||||
if save:
|
||||
if not os.path.exists(savepath):
|
||||
os.mkdir(savepath)
|
||||
torch.save(nans, savepath + model_name + date + ".pt")
|
||||
|
||||
|
||||
return dat, save_samples
|
||||
@@ -0,0 +1,306 @@
|
||||
import matplotlib.pyplot as plt
|
||||
import seaborn as sns
|
||||
import numpy as np
|
||||
import torch
|
||||
import pandas as pd
|
||||
import os
|
||||
import gpytorch
|
||||
import argparse
|
||||
import datetime
|
||||
|
||||
from botorch.models import SingleTaskGP
|
||||
from botorch.optim.fit import fit_gpytorch_torch
|
||||
from gpytorch.likelihoods import GaussianLikelihood
|
||||
from gpytorch.mlls import ExactMarginalLogLikelihood
|
||||
from gpytorch.means import ConstantMean, LinearMean
|
||||
from gpytorch.kernels import SpectralMixtureKernel, MaternKernel, RBFKernel, ScaleKernel
|
||||
from voltron.means import EWMAMean, DEWMAMean, TEWMAMean
|
||||
from voltron.train_utils import LearnGPCV, TrainVolModel, TrainVoltMagpieModel, TrainBasicModel
|
||||
from voltron.models import VoltMagpie
|
||||
from voltron.means import LogLinearMean
|
||||
|
||||
from voltron.rollout_utils import GeneratePrediction, Rollouts
|
||||
from voltron.data import make_ticker_list, DataGetter, GetStockHistory
|
||||
|
||||
|
||||
def GenerateGPCVPredictions(ticker, dat,
|
||||
forecast_horizon=20, ntimes=25,
|
||||
train_iters=400, nsample=1000,
|
||||
ntrain=400):
|
||||
|
||||
end_idxs = torch.arange(ntrain, dat.shape[0],
|
||||
int((dat.shape[0]-ntrain)/ntimes))
|
||||
ntest = forecast_horizon
|
||||
dt = 1./252
|
||||
|
||||
savepath = "./saved-outputs/" + ticker + "/"
|
||||
if not os.path.exists(savepath):
|
||||
os.mkdir(savepath)
|
||||
for last_day in end_idxs:
|
||||
date = str(dat.index[last_day.item()].date())
|
||||
print(date, ticker)
|
||||
train_y = torch.FloatTensor(dat.Close[last_day.item()-ntrain:last_day.item()].to_numpy())
|
||||
train_x = torch.arange(train_y.shape[0]-1) * dt
|
||||
test_x = torch.arange(ntest) * dt + train_x[-1] + train_x[1]
|
||||
# try:
|
||||
use_cuda = torch.cuda.is_available()
|
||||
if use_cuda:
|
||||
train_x = train_x.cuda()
|
||||
test_x = test_x.cuda()
|
||||
train_y = train_y.cuda()
|
||||
|
||||
model, likelihood = LearnGPCV(train_x, train_y,
|
||||
train_iters=train_iters, printing=False, return_model=True)
|
||||
preds = likelihood(model(test_x),
|
||||
return_gaussian=False).sample(torch.Size((nsample,)))
|
||||
preds = preds.cumsum(-1).squeeze()
|
||||
preds = preds.view(-1, preds.shape[-1])
|
||||
save_samples = preds.view(-1, preds.shape[-1]) * (dt ** 0.5) + train_y[-1].log()
|
||||
torch.save(save_samples, savepath + "gpcv_" + date + ".pt")
|
||||
|
||||
return
|
||||
|
||||
def GenerateStockPredictions(ticker, dat,
|
||||
forecast_horizon=20,
|
||||
train_iters=400, nsample=1000,
|
||||
ntrain=400, save=False, ntimes=1,
|
||||
vol_kernel='bm', gpcv_preds=False):
|
||||
|
||||
end_idxs = torch.arange(ntrain, dat.shape[0],
|
||||
int((dat.shape[0]-ntrain)/ntimes))
|
||||
ntest = forecast_horizon
|
||||
dt = 1./252
|
||||
|
||||
savepath = "./saved-outputs/" + ticker + "/"
|
||||
if not os.path.exists(savepath):
|
||||
os.mkdir(savepath)
|
||||
for last_day in end_idxs:
|
||||
date = str(dat.index[last_day.item()].date())
|
||||
print(date, ticker)
|
||||
train_y = torch.FloatTensor(dat.Close[last_day.item()-ntrain:last_day.item()].to_numpy())
|
||||
train_x = torch.arange(train_y.shape[0]-1) * dt
|
||||
test_x = torch.arange(ntest) * dt + train_x[-1] + train_x[1]
|
||||
# try:
|
||||
use_cuda = torch.cuda.is_available()
|
||||
if use_cuda:
|
||||
train_x = train_x.cuda()
|
||||
test_x = test_x.cuda()
|
||||
train_y = train_y.cuda()
|
||||
|
||||
# print("Producing " + ticker + " Forecasts.....")
|
||||
vol = LearnGPCV(train_x, train_y, train_iters=train_iters,
|
||||
printing=False, kernel=vol_kernel)
|
||||
vmod, vlh = TrainVolModel(train_x, vol,
|
||||
train_iters=train_iters, printing=False)
|
||||
# for mean in ['ewma']:#, 'dewma', 'tewma']:
|
||||
# for k in [25, 50, 100, 200, 300, 400]:
|
||||
# try:
|
||||
# voltron, lh = TrainVoltMagpieModel(train_x, train_y[1:],
|
||||
# vmod, vlh, vol,
|
||||
# printing=False,
|
||||
# train_iters=0,
|
||||
# k=k, mean_func=mean)
|
||||
# vmod.eval();
|
||||
# voltron.eval();
|
||||
# save_samples = Rollouts(train_x, train_y, test_x, voltron,
|
||||
# nsample=nsample)
|
||||
# except:
|
||||
# print("Failed: ", ticker, mean, k)
|
||||
# if save:
|
||||
# save_samples = torch.ones(nsample, ntest) * torch.nan
|
||||
|
||||
# if save:
|
||||
# model_name = "volt_"
|
||||
# if vol_kernel == 'ou':
|
||||
# model_name = 'vhgp_'
|
||||
# model_name += mean + str(k) + "_"
|
||||
# torch.save(save_samples, savepath + model_name + date + ".pt")
|
||||
# del voltron, lh, vmod, vlh, vol
|
||||
# torch.cuda.empty_cache()
|
||||
|
||||
###################
|
||||
## CONSTANT MEAN ##
|
||||
###################
|
||||
# try:
|
||||
voltron, lh = TrainVoltMagpieModel(train_x, train_y[1:],
|
||||
vmod, vlh, vol,
|
||||
printing=False,
|
||||
train_iters=0,
|
||||
k=100, mean_func='constant')
|
||||
vmod.eval();
|
||||
voltron.eval();
|
||||
save_samples = Rollouts(train_x, train_y, test_x, voltron,
|
||||
nsample=nsample)
|
||||
# except:
|
||||
# print("Failed: ", ticker, mean, k)
|
||||
# if save:
|
||||
# save_samples = torch.ones(nsample, ntest) * torch.nan
|
||||
|
||||
if save:
|
||||
model_name = "volt_"
|
||||
if vol_kernel == 'ou':
|
||||
model_name = 'vhgp_'
|
||||
model_name += "constant_"
|
||||
torch.save(save_samples, savepath + model_name + date + ".pt")
|
||||
|
||||
del voltron, lh, vmod, vlh, vol
|
||||
torch.cuda.empty_cache()
|
||||
return
|
||||
|
||||
|
||||
def GenerateOneDayPredictions(ticker, train_y, date,
|
||||
forecast_horizon=20,
|
||||
train_iters=400, nsample=1000,
|
||||
ntrain=400, save=False, mean=None):
|
||||
|
||||
ntest = forecast_horizon
|
||||
dt = 1./252
|
||||
|
||||
savepath = "./saved-outputs/" + ticker + "/"
|
||||
if not os.path.exists(savepath):
|
||||
os.mkdir(savepath)
|
||||
|
||||
train_x = torch.arange(train_y.shape[0]-1) * dt
|
||||
test_x = torch.arange(ntest) * dt + train_x[-1] + train_x[1]
|
||||
use_cuda = torch.cuda.is_available()
|
||||
if use_cuda:
|
||||
train_x = train_x.cuda()
|
||||
test_x = test_x.cuda()
|
||||
train_y = train_y.cuda()
|
||||
|
||||
vol = LearnGPCV(train_x, train_y, train_iters=train_iters,
|
||||
printing=False)
|
||||
vmod, vlh = TrainVolModel(train_x, vol,
|
||||
train_iters=train_iters, printing=False)
|
||||
|
||||
if mean=='constant':
|
||||
voltron, lh = TrainVoltMagpieModel(train_x, train_y[1:],
|
||||
vmod, vlh, vol,
|
||||
printing=False,
|
||||
train_iters=200,
|
||||
mean_func='constant')
|
||||
vmod.eval();
|
||||
voltron.eval();
|
||||
save_samples = Rollouts(train_x, train_y, test_x, voltron,
|
||||
nsample=nsample)
|
||||
|
||||
if save:
|
||||
model_name = "volt_" + mean + "_"
|
||||
torch.save(save_samples, savepath + model_name + date + ".pt")
|
||||
else:
|
||||
for mean in ['ewma', 'dewma', 'tewma']:
|
||||
for k in [25, 50, 100, 200, 300, 400]:
|
||||
try:
|
||||
voltron, lh = TrainVoltMagpieModel(train_x, train_y[1:],
|
||||
vmod, vlh, vol,
|
||||
printing=False,
|
||||
train_iters=0,
|
||||
k=k, mean_func=mean)
|
||||
vmod.eval();
|
||||
voltron.eval();
|
||||
save_samples = Rollouts(train_x, train_y, test_x, voltron,
|
||||
nsample=nsample)
|
||||
except:
|
||||
print("Failed: ", ticker, mean, k)
|
||||
if save:
|
||||
save_samples = torch.ones(nsample, ntest) * torch.nan
|
||||
|
||||
if save:
|
||||
model_name = "volt_" + mean + str(k) + "_"
|
||||
torch.save(save_samples, savepath + model_name + date + ".pt")
|
||||
|
||||
del voltron, lh, vmod, vlh, vol, save_samples
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
return
|
||||
|
||||
|
||||
|
||||
def GenerateBasicPredictions(ticker, dat, kernel_name, mean_name='ewma', k=400,
|
||||
forecast_horizon=100,
|
||||
train_iters=600, nsample=1000,
|
||||
ntrain=400, save=False, ntimes=-1):
|
||||
|
||||
if ntimes == -1:
|
||||
end_idxs = torch.arange(ntrain, dat.shape[0])
|
||||
else:
|
||||
end_idxs = torch.arange(ntrain, dat.shape[0],
|
||||
int((dat.shape[0]-ntrain)/ntimes))
|
||||
ntest = forecast_horizon
|
||||
dt = 1./252
|
||||
|
||||
savepath = "./saved-outputs/" + ticker + "/"
|
||||
if not os.path.exists(savepath):
|
||||
os.mkdir(savepath)
|
||||
|
||||
for last_day in end_idxs:
|
||||
date = str(dat.index[last_day.item()].date())
|
||||
train_y = torch.FloatTensor(dat.Close[last_day.item()-ntrain:last_day.item()].to_numpy())
|
||||
train_x = torch.arange(train_y.shape[0]-1) * dt
|
||||
test_x = torch.arange(ntest) * dt + train_x[-1] + train_x[1]
|
||||
# try:
|
||||
use_cuda = torch.cuda.is_available()
|
||||
if use_cuda:
|
||||
train_x = train_x.cuda()
|
||||
test_x = test_x.cuda()
|
||||
train_y = train_y.cuda()
|
||||
|
||||
# print("Producing " + ticker + " Forecasts.....")
|
||||
kernel_possibilities = {"sm": SpectralMixtureKernel,
|
||||
"matern": MaternKernel,
|
||||
"rbf": RBFKernel}
|
||||
kernel = kernel_possibilities[kernel_name.lower()]
|
||||
if kernel_name.lower() != 'sm':
|
||||
kernel = ScaleKernel(kernel())
|
||||
else:
|
||||
kernel = kernel(num_mixtures=15)
|
||||
kernel.initialize_from_data_empspect(train_x, train_y.log())
|
||||
|
||||
train_y = train_y[1:]
|
||||
|
||||
model = SingleTaskGP(
|
||||
train_x.view(-1,1),
|
||||
train_y.log().reshape(-1, 1),
|
||||
covar_module=kernel,
|
||||
likelihood=GaussianLikelihood()
|
||||
)
|
||||
|
||||
mean_name = mean_name.lower()
|
||||
if mean_name == "loglinear":
|
||||
model.mean_module = LogLinearMean(1)
|
||||
model.mean_module.initialize_from_data(train_x, train_y.log())
|
||||
elif mean_name == 'linear':
|
||||
model.mean_module = LinearMean(1)
|
||||
elif mean_name == "constant":
|
||||
model.mean_module = ConstantMean()
|
||||
elif mean_name == "ewma":
|
||||
model.mean_module = EWMAMean(train_x, train_y.log(), k=k).to(train_x.device)
|
||||
elif mean_name == "dewma":
|
||||
model.mean_module = DEWMAMean(train_x, train_y.log(), k=k).to(train_x.device)
|
||||
elif mean_name == "tewma":
|
||||
model.mean_module = TEWMAMean(train_x, train_y.log(), k=k).to(train_x.device)
|
||||
|
||||
if use_cuda:
|
||||
model = model.to(train_x.device)
|
||||
print("Fitting Model", ticker)
|
||||
mll = ExactMarginalLogLikelihood(model.likelihood, model)
|
||||
fit_gpytorch_torch(mll, options={'maxiter':train_iters, 'disp':False})
|
||||
|
||||
if mean_name in ["loglinear", "constant", 'linear']:
|
||||
save_samples = model.posterior(test_x).sample(torch.Size((nsample,
|
||||
))).squeeze(-1).cpu().detach()
|
||||
else:
|
||||
save_samples = Rollouts(
|
||||
train_x, train_y, test_x, model, nsample=nsample, method = "nonvol"
|
||||
).cpu().detach()
|
||||
|
||||
model_name = kernel_name + "_" + mean_name + str(k) + "_"
|
||||
torch.save(save_samples, savepath + model_name + date + ".pt")
|
||||
|
||||
model.train()
|
||||
torch.cuda.empty_cache()
|
||||
del model
|
||||
|
||||
|
||||
return dat, save_samples
|
||||
@@ -0,0 +1,116 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
import pandas as pd
|
||||
import gpytorch
|
||||
import argparse
|
||||
import datetime
|
||||
import warnings
|
||||
|
||||
from voltron.data import make_ticker_list, GetStockHistory
|
||||
from GenerateMultiMeanPreds import GenerateStockPredictions, GenerateBasicPredictions, GenerateOneDayPredictions
|
||||
from gpytorch.utils.warnings import NumericalWarning
|
||||
warnings.simplefilter("ignore", NumericalWarning)
|
||||
warnings.simplefilter("ignore", UserWarning)
|
||||
|
||||
def main(args):
|
||||
tckr = args.ticker
|
||||
## download data ##
|
||||
if args.end_date.lower() == "none":
|
||||
end_date = datetime.date.today()
|
||||
else:
|
||||
end_date = datetime.datetime.strptime(args.end_date, "%Y-%m-%d").date()
|
||||
|
||||
dat = GetStockHistory(tckr, history=args.ntrain + args.lookback,
|
||||
end_date=str(end_date))
|
||||
|
||||
## pick a day to generate forecasts ##
|
||||
end_idxs = torch.arange(args.ntrain, dat.shape[0],
|
||||
int((dat.shape[0]-args.ntrain)/args.ntimes))
|
||||
last_day = end_idxs[args.test_idx]
|
||||
date = str(dat.index[last_day.item()].date())
|
||||
|
||||
print(date, tckr)
|
||||
|
||||
train_y = torch.FloatTensor(dat.Close[last_day.item()-args.ntrain:last_day.item()].to_numpy())
|
||||
|
||||
GenerateOneDayPredictions(tckr, train_y, date,
|
||||
forecast_horizon=args.forecast_horizon,
|
||||
train_iters=args.train_iters,
|
||||
nsample=args.nsample,
|
||||
ntrain=400, save=args.save, mean=args.mean)
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--ntimes",
|
||||
type=int,
|
||||
default=25,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--test_idx",
|
||||
type=int,
|
||||
default=0,
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--forecast_horizon",
|
||||
type=int,
|
||||
default=100,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ticker",
|
||||
type=str,
|
||||
default='ADBE',
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ntrain",
|
||||
type=int,
|
||||
default=400,
|
||||
)
|
||||
parser.add_argument(
|
||||
'--kernel',
|
||||
type=str,
|
||||
default="volt",
|
||||
)
|
||||
parser.add_argument(
|
||||
'--mean',
|
||||
type=str,
|
||||
default="ewma",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--nsample",
|
||||
type=int,
|
||||
default=1000,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--printing",
|
||||
type=bool,
|
||||
default=False
|
||||
)
|
||||
parser.add_argument(
|
||||
"--train_iters",
|
||||
type=int,
|
||||
default=300,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--end_date",
|
||||
default="none",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--lookback",
|
||||
type=int,
|
||||
default=500,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--save",
|
||||
type=bool,
|
||||
default=True,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--k",
|
||||
type=int,
|
||||
default=100,
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
main(args)
|
||||
@@ -0,0 +1,69 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
import pandas as pd
|
||||
import gpytorch
|
||||
import argparse
|
||||
|
||||
def ValueFunction(pred_samples, curr_px):
|
||||
"""
|
||||
pred_samples = (num samples) x (test times) matrix of forecast paths
|
||||
curr_px = last observed price
|
||||
"""
|
||||
snr = (pred_samples.mean(0) - curr_px)/pred_samples.std(0)
|
||||
|
||||
if snr > 0.2:
|
||||
return 1.
|
||||
else:
|
||||
return 0.
|
||||
|
||||
|
||||
|
||||
def main(args):
|
||||
data_path = "../../voltron/data/"
|
||||
ticker_file = "test_tickers.txt"
|
||||
tckr_list = make_ticker_list(data_path + ticker_file)
|
||||
|
||||
if args.end_date
|
||||
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
parser.add_argument(
|
||||
"--kernel",
|
||||
type=str,
|
||||
default="volt",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--mean",
|
||||
type=str,
|
||||
default="ewma",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--k",
|
||||
type=int,
|
||||
default=100,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--nsample",
|
||||
type=int,
|
||||
default=1000,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--printing",
|
||||
type=bool,
|
||||
default=False
|
||||
)
|
||||
parser.add_argument(
|
||||
"--train_iters",
|
||||
type=int,
|
||||
default=400,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--end_date",
|
||||
default="none",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
main(args)
|
||||
@@ -0,0 +1,491 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "be6383aa",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Warning no robinhood utils.\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import torch\n",
|
||||
"import yfinance as yf\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns\n",
|
||||
"import numpy as np\n",
|
||||
"import torch\n",
|
||||
"import pandas as pd\n",
|
||||
"import os\n",
|
||||
"import gpytorch\n",
|
||||
"import argparse\n",
|
||||
"import datetime\n",
|
||||
"\n",
|
||||
"from botorch.models import SingleTaskGP\n",
|
||||
"from botorch.optim.fit import fit_gpytorch_torch\n",
|
||||
"from gpytorch.likelihoods import GaussianLikelihood\n",
|
||||
"from gpytorch.mlls import ExactMarginalLogLikelihood\n",
|
||||
"from gpytorch.means import ConstantMean, LinearMean\n",
|
||||
"from gpytorch.kernels import SpectralMixtureKernel, MaternKernel, RBFKernel, ScaleKernel\n",
|
||||
"from voltron.means import EWMAMean, DEWMAMean, TEWMAMean\n",
|
||||
"from voltron.train_utils import LearnGPCV, TrainVolModel, TrainVoltMagpieModel, TrainBasicModel\n",
|
||||
"from voltron.models import VoltMagpie\n",
|
||||
"from voltron.means import LogLinearMean\n",
|
||||
"\n",
|
||||
"from voltron.rollout_utils import GeneratePrediction, Rollouts\n",
|
||||
"from voltron.data import make_ticker_list, DataGetter, GetStockHistory"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "8cce2850",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Mon Dec 20 07:56:35 2021 \n",
|
||||
"+-----------------------------------------------------------------------------+\n",
|
||||
"| NVIDIA-SMI 465.19.01 Driver Version: 465.19.01 CUDA Version: 11.3 |\n",
|
||||
"|-------------------------------+----------------------+----------------------+\n",
|
||||
"| GPU Name Persistence-M| Bus-Id Disp.A | Volatile Uncorr. ECC |\n",
|
||||
"| Fan Temp Perf Pwr:Usage/Cap| Memory-Usage | GPU-Util Compute M. |\n",
|
||||
"| | | MIG M. |\n",
|
||||
"|===============================+======================+======================|\n",
|
||||
"| 0 NVIDIA TITAN RTX On | 00000000:1A:00.0 Off | N/A |\n",
|
||||
"| 41% 32C P8 11W / 280W | 16782MiB / 24220MiB | 0% Default |\n",
|
||||
"| | | N/A |\n",
|
||||
"+-------------------------------+----------------------+----------------------+\n",
|
||||
"| 1 NVIDIA TITAN RTX On | 00000000:1B:00.0 Off | N/A |\n",
|
||||
"| 41% 24C P8 9W / 280W | 3MiB / 24220MiB | 0% Default |\n",
|
||||
"| | | N/A |\n",
|
||||
"+-------------------------------+----------------------+----------------------+\n",
|
||||
"| 2 NVIDIA TITAN RTX On | 00000000:3D:00.0 Off | N/A |\n",
|
||||
"| 41% 28C P5 53W / 280W | 3MiB / 24220MiB | 0% Default |\n",
|
||||
"| | | N/A |\n",
|
||||
"+-------------------------------+----------------------+----------------------+\n",
|
||||
"| 3 NVIDIA TITAN RTX On | 00000000:3E:00.0 Off | N/A |\n",
|
||||
"| 41% 25C P8 16W / 280W | 3MiB / 24220MiB | 0% Default |\n",
|
||||
"| | | N/A |\n",
|
||||
"+-------------------------------+----------------------+----------------------+\n",
|
||||
"| 4 NVIDIA TITAN RTX On | 00000000:88:00.0 Off | N/A |\n",
|
||||
"| 40% 24C P8 13W / 280W | 3MiB / 24220MiB | 0% Default |\n",
|
||||
"| | | N/A |\n",
|
||||
"+-------------------------------+----------------------+----------------------+\n",
|
||||
"| 5 NVIDIA TITAN RTX On | 00000000:89:00.0 Off | N/A |\n",
|
||||
"| 41% 24C P8 19W / 280W | 3MiB / 24220MiB | 0% Default |\n",
|
||||
"| | | N/A |\n",
|
||||
"+-------------------------------+----------------------+----------------------+\n",
|
||||
"| 6 NVIDIA TITAN RTX On | 00000000:B1:00.0 Off | N/A |\n",
|
||||
"| 41% 25C P8 4W / 280W | 3MiB / 24220MiB | 0% Default |\n",
|
||||
"| | | N/A |\n",
|
||||
"+-------------------------------+----------------------+----------------------+\n",
|
||||
"| 7 NVIDIA TITAN RTX On | 00000000:B2:00.0 Off | N/A |\n",
|
||||
"| 40% 25C P8 15W / 280W | 3MiB / 24220MiB | 0% Default |\n",
|
||||
"| | | N/A |\n",
|
||||
"+-------------------------------+----------------------+----------------------+\n",
|
||||
" \n",
|
||||
"+-----------------------------------------------------------------------------+\n",
|
||||
"| Processes: |\n",
|
||||
"| GPU GI CI PID Type Process name GPU Memory |\n",
|
||||
"| ID ID Usage |\n",
|
||||
"|=============================================================================|\n",
|
||||
"| 0 N/A N/A 61360 C ...y_m/miniconda3/bin/python 16779MiB |\n",
|
||||
"+-----------------------------------------------------------------------------+\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"!nvidia-smi"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"id": "bb53d987",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"[*********************100%***********************] 1 of 1 completed\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"all_y = yf.download(\"XOM\", start = \"2020-06-01\", end=\"2020-11-01\", interval=\"1h\").Close.values"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"id": "5a9b0817",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"[<matplotlib.lines.Line2D at 0x7fe461d86fa0>]"
|
||||
]
|
||||
},
|
||||
"execution_count": 4,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAXAAAAD4CAYAAAD1jb0+AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/MnkTPAAAACXBIWXMAAAsTAAALEwEAmpwYAAA0FklEQVR4nO3dd3xb1dnA8d/R8N6xM53E2QmEDOIMSAKEEHbLHm2hlPFCKaW0lPKSt6Uv3ZSyKaNASynzhQINM4SEBAIJCc7ecfZw4pHE27I1zvuHrmTJlm05lmVd8Xw/H38s3XslPWY8Pn7uOc9RWmuEEEKYj6WnAxBCCHF8JIELIYRJSQIXQgiTkgQuhBAmJQlcCCFMyhbND8vNzdUFBQXR/EghhDC9VatWVWit81oej2oCLygooKioKJofKYQQpqeU2hvquJRQhBDCpCSBCyGESUkCF0IIkworgSul9iilNiil1iqlioxj9ymlDhrH1iqlzu/eUIUQQgTqzE3MWVrrihbHHtFaPxjJgIQQQoRHSihCCGFS4SZwDSxQSq1SSt0ccPzHSqn1Sql/KKWyQ71QKXWzUqpIKVVUXl7e5YCFEEJ4hZvAp2utTwbOA25TSp0GPA0MAyYAh4CHQr1Qa/2s1rpQa12Yl9dqHnrUfLjhEEdqG3vs84UQItLCSuBa6xLjexnwDjBFa12qtXZrrT3Ac8CU7guzaw5VNfCjV1Zzx+trezoUIYSImA4TuFIqVSmV7nsMnA1sVEr1C7jsEmBj94TYdUfrmgCokBG4ECKOhDMLpQ/wjlLKd/2rWuv5SqmXlFIT8NbH9wC3dFeQXbX/aD0AaYlR7RwghBDdqsOMprXeBYwPcfzabokowj7dWsoPX14NQHqSJHAhRPyI+2mEC7eU+R+nJdl7MBIhhIisuE/gdY0u/2MZgQsh4klcJ/D6Jhdf7mhePKq1pri0JuiYEEKYVVwn8FdX7KOitolZo7zzz5tcmjmPfM73nl/Rw5EJIUTXxXUCL69pJMFq4YXrpzAkNxWHy+0/1xjwWAghzCiuE3hNo8tf906wWvhg/SH/ucNVjp4KSwghIiKuE3itozmB220q6FxJpSRwIYS5xXcCb3SRZiTwjQerg845pIQihDC5+E7gDlebqy/dbh3laIQQIrLiMoE3utzc/toaVu45Slpi6MU7Lo8nylEJIURkxeXKllG/mu9/nNHG4h2njMCFECYXlyPwQANzUoKez7ttOgBujyRwIYS5xV0C17o5MeekJnDJxAEAnDqsl/8YgNMtJRQhhLnFXQmlwemdXfLf547m1jOG+Y+/ctNUPBrKarzTB10yAhdCmFzcJfBao3lVWovat1IKqwKrxTsfXBK4EMLs4q6EUuswEniiNeR5u8X7I7ukhCKEMLm4S+B1jd4SSmpC6D8ubFbvCFxuYgohzC7uEnhbJRQfmzECl2mEQgizi7sE3uD0JvCUDkbgLreHTSVVQbNWhBDCTOIugftG1narCnneZtzEXLi1jAse/4I3Vx2IWmxCCBFJcZfAfbVtX6mkJaUUVotiyyFvcyvfdyGEMJu4S+C+6YG+6YKh2CzKv5DHotq+TgghYlncJXC30aTK1kEC95W+27lMCCFiWtwlcJc7jBG4tfnHltkoQgizirsE7q+Bt3ETE5qX2wNUNzi7PSYhhOgOcZfAw6mBN7maV2FWSgIXQphU3CXwjmahBMpOsVMlCVwIYVJxl8DDGYH7jB+YRWV9U3eHJIQQ3SLuEng4s1B8+mUmywhcCGFacddONpwR+GNXTyA7JYHlu45QWe9Ea42S+eBCCJOJuwTu222+vRH4RRO8u/RsOVSNy6Opb3KT2sbu9UIIEaviroTSmRq4L2nXGR0MhRDCTOIugXu0xmpRYZVEfA2vnNIbXAhhQnGXwF0eHdboG8BqTDV0y2pMIYQJhVX4VUrtAWoAN+DSWhcqpXKA/wMKgD3AlVrrY90TZvjcHh3WDBQIHIHL9mpCCPPpzAh8ltZ6gta60Hh+D7BIaz0CWGQ873Eud/gjcJt/f0wZgQshzKcrJZSLgBeNxy8CF3c5mghwezxhj8B9/VKcssGxEMKEwk3gGliglFqllLrZONZHa30IwPjeO9QLlVI3K6WKlFJF5eXlXY+4A94aeHg/li/RywbHQggzCnfy83StdYlSqjfwiVJqa7gfoLV+FngWoLCwsNszZWdq4L62si6pgQshTCisoarWusT4Xga8A0wBSpVS/QCM72XdFWRndGYWit3iK6HICFwIYT4dJnClVKpSKt33GDgb2Ai8C1xnXHYdMK+7guwMt0e32ws8kH8ELglcCGFC4ZRQ+gDvGAtjbMCrWuv5SqmvgTeUUjcC+4Arui/M8HVuHrgyXiMlFCGE+XSYwLXWu4DxIY4fAWZ3R1CdUVLZwK/nbeKW04cyuSCnU7NQfPPA2xuB7yirYe3+Ki6flB+ReIUQIlJMvxLzww2HWLillCueWY7T7THmgYc7C6Xjm5jnP/YFd725LiKxCiFEJJk+gVfWN/fznr/xMAs2l5JoC+/H8q/EbGcE3mTMEW90udu8RggheoKpE7jL7eGvi3f4n9/+2hqAsBO4tRPzwOsbJYELIWKLqRP47z/YAsDkgmwCmw+GOwvFbsxCCWclZq20nBVCxBhTJ/A3ivYzqk86b9xyCr88f4z/+O8uGhvW632J3tXOCDzBGM3XNUkCF0LEFlMn8EaXhzkn9EEpxTXTBvuPD81LC+v1zc2s2h6BJxqj9DopoQghYoxpE7jL7cHt0f56d5LdyuWT8nns6glhv4fNEv4I/JfvbGDx1phYbCqEEICJE7jD5R01J9qbf4QHrxjv3+8yHLYw5oH7EvjWwzVc/8+vjydUIYToFqZN4I1Ob0kj0WY97vfw38RsZx54ZrI96PmSbTIKF0LEBvMmcN8IPMwpg6H4SijzNx5u85qMFgl89d4e33RICCGAOEjgSfbjH4HbrBYmDspi/9H6Nq9xuj1MKcjhy3vOBJp3shdCiJ5m2gTu8JdQuvYjTB3Sq90ZJi63Ji3JRt+MJADqm2Q2ihAiNpg2gTeGuIl5PNISrTS5PTS5QtfBnW4PdqvCalEk2iw0OCWBCyFig3kTeARuYkJzSeRoXVPI801uj/9mZ0qClXpZ0COEiBHmTeARuIkJkJrgTeDT/rSINfuO4QmYE37B40vZVV4XkMBtUkIRQsQM0yXwaoeTlbuPBiTwyIzAAS55ahnPf7HL/3xTSTXQ3LUwOcFK0Z5jPLGo2H+N0+3h5n8VUbTnaJfiEEKIzjJdAr/lX6u48m/LWbn7CABJXayBt2x8teFgdatrAkso+47W89An26lq8LaxLalsYMHmUi5/ZnmX4hBCiM4yXQLfZ0z5++eyPQD0y0ru0vvVOoJr2k6Xh8XbyviiuMJ/zNd2NjlgyuKu8lqO1jVRVtMY8n3dHs1NLxbx5Y6KkOeFEKKrTJfAfTVvp1uTm5ZIWhfnZZ99Yp+g5063h+tf+Jpr/r7Cfyw5wZu4R/dN9x+79eXVnPy7T9h6qHnErnVz/Xx3RR0Lt5Ry/Quy/F4I0T1Ml8AD27pmJnd9UU16kp2Xb5zqf+4M0djqmqneTodzzx/Ds9dOwmZRHK52AHDvvE3+6x5csM2/KGjrYW9ibwqj17gQQhwP0yXw+iY3V0zK58rCfB67emJE3jMtqfkXgaPFLJP3b5/BwJwUwLvq8+wT+zIqYCQe6MnFO7nt1dUA7D3SvLqzvaX6QghxvEyVwLXW1De56Z2RyAOXj2fsgMyIvG9gGeZQdYP/8cCc5JCfkZuWCMBVhQNbnfMtCCoPqI3f9upqVu6WWSpCiMgyVQJvMnqApyREth9JesAIfP/R5gTeKzUx5PWnjcwD4KKJ/QGYPryX/1yKUS8vr21kaG4qn/78dNwezefbyyMasxBCmKozk2/GiC9JRkpbDap6pSaEPH7D9AIKB2czfmAWm397DnarhRG//MiIzfteFTWN5KYnMjQvjfREm+ypKYSIOFONwH/5zkaguQ1spKS00dEwu40ErpRi/MAs72sTbNitFl65yXsj1NcdsbTaQV66dwSfmmijThK4ECLCTJPA3R7N/E3em4GRLqFYAn4h/O7i5g2R2xqBhzJ9eC6TBmdT1dDElX9bzp4j9QzplQpAaqJVluALISLONCWU99eXADBpcDYXTegf8ff/63cnsnZfJddOG8y+I3U8t3Q3mSn2jl8YIC3RxmcBte4hub4ELiUUIUTkmSaBbzxYBcDfryvEZo38Hw4XjuvPheO8vxjqjNFyZxcJ+XqG+4zu551umJogJRQhROSZooRSVu3g9ZX7OaFfBlkp4Zc1jld9o+9maecS+F3njCIrYNR+Qr8MQEbgPcHp9lBmLLYSIl6ZYgT+8CfbqWl0Max3WlQ+zzcCT+3kbJe89ES+mjsb8NbslfLW1tMSrUErSEX3e3Thdp5cvJN7zhvNDdOHkNDFtsNCxCJT/Fc93EjcXe39Ha7pw7zzuke2seKyPUl2K0l2a9DUxNREG/XtbNsmIq+4tBaA+z/ayqMLt/dwNEJ0D1OMwCcX5ADBC2a603WnFnDBuP7+aYBdJSWU6BuQ3dylcsuh1i2ChYgHpkjg4wdmsfJ/ZkcsoXZEKRXRz0pNsNHo8uBye7rlBqxozeH0kGC1MKpvOoeqpBYu4lPY2UQpZVVKrVFKvW88v08pdVAptdb4Or/7woTeGUn+mrLZpCZ6a+l1UkaJGofTTd/MJE4elMXBYw0dv0AIE+rMCPwOYAuQEXDsEa31g5ENKf74piPWNbl4dulO0hLtpCZa+WxbOU9+72T/6k0ROQ1NbpLsFgbmpFDT6OJoXRM5nViYJYQZhJXAlVL5wAXAH4A7uzWiOJTiS+CNLp5cvDPo3IFjDf6btGbgcLopLq3lpPzIdILsLg6Xm2S7lRF9vDeit5fWMG1odO6hCBEt4ZZQHgXuBlruTvBjpdRLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {
|
||||
"needs_background": "light"
|
||||
},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(all_y)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 5,
|
||||
"id": "0eec83e4",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"dt = 1 / 252 / 8"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 6,
|
||||
"id": "24f98769",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ntrain = 400"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 7,
|
||||
"id": "51c50f0b",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"device = torch.device(\"cuda:2\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 8,
|
||||
"id": "1325c8ba",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"train_y = torch.tensor(all_y[:ntrain]).view(-1).float().to(device)\n",
|
||||
"train_x = torch.linspace(0, ntrain-1, ntrain-1).view(-1).to(device) * dt\n",
|
||||
"\n",
|
||||
"test_y = torch.tensor(all_y[ntrain:]).view(-1).float().to(device)\n",
|
||||
"test_x = torch.linspace(0, test_y.shape[0], test_y.shape[0]).view(-1).to(device) * dt + train_x[-1] + dt"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 9,
|
||||
"id": "fe0b310c",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"tensor([0.1934, 0.1939, 0.1944, 0.1949, 0.1954, 0.1959, 0.1964, 0.1969, 0.1974,\n",
|
||||
" 0.1979], device='cuda:2')"
|
||||
]
|
||||
},
|
||||
"execution_count": 9,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"train_x[-10:]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 10,
|
||||
"id": "414c1737",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"tensor([0.1984, 0.1989, 0.1994, 0.1999, 0.2004, 0.2009, 0.2014, 0.2019, 0.2024,\n",
|
||||
" 0.2029], device='cuda:2')"
|
||||
]
|
||||
},
|
||||
"execution_count": 10,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"test_x[:10]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 11,
|
||||
"id": "a0812e5b",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"torch.Size([399])"
|
||||
]
|
||||
},
|
||||
"execution_count": 11,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"train_x.shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 12,
|
||||
"id": "3682fa8e",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"torch.Size([356])"
|
||||
]
|
||||
},
|
||||
"execution_count": 12,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"test_x.shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 13,
|
||||
"id": "048d01d7",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Automatic pdb calling has been turned ON\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"%pdb"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 14,
|
||||
"id": "6ae36f5d",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/home/wesley_m/gpytorch/gpytorch/utils/cholesky.py:38: NumericalWarning: A not p.d., added jitter of 1.0e-06 to the diagonal\n",
|
||||
" warnings.warn(\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Iter 1/200 - Loss: 363.720\n",
|
||||
"Iter 51/200 - Loss: 28.356\n",
|
||||
"Iter 101/200 - Loss: 12.173\n",
|
||||
"Iter 151/200 - Loss: 7.646\n",
|
||||
"Iter 1/200 - Loss: 1.053\n",
|
||||
"Iter 51/200 - Loss: 0.932\n",
|
||||
"Iter 101/200 - Loss: 0.844\n",
|
||||
"Iter 151/200 - Loss: 0.797\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/home/wesley_m/miniconda3/lib/python3.9/site-packages/torch/functional.py:445: UserWarning: torch.meshgrid: in an upcoming release, it will be required to pass the indexing argument. (Triggered internally at /opt/conda/conda-bld/pytorch_1634272204863/work/aten/src/ATen/native/TensorShape.cpp:2157.)\n",
|
||||
" return _VF.meshgrid(tensors, **kwargs) # type: ignore[attr-defined]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Iter 1/100 - Loss: 0.739\n",
|
||||
"Iter 51/100 - Loss: -1.560\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"with gpytorch.settings.max_cholesky_size(2000):\n",
|
||||
" vol = LearnGPCV(train_x, train_y, train_iters=200,\n",
|
||||
" printing=True, kernel=\"fbm\")\n",
|
||||
" vmod, vlh = TrainVolModel(train_x, vol / 8, \n",
|
||||
" train_iters=200, printing=True, kernel=\"fbm\")\n",
|
||||
" voltron, lh = TrainVoltMagpieModel(train_x, train_y[1:], \n",
|
||||
" vmod, vlh, vol / 8,\n",
|
||||
" printing=True, \n",
|
||||
" train_iters=100,\n",
|
||||
" k=200, mean_func=\"ewma\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 15,
|
||||
"id": "118e2772",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"vmod.eval();"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 16,
|
||||
"id": "15a5a6e6",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"nsample = 1000"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "e91eaf5f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"save_samples = Rollouts(train_x, train_y, test_x, voltron, \n",
|
||||
" nsample=1000)\n",
|
||||
"# voltron.vol_model.eval()\n",
|
||||
"# predvol = voltron.vol_model(test_x).sample(torch.Size((nsample, ))).exp()\n",
|
||||
"# save_samples = GeneratePrediction(train_x, train_y, test_x, \n",
|
||||
"# predvol, voltron).detach()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "0daccbfc",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"save_samples.shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "f274335f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"predictions = save_samples.exp()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "1865029c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"plt.plot(train_x.cpu(), train_y[1:].cpu())\n",
|
||||
"plt.plot(test_x.cpu(), test_y.cpu())\n",
|
||||
"plt.plot(test_x.cpu(), predictions.mean(0).cpu())\n",
|
||||
"plt.fill_between(test_x.cpu(), predictions.mean(0).cpu() - 2 * predictions.std(0).cpu(),\n",
|
||||
" predictions.mean(0).cpu() + 2 * predictions.std(0).cpu(),\n",
|
||||
" color = \"green\", alpha = 0.2\n",
|
||||
" )\n",
|
||||
"# plt.ylim((3, 8))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "7913cfc3",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"plt.plot(train_x.cpu(), vol.cpu() / 8)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "57800939",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"vmod.covar_module.vol"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "9e4d865c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.9.5"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,395 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "a2ce1a22",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import torch\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns\n",
|
||||
"import yfinance as yf\n",
|
||||
"import pandas as pd\n",
|
||||
"import numpy as np\n",
|
||||
"from pandas.tseries.holiday import USFederalHolidayCalendar\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "e977f729",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"tckrs = [\"ADBE\", \"AMAT\", \"AMZN\", \"BRK-B\", \"DAL\", \"GOOG\", \"MCD\"]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 58,
|
||||
"id": "22b13798",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"start = \"2020-03-18\"\n",
|
||||
"end = \"2021-11-01\""
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 59,
|
||||
"id": "0f530913",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"dates = pd.date_range(start, end)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 60,
|
||||
"id": "9a4d5544",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"base = \"/home/greg_b/voltron/experiments/trading/saved-outputs/\""
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 61,
|
||||
"id": "4bfb5e4b",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"full_predictions = []\n",
|
||||
"has_dates = []\n",
|
||||
"keep_dates = []\n",
|
||||
"for date in dates:\n",
|
||||
" cap = str(date)[:10]\n",
|
||||
" try:\n",
|
||||
" predictions = torch.stack([torch.load(base + tckr + \"/volt_ewma400_\"+cap+\".pt\") for tckr in tckrs])\n",
|
||||
" full_predictions.append(predictions)\n",
|
||||
" has_dates.append(cap)\n",
|
||||
" keep_dates.append(cap)\n",
|
||||
" except:\n",
|
||||
" pass\n",
|
||||
" # print(cap)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 62,
|
||||
"id": "a52b6d72",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"daily_predictions = torch.stack(full_predictions)\n",
|
||||
"num_obs = daily_predictions.shape[0]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 63,
|
||||
"id": "4f3e4de2-394f-47a1-933d-cba130fbc807",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"torch.Size([411, 7, 1000, 20])"
|
||||
]
|
||||
},
|
||||
"execution_count": 63,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"daily_predictions.shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 64,
|
||||
"id": "7fdb4753",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"[*********************100%***********************] 1 of 1 completed\n",
|
||||
"[*********************100%***********************] 1 of 1 completed\n",
|
||||
"[*********************100%***********************] 1 of 1 completed\n",
|
||||
"[*********************100%***********************] 1 of 1 completed\n",
|
||||
"[*********************100%***********************] 1 of 1 completed\n",
|
||||
"[*********************100%***********************] 1 of 1 completed\n",
|
||||
"[*********************100%***********************] 1 of 1 completed\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"bday_us = pd.offsets.CustomBusinessDay(calendar=USFederalHolidayCalendar())\n",
|
||||
"end_plus_20 = str(pd.date_range(start=end, periods=21, freq=bday_us)[-1])[:10]\n",
|
||||
"observations = torch.tensor(\n",
|
||||
" [yf.download(tckr, start=start, end=end_plus_20).Close.values for tckr in tckrs]\n",
|
||||
")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 78,
|
||||
"id": "2dd4aedf",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"lookahead = 19\n",
|
||||
"idcs = torch.arange(0, daily_predictions.shape[0] - lookahead, lookahead)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 79,
|
||||
"id": "4418f2a1-2d55-44b3-b9a8-e9e41710d33c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"returns = (observations[..., idcs+lookahead] / observations[..., idcs]) - 1."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 80,
|
||||
"id": "25080e34",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"lookahead_preds = daily_predictions[..., lookahead].permute(1, 2, 0)\n",
|
||||
"exp_means = lookahead_preds.exp().mean(1)\n",
|
||||
"exp_stds = lookahead_preds.exp().std(1)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 81,
|
||||
"id": "5e3475d1",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAABIwAAAJOCAYAAADVppwqAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAAEAAElEQVR4nOydd5gkV3W339t5ctick1arhDJCASGCACGCCCbaJhgbg42N7c8JJ5xwwAmwDTI5SYBAQgihHEE5a6XV5hwnp85dVff749yaqu7pnumZndmZ3b3v88zTPdXVFbq7bt37u+f8jtJaY7FYLBaLxWKxWCwWi8VisfhEZvsALBaLxWKxWCwWi8VisVgscwsrGFksFovFYrFYLBaLxWKxWMqwgpHFYrFYLBaLxWKxWCwWi6UMKxhZLBaLxWKxWCwWi8VisVjKsIKRxWKxWCwWi8VisVgsFoulDCsYWSwWi8VisVgsFovFYrFYyrCCkcVisVgsFovFYrFYLBaLpQwrGFkmjVLqAaXUgFIqGVr2LaVUUSk1Yv5eVEr9s1KqLbTOh5VSrlIqbf52KaU+EXp9tVJKh173/957rM/RYrHMfcZpi7RS6m0V637eLP9wxfJXm+V/av5fWdH+aKVUJvT/5cfk5CwWy3HB0bRDSqlNVfo8BaWUZ17326f/rdjOQ5VtmcViOXlRSu1RSuVMGzKglPq5UmqFec0fo6WVUv1KqbuVUqeF3vthpdRDof9blVIPK6VuVErF692P5cTFCkaWSaGUWg1cDmjgbRUvf05r3QIsAD4CXAw8rJRqCq3zqNa6WWvdDPwK8Dml1HkV22n31zF/P5yRk7FYLMctE7RF24APhdaNAe8GdlbZ1IeAfn99rfW+cPtj1jkntOyX03smFovleOVo2yGt9ZkV7c1iYBfwD6HtZIAPmn1ZLBZLLd5q2pElQBfw36HXPmdeWwYcBL5ebQNKqQ7gHmAv8F6tdWmS+7GcgFjByDJZPgg8BnyLUEcojNY6r7V+Euk8zUPEo2rrPQNsBk6fkSO1WCwnMuO1RT8DLjMdH4CrgI3AkfBKSqlGRLj+XWC9UurCmTxgi8VywnHU7VAFXwP2A38XWjZotv+Zoz5ai8VywqO1zgM/Bs6o8loOuAE4t/I1pdR84D5gE/BrWmtnqvuxnFhYwcgyWT4IXGf+3qiUWlRrRa31CHA3Mvs2BqXUy4FTgadm4DgtFsuJzXhtUR64BXhfaN3vVNnGu4A08CPgTrOexWKx1Mt0tEMAKKV+H7gM+IDW2qt4+bPAu5RSG6brwC0Wy4mJmQx7LyJmV77WBLwf2FHxUifwIPA48BtV2qBJ7cdyYmEFI0vdKKVeCawCbtBaP42EVX9ggrcdQhohn4uVUoNKqTTwBPBdYHvFe3rNOv6fjUCyWCyj1NkWfQdJ42gDrgBurrKpDwE/1Fq7wPXA+yvz9S0Wi6Ua09gOoZS6GPgn4N1a697K17XWR4Brgb+fthOwWCwnGjcrpQaBYeD1wL+FXvtj89oI8Erg1yveuwKZxP+m1lofxX4sJyBWMLJMhg8Bd4U6M9dTIy0txDLEH8TnMa11eyhX/0ykkxRmvlnH/9s8HQdvsVhOGCZsi7TWDyF+an8F3GrCsEcxJo2vQSIDAH4KpIA3z+BxWyyWE4ejbodgNA3kR8CntdbjzdT/KxLFdM50HLzFYjnheLvWuh1IAp8EHlRKLTav/bt5bTWQAyqjFZ8H/hi43feWVUr9asiM//Y692M5AYnN9gFYjg+UUg3Ae4CoUsrPv08C7bU6L0qpZuBKJJR6DFrrLqXUjcAngE9P/1FbLJYTjUm2Rd8D/gYRhir5dWTS5GdKKX9ZCkkbuXmaD9tisZxATFc7pJSKIELTw1rrcY1jtdZ9SqnPU26IbbFYLGWYqOmblFL/h0QThV/bp5T6FPBtpVSZiK21/oKp9ni3UurVWms/3bbe/fx4Js7HMvtYwchSL28HXOBlQDG0/AYqfD9MY3MWMhs2AHyz2gaVUvOAdyDmahaLxVIPb6fOtgj4IvBL4BdVtvNBxFj22tCyi4AfKaXmaa37puuALRbLCcfbmZ526G+RVJB31rnf/0SqqKmJVrRYLCcnSmbB3gZ0IMWF3hJ+XWt9t1LqEPAx4AsVr33OjOPuUUpdobXeOon9WE5QbEqapV4+hOS17tNaH/H/gP8BfhURH/9UKTWCpKB9B3gauFRrnQlt5xI/vBFpXHqA36vY12AoBDKtlPqjmT45i8Vy3FBPWwSA1rpfa31vZT6+8QtZDfxveBta61sQI8j3H7OzsVgsxyNH3Q4Z/gpYCxyp6PeklVIrK1fWWg8Dn6PcG9JisVhAIqbTiLfQZ4EPaa1rTcr/GzJuS1a+oLX+B6Ri471KqXVHuR/LCYCa2NfKYrFYLBaLxWKxWCwWi8VyMmEjjCwWi8VisVgsFovFYrFYLGVYwchisVgsFovFYrFYLBaLxVKGFYwsFovFYrFYLBaLxWKxWCxlWMHIYrFYLBaLxWKxWCwWi8VSRmziVWae+fPn69WrV8/2YVgslqPk6aef7tVaL5jt45gqti2yWE4Mjve2CGx7ZLGcCNi2yGKxzAWOpi2aE4LR6tWreeqpp2b7MCwWy1GilNo728dwNNi2yGI5MTje2yKw7ZHFciJg2yKLxTIXOJq2yKakWSwWi8VisVgsFovFYrFYyrCCkcVisVgsFovFYrFYLBaLpQwrGFksFovFYrFYLBaLxWKxWMqwgpHFYrFYLBaLxWKxWCwWi6UMKxhZLBaLxWKxWCwWy3GEUiqllHpCKfW8UmqTUurvzPJ/U0ptUUptVEr9RCnVbpavVkrllFLPmb9rZ/UELBbLcYEVjCwWi8VisVgsFovl+KIAvFZrfQ5wLnCVUupi4G7gLK312cA24NOh9+zUWp9r/j5+zI/YYrEcd1jByGKxWCwWi8VisViOI7SQNv/GzZ/WWt+ltXbM8seA5bNygBaL5YTACkYWi8VisVgsFovFcpyhlIoqpZ4DuoG7tdaPV6zyG8Dtof/XKKWeVUo9qJS6fJztfkwp9ZRS6qmenp7pP3CLxXLcYAUji8VisVgsFovFYjnO0Fq7WutzkSiii5RSZ/mvKaX+EnCA68yiw8BKrfV5wB8B1yulWmts9yta6wu11hcuWLBgRs/BYrHMbaxgZLGcCHhO9eUjO8AtHNtjmSGUUiuUUvcrpTYbc8dPmeWdSqm7lVLbzWNH6D2fVkrtUEptVUq9cfaO3mIx5Log3z3bR2GxWCxzi6EttfsylgnRWg8CDwBXASilPgS8BfhVrbU26xS01n3m+dPATuDU2Thei+WoyfeC1pDZD6Xh2T6aExorGFksxzvZg7D9f8EtQnpP+WsHb4X0rlk5rBnAAf6f1vp04GLgd5VSZwB/DtyrtV4P3Gv+x7z2PuBMpAP1JaVUdFaO3GLx2X8THLh5to/CYrEAZA+A9mb7KCwAR+6EzJ7ZPorjCqXUglAFtAbgSmCLUuoq4M+At2mtsxXrR83ztcB64ITpJFpOIrp/AVu/CPku2PdDOHLfbB/RCY0VjCyW4xknK6r68DbY8p+w+1vBDJ2TheIAqNjkt6s9eT9IRMQcQGt9WGv9jHk+AmwGlgHXAN82q30beLt5fg3wAzOjthvYAVx0TA/aYqlEO9hbr8UyR9h/s434mysUh6A0MttHcbyxBLhfKbUReBLxMLoV+B+gBbhbKfWcUupas/6rgI1KqeeBHwMf11r3z8aBWyxHRc+j4ObASUu7UTQ/4303glcK1st3B+MZy5SZwkjSYrHMCbZdC7iQWgyJDoimwBmBvidhwSXSeBYHwMtPfttdD0L/U7Dm12D/T+DU35n2wz8alFKrgfOAx4FFWuvDIKKSUmqhWW0ZUh3E54BZVrmtjwEfA1i5cuUMHrXFggi6Ns7NYpl9HDPYUFbAnXW0B24e8kdm+0iOK7TWG5G+UOXyU2qsfyNw40wfl8Uy47g5iMQkwsjNSbaFV4L0DhGJGpfJ/7u+BU2rYdV7ZvuIj2vsXdJiOR7RHpQGodALmb0Qa5Q/JwuHfi55vU7O/GUmv/2+x6Xjlt4N7txS5pVSzUiH5w+01uMlLasqy/SYBdbY0XIs0Q5VfoYWi+VYUxyQCF3rmzP7eEVpG/O9s30kFotlrqM90CWIxCF7CCJJEZzTe6DQB5l9sl6+S9r5oc2zergnAlYwsliOR9ycRBDlu2WGNNYky1tOAe2K2bVXANzJGcENbYHd18k2IykY3Cii0xxBKRVHxKLrtNY3mcVdSqkl5vUlSGlZkIiiFaG3LwcOHatjtViqoh2qa5kWi+WYUhqWCRFdmnhdy8ziFaXvUhqc7SOxWCxzHa8IKBGK8odlmVIwsh08F47cJWOXQr9EkOpieZraTND/zAlTZKgaVjCyWI5HnJw0itqBfA9EG2R5JA6JdhjeIg2qik9OMMp3ScSSBuLNYgiKOwMnMHmUUgr4OrBZa/2foZduAT5knn8I+Glo+fuUUkml1BrE3PGJY3W8FktVbDSDxTI30I7cJ+01OfscuAWIQGloto/EYrHMdUYFo4RMoKMly2Jkqzx6jjzP7BVRCTW1bIvJ0POQ7O8ExXoYWSzHI25WGsWm5aKgh4m3QmYXNK8WX6PJdMDizZL/23amRC1pd2qm2TPDZcCvAy8opZ4zy/4C+BfgBqXUR4F9wLsBtNablFI3AC8hFdZ+V2s9N9Qvy8mLtoNTi2XaKPTL4KFh8eTfqx2ZdZ7pmWfLxBT7JKXESchgLzJn+h0Wi2Wu4RblMZqE4iBoDbFmeY4HKDh4m2RiNK2UiXUnKxPqM4WTOaHvJbZFtliOR3zH/1hTkI7mo6LSeOa7RFQqTkIwcnJiFBdvkf/j7ZKeNgfQWj9E7Vye19V4z2eBz87YQVksk0FrEWGth5HFMj30PALDL8EZfzr593qO/FkRd/YpZSQ1UCmZEIu0zvYRWSyWuYpLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 1440x720 with 8 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {
|
||||
"needs_background": "light"
|
||||
},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(2, 4, figsize = (20, 10))\n",
|
||||
"ax = ax.reshape(-1)\n",
|
||||
"\n",
|
||||
"for i in range(7):\n",
|
||||
" # ax[i].plot(torch.arange(-lookahead, -lookahead+num_obs+22), observations[i])\n",
|
||||
" ax[i].plot(observations[i, lookahead:-(20+lookahead)])\n",
|
||||
" ax[i].plot(exp_means[i])\n",
|
||||
" ax[i].fill_between(\n",
|
||||
" torch.arange(num_obs), \n",
|
||||
" (exp_means - 2 * exp_stds)[i],\n",
|
||||
" (exp_means + 2 * exp_stds)[i],\n",
|
||||
" alpha = 0.4, \n",
|
||||
" color = \"orange\"\n",
|
||||
" )\n",
|
||||
" ax[i].set_title(tckrs[i])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "2d8b5bff-ea53-473b-a9cc-110845de0858",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Setup Cov and Obs"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 82,
|
||||
"id": "d0a3a4ad-33ba-45f6-b217-e0cbcb5ef272",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"returns = (observations[:, 1:] - observations[:, :-1])/(observations[:, :-1])\n",
|
||||
"cov_est = returns.cov()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 83,
|
||||
"id": "c71e3659",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def compute_strategy(preds, pxs, cov_est=None, normalize = True, interest_rate=0.03):\n",
|
||||
" excess_returns = (preds.exp().mean(-1) - pxs)/pxs - interest_rate * lookahead/252.\n",
|
||||
" if cov_est is None:\n",
|
||||
" batched_covs = torch.stack([torch.cov(excess_returns[..., i]) for i in range(excess_returns.shape[-1])])\n",
|
||||
" else:\n",
|
||||
" batched_covs = cov_est\n",
|
||||
" \n",
|
||||
" weights = torch.solve(excess_returns.t().unsqueeze(-1), batched_covs)[0]\n",
|
||||
" \n",
|
||||
" # weights = weights.clamp(min=0.)\n",
|
||||
" \n",
|
||||
" norm_constant = excess_returns.t().unsqueeze(-2).matmul(weights).sum(-1)\n",
|
||||
" \n",
|
||||
" res = weights.squeeze(-1) / norm_constant\n",
|
||||
" if normalize:\n",
|
||||
" res = res / res.abs().sum()\n",
|
||||
" \n",
|
||||
" return res\n",
|
||||
" "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "5a466c1d-9121-4a8f-9cc5-3f88d784d30e",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Loop and Run Through Strat"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 84,
|
||||
"id": "d99a2cfb-d20d-4965-87a0-b1c9921b80ba",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"lookahead = 5\n",
|
||||
"num_trading_days = daily_predictions.shape[0]\n",
|
||||
"nstock = daily_predictions.shape[1]\n",
|
||||
"volt_portfolio_value = 10000.\n",
|
||||
"hold_portfolio_value = 10000.\n",
|
||||
"start_date = pd.to_datetime(keep_dates[0]).date() - pd.to_timedelta(\"1 day\")\n",
|
||||
"logger = [[keep_dates[0], hold_portfolio_value] + [0. for _ in range(nstock)] + [\"Hold\"]]\n",
|
||||
"logger = [[keep_dates[0], volt_portfolio_value] + [0. for _ in range(nstock)] + [\"Volt\"]]\n",
|
||||
"hold_wghts = 1./nstock * torch.ones(nstock)\n",
|
||||
"for idx in range(num_trading_days):\n",
|
||||
" if idx > 0: \n",
|
||||
" ## update portfolio value\n",
|
||||
" new_px = observations[:, idx]\n",
|
||||
" holdings = volt_portfolio_value * wghts\n",
|
||||
" updates = holdings * (new_px - curr_px)/new_px\n",
|
||||
" volt_portfolio_value = volt_portfolio_value + updates.sum()\n",
|
||||
" logger.append([keep_dates[idx], volt_portfolio_value.item()] + list(updates.numpy()) + [\"Volt\"])\n",
|
||||
" old_px = curr_px\n",
|
||||
" \n",
|
||||
" holdings = hold_portfolio_value * hold_wghts\n",
|
||||
" updates = holdings * (new_px - curr_px)/new_px\n",
|
||||
" hold_portfolio_value = hold_portfolio_value + updates.sum()\n",
|
||||
" logger.append([keep_dates[idx], hold_portfolio_value.item()] + list(updates.numpy()) + [\"Hold\"])\n",
|
||||
"\n",
|
||||
" \n",
|
||||
" \n",
|
||||
" curr_px = observations[:, idx]\n",
|
||||
" pred_std = daily_predictions[idx, :, :, lookahead-1].std(-1)\n",
|
||||
" cov_est.diagonal().copy_(pred_std);\n",
|
||||
" wghts = compute_strategy(daily_predictions[idx, :, :, lookahead-1],\n",
|
||||
" curr_px,\n",
|
||||
" cov_est)\n",
|
||||
" \n",
|
||||
" \n",
|
||||
" "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 85,
|
||||
"id": "d2326570-3d05-4937-b113-267e3cb09e93",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df = pd.DataFrame(logger)\n",
|
||||
"df.columns = ['Date', 'Portfolio'] + [\"earnings\" + str(i) for i in range(daily_predictions.shape[1])] + [\"Type\"]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 86,
|
||||
"id": "2a933a5c-ff47-4645-965e-8f432e4730ec",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"Text(0.5, 1.0, 'Start Date = 2020-03-18')"
|
||||
]
|
||||
},
|
||||
"execution_count": 86,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAZIAAAEICAYAAAB1f3LfAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAABh1klEQVR4nO2dd5gURfr4P7U558BGliXDEiQHRQQUM6ioYBY9z3xfPc/T3wXDnXemOz094ymYBc98KqgoiCA557DABpbNOYep3x/VsxN2dnc2zCbq8zzzTHd1VXd1s/Q7b71JSCnRaDQajaa9uHX3BDQajUbTu9GCRKPRaDQdQgsSjUaj0XQILUg0Go1G0yG0INFoNBpNh9CCRKPRaDQdQgsSjUaj0XQILUg0XYIQ4kwhxC9CiBIhRKEQYr0QYqJx7CYhxLoOnj9JCCGFEB4t9HlUCFEnhCgzPoeFEP8WQsS04TprhBC3dmSubUUI4S2EeFMIkWbMe4cQ4gK7PrOFEAeFEJVCiNVCiP5Wx34nhNhrjD0uhPid3dgkY0ylcY45rcyn2f5CiHOEEHuEEMVCiAIhxGdCiLgWzhUjhPhSCJFl/Psl2R0PE0IsF0LkG5/3hRBBTj46TRehBYnG5Rj/8b8CXgTCgDjgMaCmk87frPBwwHIpZaAxj8uAfsC2tgiTbsADyADOBoKBPwEfmV+6QogI4FOjPQzYCiy3Gi+AG4BQ4HzgbiHEQqvjHwI7gHDgD8DHQojIFubTUv/9wFwpZQgQCxwBXmnhXCZgJXBFM8f/asw7GRgIRAOPtnA+TXcgpdQf/XHpB5gAFDdzbDhQDTQA5eZ+wEWol1Up6iX6qNWYJEACtwDpwFrjWxrnKAemOrjWo8B7dm3uwC7gWWM/FCX08oAiYzveOPaEMc9q4xr/NtqHAd8DhcAh4KoueKa7gSuM7duAX6yO+QNVwLBmxr4AvGhsD0EJ9ECr4z8Dtzcz1un+gDfwd2C/E/fjYfz7Jdm1rwDutNq/C/i2u/+m9cf2ozUSTVdwGGgQQrwthLhACBFqPiClPADcDmyQUgZI9UsWoAL1KzoEJVTuEELMtzvv2ShBNBeYYbSFGOfZ4MzEpJQNwBfAWUaTG7AU6A8kol7I/zb6/gH10rzbuMbdQgh/lBD5AIgCFgEvCyFGOrqeEOJlY9nH0We3M3MWQkSjXuj7jKaRKGFovqcKINVotx8rjHu1HntMSllm1W2Xo7HO9hdCJAohilHP7gHgaWfuqxleAi4WQoQafzdXoISLpgehBYnG5UgpS4EzUb84/wPkGevi0S2MWSOl3COlNEkpd6OWU8626/aolLJCSlnVwSlmoZaEkFIWSCk/kVJWGi/LJxxc15qLgRNSyqVSynop5XbgE2BBM/d1p5QypJnP6NYmKoTwBN4H3pZSHjSaA4ASu64lQKCDUzyKRVi2daxT/aWU6cYPggjgj8BB2s92wAsoMD4NwMsdOJ/GBWhBoukSpJQHpJQ3SSnjgRTU+vnzzfUXQkw2DLp5QogSlNYSYdcto5OmF4dalkII4SeEeM0wbJeils1ChBDuzYztD0y21iyAa1G2l05FCOEGvAvUAndbHSoH7A3QQYC11oAQ4m6UlneRlLLGmbFCiH1CiHLjc5az1wKQUhYCbwNfCCE8hBBnWZ1rn33/ZvgvSqMNNK6TCrzn5FhNF6EFiabLMX5Jv4USKKA0FXs+AL4EEqSUwcCrKKOxzama2XYa4+V8CWrJCuC3wFBgspQyCMuSmfna9tfJAH6y0ywCpJR3NHO9V61epvafZl+uxpLUmyhj8xVSyjqrw/uAMVZ9/VGG6X1WbYuBh4DZUspMu7HJQghrDWSMeayUcqRxPwFSyp9b6+8AD9SSX5CU8merczW3dGbPGOA1Q/MsR/0dXOjkWE0XoQWJxuUIIYYJIX4rhIg39hNQtoSNRpccIF4I4WU1LBAolFJWCyEmAde0cpk8lAdQspNz8hRCDEctmfUD/ml13SqgWAgRBjxiNzTH7hpfAUOEENcb5/QUQkw0zt0EKeXtVi9T+09LL9dXUPagSxws5X0GpAghrhBC+AB/Bnabl76EENcCfwPOlVIes5vPYWAn8IgQwkcIcRkwGrU852j+LfYXQlwuhBgqhHAzPLn+CewwtBOHGHP2Nna9jX0zW4BbhRC+QghflGPBLvtzaLqZ7rb260/f/6CWjj4CTqKM6CeB11C/UkGtgX+NWl7KN9oWAGmoJZOvUAbv94xjSSjNwMPuOo+jBEoxMMXBPB4F6lDLMxUo19SXgTirPrHAGqPPYeDX1tcCphrtRcALRttQY/55qHX8H4Gxnfj8+htzMHuLmT/XWvWZg7JFVBnzT7I6dtzqvs2fV62OJxljqlBeZ3NamU+z/YF7jOtVANnAMqB/K+eT9h+rYwOA/xnPtRDlKjy4u/+m9cf2I4x/LI1Go9Fo2oVe2tJoNBpNh9CCRKPRaDQdQgsSjUaj0XQILUg0Go1G0yHakuyuT3D++efLlStXdvc0NBqNprdhH8fVyGmnkeTn53f3FDQajaZPcdoJEo1Go9F0LlqQaDQajaZDaEGi0Wg0mg5x2hnbHVFXV0dmZibV1dXdPRWX4uPjQ3x8PJ6ent09FY1G04fQggTIzMwkMDCQpKQkVJLVvoeUkoKCAjIzMxkwYEB3T0ej0fQh9NIWUF1dTXh4eJ8VIgBCCMLDw/u81qXRaLoeLUgM+rIQMXM63KNGo+l6tCDRaDSansCxNZDjbOHInoW2kXQiBQUFzJ49G4Ds7Gzc3d2JjIwEYPPmzXh5ebU0XKPR9CRyD8IXd8I1/wX/cNdeq6Ee3pmnth8tce21XIAWJJ1IeHg4O3fuBODRRx8lICCABx54oHsnpdFo2sf+L+DkNsjYBMNcXN335NambQe+gm9+B4tXQmh/116/g+ilLRdSVVXFgAEDqKtT5bVLS0tJSkqirq6OmTNn8n//939MmzaNlJQUNm/eDEBFRQWLFy9m4sSJnHHGGXzxxRfdeQsazelL+gb1nbu/Y+dpqIPSLPjmQVh+PSy7Vmk7UoLJpPqk/mjpX1upjn37MJRlwfd/6tj1uwCtkbgQX19fZs6cyddff838+fNZtmwZV1xxRWMcR0VFBb/88gtr165l8eLF7N27lyeeeIJZs2axZMkSiouLmTRpEnPmzMHf37+b70ajOY1oqIfMLWo794B6sbfXWWXDS7DqEdu2I99DYD+orYDfHYX0jZZjBUfA0x+K09X+ye3tu24XojUSF3PrrbeydOlSAJYuXcrNN9/ceGzRokUAzJgxg9LSUoqLi/nuu+948sknGTt2LDNnzqS6upr09PRumbtGc9qSsxdqy8HdC/Z+DI+Fqpd/ezALiaEXwSPF8MAR8AuD4jSozIfnRsLxn6D/dNUv7xAUHFXbseOgurTDt+NqtCBxMdOnT+fEiRP89NNPNDQ0kJKS0njM3h1XCIGUkk8++YSdO3eyc+dO0tPTGT58eFdPW6M5PaitcPyiNi9rnXk/uHsDEjI2t+P8lVCeAwPOhoXvK60mIArmPmHpU3pSfY+8TH2XZEDhMbUdOxZqSpVG1BHyDsHXD0D2no6dpxlcJkiEEEuEELlCiL1WbWOFEBuFEDuFEFuFEJOsjj0shDgqhDgkhJhr1T5eCLHHOPaCMN6+QghvIcRyo32TECLJVffSUW644QYWLVpko40ALF++HIB169YRHBxMcHAwc+fO5cUXX0Qafzg7duzo8vlqNKcFeYfhb3Hw4cKmx9I3QHAinPMw/CkXAqKh7FTbzl9bCX+LgaztEJZsuzSWcgX8LtW2f/JM8A1V9pSi4+AdDKFJgFTaUXvZuhRemgTb34Ys17xPXKmRvAWcb9f2NPCYlHIs8GdjHyHECGAhMNIY87IQwt0Y8wpwGzDY+JjPeQtQJKUcBDwHPOWqG+ko1157LUVFRY1LWWZCQ0OZNm0at99+O2+++SYAf/rTn6irq2P06NGkpKTwpz/1fEObRtOrqCyEtc/CigcBCWnroSzbcrzoBBxaAYNmW9oCY2wFyaldqk9LHLYqoBfmIC2RfwTEjIVhF6vlrojBEBirBEnhMTXGO0j1rSlz7t72fgIHv27aFjkM7tsP425w7jxtxGXGdinlWgdaggSMJ0MwkGVszwOWSSlrgONCiKPAJCHECSBISrkBQAjxDjAfWGGMedQY/zHwbyGEkLKjOmDn8OijjzZur1u3jgULFhASEmLT54orruDvf/+7TZuvry+vvfZaF8xQozkNqauGty9RNhAADx+or4Yv7oar3gEvP/ULXko4+0HLuKBYKEpT25WF8NoMtf3bwxAY7fha+z61bPuEOO7z659sDflBhiCpKoK4ceAdqNqrS9Wxlqgph48Xq+0x18AFT4FXAGTthDFXQ0Bky+M7QFd7bf0f8K0Q4lmUNjTNaI8DrNwWyDTa6oxt+3bzmAwAKWW9EKIECAealEAUQtyG0mpITEzspFtxjnvuuYcVK1bwzTffdOl1NRqNA7a9pYSIhy/UV8HoqwChln3W/RNm/dGiDVi/uANj1HJXVTHs/MDS/tOTMPY6iB9vex1TAxxbq5awwgfBqCubn5P1kldQLBw1jPoTbwGfYLVd44TBfd9nlu1dH8DAWRAzGmrLIG588+M6ga4WJHcA90kpPxFCXAW8CczBcS1g2UI7rRyzbZTydeB1gAkTJnSpxvLiiy86bF+zZk1XTkOj6dvU14BwB/dWXmknfobQATByPqx7TmkK5/1Fvag3vAQzH1ZutyF2PziDYpWW8JQRGOgTAkMvhK1L1OeOX1Tg4uirwcsfsndDTQkMOd8QVk5iLbwSp4FsUNvOCJLDK8E/CnyClNdX5maoq1TH4iY4P4d20NVeWzcCZn3vv4DZ2J4JJFj1i0cte2Ua2/btNmOEEB6opbJCl8xao9H0XKSEFyfA8uuaHju5DcrzLP0yNkPCZKVFeAfBGMNumXSLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {
|
||||
"needs_background": "light"
|
||||
},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"sns.lineplot(x=df.index, y='Portfolio', hue='Type', data=df)\n",
|
||||
"sns.despine()\n",
|
||||
"plt.title(\"Start Date = \" + keep_dates[0])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "e7fd1f6a-464d-40c4-aed9-abcf0daadabb",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "84f8bdb7-820f-4cc1-8577-cffae2c8f720",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
File diff suppressed because it is too large.
Load diff
@@ -0,0 +1,513 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "be6383aa",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Warning no robinhood utils.\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import torch\n",
|
||||
"import yfinance as yf\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns\n",
|
||||
"import numpy as np\n",
|
||||
"import torch\n",
|
||||
"import pandas as pd\n",
|
||||
"import os\n",
|
||||
"import gpytorch\n",
|
||||
"import argparse\n",
|
||||
"import datetime\n",
|
||||
"\n",
|
||||
"from botorch.models import SingleTaskGP\n",
|
||||
"from botorch.optim.fit import fit_gpytorch_torch\n",
|
||||
"from gpytorch.likelihoods import GaussianLikelihood\n",
|
||||
"from gpytorch.mlls import ExactMarginalLogLikelihood\n",
|
||||
"from gpytorch.means import ConstantMean, LinearMean\n",
|
||||
"from gpytorch.kernels import SpectralMixtureKernel, MaternKernel, RBFKernel, ScaleKernel\n",
|
||||
"from voltron.means import EWMAMean, DEWMAMean, TEWMAMean\n",
|
||||
"from voltron.train_utils import LearnGPCV, TrainVolModel, TrainVoltMagpieModel, TrainBasicModel\n",
|
||||
"from voltron.models import VoltMagpie\n",
|
||||
"from voltron.means import LogLinearMean\n",
|
||||
"\n",
|
||||
"from voltron.rollout_utils import GeneratePrediction, Rollouts\n",
|
||||
"from voltron.data import make_ticker_list, DataGetter, GetStockHistory"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "8cce2850",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Mon Dec 20 07:56:41 2021 \n",
|
||||
"+-----------------------------------------------------------------------------+\n",
|
||||
"| NVIDIA-SMI 465.19.01 Driver Version: 465.19.01 CUDA Version: 11.3 |\n",
|
||||
"|-------------------------------+----------------------+----------------------+\n",
|
||||
"| GPU Name Persistence-M| Bus-Id Disp.A | Volatile Uncorr. ECC |\n",
|
||||
"| Fan Temp Perf Pwr:Usage/Cap| Memory-Usage | GPU-Util Compute M. |\n",
|
||||
"| | | MIG M. |\n",
|
||||
"|===============================+======================+======================|\n",
|
||||
"| 0 NVIDIA TITAN RTX On | 00000000:1A:00.0 Off | N/A |\n",
|
||||
"| 41% 33C P5 53W / 280W | 3MiB / 24220MiB | 0% Default |\n",
|
||||
"| | | N/A |\n",
|
||||
"+-------------------------------+----------------------+----------------------+\n",
|
||||
"| 1 NVIDIA TITAN RTX On | 00000000:1B:00.0 Off | N/A |\n",
|
||||
"| 41% 24C P8 10W / 280W | 3MiB / 24220MiB | 0% Default |\n",
|
||||
"| | | N/A |\n",
|
||||
"+-------------------------------+----------------------+----------------------+\n",
|
||||
"| 2 NVIDIA TITAN RTX On | 00000000:3D:00.0 Off | N/A |\n",
|
||||
"| 41% 29C P2 58W / 280W | 1402MiB / 24220MiB | 11% Default |\n",
|
||||
"| | | N/A |\n",
|
||||
"+-------------------------------+----------------------+----------------------+\n",
|
||||
"| 3 NVIDIA TITAN RTX On | 00000000:3E:00.0 Off | N/A |\n",
|
||||
"| 41% 25C P8 16W / 280W | 3MiB / 24220MiB | 0% Default |\n",
|
||||
"| | | N/A |\n",
|
||||
"+-------------------------------+----------------------+----------------------+\n",
|
||||
"| 4 NVIDIA TITAN RTX On | 00000000:88:00.0 Off | N/A |\n",
|
||||
"| 41% 24C P8 13W / 280W | 3MiB / 24220MiB | 0% Default |\n",
|
||||
"| | | N/A |\n",
|
||||
"+-------------------------------+----------------------+----------------------+\n",
|
||||
"| 5 NVIDIA TITAN RTX On | 00000000:89:00.0 Off | N/A |\n",
|
||||
"| 41% 24C P8 18W / 280W | 3MiB / 24220MiB | 0% Default |\n",
|
||||
"| | | N/A |\n",
|
||||
"+-------------------------------+----------------------+----------------------+\n",
|
||||
"| 6 NVIDIA TITAN RTX On | 00000000:B1:00.0 Off | N/A |\n",
|
||||
"| 41% 25C P8 4W / 280W | 3MiB / 24220MiB | 0% Default |\n",
|
||||
"| | | N/A |\n",
|
||||
"+-------------------------------+----------------------+----------------------+\n",
|
||||
"| 7 NVIDIA TITAN RTX On | 00000000:B2:00.0 Off | N/A |\n",
|
||||
"| 41% 25C P8 15W / 280W | 3MiB / 24220MiB | 0% Default |\n",
|
||||
"| | | N/A |\n",
|
||||
"+-------------------------------+----------------------+----------------------+\n",
|
||||
" \n",
|
||||
"+-----------------------------------------------------------------------------+\n",
|
||||
"| Processes: |\n",
|
||||
"| GPU GI CI PID Type Process name GPU Memory |\n",
|
||||
"| ID ID Usage |\n",
|
||||
"|=============================================================================|\n",
|
||||
"| 2 N/A N/A 61460 C ...y_m/miniconda3/bin/python 1399MiB |\n",
|
||||
"+-----------------------------------------------------------------------------+\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"!nvidia-smi"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"id": "bb53d987",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"[*********************100%***********************] 1 of 1 completed\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"all_y = yf.download(\"XOM\", start = \"2020-06-01\", end=\"2020-11-01\", interval=\"1h\").Close.values"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"id": "5a9b0817",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"[<matplotlib.lines.Line2D at 0x7fd1240f3fd0>]"
|
||||
]
|
||||
},
|
||||
"execution_count": 4,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAXAAAAD4CAYAAAD1jb0+AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/MnkTPAAAACXBIWXMAAAsTAAALEwEAmpwYAAA0FklEQVR4nO3dd3xb1dnA8d/R8N6xM53E2QmEDOIMSAKEEHbLHm2hlPFCKaW0lPKSt6Uv3ZSyKaNASynzhQINM4SEBAIJCc7ecfZw4pHE27I1zvuHrmTJlm05lmVd8Xw/H38s3XslPWY8Pn7uOc9RWmuEEEKYj6WnAxBCCHF8JIELIYRJSQIXQgiTkgQuhBAmJQlcCCFMyhbND8vNzdUFBQXR/EghhDC9VatWVWit81oej2oCLygooKioKJofKYQQpqeU2hvquJRQhBDCpCSBCyGESUkCF0IIkworgSul9iilNiil1iqlioxj9ymlDhrH1iqlzu/eUIUQQgTqzE3MWVrrihbHHtFaPxjJgIQQQoRHSihCCGFS4SZwDSxQSq1SSt0ccPzHSqn1Sql/KKWyQ71QKXWzUqpIKVVUXl7e5YCFEEJ4hZvAp2utTwbOA25TSp0GPA0MAyYAh4CHQr1Qa/2s1rpQa12Yl9dqHnrUfLjhEEdqG3vs84UQItLCSuBa6xLjexnwDjBFa12qtXZrrT3Ac8CU7guzaw5VNfCjV1Zzx+trezoUIYSImA4TuFIqVSmV7nsMnA1sVEr1C7jsEmBj94TYdUfrmgCokBG4ECKOhDMLpQ/wjlLKd/2rWuv5SqmXlFIT8NbH9wC3dFeQXbX/aD0AaYlR7RwghBDdqsOMprXeBYwPcfzabokowj7dWsoPX14NQHqSJHAhRPyI+2mEC7eU+R+nJdl7MBIhhIisuE/gdY0u/2MZgQsh4klcJ/D6Jhdf7mhePKq1pri0JuiYEEKYVVwn8FdX7KOitolZo7zzz5tcmjmPfM73nl/Rw5EJIUTXxXUCL69pJMFq4YXrpzAkNxWHy+0/1xjwWAghzCiuE3hNo8tf906wWvhg/SH/ucNVjp4KSwghIiKuE3itozmB220q6FxJpSRwIYS5xXcCb3SRZiTwjQerg845pIQihDC5+E7gDlebqy/dbh3laIQQIrLiMoE3utzc/toaVu45Slpi6MU7Lo8nylEJIURkxeXKllG/mu9/nNHG4h2njMCFECYXlyPwQANzUoKez7ttOgBujyRwIYS5xV0C17o5MeekJnDJxAEAnDqsl/8YgNMtJRQhhLnFXQmlwemdXfLf547m1jOG+Y+/ctNUPBrKarzTB10yAhdCmFzcJfBao3lVWovat1IKqwKrxTsfXBK4EMLs4q6EUuswEniiNeR5u8X7I7ukhCKEMLm4S+B1jd4SSmpC6D8ubFbvCFxuYgohzC7uEnhbJRQfmzECl2mEQgizi7sE3uD0JvCUDkbgLreHTSVVQbNWhBDCTOIugftG1narCnneZtzEXLi1jAse/4I3Vx2IWmxCCBFJcZfAfbVtX6mkJaUUVotiyyFvcyvfdyGEMJu4S+C+6YG+6YKh2CzKv5DHotq+TgghYlncJXC30aTK1kEC95W+27lMCCFiWtwlcJc7jBG4tfnHltkoQgizirsE7q+Bt3ETE5qX2wNUNzi7PSYhhOgOcZfAw6mBN7maV2FWSgIXQphU3CXwjmahBMpOsVMlCVwIYVJxl8DDGYH7jB+YRWV9U3eHJIQQ3SLuEng4s1B8+mUmywhcCGFacddONpwR+GNXTyA7JYHlu45QWe9Ea42S+eBCCJOJuwTu222+vRH4RRO8u/RsOVSNy6Opb3KT2sbu9UIIEaviroTSmRq4L2nXGR0MhRDCTOIugXu0xmpRYZVEfA2vnNIbXAhhQnGXwF0eHdboG8BqTDV0y2pMIYQJhVX4VUrtAWoAN+DSWhcqpXKA/wMKgD3AlVrrY90TZvjcHh3WDBQIHIHL9mpCCPPpzAh8ltZ6gta60Hh+D7BIaz0CWGQ873Eud/gjcJt/f0wZgQshzKcrJZSLgBeNxy8CF3c5mghwezxhj8B9/VKcssGxEMKEwk3gGliglFqllLrZONZHa30IwPjeO9QLlVI3K6WKlFJF5eXlXY+4A94aeHg/li/RywbHQggzCnfy83StdYlSqjfwiVJqa7gfoLV+FngWoLCwsNszZWdq4L62si6pgQshTCisoarWusT4Xga8A0wBSpVS/QCM72XdFWRndGYWit3iK6HICFwIYT4dJnClVKpSKt33GDgb2Ai8C1xnXHYdMK+7guwMt0e32ws8kH8ELglcCGFC4ZRQ+gDvGAtjbMCrWuv5SqmvgTeUUjcC+4Arui/M8HVuHrgyXiMlFCGE+XSYwLXWu4DxIY4fAWZ3R1CdUVLZwK/nbeKW04cyuSCnU7NQfPPA2xuB7yirYe3+Ki6flB+ReIUQIlJMvxLzww2HWLillCueWY7T7THmgYc7C6Xjm5jnP/YFd725LiKxCiFEJJk+gVfWN/fznr/xMAs2l5JoC+/H8q/EbGcE3mTMEW90udu8RggheoKpE7jL7eGvi3f4n9/+2hqAsBO4tRPzwOsbJYELIWKLqRP47z/YAsDkgmwCmw+GOwvFbsxCCWclZq20nBVCxBhTJ/A3ivYzqk86b9xyCr88f4z/+O8uGhvW632J3tXOCDzBGM3XNUkCF0LEFlMn8EaXhzkn9EEpxTXTBvuPD81LC+v1zc2s2h6BJxqj9DopoQghYoxpE7jL7cHt0f56d5LdyuWT8nns6glhv4fNEv4I/JfvbGDx1phYbCqEEICJE7jD5R01J9qbf4QHrxjv3+8yHLYw5oH7EvjWwzVc/8+vjydUIYToFqZN4I1Ob0kj0WY97vfw38RsZx54ZrI96PmSbTIKF0LEBvMmcN8IPMwpg6H4SijzNx5u85qMFgl89d4e33RICCGAOEjgSfbjH4HbrBYmDspi/9H6Nq9xuj1MKcjhy3vOBJp3shdCiJ5m2gTu8JdQuvYjTB3Sq90ZJi63Ji3JRt+MJADqm2Q2ihAiNpg2gTeGuIl5PNISrTS5PTS5QtfBnW4PdqvCalEk2iw0OCWBCyFig3kTeARuYkJzSeRoXVPI801uj/9mZ0qClXpZ0COEiBHmTeARuIkJkJrgTeDT/rSINfuO4QmYE37B40vZVV4XkMBtUkIRQsQM0yXwaoeTlbuPBiTwyIzAAS55ahnPf7HL/3xTSTXQ3LUwOcFK0Z5jPLGo2H+N0+3h5n8VUbTnaJfiEEKIzjJdAr/lX6u48m/LWbn7CABJXayBt2x8teFgdatrAkso+47W89An26lq8LaxLalsYMHmUi5/ZnmX4hBCiM4yXQLfZ0z5++eyPQD0y0ru0vvVOoJr2k6Xh8XbyviiuMJ/zNd2NjlgyuKu8lqO1jVRVtMY8n3dHs1NLxbx5Y6KkOeFEKKrTJfAfTVvp1uTm5ZIWhfnZZ99Yp+g5063h+tf+Jpr/r7Cfyw5wZu4R/dN9x+79eXVnPy7T9h6qHnErnVz/Xx3RR0Lt5Ry/Quy/F4I0T1Ml8AD27pmJnd9UU16kp2Xb5zqf+4M0djqmqneTodzzx/Ds9dOwmZRHK52AHDvvE3+6x5csM2/KGjrYW9ibwqj17gQQhwP0yXw+iY3V0zK58rCfB67emJE3jMtqfkXgaPFLJP3b5/BwJwUwLvq8+wT+zIqYCQe6MnFO7nt1dUA7D3SvLqzvaX6QghxvEyVwLXW1De56Z2RyAOXj2fsgMyIvG9gGeZQdYP/8cCc5JCfkZuWCMBVhQNbnfMtCCoPqI3f9upqVu6WWSpCiMgyVQJvMnqApyREth9JesAIfP/R5gTeKzUx5PWnjcwD4KKJ/QGYPryX/1yKUS8vr21kaG4qn/78dNwezefbyyMasxBCmKozk2/GiC9JRkpbDap6pSaEPH7D9AIKB2czfmAWm397DnarhRG//MiIzfteFTWN5KYnMjQvjfREm+ypKYSIOFONwH/5zkaguQ1spKS00dEwu40ErpRi/MAs72sTbNitFl65yXsj1NcdsbTaQV66dwSfmmijThK4ECLCTJPA3R7N/E3em4GRLqFYAn4h/O7i5g2R2xqBhzJ9eC6TBmdT1dDElX9bzp4j9QzplQpAaqJVluALISLONCWU99eXADBpcDYXTegf8ff/63cnsnZfJddOG8y+I3U8t3Q3mSn2jl8YIC3RxmcBte4hub4ELiUUIUTkmSaBbzxYBcDfryvEZo38Hw4XjuvPheO8vxjqjNFyZxcJ+XqG+4zu551umJogJRQhROSZooRSVu3g9ZX7OaFfBlkp4Zc1jld9o+9maecS+F3njCIrYNR+Qr8MQEbgPcHp9lBmLLYSIl6ZYgT+8CfbqWl0Max3WlQ+zzcCT+3kbJe89ES+mjsb8NbslfLW1tMSrUErSEX3e3Thdp5cvJN7zhvNDdOHkNDFtsNCxCJT/Fc93EjcXe39Ha7pw7zzuke2seKyPUl2K0l2a9DUxNREG/XtbNsmIq+4tBaA+z/ayqMLt/dwNEJ0D1OMwCcX5ADBC2a603WnFnDBuP7+aYBdJSWU6BuQ3dylcsuh1i2ChYgHpkjg4wdmsfJ/ZkcsoXZEKRXRz0pNsNHo8uBye7rlBqxozeH0kGC1MKpvOoeqpBYu4lPY2UQpZVVKrVFKvW88v08pdVAptdb4Or/7woTeGUn+mrLZpCZ6a+l1UkaJGofTTd/MJE4elMXBYw0dv0AIE+rMCPwOYAuQEXDsEa31g5ENKf74piPWNbl4dulO0hLtpCZa+WxbOU9+72T/6k0ROQ1NbpLsFgbmpFDT6OJoXRM5nViYJYQZhJXAlVL5wAXAH4A7uzWiOJTiS+CNLp5cvDPo3IFjDf6btGbgcLopLq3lpPzIdILsLg6Xm2S7lRF9vDeit5fWMG1odO6hCBEt4ZZQHgXuBlruTvBjpdRLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {
|
||||
"needs_background": "light"
|
||||
},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(all_y)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 5,
|
||||
"id": "0eec83e4",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"dt = 1 / 252 / 8"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 6,
|
||||
"id": "c43156f5",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ntrain = 400"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 7,
|
||||
"id": "1325c8ba",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"train_y = torch.tensor(all_y[:ntrain]).view(-1).float().cuda()\n",
|
||||
"train_x = torch.linspace(0, ntrain-1, ntrain-1).view(-1).cuda() * dt\n",
|
||||
"\n",
|
||||
"test_y = torch.tensor(all_y[ntrain:]).view(-1).float().cuda()\n",
|
||||
"test_x = torch.linspace(0, test_y.shape[0], test_y.shape[0]).view(-1).cuda() * dt + train_x[-1] + dt"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 8,
|
||||
"id": "fe0b310c",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"tensor([0.1934, 0.1939, 0.1944, 0.1949, 0.1954, 0.1959, 0.1964, 0.1969, 0.1974,\n",
|
||||
" 0.1979], device='cuda:0')"
|
||||
]
|
||||
},
|
||||
"execution_count": 8,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"train_x[-10:]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 9,
|
||||
"id": "414c1737",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"tensor([0.1984, 0.1989, 0.1994, 0.1999, 0.2004, 0.2009, 0.2014, 0.2019, 0.2024,\n",
|
||||
" 0.2029], device='cuda:0')"
|
||||
]
|
||||
},
|
||||
"execution_count": 9,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"test_x[:10]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 10,
|
||||
"id": "a0812e5b",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"torch.Size([399])"
|
||||
]
|
||||
},
|
||||
"execution_count": 10,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"train_x.shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 11,
|
||||
"id": "3682fa8e",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"torch.Size([356])"
|
||||
]
|
||||
},
|
||||
"execution_count": 11,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"test_x.shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 12,
|
||||
"id": "6ae36f5d",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/home/wesley_m/miniconda3/lib/python3.9/site-packages/torch/functional.py:445: UserWarning: torch.meshgrid: in an upcoming release, it will be required to pass the indexing argument. (Triggered internally at /opt/conda/conda-bld/pytorch_1634272204863/work/aten/src/ATen/native/TensorShape.cpp:2157.)\n",
|
||||
" return _VF.meshgrid(tensors, **kwargs) # type: ignore[attr-defined]\n",
|
||||
"/home/wesley_m/gpytorch/gpytorch/utils/cholesky.py:38: NumericalWarning: A not p.d., added jitter of 1.0e-06 to the diagonal\n",
|
||||
" warnings.warn(\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Iter 1/500 - Loss: 6.738\n",
|
||||
"Iter 51/500 - Loss: 0.607\n",
|
||||
"Iter 101/500 - Loss: 0.541\n",
|
||||
"Iter 151/500 - Loss: 0.525\n",
|
||||
"Iter 201/500 - Loss: 0.522\n",
|
||||
"Iter 251/500 - Loss: 0.520\n",
|
||||
"Iter 301/500 - Loss: 0.519\n",
|
||||
"Iter 351/500 - Loss: 0.518\n",
|
||||
"Iter 401/500 - Loss: 0.517\n",
|
||||
"Iter 451/500 - Loss: 0.516\n",
|
||||
"Iter 1/200 - Loss: 1.960\n",
|
||||
"Iter 51/200 - Loss: 1.389\n",
|
||||
"Iter 101/200 - Loss: 0.889\n",
|
||||
"Iter 151/200 - Loss: 0.395\n",
|
||||
"Iter 1/100 - Loss: 0.738\n",
|
||||
"Iter 51/100 - Loss: -1.494\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"with gpytorch.settings.max_cholesky_size(2000):\n",
|
||||
" vol = LearnGPCV(train_x, train_y, train_iters=500,\n",
|
||||
" printing=True)\n",
|
||||
" vmod, vlh = TrainVolModel(train_x, vol / 8, \n",
|
||||
" train_iters=200, printing=True)\n",
|
||||
" voltron, lh = TrainVoltMagpieModel(train_x, train_y[1:], \n",
|
||||
" vmod, vlh, vol / 8,\n",
|
||||
" printing=True, \n",
|
||||
" train_iters=100,\n",
|
||||
" k=200, mean_func=\"ewma\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 13,
|
||||
"id": "118e2772",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"vmod.eval();"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 14,
|
||||
"id": "d2c48885",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"nsample = 1000"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 15,
|
||||
"id": "e91eaf5f",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/home/wesley_m/gpytorch/gpytorch/utils/cholesky.py:38: NumericalWarning: A not p.d., added jitter of 1.0e-04 to the diagonal\n",
|
||||
" warnings.warn(\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"save_samples = Rollouts(train_x, train_y, test_x, voltron, \n",
|
||||
" nsample=1000)\n",
|
||||
"# voltron.vol_model.eval()\n",
|
||||
"# predvol = voltron.vol_model(test_x).sample(torch.Size((nsample, ))).exp()\n",
|
||||
"# save_samples = GeneratePrediction(train_x, train_y, test_x, \n",
|
||||
"# predvol, voltron).detach()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 16,
|
||||
"id": "0daccbfc",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"torch.Size([1000, 356])"
|
||||
]
|
||||
},
|
||||
"execution_count": 16,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"save_samples.shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 17,
|
||||
"id": "f274335f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"predictions = save_samples.exp()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 18,
|
||||
"id": "1865029c",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"<matplotlib.collections.PolyCollection at 0x7fd1182fafd0>"
|
||||
]
|
||||
},
|
||||
"execution_count": 18,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAXAAAAD4CAYAAAD1jb0+AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/MnkTPAAAACXBIWXMAAAsTAAALEwEAmpwYAABAkElEQVR4nO3dd3xb9bn48c9XW5b3Xtk7IdsEQsggCRsKXCjcFjpuB6PtbX90997uSdvbRUtpoQsoBTqgpYwAIZOEDCdk7zjDsR3vvbS+vz+OJI/IthwvyTzv1ysvHR2do/NYmEdfP+c7lNYaIYQQscc00gEIIYS4OJLAhRAiRkkCF0KIGCUJXAghYpQkcCGEiFGW4bxYenq6Hj9+/HBeUgghYt6uXbuqtNYZ3fcPawIfP348hYWFw3lJIYSIeUqpM+H2SwlFCCFilCRwIYSIUZLAhRAiRkWUwJVSp5VS+5VSe5RShYF931RKlQT27VFK3TC0oQohhOisPzcxr9JaV3Xb9zOt9f8NZkBCCCEiIyUUIYSIUZEmcA28rpTapZS6t9P+Tyml9iml/qCUSgl3olLqXqVUoVKqsLKycsABCyGEMESawJdorRcA1wOfVEotAx4FJgHzgDLgJ+FO1Fo/prUu0FoXZGRc0A99WNQ2u3lpX+mIXFsIIYZKRAlca10aeKwAXgAWaa3LtdY+rbUfeBxYNHRhDszHnizkU395h+qm9pEORQghBk2fCVwp5VJKJQS3gWuAA0qpnE6H3QYcGJoQB+5wWQMAPlm8QggxikTSCyULeEEpFTz+L1rrNUqpp5RS8zDq46eB+4YqyIFo8/hocfsA8PgkgQshRo8+E7jWugiYG2b/B4YkokF2x2+2hrbdXv8IRiKEEINr1HcjPFDSENr2+CSBCyFGj1GdwLsnbGmBCyFGk1GdwNccON/ludvn55X9ZVQ0tI1QREIIMXhGdQL/5brjKAVjU+MAqG/x8Imnd3PP77ePcGRCCDFwozqBVzS2c/dlY/nJncY92JOVTQAcK28aybCEEGJQjNoErrWmqc1LosOK1Wz8mN99+fAIRyWEEINn1Cbwdq8fr18T77BgM4/aH1MI8S42ajNbY5sXgAS7BZtFdXnNYlLhThFCiJgyahN4U3sggTus2MzmLq8lOIZ1LWchhBgSozKBby+q5r6nCgGIt1uwdmuBe/0ypF4IEftGXVP0nbO13PXYttDzBIcldBMTIC/ZSU2zeyRCE0KIQTXqWuCHyhq6PJ+ek9glgd88NxevX0ZkCiFi36hrgVtNRrJOclq5dHwKSU5raAj9p66ajFJSQhFCjA6jLoEHb15u+PwKUlw2AGwWEye/fwMmBQ+/eQKtwefXmKU3ihAiho26EkpzIIG77F2/m8wmhVIKi9lI2jIzoRAi1o26BN7U7sVmMWGzhP/Rgn3AfVJGEULEuFGZwOPtPVeGLIEbml5ZnUcIEeNGXQJv7iOBW4MlFOmJIoSIcaMugbe4fcTZzD2+HrxxWdvsprimZbjCEkKIQTfqeqF4/bpLv+/ugt0Mb3h4Mx6f5vRDNw5XaEIIMahGXQvc20f3wI5eKEYNXGuphQshYtOoS+A+v7/X2Qa7J/cWt2+oQxJCiCEx6hK419d7C7x7eaWu1TPUIQkhxJAYdQnc59ehMkk43Vvn9S2SwIUQsWnUJXCjBt7zj9U9ude1ysyEQojYNOoSuM+ve62BW7old2mBCyFi1ahL4JH2Qgmqlxq4ECJGjboE3lcvlGAL3GE1HuUmphAiVo26BN5XCzw4W+GqGVnYzCbqpIQihIhRoy6B91UDv2JyGvctn8j3b51NUpyVermJKYSIUaNvKL2v914odouZr1w/A4Bkp1Vq4EKImPWua4F35rJbaGqXkZhCiNg0+hK41ph7GcjTmdWs8MrKPEKIGDX6Eng/WuBmk5IFjoUQMSuiGrhS6jTQCPgAr9a6QCmVCjwHjAdOA3dqrWuHJszIeX3+iBcrtppNoV4pQggRa/rTAr9Kaz1Pa10QeP5l4E2t9RTgzcDzEdefFrhFWuBCiBg2kBLKLcATge0ngFsHHM0g6GsulM4sZlNoXnAhhIg1kSZwDbyulNqllLo3sC9La10GEHjMDHeiUupepVShUqqwsrJy4BH3od8tcLmJKYSIUZH2A1+itS5VSmUCbyiljkR6Aa31Y8BjAAUFBUPa3NVa9zkSszOL2YRPSihCiBgVUQtca10aeKwAXgAWAeVKqRyAwGPFUAUZqWAujrQFbjUpWZ1eCBGz+kzgSimXUiohuA1cAxwAXgQ+FDjsQ8C/hirISHkDyTjSfuAWs8IrNXAhRIyKpISSBbyglAoe/xet9Rql1E7gr0qpjwJngfcOXZiRCZZDIu8HLjcxhRCxq88ErrUuAuaG2V8NrBqKoPrj+d3n2HGqhi9fPx1TIHFH2gvFalahVntP/rWnhMmZ8czKTRpwrEIIMZhifjKrz/51LwBKKb547TQg8ha4xWTC10cL/DPP7gHg9EM3XnyQQggxBGJ+KH283fgOOnq+gZ+vPQaA3RJ5C1xuYgohYlVMt8APlNTTFBgKX1bfxu6zdQCMS3NFdL7ZJDcxhRCxK2Zb4BUNbdz7ZCEAV03LoKy+LfTapMzIErjFbMLr12gtSVwIEXtiNoG/XVRNaX0bP7pjDr/9QEFof3q8jYx4e0TvYQ3UynuaD0UG+QgholnMJvB2j1G7XjI5HZvFxI1zcgD48vUzCHR57JPFbPz4PSVqjwyzF0JEsdhN4F5jJZ3gDcvPXzONS8ensHpG2ClZwgr2VukpUXdumX/i6V2y/JoQIqrEbAJvC7TAgwl8QrqLv91/BclxtojfwxIYsdnTjUyPtyOxv7L/PG8cKr/YcIUQYtDFbALvaIGbL/o9giWUnroSdm+Zf/5ve/FLXVwIESViOIH7Ucroy32x7IEEvvVEddjX3WFKK60eWQRZCBEdYjqBOyzmiG9YhrMyUC8/WFof9vXgPCnfes8sPrFiEgAtbkngQojoELMDedo8PuzWgX3/pMfbSY+30dQePikHF3tIi7fhCoz4bJUELoSIEjGbwNs9/oiHzPfGZbf0uLBxsIRiNZswBVr6LR5ZBFkIER1iN4F7fQO6gRnkslkoqWsN+1qwhGIzm0KflJRQhBDRIqZr4IPRAi+tb2XXmVr+8NYpTlQ0hfZXNbVz6yNbAKO7YZzV+LKQEooQIlrEXAI/er6R4poWI4EPsAYOUNdiDM759kuHWP3TjaH9p6qaQ9tWs4k4m9EEf2xTEW8drwq99s7ZWu59slASuxBi2MVUAq9obOPan2/iQ3/YQYvbi2MQSijdBW9cdh7EYzWbcNqMa208Vsk9v98eeu1PW0/z+qFy/rH73KDHIoQQvYmpBH4+MONgUVUzu87UkpPsHPRrNLR5+fWGE5yq7miB28wm4mzmbsd5aPP4Qr1TztV2raMfKKnnnt9t7/EGqRBCDFRM3cRs7tTdz+PTTEiLG/B7/uHDBXzkT4Wh56/sL+NHa452OSYryU6S04rNbAr1TCn47lrm5CWR4jKG7tc2u7uc840XD7LrTC3bT1WzcnrWgOMUQojuYqoF3tqtC9+kzPgBv+fK6VncOi839LyosrnL6/PHJpOZ4MBuMbP2s8v5+NIJALi9fgrP1IbmRzlQWs/Dbx4PzS0eXNXteHkTQggxFGIqgQdb4N+8eSYfWjyOG2bnDMr7xjs6/hAprm0JbeclO3nhE0tCz8emxXHP5ePCvsfB0gZ++sYxTlcb5ze2GV82P3j1yAWtcyGEGAwxlcCDPT2unpXNt265BKt5cMKPt1tD251r2e8tyL/g2LGpHWWbldMvnLo2eBO0srG9431++zZtMoeKEGKQxVQCb3Ybrdpgn+zBktCpBX64rCG0nRZmZR+lFI5A98UPXTEegPctGhN6vandi8fnp6bFzf9bPYU7C/I5UdHE2ZqWC95LCCEGIqZuYgZLEXH2wU3gLlv490tzhZ9b/K0vraTV7WNMahz7vnkNNrOJZ3YUA0YCr2l2ozVkJNiZm5/MXwvPhRZfFkKIwRIzLXCfX/PwuhNAYGj7IEqKs4bdn9LD4hDp8XbGBEopiQ4rDquZZVMzAGhq84a6O2bE20PdDKU7oRBisMVMAt9xqia0PZApZMOZPyYFAIfVxA2zs0P70+IjX93n+7ddAsCGo5XcEhiCPyHdhSvw10JzDzMeCiHExYqZBP6/L+wH4HcfLOjjyP4blxbHnQX5/PYDBfz67oWh/ak9lFDCSQjcCH2usDi0b2xaHPHSAhdCDJGYqIFrrSmtb2VOfhKrZw7+oBilFD+6Y+4F+5Od4Usr4bjs5i4DfcBY7i1UQnFLAhdCDK6YaIGX1bfR5vFzZ8GYvg8eRJZ+1NotZhM/vavjS+DOQBdEV2ASLLmJObwa2zzUt3pGOgwhhlRMtMCPB6Z5nTwIIy+H0k1zclk5PROvX4cSt8NqwqSkhDLcbn90K8fKm/jtBxZy7azsvk8QIgbFRAs8OE/3lGFK4MG69cWIs1lIdFgxB8bSK6UCq/7ITczhdCwwhcF9T+1ie1H4RauFiHUx0QI/UdFISpw17MCaobDjf1cRmNJkUMT3smybGBoTM1yheW12na3lsolpIxyREIMvJhL4Z6+exn9eOnbYrhdcvGGwuOwWuYk5zExKMTc/iUNlDZT2sGSeELEu4hKKUsqslHpHKfVS4Pk3lVIlSqk9gX83DFWQGQl25o5JHqq3H3Iuu4WeVr4XQ6PV7WNSZjxTsxIoqZUELkan/tTAPwMc7rbvZ1rreYF/rwxiXKOKy2ampd1Li9vLe3+zlX3n6vjk07t5dMPJkQ5t1Gr3+nBYzYxJieOMzEMjRqmIErhSKh+4Efjd0IYzOhktcC+7z9Sx83QtH3+ykJf3l/HDNUdGOrR+q2ho42x19CfEVrcPp9XMlKx4zlS3yGyQYlSKtAX+c+CLgL/b/k8ppfYppf6Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {
|
||||
"needs_background": "light"
|
||||
},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(train_x.cpu(), train_y[1:].cpu())\n",
|
||||
"plt.plot(test_x.cpu(), test_y.cpu())\n",
|
||||
"plt.plot(test_x.cpu(), predictions.mean(0).cpu())\n",
|
||||
"plt.fill_between(test_x.cpu(), predictions.mean(0).cpu() - 2 * predictions.std(0).cpu(),\n",
|
||||
" predictions.mean(0).cpu() + 2 * predictions.std(0).cpu(),\n",
|
||||
" color = \"green\", alpha = 0.2\n",
|
||||
" )\n",
|
||||
"# plt.ylim((3, 8))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 19,
|
||||
"id": "7913cfc3",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"[<matplotlib.lines.Line2D at 0x7fd118244850>]"
|
||||
]
|
||||
},
|
||||
"execution_count": 19,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYMAAAD4CAYAAAAO9oqkAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/MnkTPAAAACXBIWXMAAAsTAAALEwEAmpwYAAA2eUlEQVR4nO3deXxU1fn48c+Tyb6ThRAIOwFEdiOLKHWXxRa1WtGvYq39IhX6q22tpbvf2sXWWq2WQhW1WlvXakVLRUTc2AOyrwEihKwEskH2Ob8/7s0wmUzIBJLMJHner1deufeec2eeGYZ5cs499xwxxqCUUqp7C/J3AEoppfxPk4FSSilNBkoppTQZKKWUQpOBUkopINjfAbRGUlKSGTBggL/DUEqpTmXz5s3HjTHJZ6vTqZLBgAEDyMzM9HcYSinVqYjIFy3V0W4ipZRSmgyUUkppMlBKKYUmA6WUUmgyUEophSYDpZRSaDJQSimFJoN2V+80vLLxCDV1Tn+HopRSzfIpGYjINBHZJyJZIrLQS7mIyJN2+XYRGW8fHyYiW91+ykTkfrssQURWisgB+3ePNn1lAeKdbbksfHMHT39y0N+hKKVUs1pMBiLiABYB04ERwG0iMsKj2nQg3f6ZCywGMMbsM8aMNcaMBS4CTgNv2ecsBFYZY9KBVfZ+p/HQsl28vyu/xXqnauoAOFBY0d4hKaXUOfOlZTAByDLGHDLG1ACvALM86swCXjSW9UC8iKR61LkKOGiM+cLtnBfs7ReAG87lBfjL39ZmM/fvm1usV1ZpJYPSylqyCisorayludXlVuzK55svbGrTOJVSyhe+zE3UBzjqtp8DTPShTh8gz+3YbOBlt/0UY0wegDEmT0R6+hq0v7Wm/7+grAqAnJOVXP3Hj13HM396NUnRYY3q3msnl6raesJDHG0QqVJK+caXloF4Oeb5p+1Z64hIKPAV4HXfQ3OdO1dEMkUks6ioqLWnt4tT1XWu7Q92F7B8h5Xzln56iOuf+rRR3cJyKxlkeXQTbT1S4tp+ctUB/rbmsGv/5Omatg5ZKaXOypeWQQ7Q120/DchtZZ3pwBZjTIHbsQIRSbVbBalAobcnN8Y8DTwNkJGR4b1/pYNVuCWD+1/dSnRYMDNGpfKr/+wBoKi8muQY66/+grJqr4+xK7eMq0ekAPDHlfsblRVX1JAaF9EeoSullFe+tAw2AekiMtD+C382sMyjzjJgjj2qaBJQ2tAFZLuNxl1EDefcZW/fBbzd6uj9xD0ZVFTXcfJ0TaPrAFuPlri280uriAlrmnN3HCuhrt7ptcvpxCltGSilOlaLycAYUwcsAFYAe4DXjDG7RGSeiMyzqy0HDgFZwDPAfQ3ni0gkcA3wpsdDPwJcIyIH7PJHzvO1tCun03Dd45/wxuacRskAoLrOyamaetf+1qMnAeseg6Lyakb0jm3yeJ/sP86Qn/yXJ1cdaFKm3URKqY7m0+I2xpjlWF/47seWuG0bYH4z554GEr0cL8YaYRSQKmvqqaiuc3X35JysZF9BOZ/sL+Iau3vH3fde3era3pdfDsAXxaeoqXcyaVAiGw6faFS/pt5qEfx5dVaTxyqu0GSglOpYegdyM+Y8t4GLf/2Baz+ryPqCX7Ytl2+//HmT+u/vti6HiECR/WW+104Kl6UnNao7Ji2uyfn9EiJd2ydO1VBb7+SpVQd4aX2LCxQppdR502TQjE3ZVldPVa3V/eM5GgggMrTp8M/L0pPZdrSEG/+yhlV7CgkSGNmn8Ze/t5bFoOQo1/aJ0zW8sukoj63cz8/f3nler0MppXyhyaAFDfcJeEsGMeGNe9k2/vgqhveKAeDzIyX8a0sOg5KjG90zsGbhlcy5ZECj84b3iiEuIsS1f6KihoP280Xo/QZKqQ6gyaAFeaVVVNXWsyarmB6RIY3KojxGCSXHhJEYFdro2PSRvRrt94mPIDbcehwReOmeifx7/hTXMYCiimrX/Qmn7GsXSinVnny6gNyd5ZVW8vbWXI6VVPLPb05kcM9oJv5mFQDBQY3vtRORRncVX31BT+ZfMQSA/3dVOsUVZ+45eHv+FBKjQ0nrYV0ruP/qdKpq6zlVU8eWL0oaPW5hWRXRydHt8fKUUgrQZOBVXf2Zsf+vZ+aw9mAx904dxCVDGl8IFi83XidGWy0DEVh618Wu49+7ZmijemP6xnucF8ajt4zhiQ/289+d+TiNITkmjKLyagrLqxmkyUAp1Y60m8iLE27j/NceLAZgwZVDmtS7JSMNgJvG9+EPt4wBcLUMwoPPra+/b49IjIHC8mpG2ReeC8u938WslFJtRVsGXnjeAdwnPoIYtz79yYMSKSyv4p5LB3LHpP6NLhAn2NcMPK8n+Cqtx5lpKEb1iePDvYUU2hexlVKqvWgy8HC6po5pT1iTzY1Ji2NbTim948Mb1Xl57iTXtufsor1iw7k1oy93Tu5/Ts+f5na/QXpKNGHBQeScrDynx1JKKV9pN5GHg4WnXNt3TxkINP3CP5ugIOF3N49ucm+Br3rFhhMfGUKoI4iJAxO5LD2Jd7blklVYwTa3OY+UUqotacvAw7ES66/wX984kllje5NXWsXMUZ7r9LQfR5Dw8QNXEBHqIDQ4iG9MGcjteza41kI4+JsZOIK8zRiulFLnTlsGHhqSwYyRqYgI37p8MP0SI1s4q23FRYYQGmz900wenOi6kQ1gV25ph8ailOoeNBl4yC2pJCLEQbzHDWb+IiJ8+8p01/7za7K59a/rKNGZTZVSbUiTgZvPDhzn2c8O0zs+HJHA6YqZOTqV7EdmMiAxkrc+P8aGwyd4Y3OOv8NSSnUhmgywbjJb8vFB7nh2A2AtSBOI3Fc/25NX7sdIlFJdjSYDYN2hYh75717X/typg/0YTfN6xZ0Z4rom6zhOZ0CsAqqU6gK6/WiivNJK3tpyzLW/8SdX0TMm/Cxn+E9K7Jm48suqyPziJBMGJvgxIqVUV+FTy0BEponIPhHJEpGFXspFRJ60y7eLyHi3sngReUNE9orIHhGZbB9/SESOichW+2dG272sszPG8Nv/7uHzIyf5xt8yefPzYyRGhbL/V9MDNhEA9Iq1prq4ZHAiESEOlm071sIZSinlmxZbBiLiABZhrVOcA2wSkWXGmN1u1aYD6fbPRGCx/RvgT8B7xpibRSQUcB+n+bgx5g/n/zJap6iimr9+fIh/bT5GrT0p3WXpSa7hnIGqoZtoYJK1EM7OY2X+DEcp1YX48u03AcgyxhwyxtQArwCzPOrMAl40lvVAvIikikgsMBV4FsAYU2OMKWm78H1jjOGk23xDe+2Lr5GhDpKiQ7mwdywP3zCyo8NqtYZuopTYcAYkRZFdfKqFM5RSyje+JIM+wFG3/Rz7mC91BgFFwPMi8rmILBWRKLd6C+xupedEpIe3JxeRuSKSKSKZRUVFPoTb1I/f2sn0P1nzDRWWVfHnD61F6HvFhlNaWcfotLhGE9EFqv6JUYQFB5HeM5qBiVGUnK7V+w2UUm3Cl2TgbcC95zCW5uoEA+OBxcaYccApoOGaw2JgMDAWyAMe8/bkxpinjTEZxpiM5ORkH8JtakjPaPLLqsgvreJ//76ZjdknXGVllbXERgR+IgBrRtR1P7qKaSN7McDuKjp8XFsHSqnz50syyAH6uu2nAbk+1skBcowxG+zjb2AlB4wxBcaYemOME3gGqzuqXYy1F5JZtu1Yo8ne8suqqKl3Nlp/ONAlRIUiIgxMsi69uHcVLfn4IFMe+dB1HUQppXzlSzLYBKSLyED7AvBsYJlHnWXAHHtU0SSg1BiTZ4zJB46KyDC73lXAbgARcZ/97UZg5/m8kLO5sHcsIQ7h+TXZAHzwvanMvrgvR06cBmi0/nBn0TchkiCBw8dP8/H+IpZty+WR/+7lWEklBbr+gVKqlVocTWSMqRORBcAKwAE8Z4zZJSLz7PIlwHJgBpAFnAbudnuIbwP/sBPJIbey34vIWKzupGzg3rZ4Qd6EhzgYmhLDrlxr9E1aj0jiI88sXN+ZWgYNwoId9I6P4GBRBU+uOtCorKCsyrW2slJK+cKnm86MMcuxvvDdjy1x2zbA/GbO3QpkeDl+Z2sCPV/pPaPZlVtGUnQY4R4T0XXGZADWENOVuwuaHC8o02UylVKtE9gD69vQYHtB+UR7Wcr4iM6fDAYkRlFT5yQq1MHkQYn89qZRQODOraSUClzdJhkM6Wklgxr74qp7N1FnGU3kqWFE0ZfH9ObluZOYfXFfQoOD9JqBUqrVuk0yGJhsfXFW19YDkBwT5irrrC2Dkb1jAbj5ojTAWvsgJTaMnJJKXlr/BRXVdf4MTynViXSbieoGJUUzOi2O719rDWwa3y+eBVcM4WBRRaMuo85k4qBE1v3oykZTW/eKDec/2/P4z/Y8DHDnpP7+C1Ap1Wl0m2QQGhzEsgWXuvZFhAeuG3aWMzoH90QAcM2IFDZlnwRAl0pWSvmq23QTdRdzpw7mqdvGAXBKu4mUUj7SZNAFzRyVighUVNVhjOHFddmNJupTSilPmgy6oKAgISo0mPLqOrbnlPLzt3fx47d2+DsspVQA02TQRUWHBXOquo7SyloAyqpq/RyRUiqQaTLooqLDg6moruN0jXXdIDzY4eeIlFKBTJNBFxUVFkx5VR3HK6xrBWEh+k+tlGqefkN0UTF2N1GxnQwEofS0dhUppbzTZNBFRYU5qKiuo/iUNWndf3bkMeaX72PNKaiUUo1pMuiiosNCyC2pYuex0kbHC8t1RlOlVFOaDLqoGPsC8pYjJY2O78kr809ASqmApsmgi4oM9T56aG9+eQdHopTqDDQZdFH59jTW/RMjSY0Ldx3fqy0DpZQXPiUDEZkmIvtEJEtEFnopFxF50i7fLiLj3criReQNEdkrIntEZLJ9PEFEVorIAft3j7Z7WapPvDWB3dN3ZnD1BSmu43m68I1SyosWk4GIOIBFwHRgBHCbiIzwqDYLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {
|
||||
"needs_background": "light"
|
||||
},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(train_x.cpu(), vol.cpu() / 8)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "57800939",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.9.5"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,104 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "9af6056f-4605-4222-894d-15319a40dad0",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import yfinance as yf\n",
|
||||
"import datetime\n",
|
||||
"import pandas as pd\n",
|
||||
"import numpy as np\n",
|
||||
"import torch\n",
|
||||
"import matplotlib.pyplot as plt"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "907eabc5-5513-4734-852a-634cc933a986",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Check Preds"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 8,
|
||||
"id": "5b831c1f-627a-47e8-8d18-1d1cff777673",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"date = \"2019-12-18\"\n",
|
||||
"mean = \"ewma\"\n",
|
||||
"k = 50\n",
|
||||
"tckr = \"AMZN\"\n",
|
||||
"fpath = \"./saved-outputs/\" + tckr + \"/\"\n",
|
||||
"fname = \"volt_\" + mean + str(k) + \"_\" + date + \".pt\"\n",
|
||||
"\n",
|
||||
"data = yf.download(tickers=tckr, period='10y', progress=False)\n",
|
||||
"preds = torch.load(fpath + fname)\n",
|
||||
"end_idx = np.where(data.index == pd.to_datetime(date))[0][0]\n",
|
||||
"data = data.iloc[end_idx-200:end_idx]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"id": "9899a715-6e1e-4db0-927a-444f2e92513a",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"train_x = torch.arange(data.shape[0])\n",
|
||||
"test_x = torch.arange(preds.shape[1]) + train_x[-1] + 1"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 6,
|
||||
"id": "1ee64027-89f7-4171-b3fe-66ac2100b0ec",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAX0AAAD4CAYAAAAAczaOAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAABT7klEQVR4nO29d3xc13nn/T3TO3oHSLCAFEVSokSqW5bkKjveyCVOrCS2UzaKHeVdx86b4s3uejfvepPNpmzsjeV1bMdxEtuRLTtSiuVuybKoQsqUxE6AAImOATCD6f28f9yCATAoRBtg5nw/H3w4OPfOnYPLO7/73Oc8RUgpUSgUCkV1YCn3BBQKhUKxeSjRVygUiipCib5CoVBUEUr0FQqFoopQoq9QKBRVhK3cE1iOxsZG2d3dXe5pKBQKxbbi5MmTk1LKpvnjW170u7u7OXHiRLmnoVAoFNsKIcSVUuPKvaNQKBRVhBJ9hUKhqCKU6CsUCkUVoURfoVAoqggl+gqFQlFFKNFXKBSKKkKJvkKhUFQRSvQVFcPJKyFeuhoq9zQUii2NEn1FxfDbj57id776crmnoVBsaZToK7Ytn/zeJZ65NAnAwGScgakEfcE4E5FUmWemUGxdlOgrtiVSSj75/V7+/jkt0/zpS0Fz2/HLU+WalkKx5VGir9iWzCSzZPIFLo5HAfjhhSBd9W78ThvPKdFXKBZlyxdcUyhKMRFNAzAwFWcmmeV43xQ/e6yToVCS431K9BWKxVCWvmJbMhHRRL8g4SsvXCWZzXPP/ibu2NPAwFSCceXXVyhKoix9xbZkIjor6n/z4wEcVgu3727Abdcu6YvjUVoCrnJNT6HYsixr6QshuoQQPxBCnBNCnBFCfEgfrxdCfEcIcUn/t67oPR8VQvQKIS4IId5cNH5UCPGqvu0TQgixMX+WotIJ6u4dq0UwFklx2+56PA4be5q9APRNxMo5PYViy7IS904O+G0p5QHgduBhIcT1wO8D35NS9gDf039H3/Ye4CBwP/ApIYRVP9YjwENAj/5z/zr+LYoqYiKaxm230tPsA+CefVqDoCafE7/LRl8wXs7pKRRblmVFX0o5KqV8SX8dBc4BHcADwN/qu/0t8Hb99QPAV6SUaSllP9AL3CqEaAMCUsrjUkoJfLHoPQrFNTERTdMccLK/1Q/Avfs10RdCsLfZR+8Slv5XTwwyEk5uyjwViq3GNS3kCiG6gZuA54EWKeUoaDcGoFnfrQMYLHrbkD7Wob+eP65QXDMTkRTNficPHGnnnTd3sKfJZ27b0+SjL1ha9E9eCfE7X3uFr7xwdbOmqlBsKVYs+kIIH/AY8FtSyshSu5YYk0uMl/qsh4QQJ4QQJ4LBYKldFFVOMJam2e/idde18Oc/e4Ti5aE9TT4momkiqeyC933umcsAjKnoHkWVsiLRF0LY0QT/H6SUX9eHx3WXDfq/E/r4ENBV9PZOYEQf7ywxvgAp5WeklMeklMeamhY0c1coCEbSNPmdJbftaSq9mDs4neDJ02MAjOkhnwpFtbGS6B0BfA44J6X886JNTwDv11+/H3i8aPw9QginEGIX2oLtC7oLKCqEuF0/5vuK3qNQrJhkJk80nVtU9Pfqi7vzF3O/ekLzOt7YWcP4jLL0FdXJSiz9u4D3Aq8TQpzSf94K/DHwRiHEJeCN+u9IKc8AjwJngSeBh6WUef1YHwQ+i7a42wd8cz3/GEV1YMToNy8i+jvqPditgkt6iQaDb50Z51h3PTd01jIeVaKvqE6WTc6SUj5DaX88wOsXec/HgY+XGD8BHLqWCSoU8zFi9Bez9G1WC9e1Bnh1eMYcG5iMc2E8yn9+2/WksnnCiSypbB6X3VryGApFpaLKMCi2HcuJPsDhzhpeHZ6hUNBiBb51RvPlv+n6FvMJQZVqUFQjSvQV247pRAaABu/ion9jZw3RVI4r0wkAvnd+goPtAbrqPbTWaOUZxpRfX1GFKNFXbDtCcU30az32Rfc53FELwCtDYQB6J2Lc0KmNteo1ecajKoJHUX0o0VdsO6bjWbwO65L++H0tPpw2C68MzRBNZZmOZ9hR7wGg2RB9ZekrqhAl+optRyiRoc7rWHIfm9XCwfYArwyFGZzWSi4Yoh9w2XDbrSpBS1GVKNFXbDum4xkalhF9gBs6azk9HGFgSovXN0RfCEFrjUst5CqqEiX6im3HSix9gBs6a0hm8/zgvJYsbog+aDH+SvQV1YgSfcW2Yzqeod6zMtEHLVyzxm2npmjht7XGpdw7iqpEib5i2xGKr8zS393ow+uwEknl5lj5oNXdn4xmNmqKCsWWRYm+YluRyuaJZ/LUr0D0LRbBoQ7N2p8v+o1+J8lsnng6tyHzVCi2Kkr0FduKcEIrl1y3AvcOwI1dtQB0lbD0YTa7V6GoFpToK7YV03piVr138cSsYg4vYekDTMaU6CuqCyX625gnT49y8kqo3NPYVEJ6CYaVWvp37mng2M467tzTMGdcWfqKakWJ/jbm4/92jj958ny5p7GpzFr6KxP9Bp+Tr33wTrobvXPGG/3a+5Wlr6g2lOhvY+LpPKcGw6Rz+eV3rhBMS3+For8YDV4nFqEsfUX1oUR/G5PI5EjnCpweXqplcWVhWPq17pX59BfDahHUex0EYypsU1FdrKRd4ueFEBNCiNNFYzcKIY4LIV4VQvyzECKgj3cLIZJFHbY+XfSeo/r+vUKIT4jiTtaKayZfkKSyBQBeHJgu82w2j1A8Q43bjs26dnul0edUlr6i6ljJN+cLwP3zxj4L/L6U8jDwDeB3irb1SSmP6D8fKBp/BHgIrWduT4ljKq6BRGY2vvxEFYn+dCK7Yn/+cjT5ncqnr6g6lhV9KeXTwHxV2Q88rb/+DvCupY4hhGgDAlLK41JKCXwRePs1z7YK+bvjA3zjJ0MLxpMZzY/vsFp4cSBkdoiqdM6NRuisc6/LsZqUpa+oQlb7jHwa+Gn99buBrqJtu4QQPxFCPCWEuFsf6wCKlWtIHyuJEOIhIcQJIcSJYDC4yilWBl88foW/+kHfgvG4LvrXtweYSWbNblKVzJWpOL0TMe7b37wux2vULX3NDlEoqoPViv6vAA8LIU4CfsBQnFFgh5TyJuAjwJd0f38p//2i3zQp5WeklMeklMeamppWOcXKIJrK0TsRM7tFGRjlA7obtKSjarBYv3dOq5b5+gPrJPo+B+lcgagqxaCoIlYl+lLK81LKN0kpjwJfBvr08bSUckp/fVIf34dm2XcWHaITGFnLxKuFaEorOzA/CSuZ1Sz9nQ1a/HlViP75cfY2+8y/ea0YjdUnq+DcKRQGqxJ9IUSz/q8F+E/Ap/Xfm4QQVv31brQF28tSylEgKoS4XY/aeR/w+DrMv6LJF6TpxnnxytxlFdPSb9Qs/YkKF66RcJLnL0+vm5UPWvQOVMcNU6EwsC23gxDiy8C9QKMQYgj4GOATQjys7/J14G/0168F/lAIkQPywAeklIZafRAtEsgNfFP/USxBLFUcoTPX0k9kqsfSzxckv/WPp3DaLPzCrTvX7bimpa9i9RVVxLKiL6V8cJFNf1li38eAxxY5zgng0DXNrsqJ6K6dBq+DV4dmSGXzZjNwQ/SbfE58ThsT0cptCPLZH13mhf5p/uzdN7KjwbP8G1bIrKVfuedOoZiPysjdwsR0F849+5vI5Au8OjxjbjPi9D0OK03+yg09HJ1J8pffu8QbDjTzzpsXDfhaFXUeB1aLUJa+oqpQor+FieruHSNEsTjzNp7WLH2Pw1bR8eZ/9G/nyRckH/t3B1nvJG6zFEOFnjuFohRK9LcwRuTOjnoPe5t9c/z6yUwOIcBlt9AUqEzRT+fyPHlmjPfc0rWgCcp60eRTWbmK6kKJ/hbGsPT9LhvHdtZxYmDazLyNZ/J4HTaEEBVr6Z8eniGTK3DHnsYN+4xGv5OgEn1FFaFEfwtjWPp+l51j3fVEUjkuTcQAzafvdmiLus0BJ9F0zizNUCm8qD/ZHN1Zt2GfoTVIV6KvqB6U6G9hIkWW/i3dmvAZfv1EJo9XF/1K7QJ1YiDErkavGVq5ETT6HUzGMqoUg6JqUKK/xfjBhQlGZ5KA5t5xWC247FZ21Huo9dg5M6JF8MTTeTwOLeK2OeACIBirnNBDKSUnr0xvqJUP2g0zky8QSapSDIrqQIn+FqJQkPz6F0/y10/3A5p7x+/ShF0IwfVtAc6OaA1TEpkcnnmW/kSkciz9vmCcUCJrPuFsFMZTRCXdMBWKpVCiv4WIpnJk8gWuTifM332u2fy569sCnB+LkssXSGTyeJzattYazdJ/+tLk5k96g7gc1NYuDrQFNvRzZhO0VKy+ojpQor+FMMojD4U00Y+lc6alD1oZ5XSuQP9knEQmZ/r0670OfuWuXXz5hav83fGBTZ/3tSCl5L//y1le6F+68YtRS6hFd11tFLOWfuU8JSkUS6FEfwth9H8dnE4gpdTcO87ZXrDXt2tW79nRCPF03ozeAfiDnzrAbbvq+cyPLm/upK+RSxMxPvtMP3/y5Pkl95uIpBBCK0GxkRiWvorgUVQLSvS3EEbN/HgmTyiRJZqaa+nvafLhsFo4OxLRLf3ZbVaL4M49jQyFkqSyWzd087vnxgE4cSXEs72T/Nm3LzCTzC7YbyKapsHrXJdeuEtR67ZjswiVoKWoGpTobyFCRd2vBqcTuujPWvp2q4V9rT7OjkZ0n751zvv3NHuREi4H45s252vl++cm2N3oxWGz8Iufe55Pfr+X430L1yImomlaAhsXqmlgsQgafI6KL02tUBgo0d9CFIv+UChJpCh6x+BwRw2nBsOkcwU89rnb9jb7AOjTF0G3GtPxDC9dDfG2G9t5182dOGwWfXyhpT8eSdG8gfH5xexu9HFxPLopn6VQlBsl+luI6XgWi15T7Op0glg6R2Ce6N++u8Esz+CdZ+l3N3gRYnHRT2XzPH5quGyJSN89N05BwhsONPOLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {
|
||||
"needs_background": "light"
|
||||
},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(train_x, data.Close)\n",
|
||||
"plt.plot(test_x, preds[:10, :].T.exp(), color='gray', alpha=0.5)\n",
|
||||
"plt.show()"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,189 @@
|
||||
import matplotlib.pyplot as plt
|
||||
import seaborn as sns
|
||||
import numpy as np
|
||||
import torch
|
||||
import pandas as pd
|
||||
import os
|
||||
import gpytorch
|
||||
import argparse
|
||||
|
||||
from botorch.models import SingleTaskGP
|
||||
from botorch.optim.fit import fit_gpytorch_torch
|
||||
from gpytorch.likelihoods import GaussianLikelihood
|
||||
from gpytorch.mlls import ExactMarginalLogLikelihood
|
||||
from gpytorch.means import ConstantMean, LinearMean
|
||||
from gpytorch.kernels import SpectralMixtureKernel, MaternKernel, RBFKernel, ScaleKernel
|
||||
|
||||
import sys
|
||||
sys.path.append("../../magpie/means/")
|
||||
from EWMA import EWMAMean, DEWMAMean, TEWMAMean
|
||||
|
||||
sys.path.append("../../magpie/")
|
||||
from train_utils import LearnGPCV, TrainVolModel, TrainVoltMagpieModel, TrainBasicModel
|
||||
|
||||
sys.path.append("../../magpie/models/")
|
||||
from VoltMagpie import VoltMagpie
|
||||
|
||||
from voltron.means import LogLinearMean
|
||||
sys.path.append("../spdr-forecasting/")
|
||||
from rollout_utils import GeneratePrediction, Rollouts
|
||||
|
||||
def main(args):
|
||||
savepath = "./saved-outputs/" + args.symbol + "/"
|
||||
if not os.path.exists(savepath):
|
||||
os.mkdir(savepath)
|
||||
|
||||
dpath = "../../magpie/data/"
|
||||
full_data = pd.read_csv(dpath + args.symbol + ".csv")
|
||||
|
||||
full_data = pd.read_csv(dpath + args.symbol + ".csv")
|
||||
full_data["Date"] = pd.to_datetime(full_data['Date'])
|
||||
full_data = full_data.set_index(["Date"])
|
||||
|
||||
ntrain = 450
|
||||
ntest = 20
|
||||
nsample = args.nsample
|
||||
train_iters = 500
|
||||
train_x = torch.arange(ntrain) * 1./252
|
||||
test_x = torch.arange(ntest) * 1./252 + train_x[-1] + train_x[1]
|
||||
dt = train_x[1] - train_x[0]
|
||||
|
||||
start_date = pd.to_datetime("2011-11-21") - pd.Timedelta(ntrain, 'd')
|
||||
px = torch.FloatTensor(full_data.loc[start_date:].Close).squeeze()
|
||||
|
||||
start_idxs = torch.arange(1, px.shape[0] - ntrain - ntest)
|
||||
|
||||
## save the typical stuff to use for plotting/validating ##
|
||||
torch.save({"ntrain":ntrain, "ntest":ntest, "start_idxs":start_idxs},
|
||||
"./saved-outputs/metadata.pt")
|
||||
|
||||
if torch.cuda.is_available():
|
||||
use_cuda = True
|
||||
train_x = train_x.cuda()
|
||||
test_x = test_x.cuda()
|
||||
else:
|
||||
use_cuda = False
|
||||
|
||||
|
||||
fname = args.kernel + "_" + args.mean + str(args.k) + ".pt"
|
||||
|
||||
save_samples = torch.zeros(start_idxs.numel(), nsample, ntest)
|
||||
for idx, start_idx in enumerate(start_idxs):
|
||||
train_y = px[start_idx-1:ntrain+start_idx].squeeze()
|
||||
test_y = px[start_idx + ntrain:start_idx + ntrain+ntest].squeeze()
|
||||
|
||||
if use_cuda:
|
||||
train_y = train_y.cuda()
|
||||
test_y = test_y.cuda()
|
||||
test_x = test_x.cuda()
|
||||
|
||||
if args.kernel.lower() == "volt":
|
||||
|
||||
dt = train_x[1] - train_x[0]
|
||||
vol = LearnGPCV(train_x, train_y, train_iters=train_iters,
|
||||
printing=False)
|
||||
vmod, vlh = TrainVolModel(train_x, vol,
|
||||
train_iters=train_iters, printing=False)
|
||||
voltron, lh = TrainVoltMagpieModel(train_x, train_y[1:],
|
||||
vmod, vlh, vol,
|
||||
printing=False,
|
||||
train_iters=train_iters,
|
||||
k=args.k, mean_func=args.mean)
|
||||
|
||||
vmod.eval();
|
||||
|
||||
if args.mean in ['ewma', 'dewma', 'tewma']:
|
||||
save_samples[idx, ::] = Rollouts(train_x, train_y, test_x, voltron,
|
||||
nsample=nsample)
|
||||
else: ## VOLT + STANDARD MEAN
|
||||
voltron.vol_model.eval()
|
||||
predvol = voltron.vol_model(test_x).sample(torch.Size((nsample, ))).exp()
|
||||
save_samples[idx, ::] = GeneratePrediction(train_x, train_y, test_x,
|
||||
predvol, voltron).detach()
|
||||
del predvol
|
||||
|
||||
del voltron, lh, vmod, vlh, vol
|
||||
|
||||
else:
|
||||
kernel_possibilities = {"sm": SpectralMixtureKernel, "matern": MaternKernel, "rbf": RBFKernel}
|
||||
kernel = kernel_possibilities[args.kernel.lower()]
|
||||
if type(kernel) is not SpectralMixtureKernel:
|
||||
kernel = ScaleKernel(kernel())
|
||||
else:
|
||||
kernel = kernel()
|
||||
kernel.initialize_from_data_empspect(train_x, train_y.log())
|
||||
|
||||
train_y = train_y[1:]
|
||||
|
||||
model = SingleTaskGP(
|
||||
train_x.view(-1,1),
|
||||
train_y.log().reshape(-1, 1),
|
||||
covar_module=kernel,
|
||||
likelihood=GaussianLikelihood()
|
||||
)
|
||||
mean_name = args.mean.lower()
|
||||
if mean_name == "loglinear":
|
||||
model.mean_module = LogLinearMean(1)
|
||||
model.mean_module.initialize_from_data(train_x, train_y.log())
|
||||
elif mean_name == 'linear':
|
||||
model.mean_module = LinearMean(1)
|
||||
elif mean_name == "constant":
|
||||
model.mean_module = ConstantMean()
|
||||
elif mean_name == "ewma":
|
||||
model.mean_module = EWMAMean(train_x, train_y.log(), k=args.k).to(train_x.device)
|
||||
elif mean_name == "dewma":
|
||||
model.mean_module = DEWMAMean(train_x, train_y.log(), k=args.k).to(train_x.device)
|
||||
elif mean_name == "tewma":
|
||||
model.mean_module = TEWMAMean(train_x, train_y.log(), k=args.k).to(train_x.device)
|
||||
|
||||
if use_cuda:
|
||||
model = model.to(train_x.device)
|
||||
|
||||
mll = ExactMarginalLogLikelihood(model.likelihood, model)
|
||||
fit_gpytorch_torch(mll, options={'maxiter':train_iters, 'disp':False})
|
||||
|
||||
if mean_name in ["loglinear", "constant", 'linear']:
|
||||
save_samples[idx] = model.posterior(test_x).sample(torch.Size((nsample, ))).squeeze(-1).cpu().detach()
|
||||
else:
|
||||
save_samples[idx] = Rollouts(
|
||||
train_x, train_y, test_x, model, nsample=nsample, method = "nonvol"
|
||||
).cpu().detach()
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
del model, kernel
|
||||
|
||||
print("Start Time = ", start_idx.item(), " out of ", len(start_idxs))
|
||||
torch.cuda.empty_cache()
|
||||
torch.save(save_samples, savepath + fname)
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
parser.add_argument(
|
||||
"--symbol",
|
||||
type=str,
|
||||
default="SPY",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--kernel",
|
||||
type=str,
|
||||
default="volt",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--mean",
|
||||
type=str,
|
||||
default="ewma",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--k",
|
||||
type=int,
|
||||
default=100,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--nsample",
|
||||
type=int,
|
||||
default=1000,
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
main(args)
|
||||
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
@@ -0,0 +1,8 @@
|
||||
#!/bin/bash
|
||||
cat ../../voltron/data/nasdaq100.txt | while read line
|
||||
do
|
||||
for test_idx in {0..9}
|
||||
do
|
||||
python TickerSingleDayGenerator.py --kernel=volt --ntimes=25 --test_idx=${test_idx} --save=True --end_date="2022-01-12" --ticker=${line} --mean=constant
|
||||
done
|
||||
done
|
||||
@@ -0,0 +1,8 @@
|
||||
#!/bin/bash
|
||||
cat ../../voltron/data/nasdaq100.txt | while read line
|
||||
do
|
||||
for test_idx in {0..9}
|
||||
do
|
||||
python TickerSingleDayGenerator.py --kernel=volt --ntimes=25 --test_idx=${test_idx} --save=True --end_date="2022-01-12" --ticker=${line} --mean=constant
|
||||
done
|
||||
done
|
||||
Vendored
BIN
Binary file not shown.
@@ -0,0 +1,158 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
import pandas as pd
|
||||
import gpytorch
|
||||
import argparse
|
||||
import datetime
|
||||
import warnings
|
||||
import copy
|
||||
import os
|
||||
from voltron.data import make_ticker_list, GetStockHistory
|
||||
import sys
|
||||
sys.path.append("../calibration")
|
||||
from torch.utils.data import DataLoader
|
||||
from voltron.models import LSTM, BasicGP, Volt
|
||||
import pickle as pkl
|
||||
|
||||
def main(args):
|
||||
|
||||
################
|
||||
## Data Setup ##
|
||||
################
|
||||
stn_names, stn_lonlat, full_data = pkl.load(open("./wind_data.p", 'rb'))
|
||||
|
||||
use_cuda = False
|
||||
if torch.cuda.is_available():
|
||||
use_cuda = True
|
||||
|
||||
stn = args.stn_idx
|
||||
ntest = args.forecast_horizon
|
||||
ntrain = args.ntrain
|
||||
n_test_times = args.n_test_times
|
||||
ntime = full_data[0].shape[0]
|
||||
|
||||
test_idxs = torch.arange(ntrain, ntime-ntest,
|
||||
int((ntime-ntest-ntrain)/n_test_times))
|
||||
|
||||
stn_idxs = list(stn_names.keys())
|
||||
stn_data = full_data[stn]
|
||||
stn_data[stn_data == -99.0] = 0.
|
||||
|
||||
train_x = torch.arange(ntrain).float()/365
|
||||
test_x = torch.arange(ntrain, ntrain + ntest).float()/365
|
||||
|
||||
if use_cuda:
|
||||
train_x, test_x = train_x.cuda(), test_x.cuda()
|
||||
|
||||
####################
|
||||
## setup filename ##
|
||||
####################
|
||||
savepath = "./saved-outputs/stn" + str(stn) + "/"
|
||||
|
||||
if args.model.lower() == 'lstm':
|
||||
modelname = "lstm_"
|
||||
else:
|
||||
if args.model.lower() == 'gp':
|
||||
modelname = "gp_" + args.kernel + "_"
|
||||
elif args.model.lower() == 'volt':
|
||||
modelname = "volt_"
|
||||
|
||||
if args.mean.lower() == 'constant':
|
||||
modelname += 'constant' + "_"
|
||||
elif args.mean.lower() in ['ewma', 'dewma', 'tewma']:
|
||||
modelname += args.mean + args.k + "_"
|
||||
|
||||
|
||||
###############
|
||||
## Main Loop ##
|
||||
###############
|
||||
|
||||
if stn_data.mean() != 0:
|
||||
if not os.path.exists(savepath):
|
||||
os.mkdir(savepath)
|
||||
|
||||
for last_day in test_idxs:
|
||||
# try:
|
||||
raw_y = stn_data[last_day-ntrain:last_day] + 1
|
||||
train_y = torch.FloatTensor(raw_y).log()
|
||||
if use_cuda:
|
||||
train_y = train_y.cuda()
|
||||
|
||||
|
||||
if args.model.lower() == 'lstm':
|
||||
model = LSTM(train_x, train_y, 10, 128, 1)
|
||||
model.Train(args.train_iters)
|
||||
elif args.model.lower() == 'gp':
|
||||
model = BasicGP(train_x, train_y, kernel=args.kernel,
|
||||
mean=args.mean, k=args.k)
|
||||
model.Train(args.train_iters)
|
||||
elif args.model.lower() == 'volt':
|
||||
model = Volt(train_x, train_y, mean=args.mean, k=args.k)
|
||||
model.Train(gpcv_iters=args.train_iters,
|
||||
vol_mod_iters=args.train_iters,
|
||||
data_mod_iters=args.train_iters)
|
||||
else:
|
||||
print("ERROR: Model not found")
|
||||
|
||||
|
||||
samples = model.Forecast(test_x).squeeze()
|
||||
torch.save(samples, savepath + modelname + str(last_day.item()) + ".pt")
|
||||
torch.cuda.empty_cache()
|
||||
# except:
|
||||
# print("### BROKEN stn", stn, " idx", last_day, " ###")
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--stn_idx",
|
||||
type=int,
|
||||
default=0,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--mean",
|
||||
type=str,
|
||||
default='constant',
|
||||
)
|
||||
parser.add_argument(
|
||||
"--n_test_times",
|
||||
type=int,
|
||||
default=20,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--forecast_horizon",
|
||||
type=int,
|
||||
default=100,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--model",
|
||||
type=str,
|
||||
default='volt',
|
||||
)
|
||||
parser.add_argument(
|
||||
"--kernel",
|
||||
type=str,
|
||||
default="matern",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ntrain",
|
||||
type=int,
|
||||
default=400,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--nsample",
|
||||
type=int,
|
||||
default=1000,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--train_iters",
|
||||
type=int,
|
||||
default=400,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--k",
|
||||
type=int,
|
||||
default=400,
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
main(args)
|
||||
@@ -0,0 +1,486 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "f7c22e90-d490-4588-95c4-5fde46204bcb",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import pickle as pkl\n",
|
||||
"import pandas as pd\n",
|
||||
"import numpy as np\n",
|
||||
"import torch\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns\n",
|
||||
"import os\n",
|
||||
"\n",
|
||||
"sns.set_style('white')\n",
|
||||
"palette = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"sns.set(palette = palette, font_scale=2.0, style=\"white\", rc={\"lines.linewidth\": 4.0})"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "1c672fb1-f4b1-49f8-9343-2fc417cd7072",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"stn_names, stn_lonlat, full_data = pkl.load(open(\"./wind_data.p\", 'rb'))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"id": "93596c44-98fd-4fbc-958f-1110dde120ab",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def ECDF(sample_pxs, true_px): \n",
|
||||
" return (torch.sum(sample_pxs < true_px, 0)/sample_pxs.shape[0])\n",
|
||||
" \n",
|
||||
"def Calibration(pcts, percentile=0.95):\n",
|
||||
" in_band = np.where((pcts < percentile))[0].shape[0]\n",
|
||||
" return in_band/pcts.shape[0]\n",
|
||||
"\n",
|
||||
"def GetCalibration(model, ema=True, k=100, theta=0.0, horizon=np.arange(75,100), \n",
|
||||
" logger=[], exp=True):\n",
|
||||
" \n",
|
||||
" ntime = full_data[0].shape[0]\n",
|
||||
" ntrain = 400\n",
|
||||
" n_test_times = 20\n",
|
||||
" ntest = 100\n",
|
||||
" test_idxs = torch.arange(ntrain, ntime-ntest, \n",
|
||||
" int((ntime-ntest-ntrain)/n_test_times))\n",
|
||||
" \n",
|
||||
" stns = list(stn_names.keys())\n",
|
||||
"# pcts = torch.zeros(len(stns), len(test_idxs), horizon.shape[0]) \n",
|
||||
" pcts = torch.tensor([])\n",
|
||||
" for stn_save_idx, stn_idx in enumerate(stns):\n",
|
||||
" for test_save_idx, test_idx in enumerate(test_idxs):\n",
|
||||
"\n",
|
||||
" fpath = \"./saved-outputs/stn\" + str(stn_idx) + \"/\"\n",
|
||||
" fname = model + \"_\"\n",
|
||||
" if model == 'volt':\n",
|
||||
" if ema:\n",
|
||||
" fname += \"ema\" + str(k) + \"_\"\n",
|
||||
" fname += \"theta\" + str(theta) + \"_\"\n",
|
||||
" \n",
|
||||
" fname += str(test_idx.item()) + \".pt\"\n",
|
||||
" \n",
|
||||
"# print(fpath + fname)\n",
|
||||
" if os.path.exists(fpath + fname): \n",
|
||||
" if model == 'volt':\n",
|
||||
" preds = torch.load(fpath + fname)[0]\n",
|
||||
" elif model == 'matern':\n",
|
||||
" preds = torch.load(fpath + fname).cpu()\n",
|
||||
" else:\n",
|
||||
" preds = torch.load(fpath + fname)\n",
|
||||
" \n",
|
||||
" preds = preds[:, horizon]\n",
|
||||
" test_y = torch.tensor(full_data[stn_idx][test_idx:]) + 1\n",
|
||||
" \n",
|
||||
"# return preds, test_y\n",
|
||||
" if exp:\n",
|
||||
" preds = preds.exp()\n",
|
||||
" pcts = torch.cat((pcts, ECDF(preds, test_y[horizon])))\n",
|
||||
" \n",
|
||||
" pcts = pcts.flatten().numpy()\n",
|
||||
" percentiles = np.linspace(0.05, 0.95, 19)\n",
|
||||
" for pct in percentiles:\n",
|
||||
" clb = Calibration(pcts, pct)\n",
|
||||
" logger.append([clb, np.round(pct, 2), model, theta, ema, k])\n",
|
||||
" \n",
|
||||
" return logger"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 19,
|
||||
"id": "cf3f611f-75eb-4081-bb3c-3e6b238e7038",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"logger = []\n",
|
||||
"\n",
|
||||
"logger = GetCalibration('lstm', theta=0.0,\n",
|
||||
" exp=True, logger=logger)\n",
|
||||
"logger = GetCalibration('matern', theta=0.0,\n",
|
||||
" exp=True, logger=logger)\n",
|
||||
"logger = GetCalibration('volt', ema=True, k=200, theta=0.025,\n",
|
||||
" exp=True, logger=logger)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 20,
|
||||
"id": "b8af92c8-4c4a-4275-9686-cdb9ffccc3f5",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df = pd.DataFrame(logger)\n",
|
||||
"df.columns = ['Calibration', 'Percentile', 'Type', 'theta', 'ema', 'k']"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 21,
|
||||
"id": "4ab50047-962d-41de-ba3a-ed9b1d9c526f",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAA80AAAKCCAYAAADrzzv4AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMSwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Z1A+gAAAACXBIWXMAABcSAAAXEgFnn9JSAAEAAElEQVR4nOzdd1hT1/8H8Hc2e8mQoQiCoCCKG/fEvbVubLXWqlVr7XC0tVs77LDWOmpbV7VVcE9U3LgVAZkCyt4rrKz7+4Mf+RKSQBLC0s/reXye5Obec05yQ7yfe875HBbDMAwIIYQQQgghhBCihN3UDSCEEEIIIYQQQporCpoJIYQQQgghhBA1KGgmhBBCCCGEEELUoKCZEEIIIYQQQghRg4JmQgghhBBCCCFEDQqaCSGEEEIIIYQQNShoJoQQQgghhBBC1KCgmRBCCCGEEEIIUYOCZkIIIYQQQgghRA0KmgkhhBBCCCGEEDUoaCaEEEIIIYQQQtSgoJkQQgghhBBCCFGDgmZCCCGEEEIIIUQNCpoJIYQQQgghhBA1uE3dAEIIaQhDhw5Famqq/PnkyZOxadOmJmxR47lz5w4CAgIUtu3duxe9e/duohYRUrt58+bh7t278ue9evXCvn37mrBFr441a9bg6NGj8ueOjo64fPlygx33Kv82N3d0bghRj4JmQgghDa64uBiRkZHIyMhAcXExhEIhuFwujIyMYGZmBgcHBzg5OcHOzq6pm0qISsnJyUhKSkJ6ejqKi4tRUVEh//6amZnB1dUVLi4uYLFYTd1UQgghekZBMyEEABAWFobXXntNYVu3bt1w8OBBncqLjY3F+PHjlbZPmzYNX3/9tU5lnjt3DitXrlTYNnLkSGzZskWn8kjDysjIQGBgIE6dOoXExEQwDFPnMRYWFvDy8kLnzp3Rv39/+Pr6gsul/6pI4xOJRLh06RIuXLiA0NBQ5Ofn13mMqakpvL29MWzYMIwdOxZWVlaN0FJCCCENja5ECCEAAC8vLxgbG6OkpES+LTw8HGVlZTA0NNS6vDt37qjcXn0Ipj7KpCHHzU9ZWRm2bNmCPXv2QCqVanVsQUEBbt68iZs3b2L79u0wMTHB9u3b0bNnzwZqLSGKxGIx9u/fj927dyM7O1urY4uLixEaGorQ0FBs2rQJgwYNwvLly9GxY8cGai15ldG0BkIaDyUCI4QAALhcLnr06KGwTSwW4+HDhzqVpy44fvHiBdLT0/VWZp8+fXQqizSMtLQ0TJo0CX/++afWAbMqQqEQhYWFemgZIXWLiIjA+PHjsWnTJq0D5pokEgkuXbqEyZMn44MPPtCop5oQQkjzRD3NhBC53r174+rVqwrb7t69i379+mlVDsMwuHfvntrX79y5g0mTJmlVZm5uLuLj4xW22djYoH379lqVQxpOZmYm5s2bh5SUFKXX2Gw2unbtis6dO6Ndu3YwNTUFl8tFYWEh8vPzERMTg4iICCQnJzdBywkBjh07ho8//hhisVjl60ZGRujVqxc8PT1hZWUFKysrcDgcCIVCpKSkIDo6Gg8ePIBQKFQ4jmEYnDhxAtOmTaORMYQQ0kJR0EwIkVN1QadumHVtYmJiFHpVOByOQq+jLkGzqnb06tVL7f6aZHEl+vX5558rBcwsFgszZszA0qVLNUrylZqaiuDgYJw/f17nUQ6EaOvff//Fhg0bVM6779q1K5YtW4Y+ffqAz+fXWo5YLMaNGzdw8OBBXLt2TaN5/C3Vpk2bKLPyS4b+3yREPQqaCSFynTp1gpmZGYqKiuTbIiIiUFpaCiMjI43LqTmMeuzYsTh58qT8AlKXec2qjqFem+bjzp07uHTpksI2NpuNzZs3Y8yYMRqX4+joiNdffx2vv/46YmJicODAAZ3m1BOiqWvXruHzzz9XCnCNjY3x9ddfY/To0RqXxePxMGTIEAwZMgSRkZH4+uuv8eDBA303mRBCSCOjOc2EEDk2m61yXrO2F301A9yRI0fC3d1d/jwlJUVhLUhdygQoaG5OTp48qbRtzpw5WgXMNXl4eOCLL77QenoAIZrKzc3FRx99pDT/3sbGBvv379cqYK7Jy8sLBw4cwIcffggej1ffphJCCGlCFDQTQhSoCkS16RmuOZ+ZxWKhR48eSkOptSkzJycHz549U9hmZ2eHdu3aaVwGaVjXrl1T2jZv3rwmaAkhmvv222+Rl5ensI3H4+GPP/5Ap06d6l0+i8XCwoULsWvXLpiamta7PEIIIU2DhmcTQhSoykatzbzmmJgYFBQUyJ936NABFhYW6NWrF/bv369Q5uTJkzUqszn2MsfGxuLp06fIycmBTCaDpaUlWrduje7du2s1lF0TDMMgIiICz58/R1ZWFiQSCczNzeHq6gofHx8IBAK91qctiUSCrKwshW0mJiZwdnZuohapV1FRgbCwMCQkJKCoqAhcLhe2trZo164dvLy8wGKxGqTeoqIihIeHIycnB/n5+RCJRLC0tISVlRU6d+4MW1vbBqkXqMxoHhMTg7y8POTl5YHD4cDS0hK2trbo2rUrjI2N9V5nfn4+wsLCkJmZiby8PBgYGMDe3h5eXl5o06aN3uvTRWxsrMoREitXroSnp6de6/Lz89P6GIZhkJqaioSEBKSnp0MoFEIsFsPU1BTm5uZo27YtOnXq9NKuYx4fH4+oqChkZ2dDLBbDysoKrVu3Rrdu3RrkO1sdwzCIi4tDXFwcsrKyUFZWBoFAAGdnZwwfPlyj41/lc1ddYmIinj17htzcXBQUFMDQ0BCtWrVC69at4ePj0yijMMrLy/H48WP5776BgQEsLS3RoUMHeHp6NtjvPnm5vPx/rYQQrXh4eMDCwkIh8I2MjERJSYlGFyo1A+yqHuaa6+xq09Osy/rMQ4cOVRgCPnny5DqT1qSkpGDYsGEK2zZu3IgpU6YAAEQiEf755x/s3btX7fByHo+HwYMHY9WqVfXO7F1aWoodO3bgxIkTSEtLU7mPkZERxowZg7fffrvJgpG8vDyl+aCNPQ/5zp07CAgIUNi2d+9e+fckOTkZ27Ztw9mzZ1FWVqayDAcHB0yYMAGLFy/Wy42P8vJyHDx4EOfPn8eTJ09qXYLL3d0dEyZMwNy5c/VSd2ZmJv7++29cuXIFCQkJavfj8Xjo0qUL5syZg9GjR9f74vHu3bvYuXMnQkNDIZFIVO7ToUMHLFiwAJMmTWrSi9W9e/dCJpMpbGvXrh0WLlzYRC0CMjIycOHCBYSGhuLBgwd1LrdmaGiInj17Yv78+ejfv38jtbLSmjVrcPToUflzR0fHeieSqvqN3bdvn8os/ABgYGCAQYMGYenSpVrf3KjrNz4/Px9//fUXgoKCVC455ujoqDZobqxz5+Hhofa1u3fv1vp6lUuXLsHJyUlpuy7/b6qSmZmJP/74A5cuXap1KpaxsTH8/PwQEBCg043wX3/9FVu3blXYFhMTI3/84sUL/P7777X+7ltbW2PWrFlYsGCB3m94k5cLDc8mhChgsVhKQ6klEgnu37+v0fE1g+GqsqysrODm5ibfnpqaqvHyQs2hpzk+Ph4TJ07Exo0ba70IEIvFCA4Oxvjx43Hw4EGd6wsNDcXYsWOxfft2tQEzUBlYHzlyBOPHj1e4gG1Mqnq68/Pz1V6kNLbAwECMGzcOQUFBtbYpLS0N27dvx9ixYxEaGlqvOg8fPozhw4dj06ZNePToUZ1rVsfFxWHz5s0YMWIEzp07p3O95eXl+PbbbzFixAj8+eeftQbMQOX39f79+1i1ahUmTpyI2NhYnetdv349AgICcP36dbUBM1DZw7tmzRoEBAQoDY1uLBUVFTh9+rTS9pkzZ4LNbppLo9mzZ2Pw4MH4+uuvcfnyZY3WJy8rK8O1a9ewcOFCzJw5ExkZGY3Q0obx4sULTJkyBRs3blQbMAOV37Xz589j6tSp+PHHH5VufOgqJCQEo0aNwo4dO7Reo/tVP3dVpFIpfvnlF4wcObLWm8tVSkpKcPHiRQQEBOCtt96q9f86bR04cECj3/2cnBz8+uuvGDduHJKSkvRWP3n5UNBMCFGi67xmhmEUguuq+cxVdOltzs7OVrrwd3R0bNRe1SdPnmDGjBl1BiDVSaVSfPbZZ/jvv/+0ri8kJASLFi3S6gKirKwMa9asqVegriszMzOlpXgkEgmCg4MbvS01/fPPP1i3bh3Ky8s1PiYtLQ2LFi3ClStXtK5PLBZj/fr1+Pjjj7W+8AYqL+Deffdd/Pbbb1ofm52djXnz5uHPP/9ERUWF1sfHxMRg5syZSmu116W8vBxvv/02jhw5otUSS3fv3sXcuXM1CjD07fbt2ygtLVXYxuPxNJ4y0hAePHhQryWqHj16hKlTpyI6OlqPrWocycnJmDVrFuLi4jQ+RiKRYMeOHVizZk29A+czZ85g6dKlCiOstPEqn7sqZWVlWLZsGbZt26bTDdOrV69ixowZevkMfvzxR3zxxRda/Q6mpqZi9uzZyMzMrHf95OVEw7MJIUp0Xa85Ojpa4aLD3d0dVlZW8ue9evVSCOru3LmDqVOn1lpmU/cyZ2Rk4LvvvoNQKAQAcLlc9OrVC71794adnR0EAgGysrJw584dXL16ValHcePGjejbt6/KoXCqPHr0CMuXL4dYLFbYzmKx0LVrVwwcOBD29vbgcDjIyMjAjRs3cP/+fXm9X375JVatWqWHd665qrbVPFffffcdfHx8mixhW1hYGH7++Wf5cy6Xiz59+sDPzw92dnYQiURIS0vD5cuX8fTpU4VjxWIxli9fjn379qFr164a1SeTybBs2TKVQaetrS38/PzQqVMnWFhYQCAQoLCwEE+fPsX169cVbpAwDIMtW7bA0tISs2fP1qjunJwczJgxQ2XPTocOHdCzZ0+4ubnBzMwMQGXW6MePH+Pq1asoKSmR71tSUoLly5fj0KFDGifCWrVqlcqeeQsLC4wYMQKenp6wsrJCQUEB4uPjceHCBfkLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 1050x600 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"def PlotCalib(df, ax, title):\n",
|
||||
"\n",
|
||||
" pal = [ palette[0], palette[4], palette[6]]\n",
|
||||
" sns.lineplot(x='Percentile', y=\"Calibration\", hue='Type', data=df, ax=ax, alpha=0.5,\n",
|
||||
" palette=pal, legend=True)\n",
|
||||
" sns.scatterplot(x='Percentile', y=\"Calibration\", hue='Type', data=df, ax=ax, s=120, legend=False, zorder=4,\n",
|
||||
" palette=pal)\n",
|
||||
" x = np.linspace(0.05,0.95)\n",
|
||||
" y = np.linspace(0, len(percentiles))\n",
|
||||
" ax.plot(x, x, color=\"gray\", lw=1., ls=\"--\")\n",
|
||||
" ax.set_title(title)\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"from matplotlib.lines import Line2D\n",
|
||||
"colors = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"fig, ax = plt.subplots(1,1,dpi=150, figsize=(7, 4))\n",
|
||||
"\n",
|
||||
"percentiles = np.linspace(0.05, 0.95, 19)\n",
|
||||
"PlotCalib(df, ax, \"Wind Speed Calibration\")\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"plt.tick_params(labelsize=16)\n",
|
||||
"sns.despine()\n",
|
||||
"\n",
|
||||
"custom_lines = [Line2D([0], [0], color=palette[0], lw=2),\n",
|
||||
" Line2D([0], [0], color=palette[4], lw=2),\n",
|
||||
" Line2D([0], [0], color=palette[6], lw=2)]\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"plt.legend(custom_lines, ['LSTM', r\"GP-Matérn\", \"Volt + Magpie\"],\n",
|
||||
" fontsize=14, frameon=False, bbox_to_anchor=(0.45, 0.6))\n",
|
||||
"# ax.legend(fontsize=14, bbox_to_anchor=(1., 0.75))\n",
|
||||
"# plt.label(\"Percentile\")\n",
|
||||
"plt.savefig(\"./wind_calibration.pdf\", bbox_inches=\"tight\")\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "c8d78872-3d84-4225-8bde-f960c8155247",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Theta Sensitivity"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 7,
|
||||
"id": "cfe39573-cac1-4091-847a-912fec719089",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"logger = []\n",
|
||||
"logger = GetCalibration('volt', theta=0.0,\n",
|
||||
" exp=True, logger=logger)\n",
|
||||
"logger = GetCalibration('volt', theta=0.01,\n",
|
||||
" exp=True, logger=logger)\n",
|
||||
"logger = GetCalibration('volt', theta=0.025,\n",
|
||||
" exp=True, logger=logger)\n",
|
||||
"logger = GetCalibration('volt', theta=0.05,\n",
|
||||
" exp=True, logger=logger)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 10,
|
||||
"id": "7c2f0a47-a9b1-415f-908b-b1d7357c75c6",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df = pd.DataFrame(logger)\n",
|
||||
"df.columns = ['Calibration', 'Percentile', 'Type', 'theta', 'ema', 'k']"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 11,
|
||||
"id": "add1176b-16c8-4932-b29c-9c28afe3259c",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAABCMAAAKCCAYAAADr3xO5AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMSwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Z1A+gAAAACXBIWXMAABcSAAAXEgFnn9JSAAEAAElEQVR4nOzdd3gU1foH8O9syab3XkmBhIQEAoTeq/QmvVrAC4iogBfRC+L9KfaCijQLIAIioUd67z0khPQE0nvPZuv8/sjNymY3m21p8n6ex8fd2TllN8vszDvnvIdhWZYFIYQQQgghhBBCSDPhtHQHCCGEEEIIIYQQ8nyhYAQhhBBCCCGEEEKaFQUjCCGEEEIIIYQQ0qwoGEEIIYQQQgghhJBmRcEIQgghhBBCCCGENCsKRhBCCCGEEEIIIaRZUTCCEEIIIYQQQgghzYqCEYQQQgghhBBCCGlWFIwghBBCCCGEEEJIs6JgBCGEEEIIIYQQQpoVBSMIIYQQQgghhBDSrCgYQQghhBBCCCGEkGZFwQhCCCGEEEIIIYQ0KwpGEEIIIYQQQgghpFlRMII8l4YMGYLAwEDFf6tXr27pLjWbmzdvKr33wMBA3Lx5s6W7RUiD5s6dq/R9nTt3bkt3iRBCngvP8/kSIaTp8Vq6A4QQUlFRgUePHiE3NxcVFRWorKwEj8eDubk5rK2t4e7uDk9PT7i4uLR0VwkhhBBCCCFGQMEIopXo6GhMmzZNaVvXrl2xZ88evepLTEzEuHHjVLa/+OKL+Oijj/Sq88SJE1i+fLnStpEjR2Ljxo161UeaVm5uLg4cOIBjx44hLS0NLMs2WsbW1hYhISEIDQ1Fv379EB4eDh6PDmOkaQQGBjb42uuvv45ly5YZ3EZWVhaGDRsGuVyu9vUNGzZg8uTJBrdDWh9N36/6+Hw+LC0tYWlpCXd3dwQHByM0NBSDBw+Gubl5E/aSEEIIaTp0Fk+0EhISAgsLC1RVVSm2xcTEQCgUwszMTOf6GpoWcOvWLb37qK7Onj176l0faRpCoRAbN27Ejh07IJPJdCpbWlqKq1ev4urVq9i8eTMsLS2xefNmRERENFFvCVHv4MGDeP3118EwjEH1REZGNhiIIKSORCJBSUkJSkpKkJGRofi9s7CwwNixY7F8+XI4ODi0cC8JIYQQ3VDOCKIVHo+H7t27K22TSCS4d++eXvU1FHR4+vQpcnJyjFZnr1699KqLNI3s7GxMnDgRP//8s86BCHUqKytRVlZmhJ4RopusrCxcv37doDpYlsXBgweN1CPyPKqqqsK+ffswZswYnD17tqW7QwghhOiERkYQrfXs2RMXL15U2nbr1i307dtXp3pYlsXt27cbfP3mzZuYOHGiTnUWFRUhOTlZaZuTkxP8/f11qoc0nby8PMydOxeZmZkqr3E4HHTp0gWhoaFo164drKyswOPxUFZWhpKSEiQkJCA2NhYZGRkt0HNC1Dtw4AD69Omjd/nr168jKyvLiD0ibRmfz2/wN0ssFqO8vByFhYVqXy8pKcHy5cvx448/on///k3ZTUIIIcRoKBhBtKZuyoM+qzAkJCSgpKRE8ZzL5SrdJdcnGKGuHz169Ghw/3PnzulUPzHc+vXrVQIRDMNg+vTpWLJkiVbJKbOysnD69GmcPHlS71E5hOjLxsZGaSTOmTNnUFFRASsrK73qO3DggNJzW1tblJaWGtJF0oY5Ozvj8OHDGvcpLi7G5cuX8csvv+Dx48dKr0kkErz99ts4ffo0bG1tm7Cn5HlC50uEkKZE0zSI1oKDg2Ftba20LTY2FtXV1TrVU386xZgxY5TmXeuTN0JdGcoX0XrcvHlTZQgxh8PBV199hfXr12u9SoaHhwcWLFiAPXv24MiRI5g+fbpeOUsI0UdwcDD8/PwUz2tqanD06FG96iorK8Pp06eVto0dO9ag/pF/Pnt7e0yYMAF//vknZs+erfJ6eXk5tmzZ0gI9I4QQQnRHwQiiNQ6HozZvxN27d3Wqp37gYOTIkWjfvr3ieWZmps5DlykY0bqpu2CbPXs2Ro8erXedgYGB+PDDD3WeJkSIIeqvbFF/dIO2jh07BpFIpHjevn17hIWFGdQ38vzg8Xj4z3/+g969e6u8dvToUa1WJyKEEEJaGgUjiE7UXeDrMpKhfr4IhmHQvXt3lSkVutRZWFiIlJQUpW0uLi5o166d1nWQpnXp0iWVbXPnzm2BnhBimIkTJyotJxsbG4vExESd66kfxKDlO4muGIbB66+/rrK9oKAACQkJLdAjQgghRDeUM4LoRN3qFLrkjUhISFCaE92hQwfY2tqiR48e+O2335TqnDRpklZ1tsZREYmJiYiLi0NhYSHkcjns7Ozg6uqKbt26GX1NeJZlERsbiydPniA/Px9SqRQ2Njbw8/NDWFgYBAKBUdvTlVQqRX5+vtI2S0tL+Pj4tFCPGiYSiRAdHY3U1FSUl5eDx+PB2dkZ7dq1Q0hIiMHLODakvLwcMTExKCwsRElJCcRiMezs7GBvb4/Q0FA4Ozs3SbtA7QonCQkJKC4uRnFxMbhcLuzs7ODs7IwuXbrAwsLC6G2WlJQgOjoaeXl5KC4uhqmpKdzc3BASEgIvLy+jt2dMTk5O6N+/P86fP6/YduDAAbz77rta1xEfH49Hjx4pnvP5fEyYMAFXrlwxal/riMVixMTEKD7vyspKWFtbw97eHgEBAQgICDBqeyzLIisrC6mpqcjJyUFlZSUkEgmsrKxgY2MDb29vBAcHKwV1mlJOTg5iYmKQnZ0NoVAIGxsbODo6omvXrnB0dGyWPjSV8PBwWFpaorKyUml7cnIygoKC9K6XZVkkJCQgIyMDxcXFKC0thZmZGezt7eHh4YHQ0NBm+/s1l8LCQjx69AglJSUoKiqCXC6Hvb09HB0d0blz52bJw5GdnY24uDhkZ2ejqqoKXC4XDg4OGD16dKNTEisrK5GYmIj09HSUl5ejuroaPB4PZmZmsLOzg7u7O9q1awd7e/smfx+NSUtLQ0pKCoqKihTfLQcHB7i6uiIsLAx8Pr/J+1BTU4MHDx4ofu9NTU1hZ2eHDh06ICgoqMl+7wkhyv5ZvySkyQUGBqokWXv06BGqqqq0umipH7ioGxERERGhtF2XkRHqgiGNBSOGDBmiNBVk0qRJ+OSTTzSWyczMxNChQ5W2bdiwQXFHUywW4/fff8fOnTsbnGbC5/MxaNAgvPXWWwav9FFdXY0tW7bgyJEjyM7OVruPubk5Ro8ejX/9618tdpFXXFysMmS4ufM83Lx5E/PmzVPatnPnTsX3JCMjA5s2bcJff/0FoVCotg53d3eMHz8er732mlECSjU1NdizZw9OnjyJhw8falzqtH379hg/fjzmzJljlLbz8vLw66+/4sKFC0hNTW1wPz6fj86dO2P27NkYNWqUwSdnt27dwtatW3H9+nVIpVK1+3To0AEvv/wyJk6c2GpPBl988UWlYMSRI0ewcuVKrU+g//zzT6XnAwcOhIODg1H7yLIsTp48iYMHD+LWrVsac/s4Oztj+PDhWLRoEVxdXfVqLzc3F6dOncL169dx9+7dRpfcNTMzQ0REBObPn49+/frp1WZjx/HTp09j+/btePDggdryDMMgNDQUr7/+OgYOHKhXH1oal8uFh4eHykiIZ5NE6yImJga7du3C1atXG1y5AwAsLCzQt29fLFy4UKvpRYWFhRg4cKDSv/uhQ4di06ZNevWzztq1a7Fv3z6lbceOHVOa+qlJWVkZdu3ahbNnz+Lx48cNTm/hcDgIDg7G1KlTMWXKFJ0vljV9VyUSCfbv3489e/Y0OMqqZ8+e8PT0VNkuk8lw5MgRHDx4ELdv34ZcLm+0L56enujWrRtGjBiB/v37N3rDQp/zJXXy8vKwfft2nD17VuNUXAsLC/Tu3Rvz5s3T68bSd999h++//15p27P/Pp4+fYoff/xR4++9o6MjZs6ciZdfftnoN5AIIcpomgbRCcMwKlMqpFIp7ty5o1X5+kGGurrq7tDVycrK0noZx9YwMiI5ORkTJkzAhg0bNP7ISiQSnD59GuPGjcOePXv0bu/69esYM2YMNm/e3GAgAqgNWPz5558YN24cDh48qHd7hlB3olNSUtLgSUBzO3DgAMaOHYvIyEiNfcrOzsbmzZsxZswYXL9+3aA29+/fj2HDhuGTTz7B/fv3NQYiACApKQlffvklhg8fjhMnTujdbk1NDT799FMMHz4cP//8s8ZABFD7fb1z5w7eeustTJgwQa/pCHXtvvfee5g3bx4uX77cYCACqB1VtHr1asybNw/FxcV6tdfUBg0apBQ8KC4uVgpOaCIWi1VyqEyZMsWo/btz5w6mTJmC5cuX48KFC40mGc7Pz8fu3bsxfPhwfPvtt1pd0Dxr1qxZGDRoED766COcO3eu0UAEAAiFQly6dAmvvPIKZsyYgdzcXJ3a1KSiogL/+te/8PrrrzcYiABqAzYPHz7EokWLsHr1ao3fy9ZM3cVS/ZESjcnKysKyZcvw4osv4vDhwxoDEQBQVVWFU6dOYerUqVi2bBnKy8s17u/o6Kiy5OilS5cM+jdeU1ODqKgopW1hYWFaBSJkMhm2bt2KYcOG4bvvvkNcXJzGPBtyuRyxsbFYt24dXnjhBb0SbauTnp6OyZMnY/369TofX+Pj4zFp0iSsXr0aN2/e1PrfbWZmJg4fPoylS5ciMjJSn27rRCaT4dtvv8XIkSM13qypU1VVhTNnzmDevHlYtGiRxnMcXe3evVur3/vCwkJ89913GDt2LNLT043WPiFEFQUjiM70zRvBsqxS0KIuX0QdfUZHFBQUqFxQeXh4NOsogIcPH2L69OmNXtg9SyaT4YMPPsAff/yhc3vnz5/HwoULdfqBFgqFWL16tUEBEH1ZW1vDxMREaZtUKlVZSaAl/P7771izZg1qamq0LpOdnY2FCxfiwoULOrcnkUjw3nvv4f3330dBQYHLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 900x600 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"def PlotCalib(df, ax, title):\n",
|
||||
"# pal = [palette[6], palette[0], palette[4]]\n",
|
||||
"# sns.lineplot(x='Percentile', y=\"Calibration\", hue='Type', data=df, ax=ax, alpha=0.25,\n",
|
||||
"# palette=pal, legend=False)\n",
|
||||
"# sns.scatterplot(x='Percentile', y=\"Calibration\", hue='Type', data=df, ax=ax, s=120, legend=False, zorder=4,\n",
|
||||
"# palette=pal, alpha=0.5)\n",
|
||||
"\n",
|
||||
"\n",
|
||||
" pal = [palette[6], palette[0], palette[4], palette[7]]\n",
|
||||
" sns.lineplot(x='Percentile', y=\"Calibration\", hue='theta', data=df, ax=ax, alpha=0.5,\n",
|
||||
" palette=pal, legend=True)\n",
|
||||
" sns.scatterplot(x='Percentile', y=\"Calibration\", hue='theta', data=df, ax=ax, s=120, legend=False, zorder=4,\n",
|
||||
" palette=pal)\n",
|
||||
" x = np.linspace(0.05,0.95)\n",
|
||||
" y = np.linspace(0, len(percentiles))\n",
|
||||
" ax.plot(x, x, color=\"gray\", lw=1., ls=\"--\")\n",
|
||||
" ax.set_title(title)\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"from matplotlib.lines import Line2D\n",
|
||||
"colors = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"fig, ax = plt.subplots(1,1,dpi=150, figsize=(6, 4))\n",
|
||||
"\n",
|
||||
"percentiles = np.linspace(0.05, 0.95, 19)\n",
|
||||
"PlotCalib(df, ax, \"Wind Speed Mean Reversion\")\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"plt.tick_params(labelsize=16)\n",
|
||||
"sns.despine()\n",
|
||||
"\n",
|
||||
"custom_lines = [Line2D([0], [0], color=palette[6], lw=2),\n",
|
||||
" Line2D([0], [0], color=palette[0], lw=2),\n",
|
||||
" Line2D([0], [0], color=palette[4], lw=2),\n",
|
||||
" Line2D([0], [0], color=palette[7], lw=2)]\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"legend = plt.legend(custom_lines,['0.0', \"0.01\", \"0.025\", \"0.05\"], \n",
|
||||
" fontsize=14, frameon=False, bbox_to_anchor=(1., 0.75))\n",
|
||||
"legend.set_title(r'$\\theta$')\n",
|
||||
"# ax.legend(fontsize=14, bbox_to_anchor=(1., 0.75))\n",
|
||||
"# plt.label(\"Percentile\")\n",
|
||||
"plt.savefig(\"./theta_sensitivity.pdf\", bbox_inches=\"tight\")\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "18c85912-e2b3-49ef-b3f2-cc62e7226ecb",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## EMA Calibration"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 12,
|
||||
"id": "1e5da442-bbe8-45dc-b13d-b42a7df4e4da",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"logger = []\n",
|
||||
"logger = GetCalibration('volt', ema=True, theta=0.025, k=50,\n",
|
||||
" exp=True, logger=logger)\n",
|
||||
"logger = GetCalibration('volt', ema=True, theta=0.025, k=100,\n",
|
||||
" exp=True, logger=logger)\n",
|
||||
"logger = GetCalibration('volt', ema=True, theta=0.025, k=200,\n",
|
||||
" exp=True, logger=logger)\n",
|
||||
"logger = GetCalibration('volt', ema=True, theta=0.025, k=400,\n",
|
||||
" exp=True, logger=logger)\n",
|
||||
"logger = GetCalibration('volt', ema=False, theta=0.025, k=400,\n",
|
||||
" exp=True, logger=logger)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 13,
|
||||
"id": "538dba70-fd2a-4292-86d3-707d045088dd",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df = pd.DataFrame(logger)\n",
|
||||
"df.columns = ['Calibration', 'Percentile', 'Type', 'theta', 'ema', 'k']"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 14,
|
||||
"id": "5f897d07-6e65-4553-a08a-e9b7dd881b26",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAA80AAAKCCAYAAADrzzv4AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMSwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Z1A+gAAAACXBIWXMAABcSAAAXEgFnn9JSAAEAAElEQVR4nOzdd3RURfvA8e/W9E4KKST00HtHmhQVUAQLSrVgV1Ss6Gv9Kdje14rYBUGwAKKC9Cq9dxJCCKT33nY3e39/xAQ2uymbhATk+ZzDOezsvTOzu8nmPndmnlEpiqIghBBCCCGEEEIIK+rG7oAQQgghhBBCCHGlkqBZCCGEEEIIIYSohATNQgghhBBCCCFEJSRoFkIIIYQQQgghKiFBsxBCCCGEEEIIUQkJmoUQQgghhBBCiEpI0CyEEEIIIYQQQlRCgmYhhBBCCCGEEKISEjQLIYQQQgghhBCVkKBZCCGEEEIIIYSohATNQgghhBBCCCFEJSRoFkIIIYQQQgghKiFBsxBCCCGEEEIIUQkJmoUQQgghhBBCiEpoG7sDQghxOQwbNoz4+Pjyx7feeitz585txB41nD179jB16lSLsoULF9KnT59G6pEQVZsyZQp79+4tf9y7d29++OGHRuzRteOFF15gxYoV5Y+DgoLYtGnTZTvvWv5uvtLJZyNE5SRoFkIIcdnl5uZy4sQJkpKSyM3NJS8vD61Wi7OzM+7u7gQGBhIcHIy/v39jd1UIm2JjY4mJiSExMZHc3FyKi4vLf37d3d1p0aIFzZs3R6VSNXZXhRBC1DMJmoUQABw5coQ77rjDoqx79+4sWbKkVvVFRkYyduxYq/LbbruNt956q1Z1rlmzhpkzZ1qUjRo1io8//rhW9YnLKykpiWXLlvHnn39y7tw5FEWp9hxPT086dOhAp06dGDhwIN26dUOrlT9VouEZDAY2btzIunXr2LVrF5mZmdWe4+bmRseOHbn++usZPXo03t7eDdBTIYQQl5tciQghAOjQoQMuLi7k5+eXlx07dozCwkKcnJzsrm/Pnj02yy+dglkfdcqU4ytPYWEhH3/8MQsWLKCkpMSuc7OystixYwc7duxg/vz5uLq6Mn/+fHr16nWZeiuEJaPRyKJFi/jmm29ITU2169zc3Fx27drFrl27mDt3LoMHD+bxxx+nXbt2l6m34lomyxqEaDiSCEwIAYBWq6Vnz54WZUajkYMHD9aqvsqC4wsXLpCYmFhvdfbt27dWdYnLIyEhgXHjxvHtt9/aHTDbkpeXR3Z2dj30TIjqHT9+nLFjxzJ37ly7A+aKTCYTGzdu5NZbb+XZZ5+t0Ui1EEKIK5OMNAshyvXp04etW7dalO3du5cBAwbYVY+iKOzbt6/S5/fs2cO4cePsqjM9PZ2oqCiLMl9fX1q2bGlXPeLySU5OZsqUKcTFxVk9p1ar6dq1K506dSIsLAw3Nze0Wi3Z2dlkZmYSERHB8ePHiY2NbYSeCwG//fYbL7/8Mkaj0ebzzs7O9O7dm/DwcLy9vfH29kaj0ZCXl0dcXBynT5/mwIED5OXlWZynKAq///47t912m8yMEUKIq5QEzUKIcrYu6CqbZl2ViIgIi1EVjUZjMepYm6DZVj969+5d6fE1yeIq6tfrr79uFTCrVCruvPNOHnnkkRol+YqPj2f9+vWsXbu21rMchLDXTz/9xKuvvmpz3X3Xrl159NFH6du3L3q9vsp6jEYjf//9N0uWLGHbtm01Wsd/tZo7d65kVv6Xkb+bQlROgmYhRLn27dvj7u5OTk5Oednx48cpKCjA2dm5xvVUnEY9evRo/vjjj/ILyNqsa7Z1jozaXDn27NnDxo0bLcrUajUffPABN910U43rCQoKYvr06UyfPp2IiAgWL15cqzX1QtTUtm3beP31160CXBcXF9566y1uvPHGGtel0+kYOnQoQ4cO5cSJE7z11lscOHCgvrsshBCigcmaZiFEObVabXNds70XfRUD3FGjRtG6devyx3FxcRZ7QdamTpCg+Uryxx9/WJVNmjTJroC5orZt2/LGG2/YvTxAiJpKT0/n+eeft1p/7+vry6JFi+wKmCvq0KEDixcv5rnnnkOn09W1q0IIIRqRBM1CCAu2AlF7RoYrrmdWqVT07NnTaiq1PXWmpaVx9uxZizJ/f3/CwsJqXIe4vLZt22ZVNmXKlEboiRA1984775CRkWFRptPp+Prrr2nfvn2d61epVNx333189dVXuLm51bk+IYQQjUOmZwshLNjKRm3PuuaIiAiysrLKH7dp0wZPT0969+7NokWLLOq89dZba1TnlTjKHBkZycmTJ0lLS8NsNuPl5UVAQAA9evSwayp7TSiKwvHjxzl//jwpKSmYTCY8PDxo0aIFnTt3xsHBoV7bs5fJZCIlJcWizNXVldDQ0EbqUeWKi4s5cuQI0dHR5OTkoNVq8fPzIywsjA4dOqBSqS5Luzk5ORw7doy0tDQyMzMxGAx4eXnh7e1Np06d8PPzuyztQmlG84iICDIyMsjIyECj0eDl5YWfnx9du3bFxcWl3tvMzMzkyJEjJCcnk5GRgaOjI02bNqVDhw6EhITUe3u1ERkZaXOGxMyZMwkPD6/Xtvr162f3OYqiEB8fT3R0NImJieTl5WE0GnFzc8PDw4NmzZrRvn37f+0+5lFRUZw6dYrU1FSMRiPe3t4EBATQvXv3y/IzeylFUThz5gxnzpwhJSWFwsJCHBwcCA0NZfjw4TU6/1r+7C517tw5zp49S3p6OllZWTg5OeHj40NAQACdO3dukFkYRUVFHD58uPx739HRES8vL9q0aUN4ePhl+94X/y7//t9WIYRd2rZti6enp0Xge+LECfLz82t0oVIxwC4bYa64z649I8212Z952LBhFlPAb7311mqT1sTFxXH99ddblM2ZM4fx48cDYDAY+PHHH1m4cGGl08t1Oh1DhgzhqaeeqnNm74KCAr744gt+//13EhISbB7j7OzMTTfdxEMPPdRowUhGRobVetCGXoe8Z88epk6dalG2cOHC8p+T2NhY5s2bx19//UVhYaHNOgIDA7n55pt58MEH6+XGR1FREUuWLGHt2rUcPXq0yi24Wrduzc0338zkyZPrpe3k5GS+//57tmzZQnR0dKXH6XQ6unTpwqRJk7jxxhvrfPG4d+9evvzyS3bt2oXJZLJ5TJs2bbj33nsZN25co16sLly4ELPZbFEWFhbGfffd10g9gqSkJNatW8euXbs4cOBAtdutOTk50atXL6ZNm8bAgQMbqJelXnjhBVasWFH+OCgoqM6JpMq+Y3/44QebWfgBHB0dGTx4MI888ojdNzeq+47PzMzku+++Y/ny5Ta3HAsKCqo0aG6oz65t27aVPrd3794qny+zceNGgoODrcpr83fTluTkZL7++ms2btxY5VIsFxcX+vXrx9SpU2t1I/yTTz7h008/tSiLiIgo//+FCxf4/PPPq/zeb9KkCXfddRf33ntvvd/wFv8uMj1bCGFBpVJZTaU2mUzs37+/RudXDIbL6vL29qZVq1bl5fHx8TXeXuhKGGmOiorilltuYc6cOVVeBBiNRtavX8/YsWNZsmRJrdvbtWsXo0ePZv78+ZUGzFAaWP/666+MHTvW4gK2Idka6c7MzKz0IqWhLVu2jDFjxrB8+fIq+5SQkMD8+fMZPXo0u3btqlObv/zyC8OHD2fu3LkcOnSo2j2rz5w5wwcffMCIESNYs2ZNrdstKirinXfeYcSIEXz77bdVBsxQ+vO6f/9+nnrqKW655RYiIyNr3e5LL73E1KlT2b59e6UBM5SO8L7wwgtMnTrVamp0QykuLmbVqlVW5RMnTkStbpxLo7vvvpshQ4bw1ltvsWnTphrtT15YWMi2bdu47777mDhxIklJSQ3Q08vjwoULjB8/njlz5lQaMEPpz9ratWuZMGEC//3vf61ufNTW5s2bueGGG/jiiy/s3qP7Wv/sypSUlPDRRx8xatSoKm8ul8nPz2fDhg1MnTqVBx54oMq/dfZavHhxjb7309LS+OSTTxgzZgwxMTH11r7495GgWQhhpbbrmhVFsQiuy9Yzl6nNaHNqaqrVhX9QUFCDjqoePXqUO++8s9oA5FIlJSW89tpr/Pzzz3a3t3nzZmbMmGHXBURhYSEvvPBCnQL12nJ3d7faisdkMrF+/foG70tFP/74I7Nnz6aoqKjG5yQkJDBjxgy2bNlid3tGo5GXXnqJl19+2e4Lbyi9gHvyySf57LPP7D43NTWVKVOm8O2331JcXGz3+REREUycONFqr/bqFBUV8dBDD/Hrr7/atcXS3r17mTx5co0CjPq2e/duCgoKLMp0Ol2Nl4xcDgcOHKjTFlWHDh1iwoQJnD59uh571TBiY2O56667OHPmTI3PMZlMfPHFF7zwwgt1DpxXr17NI488YjHDyh7X8mdXprCwkEcffZR58+bV6obp1q1bufPOO+vlPfjvf//LG2+8Ydf3YHx8PHfffTfJycl1bl/8O8n0bCGEldru13z69GmLi47WrVvj7e1d/rh3794WQd2ePXuYMGFClXU29ihzUlIS7777Lnl5eQBotVp69+5Nnz598Pf3x8HBgZSUFPbs2cPWrVutRhTnzJlD//79bU6Fs+XQoUM8/vjjGI1Gi3KVSkXXrl0ZNGgQTZs2RaPRkJSUxN9//83+/fvL233zzTd56qmn6uGV11xZ3yp+Vu+++y6dO3dutIRtR44c4cMPPyx/rNVq6du3L/369cPf3x+DwUBCQgKbNm3i5MmTFucajUYef/xxfvjhB7p27Vqj9sxmM48++qjNoNPPz49+/frRvn17PD09cXBwIDs7m5MnT7J9+3aLGySKovDxxx/j5eXF3XffXaO209LSuPPOO22O7LRp04ZevXrRqlUr3N3dgdKs0YcPH2br1q3k5+eXH5ufn8/jjz/O0qVLa5wI66mnnrI5Mu/p6cmIESMIDw/H29uLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 1050x600 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"def PlotCalib(df, ax, title):\n",
|
||||
" sub_df = df[(df['ema']==True) & (df['k']!=50)]\n",
|
||||
" pal = [palette[0], palette[4], palette[6]]\n",
|
||||
" sns.lineplot(x='Percentile', y=\"Calibration\", hue='k', data=sub_df, ax=ax, alpha=0.5,\n",
|
||||
" palette=pal, legend=True)\n",
|
||||
" sns.scatterplot(x='Percentile', y=\"Calibration\", hue='k', data=sub_df, ax=ax, s=120, legend=False, zorder=4,\n",
|
||||
" palette=pal)\n",
|
||||
" x = np.linspace(0.05,0.95)\n",
|
||||
" y = np.linspace(0, len(percentiles))\n",
|
||||
" ax.plot(x, x, color=\"gray\", lw=1., ls=\"--\")\n",
|
||||
" ax.set_title(title)\n",
|
||||
" \n",
|
||||
" \n",
|
||||
" pal = [palette[2]]\n",
|
||||
" sub_df = df[df['ema']==False]\n",
|
||||
" sns.lineplot(x='Percentile', y=\"Calibration\", hue='Type', data=sub_df, ax=ax, alpha=0.5,\n",
|
||||
" palette=pal, legend=True)\n",
|
||||
" sns.scatterplot(x='Percentile', y=\"Calibration\", hue='Type', data=sub_df, ax=ax, s=120, legend=False, zorder=4,\n",
|
||||
" palette=pal)\n",
|
||||
" x = np.linspace(0.05,0.95)\n",
|
||||
" y = np.linspace(0, len(percentiles))\n",
|
||||
" ax.plot(x, x, color=\"gray\", lw=1., ls=\"--\")\n",
|
||||
" ax.set_title(title)\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"from matplotlib.lines import Line2D\n",
|
||||
"colors = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"fig, ax = plt.subplots(1,1,dpi=150, figsize=(7, 4))\n",
|
||||
"\n",
|
||||
"percentiles = np.linspace(0.05, 0.95, 19)\n",
|
||||
"PlotCalib(df, ax, \"Wind Speed Calibration\")\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"plt.tick_params(labelsize=16)\n",
|
||||
"sns.despine()\n",
|
||||
"\n",
|
||||
"# custom_lines = [Line2D([0], [0], color=palette[0], lw=2),\n",
|
||||
"# Line2D([0], [0], color=palette[4], lw=2),\n",
|
||||
"# Line2D([0], [0], color=palette[6], lw=2)]\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"# plt.legend(custom_lines, ['LSTM', r\"GP-Matérn\", \"GP-Volt\"],\n",
|
||||
"# fontsize=14, frameon=False, bbox_to_anchor=(0.45, 0.6))\n",
|
||||
"# ax.legend(fontsize=14, bbox_to_anchor=(1., 0.75))\n",
|
||||
"# plt.label(\"Percentile\")\n",
|
||||
"# plt.savefig(\"./wind_calibration.pdf\", bbox_inches=\"tight\")\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "a8726416-26b1-4fea-900f-df2b72385619",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## NLL"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 35,
|
||||
"id": "894e358a-52d3-4e69-baa2-037669015f77",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def GetNLL(model, ema=True, k=100, theta=0.0, horizon=np.arange(75,100), \n",
|
||||
" logger=[], exp=True):\n",
|
||||
" \n",
|
||||
" ntime = full_data[0].shape[0]\n",
|
||||
" ntrain = 400\n",
|
||||
" n_test_times = 20\n",
|
||||
" ntest = 100\n",
|
||||
" test_idxs = torch.arange(ntrain, ntime-ntest, \n",
|
||||
" int((ntime-ntest-ntrain)/n_test_times))\n",
|
||||
" \n",
|
||||
" stns = list(stn_names.keys())\n",
|
||||
" nlls = torch.tensor([])\n",
|
||||
" for stn_save_idx, stn_idx in enumerate(stns):\n",
|
||||
" for test_save_idx, test_idx in enumerate(test_idxs):\n",
|
||||
"\n",
|
||||
" fpath = \"./saved-outputs/stn\" + str(stn_idx) + \"/\"\n",
|
||||
" fname = model + \"_\"\n",
|
||||
" if model == 'volt':\n",
|
||||
" if ema:\n",
|
||||
" fname += \"ema\" + str(k) + \"_\"\n",
|
||||
" fname += \"theta\" + str(theta) + \"_\"\n",
|
||||
" \n",
|
||||
" fname += str(test_idx.item()) + \".pt\"\n",
|
||||
" \n",
|
||||
"# print(fpath + fname)\n",
|
||||
" if os.path.exists(fpath + fname): \n",
|
||||
" if model == 'volt':\n",
|
||||
" preds = torch.load(fpath + fname)[0]\n",
|
||||
" elif model == 'matern':\n",
|
||||
" preds = torch.load(fpath + fname).cpu()\n",
|
||||
" else:\n",
|
||||
" preds = torch.load(fpath + fname)\n",
|
||||
" \n",
|
||||
" preds = preds[:, horizon]\n",
|
||||
" test_y = torch.tensor(full_data[stn_idx][test_idx:]) + 1\n",
|
||||
" \n",
|
||||
"# return preds, test_y\n",
|
||||
" if exp:\n",
|
||||
" preds = preds.exp()\n",
|
||||
" \n",
|
||||
" \n",
|
||||
" try:\n",
|
||||
" curr = torch.distributions.Normal(preds.mean(0), preds.std(0)).log_prob(test_y[horizon])\n",
|
||||
" nlls = torch.cat((curr, nlls))\n",
|
||||
" except:\n",
|
||||
" pass\n",
|
||||
" \n",
|
||||
"\n",
|
||||
" if nlls.numel() > 0:\n",
|
||||
" logger.append([-nlls.sum().item(), -nlls.mean().item(), nlls.std().item(), model, mean, k])\n",
|
||||
" \n",
|
||||
" return logger"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "1a435a51-c4a0-4931-bb3a-755af4cb1c9f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.8"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,450 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 25,
|
||||
"id": "df69e84a-4354-4c9a-90ea-98d52d2c06ed",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import plotly.graph_objects as go\n",
|
||||
"import pandas as pd\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns\n",
|
||||
"import pickle as pkl\n",
|
||||
"import numpy as np\n",
|
||||
"%matplotlib inline\n",
|
||||
"from IPython.display import HTML\n",
|
||||
"sns.set_style('white')\n",
|
||||
"palette = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"sns.set(palette = palette, font_scale=1.5, style=\"white\", rc={\"lines.linewidth\": 3.0})"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 10,
|
||||
"id": "6499b3d9-6fc7-456f-a204-0f5d34fb5957",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import plotly.io as pio\n",
|
||||
"pio.renderers.default = 'jupyterlab'"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 11,
|
||||
"id": "944bda45-2efa-424e-bd94-0c6500891469",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"stn_names, stn_lonlat, stn_dat = pkl.load(open(\"./wind_data.p\", 'rb'))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 12,
|
||||
"id": "cb75bd8a-7ae8-4c28-921b-65a4b104574a",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"lonlat = np.array(list(stn_lonlat.values()))\n",
|
||||
"df = pd.DataFrame(lonlat)\n",
|
||||
"df.columns = ['lon', 'lat']"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 18,
|
||||
"id": "3eddd5e7-4cb4-4cfa-b334-95e31c3a27c9",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"fig = go.Figure(data=go.Scattergeo(\n",
|
||||
" lon = df['lon'],\n",
|
||||
" lat = df['lat'],\n",
|
||||
" mode = 'markers',\n",
|
||||
" ))\n",
|
||||
"\n",
|
||||
"fig.update_layout(\n",
|
||||
" geo_scope='usa',\n",
|
||||
" )\n",
|
||||
"fig.update_yaxes(automargin=True)\n",
|
||||
"HTML(fig.to_html())\n",
|
||||
"fig.write_image(\"./map.pdf\")\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 15,
|
||||
"id": "fa2b366f-d803-485e-9b3f-9a6889caf6ac",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from IPython.display import HTML"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 16,
|
||||
"id": "ba337b77-f5b2-45d0-94bf-94aabbf4d5ab",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/html": [
|
||||
"<html>\n",
|
||||
"<head><meta charset=\"utf-8\" /></head>\n",
|
||||
"<body>\n",
|
||||
" <div> <script type=\"text/javascript\">window.PlotlyConfig = {MathJaxConfig: 'local'};</script>\n",
|
||||
" <script type=\"text/javascript\">/**\n",
|
||||
"* plotly.js v2.2.0\n",
|
||||
"* Copyright 2012-2021, Plotly, Inc.\n",
|
||||
"* All rights reserved.\n",
|
||||
"* Licensed under the MIT license\n",
|
||||
"*/\n",
|
||||
"!function(t){if(\"object\"==typeof exports&&\"undefined\"!=typeof module)module.exports=t();else if(\"function\"==typeof define&&define.amd)define([],t);else{(\"undefined\"!=typeof window?window:\"undefined\"!=typeof global?global:\"undefined\"!=typeof self?self:this).Plotly=t()}}((function(){return function t(e,r,n){function i(o,s){if(!r[o]){if(!e[o]){var l=\"function\"==typeof require&&require;if(!s&&l)return l(o,!0);if(a)return a(o,!0);var c=new Error(\"Cannot find module '\"+o+\"'\");throw c.code=\"MODULE_NOT_FOUND\",c}var u=r[o]={exports:{}};e[o][0].call(u.exports,(function(t){return i(e[o][1][t]||t)}),u,u.exports,t,e,r,n)}return r[o].exports}for(var a=\"function\"==typeof require&&require,o=0;o<n.length;o++)i(n[o]);return i}({1:[function(t,e,r){\"use strict\";var n=t(\"../src/lib\"),i={\"X,X div\":'direction:ltr;font-family:\"Open Sans\",verdana,arial,sans-serif;margin:0;padding:0;',\"X input,X button\":'font-family:\"Open Sans\",verdana,arial,sans-serif;',\"X input:focus,X button:focus\":\"outline:none;\",\"X a\":\"text-decoration:none;\",\"X a:hover\":\"text-decoration:none;\",\"X .crisp\":\"shape-rendering:crispEdges;\",\"X .user-select-none\":\"-webkit-user-select:none;-moz-user-select:none;-ms-user-select:none;-o-user-select:none;user-select:none;\",\"X svg\":\"overflow:hidden;\",\"X svg a\":\"fill:#447adb;\",\"X svg a:hover\":\"fill:#3c6dc5;\",\"X .main-svg\":\"position:absolute;top:0;left:0;pointer-events:none;\",\"X .main-svg .draglayer\":\"pointer-events:all;\",\"X .cursor-default\":\"cursor:default;\",\"X .cursor-pointer\":\"cursor:pointer;\",\"X .cursor-crosshair\":\"cursor:crosshair;\",\"X .cursor-move\":\"cursor:move;\",\"X .cursor-col-resize\":\"cursor:col-resize;\",\"X .cursor-row-resize\":\"cursor:row-resize;\",\"X .cursor-ns-resize\":\"cursor:ns-resize;\",\"X .cursor-ew-resize\":\"cursor:ew-resize;\",\"X .cursor-sw-resize\":\"cursor:sw-resize;\",\"X .cursor-s-resize\":\"cursor:s-resize;\",\"X .cursor-se-resize\":\"cursor:se-resize;\",\"X .cursor-w-resize\":\"cursor:w-resize;\",\"X .cursor-e-resize\":\"cursor:e-resize;\",\"X .cursor-nw-resize\":\"cursor:nw-resize;\",\"X .cursor-n-resize\":\"cursor:n-resize;\",\"X .cursor-ne-resize\":\"cursor:ne-resize;\",\"X .cursor-grab\":\"cursor:-webkit-grab;cursor:grab;\",\"X .modebar\":\"position:absolute;top:2px;right:2px;\",\"X .ease-bg\":\"-webkit-transition:background-color .3s ease 0s;-moz-transition:background-color .3s ease 0s;-ms-transition:background-color .3s ease 0s;-o-transition:background-color .3s ease 0s;transition:background-color .3s ease 0s;\",\"X .modebar--hover>:not(.watermark)\":\"opacity:0;-webkit-transition:opacity .3s ease 0s;-moz-transition:opacity .3s ease 0s;-ms-transition:opacity .3s ease 0s;-o-transition:opacity .3s ease 0s;transition:opacity .3s ease 0s;\",\"X:hover .modebar--hover .modebar-group\":\"opacity:1;\",\"X .modebar-group\":\"float:left;display:inline-block;box-sizing:border-box;padding-left:8px;position:relative;vertical-align:middle;white-space:nowrap;\",\"X .modebar-btn\":\"position:relative;font-size:16px;padding:3px 4px;height:22px;cursor:pointer;line-height:normal;box-sizing:border-box;\",\"X .modebar-btn svg\":\"position:relative;top:2px;\",\"X .modebar.vertical\":\"display:flex;flex-direction:column;flex-wrap:wrap;align-content:flex-end;max-height:100%;\",\"X .modebar.vertical svg\":\"top:-1px;\",\"X .modebar.vertical .modebar-group\":\"display:block;float:none;padding-left:0px;padding-bottom:8px;\",\"X .modebar.vertical .modebar-group .modebar-btn\":\"display:block;text-align:center;\",\"X [data-title]:before,X [data-title]:after\":\"position:absolute;-webkit-transform:translate3d(0, 0, 0);-moz-transform:translate3d(0, 0, 0);-ms-transform:translate3d(0, 0, 0);-o-transform:translate3d(0, 0, 0);transform:translate3d(0, 0, 0);display:none;opacity:0;z-index:1001;pointer-events:none;top:110%;right:50%;\",\"X [data-title]:hover:before,X [data-title]:hover:after\":\"display:block;opacity:1;\",\"X [data-title]:before\":'content:\"\";position:absolute;background:transparent;border:6px solid transparent;z-index:1002;margin-top:-12px;border-bottom-color:#69738a;margin-right:-6px;',\"X [data-title]:after\":\"content:attr(data-title);background:#69738a;color:#fff;padding:8px 10px;font-size:12px;line-height:12px;white-space:nowrap;margin-right:-18px;border-radius:2px;\",\"X .vertical [data-title]:before,X .vertical [data-title]:after\":\"top:0%;right:200%;\",\"X .vertical [data-title]:before\":\"border:6px solid transparent;border-left-color:#69738a;margin-top:8px;margin-right:-30px;\",\"X .select-outline\":\"fill:none;stroke-width:1;shape-rendering:crispEdges;\",\"X .select-outline-1\":\"stroke:#fff;\",\"X .select-outline-2\":\"stroke:#000;stroke-dasharray:2px 2px;\",Y:'font-family:\"Open Sans\",verdana,arial,sans-serif;position:fixed;top:50px;right:20px;z-index:10000;font-size:10pt;max-width:180px;',\"Y p\":\"margin:0;\",\"Y .notifier-note\":\"min-width:180px;max-width:250px;border:1px solid #fff;z-indLine truncated
|
||||
"/*!\n",
|
||||
" * The buffer module from node.js, for the browser.\n",
|
||||
" *\n",
|
||||
" * @author Feross Aboukhadijeh <feross@feross.org> <http://feross.org>\n",
|
||||
" * @license MIT\n",
|
||||
" */function i(t,e){if(t===e)return 0;for(var r=t.length,n=e.length,i=0,a=Math.min(r,n);i<a;++i)if(t[i]!==e[i]){r=t[i],n=e[i];break}return r<n?-1:n<r?1:0}function a(t){return r.Buffer&&\"function\"==typeof r.Buffer.isBuffer?r.Buffer.isBuffer(t):!(null==t||!t._isBuffer)}var o=t(\"util/\"),s=Object.prototype.hasOwnProperty,l=Array.prototype.slice,c=\"foo\"===function(){}.name;function u(t){return Object.prototype.toString.call(t)}function f(t){return!a(t)&&(\"function\"==typeof r.ArrayBuffer&&(\"function\"==typeof ArrayBuffer.isView?ArrayBuffer.isView(t):!!t&&(t instanceof DataView||!!(t.buffer&&t.buffer instanceof ArrayBuffer))))}var h=e.exports=y,p=/\\s*function\\s+([^\\(\\s]*)\\s*/;function d(t){if(o.isFunction(t)){if(c)return t.name;var e=t.toString().match(p);return e&&e[1]}}function m(t,e){return\"string\"==typeof t?t.length<e?t:t.slice(0,e):t}function g(t){if(c||!o.isFunction(t))return o.inspect(t);var e=d(t);return\"[Function\"+(e?\": \"+e:\"\")+\"]\"}function v(t,e,r,n,i){throw new h.AssertionError({message:r,actual:t,expected:e,operator:n,stackStartFunction:i})}function y(t,e){t||v(t,!0,e,\"==\",h.ok)}function x(t,e,r,n){if(t===e)return!0;if(a(t)&&a(e))return 0===i(t,e);if(o.isDate(t)&&o.isDate(e))return t.getTime()===e.getTime();if(o.isRegExp(t)&&o.isRegExp(e))return t.source===e.source&&t.global===e.global&&t.multiline===e.multiline&&t.lastIndex===e.lastIndex&&t.ignoreCase===e.ignoreCase;if(null!==t&&\"object\"==typeof t||null!==e&&\"object\"==typeof e){if(f(t)&&f(e)&&u(t)===u(e)&&!(t instanceof Float32Array||t instanceof Float64Array))return 0===i(new Uint8Array(t.buffer),new Uint8Array(e.buffer));if(a(t)!==a(e))return!1;var s=(n=n||{actual:[],expected:[]}).actual.indexOf(t);return-1!==s&&s===n.expected.indexOf(e)||(n.actual.push(t),n.expected.push(e),function(t,e,r,n){if(null==t||null==e)return!1;if(o.isPrimitive(t)||o.isPrimitive(e))return t===e;if(r&&Object.getPrototypeOf(t)!==Object.getPrototypeOf(e))return!1;var i=b(t),a=b(e);if(i&&!a||!i&&a)return!1;if(i)return t=l.call(t),e=l.call(e),x(t,e,r);var s,c,u=T(t),f=T(e);if(u.length!==f.length)return!1;for(u.sort(),f.sort(),c=u.length-1;c>=0;c--)if(u[c]!==f[c])return!1;for(c=u.length-1;c>=0;c--)if(s=u[c],!x(t[s],e[s],r,n))return!1;return!0}(t,e,r,n))}return r?t===e:t==e}function b(t){return\"[object Arguments]\"==Object.prototype.toString.call(t)}function _(t,e){if(!t||!e)return!1;if(\"[object RegExp]\"==Object.prototype.toString.call(e))return e.test(t);try{if(t instanceof e)return!0}catch(t){}return!Error.isPrototypeOf(e)&&!0===e.call({},t)}function w(t,e,r,n){var i;if(\"function\"!=typeof e)throw new TypeError('\"block\" argument must be a function');\"string\"==typeof r&&(n=r,r=null),i=function(t){var e;try{t()}catch(t){e=t}return e}(e),n=(r&&r.name?\" (\"+r.name+\").\":\".\")+(n?\" \"+n:\".\"),t&&!i&&v(i,r,\"Missing expected exception\"+n);var a=\"string\"==typeof n,s=!t&&i&&!r;if((!t&&o.isError(i)&&a&&_(i,r)||s)&&v(i,r,\"Got unwanted exception\"+n),t&&i&&r&&!_(i,r)||!t&&i)throw i}h.AssertionError=function(t){this.name=\"AssertionError\",this.actual=t.actual,this.expected=t.expected,this.operator=t.operator,t.message?(this.message=t.message,this.generatedMessage=!1):(this.message=function(t){return m(g(t.actual),128)+\" \"+t.operator+\" \"+m(g(t.expected),128)}(this),this.generatedMessage=!0);var e=t.stackStartFunction||v;if(Error.captureStackTrace)Error.captureStackTrace(this,e);else{var r=new Error;if(r.stack){var n=r.stack,i=d(e),a=n.indexOf(\"\\n\"+i);if(a>=0){var o=n.indexOf(\"\\n\",a+1);n=n.substring(o+1)}this.stack=n}}},o.inherits(h.AssertionError,Error),h.fail=v,h.ok=y,h.equal=function(t,e,r){t!=e&&v(t,e,r,\"==\",h.equal)},h.notEqual=function(t,e,r){t==e&&v(t,e,r,\"!=\",h.notEqual)},h.deepEqual=function(t,e,r){x(t,e,!1)||v(t,e,r,\"deepEqual\",h.deepEqual)},h.deepStrictEqual=function(t,e,r){x(t,e,!0)||v(t,e,r,\"deepStrictEqual\",h.deepStrictEqual)},h.notDeepEqual=function(t,e,r){x(t,e,!1)&&v(t,e,r,\"notDeepEqual\",h.notDeepEqual)},h.notDeepStrictEqual=function t(e,r,n){x(e,r,!0)&&v(e,r,n,\"notDeepStrictEqual\",t)},h.strictEqual=function(t,e,r){t!==e&&v(t,e,r,\"===\",h.strictEqual)},h.notStrictEqual=function(t,e,r){t===e&&v(t,e,r,\"!==\",h.notStrictEqual)},h.throws=function(t,e,r){w(!0,t,e,r)},h.doesNotThrow=function(t,e,r){w(!1,t,e,r)},h.ifError=function(t){if(t)throw t},h.strict=n((function t(e,r){e||v(e,!0,r,\"==\",t)}),h,{equal:h.strictEqual,deepEqual:h.deepStrictEqual,notEqual:h.notStrictEqual,notDeepEqual:h.notDeepStrictEqual}),h.strict.strict=h.strict;var T=Object.keys||function(t){var e=[];for(var r in t)s.call(t,r)&&e.push(r);return e}}).call(this)}).call(this,\"undefined\"!=typeof global?global:\"undefined\"!=typeof self?self:\"undefined\"!=typeof window?window:{})},{\"object-assign\":483,\"util/\":83}],81:[function(t,e,r){\"function\"==typeof Object.create?e.exports=function(t,e){t.super_=e,t.prototype=Object.create(e.prototype,{constructor:{value:t,enumerable:!1,writable:!0,configurable:!0}})}:e.exports=function(t,eLine truncated
|
||||
"/*!\n",
|
||||
" * The buffer module from node.js, for the browser.\n",
|
||||
" *\n",
|
||||
" * @author Feross Aboukhadijeh <https://feross.org>\n",
|
||||
" * @license MIT\n",
|
||||
" */\n",
|
||||
"\"use strict\";var e=t(\"base64-js\"),n=t(\"ieee754\");r.Buffer=a,r.SlowBuffer=function(t){+t!=t&&(t=0);return a.alloc(+t)},r.INSPECT_MAX_BYTES=50;function i(t){if(t>2147483647)throw new RangeError('The value \"'+t+'\" is invalid for option \"size\"');var e=new Uint8Array(t);return e.__proto__=a.prototype,e}function a(t,e,r){if(\"number\"==typeof t){if(\"string\"==typeof e)throw new TypeError('The \"string\" argument must be of type string. Received type number');return l(t)}return o(t,e,r)}function o(t,e,r){if(\"string\"==typeof t)return function(t,e){\"string\"==typeof e&&\"\"!==e||(e=\"utf8\");if(!a.isEncoding(e))throw new TypeError(\"Unknown encoding: \"+e);var r=0|f(t,e),n=i(r),o=n.write(t,e);o!==r&&(n=n.slice(0,o));return n}(t,e);if(ArrayBuffer.isView(t))return c(t);if(null==t)throw TypeError(\"The first argument must be one of type string, Buffer, ArrayBuffer, Array, or Array-like Object. Received type \"+typeof t);if(B(t,ArrayBuffer)||t&&B(t.buffer,ArrayBuffer))return function(t,e,r){if(e<0||t.byteLength<e)throw new RangeError('\"offset\" is outside of buffer bounds');if(t.byteLength<e+(r||0))throw new RangeError('\"length\" is outside of buffer bounds');var n;n=void 0===e&&void 0===r?new Uint8Array(t):void 0===r?new Uint8Array(t,e):new Uint8Array(t,e,r);return n.__proto__=a.prototype,n}(t,e,r);if(\"number\"==typeof t)throw new TypeError('The \"value\" argument must not be of type number. Received type number');var n=t.valueOf&&t.valueOf();if(null!=n&&n!==t)return a.from(n,e,r);var o=function(t){if(a.isBuffer(t)){var e=0|u(t.length),r=i(e);return 0===r.length||t.copy(r,0,0,e),r}if(void 0!==t.length)return\"number\"!=typeof t.length||N(t.length)?i(0):c(t);if(\"Buffer\"===t.type&&Array.isArray(t.data))return c(t.data)}(t);if(o)return o;if(\"undefined\"!=typeof Symbol&&null!=Symbol.toPrimitive&&\"function\"==typeof t[Symbol.toPrimitive])return a.from(t[Symbol.toPrimitive](\"string\"),e,r);throw new TypeError(\"The first argument must be one of type string, Buffer, ArrayBuffer, Array, or Array-like Object. Received type \"+typeof t)}function s(t){if(\"number\"!=typeof t)throw new TypeError('\"size\" argument must be of type number');if(t<0)throw new RangeError('The value \"'+t+'\" is invalid for option \"size\"')}function l(t){return s(t),i(t<0?0:0|u(t))}function c(t){for(var e=t.length<0?0:0|u(t.length),r=i(e),n=0;n<e;n+=1)r[n]=255&t[n];return r}function u(t){if(t>=2147483647)throw new RangeError(\"Attempt to allocate Buffer larger than maximum size: 0x\"+2147483647..toString(16)+\" bytes\");return 0|t}function f(t,e){if(a.isBuffer(t))return t.length;if(ArrayBuffer.isView(t)||B(t,ArrayBuffer))return t.byteLength;if(\"string\"!=typeof t)throw new TypeError('The \"string\" argument must be one of type string, Buffer, or ArrayBuffer. Received type '+typeof t);var r=t.length,n=arguments.length>2&&!0===arguments[2];if(!n&&0===r)return 0;for(var i=!1;;)switch(e){case\"ascii\":case\"latin1\":case\"binary\":return r;case\"utf8\":case\"utf-8\":return D(t).length;case\"ucs2\":case\"ucs-2\":case\"utf16le\":case\"utf-16le\":return 2*r;case\"hex\":return r>>>1;case\"base64\":return R(t).length;default:if(i)return n?-1:D(t).length;e=(\"\"+e).toLowerCase(),i=!0}}function h(t,e,r){var n=!1;if((void 0===e||e<0)&&(e=0),e>this.length)return\"\";if((void 0===r||r>this.length)&&(r=this.length),r<=0)return\"\";if((r>>>=0)<=(e>>>=0))return\"\";for(t||(t=\"utf8\");;)switch(t){case\"hex\":return A(this,e,r);case\"utf8\":case\"utf-8\":return T(this,e,r);case\"ascii\":return k(this,e,r);case\"latin1\":case\"binary\":return M(this,e,r);case\"base64\":return w(this,e,r);case\"ucs2\":case\"ucs-2\":case\"utf16le\":case\"utf-16le\":return S(this,e,r);default:if(n)throw new TypeError(\"Unknown encoding: \"+t);t=(t+\"\").toLowerCase(),n=!0}}function p(t,e,r){var n=t[e];t[e]=t[r],t[r]=n}function d(t,e,r,n,i){if(0===t.length)return-1;if(\"string\"==typeof r?(n=r,r=0):r>2147483647?r=2147483647:r<-2147483648&&(r=-2147483648),N(r=+r)&&(r=i?0:t.length-1),r<0&&(r=t.length+r),r>=t.length){if(i)return-1;r=t.length-1}else if(r<0){if(!i)return-1;r=0}if(\"string\"==typeof e&&(e=a.from(e,n)),a.isBuffer(e))return 0===e.length?-1:m(t,e,r,n,i);if(\"number\"==typeof e)return e&=255,\"function\"==typeof Uint8Array.prototype.indexOf?i?Uint8Array.prototype.indexOf.call(t,e,r):Uint8Array.prototype.lastIndexOf.call(t,e,r):m(t,[e],r,n,i);throw new TypeError(\"val must be string, number or Buffer\")}function m(t,e,r,n,i){var a,o=1,s=t.length,l=e.length;if(void 0!==n&&(\"ucs2\"===(n=String(n).toLowerCase())||\"ucs-2\"===n||\"utf16le\"===n||\"utf-16le\"===n)){if(t.length<2||e.length<2)return-1;o=2,s/=2,l/=2,r/=2}function c(t,e){return 1===o?t[e]:t.readUInt16BE(e*o)}if(i){var u=-1;for(a=r;a<s;a++)if(c(t,a)===c(e,-1===u?0:a-u)){if(-1===u&&(u=a),a-u+1===l)return u*o}else-1!==u&&(a-=a-u),u=-1}else for(r+l>s&&(r=s-l),a=r;a>=0;a--){for(var f=!0,h=0;h<l;h++)if(c(t,a+h)!==c(e,h)){f=!1;break}if(f)return a}return-1}function g(t,e,r,n){r=Number(r)||0;var i=t.lengLine truncated
|
||||
"/*!\n",
|
||||
" * Determine if an object is a Buffer\n",
|
||||
" *\n",
|
||||
" * @author Feross Aboukhadijeh <https://feross.org>\n",
|
||||
" * @license MIT\n",
|
||||
" */\n",
|
||||
"e.exports=function(t){return null!=t&&(n(t)||function(t){return\"function\"==typeof t.readFloatLE&&\"function\"==typeof t.slice&&n(t.slice(0,0))}(t)||!!t._isBuffer)}},{}],450:[function(t,e,r){\"use strict\";e.exports=\"undefined\"!=typeof navigator&&(/MSIE/.test(navigator.userAgent)||/Trident\\//.test(navigator.appVersion))},{}],451:[function(t,e,r){\"use strict\";e.exports=a,e.exports.isMobile=a,e.exports.default=a;var n=/(android|bb\\d+|meego).+mobile|avantgo|bada\\/|blackberry|blazer|compal|elaine|fennec|hiptop|iemobile|ip(hone|od)|iris|kindle|lge |maemo|midp|mmp|mobile.+firefox|netfront|opera m(ob|in)i|palm( os)?|phone|p(ixi|re)\\/|plucker|pocket|psp|series[46]0|symbian|treo|up\\.(browser|link)|vodafone|wap|windows (ce|phone)|xda|xiino/i,i=/(android|bb\\d+|meego).+mobile|avantgo|bada\\/|blackberry|blazer|compal|elaine|fennec|hiptop|iemobile|ip(hone|od)|iris|kindle|lge |maemo|midp|mmp|mobile.+firefox|netfront|opera m(ob|in)i|palm( os)?|phone|p(ixi|re)\\/|plucker|pocket|psp|series[46]0|symbian|treo|up\\.(browser|link)|vodafone|wap|windows (ce|phone)|xda|xiino|android|ipad|playbook|silk/i;function a(t){t||(t={});var e=t.ua;if(e||\"undefined\"==typeof navigator||(e=navigator.userAgent),e&&e.headers&&\"string\"==typeof e.headers[\"user-agent\"]&&(e=e.headers[\"user-agent\"]),\"string\"!=typeof e)return!1;var r=t.tablet?i.test(e):n.test(e);return!r&&t.tablet&&t.featureDetect&&navigator&&navigator.maxTouchPoints>1&&-1!==e.indexOf(\"Macintosh\")&&-1!==e.indexOf(\"Safari\")&&(r=!0),r}},{}],452:[function(t,e,r){\"use strict\";e.exports=function(t){var e=typeof t;return null!==t&&(\"object\"===e||\"function\"===e)}},{}],453:[function(t,e,r){\"use strict\";var n=Object.prototype.toString;e.exports=function(t){var e;return\"[object Object]\"===n.call(t)&&(null===(e=Object.getPrototypeOf(t))||e===Object.getPrototypeOf({}))}},{}],454:[function(t,e,r){\"use strict\";e.exports=function(t){for(var e,r=t.length,n=0;n<r;n++)if(((e=t.charCodeAt(n))<9||e>13)&&32!==e&&133!==e&&160!==e&&5760!==e&&6158!==e&&(e<8192||e>8205)&&8232!==e&&8233!==e&&8239!==e&&8287!==e&&8288!==e&&12288!==e&&65279!==e)return!1;return!0}},{}],455:[function(t,e,r){\"use strict\";e.exports=function(t){return\"string\"==typeof t&&(t=t.trim(),!!(/^[mzlhvcsqta]\\s*[-+.0-9][^mlhvzcsqta]+/i.test(t)&&/[\\dz]$/i.test(t)&&t.length>4))}},{}],456:[function(t,e,r){e.exports=function(t,e,r){return t*(1-r)+e*r}},{}],457:[function(t,e,r){!function(t,n){\"object\"==typeof r&&void 0!==e?e.exports=n():(t=t||self).mapboxgl=n()}(this,(function(){\"use strict\";var t,e,r;function n(n,i){if(t)if(e){var a=\"var sharedChunk = {}; (\"+t+\")(sharedChunk); (\"+e+\")(sharedChunk);\",o={};t(o),(r=i(o)).workerUrl=window.URL.createObjectURL(new Blob([a],{type:\"text/javascript\"}))}else e=i;else t=i}return n(0,(function(t){function e(t,e){return t(e={exports:{}},e.exports),e.exports}var r=n;function n(t,e,r,n){this.cx=3*t,this.bx=3*(r-t)-this.cx,this.ax=1-this.cx-this.bx,this.cy=3*e,this.by=3*(n-e)-this.cy,this.ay=1-this.cy-this.by,this.p1x=t,this.p1y=n,this.p2x=r,this.p2y=n}n.prototype.sampleCurveX=function(t){return((this.ax*t+this.bx)*t+this.cx)*t},n.prototype.sampleCurveY=function(t){return((this.ay*t+this.by)*t+this.cy)*t},n.prototype.sampleCurveDerivativeX=function(t){return(3*this.ax*t+2*this.bx)*t+this.cx},n.prototype.solveCurveX=function(t,e){var r,n,i,a,o;for(void 0===e&&(e=1e-6),i=t,o=0;o<8;o++){if(a=this.sampleCurveX(i)-t,Math.abs(a)<e)return i;var s=this.sampleCurveDerivativeX(i);if(Math.abs(s)<1e-6)break;i-=a/s}if((i=t)<(r=0))return r;if(i>(n=1))return n;for(;r<n;){if(a=this.sampleCurveX(i),Math.abs(a-t)<e)return i;t>a?r=i:n=i,i=.5*(n-r)+r}return i},n.prototype.solve=function(t,e){return this.sampleCurveY(this.solveCurveX(t,e))};var i=a;function a(t,e){this.x=t,this.y=e}function o(t,e,n,i){var a=new r(t,e,n,i);return function(t){return a.solve(t)}}a.prototype={clone:function(){return new a(this.x,this.y)},add:function(t){return this.clone()._add(t)},sub:function(t){return this.clone()._sub(t)},multByPoint:function(t){return this.clone()._multByPoint(t)},divByPoint:function(t){return this.clone()._divByPoint(t)},mult:function(t){return this.clone()._mult(t)},div:function(t){return this.clone()._div(t)},rotate:function(t){return this.clone()._rotate(t)},rotateAround:function(t,e){return this.clone()._rotateAround(t,e)},matMult:function(t){return this.clone()._matMult(t)},unit:function(){return this.clone()._unit()},perp:function(){return this.clone()._perp()},round:function(){return this.clone()._round()},mag:function(){return Math.sqrt(this.x*this.x+this.y*this.y)},equals:function(t){return this.x===t.x&&this.y===t.y},dist:function(t){return Math.sqrt(this.distSqr(t))},distSqr:function(t){var e=t.x-this.x,r=t.y-this.y;return e*e+r*r},angle:function(){return Math.atan2(this.y,this.x)},angleTo:function(t){return Math.atan2(this.y-t.y,this.x-t.x)},angleWith:function(t){return this.angleWithSep(t.x,t.y)},angleWithSep:function(t,e){return Math.atan2(this.x*e-this.y*tLine truncated
|
||||
"/*! Native Promise Only\n",
|
||||
" v0.8.1 (c) Kyle Simpson\n",
|
||||
" MIT License: http://getify.mit-license.org\n",
|
||||
"*/\n",
|
||||
"!function(t,r,n){r[t]=r[t]||n(),void 0!==e&&e.exports&&(e.exports=r[t])}(\"Promise\",void 0!==t?t:this,(function(){\"use strict\";var t,e,n,i=Object.prototype.toString,a=void 0!==r?function(t){return r(t)}:setTimeout;try{Object.defineProperty({},\"x\",{}),t=function(t,e,r,n){return Object.defineProperty(t,e,{value:r,writable:!0,configurable:!1!==n})}}catch(e){t=function(t,e,r){return t[e]=r,t}}function o(t,r){n.add(t,r),e||(e=a(n.drain))}function s(t){var e,r=typeof t;return null==t||\"object\"!=r&&\"function\"!=r||(e=t.then),\"function\"==typeof e&&e}function l(){for(var t=0;t<this.chain.length;t++)c(this,1===this.state?this.chain[t].success:this.chain[t].failure,this.chain[t]);this.chain.length=0}function c(t,e,r){var n,i;try{!1===e?r.reject(t.msg):(n=!0===e?t.msg:e.call(void 0,t.msg))===r.promise?r.reject(TypeError(\"Promise-chain cycle\")):(i=s(n))?i.call(n,r.resolve,r.reject):r.resolve(n)}catch(t){r.reject(t)}}function u(t){var e,r=this;if(!r.triggered){r.triggered=!0,r.def&&(r=r.def);try{(e=s(t))?o((function(){var n=new p(r);try{e.call(t,(function(){u.apply(n,arguments)}),(function(){f.apply(n,arguments)}))}catch(t){f.call(n,t)}})):(r.msg=t,r.state=1,r.chain.length>0&&o(l,r))}catch(t){f.call(new p(r),t)}}}function f(t){var e=this;e.triggered||(e.triggered=!0,e.def&&(e=e.def),e.msg=t,e.state=2,e.chain.length>0&&o(l,e))}function h(t,e,r,n){for(var i=0;i<e.length;i++)!function(i){t.resolve(e[i]).then((function(t){r(i,t)}),n)}(i)}function p(t){this.def=t,this.triggered=!1}function d(t){this.promise=t,this.state=0,this.triggered=!1,this.chain=[],this.msg=void 0}function m(t){if(\"function\"!=typeof t)throw TypeError(\"Not a function\");if(0!==this.__NPO__)throw TypeError(\"Not a promise\");this.__NPO__=1;var e=new d(this);this.then=function(t,r){var n={success:\"function\"!=typeof t||t,failure:\"function\"==typeof r&&r};return n.promise=new this.constructor((function(t,e){if(\"function\"!=typeof t||\"function\"!=typeof e)throw TypeError(\"Not a function\");n.resolve=t,n.reject=e})),e.chain.push(n),0!==e.state&&o(l,e),n.promise},this.catch=function(t){return this.then(void 0,t)};try{t.call(void 0,(function(t){u.call(e,t)}),(function(t){f.call(e,t)}))}catch(t){f.call(e,t)}}n=function(){var t,r,n;function i(t,e){this.fn=t,this.self=e,this.next=void 0}return{add:function(e,a){n=new i(e,a),r?r.next=n:t=n,r=n,n=void 0},drain:function(){var n=t;for(t=r=e=void 0;n;)n.fn.call(n.self),n=n.next}}}();var g=t({},\"constructor\",m,!1);return m.prototype=g,t(g,\"__NPO__\",0,!1),t(m,\"resolve\",(function(t){return t&&\"object\"==typeof t&&1===t.__NPO__?t:new this((function(e,r){if(\"function\"!=typeof e||\"function\"!=typeof r)throw TypeError(\"Not a function\");e(t)}))})),t(m,\"reject\",(function(t){return new this((function(e,r){if(\"function\"!=typeof e||\"function\"!=typeof r)throw TypeError(\"Not a function\");r(t)}))})),t(m,\"all\",(function(t){var e=this;return\"[object Array]\"!=i.call(t)?e.reject(TypeError(\"Not an array\")):0===t.length?e.resolve([]):new e((function(r,n){if(\"function\"!=typeof r||\"function\"!=typeof n)throw TypeError(\"Not a function\");var i=t.length,a=Array(i),o=0;h(e,t,(function(t,e){a[t]=e,++o===i&&r(a)}),n)}))})),t(m,\"race\",(function(t){var e=this;return\"[object Array]\"!=i.call(t)?e.reject(TypeError(\"Not an array\")):new e((function(r,n){if(\"function\"!=typeof r||\"function\"!=typeof n)throw TypeError(\"Not a function\");h(e,t,(function(t,e){r(e)}),n)}))})),m}))}).call(this)}).call(this,\"undefined\"!=typeof global?global:\"undefined\"!=typeof self?self:\"undefined\"!=typeof window?window:{},t(\"timers\").setImmediate)},{timers:593}],471:[function(t,e,r){\"use strict\";var n=t(\"typedarray-pool\");function i(t){return\"a\"+t}function a(t){return\"d\"+t}function o(t,e){return\"c\"+t+\"_\"+e}function s(t){return\"s\"+t}function l(t,e){return\"t\"+t+\"_\"+e}function c(t){return\"o\"+t}function u(t){return\"x\"+t}function f(t){return\"p\"+t}function h(t,e){return\"d\"+t+\"_\"+e}function p(t){return\"i\"+t}function d(t,e){return\"u\"+t+\"_\"+e}function m(t){return\"b\"+t}function g(t){return\"y\"+t}function v(t){return\"e\"+t}function y(t){return\"v\"+t}e.exports=function(t){function e(t){throw new Error(\"ndarray-extract-contour: \"+t)}\"object\"!=typeof t&&e(\"Must specify arguments\");var r=t.order;Array.isArray(r)||e(\"Must specify order\");var b=t.arrayArguments||1;b<1&&e(\"Must have at least one array argument\");var _=t.scalarArguments||0;_<0&&e(\"Scalar arg count must be > 0\");\"function\"!=typeof t.vertex&&e(\"Must specify vertex creation function\");\"function\"!=typeof t.cell&&e(\"Must specify cell creation function\");\"function\"!=typeof t.phase&&e(\"Must specify phase function\");for(var w=t.getters||[],T=new Array(b),k=0;k<b;++k)w.indexOf(k)>=0?T[k]=!0:T[k]=!1;return function(t,e,r,b,_,w){var T=w.length,k=_.length;if(k<2)throw new Error(\"ndarray-extract-contour: Dimension must be at least 2\");for(var M=\"extractContour\"+_.join(\"_\"),A=[],S=[],E=[],L=0;L<T;++LLine truncated
|
||||
"/*\n",
|
||||
"object-assign\n",
|
||||
"(c) Sindre Sorhus\n",
|
||||
"@license MIT\n",
|
||||
"*/\n",
|
||||
"\"use strict\";var n=Object.getOwnPropertySymbols,i=Object.prototype.hasOwnProperty,a=Object.prototype.propertyIsEnumerable;function o(t){if(null==t)throw new TypeError(\"Object.assign cannot be called with null or undefined\");return Object(t)}e.exports=function(){try{if(!Object.assign)return!1;var t=new String(\"abc\");if(t[5]=\"de\",\"5\"===Object.getOwnPropertyNames(t)[0])return!1;for(var e={},r=0;r<10;r++)e[\"_\"+String.fromCharCode(r)]=r;if(\"0123456789\"!==Object.getOwnPropertyNames(e).map((function(t){return e[t]})).join(\"\"))return!1;var n={};return\"abcdefghijklmnopqrst\".split(\"\").forEach((function(t){n[t]=t})),\"abcdefghijklmnopqrst\"===Object.keys(Object.assign({},n)).join(\"\")}catch(t){return!1}}()?Object.assign:function(t,e){for(var r,s,l=o(t),c=1;c<arguments.length;c++){for(var u in r=Object(arguments[c]))i.call(r,u)&&(l[u]=r[u]);if(n){s=n(r);for(var f=0;f<s.length;f++)a.call(r,s[f])&&(l[s[f]]=r[s[f]])}}return l}},{}],484:[function(t,e,r){\"use strict\";e.exports=function(t,e,r,n,i,a,o,s,l,c){var u=e+a+c;if(f>0){var f=Math.sqrt(u+1);t[0]=.5*(o-l)/f,t[1]=.5*(s-n)/f,t[2]=.5*(r-a)/f,t[3]=.5*f}else{var h=Math.max(e,a,c);f=Math.sqrt(2*h-u+1);e>=h?(t[0]=.5*f,t[1]=.5*(i+r)/f,t[2]=.5*(s+n)/f,t[3]=.5*(o-l)/f):a>=h?(t[0]=.5*(r+i)/f,t[1]=.5*f,t[2]=.5*(l+o)/f,t[3]=.5*(s-n)/f):(t[0]=.5*(n+s)/f,t[1]=.5*(o+l)/f,t[2]=.5*f,t[3]=.5*(r-i)/f)}return t}},{}],485:[function(t,e,r){\"use strict\";e.exports=function(t){var e=(t=t||{}).center||[0,0,0],r=t.rotation||[0,0,0,1],n=t.radius||1;e=[].slice.call(e,0,3),u(r=[].slice.call(r,0,4),r);var i=new f(r,e,Math.log(n));i.setDistanceLimits(t.zoomMin,t.zoomMax),(\"eye\"in t||\"up\"in t)&&i.lookAt(0,t.eye,t.center,t.up);return i};var n=t(\"filtered-vector\"),i=t(\"gl-mat4/lookAt\"),a=t(\"gl-mat4/fromQuat\"),o=t(\"gl-mat4/invert\"),s=t(\"./lib/quatFromFrame\");function l(t,e,r){return Math.sqrt(Math.pow(t,2)+Math.pow(e,2)+Math.pow(r,2))}function c(t,e,r,n){return Math.sqrt(Math.pow(t,2)+Math.pow(e,2)+Math.pow(r,2)+Math.pow(n,2))}function u(t,e){var r=e[0],n=e[1],i=e[2],a=e[3],o=c(r,n,i,a);o>1e-6?(t[0]=r/o,t[1]=n/o,t[2]=i/o,t[3]=a/o):(t[0]=t[1]=t[2]=0,t[3]=1)}function f(t,e,r){this.radius=n([r]),this.center=n(e),this.rotation=n(t),this.computedRadius=this.radius.curve(0),this.computedCenter=this.center.curve(0),this.computedRotation=this.rotation.curve(0),this.computedUp=[.1,0,0],this.computedEye=[.1,0,0],this.computedMatrix=[.1,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0],this.recalcMatrix(0)}var h=f.prototype;h.lastT=function(){return Math.max(this.radius.lastT(),this.center.lastT(),this.rotation.lastT())},h.recalcMatrix=function(t){this.radius.curve(t),this.center.curve(t),this.rotation.curve(t);var e=this.computedRotation;u(e,e);var r=this.computedMatrix;a(r,e);var n=this.computedCenter,i=this.computedEye,o=this.computedUp,s=Math.exp(this.computedRadius[0]);i[0]=n[0]+s*r[2],i[1]=n[1]+s*r[6],i[2]=n[2]+s*r[10],o[0]=r[1],o[1]=r[5],o[2]=r[9];for(var l=0;l<3;++l){for(var c=0,f=0;f<3;++f)c+=r[l+4*f]*i[f];r[12+l]=-c}},h.getMatrix=function(t,e){this.recalcMatrix(t);var r=this.computedMatrix;if(e){for(var n=0;n<16;++n)e[n]=r[n];return e}return r},h.idle=function(t){this.center.idle(t),this.radius.idle(t),this.rotation.idle(t)},h.flush=function(t){this.center.flush(t),this.radius.flush(t),this.rotation.flush(t)},h.pan=function(t,e,r,n){e=e||0,r=r||0,n=n||0,this.recalcMatrix(t);var i=this.computedMatrix,a=i[1],o=i[5],s=i[9],c=l(a,o,s);a/=c,o/=c,s/=c;var u=i[0],f=i[4],h=i[8],p=u*a+f*o+h*s,d=l(u-=a*p,f-=o*p,h-=s*p);u/=d,f/=d,h/=d;var m=i[2],g=i[6],v=i[10],y=m*a+g*o+v*s,x=m*u+g*f+v*h,b=l(m-=y*a+x*u,g-=y*o+x*f,v-=y*s+x*h);m/=b,g/=b,v/=b;var _=u*e+a*r,w=f*e+o*r,T=h*e+s*r;this.center.move(t,_,w,T);var k=Math.exp(this.computedRadius[0]);k=Math.max(1e-4,k+n),this.radius.set(t,Math.log(k))},h.rotate=function(t,e,r,n){this.recalcMatrix(t),e=e||0,r=r||0;var i=this.computedMatrix,a=i[0],o=i[4],s=i[8],u=i[1],f=i[5],h=i[9],p=i[2],d=i[6],m=i[10],g=e*a+r*u,v=e*o+r*f,y=e*s+r*h,x=-(d*y-m*v),b=-(m*g-p*y),_=-(p*v-d*g),w=Math.sqrt(Math.max(0,1-Math.pow(x,2)-Math.pow(b,2)-Math.pow(_,2))),T=c(x,b,_,w);T>1e-6?(x/=T,b/=T,_/=T,w/=T):(x=b=_=0,w=1);var k=this.computedRotation,M=k[0],A=k[1],S=k[2],E=k[3],L=M*w+E*x+A*_-S*b,C=A*w+E*b+S*x-M*_,P=S*w+E*_+M*b-A*x,I=E*w-M*x-A*b-S*_;if(n){x=p,b=d,_=m;var O=Math.sin(n)/l(x,b,_);x*=O,b*=O,_*=O,I=I*(w=Math.cos(e))-(L=L*w+I*x+C*_-P*b)*x-(C=C*w+I*b+P*x-L*_)*b-(P=P*w+I*_+L*b-C*x)*_}var z=c(L,C,P,I);z>1e-6?(L/=z,C/=z,P/=z,I/=z):(L=C=P=0,I=1),this.rotation.set(t,L,C,P,I)},h.lookAt=function(t,e,r,n){this.recalcMatrix(t),r=r||this.computedCenter,e=e||this.computedEye,n=n||this.computedUp;var a=this.computedMatrix;i(a,e,r,n);var o=this.computedRotation;s(o,a[0],a[1],a[2],a[4],a[5],a[6],a[8],a[9],a[10]),u(o,o),this.rotation.set(t,o[0],o[1],o[2],o[3]);for(var l=0,c=0;c<3;++c)l+=Math.pow(r[c]-e[c],2);this.radius.set(t,.5*Math.log(Math.max(l,1e-6))),this.center.set(t,r[0],r[1],r[2])},h.translate=function(t,e,r,n){this.center.move(t,e||0,r||0,n||0)},h.setMatrix=function(t,e){var r=thLine truncated
|
||||
"/*!\n",
|
||||
" * pad-left <https://github.com/jonschlinkert/pad-left>\n",
|
||||
" *\n",
|
||||
" * Copyright (c) 2014-2015, Jon Schlinkert.\n",
|
||||
" * Licensed under the MIT license.\n",
|
||||
" */\n",
|
||||
"\"use strict\";var n=t(\"repeat-string\");e.exports=function(t,e,r){return n(r=void 0!==r?r+\"\":\" \",e)+t}},{\"repeat-string\":537}],487:[function(t,e,r){\"use strict\";function n(t,e){if(\"string\"!=typeof t)return[t];var r=[t];\"string\"==typeof e||Array.isArray(e)?e={brackets:e}:e||(e={});var n=e.brackets?Array.isArray(e.brackets)?e.brackets:[e.brackets]:[\"{}\",\"[]\",\"()\"],i=e.escape||\"___\",a=!!e.flat;n.forEach((function(t){var e=new RegExp([\"\\\\\",t[0],\"[^\\\\\",t[0],\"\\\\\",t[1],\"]*\\\\\",t[1]].join(\"\")),n=[];function a(e,a,o){var s=r.push(e.slice(t[0].length,-t[1].length))-1;return n.push(s),i+s+i}r.forEach((function(t,n){for(var i,o=0;t!=i;)if(i=t,t=t.replace(e,a),o++>1e4)throw Error(\"References have circular dependency. Please, check them.\");r[n]=t})),n=n.reverse(),r=r.map((function(e){return n.forEach((function(r){e=e.replace(new RegExp(\"(\\\\\"+i+r+\"\\\\\"+i+\")\",\"g\"),t[0]+\"$1\"+t[1])})),e}))}));var o=new RegExp(\"\\\\\"+i+\"([0-9]+)\\\\\"+i);return a?r:function t(e,r,n){for(var i,a=[],s=0;i=o.exec(e);){if(s++>1e4)throw Error(\"Circular references in parenthesis\");a.push(e.slice(0,i.index)),a.push(t(r[i[1]],r)),e=e.slice(i.index+i[0].length)}return a.push(e),a}(r[0],r)}function i(t,e){if(e&&e.flat){var r,n=e&&e.escape||\"___\",i=t[0];if(!i)return\"\";for(var a=new RegExp(\"\\\\\"+n+\"([0-9]+)\\\\\"+n),o=0;i!=r;){if(o++>1e4)throw Error(\"Circular references in \"+t);r=i,i=i.replace(a,s)}return i}return t.reduce((function t(e,r){return Array.isArray(r)&&(r=r.reduce(t,\"\")),e+r}),\"\");function s(e,r){if(null==t[r])throw Error(\"Reference \"+r+\"is undefined\");return t[r]}}function a(t,e){return Array.isArray(t)?i(t,e):n(t,e)}a.parse=n,a.stringify=i,e.exports=a},{}],488:[function(t,e,r){\"use strict\";var n=t(\"pick-by-alias\");e.exports=function(t){var e;arguments.length>1&&(t=arguments);\"string\"==typeof t?t=t.split(/\\s/).map(parseFloat):\"number\"==typeof t&&(t=[t]);t.length&&\"number\"==typeof t[0]?e=1===t.length?{width:t[0],height:t[0],x:0,y:0}:2===t.length?{width:t[0],height:t[1],x:0,y:0}:{x:t[0],y:t[1],width:t[2]-t[0]||0,height:t[3]-t[1]||0}:t&&(t=n(t,{left:\"x l left Left\",top:\"y t top Top\",width:\"w width W Width\",height:\"h height W Width\",bottom:\"b bottom Bottom\",right:\"r right Right\"}),e={x:t.left||0,y:t.top||0},null==t.width?t.right?e.width=t.right-e.x:e.width=0:e.width=t.width,null==t.height?t.bottom?e.height=t.bottom-e.y:e.height=0:e.height=t.height);return e}},{\"pick-by-alias\":494}],489:[function(t,e,r){e.exports=function(t){var e=[];return t.replace(i,(function(t,r,i){var o=r.toLowerCase();for(i=function(t){var e=t.match(a);return e?e.map(Number):[]}(i),\"m\"==o&&i.length>2&&(e.push([r].concat(i.splice(0,2))),o=\"l\",r=\"m\"==r?\"l\":\"L\");;){if(i.length==n[o])return i.unshift(r),e.push(i);if(i.length<n[o])throw new Error(\"malformed path data\");e.push([r].concat(i.splice(0,n[o])))}})),e};var n={a:7,c:6,h:1,l:2,m:2,q:4,s:4,t:2,v:1,z:0},i=/([astvzqmhlc])([^astvzqmhlc]*)/gi;var a=/-?[0-9]*\\.?[0-9]+(?:e[-+]?\\d+)?/gi},{}],490:[function(t,e,r){e.exports=function(t,e){e||(e=[0,\"\"]),t=String(t);var r=parseFloat(t,10);return e[0]=r,e[1]=t.match(/[\\d.\\-\\+]*\\s*(.*)/)[1]||\"\",e}},{}],491:[function(t,e,r){(function(t){(function(){(function(){var r,n,i,a,o,s;\"undefined\"!=typeof performance&&null!==performance&&performance.now?e.exports=function(){return performance.now()}:null!=t&&t.hrtime?(e.exports=function(){return(r()-o)/1e6},n=t.hrtime,a=(r=function(){var t;return 1e9*(t=n())[0]+t[1]})(),s=1e9*t.uptime(),o=a-s):Date.now?(e.exports=function(){return Date.now()-i},i=Date.now()):(e.exports=function(){return(new Date).getTime()-i},i=(new Date).getTime())}).call(this)}).call(this)}).call(this,t(\"_process\"))},{_process:524}],492:[function(t,e,r){\"use strict\";e.exports=function(t){var e=t.length;if(e<32){for(var r=1,i=0;i<e;++i)for(var a=0;a<i;++a)if(t[i]<t[a])r=-r;else if(t[i]===t[a])return 0;return r}var o=n.mallocUint8(e);for(i=0;i<e;++i)o[i]=0;for(r=1,i=0;i<e;++i)if(!o[i]){var s=1;o[i]=1;for(a=t[i];a!==i;a=t[a]){if(o[a])return n.freeUint8(o),0;s+=1,o[a]=1}1&s||(r=-r)}return n.freeUint8(o),r};var n=t(\"typedarray-pool\")},{\"typedarray-pool\":613}],493:[function(t,e,r){\"use strict\";var n=t(\"typedarray-pool\"),i=t(\"invert-permutation\");r.rank=function(t){var e=t.length;switch(e){case 0:case 1:return 0;case 2:return t[1]}var r,a,o,s=n.mallocUint32(e),l=n.mallocUint32(e),c=0;for(i(t,l),o=0;o<e;++o)s[o]=t[o];for(o=e-1;o>0;--o)a=l[o],r=s[o],s[o]=s[a],s[a]=r,l[o]=l[r],l[r]=a,c=(c+r)*o;return n.freeUint32(l),n.freeUint32(s),c},r.unrank=function(t,e,r){switch(t){case 0:return r||[];case 1:return r?(r[0]=0,r):[0];case 2:return r?(e?(r[0]=0,r[1]=1):(r[0]=1,r[1]=0),r):e?[0,1]:[1,0]}var n,i,a,o=1;for((r=r||new Array(t))[0]=0,a=1;a<t;++a)r[a]=a,o=o*a|0;for(a=t-1;a>0;--a)e=e-(n=e/o|0)*o|0,o=o/a|0,i=0|r[a],r[a]=0|r[n],r[n]=0|i;return r}},{\"invert-permutation\":446,\"typedarray-pool\":613}],494:[function(t,e,r){\"use strict\";e.exports=functLine truncated
|
||||
"/*\n",
|
||||
" * @copyright 2016 Sean Connelly (@voidqk), http://syntheti.cc\n",
|
||||
" * @license MIT\n",
|
||||
" * @preserve Project Home: https://github.com/voidqk/polybooljs\n",
|
||||
" */\n",
|
||||
"var n,i=t(\"./lib/build-log\"),a=t(\"./lib/epsilon\"),o=t(\"./lib/intersecter\"),s=t(\"./lib/segment-chainer\"),l=t(\"./lib/segment-selector\"),c=t(\"./lib/geojson\"),u=!1,f=a();function h(t,e,r){var i=n.segments(t),a=n.segments(e),o=r(n.combine(i,a));return n.polygon(o)}n={buildLog:function(t){return!0===t?u=i():!1===t&&(u=!1),!1!==u&&u.list},epsilon:function(t){return f.epsilon(t)},segments:function(t){var e=o(!0,f,u);return t.regions.forEach(e.addRegion),{segments:e.calculate(t.inverted),inverted:t.inverted}},combine:function(t,e){return{combined:o(!1,f,u).calculate(t.segments,t.inverted,e.segments,e.inverted),inverted1:t.inverted,inverted2:e.inverted}},selectUnion:function(t){return{segments:l.union(t.combined,u),inverted:t.inverted1||t.inverted2}},selectIntersect:function(t){return{segments:l.intersect(t.combined,u),inverted:t.inverted1&&t.inverted2}},selectDifference:function(t){return{segments:l.difference(t.combined,u),inverted:t.inverted1&&!t.inverted2}},selectDifferenceRev:function(t){return{segments:l.differenceRev(t.combined,u),inverted:!t.inverted1&&t.inverted2}},selectXor:function(t){return{segments:l.xor(t.combined,u),inverted:t.inverted1!==t.inverted2}},polygon:function(t){return{regions:s(t.segments,f,u),inverted:t.inverted}},polygonFromGeoJSON:function(t){return c.toPolygon(n,t)},polygonToGeoJSON:function(t){return c.fromPolygon(n,f,t)},union:function(t,e){return h(t,e,n.selectUnion)},intersect:function(t,e){return h(t,e,n.selectIntersect)},difference:function(t,e){return h(t,e,n.selectDifference)},differenceRev:function(t,e){return h(t,e,n.selectDifferenceRev)},xor:function(t,e){return h(t,e,n.selectXor)}},\"object\"==typeof window&&(window.PolyBool=n),e.exports=n},{\"./lib/build-log\":501,\"./lib/epsilon\":502,\"./lib/geojson\":503,\"./lib/intersecter\":504,\"./lib/segment-chainer\":506,\"./lib/segment-selector\":507}],501:[function(t,e,r){e.exports=function(){var t,e=0,r=!1;function n(e,r){return t.list.push({type:e,data:r?JSON.parse(JSON.stringify(r)):void 0}),t}return t={list:[],segmentId:function(){return e++},checkIntersection:function(t,e){return n(\"check\",{seg1:t,seg2:e})},segmentChop:function(t,e){return n(\"div_seg\",{seg:t,pt:e}),n(\"chop\",{seg:t,pt:e})},statusRemove:function(t){return n(\"pop_seg\",{seg:t})},segmentUpdate:function(t){return n(\"seg_update\",{seg:t})},segmentNew:function(t,e){return n(\"new_seg\",{seg:t,primary:e})},segmentRemove:function(t){return n(\"rem_seg\",{seg:t})},tempStatus:function(t,e,r){return n(\"temp_status\",{seg:t,above:e,below:r})},rewind:function(t){return n(\"rewind\",{seg:t})},status:function(t,e,r){return n(\"status\",{seg:t,above:e,below:r})},vert:function(e){return e===r?t:(r=e,n(\"vert\",{x:e}))},log:function(t){return\"string\"!=typeof t&&(t=JSON.stringify(t,!1,\" \")),n(\"log\",{txt:t})},reset:function(){return n(\"reset\")},selected:function(t){return n(\"selected\",{segs:t})},chainStart:function(t){return n(\"chain_start\",{seg:t})},chainRemoveHead:function(t,e){return n(\"chain_rem_head\",{index:t,pt:e})},chainRemoveTail:function(t,e){return n(\"chain_rem_tail\",{index:t,pt:e})},chainNew:function(t,e){return n(\"chain_new\",{pt1:t,pt2:e})},chainMatch:function(t){return n(\"chain_match\",{index:t})},chainClose:function(t){return n(\"chain_close\",{index:t})},chainAddHead:function(t,e){return n(\"chain_add_head\",{index:t,pt:e})},chainAddTail:function(t,e){return n(\"chain_add_tail\",{index:t,pt:e})},chainConnect:function(t,e){return n(\"chain_con\",{index1:t,index2:e})},chainReverse:function(t){return n(\"chain_rev\",{index:t})},chainJoin:function(t,e){return n(\"chain_join\",{index1:t,index2:e})},done:function(){return n(\"done\")}}}},{}],502:[function(t,e,r){e.exports=function(t){\"number\"!=typeof t&&(t=1e-10);var e={epsilon:function(e){return\"number\"==typeof e&&(t=e),t},pointAboveOrOnLine:function(e,r,n){var i=r[0],a=r[1],o=n[0],s=n[1],l=e[0];return(o-i)*(e[1]-a)-(s-a)*(l-i)>=-t},pointBetween:function(e,r,n){var i=e[1]-r[1],a=n[0]-r[0],o=e[0]-r[0],s=n[1]-r[1],l=o*a+i*s;return!(l<t)&&!(l-(a*a+s*s)>-t)},pointsSameX:function(e,r){return Math.abs(e[0]-r[0])<t},pointsSameY:function(e,r){return Math.abs(e[1]-r[1])<t},pointsSame:function(t,r){return e.pointsSameX(t,r)&&e.pointsSameY(t,r)},pointsCompare:function(t,r){return e.pointsSameX(t,r)?e.pointsSameY(t,r)?0:t[1]<r[1]?-1:1:t[0]<r[0]?-1:1},pointsCollinear:function(e,r,n){var i=e[0]-r[0],a=e[1]-r[1],o=r[0]-n[0],s=r[1]-n[1];return Math.abs(i*s-o*a)<t},linesIntersect:function(e,r,n,i){var a=r[0]-e[0],o=r[1]-e[1],s=i[0]-n[0],l=i[1]-n[1],c=a*l-o*s;if(Math.abs(c)<t)return!1;var u=e[0]-n[0],f=e[1]-n[1],h=(s*f-l*u)/c,p=(a*f-o*u)/c,d={alongA:0,alongB:0,pt:[e[0]+h*a,e[1]+h*o]};return d.alongA=h<=-t?-2:h<t?-1:h-1<=-t?0:h-1<t?1:2,d.alongB=p<=-t?-2:p<t?-1:p-1<=-t?0:p-1<t?1:2,d},pointInsideRegion:function(e,r){for(var n=e[0],i=e[1],a=r[r.length-1][0],o=r[r.length-1][1],s=!1,l=0;l<r.length;l++){var c=r[l][0],u=r[l][1];u-i>t!=o-i>t&&(a-c)*(i-u)/(o-u)+c-n>t&&(s=!s),a=c,o=u}return s}};rLine truncated
|
||||
"/*!\n",
|
||||
" * repeat-string <https://github.com/jonschlinkert/repeat-string>\n",
|
||||
" *\n",
|
||||
" * Copyright (c) 2014-2015, Jon Schlinkert.\n",
|
||||
" * Licensed under the MIT License.\n",
|
||||
" */\n",
|
||||
"\"use strict\";var n,i=\"\";e.exports=function(t,e){if(\"string\"!=typeof t)throw new TypeError(\"expected a string\");if(1===e)return t;if(2===e)return t+t;var r=t.length*e;if(n!==t||void 0===n)n=t,i=\"\";else if(i.length>=r)return i.substr(0,r);for(;r>i.length&&e>1;)1&e&&(i+=t),e>>=1,t+=t;return i=(i+=t).substr(0,r)}},{}],538:[function(t,e,r){(function(t){(function(){e.exports=t.performance&&t.performance.now?function(){return performance.now()}:Date.now||function(){return+new Date}}).call(this)}).call(this,\"undefined\"!=typeof global?global:\"undefined\"!=typeof self?self:\"undefined\"!=typeof window?window:{})},{}],539:[function(t,e,r){\"use strict\";e.exports=function(t){for(var e=t.length,r=t[t.length-1],n=e,i=e-2;i>=0;--i){var a=r,o=t[i];(l=o-((r=a+o)-a))&&(t[--n]=r,r=l)}var s=0;for(i=n;i<e;++i){var l;a=t[i];(l=(o=r)-((r=a+o)-a))&&(t[s++]=l)}return t[s++]=r,t.length=s,t}},{}],540:[function(t,e,r){\"use strict\";var n=t(\"two-product\"),i=t(\"robust-sum\"),a=t(\"robust-scale\"),o=t(\"robust-compress\");function s(t,e){for(var r=new Array(t.length-1),n=1;n<t.length;++n)for(var i=r[n-1]=new Array(t.length-1),a=0,o=0;a<t.length;++a)a!==e&&(i[o++]=t[n][a]);return r}function l(t){for(var e=new Array(t),r=0;r<t;++r){e[r]=new Array(t);for(var n=0;n<t;++n)e[r][n]=[\"m[\",r,\"][\",n,\"]\"].join(\"\")}return e}function c(t){if(2===t.length)return[\"sum(prod(\",t[0][0],\",\",t[1][1],\"),prod(-\",t[0][1],\",\",t[1][0],\"))\"].join(\"\");for(var e=[],r=0;r<t.length;++r)e.push([\"scale(\",c(s(t,r)),\",\",(n=r,1&n?\"-\":\"\"),t[0][r],\")\"].join(\"\"));return function t(e){if(1===e.length)return e[0];if(2===e.length)return[\"sum(\",e[0],\",\",e[1],\")\"].join(\"\");var r=e.length>>1;return[\"sum(\",t(e.slice(0,r)),\",\",t(e.slice(r)),\")\"].join(\"\")}(e);var n}function u(t){return new Function(\"sum\",\"scale\",\"prod\",\"compress\",[\"function robustDeterminant\",t,\"(m){return compress(\",c(l(t)),\")};return robustDeterminant\",t].join(\"\"))(i,a,n,o)}var f=[function(){return[0]},function(t){return[t[0][0]]}];!function(){for(;f.length<6;)f.push(u(f.length));for(var t=[],r=[\"function robustDeterminant(m){switch(m.length){\"],n=0;n<6;++n)t.push(\"det\"+n),r.push(\"case \",n,\":return det\",n,\"(m);\");r.push(\"}var det=CACHE[m.length];if(!det)det=CACHE[m.length]=gen(m.length);return det(m);}return robustDeterminant\"),t.push(\"CACHE\",\"gen\",r.join(\"\"));var i=Function.apply(void 0,t);for(e.exports=i.apply(void 0,f.concat([f,u])),n=0;n<f.length;++n)e.exports[n]=f[n]}()},{\"robust-compress\":539,\"robust-scale\":546,\"robust-sum\":549,\"two-product\":600}],541:[function(t,e,r){\"use strict\";var n=t(\"two-product\"),i=t(\"robust-sum\");e.exports=function(t,e){for(var r=n(t[0],e[0]),a=1;a<t.length;++a)r=i(r,n(t[a],e[a]));return r}},{\"robust-sum\":549,\"two-product\":600}],542:[function(t,e,r){\"use strict\";var n=t(\"two-product\"),i=t(\"robust-sum\"),a=t(\"robust-subtract\"),o=t(\"robust-scale\");function s(t,e){for(var r=new Array(t.length-1),n=1;n<t.length;++n)for(var i=r[n-1]=new Array(t.length-1),a=0,o=0;a<t.length;++a)a!==e&&(i[o++]=t[n][a]);return r}function l(t){if(1===t.length)return t[0];if(2===t.length)return[\"sum(\",t[0],\",\",t[1],\")\"].join(\"\");var e=t.length>>1;return[\"sum(\",l(t.slice(0,e)),\",\",l(t.slice(e)),\")\"].join(\"\")}function c(t,e){if(\"m\"===t.charAt(0)){if(\"w\"===e.charAt(0)){var r=t.split(\"[\");return[\"w\",e.substr(1),\"m\",r[0].substr(1)].join(\"\")}return[\"prod(\",t,\",\",e,\")\"].join(\"\")}return c(e,t)}function u(t){if(2===t.length)return[[\"diff(\",c(t[0][0],t[1][1]),\",\",c(t[1][0],t[0][1]),\")\"].join(\"\")];for(var e=[],r=0;r<t.length;++r)e.push([\"scale(\",l(u(s(t,r))),\",\",(n=r,!0&n?\"-\":\"\"),t[0][r],\")\"].join(\"\"));return e;var n}function f(t,e){for(var r=[],n=0;n<e-2;++n)r.push([\"prod(m\",t,\"[\",n,\"],m\",t,\"[\",n,\"])\"].join(\"\"));return l(r)}function h(t){for(var e=[],r=[],c=function(t){for(var e=new Array(t),r=0;r<t;++r){e[r]=new Array(t);for(var n=0;n<t;++n)e[r][n]=[\"m\",n,\"[\",t-r-2,\"]\"].join(\"\")}return e}(t),h=0;h<t;++h)c[0][h]=\"1\",c[t-1][h]=\"w\"+h;for(h=0;h<t;++h)0==(1&h)?e.push.apply(e,u(s(c,h))):r.push.apply(r,u(s(c,h)));var p=l(e),d=l(r),m=\"exactInSphere\"+t,g=[];for(h=0;h<t;++h)g.push(\"m\"+h);var v=[\"function \",m,\"(\",g.join(),\"){\"];for(h=0;h<t;++h){v.push(\"var w\",h,\"=\",f(h,t),\";\");for(var y=0;y<t;++y)y!==h&&v.push(\"var w\",h,\"m\",y,\"=scale(w\",h,\",m\",y,\"[0]);\")}return v.push(\"var p=\",p,\",n=\",d,\",d=diff(p,n);return d[d.length-1];}return \",m),new Function(\"sum\",\"diff\",\"prod\",\"scale\",v.join(\"\"))(i,a,n,o)}var p=[function(){return 0},function(){return 0},function(){return 0}];function d(t){var e=p[t.length];return e||(e=p[t.length]=h(t.length)),e.apply(void 0,t)}!function(){for(;p.length<=6;)p.push(h(p.length));for(var t=[],r=[\"slow\"],n=0;n<=6;++n)t.push(\"a\"+n),r.push(\"o\"+n);var i=[\"function testInSphere(\",t.join(),\"){switch(arguments.length){case 0:case 1:return 0;\"];for(n=2;n<=Line truncated
|
||||
"</body>\n",
|
||||
"</html>"
|
||||
],
|
||||
"text/plain": [
|
||||
"<IPython.core.display.HTML object>"
|
||||
]
|
||||
},
|
||||
"execution_count": 16,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"HTML(fig.to_html())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 33,
|
||||
"id": "94f1e3ab-903a-48c9-be81-f5bfdcf130ea",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"[-105.54, 40.04]"
|
||||
]
|
||||
},
|
||||
"execution_count": 33,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"stn_lonlat[52]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 114,
|
||||
"id": "0d589ec3-4da7-480b-a848-ec11be69eb8d",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAY0AAAEtCAYAAAD0uzw/AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMSwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Z1A+gAAAACXBIWXMAAAsTAAALEwEAmpwYAABm8klEQVR4nO2dd3wT9f/HX2mapHvSQgejrDLK3kVEKfCFKkPZCIgsURygsh2IKCJDpsCPLYhsKENBlrJbBGVvCpTRUrp32iS/P2rS3OUuucto0vb9fDx40Nx97u6TXPJ533tLNBqNBgRBEAQhACd7T4AgCIIoO5DQIAiCIARDQoMgCIIQDAkNgiAIQjAkNAiCIAjBkNAgCIIgBENCgyDMoFOnThg6dKigsbt27UJ4eDhiY2NtPCuCsD0kNAiT5OXlYf369Rg8eDBat26Nhg0bIjIyEqNHj8auXbtQVFRkcMz58+fx0Ucf4aWXXkJERATatWuH0aNH48iRI2bPY8mSJQgPD2f8a968OaKjo/Hjjz8iPT3dgndZvrl//z7mzJmDYcOGoWXLlggPD8eSJUsEHZuXl4eoqCiEh4dj5syZoq89YcIEhIeH4+HDhwb7vvzyS4SHh+Ozzz4z2JecnIzw8HCMHTtW9DUJ2+Fs7wkQjs3Dhw8xZswYPHjwAJGRkRgzZgx8fX2RkpKCs2fPYurUqbh79y4mTZqkO2bBggVYuXIlQkJC0LdvX4SGhuLFixfYv38/xo0bh169emH27NmQSqVmzemjjz5CaGgoACArKwuxsbFYsWIF/vrrL+zatQtOTvQsxObff//FunXrUK1aNTRs2BDnzp0TfOzixYuRmppq9rXbtGmD3377DXFxcahevTpjX2xsLJydnXH+/HmD47SaWZs2bcy+NmF9SGgQvOTn5+Pdd9/F48ePsWTJEnTt2pWxf8yYMbh8+TKuXLmi27Z9+3asXLkSkZGR+Omnn+Dq6qrbN2rUKEyfPh179uxBSEgIPv74Y7Pm9fLLL6NRo0a610OGDMEHH3yAw4cP4+bNm2jQoIFZ5y1LqFQqKJVKxudrjE6dOiEuLg5eXl64cuUK+vbtK+i4a9euYcOGDZg4cSK+//57s+aqXfRjY2PRr18/3fbnz5/jwYMHePPNN7Fr1y48fPiQIVTi4uIAAK1btzbruoRtoEcygpft27cjPj4e77zzjoHA0NK4cWO89dZbAAClUolFixbBzc0N8+bNM1jQnJ2dMXPmTAQHB2Pt2rUWPb2yCQwMBADIZDLG9tTUVHz99dfo2LEjIiIi0LFjR3z99ddIS0tjjNOavh4/fmxwbjH+i23btqFbt26IiIhAly5dsH79evBV6snKysLcuXPRpUsXREREoG3btvjkk0+QkJDAGKf1iZw5cwbLli1D586d0bhxY/z++++C5gQAPj4+8PLyEjweKBZMX3zxBTp06IAuXbqIOlafsLAwBAYG6oSAFu3rsWPHwtnZ2cDnExsbCy8vL9SvX9/saxPWhzQNgpdDhw4BAAYMGCBo/MWLF5GcnIwePXrA39+fc4xCoUDPnj115qQ33nhD9Lyys7N1Aic7OxtxcXHYtWsXWrRogdq1a+vGZWVlYdCgQXj48CH69OmDBg0a4MaNG/j1119x7tw5bN++HR4eHqKvz8f69esxe/Zs1KtXD5988gny8vKwdu1azs8iKysLAwcOxNOnT9GnTx/UqVMHycnJ2Lx5M/r164edO3ciJCSEccycOXNQVFSE/v37w93dHWFhYVabO9/7uX//PhYvXmzxudq0aYN9+/bhwYMHqFGjBoBioREWFobq1aujQYMGiI2NRf/+/QGUaCFRUVFkbnQwSGgQvNy5cwceHh6oWrWq4PEA0LBhQ6PjtPtv375t1ryGDx9usC0qKgpz586FRCLRbVu9ejUePHiAL7/8UqcNAUD9+vUxc+ZMrF69GuPHjzdrDmwyMzOxcOFC1KpVC1u2bNFpWX369EH37t0Nxi9atAgJCQnYtm0b6tWrp9v+xhtvoEePHliyZImBOSg/Px979uwRbJKyhISEBCxZsgTvv/8+QkNDOTUwMWiFRmxsrE5oxMbG6kxXrVu3xt69e3XjyTTluJAIJ3jJzs6Gu7u7qPEATD69a/dnZWWZNa8vv/wS69atw7p167B48WIMHz4cJ06cwEcffQSlUqkbd/jwYfj5+RloSgMGDICfn59FkVxsTp06hby8PLz11luMRb1KlSro0aMHY6xGo8G+ffvQqlUrBAYGIjU1VffP1dUVTZs2xalTpwyuMWjQoFIRGAAwY8YMVK1aFe+8845Vzte2bVsAJcJAq0lohULr1q3x/PlzxMfHM8ZpjyMcB9I0CF48PDyQk5MjajxQIjz40O739PQ0a16NGzdmOML/97//wd/fH/Pnz8fOnTsxaNAgAMDjx48REREBZ2fm19zZ2Rk1atTA9evXzbo+F9on8Zo1axrsq1WrFuN1amoq0tPTcerUKbRr147zfFwmGVubo7TExMTg9OnT2LRpk4GPyFyqVq2K4OBgnd9CKxRatWoFAGjRogWkUiliY2MRFhaG2NhY+Pj4IDw83CrXJ6wHCQ2Clzp16uD8+fNISEgQZKKqU6cOgOKIG2No99etW9fySf5Hhw4dMH/+fJw7d04nNMSgb9Ziw5WHYglax7g210UoLi4uVp0HF0qlEt9//z06duyIgIAAXW5FUlISgGLt8OHDh/D19RXtWG/Tpg12796N+/fv68JvK1euDKD4gaNevXqIi4tDp06d8ODBA3Tp0sXofSHsAwkNgpeuXbvi/Pnz2L59Oz755BOT45s3b45KlSrh6NGjSE1NhZ+fn8GYgoIC7Nu3DwqFAi+//LLV5lpYWAgADM2oatWqiI+PR1FREUPbKCoqwoMHDxiC0NvbGwCQkZGhywHRzjc5Odkgv4CN9pj79+8baA/37t1jvPbz84OXlxeys7MRGRkp5m3anPz8fKSmpuLPP//En3/+abB/79692Lt3LyZNmoSRI0eKOrdWaMTGxiI2NlanZWhp3bo19u/fr9NCKD/DMSGfBsFLv379EBYWhrVr1/La/69evYpffvkFACCXy/HRRx8hNzcXEydORH5+PmOsSqXCjBkz8OTJE4wcOZI3wsocjh49CoDphO/cuTNSU1Oxfft2xtht27YhNTUVnTt31m3TOmfPnDnDGLt+/Xqo1WqT12/fvj1cXFzwyy+/IC8vT7c9MTER+/btY4x1cnJCjx49cPnyZRw8eJDzfCkpKSavaQtcXV2xaNEig39fffUVgGKNbtGiRejUqZPoc2uFwIEDB/DgwQMDodGqVSskJydjy5YtAMgJ7qiQpkHw4urqipUrV2LMmDEYN24cXnrpJURGRsLHxwepqamIjY3FqVOnMGrUKN0xAwYMwMOHD7FmzRpER0ejd+/eCAkJ0WWE3759Gz179sQHH3xg9rxOnDiB+/fvAyj2j1y8eBEHDhxAlSpVMGzYMN24UaNG4eDBg5g5cyauX7+O+vXr48aNG9ixYwfCwsIY846MjERYWBgWL16M9PR0hIaG4sKFC7h06RJ8fX1Nzsnb2xsff/wx5syZg4EDB6J3797Iy8vDli1bOP0nEyZMwMWLFzF+/Hh0794dTZo0gUwmw9OnT3HixAk0bNjQ7GQ6LrKysrBx40YAxU5ooLjUy08//QSgOBelXr16kMlk6Natm8HxWp9NtWrVOPcLITg4GFWrVtVlf7OFQsuWLeHk5ITz58/D19fXquZLwnqQ0CCMUr16dezZswdbt27FoUOHsGLFCuTm5sLb2xsRERH4/vvvDaKDJk2ahI4dO2LTpk3Ytm0b0tPT4eHhgYiICHz00UcWJYoBYOQNODs7o3LlyhgwYADGjRvH0F48PT3x66+/YvHixTh27Bh27doFf39/DBw4EB9++CEjyksqlWL58uWYNWuWzgHcvn17bNq0SbCPZMSIEXBzc8O6deswf/58BAUFYcSIEfD09MS0adMYY7VzW7t2LQ4ePIijR49CKpWiSpUqaNGiBSNz2hpkZGRg0aJFjG1aMxFQHOWlH/prK9q0aYOEhASEhIQgODiYsc/b2xt169bFzZs30bp1a/JnOCgSDV+6KkEQBEGwIJ8GQRAEIRgyTxF2JScnB7m5uUbHSKVSzkisik5WVpZBsAEbmUwGHx+fcnl9wj6Q0CDsytq1a7F06VKjY0JCQnDs2LFSmlHZ4dtvv8Xu3buNjmndurXOAV7erk/YB/JpEHYlISHBoKorG4VCgRYtWpTSjMoOd+/e1UVC8eHl5YWIiIhyeX3CPpR5oVFUVITExERUqVLFoFwEQRAEYV3KvCM8MTERUVFRSExMtPdUCIIgyj1lXmgQBEEQpQcJDYIgCEIwJDQIgiAIwZDQIAiCIARDQoMgCIIQDAkNgiAIQjAkNAiCIAjBkNCwAncePseH327DL/vP23sqBEEQNoVSqK3AsCkbkJCYht1H/kXrRtVRp3qgvadEEARhE0jTsAIJiWm6v09fvGdkJEEQRNmGhAZBEAQhGBIaVoZaVBIEUZ4hoUEQBEEIhoSGlSFFgyCI8gwJDStD5imCIMozJDQIgiAIwZDQIAiCIARDQsPKkHGKIIjyDAkNa0M+DYIgyjEkNAiCIAjBkNCwMhQ9RRBEeYaEBkEQBCEYEhoWol+sECCXBkEQ5RsSGhbywaytjNcajZ0mQhAEUQqQ0LCQC9ceMV6r1Wo7zYQgCML2kNCwMioVCQ2CIMovJDSsjEpN9imCIMovJDSsDGkaBEGUZ0hoWBkV+TQIgijHkNCwMmSeIgiiPOMsdGB8fDzi4uJw584dpKamQiKRwNfXF3Xr1kWrVq0QFhZmy3mWGSh6iiCI8oxRoVFQUICdO3di69atuH37NjQ8SQgSiQR169bFwIED8eabb0KhUJg1mVWrVmHevHmoV68eYmJizDqHvSGfBkEQ5RleobFnzx4sXLgQSUlLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"ntrain = 500\n",
|
||||
"start_idx = 18000\n",
|
||||
"stn_idx = 52\n",
|
||||
"ntest = 100\n",
|
||||
"wind = stn_dat[stn_idx][start_idx:start_idx + ntrain]\n",
|
||||
"test_wind = stn_dat[stn_idx][start_idx + ntrain:start_idx+ntrain+ntest]\n",
|
||||
"x = np.arange(ntrain) * 5.\n",
|
||||
"test_x = np.arange(ntest)*5. + x[-1] + x[1]\n",
|
||||
"plt.plot(x, wind)\n",
|
||||
"plt.plot(test_x, test_wind)\n",
|
||||
"plt.ylabel(\"Wind Speed (m/s)\")\n",
|
||||
"plt.xlabel(\"Minutes\")\n",
|
||||
"plt.title(stn_names[stn_idx])\n",
|
||||
"sns.despine()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 115,
|
||||
"id": "9db49bb2-502d-407f-a69d-7d4453a6ed7c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import torch\n",
|
||||
"import pandas as pd\n",
|
||||
"import gpytorch\n",
|
||||
"import argparse\n",
|
||||
"import datetime\n",
|
||||
"import warnings\n",
|
||||
"import copy\n",
|
||||
"import os\n",
|
||||
"from voltron.data import make_ticker_list, GetStockHistory\n",
|
||||
"import sys\n",
|
||||
"sys.path.append(\"../calibration\")\n",
|
||||
"from LSTMUtils import SequenceDataset, LSTM, TrainLSTM, LSTMRollouts, NLL\n",
|
||||
"from torch.utils.data import DataLoader\n",
|
||||
"from voltron.train_utils import LearnGPCV, TrainVolModel, TrainVoltMagpieModel, TrainBasicModel\n",
|
||||
"from voltron.rollout_utils import Rollouts\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 116,
|
||||
"id": "70b2bd65-6637-49e6-b3c6-47c059bac888",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"\n",
|
||||
"train_x = torch.arange(ntrain-1).float()/365\n",
|
||||
"train_y = torch.FloatTensor(wind[:ntrain]) + 1\n",
|
||||
"test_x = torch.arange(ntrain, ntrain+ntest).float()/365\n",
|
||||
"test_y = torch.FloatTensor(test_wind) + 1\n",
|
||||
"\n",
|
||||
"if torch.cuda.is_available():\n",
|
||||
" use_cuda = True,\n",
|
||||
" train_x, train_y = train_x.cuda(), train_y.cuda()\n",
|
||||
" test_x, test_y = test_x.cuda(), test_y.cuda()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 117,
|
||||
"id": "dec612fd-edf5-45f1-b510-de27089249aa",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/home/greg_b/miniconda3/envs/rpp/lib/python3.8/site-packages/gpytorch/utils/cholesky.py:40: NumericalWarning:\n",
|
||||
"\n",
|
||||
"A not p.d., added jitter of 1.0e-06 to the diagonal\n",
|
||||
"\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Iter 1/200 - Loss: 18.103\n",
|
||||
"Iter 51/200 - Loss: 2.369\n",
|
||||
"Iter 101/200 - Loss: 2.264\n",
|
||||
"Iter 151/200 - Loss: 2.250\n",
|
||||
"Iter 1/1000 - Loss: 0.792\n",
|
||||
"Iter 51/1000 - Loss: 0.602\n",
|
||||
"Iter 101/1000 - Loss: 0.391\n",
|
||||
"Iter 151/1000 - Loss: 0.165\n",
|
||||
"Iter 201/1000 - Loss: -0.068\n",
|
||||
"Iter 251/1000 - Loss: -0.299\n",
|
||||
"Iter 301/1000 - Loss: -0.523\n",
|
||||
"Iter 351/1000 - Loss: -0.735\n",
|
||||
"Iter 401/1000 - Loss: -0.930\n",
|
||||
"Iter 451/1000 - Loss: -1.104\n",
|
||||
"Iter 501/1000 - Loss: -1.251\n",
|
||||
"Iter 551/1000 - Loss: -1.370\n",
|
||||
"Iter 601/1000 - Loss: -1.460\n",
|
||||
"Iter 651/1000 - Loss: -1.522\n",
|
||||
"Iter 701/1000 - Loss: -1.562\n",
|
||||
"Iter 751/1000 - Loss: -1.585\n",
|
||||
"Iter 801/1000 - Loss: -1.597\n",
|
||||
"Iter 851/1000 - Loss: -1.602\n",
|
||||
"Iter 901/1000 - Loss: -1.605\n",
|
||||
"Iter 951/1000 - Loss: -1.606\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"with gpytorch.settings.max_cholesky_size(2000):\n",
|
||||
" vol = LearnGPCV(train_x, train_y, train_iters=200,\n",
|
||||
" printing=True)\n",
|
||||
" vmod, vlh = TrainVolModel(train_x, vol, \n",
|
||||
" train_iters=1000, printing=True)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 118,
|
||||
"id": "36e4205b-4a10-4555-9903-45f69ee5c965",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Iter 1/100 - Loss: 0.829\n",
|
||||
"Iter 51/100 - Loss: -0.682\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"voltron, lh = TrainVoltMagpieModel(train_x, train_y[1:], \n",
|
||||
" vmod, vlh, vol,\n",
|
||||
" printing=True, \n",
|
||||
" train_iters=100,\n",
|
||||
" k=200, mean_func=\"constant\", theta=0.75)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 119,
|
||||
"id": "7b592f9c-6ac4-4591-af30-1e3484129cc5",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"vmod.eval();\n",
|
||||
"voltron.eval();\n",
|
||||
"voltron.vol_model.eval();\n",
|
||||
"thetas = [0., 0.01, 0.025]\n",
|
||||
"full_samples = torch.zeros(len(thetas), 10, ntest)\n",
|
||||
"for theta_idx, theta in enumerate(thetas):\n",
|
||||
" temp_model = copy.deepcopy(voltron)\n",
|
||||
" \n",
|
||||
" with torch.no_grad():\n",
|
||||
" save_samples, pred_vol = Rollouts(train_x, train_y, test_x, temp_model, \n",
|
||||
" nsample=10, theta=theta)\n",
|
||||
" full_samples[theta_idx, ...] = save_samples\n",
|
||||
" torch.cuda.empty_cache()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 120,
|
||||
"id": "874c24aa-e8b4-44ac-be7c-58e8cec7a468",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# pred_vol = vmod(test_x).sample(torch.Size((100, ))).exp()\n",
|
||||
"# predictions = save_samples.exp().detach()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 125,
|
||||
"id": "a29609c9-e768-4bf6-aa34-d7fdd1cbf88b",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAfgAAAFoCAYAAAC7Tuk8AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMSwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Z1A+gAAAACXBIWXMAAAsTAAALEwEAmpwYAADPeElEQVR4nOydd3QU1dvHvzNb0wtJ6L2ELqE3EQgoIChSpKiIIIIiNkQQ9SeW165YUFEURBREFOlFmkqT3pEOoZf0sn1m3j9mZ3ZmdrYlu9kk3M85HLJT72y5z306xXEcBwKBQCAQCBUKOtwDIBAIBAKBEHyIgCcQCAQCoQJCBDyBQCAQCBUQIuAJBAKBQKiAEAFPIBAIBEIFhAh4AoFAIBAqIETAEwgVmKVLlyI1NRWPPPJIsc5PTU1FamoqLl++HOSREQiEUKMN9wAIhPLAtGnT8Mcff7htp2kaMTExqF+/Pu6++26MGDECRqMxDCMkeMLhcGDVqlXYsmULjhw5gpycHDAMg4SEBDRp0gTdunVD//79ERsb6/Ea169fxy+//IJt27bh8uXLKCwsRHx8PBo0aIAePXpg6NChiIyMLMWnIhB8QwQ8gRAAOp0OcXFx4mur1Yq8vDzs378f+/fvx2+//YYFCxYgMTExjKMkCBw8eBBTpkzBxYsXxW1GoxF6vR7Xr1/H9evXsWXLFsycORMzZszAvffe63aNuXPn4rPPPoPFYgEAaDQaREdHIzMzE7du3cLOnTsxZ84cfPDBB+jcuXOpPRuB4AtioicQAiAtLQ3bt28X/+3duxd79+7F1KlTQdM0zpw5g48++ijcwyQA2LZtG0aNGoWLFy+icuXKeP311/HPP//g0KFD2Lt3Lw4ePIhvv/0WvXv3Rn5+PrZu3ep2jU8++QTvv/8+LBYLunTpgp9++glHjhzB7t27cfDgQXz++eeoU6cObt26hSeeeAJbtmwJw5MSCOoQAU8glJCYmBiMGTMGQ4YMAQAyyZcBbty4gcmTJ8NqtaJ58+ZYvnw5Ro4cicqVK4vHRERE4K677sKsWbMwf/58JCUlya7x119/4ZtvvgEAPPTQQ5g7dy7atWsHjUYDgLcE3HPPPfj999/RqlUr2O12TJ06FTdu3Ci9ByUQvEAEPIEQJFJTUwEAZrPZ4zGZmZl477330KdPH9xxxx1o06YNhgwZgrlz58Jms6me88gjjyA1NRVLly71eN2ePXsiNTUVu3btCmjMLMtiwYIFuO+++9CyZUt07NgREyZMwIEDB/w6Pzs7Gx9//DEGDBiAtLQ0tGrVCv3798fMmTORm5vrc6w3btzAjBkzkJ6ejubNm+P+++8PaPye+Pbbb5Gbm4vIyEh8/vnnSEhI8Hp8x44dMXnyZNk2wRLTuHFjTJ8+3eO50dHR+OSTTxAREYG8vDx8++23JX8AAiEIEB88gRAkTp06BQCoVauW6v7Dhw9j3LhxouCLioqC3W7HkSNHcOTIESxfvhxz585FpUqVSmW8DocDzzzzDDZt2gQA0Gq1YBgGW7ZswdatWzFz5kyv5+/duxcTJ04Un0en04GmaZw+fRqnT58Wn6devXqq51+4cAHPPvsscnJyEBERAZ1OF5Tnstls4mLo/vvvR/Xq1f06j6Io8e99+/bh9OnTAIBx48ZBq/U+VVavXh0DBgzAr7/+iqVLl2Lq1KnQ6/XFfAICITgQDZ5AKCGFhYX44YcfsGTJEgDA6NGj3Y7Jy8vDU089hdzcXDRq1AhLlizB/v37ceDAAXz22WeIi4vDiRMn8OKLL5bauOfMmYNNmzaBpmm89NJL2Lt3L/bs2YONGzeiU6dOXrXWK1eu4Mknn0Rubi5GjBiBP//8E4cPH8bBgwexcuVKdO3aFdeuXcOkSZPAMIzqNd577z0kJydj0aJFOHjwIA4cOIDPP/+8xM915MgRmEwmALy1oDjs3r0bAJ8l0aNHD7/O6dWrFwDAZDLh6NGjxbovgRBMiAZPIATAgQMH0KVLF/G11WpFQUEBAKBp06Z49NFHMXDgQLfzfvrpJ9y6dQuxsbGYO3cukpOTAfAR2X369EF0dDTGjh2LHTt2YOfOnejUqVNIn8NkMmHOnDkAgKeeegpjx44V99WsWRNfffUVHnjgAfHZlMycORP5+fl44okn3EzbjRo1wtdff40hQ4bg5MmT2LBhA/r06eN2Da1Wi3nz5sl837Vr1y7xs507d078u3HjxsW6xpkzZwDw1pioqCi/zhFcNABw9uxZtG7dulj3JhCCBdHgCYQAsNvtyMzMFP9JBWBeXh6ys7PBcZzbeevXrwcADBkyRBTuUrp27Yq0tDQAwNq1a0M0ehfbt29HUVER9Hq9qsVBr9djzJgxqueazWasW7cONE3jscceUz1Gr9fjnnvuAQDs2LFD9Zj777/fLbAtGOTk5Ih/x8fHF+saeXl5AZ8v9fN7ij8gEEoTosETCAHQvn17LFiwQHzNMAyuXr2KrVu34rPPPsP777+PM2fO4J133hGPsdlsoj+3Y8eOHq/doUMHHDhwAMePHw/dAzg5duwYAKBJkyaIiYlRPaZ9+/Yez7Xb7aAoCgMGDPB4DyFv/Nq1a6r7hQUNgUAIDUTAEwglQKPRoGbNmhg5ciRq1qyJxx9/HL///jsGDRqEtm3bAuC1QZZlAUCWpqWkSpUqAPjI9FAj3CMlJcXjMZ7GevPmTQAAx3HIzMz0eS9B0CsJVTEgpSbt7Rk9IRQzCkQTD4blgEAIJkTAEwhB4s4770RycjJu3bqFdevWiQJeitVqDcPIgovggoiJicHevXuLfR2aDo2HUBq1f+LEiWIJ+Pr16wMALl68iKKiIr/88CdPnnQ7n0AIJ8QHTyAEkapVqwIALl26JG6Li4sThdnVq1c9nnv9+nUA7pqtUFjF2+LAUzCcJ4R7CNq4Gp4KtghpfIWFhQHftzRo0aKFWBd+8+bNxbpGhw4dAPB1AvwtXLRx40YAQGRkJJo3b16s+xIIwYQIeAIhiAhCUZo3rdfr0bBhQwDwWohG2Ne0aVPZdqEJirAAUJKRkYH8/PyAxtmsWTMAwH///YfCwkLVY/bs2aO6vXnz5tBqteA4TrW8a7jR6/V44IEHAADLly/3uqiSIg2ObNOmjfiZfffdd3A4HF7PvXLlClauXAkAeOCBB0gOPKFMQAQ8gRAk9u3bJwp4pZAWIsr/+OMPVa1527ZtYvW4vn37yvY1atQIgGdttDiV07p06YLo6GjYbDbMnz/fbb/NZsO8efNUz42Ojsbdd98NAPjss888LhAAvphOUVFRwOMrKU888QTi4+NhMpnwzDPP+PSl//vvv/j4449l21544QUA/CJIGjSppKioCC+88ALMZjNiY2Mxbty4Eo+fQAgGRMATCCXEYrFg48aNYj54RESEWJde4OGHH0ZycjIsFgsef/xxHDlyBAAfhb9+/XpRmHTu3NktB/6ee+4BRVE4deoU3n77bVFbz8rKwttvv43ly5cjIiIioDFHRkbi8ccfBwB8+eWXmDdvnhgMd/nyZTz99NMeo98BYPLkyYiPj8eFCxcwYsQI/PPPP7Db7QB4TfjChQuYN28e+vbtW+yiL0Iv++L0o69SpQo+/PBD6PV6HDlyBPfffz8WLlwoW1yZzWb8/fffePrpp/Hoo4+6BQz27NlTrA/w888/Y+zYsdi7d68YMGmxWLB+/XoMHjwYBw8ehE6nw7vvviu6aQiEcEOC7AiEAFAWumEYRhY9HRkZiU8++cQtAj0uLg5fffUVHn/8cZw8eRJDhgxBVFQUHA6H6FtPTU1V7UTXsGFDPProo/jhhx+wYMECLFiwALGxsSgoKABN03j77bcxa9YsXLlyJaBnGTduHI4cOYJNmzbhvffew0cffYTIyEjk5+dDq9Vi5syZmDRpkuq5NWrUwJw5czBx4kScOnUK48aNg06nQ1RUFIqKikRhD8hLwJYm3bp1ww8//ICXXnoJly9fxhtvvIE33ngDERER0Gq1sviBhIQE1Yp1L730EhISEvDFF19g27Zt2LZtG7RaLaKiopCfny+a9ZOTk/Hee++ha9eupfZ8BIIviIAnEAJAKHQjJTIyEjVr1kSXLl3w8MMPe6x93rJlS6xevRrfffcd/vrrL1y7dg0ajQbNmzdHv3798PDDD8NgMKieO23aNNSuXRuLFy/G+fPnQVEUunbtivHjx6Ndu3aYNWtWwM+i1WrxxRdfYOHChViyZAkuXLgAmqbRvXt3jB8/3mcltpYtW2Lt2rVYtGgRNm3ahHPnzqGgoABRUVFITU1FWloa7rnnHrRr1y7gsQHArVu3APDpesWJhAd4X/q6deuwatUqbNmyBUePHkV2djasViuqVKmCJk2aoGfPnrj33ns9RsqPGzcO9957LxYvXoytW7fiypUrKCoqQqVKldCgQQP06NEDQ4cO9bviHYFQWlCcWtktAoFACDNjx47Ftm3b8Nprr+Hhhx8O93AIhHIHEfAEAqHMwTAM2rZti+joaGzatIlEpRMIxYAE2REIhDLHsWPHYDKZ8PjjjxPhTiAUE6LBEwgEAoFQASEaPIFAIBAIFZByL+AdDgcuX77ss9IUgUAgEAi3E+VewF+/fh3p6ekey3gSCAQCgXA7Uu4FPIFAIBAIBHeIgCcQCAQCoQJCBDyBQCAQCBUQIuAJBAKBQKiAEAFPIBAIBEIFhAh4AoFAIBAqIETAEwgEAoFQASECnkAgEAiECggR8AQCgUAgVECIgCcQCAQCoQJCBDyBQCAQyjwcx4E5+HG4h1GuIAKeQCAQCGUfxgLuwqpwj6JcQQQ8gUAgEMo+LN8xlOO4MA+k/EAEPIFAIBDKPoyV/59jwjuOcgQR8KWM2WLDuUuZ4R4GoYIzbdo0tG3bNtzDKHWmTZuGnj17hvw+PXv2xLRp08r8NSsUrI3/nwh4vyECvpS5e9wX6DbqExw6eTncQyGUM06cOIHnn38eXbt2RfPmzdGLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 576x360 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(1, 1, figsize=(8, 4))\n",
|
||||
"\n",
|
||||
"train_x = torch.arange(train_y.shape[0]-1) * 5.\n",
|
||||
"test_x = torch.arange(1, full_samples.shape[-1]+1) * 5. + train_x[-1]\n",
|
||||
"\n",
|
||||
"ax.plot(train_x.cpu(), train_y[1:].cpu(), lw=2., label=\"Observed Wind\")\n",
|
||||
"ax.plot(test_x.cpu(), full_samples[2, 5].T.exp().cpu(), alpha=0.75, lw=1., color=palette[7],\n",
|
||||
" label='Volt Samples')\n",
|
||||
"ax.plot(test_x.cpu(), full_samples[2, 6:12].T.exp().cpu(), alpha=0.75, lw=1., color=palette[7])\n",
|
||||
"ax.plot(test_x.cpu(), test_y.cpu(), color=palette[2], label=\"Truth\")\n",
|
||||
"ax.set_ylabel(\"Wind Speed (m/s)\")\n",
|
||||
"ax.set_xlabel(\"Minutes\")\n",
|
||||
"ax.set_title(\"Boulder, CO\", fontsize=24)\n",
|
||||
"ax.legend(frameon=False)\n",
|
||||
"sns.despine()\n",
|
||||
"plt.savefig(\"./wind_ex.pdf\", bbox_inches='tight')"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "9d5eb8ef-b222-4aaf-85ae-1a144a4456af",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.12"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,247 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "21eabbd3-cd2d-48dc-ae08-e4088abf5609",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"ename": "ModuleNotFoundError",
|
||||
"evalue": "No module named 'bs4'",
|
||||
"output_type": "error",
|
||||
"traceback": [
|
||||
"\u001b[0;31m---------------------------------------------------------------------------\u001b[0m",
|
||||
"\u001b[0;31mModuleNotFoundError\u001b[0m Traceback (most recent call last)",
|
||||
"\u001b[0;32m/var/folders/9y/fm22mhc16dn33r8h87_hhvtm0000gn/T/ipykernel_76886/708094613.py\u001b[0m in \u001b[0;36m<module>\u001b[0;34m\u001b[0m\n\u001b[1;32m 1\u001b[0m \u001b[0;32mimport\u001b[0m \u001b[0mnumpy\u001b[0m \u001b[0;32mas\u001b[0m \u001b[0mnp\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 2\u001b[0m \u001b[0;32mimport\u001b[0m \u001b[0mpandas\u001b[0m \u001b[0;32mas\u001b[0m \u001b[0mpd\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m----> 3\u001b[0;31m \u001b[0;32mfrom\u001b[0m \u001b[0mbs4\u001b[0m \u001b[0;32mimport\u001b[0m \u001b[0mBeautifulSoup\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m 4\u001b[0m \u001b[0;32mimport\u001b[0m \u001b[0mrequests\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 5\u001b[0m \u001b[0;32mimport\u001b[0m \u001b[0mmatplotlib\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mpyplot\u001b[0m \u001b[0;32mas\u001b[0m \u001b[0mplt\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;31mModuleNotFoundError\u001b[0m: No module named 'bs4'"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import pandas as pd\n",
|
||||
"from bs4 import BeautifulSoup\n",
|
||||
"import requests\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import torch"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "e3b55ef7-21f0-4265-a6cc-4ba4fb905024",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"base_url = 'https://www.ncei.noaa.gov/pub/data/uscrn/products/subhourly01/2021/'"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "acce1d74-81f9-490e-88d6-f351cc29f489",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"grab = requests.get(base_url)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "eb774d40-fb46-4c11-80b6-0fadb348a7bd",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"stn_names = {}\n",
|
||||
"stn_lonlat = {}\n",
|
||||
"stn_data = {}\n",
|
||||
"ndata = 105120\n",
|
||||
"stn_id = 0\n",
|
||||
"\n",
|
||||
"soup = BeautifulSoup(grab.text, 'html.parser')\n",
|
||||
"for link in soup.find_all('a'):\n",
|
||||
" url = link.get('href')\n",
|
||||
" if url[-4:] == '.txt':\n",
|
||||
" dat = pd.read_csv(base_url + url,\n",
|
||||
" header=None, delim_whitespace=True)\n",
|
||||
" if dat.shape[0] == ndata:\n",
|
||||
" stn_names[stn_id] = url[17:-4]\n",
|
||||
" stn_lonlat[stn_id] = [dat.iloc[0, 6], dat.iloc[0, 7]]\n",
|
||||
" stn_data[stn_id] = dat.iloc[:, 21].to_numpy()\n",
|
||||
" stn_id += 1"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "81464613-f48c-4341-9c67-7a35c1740b65",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"ename": "NameError",
|
||||
"evalue": "name 'stn_data' 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/var/folders/9y/fm22mhc16dn33r8h87_hhvtm0000gn/T/ipykernel_76886/111252756.py\u001b[0m in \u001b[0;36m<module>\u001b[0;34m\u001b[0m\n\u001b[0;32m----> 1\u001b[0;31m \u001b[0mfull_dat\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mnp\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mstack\u001b[0m\u001b[0;34m(\u001b[0m \u001b[0mlist\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mstn_data\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mvalues\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m)\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 2\u001b[0m \u001b[0mlonlat\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mnp\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0marray\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mlist\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mstn_lonlat\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mvalues\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;31mNameError\u001b[0m: name 'stn_data' is not defined"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"full_dat = np.stack( list(stn_data.values()))\n",
|
||||
"lonlat = np.array(list(stn_lonlat.values()))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 6,
|
||||
"id": "82066517-6625-4b6c-ad77-73ed32c9525c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"full_dat[full_dat == -99.0] = 0."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 7,
|
||||
"id": "29fd2ffd-57c5-4aed-b321-7a9dc3b0b5a4",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"<matplotlib.lines.Line2D at 0x7f765db6d250>"
|
||||
]
|
||||
},
|
||||
"execution_count": 7,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAwAAAAIFCAYAAABlO5qJAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMSwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Z1A+gAAAACXBIWXMAABcSAAAXEgFnn9JSAADRRklEQVR4nOzdd1iUV/YH8O+dGYZh6L0jiNIV7GIvsXdjj0lM78luskn2t8kmu0k2m01vm5411cTeezf2iiAgFkQR6b1Ovb8/AKUMM+87zFDP53l8iDP3zHs0qPe8773nMs45CCGEEEIIId2DpL0TIIQQQgghhLQdKgAIIYQQQgjpRqgAIIQQQgghpBuhAoAQQgghhJBuhAoAQgghhBBCuhEqAAghhBBCCOlGqAAghBBCCCGkG6ECgBBCCCGEkG6ECgBCCCGEEEK6ESoACCGEEEII6UaoACCEEEIIIaQboQKAEEIIIYSQbkTW3glYE2MsB4ASQGZ750IIIYQQQogFBQKo4pz7iA1knHMr5NMxMMbKbG1tHUNDQ9s7FUJIN6fnHFfzKm//PNTLHhLG2jEjQgghndnVq1ehUqnKOedOYmO79BMAAJmhoaFRycnJ7Z0HIaSbK63WIPafu27//PjrE+FsZ9OOGRFCCOnMoqOjkZKSYtYqF9oDQAghhBBCSDdCBQAhhBBCCCHdCBUAhBBCCCGEdCNUABBCCCGEENKNUAFACCGEEEJIN0IFACGEEEIIId0IFQCEEEIIIYR0I1QAEEIIIYQQ0o1QAUAIIYQQQkg3QgUAIYQQQggh3QgVAIQQQgghhHQjsvZOgBAh8qoqcK20GBq9Dh529gh39QBjrL3TIoQQQgjpdKgAIB0W5xxHbl3Hj8nnsPfGVeg5v/1eqIsb7o3sh/nhMbC3kbdjloQQQgghnQsVAKRD0nOON47tww/JZw2+f7WkCP84thcrLp7H8sl3w9/BqY0zJIQQQgjpnGgPAOmQ/nPyYIuT/4YuFRfg3m2rUVJT3QZZEUIIIYR0flQAkA4nrSgfXyeeEjw+vbQIX54/YcWMCCGEEEK6DioASIfzc0qC6JiVaUmo0WosnwwhhBBCSBdDBQDpUHR6PTZcSREdV6Kqwb4b6VbIiBBCCCGka6ECgHQopeoaVGjUZsVmVZRZOBtCCCGEkK6HCgBCCCGEEEK6ESoASIfiLFfA3sbGrFi/TtIKlHOOKq0aOr2+vVMhhBBCSDdE5wCQDkUqkWB2ryj8mnpeVJyLrQLjg3paKavW45zjWO51/HLlDPZlXYFarwMARLh4YUmvfpgdHAMHG9t2zpIQQggh3YFFngAwxsYwxriAH68ZiF3GGDvJGKtgjBUxxrYxxoZZIi/SOS2NjBMdMz8sBgqZeU8OrK1UXY2l+1Zg6f4V2JGZdnvyDwAXS/Lw2umdGLXpvziWm9F+SRJCCCGk27DUE4AcAD+28J4UwNK6//6j4RuMsY8BPAegGsAuAAoAEwBMZIzN45xvsFB+pBOJdPfCQzED8P2FM4LGhzi54sm4oVbOyjyVGjXu3f8bLhTlGB1Xoq7BsgMr8dPYRRji1aONsiOEEEJId2SRAoBzfhHAMkPvMcamoLYAyARwoMHrd6F28l8IIJ5zfrnu9fi6ccsZYwc45yWWyJF0Ln8bMgYavR4/pZwzOi7UxQ0/Tp4HV4VdG2UmzgeJB01O/utp9Do8c2QD/pj5FGyltDqPEEIIIdbRFpuA6+/+/8o55w1ef77u61v1k38A4JwfA/AVABcAD7VBfqQDkkok+Oew8fhx8jyMDwoFa/J+iJMr/j50LDbOuhcBjs7tkqMplRo11qSL28tQUFOJ7TcuWikjQgghhBArbwJmjNkDmFX3058bvG4HYFzdT9cYCF0D4FkAMwB8YM0cScfFGMPowBCMDgxBdkU5rpUVQ63TwcNOiSh3L0hY07KgY9lyIwUVWvFnGqy4ehazQ2KskBEhhBBCiPW7AM0FYA/gHOe84fGu4QBsAeRzzm8aiDtb97WvlfMjnYSvgyN8HRzbOw1R0kryzYq7VFJg4UwIIYQQQu6w9hKg+uU/Pzd5Pajuq6HJPzjnlQBKALgyxjrXrI+QOpoG3X7EUOu1Fs6EEEIIIeQOqz0BYIz5AhgPQAfgtyZvO9R9rTLyEZWo3QfgCKDcxLWSW3gr1GSihFiJm0JpXpyteXGEEEIIIUJY8wnAYtS2AN3NORfWBoWQLmRyQLhZcVMCIyycCSGEEELIHdbcA9DS8h8AqKj7auxWp33dV6N3/wGAcx5t6PW6JwNRpuIJsYZIV28M8AjAmQKDK91atKR3fytlRAghhBBipScAjLFIAP1QO9HfYGDIjbqvAS3E26N2+U8x59xkAUBIR/Vs9AhIRfwxWxgahxBHNytmRAghhJDuzlpPAO6t+7qOc25onX8aABUAT8aYP+c8q8n79bdAE62UHyFWc7EgH78mJmJz2kWUqVQA5ICUA3ZawE7XYtk9MSAM/xwwqU1zJYQQQkj3Y/ECgDHGACyp+6mh5T/gnFczxvYBmAJgPoCPmwyZV/d1s6XzI8RaNDodXt+/D79fSGr+po4BFTZAhQxwqisE6gTau+C+sIFYFjYQUklbnM1HCCGEkO7MGk8ARgLoASALwD4j4z5EbQHwKmNsa/1pwIyxeACPobYN6PdWyI8Qi9Nzjhd27sCWS2kmRjKgzAbj/Hqhf6Avol29MdK3Z4c/1IwQQgghXYc1bjfWb/5dwTnXtzSIc74HwCcA3AEkMMY2MMa2ATiE2sLkAc55iRXyI8Ti1qQkC5j833HoUiZmBERjtF8oTf4JIYQQ0qYsWgAwxmxxZ/nOL6bGc87/BOABAKkAJgCIB7AHwCjO+QZL5kZM0+r1KKipRF51BdQ68w6x6o445/gp4ZyoGK1ejxVJtMWFEEIIIW3PokuAOOcqAKJamHDOfwDwgyXzIOKklxXil8tnsTY9EeUaFQBALpFialAklvbuj34e/mB0l7pFibk5SMnPFx23OvkCXhg2HDJa908IIYSQNmTNcwBIB8c5x+cXjuDjpEPgTd5T63XYkHEBGzIuYG5IH7w9eCrkUmm75NnRXSwoMCuuqLoa+ZWV8HV0tHBGhBBCCCEto1uP3dinFw7jIwOT/6bWXUvC88c2Qs9NjeyeWrNcipZaEUIIIaStUQHQTV0oysEnSX8IHr/txkVszLhgxYw6L3c7YwdaG+dqZ2fBTAghhBBCTKMCoJv6+dKZNonpDkb06AGljY3ouOFBQXCytbVCRoQQQgghLaMCoBuq1Kix6Xqy6LiEwltILc6zQkadm5OtLWZFRIqOW9o31grZEEIIIYQYRwVAN5RVWQqVTmtWbHqZeRteu7rHBg6Es4i7+f19fTG+Z6gVMyKEEEIIMYwKgG5I2/L5bCZp9ObHdmVBzi74btYcQUt6ojw98fWMWdT+kxBCCCHtgmYg3ZCHwt78WDvzY7u6AX5+WLdwMaaHhRuc3LsoFHh0wED8Pn8h3JXmbxwmhBBCCGkNOgegG/Kyc8Agz0Ccys8UFeehsMdgzyArZdU19HRzw6dTpyG/shI7rlxGfmUlZBIJgl1cMbFXKBQy8ZuFCSGEEEIsiQqAbmpp7/6iC4BFoXF0GJhAnvb2uDc2rr3TIIQQQghphgqAbmpyUAQGXDqDMwU3BY33t3fG/eEDW3XNK3mFSMzKQY1GC0eFLeJ7BsLDgZYUEUIIIYS0JSoAuikbiRTfjJ6HZftXIqko2+hYX6UTfhizEO5m7h3Ye/Eqlh89g9PXsxq9LpNIMCm6Nx4ePhCRvl5mfTYhhBBCCBGHNgF3Y662Svx21z34U5+R8LZzaPa+g40t7g8biA2TliHU2UP053PO8f6uP/DUb5uaTf4BQKvXY2tSGhZ++zt2pVw269dACCGEEELEoScA3ZxSJsezfUbiiehh+CP7GjIriqHlHD52Dhjr3wtKmdzsz/7u8Gl8d+S0yXFqnQ7Pr96G5fffjUHBAWZfjxBCCCGEmEYFAAFQuyRonH8vi31ecVU1Ptt/TPB4rV6Pd3cdwupHl1gsB0IIIYQQ0hwtASJWse5cMtQ6naiYpKxcJGXlWCkjQgghhBACUAFArGRrUlqbxhFCCCGEEGGoACBWUVBRaWZclYUzIYQQQgghDVEBQKxCJjHvW0smpW9JQgghhBBrotkWsYpgd1ez4nq4uVg2EUIIIYQQ0ggVAMQq5g2IER0jZQxz4qKskA0hhBBCCKlHBQCxirsiesHTUdzJweMjQ+Hj7GiljAghhBBCCEAFALESuUyKd+dOhlTgXgAfJwe8MmWslbMihBBCCCF0EBixmvieQfhyyUz8adVWVKk1LY4LdnfFN0tnw9vJoQ2z6zjKy2uwfU8SziRcR2WlCgqFDWIi/TFtUl94etATEUIIIYRYFhUAxKpG9Q7BrucewJozF7DyTBKyS8tvvxfj540lg2MxNSYcCpvu962o1enx7Q8HsX7zOajU2kbvnTqbgR9/O4rxoyPx/FMToFTatlOWhBBCCOlqut+si7Q5Dwd7PD56CB4dORj5FZWo1mjgrFDA1d6uvVNrN1qdHq+/vRGHj11ucYxez7F7fwpu3CzCx/9eSEUAIYQQQiyC9gCQNiORMHg7OSDY3bXNJ/+cc6jVWnDO2/S6LVn+y2Gjk/+G0i7n4D8f77ByRoQQQgjpLugJAOmydDo9Th27gs3rziDhzDVo1DpIpRJE9QnAjLsHYvjoCNjYSNs8r6pqNdZvPisq5sDhNGRlF8Pf17zzFQghhBBC6lEBQLqkvNxSvP7iSly9nNvodZ1Oj6SEG0hKuAFff1e88d5C9AjxbNPc9hxIQWWVWnTcxq0JePJh6pRECCGEkNahJUCkyykqrMALT/zUbPLfVHZWMV544kfcvFHYRpnVOpd4w6y4BDPjCCGEEEIaogKAdDmfvLsVudklgsaWlVbjnX+sb9O9AdVm3P0HYNZTA0IIIYSQpqgLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 900x600 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {
|
||||
"needs_background": "light"
|
||||
},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.figure(dpi=150)\n",
|
||||
"plt.scatter(lonlat[:, 0], lonlat[:, 1], c=full_dat.mean(-1))\n",
|
||||
"plt.axvline(-128)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 8,
|
||||
"id": "03d676dd-4cff-4a12-9b1e-3877242b1ecc",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import pickle as pkl"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 9,
|
||||
"id": "e07bed54-2fe3-4c05-a530-4ae1032b1f60",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"dicts = [stn_names, stn_lonlat, stn_data]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 10,
|
||||
"id": "fbd09e88-8c49-4490-a0a7-f26985ae1366",
|
||||
"metadata": {
|
||||
"tags": []
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"pkl.dump(dicts, open(\"./wind_data.p\", \"wb\"))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 11,
|
||||
"id": "1c435ca6-a562-4351-8271-865309cf4925",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"stn_names, stn_lonlat, stn_dat = pkl.load(open(\"./wind_data.p\", \"rb\"))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 12,
|
||||
"id": "a1e96c14-3a79-410d-ad50-496f491b7769",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"lonlat = np.array(list(stn_lonlat.values()))\n",
|
||||
"full_dat = np.stack(list(stn_dat.values()), -1)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 13,
|
||||
"id": "514de640-a211-40f1-912e-951db384e5a0",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"conus_idxs = np.where(lonlat[:, 0] > -128)[0]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 14,
|
||||
"id": "42cf1c4b-8398-4d63-a1c5-106f8a123aea",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"full_dat[full_dat == -99.] = 0."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "51846f2b-4f1f-4bc6-939d-b5e3f02c1560",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.12"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,529 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 54,
|
||||
"id": "886427e7",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import torch\n",
|
||||
"import seaborn as sns\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"\n",
|
||||
"sns.set_style('white')\n",
|
||||
"palette = [\"#1b4079\", \"#C6DDF0\", \"#048A81\", \"#B9E28C\", \"#8C2155\", \"#AF7595\", \"#E6480F\", \"#FA9500\"]\n",
|
||||
"sns.set(palette = palette, font_scale=1.5, style=\"white\", rc={\"lines.linewidth\": 3.0})"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 55,
|
||||
"id": "f7e159c2",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"mtgp_magpie = torch.load(\"./full_ewma400_theta005_.pt\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 56,
|
||||
"id": "5eab3be8",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"['quantiles',\n",
|
||||
" 'x_paths',\n",
|
||||
" 'v_paths',\n",
|
||||
" 'covar',\n",
|
||||
" 'train_v_list',\n",
|
||||
" 'names_list',\n",
|
||||
" 'pars']"
|
||||
]
|
||||
},
|
||||
"execution_count": 56,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"list(mtgp_magpie.keys())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 57,
|
||||
"id": "2b7ca372",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"inds = [3, 10, 30, 50, 70, 90]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 58,
|
||||
"id": "2883a443",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"torch.Size([111, 100, 126])"
|
||||
]
|
||||
},
|
||||
"execution_count": 58,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"mtgp_magpie[\"x_paths\"][0].shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 59,
|
||||
"id": "ed7fd1e4",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"array(['AR_Batesville_8_WNW', 'CA_Merced_23_WSW', 'IL_Champaign_9_SW',\n",
|
||||
" 'ND_Jamestown_38_WSW', 'ON_Egbert_1_W', 'UT_Torrey_7_E'],\n",
|
||||
" dtype='<U34')"
|
||||
]
|
||||
},
|
||||
"execution_count": 59,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"mtgp_magpie[\"names_list\"][-10][inds]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 60,
|
||||
"id": "b83b62de",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"paths = mtgp_magpie[\"x_paths\"][0]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 61,
|
||||
"id": "660a5dff",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"path_lmeans = paths.log().mean(1)\n",
|
||||
"path_lstds = paths.log().std(1)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 62,
|
||||
"id": "0db951cc",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"lower = paths.sort(1)[0][inds, 33]\n",
|
||||
"upper = paths.sort(1)[0][inds, 67]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 63,
|
||||
"id": "c5a35f3f",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"[<matplotlib.collections.PolyCollection at 0x7f6af859f5e0>,\n",
|
||||
" <matplotlib.collections.PolyCollection at 0x7f6af859f9d0>,\n",
|
||||
" <matplotlib.collections.PolyCollection at 0x7f6af859fd90>,\n",
|
||||
" <matplotlib.collections.PolyCollection at 0x7f6af85ad190>,\n",
|
||||
" <matplotlib.collections.PolyCollection at 0x7f6af85ad550>,\n",
|
||||
" <matplotlib.collections.PolyCollection at 0x7f6af85ad910>]"
|
||||
]
|
||||
},
|
||||
"execution_count": 63,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAXEAAAEACAYAAABF+UbAAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjUuMCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8/fFQqAAAACXBIWXMAAAsTAAALEwEAmpwYAAB8nElEQVR4nOy9ebCta17X93med17z2sMZ79T3dt+WSzMYuiEaqhWIQzVUuiRYhQERA4EYUVtFTRWCKcuKiaId7EgoRIrqxlidiqIVoUQojRWjDRpx6G6g5zucYU9rr+mdnyF/vGuvvdfea0/n7L3P2ee+H+o2Z6/xWe/wfZ7nNwprraWmpqam5loin/QAampqamoenVrEa2pqaq4xtYjX1NTUXGNqEa+pqam5xtQiXlNTU3ONca/qi7Is45Of/CTr6+s4jnNVX1tTU1NzrdFas7W1xXve8x7CMDzy/JWJ+Cc/+Um+4zu+46q+rqampuaZ4u/8nb/De9/73iOPX5mIr6+vzwdy69atq/rampqammvNw4cP+Y7v+I65hh7mykR8z4Ry69Ytnnvuuav62pqamppnguPM0LVjs6ampuYaU4t4TU1NzTWmFvGampqaa0wt4jU1NTXXmFrEa2pqaq4xtYjX1NTUXGNqEa+pqam5xlwbETfWYK150sOoqampeaq4NiIOltJkT3oQNTU1NU8V10jEoTT5kx5CTU1NzVPFNRPxDGP1kx5GTU1NzVPDtRJxbRWqNqnU1NTUzLlWIm4x5LoW8Zqampo9rpWIA5Qmrk0qNTU1NTOunYhXJpXawVlTU1MD11DELRZt1ZMeRk1NTc1TwbUTcQBliic9hJqampqngmsp4trWIl5TU1MD11TElS1r52ZNTU0N11TErTUYU9vFa2pqaq6niGPR1CJeU1NTcy1FHCzalE96EDU1NTVPnGsj4tZayCzCuEgctK1FvKampsZ90gM4K0YZtj91D1UWtG+t4LwYPOkh1dTU1DxxzrwS/5Vf+RX+6//6v+a9730vX/VVX8UHPvABPv7xj1/m2I6gy5I8iRm8fp90Y1wn/dTU1LztOdNK/Od+7uf4oR/6IX7/7//9fPd3fzee5/GFL3yBsrxKk4ZAzP5ltGb7c2/QaqzQubF+hWOoqampebo4VcQfPHjA//A//A/8yT/5J/lv/pv/Zv74b/ttv+1SB3YYKQRC7A+3LHK2PvM6QdQkaDeudCw1NTU1TwunmlP+z//z/wTgD/7BP3jpgzkNT3gLf+fxlO3PvIHKaydnTU3N25NTRfxf/+t/zSuvvMI/+Sf/hN/ze34PX/ZlX8b73/9+fvRHf5SiuNr0d4F7wKgCymqy4ZSdz72F0XUT5Zqamrcfp5pTNjc32dzc5C/9pb/En/gTf4J3vvOdfOITn+Anf/InefDgAX/tr/21qxgnAFK6SCHRs5R7bQuM1cSbu7Ru9mmu9a5sLDU1NTVPA6eKuLWWOI7563/9r/PN3/zNAHzd130dWZbx0z/90/zxP/7HefHFFy99oFDZxR3ho206G5vGGIUUDulgUot4TU3N245TzSm9Xg+Ar//6r194/P3vfz8An/rUpy5+VCfgyn27uKUqhgWQbI/QZR1yWFNT8/biVBF/9dVXT/4AebVJnxJv0S5ucqwFlRdk4/hKx1JTU1PzpDlVgX/X7/pdAPzzf/7PFx7/5//8nyOE4Cu+4isuZ2THsGcX30NbhTYl1hiy3emVjqWmpqbmSXOqTfz9738/73//+/mLf/Evsru7y7ve9S4+8YlP8NGPfpRv//Zv5+7du1cxzjlH7OJYDCXgkWwP6b90C+k6VzqmmpqamifFmTI2f+zHfoyPfOQj/NRP/RS7u7vcvn2bD33oQ3zv937vZY9vKY5YHHZpCnynQZnlZOOYxkrniYyrpqam5qo5k4g3Gg3+3J/7c/y5P/fnLns8Z0KIxZW2tiXGaCQQbw2J+m2EEMvfXFNTU/MMcW1K0R5EIjko0cZq1Kzv5uTBNuN7209mYDU1NTVXzLUU8WqVvTj00qRYC1YbBl+4R7w9PPEztKrDEWtqaq4/11TEnSPmEmVKtKlW46ZUbP/mm2Sj5SGHRmnirdGlj7OmpqbmsrmWIi6FRBwausVSmGz+t8pytn/zdYo4O/x2ksGYfJJc+jhrampqLptrKeIAUhwNI1QmQ5t9M0k+Sdj+zddR2WKhrunDAXmdGFRTU/MMcH1FnKMibjAoky88lu5O2PqN11FZ9Xg2mpIOJ+isoKxL2NbU1Fxzrq2IO3J5Qo9a0kA52Rmx8akvUsQZ8dYQUyqM1qj0qKmlpqam5jpxbRolH2X5/GNsibX2iOMzG07Z+OQX0EUl8kZpVFavxGtqaq4311bEq1hxgcUuPG6sxRiF43hH3lNMF52ZZVKvxGtqaq4319acIoRcqGa4h8WgOVuXn8OiXlNTU3PduMYifjRWfA+zxC6+jGKaYU3d1q2mpub6cm1FXApxJFZ8D23Plo1ptKJI89NfWFNTU/OUcm1FHI4Wwtpjz7l5GkYZVHq1zZ5rampqLpJrLeJSLPfLGmsxZ1iNW2Pm8eM1NTU115FrLeKOWD58i0Hbszo36wiVmpqa68u1FnGxJGtzj7M6N5OdEWVtF6+pqbmmXHMRXx5mCGd3bqosZ3xv6yKHVVNTU3NlXGsRl0KeGGZ4FucmVI0kjitbW1NTU/M0c71FXDrHmlSMNWhzNpOKLhSjNzfOLPo1NTU1TwvXWsQBHHE0vR6q+uLLimEdR7w9JBmML2pYNTU1NVfCtRFx6Tr4rcaRxw93vj+IMhlnXVxbbZg+2HnU4dXU1NQ8Ea6NiAsh6L1wA+kumk+cYxJ+oHJuntWkAlXHn9o2XlNTc524NiIOEPU7NG+sLDwmhIM85mdYLNqePSPTlIrp5uCxxlhTU1NzlVwrEQfoPr+OE+zbwaVwEcck/QAoe74Y8OnGoC5RW1NTc224diIetBq0b63N/xbi+PR7AG3OZ1LReUm8PXycIdbU1NRcGddOxAH8dsTBHJ+TnJsGQ6FTchWTq/hMjs50MLmAUdbU1NRcPteys48X+EjHwSgNgDwmzHCP3FTNHwQSV/rHhiXuUcQpulQ43rU8PDU1NW8jTlWpX/mVX+G7vuu7lj73C7/wC7zyyisXPqjTcBsB0t0XceeYVm2HqQpjaRxOFnGjNEWcEvXaFzbmmpqamsvgzEvNH/zBH+R973vfwmPPPffchQ/oLLi+hxsFqKyKPBHSRQqJtvrU956lRK1Ruqoz3nvckdbU1NRcLmcW8Xe84x189Vd/9SUO5XwErQbZbmW7lkIghXsmET9ryGHVf3P1kcdnjUHIa+lyqKmpuUZcW5XxW9HC36fZxffQVmPOIPbZ+PGSfuqkoZqamqvgzCL+Iz/yI7z22mt8zdd8Dd///d/PJz/5ycsc16m4oY9w9od/UubmQazVGHO6iJdpTpk9euu2dHdSN2Guqam5dE41p7Tbbf7QH/pDfO3Xfi29Xo/Pf/7z/ORP/iR/4A/8AX72Z3+Wr/qqr7qKcR7BiyrnptaVUErcMzo3waAA/8TXGaUpkwwvPPl1S7/DGLJxjMoKvEZ47vfX1NTUnJVTRfy1117jtddem//93ve+l2/8xm/kW77lW/jwhz/Mz/zMz1zm+I7FDX0cz0PnVSKPlBIhJPYMphJlSvxTFu5WG8okh5WTX7eMIs4okwxVKLyjNbtqampqLoxHsomvr6/z9V//9fz7f//vL3o8Z0YIQdDeV0gpHOQJ6fcH0bbEWrDWkqvpsa8rJsmJn3OcuaSYpqisQOVnzxStqampeRQeOZvFPAX23sPOTQcPxenCudcwojQpyua4JsKRR5fm2TjGWosQYmYiSUi2R6S7E4xSSNeh+9wNGmvdhcSgbBSDtais7t1ZU1NzuTySiG9tbfEv/+W/fOIhh14jQEiBNZUd3JEunGFusRhSPUJbhYCZCeaoiKssZ+vTX0IrXa2ss3yeYLTH1m98ibDTYv3LXsJrBBhtyIZV6KOqC2nV1NRcMqeK+J/+03+a559/ni//8i+n0+nwhS98gb/1t/4WWZbxp/7Un7qKMR5L0Irov3ibIslId6fItKhE+Qzv3WukfJKj0yjN5OHJjSKssaTDCduffZMbr71EmeSovIpqKWoRr6mpuWROFfF3v/vd/PzP/zw/+7M/S5qm9Ho9vvZrv5Y/8kf+CK+++upVjPFY3DCg//IdAHa/9ICtz6YIJPYsy/EDaKOWLcTPRbIzZPCF+wStxny1rrKirsFSU1NzqZyqLt/3fd/H933f913FWB4Lv9XAcTyEcsCeU8TP0YvzWCxM7m+T+PtJR1YbVF7UIl5TU3NpXNuMzcP4rQjXd0+tULgMYzX6DAlAp2GNWXBmGm3mIZA1NTU1l8EzI+Je6OM3ozNnbh7EYrBnKIx17s81BlXUIl5TU3N5PDMiDhCutE9sEHEclXPz8Vfiy1Dpo6fu19TU1JzGMyXifivCcUPcRzCpaHPxK3Gg7tdZU1NzqTxjIt7ACzx85/y57uoinJtLKOJaxGtqai6PZ0rEvcDDb0a4Mji3WcVekHPzMKZUlLVzs6am5pJ4pkQcIFppI4UkOOdq3GAoTXrh49FlOc/grKmpqblonjkR91sNhCNnq/HzRarkOqbQFyvk1ljGb21i9JOvNVNTU/Ps8cyJeNhp4oUBUjj48nyrcYslVxOUvtiIknwcE28PL/Qza2pqauAZFHHpOjTX+wC4Mjz3alxjKC7YrFKvxmtqai6LZ07EAaJ+C+k4ONLBk9HpbzjEXr3xiyQfxyT1arympuaCeSZFPOi25m3RPBkiz/kzrTVnaqZ8rs80lmRnfKGfWVNTU/NMirh0JM31HlDVGPfk+fpcXlYafj6Ka5NKTU3NhfJMijhA2G8j3coe7sngXO+9rDR8VRQU8cWHMdbU1Lx9eXZFvNMg7LVBCKRwked1cF5CGr5RmrLO4KypqblLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"[plt.fill_between(torch.arange(126), lower[i], upper[i], alpha = 0.3) for i in range(6)]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 64,
|
||||
"id": "26db9a7c",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAA64AAAE2CAYAAABhrM91AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjUuMCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8/fFQqAAAACXBIWXMAAAsTAAALEwEAmpwYAAEAAElEQVR4nOy9e7BmV1nn/33v98u5dXcuTYiBREIgYYwgkEIg3io6EhDGSxJAwkQMIogg1iD4G8cqwbEADZCICVpCYFIz3EQtFQwzJUwIMkCVDCWJF0DS6e5ze+/3y++PM591nnef/Z5bn9N9+vT6VnV19/u+e++11157ref7PN/nWZHxeDyWh4eHh4eHh4eHh4eHh8cBRfRcN8DDw8PDw8PDw8PDw8PDYzN44urh4eHh4eHh4eHh4eFxoOGJq4eHh4eHh4eHh4eHh8eBhieuHh4eHh4eHh4eHh4eHgcanrh6eHh4eHh47AqDwUDf/e53NRgMznVTPDw8PDwOATZbV+LnoD27QqfT0de//nUtLCwoFoud6+Z4eHh4eJznGA6HWlxc1DXXXKN0On2um3Ne4rHHHtOP/MiP6P7779exY8fOdXM8PDw8PM5znDx5Urfccov+5m/+RpdddtnEd+cNcf3617+uW2655Vw3w8PDw8PjkOH+++/X9ddff66bcV5icXFRkvz67OHh4eGxp1hcXDx/ievCwoIkea+uh4fHeYnhcKhoNKpIJHKum+Lx/4BXl/XFY+e4kNbm8Xis0Wjk32MPDw+PfcRma/N5Q1yRBx87dkyXXnrpOW6Nh8f+AcPI4/Cg0+mo3+8rEokol8t5o3cTjMfjs94/Pv1k9zjf1ubRaKRWq6VIJKJsNrujsebfYw8PD4+zh7C12VvHHh4HCN1uV81mU61W61w3xWMPMRwOJa1HbDzC0W631Wg01Ol0znVTPA4per2ehsOhBoOB+v3+jo6lUMh4PHbvtIeHh4fH2YMnrh4eBwgYUsPh0BOcQ4TxeBz678OK3dzjeDx2xGCnhMLj/MdwOFS3293wrqyurur06dN75szo9/tqtVpqtVo7qoQ8Go0m2rZf87Of9z08PDym47yRCnt4nCuMx2P1ej2Nx2OlUql9k4cFDaMLgeBcKLiQnmu/31en01EsFlM2m932ccF+8ZL5CwsnT55Ur9dTsVjU3NycpDVpbq1Wc3PjXlR+7vV67t87cZAECWXw/1biPhqNFIlEdrxWNJtNjUYjpVIpJZPJHR3r4eHhcSHAWwUeHltgMBio1+up3+/vayTIe9oPJ4KE7LATV4jBcDickFMOBgM1m80J4mARRgQ8LgyMx2O1Wi31+31VKhU3blqtltrttlqtlur1+p5cy0ZZh8OhxuPxlmMtTBpsx2uv11Oj0VC73Va73Vaz2VS73d5Ru6zKxisOPDw8PMLhI64eHlvAGij7SS7PF8N9NBqp0+koEokonU4f+gIlPIfd3udWxPWwRRbtOLb32u12NRqN1O12lUgkNvTnhUbwPdYRiUQ0HA7VarWUSqVUq9WUy+VUr9ddFHI0GqnX67lIJJ/F49PNGMhjJpORtJ5jztiCcFKoKew9REHQ7XaVTCYnoqoAZ0y/33eRV4jodt/ts7XOeHh4eJzP8MTV49CBwhuJRGLPCcF+GtPnC3Ht9/su+kA/ny2cbZJHBVJJUw3brRBGVAFVSncqqz2o2GwMBw3zYLXAzfrJ43CDqGe/39doNNLKyoq63a7q9bojg91uV+12W8lkUqPRSNVqdUI+PB6Plclk3Dva6XRUqVTcd9lsVoPBwEX+JSmVSrl1otVqKRaLKRaLTch0IaWDwUCxWMwRZaKw0WjUjd3xeKxOp6PRaKREIqF4PK5IJKJEIrFl5eqw8b8bubGHh4fHYYYnrh6HCuPxWO122xkVe0EGzkZ+YlBWuZ/XOlNsp0AJRmg0Gt00IrITtNttR5T3ItdtO8BoltYM193knW0WSQwW4zqIkdfhcKh2u72t7UOm3et2oqk+4nphIxaLaTAYqNvtOgLLfD4cDpXL5dRqtVQoFNTpdHT69GmNRiNls1lHOIfDoRKJhJMeQ1wTiYQSiYSrWs1712q11O12FY1GlcvlNBwOlU6nNTMz485j2xSNRtXr9RSLxRSNRlWr1SbmN6LC8Xhc/X5fjUZD6XRa/X5fmUzGRYjD3vPgOkMkOJ1O74lzcDgcqlKpKBKJaGZmxhNiDw+P8xKeuHocKlgZ2F5tV7DfxLXT6bjcqHw+7wyKfr+vXq+nRCKhVCq1o3N2u10NBgOlUqk9I45gO8S11+u5SAWRjlgspkgkok6no3g8rlgspm63q3g8vuX9jUajiYqz2y2SdaZ7gk6Tve4EZ0LmDgIg7xjxmxnRwfFApdigoR42bnzE9cLG6dOnderUKY3HY0WjUaecSaVSbvy1Wi01Gg01m01Vq1UNBgNVKhWl02mlUik3XyIBXl5elrQ+DldXV53CQZIqlYoGg4Gi0ag6nY4jwcjZh8Ohms2marWaer2e0un0RNSVaGomk3HEGdIdiUQ0Go2UTqdddWRIbC6X23D/drzjyEomk+r3+7smrnbPWtYTaY2wh7XBw8PD46BjS4v2K1/5it73vvfpkUceUaVSUS6X05VXXqnbb79dP/iDP7jpsXfddZfe+973bvh8fn5eX/jCF3bfao8LChjMeLm3+m3w/2eam7hXxHUwGKhWqymRSKhQKLjPm82mOp2OOp2OMpmMM4o6nY6SyaTL69rufeD1l+SI4V5iO/1hnQZ2T9pIJOKqNAOMzc2e7WYVPKeh1Wq5CMpuDb+9ePbnO3HdSe5d8HucFEgqeWZh93q+SOX3Gw899JA+9alP6atf/apOnjypUqmkpz/96Xrd616nq666atNjd7rm/umf/qnuv/9+PfbYYzp27Jh++qd/WrfffvtZj/wPBgOdOnVKlUrFzVmRSMTNDbFYzOW/ttttVSoVJyeORqPKZDJKJBLK5XKamZnReDzWd77zHdXrdY1GI1UqFR07dky1Wk2pVErFYtGR0lgs5irGJ5NJNZtNJZNJJ0+u1WpqNpvq9/tqNpsuXzYajSqVSmkwGGg0GjnVDMWlEomEZmdnHQmlyB/n5l2oVCruPcGBR9Eoae29gAyzHozHYyWTyS2fE78lcg2249QlHz0onfbw8PA4l9jSoq3Varr88sv1kpe8RPPz86rVanrggQd0xx136F3vepd+/Md/fMuL/PEf//GEZPNs5sR5nP/odDoaDAaKRCLK5/Ob/nY3BGfaeVqtllv0NzO4t4t6va5ut6tutzux3YGVxQ2HQ8XjcVdEJOw++A4iEJSYBonGmUYdg8Dgi0QiU732wegBEY5+v+/uj3vl9zshrvb3wXNJa4YwxtmZRCyC0WUkhjs5X3DMDIfD0D0pp40t7o+I9dnGbokrDic7nhkHXio8HR/96EdVqVT0yle+UldccYWWlpZ077336qUvfak+9KEP6brrrtvyHNtZc9///vfrrrvu0mte8xr9wA/8gL761a/qPe95j6rVqt70pjft5S1tC0RTIaPJZNLlnaZSKRcRHY1GOnHihPr9vnvPy+WyksmkqtWqGo2GJGllZcURxpWVFdXrdY3HY+VyOTUaDRfBhWRCTolOzs7OqlKpqN1uO2LabrddWwaDgUqlkgaDger1ujtvo9FQKpVykc5Op+PSG5AR48gZjUaOUNbrdUdceY9492kvUeXhcKhSqeTOMS3/3r6PvIvBzy1sigfX4biDmMbg4eFx4WFL4vr85z9fz3/+8yc+e8ELXqAbb7xRDzzwwLaI6zXXXKNisbjrRnpc2EAiagnbNFgCiCd9t9e0lSetlBUiSCGOSCTiJF38OwgMAgCpghhjSAT/tscDcr4gARhetkhJsE/CCoNgkCF/2y6sFz9saxPulb08+WPvy0YUwtocds6w/2N8UgyFPtirCp32uhhy0lrkeLeR7G6366qOWjI6rQ+4v3g87qI9/H7as7VtJlKzHbUCY5QxzHO2v9kMYU6TsOO2Q1wvVKnwb/7mb7p9TMENN9ygG2+8Uffdd5/uuuuuLc+x1Zq7urqqe+65R7fccote//rXS5Ke9axnqd1u695779Wtt96qY8eOndmN7ADRaNQ59cghtZFGxvnJkydVrVZVrVYlrdcGYIyTjyqtRTIhroPBwBV2arVaLuLaarWUyWSchDcajSqRSGhpaUnFYlHVatVFfvv9visQ1W63J5Qw7L3aarVcLvxwONTy8rKy2aySyaSOHj3q7jPomOReiDATneX++H2j0XD/7/f7ToYsacKJCNGV5Jy+pI8gb2a9svOYzS8GYc5RgPx4s3SPXq/n+uRM1D/cj8/N9fC4sLGrWSQej6tQKPjIqce+Y6dRGDzYGDrkHp3JdafJj1mQpUmiG1aQCDkZgABZz7r9fBphgJRLcoYHEQL7WbBPgn2AMSetGRaZTGbbkVkrM8OIsschS2u1Wk7GhzwueN82F2wzhEVc+Zt/Y5zafRppz24QfAbI1SVtiPBudZ5g2zFEI5HIplHI4P1Zp0mz2dR4PHY50MHnTF4f592soBVOGPoO43an75+VN9p3QlrfViQej4cWTZtGZu27dqbG7/mAIGmVpGKxqMsuu0wnT57ck2v83d/9nbrdrl784hdPfP7iF79Y99xzj/72b/9Wt9xyy55caztgDqNiuZ1j7FiCPKKG4XvrUIKc1Wo157zkeIgx7x/kEEns8vKyksmkUqmUHnnkEUeCm82mI27Ly8vqdDpunk8mk+531WrVEVDmI/JJIcixWEyLi4uS1m0p6yQlvxd1Ckod2pFIJJRMJtXpdJxzgYKEoF6vazgcuvmFLcyoDdBsNh3RJK9X0sR6Zrf2mVYrgfllWjoLfS6tKafy+XzoOsNcO825Zuc7m07j4eFx4WHbbz8G1PLysh544AF961vf0q/92q9t69ibbrpJy8vLmpub0/Of/3z9yq/8SujLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 1152x360 with 6 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(3, 2, figsize = (16, 5))\n",
|
||||
"ax = ax.reshape(-1)\n",
|
||||
"\n",
|
||||
"for i in range(6):\n",
|
||||
" ax[i].plot(paths[inds[i]].t(), alpha = 0.1, color = \"grey\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 65,
|
||||
"id": "25e263c5",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ntrain = 252\n",
|
||||
"ntest = 126\n",
|
||||
"neval = 25"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 66,
|
||||
"id": "3bfde56e",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"[[0, 252], [4194, 4446], [8388, 8640], [12582, 12834], [16776, 17028], [20970, 21222], [25164, 25416], [29358, 29610], [33552, 33804], [37746, 37998], [41940, 42192], [46134, 46386], [50328, 50580], [54522, 54774], [58716, 58968], [62910, 63162], [67104, 67356], [71298, 71550], [75492, 75744], [79686, 79938], [83880, 84132], [88074, 88326], [92268, 92520], [96462, 96714], [100656, 100908], [104850, 105102]] [378, 4572, 8766, 12960, 17154, 21348, 25542, 29736, 33930, 38124, 42318, 46512, 50706, 54900, 59094, 63288, 67482, 71676, 75870, 80064, 84258, 88452, 92646, 96840, 101034, 105119]\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
" import pickle as pkl\n",
|
||||
" import numpy as np\n",
|
||||
" \n",
|
||||
" stn_names, stn_lonlat, stn_data = pkl.load(open(\"./wind_data.p\", 'rb'))\n",
|
||||
" lonlat = np.array(list(stn_lonlat.values()))\n",
|
||||
" full_dat = np.stack(list(stn_data.values()), -1).T\n",
|
||||
" full_dat[full_dat == -99.0] = 0.\n",
|
||||
" names = np.array(list(stn_names.values()))\n",
|
||||
"\n",
|
||||
" conus_idx = np.where(lonlat[:, 0] > -128)[0]\n",
|
||||
" lonlat = lonlat[conus_idx]\n",
|
||||
" full_dat = full_dat[conus_idx]\n",
|
||||
" names = names[conus_idx]\n",
|
||||
"\n",
|
||||
" full_data = torch.tensor(full_dat).float() + 1 # wind speeds can be zero.\n",
|
||||
" ndata = full_data.shape[-1]\n",
|
||||
" full_returns = torch.log(full_data[..., 1:] / full_data[..., :-1])\n",
|
||||
"\n",
|
||||
" T = 52\n",
|
||||
" ts = torch.linspace(0, T, ndata) + 1\n",
|
||||
"\n",
|
||||
" device = torch.device(\"cpu\")\n",
|
||||
"\n",
|
||||
" # train on one full year worht of data and test on 6 mos for defaults\n",
|
||||
" # we can do logner\n",
|
||||
" split_diff = int((ndata - ntrain) / neval)\n",
|
||||
" train_starts = list(range(0, ndata - ntrain, split_diff))\n",
|
||||
" train_splits = [[x, x + ntrain] for x in train_starts]\n",
|
||||
" if train_splits[-1][-1] > (ndata-1):\n",
|
||||
" train_splits.pop(-1)\n",
|
||||
" eval_splits = [min(ndata-1, x + ntest + ntrain) for x in train_starts]\n",
|
||||
" print(train_splits, eval_splits)\n",
|
||||
"\n",
|
||||
" ts = ts.to(device)\n",
|
||||
" full_returns = full_returns.to(device)\n",
|
||||
" full_data = full_data.to(device)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 67,
|
||||
"id": "7db3e048",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"i = 0"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 68,
|
||||
"id": "6db14348",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"[ 4 5 14 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33\n",
|
||||
" 34 35 36 37 38 39 40 41 42 43 44 45 46 47 49 50 51 52\n",
|
||||
" 53 54 55 56 57 58 59 60 61 63 64 65 67 68 69 70 71 72\n",
|
||||
" 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90\n",
|
||||
" 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108\n",
|
||||
" 109 110 111 112 113 114 115 116 117 118 119 120 121 122 124 125 126 127\n",
|
||||
" 128 129 130]\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"train_start, train_end = train_splits[i][0], train_splits[i][1]\n",
|
||||
"test_end = eval_splits[i]\n",
|
||||
"\n",
|
||||
"# filter out any tses w/ only missing data\n",
|
||||
"keep_idx = np.where((full_data[:, train_start:train_end]-1).mean(-1).cpu() > 0.)[0]\n",
|
||||
"print(keep_idx)\n",
|
||||
"clonlat = lonlat[keep_idx]\n",
|
||||
"full_data = full_data[keep_idx]\n",
|
||||
"names = names[keep_idx]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 69,
|
||||
"id": "aad7f0ea",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"[67104, 67356]"
|
||||
]
|
||||
},
|
||||
"execution_count": 69,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"train_splits[-10]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 70,
|
||||
"id": "2c648c98",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"train_y = full_data[:, train_splits[-10][0]:train_splits[-10][1]]\n",
|
||||
"test_y = full_data[:, train_splits[-10][1]:eval_splits[-10]]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 71,
|
||||
"id": "ae77ba82",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"train_x = ts[train_splits[0][0]:train_splits[0][1]]\n",
|
||||
"test_x = ts[train_splits[0][1]:eval_splits[0]]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 72,
|
||||
"id": "43998749",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAA6kAAAE6CAYAAAD0shR3AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjUuMCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8/fFQqAAAACXBIWXMAAAsTAAALEwEAmpwYAAEAAElEQVR4nOydd5hU1fnHP9PLdpaFpTcBRRQUxIa9RTQGey9RY40lliRqNIlJfkaTWKKxghoVjUlUNGqMoiY2BBEsWABRetk+u9Nn7tzfH4dz987MndnZZbYA5/M8+7Bz65lZds/9nvd9v69N13UdhUKhUCgUCoVCoVAo+gD23h6AQqFQKBQKhUKhUCgUEiVSFQqFQqFQKBQKhULRZ1AiVaFQKBQKhUKhUCgUfQYlUhUKhUKhUCgUCoVC0WdQIlWhUCgUCoVCoVAoFH0GJVIVCoVCoVAoFAqFQtFncPb2ABQKhUKhUOTm5z//OS+88ELO/e+99x41NTWW++69917uu+++rO39+/fn/fffL9oYFQqFQqEoJkqkKhQKhULRh7nssss47bTT0rYlk0kuuOACxo8fn1Ogmnnsscfw+/3Ga5fLVfRxKhQKhUJRLJRIVSgUCoWiDzN8+HCGDx+etu31118nGo1y0kknFXSNiRMnUl5evlXjSCaTbNq0idraWpxO9figUCgUiq0j37zSJ2eZaDTK0qVLqampweFw9PZwFAqFQrGNo2ka9fX1TJw4Ea/X29vD2Wqee+45fD4fM2bM6LF7rl+/niOPPJI5c+ZQW1vbY/dVKBQKxfbJpk2bOPPMM3n99dcZMWJE2r4+KVKXLl3KmWee2dvDUCgUCsV2xpw5c5g6dWpvD2OrqKur49133+X73/8+paWlBZ0zY8YMGhsbqa6u5uCDD+YnP/kJ1dXVnbpvfX09gJqfFQqFQlFU6uvrtw2RKutr1GqtQqFQ9C00TcNut2Oz2Xp7KJ1CrtYWUr/Z15k7dy6aphWU6jts2DCuueYadtllF1wuF4sXL2bWrFnMnz+f559/noqKioLv291z89Kl8Pvfg6bBUUfB2WeDSqZSKBSK7Zd8c3OfFKkyxbe2tpahQ4f28mgUCoWi50gmk8TjcVwuV58zt4lGoyQSCWw2GyUlJd0iVGOxGKlUCo/Hg91e/C5p20MJyfPPP8+IESPYa6+9Ojx25syZaa/33XdfJk+ezPnnn8+cOXO47LLLCr5vd87NH34IRx/d/vq992DVKpg1q6i3USgUCkUfxGpuVn1SFQqFoodIpVKG0MtFJBJB0zSi0WgPjqww5Lh1XUfTtK2+XiwWIxwOk0wmSaVSJBIJ4vE4yWSSWCy21dffHlm0aBHfffcdJ5xwQpevsf/++1NTU8Mnn3xSvIFtBbfcAvvum719zhwRVVUoFArFjkefjKQqFArFtoyu6wBZkcZoNIqmaSQSCRwOR1akUJ5nfl1otLKYEUhd10kmk9jtdsvVTV3XicVi2O32Lt9LRox1XScYDOLz+UilUsb1ksnkVr2H7ZXnnnsOh8PB8ccfv1XX0XW9WyLVnUXT4M47rfdFo/Ddd7B4MVRXw+DB8Je/wDffiP3jx8OFF8KIEbCVxsUKhUKh6GMokapQKBRFJJlMEolELFNizdHHZDKJ2+22PFfasGuaZnyfSqWM15nCVdM04vG4cZzL5TJqR0Gk0TidTlKpFJFIBACfz5dTpMTjceN6JSUlWcdFo1GcTie6rncq7TcWi5FMJvF4PIYIjcfjJBIJfD4fsVgMn89X0LV2RMLhMK+99hrTp09n4MCBXb7Oe++9R0NDA5MmTSri6LrG+vUQCuXeP2EC5Eo8+M9/4M9/FgL2tddgG/fDUigUCoUJJVIVCoViK5BRRZvNhtvtNkSg3C7bnaRSKWKxmCEYzUQiEYLBIMFgEE3TSCaTVFRUEIlEDBEYDofRdR2n05kl5JLJpJEya7PZSCaTtLW1EYvFcLlclJaWUlZWZhwDQmh6PB5DUPv9fkNsSoEqxa+5ZYt8Xw6HA13XjXt2FJUzC+lIJILL5ULXdUKhEIlEYrtoC9PdvPrqq4TDYU488UTL/WeffTYLFy5k2bJlxraZM2cyc+ZMRo0ahdPpZMmSJcyePZsRI0b0CZfelSvz78+TGW/Q2Ai/+hW8/HJRhqRQKBSKPoASqQqFQrEVyEggiPReXdeJRqNZqbKxWIxEIkEikTDErMPhoLS0lJaWFgCam5spKSkhHo+jaRqpVIqGhgbj+9LSUiPamkgk8Hg8eDweo9YV2sVga2srdrsdv99viEKJpmm0tbWRSqXwer14PJ40QS3fl6wP9Xg8ae9D1ou6XC4jhVm+ls6/sVgMj8djvNdYLGaIbGgXu6lUimQyaaT8SswGUtuak3B38fzzz1NVVcWhhx5a8DmjR4/m6aefpq6ujmQySW1tLSeffDKXXXYZ5b2cI9vcDDn0dqd55ZXiXEehUCgUfQMlUhUKxY6DFgdHe4qtFGFutzsruilpa2tD0zQjPdYs9nRdJxAIGPV9MhLqcDjQNC0tvVeaAtntdkOE6rpOPB43optmQyIZpZRRURm5tNlsaJpGOBxG0zQGDBhgHCvHG41GiUajRt2oFKAyvViaN4VCIUKhEDU1NdjtdlKpFOFwmGg0mlYTmkgkjM9HCvJEIkEgEDCitbqu4/V6jTHL9yGFtRSa8jqRSITGxkbjs3A6nUZEuK2tjXA4TE1NDdXV1SoFeAtPP/103v1PPvlk1rY7cxV89jK6LgRqc3PxrtnQAP37F+96CoVCoeg9lEhVKBTbP4lWeGcm1L0DQ46D6c+STNmMiGAkEqGsrCzrtGg0SjAYJBKJoOs61dXVeL1eI7IpRV08HicSiRgRx5KSEkBEHcvKykilUjQ2NhIIBIx032Qyic/nMyKqDoeDRCJBKpXC4XAY6cFtbW2GmA4GgzgcDkKhkJH+KyOzsViMaDRKIBDA6XQa1/L5fLjdblpaWujfvz+BQMAYcygUMsRtdXU1ra2ttLa20tTUZKQVV1RUUFZWZtwrkUjQ2tpKKBQyzJ8qKyuNlF+73W5EemXUt6qqynjvMjIq31cgEADA4/GQSCSIRqM0NDTgdrtpampC13WGDh2qoqnbGa+9Bm+/nb29f38hNrvCp5/CYYdt3bgUCoVC0TdQIlWhUGz/rHwMNm95Il73ApFP/0DTgHOx2+2UlpYCGAIPoLKyErvdTigUIplMGtFFaV5UXl6Oy+UiGAzS1tZGKBQiGo3i9/uNSKHD4cBmsxlicOPGjUZ9p8fjMYyNpPjyeDzE43F8Pp8R6ZSiUdaNOhwOPB4PgUDAEIlut5uSkhICgQBNTU1GBNTj8RhpxzKN1+l0GsK7oaHBOL+urg6Px0NjYyNtbW00NjaiaRqDBw82amRlJLeuro7GxkZaWlqM8UhhLUVoJBIx6l81TcPv99PW1maI59bWVhwOh1GH6/F42LBhA8lkEq/XSygUMsyVXC4X8Xg8LeVYsW2j63Dbbdb77rsPTjuta9edPRtefRX22gtOPRXUuoZCoVBsuyiRqlAoti0SQdj4HygdBf32TNuVSqWIx+M4nc709N11c9OO8311E83O7xOOROjfv79RS9rY2IjD4WDz5s1GZE/XdcLhMI2NjYa4klHTZDJJKBSisbHRqMN0uVzEYjHcbjcNDQ00NzcTCoVoaGgwakD9fj92u51EImFEcZPJJNFolObmZoLBIJs2bTJSZ6PRKBUVFXi9XiONV7ZvkdFLmSKbWYuaTCYpLy/H6/Wyfv36tEio2+02UojtdjstLS00NzcbNazyms3NzdhsNmw2Gw0NDQQCARKJBHa73Ugh9vv9Rg3uhg0bjO+TySTNzc1GxDYWixnOwDLVV36ecqHA7XbT2tqK3+83anSVSN0+WLQIjjsONm603r/vvlBTA/X1nb/2M8+0f6/rcPrpXRujQqFQKHofJVIVCsU2QSoFpDTsb+wPLZ8BNpj+LAw/2ThGmvjIGkyn0ynEUjK7x4Wt/n+EbLsSDofx+/1GaqmMDNpstrToaTKZZPPmzUZ6q+wjGg6HaWtrA0RbF1mzKo1/ZGRUpge73W6qq6ux2Ww4nU4jEivrW5ubm426TxlJlHWssh2MFLeypUxbWxuJRIJwOEw4HAba+6bKCKiskfV6vcRiMUKhkJFyK99HMBg0eqQmEgmam5sNA6hUKoXT6aSxsTGt9tYsuGUNbDQaNZx84/E4LS0t+P1+oy5XCl6Z+isjp9IZ2e12G9cMBoPd9V9K0Qv87Ge5BerYsTB0qIiEvvrq1t3nkUeUSFUoFIptGSVSFQpFn+d3v4NbboEjpn7Oa1d9tmWrjr7yUaI1x+J0Oo0IpqyblI6yPp+PqkQcV8Y13Zv/xcqQl3g8Tnl5OdFoNK2FjN1uT6s9Nbv4NjU1GeJPjweY6P8ITXfy6cYp6DZhIgTCRVeaCknM6bDmaK/H46G1tZVwOGzcEzDavWialiZ2E4mEIXSDwaBRgyqjv3Ksuq7T1NREeXm5Ef2U+zVNM7bJKLFEClu73W4YH8nzpIiVx0thKz+7lpYWI7IsW9O0tbUZx8nP126343Q6icViRCIRw0TK6XTidrvRNA2Px5P2+Sm2bXKl4O63H9xxB9jtot/p1opUq3pXhUKhUGw7KJGqUCj6NBs3CoGaSkFTQzxtnxbaRDQaxel0Gu66wWCQpqYmQ4BpmsYBwfVZInVgYgENDQej69Da2orT6UwTeVLMAWlOt4Bxz2QyyZmDn2SnEtHscaDzO17cfHzasVbnysgrtNeumoWpGSnsgsEgLpfLiOx2hDzG7MZrdgHuCJvNZrj6SmGaify8ZHpw5vuQ5+i6bnyW0kVYphhLAe5wOIzoq6SiogKn00k4HO71dimK4nDccfDmm+nbHnoILrqo/fXUqdbnDhsGa9cWdh/TeotCoVAotkGUSFUoFH2a11/fkuoLBGOlafsSkSa++eYbI11148aN1NfLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 1152x360 with 6 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(3, 2, figsize = (16, 5))\n",
|
||||
"ax = ax.reshape(-1)\n",
|
||||
"\n",
|
||||
"for i in range(6):\n",
|
||||
" ax[i].plot(test_x, paths[inds[i]].t(), alpha = 0.1, color = \"grey\")\n",
|
||||
" ax[i].plot(train_x, train_y[inds[i]], color = \"blue\", linewidth =4)\n",
|
||||
" ax[i].plot(test_x, test_y[inds[i]], color = \"orange\", linewidth=4)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 73,
|
||||
"id": "d5f2d96e",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"50"
|
||||
]
|
||||
},
|
||||
"execution_count": 73,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"inds[3]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 74,
|
||||
"id": "7eff5601",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"70"
|
||||
]
|
||||
},
|
||||
"execution_count": 74,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"inds[4]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 83,
|
||||
"id": "126730ba",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"train_x = torch.arange(0, 252*5, 5)\n",
|
||||
"test_x = torch.arange(252 * 5, 252 * 5 + 126 * 5, 5)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 93,
|
||||
"id": "4c6b96f4",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAicAAAGXCAYAAABocvA1AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjUuMCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8/fFQqAAAACXBIWXMAAAsTAAALEwEAmpwYAAEAAElEQVR4nOydd3gU5dqH79mS3nshQCCQhBZ66EgTLCBNVARFUOy9oaKez949HrArqNilK6A0Cx2k95IC6X3TN8nuzvfHuzW7CQFC07mvi4vszDsz72yZ+c1TJVmWZRQUFBQUFBQULhFUF3sCCgoKCgoKCgr2KOJEQUFBQUFB4ZJCEScKCgoKCgoKlxSKOFFQUFBQUFC4pFDEiYKCgoKCgsIlhSJOFBQUFBQUFC4pFHGioKCgoKCgcEmhudgTUFBQOP9kZGTwySefsGPHDnJycnBzcyM0NJTOnTszbtw4+vTpYx07Z84cEhMTGT58+Dkdc/HixTz11FMAPPHEE8yYMcNpzOHDhxk7diwA48aN47XXXjunY15IMjMzGTZsGABXXHEFH3/8sdOYuro6Bg4cSElJCdHR0axfv97hfTkdvXv3ZsGCBc06bwWFy4Emi5O0tDS2b9/O8ePHKS4uRpIkAgMDad++Pb169SI2NvZ8zlNBQeEs2b9/P1OnTkWj0TB27Fji4uLQ6/Wkp6fz+++/4+3t7SBO5s6dy7hx485ZnFhwd3dn8eLFLsXJwoULcXd3p6amplmOdTFwd3dnw4YN5OfnExYW5rBu/fr1lJSU4O7ubl3Wq1cv3njjDYdxH330EampqU7LQ0JCzt/EFRQuYRoVJzU1NSxatIgffviBY8eO0VAxWUmSaN++PTfeeCPjx493+CEqKChcXN5//32qq6tZunQpiYmJDuuee+45CgoKzuvxR4wYwS+//MK+ffvo0qWLdXltbS2//PKLdX1zU1FRgY+PT7Pvtz5Dhgxh7dq1LFu2jDvuuMNh3aJFi4iPj8dkMlFVVQVATEwMMTExDuMWLlxIamoq11133Xmfr4LC5UCDMSdLly5l5MiRvPjii/j5+fHwww+zYMEC/vzzT/bu3cuePXv4888/+eqrr3j44Yfx8fHhhRdeYOTIkSxbtuxCnoOCgkIjpKenExAQ4CRMAFQqFeHh4YBwU8THxwOwZMkS4uPjrf/OhSFDhhAUFMSiRYsclq9duxadTseECRNcbrdy5UruuusurrjiCjp16kRycjL33HMPR44ccRo7dOhQpk6dyqFDh5gxYwY9evRgzJgxrF69mvj4eH766SeXx7jmmmsYMWJEgw9eTSE4OJhBgwaxePFih+X5+fls3LiR8ePHn/W+FRT+rTQoTv7zn/8watQo1q5dy4IFC5g5cya9evUiPDwcd3d3PDw8CA8Pp3fv3sycOZNvvvmGtWvXcuWVV/Kf//znAp6CgoJCY7Rs2RKdTsfq1asbHRcUFGR1K/Ts2ZM33njD+u9c0Gg0jB49mhUrVqDX663LFy1aRIcOHUhISHC53ddff40kSUyaNInnn3+eSZMmsXPnTm666SbS09OdxmdnZ3PrrbcSFRXFE088wdSpUxk6dCihoaEsXLjQafyePXs4ceIEEyZMQJKkczrHCRMmkJqayu7du63Lli5dikqlYsyYMee0bwWFfyMNunXWrFlDaGjoGe0sOjqap59+mpkzZ57zxBQUFJqHu+++m82bN3P//ffTunVrunfvTufOnUlOTqZt27bWcV5eXlx33XU88cQTxMTENKuLYcKECXz55ZesWbOG0aNHk5uby+bNm3nmmWca3Oazzz7Dy8vLYdnYsWO57rrr+OKLL5wegjIzM3nppZe4/vrrHZaPHz+ejz/+mBMnThAXF2ddvnDhQtRqNePGjTvn87viiisICQlh8eLFdOvWDRABwUOHDiUoKOic96+g8G+jQcvJmQoTe5QgLgWFS4du3bqxaNEixo0bR3l5OYsXL+b//u//uPrqq5k8eTIZGRnnfQ7x8fF06tTJ6vpYsmQJGo2Ga6+9tsFtLMJElmUqKiooLi4mMDCQ2NhY9u3b5zQ+ICDApQvl+uuvR5IkB+tJVVUVK1euZNCgQVa31rmg0WgYM2YMK1eupLq6mp07d5KWltagy0pBQaFxzrnOyYEDB9i0adNlHW2voPBPJz4+ntdee43Nmzezfv16Xn/9dXr27MnOnTu55557qK2tPe9zGD9+PFu2bCErK4slS5YwbNgwAgICGhx/6NAh7rzzTrp3706PHj3o27cvffv25dixY5SWljqNj4mJQa1Wu1zer18/li1bRl1dHQCrVq2isrKSiRMnNtv5TZw4kYqKClavXs2iRYsICwtjwIABzbZ/BYV/E00WJ59//jl33XWXw7JHH32U66+/nttvv53Ro0dTWFjY7BNUUFBoXqKjoxk7dixff/013bt359ixYy4tEc3N6NGjcXNz49lnn+XkyZONWhWys7O5+eabOXToEHfffTfvv/8+8+bNY/78+bRr185lAKunp2eD+5s0aRLFxcWsX78eEC6d0NBQrrjiinM+Lwtt27YlKSmJb7/9llWrVjF27FiXYklBQeH0NFmcrFixgsjISOvrLVu2sGLFCq6++moefvhhCgoK+Oyzz87LJBUUFJofSZJISkoCRGbJ+cbPz48RI0awadMmIiMj6d+/f4Nj16xZQ1VVFW+++SYzZ85k+PDh9O/fn379+qHT6c742MOGDSM4ONiasrtr1y7Gjh2LRtO8dSgnTJjAnj17qKqqUrJ0FBTOgSb/MrOyshwCx9atW0doaChvvfUWkiRRUlLC+vXrmTVr1nmZqIKCwtmxadMmkpOTnW7Eer2eTZs2ATgFxjYkALKzs6murqZly5Zotdoznssdd9xB69at6dChAypVw89GFotDfQvJjz/+SEFBAdHR0Wd0XK1Wy7hx45g3bx7vv/8+QLO6dCxcc8015Ofn4+/vrxSmVFA4B5osTqqrq/Hw8LC+3rp1K/369bOm4LVt25bvvvuu+WeooKBwTrz66qvodDqGDh1K+/bt8fDwIDc3l59//pn09HTGjh3rUMuka9eubNmyhU8++YSoqCgkSeKaa64B4Mknn2T79u2sW7eOFi1anPFcEhISGkwdtmfQoEF4enryxBNPMGXKFPz8/Ni1axd//fUXLVu2xGg0nvGxJ02axOeff84vv/xC7969ad26tdOYbdu2ccstt5x1KX0fHx/uv//+M95OQUHBkSaLk/DwcI4ePQoIK8qJEyeYNm2adX1ZWRlubm7NPkEFBYVzY9asWaxbt46dO3fy22+/UV5ejq+vL+3bt+eOO+5wcj88//zzvPDCC3z00UdUVlYCWMXJhaJly5Z8+umnvPPOO3z00Ueo1Wq6d+/OggULePHFF8nKyjrjfbZq1Yrk5GS2bt3aYLyL5XybI4NHQUHh7JHkJpZGfOWVV/j222+ZNGkSe/fu5fjx46xfv96aNvzUU09x5MgRlixZcl4nrKCgoHC23HHHHezZs4cNGzY4WIItvPrqqyxevJg1a9Y0mkmkoKBwfmmy5eTee+/l6NGjfPvtt7i5ufH0009bhYler2fNmjXnxYeroKCg0BycPHmSjRs3cvPNN7sUJgAbN27k7rvvVoSJgsJFpsmWEwsVFRW4u7s7BMNZOpxGREQoP2oFBYVLir1795KSksKCBQtISUlh5cqVZxUvo6CgcOFoNJX40Ucf5bfffrN20wQR8FU/St/Dw4OEhARFmCgoKFxyfPfddzz99NNUVFTw1ltvKcJEQeEyoFHLyTXXXENKSgru7u707duX4cOHK70iFBQUFBQUFM4rp3XrnDx5ktWrV7Nu3Tr27t2LSqWia9eujBgxgmHDhhETE3Oh5npeMBgM5ObmEhER0ewFmRQUFBQUFBTOnDOKOSkoKGDNmjWsW7eObdu2YTQaad++PSNGjGD48OFNql9wqZGZmcmwYcPOum6DgoKCgoKCQvNyxgGxFioqKli/fj1r165lw4YN6PV6oqKiGDFiBNdff71DxclLGUWcKCgoKCgoXFqctR/Dx8eHMWPGMGbMGGpra9mwYQNr165l2bJl+Pj4cN999zXnPBUUFBQUFBT+JTRLkIWbmxvDhg1j2LBhmEyms2rMpaCgoKBw/pDrqqA6H3xbIklN7vmqoHBROGNxUl1dTVZWFjqdzmXb8l69einZPAoKCgqXELKxDtOhT6GmBCmsJ1Lray/2lBQUGqXJ4qSqqopXX32VpUuXYjAYnNbLsowkSRw+fLjJB581a1aj5e43btxIaGioy3Vz5sxh7ty5TstDQkKsnVYVFBQUFICSg1BTAoCc/zco4kThEqfJ4uT555/n559/ZsSIEfTo0QN/f/9zPvg999zDjTfe6LDMYDAwY8YM4uPjGxQm9syfPx8vLy/r67Np466goKDwT0auyLzYU1BQOCOaLE7WrVvHxIkTeemll5rt4C1btqRly5YOy1avXo1er29yn55OnTrh5+fXbHNSUFBQ+MdRlXuxZ6CgcEY0OSpKq9XSuXPn8zkXABYtWoSnpydXX331eT+WgoKCwj8dWTYhV+XVW3ZWFSQUFC4YTRYnycnJ7N2793zOhfz8fDZs2MDIkSPx8fFp0jZXX301iYmJDBgwgNmzZ1NUVHRe56igoKBwWVFdAKY6x2Wy6eLMRUGhiTTZrTNr1iymTJnCl19+yeTJk89LbMfSpUsxGo1NcunExMTwyCOPkJiYiFarZdeuXXz22Wds2bKFxYsXN0tMjIKCgsJlT3WB8zJTLag8L/xcFBSayBlViP3555958sknUalUhIaGolI5Gl4kSWLt2rVnPZlRo0ZhMplYvXr1WW2/adMmpk+fzoMPPsg999zTpG2UCrEKCgr/ZEy5W5FP/eqwTNX1ESQ3JVZP4dKlyZaTxYsX88wzz6DVaomNjW32INS///6btLQ0Hn744bPeR//+/QkNDWXPnj3NNzELine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 576x432 with 2 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(2, 1, figsize = (8, 6))\n",
|
||||
"ax[0].plot(train_x, train_y[inds[3]],linewidth=2, label = \"Observed Wind\")\n",
|
||||
"ax[1].plot(train_x, train_y[inds[5]],linewidth=2)\n",
|
||||
"\n",
|
||||
"ax[0].plot(test_x, paths[inds[3]][5:10].t(), color = palette[7], alpha = 0.5, \n",
|
||||
" label = [None, None, None, None, \"Volt Samples\"])\n",
|
||||
"ax[1].plot(test_x, paths[inds[5]][5:10].t(), color = palette[7], alpha = 0.5)\n",
|
||||
"\n",
|
||||
"ax[0].plot(test_x, test_y[inds[3]], color = palette[2], linewidth=2, label = \"Truth\")\n",
|
||||
"ax[1].plot(test_x, test_y[inds[5]], color = palette[2], linewidth=2)\n",
|
||||
"\n",
|
||||
"ax[0].set_title(\"St. Mary, MT\")\n",
|
||||
"ax[1].set_title(\"Sioux Falls, SD\")\n",
|
||||
"\n",
|
||||
"ax[0].legend(frameon=False)\n",
|
||||
"ax[0].set_xlabel(\"Minutes\")\n",
|
||||
"ax[1].set_xlabel(\"Minutes\")\n",
|
||||
"ax[0].set_ylabel(\"Wind Speed (m/s)\")\n",
|
||||
"ax[1].set_ylabel(\"Wind Speed (m/s)\")\n",
|
||||
"plt.tight_layout()\n",
|
||||
"sns.despine()\n",
|
||||
"plt.savefig(\"mt_wind_modelling.pdf\", bbox_inches=\"tight\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 104,
|
||||
"id": "9778360a",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAggAAAE7CAYAAAC4+nn9AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjUuMCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8/fFQqAAAACXBIWXMAAAsTAAALEwEAmpwYAACKD0lEQVR4nO3dd3wU1drA8d/MbEsvJBAgdAi9SZMiUhQUlSaiiCiK4mtBUSx47e0iiqJiwc4V8YoFBBVBKeoVKTakCYkU6RAgCalbZs77x0pg2Q3JJpvdTXK+fvZzb2Zmzz4DZPfZU56jCCEEkiRJkiRJp1FDHYAkSZIkSeFHJgiSJEmSJHmRCYIkSZIkSV5kgiBJkiRJkheZIEiSJEmS5EUmCJIkSZIkeTGVdsG6deu49tprfZ5bsmQJzZo1K/559erVvPTSS2zbto2oqCguvPBC7rnnHmJjYwMXsSRJkiRJla7UBOGke+65h27dunkcS01NLf7/69atY+LEiQwcOJDJkydz5MgRZsyYQXp6Oh9++CGqKjsrJEmSJKmqKHOC0KRJEzp16lTi+eeee44WLVrw4osvFicDycnJ3HDDDSxdupQhQ4ZUOFhJkiRJkoIjIF/rDx8+zKZNmxg2bJhHT0Hv3r2pU6cOy5Yt86s9l8vFvn37cLlcgQhPkiRJqgD5nlwzlbkH4ZFHHuGOO+4gIiKCrl27MmnSJNq1awdAeno6AC1atPB6XlpaGhkZGX4FtX//fgYNGsS8efNISUnx67mSJElSYB06dIixY8fyzTff0KhRo4C2nZ2dTV5ent/Pi46OJj4+PqCxSJ5KTRBiYmK47rrr6N69O/Hx8ezYsYM333yTMWPG8MEHH9CxY0eys7MBiIuL83p+XFwcW7du9SuozMxMAMaOHevX8yRJkqTKk5mZGdAEITs7mwsvGMiJXP8ThLi4OL755huZJFSiUhOENm3a0KZNm+Kfu3btyoABA7j00kuZOXMmc+bMKT6nKIrPNko6XpLk5GQA2YMgSZIUBk72IJx8bw6UvLw8TuTm8f6r06mTnFTm5x3OPMq1t91PXl6eTBAqUZmHGE6XnJxMnz59WLlyJUDxX9DJnoTT5eTk+OxZOBtN0wBISUnxWCkhSZIkhc7J9+ZAq5OUSP0UP5IPYVRKHJKnck9SNIxTf0En5x74mmuQnp7uc26CJEmSJAFgGP4/pEpXrgQhMzOTn376qXjZY0pKCu3ateOLL77wSBzWrFnD4cOHGTRoUECClSRJkqofgUAIo+wPRKhDrhFKHWKYMmUKDRo0oG3btsTGxrJz507eeustioqKuPvuu4uvu+eee5gwYQJ33303V155JYcPH2bGjBl07NiRiy66qFJvQpIkSarC/O0VkD0IQVFqgtCyZUu++uorPvjgAwoLC4mPj6d79+7ccsstpKWlFV/Xs2dPZs+ezaxZs5g4cSJRUVFccMEF3HvvvZU2biVJkiRVA8Lwb16BnIMQFKUmCBMnTmTixIllaqxv37707du3wkFJkiRJNYhhgKH7d71U6cq1ikGSwpFRlIt+MANRlIcaVwetTlMUzRzqsCRJKo3sQQhLMkGQqgU9+xDOP7/755uFwMg+gH7gTywdL0YxW0MdniRJZyP8nIMgE4SgkFssViPCacc4dhBRVBDqUIJKCIEz46d/uij/md1s6AhHIa59m0MamyRJpfNrBcM/D6nyyR6EakAIgfPnJbh+Xw6KCoaOltYVy/lXoWjV/69Y2PPB5fBxwkA/thdzky7BD0qSpLIzhJ+rGOQyx2Co/p8eNYBr62pcv68Al7P4mJ7xK06zDct5o0IYWXAoqgbC9xtGTUiQJKnKk3MQwpIcYqgGXL8t9/4G7XLi2roaofsxMziEjMJcXEd34zq+F+G0+/VcxRKBElMLOGPPD1VDS0nz+RxJksKIofv/kCqd/HpVDYjCXN8nDMOdOGgRwQ3ID0IIXIfSMXIOub8VKCr6kZ2Y6rVGi61d5nYsaX1wbFmBsBecbBg1qRFaiizzLUlhT/YghCWZIFQDap3GGPu2ex1XImPBYgtBRGUn8rNOJQdQ/L+uA3+iRieiqGX7J6pYI7F0vhSRm4lwFKJEJ6LaYior7DIRTgeurWsx9u9ASaqLuX0flIjokMYkSWFJ+DkHoYQhRSmwZIJQDVh6jaBowQugO0/94pjMmPuO9nur7WDTT08OTqcoGPlZaDFl3+FNURQUP3odKpPIP0Hhu48gCnLBaQeTBef3C4gY/whqstyhVJI8yB6EsCTnIFQDanIqtivuRWvWGSU2CbVhG6xDb8fUpH2oQ6uxHKvmI3Kz3MkBuId67AXYF78Z2sAkSZLKSPYgVBNqYl2sg28IdRh+0+JTMHIzvb8RCIEalRiaoALAtf1XnxOpjMN/I+yFKNbwnRciSUEnN2sKSzJBkEJKiUxAjUvxmKQIYKrf1r18sYpSVK3kDWlV2XEnSadzFz8q+8oEWSgpOGSCIIWUoiiY67bESKiPkXcMRdVQY2ujmCyhDq1CtI59ca1f6lGbAkVFbdJWln6WpDPJOQhhSSYIUlhQbdGotuozw99y3giMfRkYB3e538xUDSUqDuulN4U6NEkKP3IvhrAkEwRJqgSK2YJt3IMY+//COLwHJT4ZrUk7FDm8IEneZA9CWKr2CYL4Z9lfuC/3k6ofRVHQUlugpcpiTZJ0VobhX3VEOUkxKKptglC0Yx97H3qDvDWbUSwmEkf2o/7DN6BFhXb2uHC5QFFQtKo7AU+SJCmgZA9CWKqWCYLzWA7pw+5Dzy0AIRBFDo5/uoqijL2kffZMSGIyjh+h8P2Z6H/+Dihobc4h4tq7UBOSQhKPJElS2JDLHMNStRwQPfbfbzAcTo9ynMLhpGDzTgo27Qh6PMLpIP/fd7qTg3+60vStv5I/7U7E6bPcJUmSaqKTPQj+PKRKVy0ThMItOxFFDq/jiqpStGNf0ONx/fYjoqjAM+s1DERBHq4Na4IejyRJUlg5uRdDWR9yL4agqJYJQmT75ig273X0wjCwtWgQ9HiMw/vBXuh9wmF3n5MkSarJ/EkO/B2OkMqtWiYIta66ENVmgdNWLihWM5EdmxPZtmnQ41HrNwZfpXUtVtTUxsEOR5KkGsTlcrF//0GKiopCHUqJhND9fkiVr1omCKbEWFounkFM305g0lCjbNS66kKazXkkNPF07IkSlwDaaXNCNRNqQhKm9t1DEtPZCEch+vH9GAU5oQ7FixACI/84etZ+jKK8UIcjSWHt9dlzSKnXgVZt+pBcpy333vc4uh6GH66Gn0MMhhxiCIZquYoBwNq4Ls3nPhbqMABQTCaiHngZ+6dv4fz1fwCYu52P7fIbw2q/ASEEzvQ16Hv+AFUDQ0eNT8HSaUhYlAcWTjuOHesQzn+Ga4RAja2NuVEnFCUwua5wFoFqQtGq7a+GVEN88sli7p/6FAUFp4Y333hzLpqm8cy0h0IYmQ9huMxx2bJlfP3112zatInMzEySkpLo1q0bkyZNIjW19C3bp06dysKFC72Od+zYkY8//rgyQg44+S4YJGp0LBHjpxAxfkqoQymRvn8b+t6N7oIl/xQtMbIO4ti8AmvnISGODhx7/kDY8+G0bZCME0dwZe7CXLtZhdrWj+3D+dtXiLwsUBS0ummYz7kYxWyrYNSSFBpPPDXTIzkAKCgo5LXX5/DkE/djNptDFJkPYVhq+e233yYpKYnbbruN1NRU9u/fz+uvv87IkSP57LPPaNCg9PlskZGRvPfeex7HoqKiKivkgJMJglTM9ffvoLs8DwoDI/NvhNMe0l4EoTsR+cfgzD0ShYF+bE+FEgQjPwvHj/8F/Z8lpwL0g+kYP+ViO//a8gctSSF04MBBn8ddLp0TJ3KpVSuMtlMPwx6E2bNnU6tWLY9jXbt25cILL2TevHlMnTq11DY0TaNTp06VFGHlq5ZzEKTyEU677xOKgnB5LxsNKsMASiiX7U+JVh9cf/3i3YahI7IPYeQcqVDbkhQqnTq283k8IT6WhIT44AZTBZ2ZHAA0aNCAhIQEDh06FIKIgk/2IEjFtFoN0A9sx+tbusmCEuqdFk0WFLMN4Sg444SCFlenQk2L3KO+v5EoGqIgB+JqV6h9SQqFadMe5IILr/AYZoiMjGD6Mw+jhtumYScnKfpzPfj8oI6NjSU2NjZQkXlIT0/n+PHjtGhRtv1VCgoK6NWrF1lZWaSkpDB48GAmTZpUZYYZZIIgFTM174GeuRtcjn8+MBVQNSxt+4d8sytFUTA37Ihj53p3kRRhgKKByYypTlqF2lZrpWIc3eOzF0GJTa5Q29WJUWjnxPe/I1w6MX06YoqvPttzV0fdu3Vm5YrPePiR6Wz4fTONGqXyyCNTGHLxwFCH5q2cQwxjx471OnX77bczadKkQEVWzOFw8OCDDxIfH8+YMWNKvb5Vq1a0atWKtLQ0dF3np59+Yu7cufzyyy/897//Da85ICWQCYJUTI2IwdZ7DM7dGzCyDqBGxGJq0hk1Njy+QatRCVhb9sV1bA/Cno8alYiWmFrhFQempufg2vELOIoo7j3RTGh101Cj4iscd3Vw4ocN7Jw4zV1bRIDQdRo8fTNJoy8IdWjSWXTt0pGvv/ow1GGUrpx7McybN4+UlBSPU756D9atW8e115ZtPtGaNWtITPScn6HrOvfddx9//vknb7zxhtd5X8aPH+/x83nnnUeTJk14+OGHWbJkCcOGDStTPKEkEwTJg2KNwtKyd6jDKJFiicBct2Vg27RGYR1wA67Nq9AP7wCTBVPTLphanBvQ1ymvUG9Zrp/IZ+eN/8Yo9JyjsvfBN4ju1gZbk3ohiUuqRsq5iiElJaVMSw6bNm3KtGnTytR0dLRnz5hhGDzwwAN8++23zJw5k969y//+OHToUB599FE2bNggEwRJqirUyDgs3YeHOgwPoqiQwo9ex/nTcnA50dLaEzHuTrT6jYIaR/a360H1Tk6Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 576x360 with 2 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"fig, ax = plt.subplots(figsize = (8, 5))\n",
|
||||
"f = plt.scatter(*lonlat[keep_idx].T, c=mtgp_magpie[\"covar\"][0][15].log())\n",
|
||||
"plt.colorbar(f, label = \"Log Covariance\")\n",
|
||||
"plt.savefig(\"mt_wind_usa.pdf\", bbox_inches = \"tight\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "25bf1a12",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.9.5"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,692 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"id": "bad289dd-bf88-4240-a981-c8aa438e8385",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Warning no robinhood utils.\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import voltron \n",
|
||||
"import seaborn as sns\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import torch\n",
|
||||
"import pandas as pd\n",
|
||||
"import pickle as pkl\n",
|
||||
"import numpy as np\n",
|
||||
"\n",
|
||||
"import datetime\n",
|
||||
"import gpytorch\n",
|
||||
"from botorch.models import SingleTaskGP\n",
|
||||
"from botorch.optim.fit import fit_gpytorch_torch\n",
|
||||
"from gpytorch.likelihoods import GaussianLikelihood\n",
|
||||
"from gpytorch.mlls import ExactMarginalLogLikelihood\n",
|
||||
"from gpytorch.means import ConstantMean, LinearMean, Mean\n",
|
||||
"from gpytorch.kernels import SpectralMixtureKernel, MaternKernel, RBFKernel, ScaleKernel\n",
|
||||
"from voltron.train_utils import LearnGPCV, TrainVolModel, TrainVoltMagpieModel, TrainBasicModel\n",
|
||||
"from voltron.models import VoltMagpie\n",
|
||||
"from voltron.means import LogLinearMean, EWMAMean\n",
|
||||
"\n",
|
||||
"from voltron.rollout_utils import GeneratePrediction, Rollouts\n",
|
||||
"from voltron.data import make_ticker_list, DataGetter, GetStockHistory\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"id": "1f101591-6812-48a7-b025-daac0dbd0c81",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from voltron.likelihoods import VolatilityGaussianLikelihood\n",
|
||||
"from voltron.kernels import BMKernel, VolatilityKernel\n",
|
||||
"from voltron.models import BMGP, VoltronGP\n",
|
||||
"from gpytorch.kernels import ScaleKernel, RBFKernel, MaternKernel\n",
|
||||
"from voltron.models import MultitaskVariationalGP\n",
|
||||
"from gpytorch.priors import LKJCovariancePrior, SmoothedBoxPrior\n",
|
||||
"from voltron.models import MultitaskBMGP\n",
|
||||
"from gpytorch.likelihoods import MultitaskGaussianLikelihood\n",
|
||||
"\n",
|
||||
"def get_and_fit_volmodel(train_x, pred_scale, train_iter=200,\n",
|
||||
" printing=False):\n",
|
||||
" prior = LKJCovariancePrior(eta=5.0, n=pred_scale.shape[-1], sd_prior=SmoothedBoxPrior(0.05, 1.0))\n",
|
||||
"\n",
|
||||
" vol_lh = MultitaskGaussianLikelihood(num_tasks=pred_scale.shape[-1])\n",
|
||||
" vol_lh.noise.data = torch.tensor([1e-6])\n",
|
||||
" vol_model = MultitaskBMGP(train_x, pred_scale.log(), vol_lh, prior=prior).to(train_x.device)\n",
|
||||
"\n",
|
||||
" optimizer = torch.optim.Adam([\n",
|
||||
" {'params': vol_model.parameters()}, # Includes GaussianLikelihood parameters\n",
|
||||
" ], lr=0.01)\n",
|
||||
"\n",
|
||||
" # \"Loss\" for GPs - the marginal log likelihood\n",
|
||||
" mll = gpytorch.mlls.ExactMarginalLogLikelihood(vol_lh, vol_model)\n",
|
||||
"\n",
|
||||
" for i in range(train_iter):\n",
|
||||
" # Zero gradients from previous iteration\n",
|
||||
" optimizer.zero_grad()\n",
|
||||
" # Output from model\n",
|
||||
" output = vol_model(train_x)\n",
|
||||
" # Calc loss and backprop gradients\n",
|
||||
" loss = -mll(output, pred_scale.log())\n",
|
||||
" loss.backward()\n",
|
||||
" if printing:\n",
|
||||
" if i % 50 == 0:\n",
|
||||
" print(loss.item(), vol_model.covar_module.data_covar_module.raw_vol.item())\n",
|
||||
" optimizer.step()\n",
|
||||
" \n",
|
||||
" \n",
|
||||
" return vol_model, vol_lh\n",
|
||||
"\n",
|
||||
"def get_and_fit_mtgpcv(train_x, log_returns, train_iter=200, printing=False):\n",
|
||||
" likelihood = VolatilityGaussianLikelihood(batch_shape=[log_returns.shape[0]], param=\"exp\")\n",
|
||||
" dt = train_x[1] - train_x[0]\n",
|
||||
" # corresponds to ICM\n",
|
||||
" model = MultitaskVariationalGP(\n",
|
||||
" inducing_points=train_x, \n",
|
||||
" covar_module=BMKernel().to(train_x.device), learn_inducing_locations=False,\n",
|
||||
" num_tasks = log_returns.shape[0], \n",
|
||||
" prior=LKJCovariancePrior(eta=5.0, n=log_returns.shape[0], sd_prior=SmoothedBoxPrior(0.05, 1.0))\n",
|
||||
" )\n",
|
||||
" model = model.to(train_x.device)\n",
|
||||
" model.initialize_variational_parameters(likelihood=likelihood, x=train_x, y=log_returns.t())\n",
|
||||
" \n",
|
||||
" model = model.to(train_x.device)\n",
|
||||
" likelihood = likelihood.to(train_x.device)\n",
|
||||
"\n",
|
||||
"\n",
|
||||
" # Find optimal model hyperparameters\n",
|
||||
" model.train()\n",
|
||||
" likelihood.train()\n",
|
||||
"\n",
|
||||
" # Use the adam optimizer\n",
|
||||
" optimizer = torch.optim.Adam([\n",
|
||||
" {\"params\": model.parameters()}, \n",
|
||||
" ], lr=0.01)\n",
|
||||
"\n",
|
||||
" # \"Loss\" for GPs - the marginal log likelihood\n",
|
||||
" # num_data refers to the number of training datapoints\n",
|
||||
" mll = gpytorch.mlls.VariationalELBO(likelihood, model, train_x.shape[0])\n",
|
||||
" \n",
|
||||
" batched_train_x = train_x#[:-1]\n",
|
||||
" \n",
|
||||
" print_every = 50\n",
|
||||
" for i in range(train_iter):\n",
|
||||
" # Zero backpropped gradients from previous iteration\n",
|
||||
" optimizer.zero_grad()\n",
|
||||
" # Get predictive output\n",
|
||||
" output = model(batched_train_x)\n",
|
||||
" # Calc loss and backprop gradients\n",
|
||||
" loss = -mll(output, log_returns.t())\n",
|
||||
" loss.backward()\n",
|
||||
" if printing:\n",
|
||||
" if i % print_every == 0:\n",
|
||||
" print('Iter %d/%d - Loss: %.3f' % (i + 1, train_iter, loss.item()))\n",
|
||||
" optimizer.step()\n",
|
||||
" \n",
|
||||
" model.eval();\n",
|
||||
" likelihood.eval();\n",
|
||||
" predictive = model(train_x)\n",
|
||||
" # pred_scale = likelihood(predictive).scale.mean(0).detach()\n",
|
||||
" samples = likelihood(predictive).scale.detach()\n",
|
||||
" return samples.mean(0) / dt**0.5\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"def get_and_fit_datamods(train_x, train_y, vols, vmod, vlh,\n",
|
||||
" train_iter=200, printing=False, k=200, mean_func='ewma'):\n",
|
||||
" model_list = []\n",
|
||||
" nmodel = vols.shape[-1]\n",
|
||||
" for mdl_idx in range(nmodel):\n",
|
||||
" mod = TrainVoltMagpieModel2(train_x, train_y[mdl_idx, :], vmod, vlh, vols[:, mdl_idx],\n",
|
||||
" train_iters=train_iter, k=k, mean_func=mean_func)\n",
|
||||
" model_list.append(mod)\n",
|
||||
" \n",
|
||||
" return model_list\n",
|
||||
" \n",
|
||||
"\n",
|
||||
"def TrainVoltMagpieModel2(train_x, train_y, vol_model, vol_lh, vol_path,\n",
|
||||
" train_iters=1000, printing=False, k=25,\n",
|
||||
" mean_func=\"ewma\"):\n",
|
||||
" \n",
|
||||
" voltron_lh = gpytorch.likelihoods.GaussianLikelihood().to(train_x.device)\n",
|
||||
" voltron = VoltMagpie(train_x, train_y.log(),\n",
|
||||
" voltron_lh, vol_path, k=k).to(train_x.device)\n",
|
||||
" \n",
|
||||
" if mean_func.lower() in [\"ewma\", \"dewma\", \"tewma\", \"meanrevert\"]:\n",
|
||||
" # default voltmagpie is an ewma mean so we don't need to redefine anything\n",
|
||||
" grad_flags = [True, False, False, False]\n",
|
||||
" \n",
|
||||
" if mean_func.lower() == \"dewma\":\n",
|
||||
" voltron.mean_module = DEWMAMean(train_x, train_y.log(), k).to(train_x.device)\n",
|
||||
" elif mean_func.lower() == 'tewma':\n",
|
||||
" voltron.mean_module = TEWMAMean(train_x, train_y.log(), k).to(train_x.device)\n",
|
||||
" \n",
|
||||
" elif mean_func.lower()=='constant':\n",
|
||||
" voltron.mean_module = gpytorch.means.ConstantMean().to(train_x.device)\n",
|
||||
" grad_flags = [True, True, False, False, False]\n",
|
||||
" elif mean_func.lower()=='loglinear':\n",
|
||||
" voltron.mean_module = LogLinearMean(1).to(train_x.device)\n",
|
||||
" voltron.mean_module.initialize_from_data(train_x, train_y.log())\n",
|
||||
" grad_flags = [True, True, True, False, False, False]\n",
|
||||
" elif mean_func.lower()=='linear':\n",
|
||||
" voltron.mean_module = gpytorch.means.LinearMean(1).to(train_x.device)\n",
|
||||
" grad_flags = [True, True, True, False, False, False]\n",
|
||||
"\n",
|
||||
" voltron.likelihood.raw_noise.data = torch.tensor([1e-5]).to(train_x.device)\n",
|
||||
"# voltron.vol_lh = vol_lh.to(train_x.device)\n",
|
||||
"# voltron.vol_model = vol_model.to(train_x.device)\n",
|
||||
" \n",
|
||||
" for idx, p in enumerate(voltron.parameters()):\n",
|
||||
" p.requires_grad = grad_flags[idx]\n",
|
||||
"\n",
|
||||
" voltron.train();\n",
|
||||
" voltron_lh.train();\n",
|
||||
" voltron.vol_lh.train();\n",
|
||||
" voltron.vol_model.train();\n",
|
||||
"\n",
|
||||
" # Use the adam optimizer\n",
|
||||
" optimizer = torch.optim.Adam([\n",
|
||||
" {'params': voltron.parameters()}, # Includes GaussianLikelihood parameters\n",
|
||||
" ], lr=0.1)\n",
|
||||
"\n",
|
||||
" # \"Loss\" for GPs - the marginal log likelihood\n",
|
||||
" mll = gpytorch.mlls.ExactMarginalLogLikelihood(voltron_lh, voltron)\n",
|
||||
"\n",
|
||||
" print_every = 50\n",
|
||||
" for i in range(train_iters):\n",
|
||||
" # Zero gradients from previous iteration\n",
|
||||
" optimizer.zero_grad()\n",
|
||||
" # Output from model\n",
|
||||
" output = voltron(train_x)\n",
|
||||
" # Calc loss and backprop gradients\n",
|
||||
" loss = -mll(output, train_y.log())\n",
|
||||
" loss.backward()\n",
|
||||
" if printing:\n",
|
||||
" if i % print_every == 0:\n",
|
||||
" print('Iter %d/%d - Loss: %.3f' % (i + 1, train_iters, loss.item()))\n",
|
||||
" optimizer.step()\n",
|
||||
" \n",
|
||||
" \n",
|
||||
" return voltron, voltron_lh"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "68e7e0da-c65e-4943-b613-5d3ff02ab7d7",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Load In, Set Up Data"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 24,
|
||||
"id": "2901a851-07fc-4877-bb4e-319acb740ea4",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"stn_names, stn_lonlat, stn_data = pkl.load(open(\"./wind_data.p\", 'rb'))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 46,
|
||||
"id": "718c74a5-aed7-4efb-930c-e164227540ac",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"lonlat = np.array(list(stn_lonlat.values()))\n",
|
||||
"full_dat = np.stack(list(stn_data.values()), -1).T\n",
|
||||
"full_dat[full_dat == -99.0] = 0.\n",
|
||||
"names = np.array(list(stn_names.values()))\n",
|
||||
"\n",
|
||||
"conus_idx = np.where(lonlat[:, 0] > -128)[0]\n",
|
||||
"lonlat = lonlat[conus_idx]\n",
|
||||
"full_dat = full_dat[conus_idx]\n",
|
||||
"names = names[conus_idx]\n",
|
||||
"\n",
|
||||
"keep_idx = np.where(full_dat.mean(-1) > 0.)[0]\n",
|
||||
"lonlat = lonlat[keep_idx]\n",
|
||||
"full_dat = full_dat[keep_idx]\n",
|
||||
"names = names[keep_idx]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 47,
|
||||
"id": "c6465482-f591-48aa-a320-02485165e410",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"full_dat += 1.\n",
|
||||
"returns = np.log(full_dat[..., 1:]/full_dat[...,:-1])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 48,
|
||||
"id": "4fbcc881-0bbe-46de-960b-4dba6883b1b2",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"<matplotlib.collections.PathCollection at 0x7f9c006ae1f0>"
|
||||
]
|
||||
},
|
||||
"execution_count": 48,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAXAAAAD7CAYAAABzGc+QAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMSwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Z1A+gAAAACXBIWXMAAAsTAAALEwEAmpwYAABa/klEQVR4nO2dd3hUVfrHP2d6eiEhoffeJfQuoICo2LAjdte1rLr28rOtu+5ad+0dKwiIAtKkKR1C74QOCSUhpE+7M+f3R0IkZAIpU5PzeR4eknNzz/3eZOadc9/zFiGlRKFQKBShhy7QAhQKhUJRPZQBVygUihBFGXCFQqEIUZQBVygUihBFGXCFQqEIUZQBVygUihDFUJkfEkIcBPIBF6BJKVOEEPHAFKA5cBAYL6U87RuZCoVCoTiXqqzAh0kpu0spU0q+fwpYJKVsAywq+V6hUCgUfkJUJpGnZAWeIqXMOmtsNzBUSnlMCNEAWCqlbHe+eRISEmTz5s1rplihUCjqGOvXr8+SUiaeO14pFwoggQVCCAl8LKX8BEiSUh4rOX4cSPJ0ohDiHuAegKZNm5Kamlpl8QqFQlGXEUIc8jReWQM+UEqZLoSoD/wmhNh19kEppSwx7uUoMfafAKSkpKi8fYVCofASlfKBSynTS/4/CcwAegMnSlwnlPx/0lciFQqFQlGeCxpwIUSEECLqzNfAJcA2YCZwW8mP3Qb84iuRCoVCoShPZVwoScAMIcSZn/9eSjlPCLEO+FEIcSdwCBjvO5kKhUKhOJcLGnAp5X6gm4fxU8BwX4hS1B1O24tIy8ukYXgMjSNiAy1HoQgpKruJqVB4FSkl/9q8iG/3pmLS63G4XPRKbMJ7/a8l0mgOtDyFIiRQqfRBhM3lZHXmATaeOoK7ljfa+HH/Rr7ftx67WyPfacfu1libeZhnUmcHWppCETKoFXiQMPfodp5ZPwudEEgkEQYTn/S/iQ6xyYGW5hM+37MGq8tZZszhdrEwfQ+FTgcRRlOAlCkUoUOtXoE7HRpLJq/kfw99ydS3Z5OTmRdoSR45kH+Kp9b/QpHLQYFmp1BzcNJWwMTl3+BwuwIt77w4XblYtXSkdFfpvFyH1eO4AAo1uxeUKRS1n1q7Ai/MK+LhQf/HySNZ2ArsmMJMfPfqT7w+/1napbQKtLwyTD+0Ec1d3gA63S6WHd/L8IbnrVAQEJyuXHZkPUG2dQVC6DGISNonvEpC+LBKnd8/qQW/Ht6Bm7KuojhzOImWSF9IVihqHbV2BT759V84tv8ktoLi1ZzD6qAo38a/bnufYGvkfMpWiOZhBSsl5Do9r1QDzZaTfyHbugKJE7e04XBnsS3zb+Q7dl34ZODRzkOJNJox6vQA6BBY9EZeTbmMkpBVv2G1Ofll4WZe+WAu3/y8ltN5RX69vkJRXWrtCnzplFU47c5y45lHTpGVnk1i43oBUOWZYQ3aMi99B0Xn+IRd0k2fxOaBEXUeCp37yXdsR1JWr1s6OJL7JR0TX7/gHE0i45hz6T18uWcNazMP0zwqnrva9aNTnH99/tm5Rdzx9LfkFdiw2Z2YjHq+/nkN7794PW2b1/erFoWiqtRaA24web41KSUGY3Dd9sUN2tExtgHbc46VbuyF6Y3c1DKFRuGxgRXnAbt2HIERsJ1zxI1VO1zpeZLDo3m6+0ivaqsqH09exqmcQlyu4icgh9OFw+niHx/MY9K/JwRUm0JxIYLLknmR0XcO45uXp2O3OkrHdDpBy67NiEuKCaCy8hh0Or4cdCu/HN7C7CNbCdObuLFlTwYntQ60NI9EmtojZfmNRoGJOEvfACiqPn+s3VtqvM/mwNFT5BfaiIqwBECVQlE5aq0Bv/qh0Wz+fSdb/tiBlGAw6AiLCuOZbx8MtDSPmHR6rmveg+ua9wi0lAti0sfTKPoW0vN/wC3P+OgNGHQRNI6+NaDaqorRqK/wmEFf8TGFIhiotQbcYDTwj5lPsGf9fnan7iOhUTy9R3VHb1BvSm/QOu4JIk1tOZz7JZo7h/iwwbSMfRCTPj7Q0qrE5Rd34buZ63A4/wzX1OsFKV2aEWYxBlCZQnFhKtWRx1ukpKRI1dBBEUw4nBpP/PtntuxKBwE6IUiIj+SD/7ue+NiIQMtTKAAQQqw/q51lKSGzAnfYncz/aimLf1iB0WLksruGM/iaPn4POVPULkxGA+88ey27958g7dBJGiTG0KNjE3Q69bpSBD8hYcBdLjdPjX6NvZsOYi8q3pTcvW4fGxdv428f3OVXLfsPZfLfTxezZcdRwiwmrhrTnYk39MegXDMhTbuWSbRr6bEroEIRtIREIs+aXzewb/OhUuMNYCu0s+j75RzZk+E3HScy87j/ie/ZsOUwmuYmv8DGlJ9T+cfbc/ymQaFQKM4QEgY8deEWbIUewtZ0gq1/VC7zzxtMnZmKw6mVGbM7NJatTuNEkNZZUSgUtZeQMODxSTEeE3N0eh3RCVF+07Ez7TiaVj5m2Gg0cPhott90KBSK0ELTXD4p4RESBnzkLYPR68tLNRgN9B7d3W862rSo71GHU9No3DDObzrOR7bVyty0PSw7dMhjgaxAkOuwMuPQZibvX8+xotxAy1Eo/Mbcpdu44u4PGXrj24y98wN+mrfRq4Y8JDYxk5ol8vzkh/nXbR/gdrlwuyVR8ZG8NP0xTGb/xeped0VP5izaViZzz2Qy0KtHcxoEQXbnJ+vX8daqlZh0ehBg1uuZdNU1dEwMXE2PJcf28MjaaegQuIF/bpnPXzsM4Z52A2o0r91lY0feRhxuO+2iuhJrCq34c0XtZ+Hynfznk4XYHcVu15w8K+9/8zsguHpUd69cI6TiwDWnRtqGAxjNRlp1axaQEMJdacd566Pf2L33OCajgbGXdOW+iUMwV1B7xV+kZqRz24zpWLWyPvqE8HBW3XkPep3/H7bynTYGzXkLm6usJovewPdDbqdjbINqzZuWv53PD/wHAInELd1cmnwtI5KurLFmhcJbjH/gM9KP55Qbj4sJZ/bn91dprpCPA4dil0mHPm0CqqF9m2Q+efNWXC43Op0Imjj0H7ZuwXaO8QawOjXWZaTTt3ETv2taejwNvSj/weFwuZh1eGu1DLjDbefzA//B7i5bSGvB8em0jepM0/DgqvWuqLucyPIc2JCTW4SmubwSehwSPvBgRK/XBY3xBsiz2/H0LCUEFDocHo74Hs3t9ujvk8hqdxranbeF4r4951xLOll7amm15lQofEGjpFiP4wnxkV7LG1EGvJYwpk1bwo3l9wOcLhe9GjUKgCIYlNQKlwcDbtEbubRRh2rN6ZRO8PBRJZE4ZGA+qAKNy+1mydo9vPDer/z7i4Xs2n8i0JIUwF9vLe9atZgN3HvTIK9dQxnwWsLYtu3omFi/1IjrhMBiMPDs4KFEmwNTEjXBEskTXUZi1hkwCB2C4jrnlzXuTK+EZtWas21UZ1yy/OrdpDPTPTa0Stl6A83l5tHXf+KVD+fx28pd/LJoC/e9PJmp8zcEWlqdZ0BKK15+9HKaN47HaNDRuEEsT98/itFDO3ntGiG1iak4P06Xi7l705iXlkZsmIUbO3elS1Lg08P352cx6/BWbC6NkY3a0SO+SY3cTyuyFvBL+re4pIYbNyadmQ5R3ZnQ/GF0HnzuvkRKyYolu5g1bR02q4Ohl3RmzFU9MfupkuHiNXt49aN5WM/pPmUy6pn53r3ERIX5RYfCt1S0iakMuCIkOWY9wrrsP7C7rXSJ6UW7qK4B2ZP48M15zPt5AzZbsQE1m400aZHAO1/ced5a497iuf/OYtHqPeXGwy0mnr33Ui7u09bnGhS+p1ZEoSgUZ2gQ1oQrGt0cUA0njuXw60+pOB1/unTsdidHD2WxbNEOLh7VxecaIsLMCCHKbRYLgapnXgdQPnBFSKO5HRwoWM+BgvVobv9uYm7deMhjNIHN6mTdijS/aLhiWBdMHlb6ep2OlE5N/aJBETjUClwRshwoWM8vR187a0RyeaOnaBXV2y/Xj6mg4YPeoCPeTzV6OrVuwL3jB/DRlOWlHyZ6neCtJ6/GqEoc13qUAVf4nONF+Xy5M5Utp47RMS6J2zuk0DiyZqUHirRcZhx5Be2c5sq/HH2Ne9t8SYTB97VpLurdEkuYCZvVwdkeDINex5irLvL59c9w45gURg3syPodRwg3G+nVpZky3nUE5UJR+JS0nCxG/PIpX+xMZdXxw3y9ez2XzvyMbaeO12je3XnLKjgi2ZX7R43mrix6g45/f3QbyQ3jsIQZCY8wExFp5qlXr6FR03p+0XCGuOhwRvRtR/8eLZXxrkOoFbjCp7y49jcKnY7S1Bun243T7eb5NQuYMWZCted1uItwy/KlA1zSid1dVO15q0rT5gl8OeNBDuw9id3mpHX7Bn6JPlEoQBnwoGRPZhZTN28n32ZjRNvWDGvdIiDFqLzB2pNHPKb4b8rKwOV2V/u+mkVcxArxPe5zknr0wkSLSP+5LwCEELRsE/h4e0XdQxnwIOPHTVt55belOF0uXFIyd1caKU0a8cl1V4akEQ83mMh12MqNm/UGdDWI204Oa02H6CHsyvsDpyye3ygstInuT4OwdtWeV6EIJSptwIUQeiAVSJdSjhVCfAUMAc5U6J8opdzkdYV1iHybnZd/W4r9rKqCRU4nqUfSWbB7L6M7hF5Sxo1tuvHVrvVlSsqadXqubd2lxok3oxr+jTbR/dmW8xsAnWJH0DqyT43mrAtIKdmx7zibdh0lPiaCob3aqJhxL7FtTwazF2/DancyvF9bBqa0RqfzXYJZVVbgDwM7geizxh6XUk7zrqS6y5rDRzHqdJzb/bPI6eTXnbtD0oA/2n0wB/NPsyR9PyadHqfbRb/kpjzX8+Iazy2EoHVUH1pHKaNdWTSXm2ffmcnabYfRNBcmo563Ji3mveeuo11z5QaqCd/MWMtX01Zhd2pICSvW7aNH5ya8/sQ4nxnxShlwIURLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {
|
||||
"needs_background": "light"
|
||||
},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.scatter(lonlat[:, 0], lonlat[:, 1], c=full_dat.mean(-1))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 51,
|
||||
"id": "0b555dc8-209f-4231-b485-df5c9f8507eb",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"64"
|
||||
]
|
||||
},
|
||||
"execution_count": 51,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"np.argmax(full_dat.max(1))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 52,
|
||||
"id": "4df7026b-64c9-49c4-9986-0d3bda5e78f6",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"array([-101.44, 42.07])"
|
||||
]
|
||||
},
|
||||
"execution_count": 52,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"lonlat[64]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 53,
|
||||
"id": "5afac95d-25e7-4845-98a9-a993183aa7f5",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"'NE_Whitman_5_ENE'"
|
||||
]
|
||||
},
|
||||
"execution_count": 53,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"names[64]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "9d34b1b4-2d9c-40c9-91bf-a54306935223",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## subset just for testing"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 9,
|
||||
"id": "b7f1b4fe-19b3-4cd8-adeb-bd136cad7f8a",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"lon_keepers = np.logical_and(lonlat[:, 0] > -120, lonlat[:, 0] < -110)\n",
|
||||
"lat_keepers = np.logical_and(lonlat[:, 1] > 30, lonlat[:, 1] < 40)\n",
|
||||
"keepers = np.where(np.logical_and(lon_keepers, lat_keepers))[0]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 10,
|
||||
"id": "1bbe2fc3-f5ae-4548-9d05-8369690a6d49",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"<matplotlib.collections.PathCollection at 0x7f9bfa9c9ac0>"
|
||||
]
|
||||
},
|
||||
"execution_count": 10,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAXAAAAD4CAYAAAD1jb0+AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMSwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/Z1A+gAAAACXBIWXMAAAsTAAALEwEAmpwYAAAYvUlEQVR4nO3de5xV9X3u8c8ze/ZcAAEvg2CI0IqKSAR0BC/HG4aUeGqaqH3F2ObSxJiY2vaYxhOtSY3Hpm0abY4maRp6PGpaY0ONmgZDEpOgSKKSwQCKUhTxbmRQBwSGuX77x14gMjPMBvbeay/meb9e++XMb+2197OXzMPiN+uiiMDMzLKnJu0AZma2d1zgZmYZ5QI3M8soF7iZWUa5wM3MMqq2km92yCGHxMSJEyv5lmZmmbds2bINEdG063hFC3zixIm0tLRU8i3NzDJP0nP9jXsKxcwso1zgZmYZ5QI3M8soF7iZWUa5wM1S9tza9Tz+m+fY1t6ZdhTLmEGPQpHUACwG6pPn3xkR10iaDVwP1AHLgE9ERHc5w5rtT1p/u5Ev/tm/8coLr5PL1dDT08unPjeXcy44Me1olhHF7IF3ALMjYhowHZgr6RTgNuDCiJgKPAd8tGwpzfYzEcEXPvMdnl+7no5tXWzd0kHHti6+ff2PWfWbfo8YM+tj0AKPgs3Jt/nk0QN0RsSaZPw+4PzyRDTb/6xb8yqvvtJGb+/bL+fc2dHFPd99OKVUljVFzYFLyklaDqynUNZLgVpJzclTLgDeWZaEZvuhTRu3UpPr++MXAW+8trmfNcz6KqrAI6InIqYD44GZwLHAhcDXJC0F3qSwV96HpEsktUhqaW1tLU1qs4w78pjD6O7q+yNTV1/LrNOPTiGRZdEeHYUSEW3AImBuRDwUEadFxEwKv+RcM8A68yKiOSKam5r6nMpvNiQNP6CBj152NvUN+R1jdfW1HNx0AOdc0LybNc3eUsxRKE1AV0S0SWoE5gBfkTQmItZLqgc+D3y5zFnN9ivnf/hUfveosdxz+8NsfGMLJ51xNOdeOIvhIxrSjmYZUczFrMYBt0nKUdhjnx8RCyR9VdLvJ2PfiohflDOo2f5oxqwjmDHriLRjWEYNWuARsRKY0c/4FcAV5QhlZmaD85mYZmYZ5QI3M8soF7iZWUa5wM3MMsoFbmaWUS5wM7OMcoGbmWWUC9zMLKNc4GZmGVXMqfRmZvudzq5uWla9QHdPDycc806GD6tPO9Iec4Gb2ZDz6JMvcMUN9+z4vrunl6sufg9zTz0mvVB7wVMoZjakbGnv5C+vv5st7Z07Hh2d3fzd//spL77alna8PeICN7Mh5cFH16J+xnt6eln44KqK59kXLnAzG1K2tnf2uRcpFKZR3tzakUKivecCN7MhZdZxE+inv2msz3P6CZMqH2gfuMDNbEh5x5jRXDh3Bg31+R1TKY31eWYdN4ETpmTr3uw+CsXMhpzPXHg6s941kQWLV9HZ1c2ckydz+gmTkPqbHa9eLnAzG5JOOPZwTjj28LRj7JNBp1AkNUhaKmmFpFWSrk3Gz5b0qKTlkpZIytbkkZlZxhUzB94BzI6IacB0YK6kk4BvAX8UEdOB7wJfKFdIMzPrq5ibGgewOfk2nzwieYxMxkcBL5cjoJmZ9a+oOXBJOWAZMAn4ZkQ8Iuli4EeS2oFNwEkDrHsJcAnA4Ydne77JzKyaFHUYYUT0JFMl44GZkqYClwPnRMR44BbgHwdYd15ENEdEc1NTU4lim5nZHh0HHhFtwCLgvcC0iHgkWfQ94JTSRjMzs90ZdApFUhPQFRFtkhqBOcBXgFGSjoqINcnYk+UK2ba5nRvnL+YXy54C4N0nHsWfX3Aao0Y0lustzcyqXjFz4OOA25J58BpgfkQskPRJ4PuSeoE3gI+XI2B3dw8f/9s7eHnDJrp7egG491dPsOKpl/nedR8hV+OTSc1saCrmKJSVwIx+xu8G7i5HqJ0tXvEMG9q27ChvKFx0Zn3bm/xy5TpOn35EuSOYmVWlqt99ffrFDWzt6Oozvq2jm6de3JBCIjOz6lD1BX742AMZ1pDvM95Qn2fCoQemkMjMrDpUfYGfdfwkhjfUkat56yIzuRoxorGOM2Z4+sTMhq6qL/D6fC23XH0RM6dMIFcjcjVi1rETuPXqD5GvzaUdz8wsNZm4GuHYgw7g65efR3d3D0jU5qr+7x0zs7LLRIFvV+s9bjOzHbwra2aWUS5wM7OMcoGbmWWUC9zMLKNc4GZmGeUCNzPLKBe4mVlGucDNzDLKBW5mllEucDOzjHKBm5lllAvczCyjirmpcQOwGKhPnn9nRFwj6UHggORpY4ClEfH+cgU1M7O3K+ZqhB3A7IjYLCkPLJG0MCJO2/4ESd8HflCukGZm1tegUyhRsDn5Np88YvtySSOB2cA95QhoZmb9K2oOXFJO0nJgPXBfRDyy0+L3Az+PiE0DrHuJpBZJLa2trfua18zMEkUVeET0RMR0YDwwU9LUnRZ/CLhjN+vOi4jmiGhuamrap7BmZvaWPToKJSLagEXAXABJhwAzgXtLnszMzHZr0AKX1CRpdPJ1IzAHWJ0svgBYEBHbypbQzMz6VcxRKOOA2yTlKBT+/IhYkCy7EPj7coUzM7OBDVrgEbESmDHAsjNLHcjMzIqTqbvSW/l0d/cQAfl8Lu0obNyyjR8+/ATrXn2dqRPGMvfEo2msy6cdy6zquMCHuNde38xXb/wxv172LBHBcVPHc8VfzOUdhx2YSp61L2/gT26YT2d3Dx1d3Sxcuppv/+ghbv/8RRw8cngqmcyqla+FMoR19/Ry2V/ezq+XraOnp5fe3mDFYy/yp5+9na3tnalkuuZff8rm9g46uroBaO/sYsOmrdx0z5JU8phVMxf4ELb018/QtrGdnp4dJ9YSEWzr6OIXDzxZ8TztHV2sfmH9W6f5Jnp6elm0Ym3F85hVOxf4EPbSK2/Q3d3TZ3zbti5eeOH1iuepqRFC/S7L16Y/N29WbVzgQ9jvThxDbW3fPwKNjXmOnHRoxfPU52s5ecoEcjVvL/G62hzvO3lKxfOYVTsX+BA2Y9rhjD/swLcdeZLL1TBq5DBOP/WoVDJd88dzGH/IaIbV52nI19JYV8vUiWP51Dknp5LHrJopYtcZx/Jpbm6OlpaWir2fDW7Llg7+5dbF/Pz+J+np7eW0U47k0584kwNHp3fER29vsHTN87y0YSNHvaOJqRPHIvU/tWI2FEhaFhHNfcZd4GZm1W2gAvcUiplZRrnAzcwyygVuZpZRLnAzs4xygZuZZZQL3Mwso1zgZmYZ5QI3M8uoYu6J2SBpqaQVklZJujYZl6QvS1oj6UlJf17+uGZmtl0xN3ToAGZHxGZJeWCJpIXAMcA7gckR0StpTDmDmpnZ2xVzT8wANiff5pNHAJcCF0VEb/K89eUKaWZmfRU1By4pJ2k5sB64LyIeAY4APiipRdJCSUcOsO4lyXNaWltbSxbczGyoK6rAI6InIqYD44GZkqYC9cC25AIr/wL8/wHWnRcRzRHR3NTUVKLYZma2R0ehREQbsAiYC7wI3JUsuhs4rqTJzMxst4o5CqVJ0ujk60ZgDrAauAc4K3naGcCa8kQ0s/1FRC+x7af0vvGn9L7xF0THA1Tyktb7m2KOQhkH3CYpR6Hw50fEAklLgNslXU7hl5wXlzGnmWVcRBBtl0PHA8DWwljH/dD4ATTqS2lGy6xijkJZCczoZ7wN+J9lyGRm+6OuFui8H2jfabAd2u8ihv8xqp2UUrDs8pmYZlYR0fEgxLZ+lvRCx5KK59kfuMDNrDJ0AIXTSHZVCxpR6TT7BRe4mVWEGs+l38pRQMN7Kp5nf+ACN7OKUG4sjLoBNKywx508NPqfUc3ItONlUjFHoZiZlURN4xyi4SHoXArUQt2JSHVpx8osF7iZVZTUCPVnpB1jv+ApFDOzjHKBm5lllAvczCyjXOBmZhnlAjczyygXuJlZRrnAzcwyygVuZpZRLnAzs4xygZuZZZQL3Mwso1zgZmYZVcxNjRskLZW0QtIqSdcm47dKWidpefKYXva0Zma2QzFXI+wAZkfEZkl5YImkhcmyKyLizvLFMzOzgRRzU+OgcNd5KNwPKQ9EOUOZmdngipoDl5STtBxYD9wXEY8ki74saaWkr0mqH2DdSyS1SGppbW0tTWozMyuuwCOiJyKmA+OBmZKmAlcBk4ETgYOAzw+w7ryIaI6I5qamptKkNjOzPTsKJSLagEXA3Ih4JQo6gFuAmWXIZ2ZmAyjmKJQmSaOTrxuBOcBqSeOSMQHvBx4vX0wzM9tVMUehjANuk5SjUPjzI2KBpF9IagIELAc+Xb6YZma2q2KOQlkJzOhnfHZZEpmZWVF8JqaZWUa5wM3MMsoFbmaWUS5wM7OMcoGbmWWUC9zMLKOKOQ48NRHBr37+BD+5s4Wenl7Oft8MzjjnOHI5/71jZlbVBX7jNXfzwL0r2NbeBcATjz7LAwtX8KVvfoTCCaBmZkNX1e7KrlvzW+5f8FZ5A2xr72Ll0nWsXPpMisnMzKpD1Rb48ofX0tvb22d829ZOli15KoVEZmbVpWoLfMTIRnK1uT7j+bocIw8clkIiM7PqUrUFfuqcY/ud566pqeGs359e+UBmZlWmagt82PB6rvv2xxg5ehjDhtcXHiPq+auvXcTBY0amHc/MLHVVfRTKscdP4LuLr+LJFS/Q29PLMdMPJ19X1ZHNzCqm6tswV5tj6gkT045hZlZ1qnYKxczMds8FbmaWUS5wM7OMKuamxg2SlkpaIWmVpGt3WX6TpM3li2hmZv0p5peYHcDsiNgsKQ8skbQwIh6W1AwcWN6IZmbWn0H3wKNg+x52PnlEcpf6rwL/u4z5zMxsAEXNgUvKSVoOrAfui4hHgMuA/4yIVwZZ9xJJLZJaWltb9zmwmZkVFFXgEdETEdOB8cBMSacDfwh8vYh150VEc0Q0NzU17VNLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {
|
||||
"needs_background": "light"
|
||||
},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.scatter(lonlat[keepers, 0], lonlat[keepers, 1], c=full_dat[keepers, 500])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "09c1471c-13ef-4e3f-99bc-e6a7098e2701",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Train "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 11,
|
||||
"id": "8c47c5e9-3c72-41ed-a409-f57e1e7d571a",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ntrain = 300\n",
|
||||
"ntest = 100\n",
|
||||
"train_x = torch.arange(ntrain).type(torch.FloatTensor).cuda()\n",
|
||||
"test_x = torch.arange(ntrain, ntrain + ntest).type(torch.FloatTensor).cuda()\n",
|
||||
"train_y = torch.FloatTensor(full_dat[keepers, :ntrain+1]).cuda()\n",
|
||||
"test_y = torch.FloatTensor(full_dat[keepers, (ntrain):(ntrain + ntest)]).cuda()\n",
|
||||
"\n",
|
||||
"returns = torch.log(train_y[:, 1:]/train_y[:, :-1])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 12,
|
||||
"id": "8770dd88-95f7-4e91-b88c-6f921041d0cf",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/home/greg_b/miniconda3/envs/rpp/lib/python3.8/site-packages/torch/functional.py:445: UserWarning: torch.meshgrid: in an upcoming release, it will be required to pass the indexing argument. (Triggered internally at /opt/conda/conda-bld/pytorch_1634272068694/work/aten/src/ATen/native/TensorShape.cpp:2157.)\n",
|
||||
" return _VF.meshgrid(tensors, **kwargs) # type: ignore[attr-defined]\n",
|
||||
"/home/greg_b/miniconda3/envs/rpp/lib/python3.8/site-packages/gpytorch/utils/cholesky.py:40: NumericalWarning: A not p.d., added jitter of 1.0e-06 to the diagonal\n",
|
||||
" warnings.warn(\n",
|
||||
"/home/greg_b/miniconda3/envs/rpp/lib/python3.8/site-packages/gpytorch/distributions/multivariate_normal.py:259: NumericalWarning: Negative variance values detected. This is likely due to numerical instabilities. Rounding negative variances up to 1e-06.\n",
|
||||
" warnings.warn(\n",
|
||||
"/home/greg_b/miniconda3/envs/rpp/lib/python3.8/site-packages/gpytorch/utils/cholesky.py:40: NumericalWarning: A not p.d., added jitter of 1.0e-08 to the diagonal\n",
|
||||
" warnings.warn(\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Iter 1/100 - Loss: 1598399.875\n",
|
||||
"Iter 51/100 - Loss: 19815.930\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/home/greg_b/miniconda3/envs/rpp/lib/python3.8/site-packages/gpytorch/utils/cholesky.py:40: NumericalWarning: A not p.d., added jitter of 1.0e-07 to the diagonal\n",
|
||||
" warnings.warn(\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"vol = get_and_fit_mtgpcv(train_x, returns, printing=True, train_iter=100)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 13,
|
||||
"id": "9321bf48-57cd-4e7e-b937-537fe211ab21",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"1.310472846031189 -1.3862943649291992\n",
|
||||
"1.0877968072891235 -1.8699604272842407\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"vmod, vlh = get_and_fit_volmodel(train_x, vol, printing=True, train_iter=100)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 19,
|
||||
"id": "326785cc-f1ec-4496-96b2-79fd8bc2af35",
|
||||
"metadata": {
|
||||
"tags": []
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"model_list = get_and_fit_datamods(train_x, train_y[:, 1:], vol, vmod, vlh,\n",
|
||||
" train_iter=200)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "76de2aa6-f867-4e96-9a6d-0cdc377d6a28",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Simulate"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 20,
|
||||
"id": "afc89d71-671f-4364-ae1f-14b14e313068",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from gpytorch.utils.cholesky import psd_safe_cholesky\n",
|
||||
"from gpytorch.utils.cholesky import psd_safe_cholesky\n",
|
||||
"\n",
|
||||
"def GeneratePrediction(train_x, train_y, test_x, pred_vol, model, latent_mean=None, theta=0.5):\n",
|
||||
" vol = model.log_vol_path.exp()\n",
|
||||
" if model.train_x.ndim != test_x.ndim:\n",
|
||||
" test_x_for_stack = test_x.unsqueeze(0).repeat(model.train_x.shape[0], 1)\n",
|
||||
" else:\n",
|
||||
" test_x_for_stack = test_x\n",
|
||||
" if vol.ndim == 1:\n",
|
||||
" vol_for_stack = vol.unsqueeze(0).repeat(pred_vol.shape[0], 1)\n",
|
||||
" else:\n",
|
||||
" vol_for_stack = vol\n",
|
||||
"\n",
|
||||
" full_x = torch.cat((model.train_x, test_x_for_stack),dim=-1)\n",
|
||||
" # print(\"vol stack = \", vol_for_stack.shape)\n",
|
||||
" # print(\"pred_vol = \", pred_vol.shape)\n",
|
||||
" full_vol = torch.cat((vol_for_stack, pred_vol),dim=-1)\n",
|
||||
"\n",
|
||||
" test_x.repeat(2, test_x.numel())\n",
|
||||
"\n",
|
||||
" idx_cut = model.train_x.shape[-1]\n",
|
||||
" \n",
|
||||
" cov_mat = model.covar_module(full_x.unsqueeze(-1), full_vol.unsqueeze(-1)).evaluate()\n",
|
||||
" K_tr = cov_mat[..., :idx_cut, :idx_cut]\n",
|
||||
" K_tr_te = cov_mat[..., :idx_cut, idx_cut:]\n",
|
||||
" K_te = cov_mat[..., idx_cut:, idx_cut:]\n",
|
||||
"\n",
|
||||
" train_mean = model.mean_module(model.train_x)\n",
|
||||
" train_diffs = model.train_y.unsqueeze(-1) - train_mean.unsqueeze(-1)\n",
|
||||
" # use psd cholesky if you must evaluate\n",
|
||||
" K_tr_chol = psd_safe_cholesky(K_tr, jitter=1e-4)\n",
|
||||
" pred_mean = K_tr_te.transpose(-1, -2).matmul(torch.cholesky_solve(train_diffs, K_tr_chol))\n",
|
||||
" # print(voltron.mean_module(test_x).detach().T.shape)\n",
|
||||
" # print(pred_mean.shape)\n",
|
||||
" pred_mean += model.mean_module(test_x).detach().T.unsqueeze(-1)\n",
|
||||
" if latent_mean is not None:\n",
|
||||
" pred_mean -= theta * (pred_mean - latent_mean)\n",
|
||||
"\n",
|
||||
" pred_cov = K_te - K_tr_te.transpose(-1, -2).matmul(torch.cholesky_solve(K_tr_te, K_tr_chol))\n",
|
||||
"\n",
|
||||
" pred_cov_L = psd_safe_cholesky(pred_cov, jitter=1e-4)\n",
|
||||
" samples = torch.randn(*cov_mat.shape[:-2], test_x.shape[0], 1).to(test_x.device)\n",
|
||||
" samples = pred_cov_L @ samples\n",
|
||||
"\n",
|
||||
" if pred_mean.ndim == 1:\n",
|
||||
" return samples + pred_mean.unsqueeze(-1)\n",
|
||||
" else:\n",
|
||||
" return (samples + pred_mean).squeeze(-1)\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"def SimulateMultiWind(train_x, train_y, test_x, model_list, vmod, nsample=10):\n",
|
||||
" vmod.eval()\n",
|
||||
" pred_vols = vmod(test_x).sample(sample_size=torch.Size((nsample)))\n",
|
||||
" "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 23,
|
||||
"id": "c29b2a0f-1a99-4da2-a8fe-2bc02e05493e",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"ename": "RuntimeError",
|
||||
"evalue": "Sizes of tensors must match except in dimension 1. Expected size 25 but got size 5 for tensor number 1 in the list.",
|
||||
"output_type": "error",
|
||||
"traceback": [
|
||||
"\u001b[0;31m---------------------------------------------------------------------------\u001b[0m",
|
||||
"\u001b[0;31mRuntimeError\u001b[0m Traceback (most recent call last)",
|
||||
"\u001b[0;32m<ipython-input-23-2fe6eec1fa70>\u001b[0m in \u001b[0;36m<module>\u001b[0;34m\u001b[0m\n\u001b[1;32m 19\u001b[0m stack_y = torch.cat((train_stack_y, \n\u001b[1;32m 20\u001b[0m samples[:, :idx, mdl_idx].to(train_stack_y.device)), -1)\n\u001b[0;32m---> 21\u001b[0;31m stack_vol = torch.cat((train_stack_vol, \n\u001b[0m\u001b[1;32m 22\u001b[0m pred_vol[:, :idx, mdl_idx].to(train_stack_vol.device)), -1)\n\u001b[1;32m 23\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n",
|
||||
"\u001b[0;31mRuntimeError\u001b[0m: Sizes of tensors must match except in dimension 1. Expected size 25 but got size 5 for tensor number 1 in the list."
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"nsample = 5\n",
|
||||
"nmodel = len(model_list)\n",
|
||||
"vmod.eval()\n",
|
||||
"pred_vol = vmod(test_x).sample(torch.Size((nsample,))).exp()\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"ntest = test_x.numel()\n",
|
||||
"samples = torch.zeros(nsample, ntest, nmodel)\n",
|
||||
"for mdl_idx, (model, lh) in enumerate(model_list):\n",
|
||||
" latent_mean = train_y.log()[mdl_idx, :].mean()\n",
|
||||
" samples[:, 0, mdl_idx] = GeneratePrediction(train_x, train_y[mdl_idx, :], \n",
|
||||
" test_x[0].unsqueeze(0), \n",
|
||||
" pred_vol[:, 0, mdl_idx].unsqueeze(1),\n",
|
||||
" model, latent_mean=None, theta=0.04).squeeze()\n",
|
||||
" train_stack_y = train_y[mdl_idx, 1:].log().repeat(nsample, 1)\n",
|
||||
" train_stack_vol = model.log_vol_path.repeat(nsample, 1)\n",
|
||||
"\n",
|
||||
" for idx in range(1, ntest):\n",
|
||||
" stack_y = torch.cat((train_stack_y, \n",
|
||||
" samples[:, :idx, mdl_idx].to(train_stack_y.device)), -1)\n",
|
||||
" stack_vol = torch.cat((train_stack_vol, \n",
|
||||
" pred_vol[:, :idx, mdl_idx].to(train_stack_vol.device)), -1)\n",
|
||||
" \n",
|
||||
" rolling_x = torch.cat((train_x, test_x[:idx]))\n",
|
||||
" model.mean_module.train_y = stack_y\n",
|
||||
" model.mean_module.train_x = rolling_x\n",
|
||||
"\n",
|
||||
" model.train_x = rolling_x\n",
|
||||
" model.train_y = stack_y\n",
|
||||
" model.log_vol_path = stack_vol\n",
|
||||
" samples[:, idx, mdl_idx] = GeneratePrediction(train_x, train_y[mdl_idx, :], \n",
|
||||
" test_x[idx].unsqueeze(0), \n",
|
||||
" pred_vol[:, idx, mdl_idx].unsqueeze(-1),\n",
|
||||
" model, latent_mean=None, theta=0.04).squeeze()\n",
|
||||
"\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "a9c0bc48-8a4e-4dfa-8d71-abd32f40d235",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"stn_idx = 0\n",
|
||||
"plt.plot(train_x.cpu(), train_y[stn_idx, 1:].cpu())\n",
|
||||
"plt.plot(test_x.cpu(), samples[..., stn_idx].cpu().T.exp(), color='gray')\n",
|
||||
"plt.ylim(0, 10)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "98d6f212-06bd-4e86-80e5-9689e30982e7",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.8"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,121 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "f2f51e59-0cf2-4022-9c78-d7c8fb287a2c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import pickle as pkl\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns\n",
|
||||
"import numpy as np\n",
|
||||
"import torch"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "ba5224a9-1f0f-4bc5-8c57-617f451e6777",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"names, locs, dat = pkl.load(open(\"./wind_data.p\", 'rb'))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 6,
|
||||
"id": "c09edb24-6efe-4942-b291-6af5a5a20e1a",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"wind = dat[0][:400]\n",
|
||||
"tr_x = torch.arange(wind.shape[0])\n",
|
||||
"te_x = torch.arange(wind.shape[0], wind.shape[0]+100)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 8,
|
||||
"id": "ea756bc0-74ff-4292-b423-200a2ed01b90",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"volt_preds = torch.load(\"./saved-outputs/stn0/volt_constant_400.pt\")\n",
|
||||
"gp_preds = torch.load(\"./saved-outputs/stn0/gp_matern_constant_400.pt\")\n",
|
||||
"lstm_preds = torch.load(\"./saved-outputs/stn0/lstm_400.pt\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 16,
|
||||
"id": "7cda0385-7859-4364-9831-6bc5274fe941",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"[<matplotlib.lines.Line2D at 0x7fced476e730>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7fced476e760>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7fced476e880>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7fced476e9a0>,\n",
|
||||
" <matplotlib.lines.Line2D at 0x7fced476eac0>]"
|
||||
]
|
||||
},
|
||||
"execution_count": 16,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAXQAAAD4CAYAAAD8Zh1EAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/MnkTPAAAACXBIWXMAAAsTAAALEwEAmpwYAABVM0lEQVR4nO2debgcVZn/v6d6vVvWm5VAdsIqAULYNSKGXVQYf6AjjqMyIo77jLiLjsqIOCIoDLjhhjgKiBJkJwIhhCSQjSSQjZB9vblrL1V1fn9UnepTp091V/ftvr3k/TzPfW53dXX1qequb731Pe95D+OcgyAIgmh8jFo3gCAIgqgMJOgEQRBNAgk6QRBEk0CCThAE0SSQoBMEQTQJ0Vp9cGdnJ58yZUqtPp4gCKIhWbZs2T7O+RjdazUT9ClTpmDp0qW1+niCIIiGhDH2RtBrZLkQBEE0CSToBEEQTQIJOkEQRJNAgk4QBNEkkKATBEE0CSToBEEQTQIJOkEQRJNAgk6EZvnWg1iz41Ctm0EQRAA1G1hENB7v/ekiAMCWmy6pcUsIgtBBETpBEESTQIJOEATRJJCgEwRBNAkk6ARBEE0CCTpBEESTQIJOEATRJJCgEwRBNAkk6ARBEE0CCTpBEESTQIJOEATRJJCgEwRBNAkk6ETJcM5r3QSCIDSQoBMlk7VI0AmiHikq6IyxJGNsCWNsBWNsDWPsRs068xhjhxhjr7h/X69Oc4laIUflWcuuYUsIgggiTPncNIDzOOe9jLEYgOcYY49wzhcr6z3LOb+08k0k6gHLJkEniHqnqKBzJzTrdZ/G3D+65z7MMCVBz5CgE0RdEspDZ4xFGGOvANgD4HHO+Yua1c50bZlHGGPHB2znWsbYUsbY0r1795bfamLI8UfodD0niHoklKBzzi3O+WwAkwDMZYydoKyyHMBkzvlJAG4D8GDAdu7inM/hnM8ZM2ZM+a0mhhxL9tBNitAJoh4pKcuFc94F4BkAFyrLuznnve7jBQBijLHOCrWRqAMsizx0gqh3wmS5jGGMjXAftwA4H8A6ZZ3xjDHmPp7rbnd/xVtL1IwgD31fbxoDGasWTSIIQiFMlssEAPcwxiJwhPqPnPO/McY+DgCc8zsBXAngOsaYCWAAwFWcRp80FUEe+pz/egLHTxyGhz91bi2aRRCERJgsl5UATtYsv1N6fDuA2yvbNKKeMO1cVK5aLmt2dA91cwiC0EAjRYlQSHpOnaIEUaeQoBOhkCN0ykMniPqEBJ0IBeWhE0T9Q4JOhELOckmblNVCEPUICToRCjlCz5CHThA+1q9fj4MHD9a6GSToRDgsX4TuCDplphKEcx784Q9/wM9//vNaN4UEnQiHz3LJOpaLTXpOEOjv7wcA9PX11bglJOhESHQRukWKThDo7nbGYbS0tNS4JSToREhI0AlCT09PDwCgra2txi0hQSdCYmmyXCzy0AnCi9BbW1tr3BISdCIk8sCidJYidIIQiAg9mUzWuCUk6ERIdJaLTYJOEBgYGABQH1lfJOhEKHQDi8hyIQggm80CACyr9gPuSNCJUNgUoROEFtM0ARQX9CVLlmDVqlVVbUuYeugE4UXoBst56CYJOkGEFvRHHnkEAHDiiSdWrS0UoROhEB56Wzyas1xI0InDCNu2PfGWCWO5iHWqDQk6EQoRjbfEIznLhTx04jDil7/8Jb7zne/kLRcirxN7wVDVeSFBJ0Ih/PK2RFQ7sKgeevgJopps27ZNuzxMhL5/vzPFsjv1ctUgQSdC4UXosYhnucgROtVIJw4XVOEO46H39vYCANrb26vXMJCgEyER4t0aj0gDi3KvywOPCKKZEcW4BGEidCH6NY/QGWNJxtgSxtgKxtgaxtiNmnUYY+zHjLENjLGVjLFTqtNcolYIQW+JR5AyLWza2+ufONqkCJ04PBDRtqChBB1AGsB5nPOTAMwGcCFj7AxlnYsAzHT/rgVwRyUbSdQe4aEnYxG8eWAA592yENsPDnivZylCJ5ocw3DkUi2TG8ZyGapBR0UFnTuIS1LM/VPDscsB/NpddzGAEYyxCZVtKlFLRP/n8JaYt6yrP5eKlaWJo4kqkclk8MADD9S83rgoj6u2QxehP/bYY3jggQe852Fz1QdLKA+dMRZhjL0CYA+AxznnLyqrHAHgTen5NneZup1rGWNLGWNL9+7dW2aTiVogLJfr5k3H/OPGAQBS0tyiJnWKEhUmk8mgv78fK1aswMqVK/H000/XtD2i+JYs6JxzT6Qty/KyvV544QWsXLnSW6+uBJ1zbnHOZwOYBGAuY+wEZRWdMZR3hnPO7+Kcz+GczxkzZkzJjSVqh0hoGTcsifeffhQAIJXN/TgzFKETFeaWW27BzTffXOtmeESjzsD6dDrtLctkMgCARCIBwBl8JKoviucAfKJfTUrKcuGcdwF4BsCFykvbABwpPZ8EYMdgGkbUFyJCNxgQjzg/m1Q2J+IUoROVRohltTsSwyKib9EuAPjRj34EICfolmVh165d3uuiA7VuInTG2BjG2Aj3cQuA8wGsU1Z7CMA1brbLGQAOcc53VrqxRO0QHrrBGGJR52czIEXo5KET1UKIYKGRmGHp6enxCXIpiGhbjtBTqRSAnB1jWZavnWLyC7EPdpWTB8IU55oA4B7GWATOBeCPnPO/McY+DgCc8zsBLABwMYANAPoBfLhK7SVqhIjQGQNiXoSeE3Sq60JUCzFCc8eOwd/0//CHP8TYsWNx3XXXlfxeIcq6C4IQdNM0fVG4sF/kmum2bXsZM5WmqKBzzlcCOFmz/E7pMQdwfWWbRtQT3LNcGGIR5xZYtlyorgtRLYRAVqq8xJ49e8p6n4iuhaDLwi5bLqqgm6aJDRs2eMssy6qaoNNIUSIUsuUiPPS0FKGToBPVQtgWstVRC1RBl7NdREqjKujZbDYvov/ud7+bNzipUpCgE6GwpHronuViyoJek2YRhwHCrhishz5Y/1oItbiwyKLc0dEBwKmqKAu6aZra0rmLFy8eVFuCIEEnQsE5B2NOxoHoFJUtF/LQiWpRTidmJpPBunX+3I3BZpjIEfr+/fuxceNG77WRI0fCMAxs3rzZ9zlr167V3lmIFMhKQ4JOhMLmjt0CwPPQBzJkuRDVQfbLhUCWkr74l7/8Bffddx/27dvnLatUhJ9Op3H77bdj4cKF3murVq3CxIkTsW3bNp+g7969G11dXXnbIg+dqCk25zDc8ymus1woa5GoILJNEVbQV69eje9///uwLMsTctXPDkNQ52uhLBfTNNHR0YFDhw5h7dq1vteGsmQBCToRCpvnTqiYZmARRehEJZFFM8yMQACwYMECDAwMIJVKeeIbiUTytlOIl19+Gd/61re8jlgZEaHrLgyccyQSCXR3d2P79u2+13QdoNUaLEWCToSCSxF6TJPlYpGgExVEjmrV4fNByEIrxFvuCA0j6KtWrQIAn1UDOL9/IdpByK/JlooQdFnESdCJmuJYLn4PXR5YFDZHuDs1NJPlEo1DOp3GzTff7OvElOuhhEUWcd3IzMF46GI7c+bM0c46xDlHMpn0zgNZsHWCXi1I0IlQyJ2izM1FH/CNFC2+jQde3oa3fPMxrN2Zfzt7OGPbHHcu3Ije9OCHtjcihw4dQn9/P+677z5vmc7yCIss6GoKYVjUAEUIejKZxIc/nD8Q3rbtwOhdXJxkQa/WHLwk6EQobDdtURCLMJ+gh/HQn1nvlExet4sEXeaxV3fhpkfW4XsL1hZfuQmRxU08PnToUNnbCxOhB9k3QVG07MmPGjUK11xzje91zjk2b96sfa/OQydBJ2oKlyJ0AIhFDX+naIg8dPFustv9pE3nOHanTGW5hUt+/CzuX74Nh/qb16qS87QPHjwIQG+5mKaJ/fv35y1/7rnn8Prrr3vPlyxZ4rNfbrnlFjz77LM+Qd+9ezeef/750G0UFwbhjavROOfc1wZZsHWCHjbjplRI0IlQWHauUxTIdYwKwowrEtEPCbqf3HHxH5g3DwxgzY5ufO6PKzDvB7Wd3EGmO5XFl+5fhb4KWURyudmdO50irUGpfnfccQcWLFjgW/bkk0/i97//vfd88eLFngCbpone3l489dRTPkG/++678cQTT2DhwoX44x//mPc56qjSYoIuLkS69+vEW12/UpCgE6GwOUdEUvS4Iuhhsly8CL2SDWsCgo6LfLwPFojQH165Exv2lN6JWC7/u3Aj7l2yFb9+4Y2KbE/OKBEzmR04cEC7rmVZeOmll0L74XL0v2nTprzXn3nmmby8cfE5uuf/+Mc/sGrVqryou9QO18FYSoWozvhToumQ89CBXKaLIJQnWB/zFNQd4rCqx9AIebyu//1yAMCWmy6pZLMCMd3bMV6hS7Msjlu3bsXmzZu1oytltm7dimnTphXdtqhXDgDLli0rur74jauCLiLu3t5e3H///UW3o0P+fkeNGlXWNopBgk6EQs5DB/Itl1JquVSrQ6hRYdBbUWaBY2rbHG8c6MfUzraKtGFPTwrJWATDkrHiK7vNYhW6QsuDiDZv3hzYuShz4MABTJs2TftbYox5y2VBL4Rao1yNuCsxMYW8jUmTJg16ezrIciFCIeehA2V66EK4KtqyxicXofuXF5rW78dPvY63/+AZbNhTmTKsc7/zJM7973A+vWhV2DuIYpRTfEt47DqhbW1tzVsvTBvWr1/vXQiCLBeVcotszZkzp6z3FYMidCIUch46AK/iovd6mCwXMtG15A6L/8AUmtZvyWbHY97eNVCxdhwaCJd5Ib7rSo2TKWfAz969e/Nqjwvk4f6LFi0Ktb1XXnkFjz76KEaPHg3AL+DLly/HE088oX1fJBIpq/3VGmREgk6EQs1Djyseepg8dLLQ9QRl/xSyXMR3kTWHviqaaJVsuSzfehDTO9sxvDWEZaNQjiCuWbMGAwMDOO+88/JeC2uzyCxf7vRDyPOXdnd34+9//7u201RLine truncated
|
||||
"text/plain": [
|
||||
"<Figure size 432x288 with 1 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {
|
||||
"needs_background": "light"
|
||||
},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"plt.plot(tr_x, wind)\n",
|
||||
"# plt.plot(te_x, volt_preds[:5, :].T.exp().detach(), color='gray')\n",
|
||||
"# plt.plot(te_x, gp_preds[:5, :].T.exp().detach(), color='gray')\n",
|
||||
"plt.plot(te_x, lstm_preds[:5, :].T.exp().detach(), color='gray')"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "9723695b-71fc-4963-a35a-06be9a5f43bf",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.12"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
Loaded 100 of 721 files, more files were not shown because too many files have changed in this diff.
Show more
Reference in new issue
Block a user