| 232 | |
| 233 | |
| 234 | class BatchGDManager: |
| 235 | def __init__( |
| 236 | self, |
| 237 | experiment_id, |
| 238 | model, |
| 239 | train_loader, |
| 240 | val_loader, |
| 241 | test_loader, |
| 242 | epochs, |
| 243 | learning_rate, |
| 244 | patience, |
| 245 | general_hyperparameters, |
| 246 | model_hyperparameters, |
| 247 | ): |
| 248 | self.experiment_id = experiment_id |
| 249 | self.model = model |
| 250 | self.train_loader = train_loader |
| 251 | self.val_loader = val_loader |
| 252 | self.test_loader = test_loader |
| 253 | self.epochs = epochs |
| 254 | self.learning_rate = learning_rate |
| 255 | self.patience = patience |
| 256 | self.general_hyperparameters = general_hyperparameters |
| 257 | self.model_hyperparameters = model_hyperparameters |
| 258 | self.lob_lightning_module = None |
| 259 | self.trainer = None |
| 260 | self.deleted_run = None |
| 261 | |
| 262 | def delete_run(self): |
| 263 | api = wandb.Api() |
| 264 | project_path = "<Specify here the name of WB project>" # TODO: Specify here the name of WB project. |
| 265 | runs = api.runs(path=project_path) |
| 266 | print('Deleting runs...') |
| 267 | while len(runs) < 1: |
| 268 | runs = api.runs(path=project_path) |
| 269 | for run in runs: |
| 270 | input_list = run.metadata |
| 271 | if input_list is not None: |
| 272 | input_list = input_list['args'] |
| 273 | result_dict = {input_list[i][2:]: input_list[i + 1] for i in range(0, len(input_list), 2)} |
| 274 | modified_dict = result_dict |
| 275 | if modified_dict['model'] == str(self.general_hyperparameters['model']) and modified_dict['prediction_horizon'] == str(self.model_hyperparameters['prediction_horizon']) and modified_dict['training_stocks'] == str(self.general_hyperparameters['training_stocks'][0]) and modified_dict['target_stocks'] == str(self.general_hyperparameters['target_stocks'][0]): |
| 276 | self.deleted_run = run.name |
| 277 | run.delete() |
| 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), |