Skip to content

Commit 2874fa9

Browse files
committed
Implemented it for all architectures and added tests
1 parent 5bcb6e1 commit 2874fa9

17 files changed

Lines changed: 202 additions & 22 deletions

File tree

src/metatrain/composition/model.py

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,6 @@
11
import logging
22
import warnings
3-
from typing import Dict, List, Literal, Optional, Union
3+
from typing import Any, Dict, List, Literal, Optional, Union
44

55
import metatensor.torch as mts
66
import torch
@@ -22,6 +22,7 @@
2222
sparsify_atomic_basis_target,
2323
)
2424
from metatrain.utils.dtype import dtype_to_str
25+
from metatrain.utils.hypers import raise_if_hypers_mismatch
2526
from metatrain.utils.metadata import merge_metadata
2627

2728
from . import checkpoints
@@ -204,7 +205,9 @@ def train_model(
204205
checkpoint_dir="",
205206
)
206207

207-
def restart(self, dataset_info: DatasetInfo) -> "CompositionModel":
208+
def restart(
209+
self, dataset_info: DatasetInfo, model_hypers: Optional[dict[str, Any]] = None
210+
) -> "CompositionModel":
208211
"""
209212
Update the model to continue training, possibly with new targets.
210213
@@ -214,8 +217,13 @@ def restart(self, dataset_info: DatasetInfo) -> "CompositionModel":
214217
215218
:param dataset_info: Information about the new dataset, including the
216219
targets that will be used for training.
220+
:param model_hypers: New hyperparameters for the model. They must match
221+
the existing hyperparameters, otherwise an error is raised.
217222
:return: The updated model.
218223
"""
224+
if model_hypers is not None:
225+
raise_if_hypers_mismatch(self.hypers, model_hypers)
226+
219227
raw_targets = {}
220228
for target_name in dataset_info.targets:
221229
target_info = dataset_info.targets[target_name]

src/metatrain/experimental/classifier/model.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -123,7 +123,9 @@ def build_mlp(self, feature_size: int, num_classes: int) -> None:
123123
# Final classification layer
124124
self.linear = torch.nn.Linear(current_size, num_classes, bias=False)
125125

126-
def restart(self, dataset_info: DatasetInfo) -> "Classifier":
126+
def restart(
127+
self, dataset_info: DatasetInfo, model_hypers: Optional[dict[str, Any]] = None
128+
) -> "Classifier":
127129
raise ValueError("Restarting from a Classifier model is not supported.")
128130

