This commit is contained in:
wassname
2020-04-09 14:15:38 +08:00
parent 629c0164d7
commit c3fd09bc43
6 changed files with 685 additions and 470 deletions
+1 -1
View File
@@ -2821,7 +2821,7 @@
" name=\"anp-rnn-mcdropout\",\n",
" params={\n",
" **default_params, \n",
"# 'det_enc_cross_attn_type': 'ptmultihead',\n",
" 'det_enc_cross_attn_type': 'ptmultihead',\n",
" 'latent_enc_self_attn_type': 'ptmultihead',\n",
" 'dropout': 0.3,\n",
" 'attention_dropout': 0.1,\n",
+23 -65
View File
@@ -89,7 +89,9 @@
"from src.data.smart_meter import collate_fns, SmartMeterDataSet, get_smartmeter_df\n",
"# from src.plot import plot_from_loader\n",
"from src.models.lstm import SequenceDfDataSet, LSTM_PL, plot_from_loader\n",
"from src.dict_logger import DictLogger"
"from src.dict_logger import DictLogger\n",
"from src.utils import PyTorchLightningPruningCallback\n",
"from src.train import main, objective, add_number, run_trial"
]
},
{
@@ -215,34 +217,7 @@
}
},
"outputs": [],
"source": [
"def main(trial, train=True):\n",
" # PyTorch Lightning will try to restore model parameters from previous trials if checkpoint\n",
" # filenames match. Therefore, the filenames for each trial must be made unique.\n",
" \n",
" checkpoint_callback = pl.callbacks.ModelCheckpoint(\n",
" os.path.join(MODEL_DIR, name, 'version_{}'.format(trial.number), \"chk\"), monitor='val_loss', mode='min')\n",
"\n",
" # The default logger in PyTorch Lightning writes to event files to be consumed by\n",
" # TensorBoard. We create a simple logger instead that holds the log in memory so that the\n",
" # final accuracy can be obtained after optimization. When using the default logger, the\n",
" # final accuracy could be stored in an attribute of the `Trainer` instead.\n",
" logger = DictLogger(MODEL_DIR, name=name, version=trial.number)\n",
"# print(\"log_dir\", logger.experiment.log_dir)\n",
"\n",
" trainer = pl.Trainer(\n",
" logger=logger,\n",
" val_percent_check=PERCENT_TEST_EXAMPLES,\n",
" checkpoint_callback=checkpoint_callback,\n",
" max_epochs=trial.params['max_nb_epochs'],\n",
" gpus=-1 if torch.cuda.is_available() else None,\n",
" early_stop_callback=PyTorchLightningPruningCallback(trial, monitor='val_loss')\n",
" )\n",
" model = LSTM_PL(trial.params)\n",
" if train:\n",
" trainer.fit(model)\n",
" return model, trainer"
]
"source": []
},
{
"cell_type": "code",
@@ -255,23 +230,23 @@
},
"outputs": [],
"source": [
"def add_suggest(trial):\n",
" trial.suggest_loguniform(\"learning_rate\", 1e-5, 1e-2)\n",
" trial.suggest_uniform(\"lstm_dropout\", 0, 0.75)\n",
" trial.suggest_categorical(\"hidden_size\", [1, 2, 4, 8, 16, 32, 64, 128]) \n",
" trial.suggest_categorical(\"lstm_layers\", [1, 2, 4, 8]) \n",
" trial.suggest_categorical(\"bidirectional\", [False, True]) \n",
"# def add_suggest(trial):\n",
"# trial.suggest_loguniform(\"learning_rate\", 1e-5, 1e-2)\n",
"# trial.suggest_uniform(\"lstm_dropout\", 0, 0.75)\n",
"# trial.suggest_categorical(\"hidden_size\", [1, 2, 4, 8, 16, 32, 64, 128]) \n",
"# trial.suggest_categorical(\"lstm_layers\", [1, 2, 4, 8]) \n",
"# trial.suggest_categorical(\"bidirectional\", [False, True]) \n",
" \n",
" # constants\n",
" trial.suggest_int(\"window_length\", 24 * 4, 24 * 4)\n",
" trial.suggest_int(\"target_length\", 24*4, 24*4)\n",
" trial.suggest_int(\"max_nb_epochs\", 20, 20)\n",
" trial.suggest_int(\"num_workers\", 4, 4)\n",
" trial.suggest_int(\"grad_clip\", 40, 40)\n",
" trial.suggest_int(\"vis_i\", 670, 670)\n",
" trial.suggest_int(\"input_size\", 17, 17)\n",
" trial.suggest_int(\"batch_size\", 16, 16) \n",
" return trial"
"# # constants\n",
"# trial.suggest_int(\"window_length\", 24 * 4, 24 * 4)\n",
"# trial.suggest_int(\"target_length\", 24*4, 24*4)\n",
"# trial.suggest_int(\"max_nb_epochs\", 20, 20)\n",
"# trial.suggest_int(\"num_workers\", 4, 4)\n",
"# trial.suggest_int(\"grad_clip\", 40, 40)\n",
"# trial.suggest_int(\"vis_i\", 670, 670)\n",
"# trial.suggest_int(\"input_size\", 17, 17)\n",
"# trial.suggest_int(\"batch_size\", 16, 16) \n",
"# return trial"
]
},
{
@@ -284,24 +259,7 @@
}
},
"outputs": [],
"source": [
"\n",
"def objective(trial):\n",
" # see https://github.com/optuna/optuna/blob/cf6f02d/examples/pytorch_lightning_simple.py\n",
" trial = add_suggest(trial)\n",
"\n",
" \n",
" print('trial', trial.number, 'params', trial.params)\n",
" \n",
" model, trainer = main(trial)\n",
" \n",
" # also report to tensorboard & print\n",
" print('logger.metrics', model.logger.metrics[-1:])\n",
" model.logger.experiment.add_hparams(trial.params, logger.metrics[-1])\n",
" model.logger.save()\n",
" \n",
" return model.logger.metrics[-1]['val_loss']\n"
]
"source": []
},
{
"cell_type": "markdown",
@@ -424,8 +382,8 @@
" 'vis_i': '670',\n",
" 'window_length': 24*4\n",
" })\n",
"trial = add_suggest(trial)\n",
"trial.number = 109\n",
"trial = LSTM_PL.add_suggest(trial)\n",
"trial = add_number(trial, MODEL_DIR/name)\n",
"model, trainer = main(trial, train=False)\n",
"trainer.fit(model)"
]
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
+12 -8
View File
@@ -60,7 +60,6 @@ class TransformerSeq2SeqNet(nn.Module):
norm=encoder_norm
)
self.dec_norm = BatchNormSequence(self.hparams.input_size_decoder)
self.dec_emb = nn.Linear(self.hparams.input_size_decoder, hidden_out_size)
layer_dec = nn.TransformerDecoderLayer(
@@ -89,20 +88,21 @@ class TransformerSeq2SeqNet(nn.Module):
torch.nn.init.xavier_uniform_(p)
def forward(self, context_x, context_y, target_x, target_y=None):
device = next(self.parameters()).device
x = torch.cat([context_x, context_y], -1)
# Size([B, C, input_dim])
x = self.enc_emb(self.enc_norm(x))
# Size([B, C, emb_dim])
x = self.enc_emb(self.enc_norm(x)).permute(1, 0, 2)
# Size([C, B, emb_dim])
memory = self.encoder(x)
# Size([B, C, emb_dim])
target_x = self.dec_emb(self.dec_norm(target_x))
# Size([B, T, input_target_dim]) -> Size([B, T, emb_dim])
# Size([C, B, emb_dim])
target_x = self.dec_emb(self.dec_norm(target_x)).permute(1, 0, 2)
# Size([T, B, input_target_dim]) -> Size([B, T, emb_dim])
# In transformers the memory and target_x need to be the same length. Lets use a permutation invariant agg on the context
# Then expand it, so it's available as we decode, conditional on target_x
memory = memory.max(dim=1, keepdim=True)[0].expand_as(target_x)
memory = memory.max(dim=0, keepdim=True)[0].expand_as(target_x)
outputs = self.decoder(target_x, memory)
outputs = self.decoder(target_x, memory).permute(1, 0, 2).contiguous()
# Size([B, T, emb_dim])
mean = self.mean(outputs)
log_sigma = self.std(outputs)
@@ -128,6 +128,10 @@ class TransformerSeq2SeqNet(nn.Module):
# mean = mean[:, self.hparams.num_context:]
# log_sigma = log_sigma[:, self.hparams.num_context:]
# Weight loss nearer to prediction time?
weight = (torch.arange(loss_p.shape[1])+1).float().to(device)[None, :]
loss_p = loss_p / torch.sqrt(weight) # We want to weight nearer stuff more
y_pred = y_dist.rsample if self.training else y_dist.loc
return y_pred, dict(loss_p=loss_p.mean(), loss_mse=loss_mse.mean()), dict(log_sigma=log_sigma, dist=y_dist)
+18 -10
View File
@@ -102,6 +102,17 @@ def run_trial(
print('KeyboardInterrupt, skipping rest of training')
pass
# Plot
loader = model.val_dataloader()
dset_test = loader.dataset
label_names = dset_test.label_names
plot_from_loader(model.val_dataloader(), model, i=670, title='overfit val 670')
plt.show()
plot_from_loader(model.train_dataloader(), model, i=670, title='overfit train 670')
plt.show()
plot_from_loader(model.test_dataloader(), model, i=670, title='overfit test 670')
plt.show()
# Load checkpoint
checkpoints = sorted(Path(trainer.checkpoint_callback.dirpath).glob("*.ckpt"))
if len(checkpoints):
@@ -110,16 +121,13 @@ def run_trial(
print(f"Loading checkpoint {checkpoint}")
model = model.load_from_checkpoint(checkpoint).to(device)
# Plot
loader = model.val_dataloader()
dset_test = loader.dataset
label_names = dset_test.label_names
plot_from_loader(model.val_dataloader(), model, i=670, title='val 670')
plt.show()
plot_from_loader(model.train_dataloader(), model, i=670, title='train 670')
plt.show()
plot_from_loader(model.test_dataloader(), model, i=670, title='test 670')
plt.show()
# Plot
plot_from_loader(model.val_dataloader(), model, i=670, title='val 670')
plt.show()
plot_from_loader(model.train_dataloader(), model, i=670, title='train 670')
plt.show()
plot_from_loader(model.test_dataloader(), model, i=670, title='test 670')
plt.show()
try:
trainer.test(model)