MCPcopy Create free account
hub / github.com/FinancialComputingUCL/LOBFrame / train

Method train

optimizers/lightning_batch_gd.py:280–345  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

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

Callers 1

execute_trainingMethod · 0.80

Calls 3

delete_runMethod · 0.95
LOBLightningModuleClass · 0.85

Tested by

no test coverage detected