129131
def forward(

src/metatrain/experimental/dpa3/model.py

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,7 @@
2525
from metatrain.utils.data.atom_pair_helpers import check_no_atom_pair_targets
2626
from metatrain.utils.data.dataset import DatasetInfo
2727
from metatrain.utils.dtype import dtype_to_str
28+
from metatrain.utils.hypers import raise_if_hypers_mismatch
2829
from metatrain.utils.metadata import merge_metadata
2930
from metatrain.utils.scaler import Scaler
3031
from metatrain.utils.sum_over_atoms import sum_over_atoms
@@ -391,7 +392,13 @@ def forward(
391392

392393
return return_dict
393394

394-
def restart(self, dataset_info: DatasetInfo) -> "DPA3":
395+
def restart(
396+
self, dataset_info: DatasetInfo, model_hypers: Optional[dict[str, Any]] = None
397+
) -> "DPA3":
398+
399+
if model_hypers is not None:
400+
raise_if_hypers_mismatch(self.hypers, model_hypers)
401+
395402
# merge old and new dataset info
396403
merged_info = self.dataset_info.union(dataset_info)
397404
new_atomic_types = [

src/metatrain/experimental/flashmd/model.py

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,7 @@
2828
from metatrain.utils.data import DatasetInfo, TargetInfo
2929
from metatrain.utils.data.atom_pair_helpers import check_no_atom_pair_targets
3030
from metatrain.utils.dtype import dtype_to_str
31+
from metatrain.utils.hypers import raise_if_hypers_mismatch
3132
from metatrain.utils.long_range import DummyLongRangeFeaturizer, LongRangeFeaturizer
3233
from metatrain.utils.metadata import merge_metadata
3334
from metatrain.utils.scaler import Scaler
@@ -232,7 +233,13 @@ def set_timestep(self, timestep: float):
232233
def supported_outputs(self) -> Dict[str, ModelOutput]:
233234
return self.outputs
234235

235-
def restart(self, dataset_info: DatasetInfo) -> "FlashMD":
236+
def restart(
237+
self, dataset_info: DatasetInfo, model_hypers: Optional[dict[str, Any]] = None
238+
) -> "FlashMD":
239+
240+
if model_hypers is not None:
241+
raise_if_hypers_mismatch(self.hypers, model_hypers)
242+
236243
# merge old and new dataset info
237244
merged_info = self.dataset_info.union(dataset_info)
238245
new_atomic_types = [

src/metatrain/experimental/flashmd_symplectic/model.py

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,7 @@
3131
from metatrain.utils.data.atom_pair_helpers import check_no_atom_pair_targets
3232
from metatrain.utils.data.target_info import get_energy_target_info
3333
from metatrain.utils.dtype import dtype_to_str
34+
from metatrain.utils.hypers import raise_if_hypers_mismatch
3435
from metatrain.utils.long_range import DummyLongRangeFeaturizer, LongRangeFeaturizer
3536
from metatrain.utils.metadata import merge_metadata
3637
from metatrain.utils.scaler import Scaler
@@ -222,7 +223,13 @@ def set_timestep(self, timestep: float):
222223
def supported_outputs(self) -> Dict[str, ModelOutput]:
223224
return self.outputs
224225

225-
def restart(self, dataset_info: DatasetInfo) -> "FlashMDSymplectic":
226+
def restart(
227+
self, dataset_info: DatasetInfo, model_hypers: Optional[dict[str, Any]] = None
228+
) -> "FlashMDSymplectic":
229+
230+
if model_hypers is not None:
231+
raise_if_hypers_mismatch(self.hypers, model_hypers)
232+
226233
# merge old and new dataset info
227234
merged_info = self.dataset_info.union(dataset_info)
228235
new_atomic_types = [

src/metatrain/experimental/mace/model.py

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,7 @@
3030
sparsify_atomic_basis_target,
3131
)
3232
from metatrain.utils.dtype import dtype_to_str
33+
from metatrain.utils.hypers import raise_if_hypers_mismatch
3334
from metatrain.utils.metadata import merge_metadata
3435
from metatrain.utils.scaler import Scaler
3536
from metatrain.utils.sum_over_atoms import sum_over_atoms
@@ -306,7 +307,13 @@ def __init__(self, hypers: ModelHypers, dataset_info: DatasetInfo) -> None:
306307

307308
self.finetune_config: Dict[str, Any] = {}
308309

309-
def restart(self, dataset_info: DatasetInfo) -> "MetaMACE":
310+
def restart(
311+
self, dataset_info: DatasetInfo, model_hypers: Optional[dict[str, Any]] = None
312+
) -> "MetaMACE":
313+
314+
if model_hypers is not None:
315+
raise_if_hypers_mismatch(self.hypers, model_hypers)
316+
310317
# Check that the new dataset info does not contain new atomic types
311318
if new_atomic_types := set(dataset_info.atomic_types) - set(
312319
self.dataset_info.atomic_types

src/metatrain/experimental/space/model.py

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -35,6 +35,7 @@
3535
)
3636
from metatrain.utils.data.dataset import DatasetInfo, TargetInfo
3737
from metatrain.utils.dtype import dtype_to_str
38+
from metatrain.utils.hypers import raise_if_hypers_mismatch
3839
from metatrain.utils.metadata import merge_metadata
3940
from metatrain.utils.scaler import Scaler
4041

@@ -166,7 +167,13 @@ def __init__(self, hypers: ModelHypers, dataset_info: DatasetInfo) -> None:
166167
def supported_outputs(self) -> Dict[str, ModelOutput]:
167168
return self.outputs
168169

169-
def restart(self, dataset_info: DatasetInfo) -> "SPACE":
170+
def restart(
171+
self, dataset_info: DatasetInfo, model_hypers: Optional[dict[str, Any]] = None
172+
) -> "SPACE":
173+
174+
if model_hypers is not None:
175+
raise_if_hypers_mismatch(self.hypers, model_hypers)
176+
170177
# merge old and new dataset info
171178
merged_info = self.dataset_info.union(dataset_info)
172179
new_atomic_types = [

src/metatrain/experimental/space/tests/test_basic.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@
1111
from metatrain.utils.testing import (
1212
ArchitectureTests,
1313
CheckpointTests,
14+
InputTests,
1415
OutputTests,
1516
TorchscriptTests,
1617
)
@@ -31,6 +32,9 @@ def minimal_model_hypers(self):
3132
return hypers
3233

3334

35+
class TestInput(InputTests, SPACETests): ...
36+
37+
3438
class TestOutput(OutputTests, SPACETests):
3539
is_equivariant_reflections = False
3640
equivariance_error_tolerance = 1e-4 # due to many layers in the default hypers

src/metatrain/gap/model.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -169,7 +169,9 @@ def __init__(self, hypers: ModelHypers, dataset_info: DatasetInfo) -> None:
169169
def supported_outputs(self) -> Dict[str, ModelOutput]:
170170
return self.outputs
171171

172-
def restart(self, dataset_info: DatasetInfo) -> "GAP":
172+
def restart(
173+
self, dataset_info: DatasetInfo, model_hypers: Optional[dict[str, Any]] = None
174+
) -> "GAP":
173175
raise NotImplementedError("GAP does not allow restarting training")
174176

175177
@classmethod

src/metatrain/llpr/model.py

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,7 @@
2626
from metatrain.utils.data.target_info import (
2727
is_auxiliary_output,
2828
)
29+
from metatrain.utils.hypers import raise_if_hypers_mismatch
2930
from metatrain.utils.io import model_from_checkpoint
3031
from metatrain.utils.metadata import merge_metadata
3132
from metatrain.utils.neighbor_lists import (
@@ -250,7 +251,13 @@ def set_wrapped_model(self, model: ModelInterface) -> None:
250251
bias=False,
251252
)
252253

253-
def restart(self, dataset_info: DatasetInfo) -> "LLPRUncertaintyModel":
254+
def restart(
255+
self, dataset_info: DatasetInfo, model_hypers: Optional[dict[str, Any]] = None
256+
) -> "LLPRUncertaintyModel":
257+
258+
if model_hypers is not None:
259+
raise_if_hypers_mismatch(self.hypers, model_hypers)
260+
254261
# merge old and new dataset info
255262
merged_info = self.dataset_info.union(dataset_info)
256263
new_atomic_types = [

0 commit comments

Comments
 (0)