MCPcopy Create free account
hub / github.com/MarcCoru/locationencoder / get_hyperparameter

Function get_hyperparameter

tune.py:15–47  ·  view source on GitHub ↗
(trial: optuna.trial.Trial, positional_encoding_name, neural_network_name)

Source from the content-addressed store, hash-verified

13logging.getLogger("lightning").setLevel(logging.ERROR)
14
15def get_hyperparameter(trial: optuna.trial.Trial, positional_encoding_name, neural_network_name):
16
17 hparams_pe = {}
18 if positional_encoding_name in ["theory", "grid", "spherec", "spherecplus", "spherem", "spheremplus"]:
19 hparams_pe["min_radius"] = trial.suggest_int("min_radius", 1, 90, step=9)
20 hparams_pe["max_radius"] = 360
21 hparams_pe["frequency_num"] = trial.suggest_int("frequency_num", 16, 64, step=16)
22 elif positional_encoding_name == "sphericalharmonics":
23 hparams_pe["legendre_polys"] = trial.suggest_int("legendre_polys", 10, 30, step=5)
24 hparams_pe["embedding_dim"] = trial.suggest_int("embedding_dim", 16, 128, step=16)
25
26 hparams_nn = {}
27 if neural_network_name == "mlp":
28 hparams_nn["dim_hidden"] = trial.suggest_int("dim_hidden", 32, 128, step=32)
29 hparams_nn["num_layers"] = trial.suggest_int("num_layers", 1, 3)
30 elif neural_network_name == "fcnet":
31 hparams_nn["dim_hidden"] = trial.suggest_int("dim_hidden", 32, 128, step=32)
32 elif neural_network_name == "siren":
33 hparams_nn["dim_hidden"] = trial.suggest_int("dim_hidden", 32, 128, step=32)
34 hparams_nn["num_layers"] = trial.suggest_int("num_layers", 1, 3)
35
36 hparams_opt = {}
37 hparams_opt["lr"] = trial.suggest_float("lr", 1e-4, 1e-1, log=True)
38 hparams_opt["wd"] = trial.suggest_float("wd", 1e-8, 1e-1, log=True)
39
40 hparams = {}
41 hparams.update(hparams_pe)
42 hparams.update(hparams_nn)
43 hparams["optimizer"] = hparams_opt
44
45 hparams['harmonics_calculation'] = "analytic"
46
47 return hparams
48
49def tune(positional_encoding_name, neural_network_name, dataset="landoceandataset"):
50 n_trials = 100

Callers 1

objectiveFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected