Skip to content

Commit 8bcb779

Browse files
committed
Add train_hyps to the on-the-fly settings
A variable controlling when the training is performed is added. This can give improvements if the hyperparameter training gets stucked into local minima due to few environments in the dataset.
1 parent e85e810 commit 8bcb779

3 files changed

Lines changed: 12 additions & 1 deletion

File tree

Modules/Ensemble.py

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -194,6 +194,7 @@ def __init__(self, dyn0, T0, supercell = None, **kwargs):
194194
self.checkpt_files = None
195195
self.write_model = None
196196
self.init_atoms = None
197+
self.train_hyps = None
197198

198199
# The frequencies and polarizations of the ensemble
199200
# In q space
@@ -4239,6 +4240,7 @@ def set_otf(
42394240
update_threshold: float | None = None,
42404241
# other args
42414242
build_mode="bayesian",
4243+
train_hyps: tuple = (100,120),
42424244
):
42434245
"""Set on-the-fly training.
42444246
@@ -4299,6 +4301,12 @@ def set_otf(
42994301
]
43004302

43014303
self.write_model = write_model
4304+
4305+
if train_hyps[0] == "inf":
4306+
train_hyps[0] = np.inf
4307+
if train_hyps[1] == "inf":
4308+
train_hyps[1] = np.inf
4309+
self.train_hyps = train_hyps
43024310

43034311

43044312
#-------------------------------------------------------------------------------

Modules/aiida_ensemble.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -171,7 +171,8 @@ def compute_ensemble( # pylint: disable=arguments-renamed
171171
# ================ TRAIN SECTION ================ #
172172
if self.gp_model is not None:
173173
if dft_counts > 0:
174-
self._train_gp()
174+
if self.train_hyps[0] <= len(self.gp_model.training_data) <= self.train_hyps[1]:
175+
self._train_gp()
175176
self._write_model()
176177

177178
sys.stdout.flush()

tests/aiida_ensemble/test_otf_flare.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -74,6 +74,7 @@ def test_no_otf(generate_ensemble):
7474
assert ensemble.checkpt_files is None
7575
assert ensemble.write_model is None
7676
assert ensemble.init_atoms is None
77+
assert ensemble.train_hyps is None
7778

7879

7980
def test_set_otf(generate_ensemble):
@@ -83,6 +84,7 @@ def test_set_otf(generate_ensemble):
8384
ensemble.set_otf(flare_calc, max_atoms_added=-1)
8485

8586
assert ensemble.gp_model is not None
87+
assert ensemble.train_hyps == (100,120)
8688

8789

8890
def test_compute_properties(generate_ensemble):

0 commit comments

Comments
 (0)