diff --git a/mlatom/addons/omnip2x/vecmsani.py b/mlatom/addons/omnip2x/vecmsani.py index a26efeb..d31d502 100644 --- a/mlatom/addons/omnip2x/vecmsani.py +++ b/mlatom/addons/omnip2x/vecmsani.py @@ -429,6 +429,11 @@ def train( if reset_optimizer: self.optimizer_setup(**self.hyperparameters) + # fix layers + if 'fixed_layers' in hyperparameters: + self.fix_layers(getattr(hyperparameters['fixed_layers'], 'value', + hyperparameters['fixed_layers'])) + self.model.train() if self.verbose: print(self.model) @@ -1000,4 +1005,4 @@ def data_setup(self, molecular_database, validation_molecular_database, spliting self.subtraining_set = self.subtraining_set.collate(self.hyperparameters.batch_size, padding=PADDING) self.validation_set = self.validation_set.collate(self.hyperparameters.batch_size, padding=PADDING) - self.argsdict.update({'self_energies': self.energy_shifter.self_energies, 'property': self.property_name}) \ No newline at end of file + self.argsdict.update({'self_energies': self.energy_shifter.self_energies, 'property': self.property_name}) diff --git a/mlatom/interfaces/torchani_interface.py b/mlatom/interfaces/torchani_interface.py index 865dcbc..4b36e3a 100644 --- a/mlatom/interfaces/torchani_interface.py +++ b/mlatom/interfaces/torchani_interface.py @@ -1272,6 +1272,11 @@ def train( if reset_optimizer: self.optimizer_setup(**self.hyperparameters) + # fix layers + if 'fixed_layers' in hyperparameters: + self.fix_layers(getattr(hyperparameters['fixed_layers'], 'value', + hyperparameters['fixed_layers'])) + self.model.train() if self.verbose: print(self.model)