(positional_encoding_name, neural_network_name, dataset="landoceandataset")
| 47 | return hparams |
| 48 | |
| 49 | def tune(positional_encoding_name, neural_network_name, dataset="landoceandataset"): |
| 50 | n_trials = 100 |
| 51 | timeout = 4 * 60 * 60 # seconds |
| 52 | epochs = 100 |
| 53 | |
| 54 | if dataset == "landoceandataset": |
| 55 | datamodule = LandOceanDataModule() |
| 56 | num_classes = 1 |
| 57 | regression = False |
| 58 | presence_only = False |
| 59 | loss_bg_weight = False |
| 60 | if dataset == "checkerboard": |
| 61 | datamodule = CheckerboardDataModule() |
| 62 | num_classes = 16 |
| 63 | regression = False |
| 64 | presence_only = False |
| 65 | loss_bg_weight = False, |
| 66 | elif dataset == "inat2018": |
| 67 | datamodule = Inat2018DataModule("/data/sphericalharmonics/inat2018/") |
| 68 | num_classes = 8142 |
| 69 | regression = False |
| 70 | presence_only = True |
| 71 | loss_bg_weight = 5 |
| 72 | |
| 73 | def objective(trial: optuna.trial.Trial) -> float: |
| 74 | |
| 75 | hparams = get_hyperparameter(trial, positional_encoding_name, neural_network_name) |
| 76 | hparams["num_classes"] = num_classes |
| 77 | hparams["presence_only_loss"] = presence_only |
| 78 | hparams["loss_bg_weight"] = loss_bg_weight |
| 79 | hparams["regression"] = regression |
| 80 | |
| 81 | spatialencoder = LocationEncoder( |
| 82 | positional_encoding_name, |
| 83 | neural_network_name, |
| 84 | hparams=hparams |
| 85 | ) |
| 86 | |
| 87 | trainer = pl.Trainer( |
| 88 | max_epochs=epochs, |
| 89 | log_every_n_steps=5, |
| 90 | accelerator='gpu', |
| 91 | callbacks=[EarlyStopping(monitor="val_loss", mode="min", patience=30)]) |
| 92 | |
| 93 | trainer.logger.log_hyperparams(hparams) |
| 94 | |
| 95 | trainer.fit(model=spatialencoder, datamodule=datamodule) |
| 96 | |
| 97 | return trainer.callback_metrics["val_loss"].item() |
| 98 | |
| 99 | pruner = optuna.pruners.MedianPruner() |
| 100 | study_name = f"{dataset}-{positional_encoding_name}-{neural_network_name}" |
| 101 | os.makedirs(f"{TUNE_RESULTS_DIR}/{dataset}/runs/", exist_ok=True) |
| 102 | storage_name = f"sqlite:///{TUNE_RESULTS_DIR}/{dataset}/runs/{study_name}.db" |
| 103 | study = optuna.create_study(study_name=study_name, direction="minimize", |
| 104 | storage=storage_name, load_if_exists=True, |
| 105 | pruner=pruner) |
| 106 | study.optimize(objective, n_trials=n_trials, timeout=timeout) |
no test coverage detected