(self)
| 278 | print(f"Run succesfully deleted from WanDB: {run.name}.") |
| 279 | |
| 280 | def train(self): |
| 281 | self.lob_lightning_module = LOBLightningModule( |
| 282 | self.model, |
| 283 | experiment_id=self.experiment_id, |
| 284 | learning_rate=self.learning_rate, |
| 285 | general_hyperparameters=self.general_hyperparameters, |
| 286 | model_hyperparameters=self.model_hyperparameters, |
| 287 | ) |
| 288 | |
| 289 | checkpoint_callback = ModelCheckpoint( |
| 290 | monitor="val_loss", |
| 291 | dirpath=logger.find_save_path(self.experiment_id), |
| 292 | filename="best_val_model", |
| 293 | save_top_k=1, |
| 294 | mode="min", |
| 295 | ) |
| 296 | early_stopping_callback = EarlyStopping("val_loss", patience=self.patience, min_delta=0.003) |
| 297 | |
| 298 | os.environ["WANDB_API_KEY"] = "" # TODO: Insert API key |
| 299 | os.environ["WANDB__SERVICE_WAIT"] = "300" |
| 300 | try: |
| 301 | wandb_logger = WandbLogger( |
| 302 | project="Limit_Order_Book", |
| 303 | name=self.experiment_id, |
| 304 | save_dir=logger.find_save_path(self.experiment_id), |
| 305 | ) |
| 306 | wandb_hyperparameters_saving( |
| 307 | wandb_logger=wandb_logger, |
| 308 | general_hyperparameters=self.general_hyperparameters, |
| 309 | model_hyperparameters=self.model_hyperparameters, |
| 310 | ) |
| 311 | self.trainer = pl.Trainer( |
| 312 | max_epochs=self.epochs, |
| 313 | callbacks=[checkpoint_callback, early_stopping_callback], |
| 314 | logger=wandb_logger, |
| 315 | num_sanity_val_steps=0, |
| 316 | ) |
| 317 | self.trainer.fit(self.lob_lightning_module, self.train_loader, self.val_loader) |
| 318 | wandb.finish() |
| 319 | except: |
| 320 | root_path = sys.path[0] |
| 321 | dir_path = f"{root_path}/loggers/results/{self.experiment_id}" |
| 322 | if os.path.exists(dir_path): |
| 323 | shutil.rmtree(dir_path) |
| 324 | print(f"Folder {self.experiment_id} deleted successfully.") |
| 325 | else: |
| 326 | print(f"Unable to delete folder {self.experiment_id}.") |
| 327 | |
| 328 | self.delete_run() |
| 329 | |
| 330 | model = self.general_hyperparameters['model'] |
| 331 | horizon = self.model_hyperparameters['prediction_horizon'] |
| 332 | training_stocks = self.general_hyperparameters['training_stocks'] |
| 333 | target_stocks = self.general_hyperparameters['target_stocks'] |
| 334 | errors_string = f"{model} {horizon} {training_stocks} {target_stocks} {self.deleted_run}\n" |
| 335 | with open("errors.txt", 'r+') as file: |
| 336 | content = file.read() |
| 337 |
no test coverage detected