From 7b5185f7eedfe5cb095211a30cb0d4072d006402 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Janko=20=C4=8Civi=C4=87?= <107419364+jankocivic@users.noreply.github.com> Date: Fri, 11 Sep 2026 18:47:56 +0200 Subject: [PATCH 1/2] Implement layer fixing based on hyperparameters Added functionality to fix layers based on hyperparameters like how standard ani model does. --- mlatom/addons/omnip2x/vecmsani.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) 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}) From 3df49625aae2e4e641fbc903eade412b650a1489 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Janko=20=C4=8Civi=C4=87?= <107419364+jankocivic@users.noreply.github.com> Date: Fri, 11 Sep 2026 19:02:26 +0200 Subject: [PATCH 2/2] Implement layer fixing in msani Added functionality to fix layers based on hyperparameters the same as standard ani model does. --- mlatom/interfaces/torchani_interface.py | 5 +++++ 1 file changed, 5 insertions(+) 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)