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

Class BatchGDManager

optimizers/lightning_batch_gd.py:234–381  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

232
233
234class 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),

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected