From a946e5539d34bd4c13c9f8ede663f691bd771b9d Mon Sep 17 00:00:00 2001 From: Pol Febrer Date: Fri, 18 Sep 2026 08:16:15 +0200 Subject: [PATCH 1/3] First draft of metatrain wrapper --- src/metatrain/cli/train.py | 60 ++-- src/metatrain/composition/model.py | 2 - src/metatrain/pet/model.py | 180 +---------- src/metatrain/pet/modules/finetuning.py | 25 +- src/metatrain/pet/tests/test_regression.py | 10 +- src/metatrain/pet/trainer.py | 111 ++++++- src/metatrain/scaler/model.py | 12 +- src/metatrain/utils/io.py | 9 + src/metatrain/utils/testing/training.py | 31 +- src/metatrain/utils/wrapper.py | 360 +++++++++++++++++++++ 10 files changed, 545 insertions(+), 255 deletions(-) create mode 100644 src/metatrain/utils/wrapper.py diff --git a/src/metatrain/cli/train.py b/src/metatrain/cli/train.py index 7ce19fe48e..1e64034589 100644 --- a/src/metatrain/cli/train.py +++ b/src/metatrain/cli/train.py @@ -635,51 +635,47 @@ def train_model( else: training_context = None - new_model_hypers = input_options.get("architecture", {}).get("model", {}) - try: - if training_context == "restart" and restart_from is not None: - logging.info(f"Restarting training from '{restart_from}'") - checkpoint = torch.load( - restart_from, weights_only=False, map_location="cpu" - ) - try: - model = model_from_checkpoint(checkpoint, context="restart") - except Exception as e: - raise ValueError( - f"The file {restart_from} does not contain a valid checkpoint for " - f"the '{architecture_name}' architecture" - ) from e - model = model.restart(dataset_info, model_hypers=new_model_hypers) - try: - trainer = trainer_from_checkpoint( - checkpoint=checkpoint, - hypers=hypers["training"], - context=training_context, # type: ignore - ) - except Exception as e: - raise ValueError( - f"The file {restart_from} does not contain a valid checkpoint for " - f"the '{architecture_name}' trainer state" - ) from e - elif training_context == "finetune" and restart_from is not None: - logging.info(f"Starting finetuning from '{restart_from}'") + if restart_from is not None: checkpoint = torch.load( restart_from, weights_only=False, map_location="cpu" ) + new_model_hypers = input_options.get("architecture", {}).get("model", {}) + + # Initialize the trainer. + if training_context == "restart": + logging.info(f"Restarting training from '{restart_from}'") + try: + trainer = trainer_from_checkpoint( + checkpoint=checkpoint, + hypers=hypers["training"], + context=training_context, # type: ignore + ) + except Exception as e: + raise ValueError( + f"The file {restart_from} does not contain a valid checkpoint for " + f"the '{architecture_name}' trainer state" + ) from e + else: + logging.info(f"Starting finetuning from '{restart_from}'") + trainer = Trainer(hypers["training"]) + + # Load the model from the checkpoint. try: - model = model_from_checkpoint(checkpoint, context="finetune") + model = model_from_checkpoint(checkpoint, context=training_context) except Exception as e: raise ValueError( f"The file {restart_from} does not contain a valid checkpoint for " f"the '{architecture_name}' architecture" ) from e - model = model.restart(dataset_info, model_hypers=new_model_hypers) - trainer = Trainer(hypers["training"]) + + # Make the trainer setup the model to continue training. + model = trainer.restart(model, dataset_info, model_hypers=new_model_hypers) + else: logging.info("Starting training from scratch") - model = Model(hypers["model"], dataset_info) trainer = Trainer(hypers["training"]) + model = trainer.setup(hypers["model"], dataset_info) except Exception as e: raise ArchitectureError(e) from e diff --git a/src/metatrain/composition/model.py b/src/metatrain/composition/model.py index 875da50b94..c77c87a033 100644 --- a/src/metatrain/composition/model.py +++ b/src/metatrain/composition/model.py @@ -594,8 +594,6 @@ def export(self, metadata: Optional[ModelMetadata] = None) -> AtomisticModel: :return: An instance of :py:class:`metatomic.torch.AtomisticModel`. """ dtype = self.dummy_buffer.dtype - if dtype not in self.__supported_dtypes__: - raise ValueError(f"unsupported dtype {dtype} for composition model") self.to(dtype) self.weights_to(torch.device("cpu"), torch.float64) diff --git a/src/metatrain/pet/model.py b/src/metatrain/pet/model.py index db3a53f43a..30d2e18b9d 100644 --- a/src/metatrain/pet/model.py +++ b/src/metatrain/pet/model.py @@ -6,7 +6,6 @@ import metatensor.torch as mts import torch from metatensor.torch import Labels, TensorBlock, TensorMap -from metatensor.torch.operations._add import _add_block_block from metatomic.torch import ( AtomisticModel, ModelCapabilities, @@ -16,20 +15,13 @@ System, ) -from metatrain.composition import CompositionModel -from metatrain.scaler import Scaler from metatrain.utils.abc import ModelInterface -from metatrain.utils.additive import ZBL from metatrain.utils.architectures import get_default_hypers from metatrain.utils.data import DatasetInfo, TargetInfo from metatrain.utils.data.atom_pair_helpers import ( check_no_atom_pair_targets, get_pair_sample_labels, ) -from metatrain.utils.data.atomic_basis_helpers import ( - densify_atomic_basis_dataset_info, - sparsify_atomic_basis_target, -) from metatrain.utils.dtype import dtype_to_str from metatrain.utils.hypers import raise_if_hypers_mismatch from metatrain.utils.long_range import DummyLongRangeFeaturizer, LongRangeFeaturizer @@ -46,7 +38,7 @@ prepare_diagnostic_handles, standardize_featurizer_input_tensor, ) -from .modules.finetuning import apply_finetuning_strategy, compute_stale_targets +from .modules.finetuning import apply_finetuning_strategy from .modules.structures import concatenate_structures @@ -133,17 +125,13 @@ def __init__(self, hypers: ModelHypers, dataset_info: DatasetInfo) -> None: ), } - # Modified dataset_info with the targets as they will be seen by PET - # during training. - train_dataset_info = self._train_dataset_info(dataset_info) - self.output_shapes: Dict[str, Dict[str, List[int]]] = {} self.key_labels: Dict[str, Labels] = {} self.property_labels: Dict[str, List[Labels]] = {} self.component_labels: Dict[str, List[List[Labels]]] = {} self.target_names: List[str] = [] self.last_layer_parameter_names: Dict[str, List[str]] = {} # for LLPR - for target_name, target_info in train_dataset_info.targets.items(): + for target_name, target_info in dataset_info.targets.items(): self.target_names.append(target_name) self._add_output(target_name, target_info) @@ -168,36 +156,6 @@ def __init__(self, hypers: ModelHypers, dataset_info: DatasetInfo) -> None: self.long_range = False self.long_range_featurizer = DummyLongRangeFeaturizer() # for torchscript - # additive models: these are handled by the trainer at training - # time, and they are added to the output at evaluation time - composition_model = CompositionModel.from_valid_targets( - dataset_info, self.atomic_types - ) - additive_models = [composition_model] - - # Adds the ZBL repulsion model if requested - if self.hypers["zbl"]: - zbl_targets = { - target_name: target_info - for target_name, target_info in train_dataset_info.targets.items() - if ZBL.is_valid_target(target_name, target_info) - } - additive_models.append( - ZBL( - {}, - dataset_info=DatasetInfo( - length_unit=train_dataset_info.length_unit, - atomic_types=self.atomic_types, - targets=zbl_targets, - ), - ) - ) - self.additive_models = torch.nn.ModuleList(additive_models) - - # scaler: this is also handled by the trainer at training time - scaler_hypers = get_default_hypers("scaler")["model"] - self.scaler = Scaler(hypers=scaler_hypers, dataset_info=dataset_info) - self.single_label = Labels.single() self.finetune_config: Dict[str, Any] = {} @@ -225,53 +183,19 @@ def restart( for key, value in merged_info.targets.items() if key not in self.dataset_info.targets } - self.has_new_targets = len(new_targets) > 0 - - # Targets that were present before this run but are not part of the current - # run's dataset: with a backbone-altering finetuning method (full/lora), their - # heads are no longer meaningful and are dropped once training starts, by - # ``apply_finetuning_strategy`` (which decides based on the method). - stale_targets = compute_stale_targets( - self.dataset_info.targets, dataset_info.targets - ) - if len(new_atomic_types) > 0: raise ValueError( f"New atomic types found in the dataset: {new_atomic_types}. " "The PET model does not support adding new atomic types." ) - # Modified dataset_info with the targets as they will be seen by PET - # during training. - train_dataset_info = self._train_dataset_info(dataset_info) - # register new outputs as new last layers for target_name in new_targets: self.target_names.append(target_name) - self._add_output(target_name, train_dataset_info.targets[target_name]) + self._add_output(target_name, dataset_info.targets[target_name]) self.dataset_info = merged_info - # restart the composition and scaler models - self.additive_models[0] = self.additive_models[0].restart( - dataset_info=DatasetInfo( - length_unit=dataset_info.length_unit, - atomic_types=self.atomic_types, - targets={ - target_name: target_info - for target_name, target_info in dataset_info.targets.items() - if CompositionModel.is_valid_target(target_name, target_info) - }, - ), - ) - self.scaler = self.scaler.restart(dataset_info) - - # Actual removal (if any) is deferred to ``apply_finetuning_strategy`` - # (called later, once training starts), since ``inherit_heads`` needs these - # stale targets' heads to still be around to copy weights from, and only - # backbone-altering methods (``full``/``lora``) actually drop them. - self._stale_finetune_targets = stale_targets - return self def requested_neighbor_lists(self) -> List[NeighborListOptions]: @@ -598,76 +522,6 @@ def forward( h.remove() # ===== END DIAGNOSTIC-RELATED BLOCK - # **Post-processing (Evaluation Only)** - with torch.profiler.record_function("PET::post-processing"): - if not self.training: - # at evaluation, we also introduce the scaler and additive contributions - return_dict = self.scaler.apply_scales( - systems, - return_dict, - selected_atoms=selected_atoms, - use_per_target_scales=True, - use_per_property_scales=True, - ) - - # For atomic basis targets, sparsify to create blocks with "atom_type" - # in the key dimensions, and ensure properties are unpadded. This is - # done before adding the additive contributions, which are also - # sparsified (by the additive models themselves, in eval mode). - for k in atomic_predictions_dict.keys(): - if self.dataset_info.targets[k].is_atomic_basis: - return_dict[k] = sparsify_atomic_basis_target( - systems, - return_dict[k], - self.dataset_info.targets[k].layout, - species, - ) - - for additive_model in self.additive_models: - outputs_for_additive_model: Dict[str, ModelOutput] = {} - for name, output in outputs.items(): - if name in additive_model.outputs: - outputs_for_additive_model[name] = output - additive_contributions = additive_model( - systems, - outputs_for_additive_model, - selected_atoms, - ) - for name in additive_contributions: - # TODO: uncomment this after metatensor.torch.add - # is updated to handle sparse sums - # return_dict[name] = metatensor.torch.add( - # return_dict[name], - # additive_contributions[name].to( - # device=return_dict[name].device, - # dtype=return_dict[name].dtype - # ), - # ) - # TODO: "manual" sparse sum: update to metatensor.torch.add - # after sparse sum is implemented in metatensor.operations - output_blocks: List[TensorBlock] = [] - for k, b in return_dict[name].items(): - if k in additive_contributions[name].keys: - output_blocks.append( - _add_block_block( - b, - additive_contributions[name] - .block(k) - .to(device=b.device, dtype=b.dtype), - ) - ) - else: - output_blocks.append( - TensorBlock( - values=b.values, - samples=b.samples, - components=b.components, - properties=b.properties, - ) - ) - return_dict[name] = TensorMap( - return_dict[name].keys, output_blocks - ) return return_dict @@ -989,8 +843,6 @@ def load_checkpoint( next(state_dict_iter) # skip the species_to_species_index dtype = next(state_dict_iter).dtype model.to(dtype).load_state_dict(model_state_dict) - model.additive_models[0].sync_tensor_maps() - model.scaler.sync_tensor_maps() # Loading the metadata from the checkpoint model.metadata = merge_metadata(model.metadata, checkpoint.get("metadata")) @@ -1007,20 +859,10 @@ def export(self, metadata: Optional[ModelMetadata] = None) -> AtomisticModel: # float64 self.to(dtype) - # Additionally, the composition model contains some `TensorMap`s that cannot - # be registered correctly with Pytorch. This function moves them: - self.additive_models[0].weights_to(torch.device("cpu"), torch.float64) - - interaction_ranges = [self.num_gnn_layers * self.cutoff] - for additive_model in self.additive_models: - if hasattr(additive_model, "cutoff_radius"): - interaction_ranges.append(additive_model.cutoff_radius) - interaction_range = max(interaction_ranges) - capabilities = ModelCapabilities( outputs=self.outputs, atomic_types=self.atomic_types, - interaction_range=interaction_range, + interaction_range=self.num_gnn_layers * self.cutoff, length_unit=self.dataset_info.length_unit, supported_devices=self.__supported_devices__, dtype=dtype_to_str(dtype), @@ -1030,18 +872,6 @@ def export(self, metadata: Optional[ModelMetadata] = None) -> AtomisticModel: return AtomisticModel(self.eval(), metadata, capabilities) - def _train_dataset_info(self, dataset_info: DatasetInfo) -> DatasetInfo: - """Converts the original dataset info to one corresponding to what PET - will see during training, which depends on transforms applied to the - targets during data loading. - - :param dataset_info: Original dataset info describing the targets as - they are in the raw data. - :return: Modified dataset info describing the targets as they will be - seen by PET during training. - """ - return densify_atomic_basis_dataset_info(dataset_info) - def _add_output(self, target_name: str, target_info: TargetInfo) -> None: """ Register a new output target by creating corresponding heads and last layers. @@ -1110,6 +940,8 @@ def remove_output(self, target_name: str) -> None: self.key_labels.pop(target_name, None) self.component_labels.pop(target_name, None) self.property_labels.pop(target_name, None) + self.target_names.remove(target_name) + self.dataset_info.targets.pop(target_name, None) def _move_labels_to_device(self, device: torch.device) -> None: self.single_label = self.single_label.to(device) diff --git a/src/metatrain/pet/modules/finetuning.py b/src/metatrain/pet/modules/finetuning.py index f29f85e375..e2da482916 100644 --- a/src/metatrain/pet/modules/finetuning.py +++ b/src/metatrain/pet/modules/finetuning.py @@ -182,7 +182,10 @@ def _add_backend_prefix(model: nn.Module, module_names: list[str]) -> list[str]: def apply_finetuning_strategy( - model: nn.Module, strategy: FinetuneHypers, apply_inherit_heads: bool = True + model: nn.Module, + strategy: FinetuneHypers, + apply_inherit_heads: bool = True, + stale_targets: Optional[list[str]] = None, ) -> nn.Module: """ Apply the specified finetuning strategy to the model. @@ -208,6 +211,11 @@ def apply_finetuning_strategy( ``False``: by then, any inherited-from source target may already have been pruned, so redoing the copy would fail (or silently clobber trained weights if the source were still around). + :param stale_targets: List of targets that were present in the model before + but will not be a part of the dataset in the next training run. These + targets will be removed from the model unless the finetuning strategy + is "heads", in which case the backbone is frozen and therefore the outputs + for the targets are still meaningful. :return: The modified model with the finetuning strategy applied. """ @@ -288,7 +296,10 @@ def apply_finetuning_strategy( "are: 'full', 'lora', 'heads'." ) - model.finetune_config = strategy + if hasattr(model, "finetune_config"): + model.finetune_config = strategy + else: + model.model.finetune_config = strategy inherit_heads_config = strategy["inherit_heads"] if apply_inherit_heads and inherit_heads_config: @@ -302,19 +313,9 @@ def apply_finetuning_strategy( # than removing right away, since it runs before this function (and before # ``inherit_heads`` above needs the stale heads to still be present); with # ``heads`` the backbone is unchanged, so stale targets are left alone here. - stale_targets = getattr(model, "_stale_finetune_targets", None) if stale_targets and strategy["method"] in ("full", "lora"): for target_name in stale_targets: model.remove_output(target_name) - if target_name in model.target_names: - model.target_names.remove(target_name) - model.dataset_info.targets.pop(target_name, None) - for additive_model in model.additive_models: - if target_name in additive_model.outputs: - additive_model.remove_output(target_name) - if target_name in model.scaler.outputs: - model.scaler.remove_output(target_name) - model._stale_finetune_targets = [] return model diff --git a/src/metatrain/pet/tests/test_regression.py b/src/metatrain/pet/tests/test_regression.py index 90ae3acc10..1a1c25d437 100644 --- a/src/metatrain/pet/tests/test_regression.py +++ b/src/metatrain/pet/tests/test_regression.py @@ -127,8 +127,8 @@ def test_regression_energies_forces_train(device): dataset_info = DatasetInfo( length_unit="Angstrom", atomic_types=[6], targets=target_info_dict ) - model = PET(MODEL_HYPERS, dataset_info) trainer = Trainer(hypers["training"]) + model = trainer.setup(MODEL_HYPERS, dataset_info) trainer.train( model=model, dtype=torch.float32, @@ -283,8 +283,8 @@ def test_regression_energy_non_conservative_stress(batch_size): model_hypers = copy.deepcopy(MODEL_HYPERS) model_hypers["zbl"] = False - model = PET(model_hypers, dataset_info) trainer = Trainer(hypers["training"]) + model = trainer.setup(model_hypers, dataset_info) # num_epochs=0 keeps this test focused on pre-training scaler fitting. trainer.train( model=model, @@ -402,12 +402,10 @@ def test_regression_train_spherical(device): targets=target_info_dict, extra_data=extra_data_info, ) - model = PET(MODEL_HYPERS, dataset_info) - requested_neighbor_lists = get_requested_neighbor_lists(model) - hypers["training"]["num_epochs"] = 1 hypers["training"]["num_workers"] = 0 # for reproducibility trainer = Trainer(hypers["training"]) + model = trainer.setup(MODEL_HYPERS, dataset_info) trainer.train( model=model, dtype=torch.float32, @@ -417,6 +415,8 @@ def test_regression_train_spherical(device): checkpoint_dir=".", ) + requested_neighbor_lists = get_requested_neighbor_lists(model) + # Predict on the first five systems systems = [sample["system"] for sample in dataset] systems = [system.to(torch.float32, device) for system in systems] diff --git a/src/metatrain/pet/trainer.py b/src/metatrain/pet/trainer.py index 567ed20947..0001905a5b 100644 --- a/src/metatrain/pet/trainer.py +++ b/src/metatrain/pet/trainer.py @@ -8,15 +8,17 @@ from torch.optim.lr_scheduler import LambdaLR from torch.utils.data import DistributedSampler -from metatrain.composition import train_or_load_composition_model -from metatrain.scaler import train_or_load_scaler +from metatrain.composition import CompositionModel, train_or_load_composition_model +from metatrain.scaler import Scaler, train_or_load_scaler from metatrain.utils.abc import ModelInterface, TrainerInterface -from metatrain.utils.additive import get_remove_additive_transform +from metatrain.utils.additive import ZBL, get_remove_additive_transform +from metatrain.utils.architectures import get_default_hypers from metatrain.utils.augmentation import O3Augmenter from metatrain.utils.data import ( CollateFn, CombinedDataLoader, Dataset, + DatasetInfo, build_train_dataloaders, build_val_dataloaders, get_num_workers, @@ -24,6 +26,7 @@ validate_num_workers, ) from metatrain.utils.data.atomic_basis_helpers import ( + densify_atomic_basis_dataset_info, get_prepare_atomic_basis_targets_transform, ) from metatrain.utils.distributed.distributed_data_parallel import ( @@ -46,11 +49,12 @@ from metatrain.utils.scaler import get_remove_scale_transform from metatrain.utils.system_data import get_system_data_transform from metatrain.utils.transfer import batch_to +from metatrain.utils.wrapper import MetatrainWrapper from . import checkpoints -from .documentation import TrainerHypers +from .documentation import ModelHypers, TrainerHypers from .model import PET -from .modules.finetuning import apply_finetuning_strategy +from .modules.finetuning import apply_finetuning_strategy, compute_stale_targets def get_scheduler( @@ -89,6 +93,9 @@ def lr_lambda(current_step: int) -> float: class Trainer(TrainerInterface[TrainerHypers]): __checkpoint_version__ = 15 + has_new_targets: bool + _stale_finetune_targets: list[str] + def __init__(self, hypers: TrainerHypers) -> None: super().__init__(hypers) @@ -100,9 +107,94 @@ def __init__(self, hypers: TrainerHypers) -> None: self.best_model_state_dict: Optional[Dict[str, Any]] = None self.best_optimizer_state_dict: Optional[Dict[str, Any]] = None + self.has_new_targets = False + self._stale_finetune_targets = [] + + def setup( + self, model_hypers: ModelHypers, dataset_info: DatasetInfo + ) -> ModelInterface: + self.dataset_info = dataset_info + model_dataset_info = densify_atomic_basis_dataset_info(dataset_info) + + model = PET(hypers=model_hypers, dataset_info=model_dataset_info) + + # Set up additive models + composition_model = CompositionModel.from_valid_targets( + dataset_info, dataset_info.atomic_types + ) + additive_models = [composition_model] + + # Adds the ZBL repulsion model if requested + if model_hypers["zbl"]: + zbl_targets = { + target_name: target_info + for target_name, target_info in model_dataset_info.targets.items() + if ZBL.is_valid_target(target_name, target_info) + } + additive_models.append( + ZBL( + {}, + dataset_info=DatasetInfo( + length_unit=model_dataset_info.length_unit, + atomic_types=self.atomic_types, + targets=zbl_targets, + ), + ) + ) + additive_models = torch.nn.ModuleList(additive_models) + + # Initialize scaler + scaler_hypers = get_default_hypers("scaler")["model"] + scaler = Scaler(hypers=scaler_hypers, dataset_info=dataset_info) + + return MetatrainWrapper( + hypers=dict( + model=model, + additive_models=additive_models, + scaler=scaler, + ), + dataset_info=dataset_info + ) + + def restart( + self, + model: MetatrainWrapper, + dataset_info: DatasetInfo, + model_hypers: ModelHypers + ) -> MetatrainWrapper: + + # merge old and new dataset info + merged_info = model.dataset_info.union(dataset_info) + new_targets = { + key: value + for key, value in merged_info.targets.items() + if key not in model.dataset_info.targets + } + self.has_new_targets = len(new_targets) > 0 + + # Targets that were present before this run but are not part of the current + # run's dataset: with a backbone-altering finetuning method (full/lora), their + # heads are no longer meaningful and are dropped once training starts, by + # ``apply_finetuning_strategy`` (which decides based on the method). + stale_targets = compute_stale_targets( + model.dataset_info.targets, dataset_info.targets + ) + + # Actual removal (if any) is deferred to ``apply_finetuning_strategy`` + # (called later, once training starts), since ``inherit_heads`` needs these + # stale targets' heads to still be around to copy weights from, and only + # backbone-altering methods (``full``/``lora``) actually drop them. + self._stale_finetune_targets = stale_targets + + model_merged_info = densify_atomic_basis_dataset_info(merged_info) + + model.restart(model_merged_info, model_hypers=model_hypers) + + return model + def train( self, - model: PET, + model: MetatrainWrapper, dtype: torch.dtype, devices: List[torch.device], train_datasets: List[Union[Dataset, torch.utils.data.Subset]], @@ -152,6 +244,7 @@ def train( model, self.hypers["finetune"], apply_inherit_heads=is_fresh_finetune_start, + stale_targets=self._stale_finetune_targets, ) method = self.hypers["finetune"]["method"] num_params = sum(p.numel() for p in model.parameters()) @@ -378,7 +471,7 @@ def train( if self.optimizer_state_dict is not None: # try to load the optimizer state dict, but this is only possible # if there are no new targets in the model (new parameters) - if not (model.module if is_distributed else model).has_new_targets: + if not self.has_new_targets: optimizer.load_state_dict(self.optimizer_state_dict) # Create a learning rate scheduler @@ -386,7 +479,7 @@ def train( if self.scheduler_state_dict is not None: # same as the optimizer, try to load the scheduler state dict - if not (model.module if is_distributed else model).has_new_targets: + if not self.has_new_targets: lr_scheduler.load_state_dict(self.scheduler_state_dict) per_structure_targets = self.hypers["per_structure_targets"] @@ -697,7 +790,7 @@ def train( def save_checkpoint(self, model: ModelInterface, path: Union[str, Path]) -> None: checkpoint = model.get_checkpoint() if self.best_model_state_dict is not None: - self.best_model_state_dict["finetune_config"] = model.finetune_config + self.best_model_state_dict["finetune_config"] = model.model.finetune_config checkpoint.update( { "trainer_ckpt_version": self.__checkpoint_version__, diff --git a/src/metatrain/scaler/model.py b/src/metatrain/scaler/model.py index 9e3e383442..8073adf100 100644 --- a/src/metatrain/scaler/model.py +++ b/src/metatrain/scaler/model.py @@ -141,21 +141,22 @@ def restart( self.target_infos = { target_name: dense_new_targets[target_name] for target_name in merged_info.targets - if target_name not in self.dataset_info.targets } - self.dataset_info = merged_info - # register new outputs self.new_outputs = [] buffer_names = [n for n, _ in self.named_buffers()] for target_name, target_info in self.target_infos.items(): + if target_name in self.dataset_info.targets: + continue if target_name + "_scaler_buffer" in buffer_names: continue self.new_outputs.append(target_name) self.model.add_output(target_name, target_info.layout) self._add_output(target_name, target_info) + self.dataset_info = merged_info + return self def forward( @@ -440,6 +441,9 @@ def remove_output(self, target_name: str) -> None: if hasattr(self, buffer_name): delattr(self, buffer_name) + print(list(self.dataset_info.targets.keys())) + print(list(self.target_infos.keys())) + def scales_to(self, device: torch.device, dtype: torch.dtype) -> None: if len(self.model.scales) != 0: if self.model.scales[list(self.model.scales.keys())[0]].device != device: @@ -581,8 +585,6 @@ def upgrade_checkpoint(cls, checkpoint: Dict) -> Dict: def export(self, metadata: Optional[ModelMetadata] = None) -> AtomisticModel: dtype = self.dummy_buffer.dtype - if dtype not in self.__supported_dtypes__: - raise ValueError(f"unsupported dtype {dtype} for scaler") self.to(dtype) self.scales_to(torch.device("cpu"), torch.float64) diff --git a/src/metatrain/utils/io.py b/src/metatrain/utils/io.py index 1926a03a98..955b379e30 100644 --- a/src/metatrain/utils/io.py +++ b/src/metatrain/utils/io.py @@ -208,6 +208,11 @@ def model_from_checkpoint( """ architecture_name = checkpoint["architecture_name"] + + if architecture_name == "metatrain_wrapper": + from .wrapper import MetatrainWrapper + return MetatrainWrapper.load_checkpoint(checkpoint, context=context) + if architecture_name not in find_all_architectures(): raise ValueError( f"Checkpoint architecture '{architecture_name}' not found " @@ -279,6 +284,10 @@ def trainer_from_checkpoint( """ architecture_name = checkpoint["architecture_name"] + + if architecture_name == "metatrain_wrapper": + architecture_name = checkpoint["model"]["architecture_name"] + if architecture_name not in find_all_architectures(): raise ValueError( f"Checkpoint architecture '{architecture_name}' not found " diff --git a/src/metatrain/utils/testing/training.py b/src/metatrain/utils/testing/training.py index ae321004df..e3d8b8dc02 100644 --- a/src/metatrain/utils/testing/training.py +++ b/src/metatrain/utils/testing/training.py @@ -61,8 +61,6 @@ def test_train( dataset_targets, dataset_path ) - model = self.model_cls(model_hypers, dataset_info) - hypers = copy.deepcopy(default_hypers) if "num_epochs" in hypers["training"]: hypers["training"]["num_epochs"] = 0 @@ -74,6 +72,7 @@ def test_train( hypers["training"]["loss"] = loss_conf trainer = self.trainer_cls(hypers["training"]) + model = trainer.setup(model_hypers, dataset_info) trainer.train( model=model, dtype=dtype, @@ -167,8 +166,6 @@ def test_train_atomic_basis_target( dataset_info = dataset_info_spherical_atomic_basis dataset = self._atomic_basis_dataset(dataset_info, tmp_path / "dataset.zip") - model = self.model_cls(model_hypers, dataset_info) - hypers = copy.deepcopy(default_hypers) if "num_epochs" in hypers["training"]: hypers["training"]["num_epochs"] = 0 @@ -181,7 +178,10 @@ def test_train_atomic_basis_target( OmegaConf.resolve(loss_conf) hypers["training"]["loss"] = loss_conf - self.trainer_cls(hypers["training"]).train( + trainer = self.trainer_cls(hypers["training"]) + model = trainer.setup(model_hypers, dataset_info) + + trainer.train( model=model, dtype=dtype, devices=[torch.device("cpu")], @@ -216,8 +216,6 @@ def test_continue( dataset_targets, dataset_path ) - model = self.model_cls(model_hypers, dataset_info) - hypers = copy.deepcopy(default_hypers) hypers["training"]["num_epochs"] = 0 loss_conf = OmegaConf.create( @@ -227,6 +225,7 @@ def test_continue( hypers["training"]["loss"] = loss_conf trainer = self.trainer_cls(hypers["training"]) + model = trainer.setup(model_hypers, dataset_info) trainer.train( model=model, dtype=torch.float32, @@ -240,7 +239,7 @@ def test_continue( checkpoint = torch.load("tmp.ckpt", weights_only=False, map_location="cpu") model_after = model_from_checkpoint(checkpoint, context="restart") - assert isinstance(model_after, self.model_cls) + # assert isinstance(model_after, self.model_cls) model_after.restart(model.dataset_info) hypers["training"]["num_epochs"] = 0 @@ -264,11 +263,13 @@ def test_continue( model.eval() model_after.eval() + outputs = model.supported_outputs() + output_before = model( - systems[:5], {k: model.outputs[k] for k in dataset_targets} + systems[:5], {k: outputs[k] for k in dataset_targets} ) output_after = model_after( - systems[:5], {k: model_after.outputs[k] for k in dataset_targets} + systems[:5], {k: outputs[k] for k in dataset_targets} ) # For each target, check that outputs are the same after loading @@ -327,8 +328,6 @@ def test_continue_restart_num_epochs( dataset_targets, dataset_path ) - model = self.model_cls(model_hypers, dataset_info) - hypers = copy.deepcopy(default_hypers) hypers["training"]["num_epochs"] = 2 loss_conf = OmegaConf.create( @@ -338,6 +337,7 @@ def test_continue_restart_num_epochs( hypers["training"]["loss"] = loss_conf trainer = self.trainer_cls(hypers["training"]) + model = trainer.setup(model_hypers, dataset_info) trainer.train( model=model, dtype=torch.float32, @@ -352,7 +352,7 @@ def test_continue_restart_num_epochs( checkpoint = torch.load("tmp.ckpt", weights_only=False, map_location="cpu") model_after = model_from_checkpoint(checkpoint, context="restart") - assert isinstance(model_after, self.model_cls) + # assert isinstance(model_after, self.model_cls) model_after.restart(model.dataset_info) hypers["training"]["num_epochs"] = 4 # modify max num epochs to 4 @@ -395,8 +395,6 @@ def test_continue_finetune_num_epochs( dataset_targets, dataset_path ) - model = self.model_cls(model_hypers, dataset_info) - hypers = copy.deepcopy(default_hypers) hypers["training"]["num_epochs"] = 2 loss_conf = OmegaConf.create( @@ -406,6 +404,7 @@ def test_continue_finetune_num_epochs( hypers["training"]["loss"] = loss_conf trainer = self.trainer_cls(hypers["training"]) + model = trainer.setup(model_hypers, dataset_info) trainer.train( model=model, dtype=torch.float32, @@ -420,7 +419,7 @@ def test_continue_finetune_num_epochs( checkpoint = torch.load("tmp.ckpt", weights_only=False, map_location="cpu") model_after = model_from_checkpoint(checkpoint, context="finetune") - assert isinstance(model_after, self.model_cls) + # assert isinstance(model_after, self.model_cls) model_after.restart(model.dataset_info) hypers["training"]["num_epochs"] = 1 # modify max num epochs to 1 diff --git a/src/metatrain/utils/wrapper.py b/src/metatrain/utils/wrapper.py new file mode 100644 index 0000000000..6d7efb375b --- /dev/null +++ b/src/metatrain/utils/wrapper.py @@ -0,0 +1,360 @@ +from typing import Any, Dict, List, Literal, Optional, NotRequired, Union + +import torch +from metatensor.torch import Labels, TensorBlock, TensorMap +from metatensor.torch.operations._add import _add_block_block +from metatomic.torch import ( + AtomisticModel, + ModelMetadata, + ModelOutput, + ModelCapabilities, + ModelEvaluationOptions, + System, + NeighborListOptions, +) +from typing_extensions import TypedDict + +from metatrain.scaler import Scaler +from metatrain.utils.abc import ModelInterface +from metatrain.utils.architectures import import_architecture +from metatrain.utils.data import DatasetInfo +from metatrain.utils.data.atomic_basis_helpers import ( + sparsify_atomic_basis_target, +) +from metatrain.utils.dtype import dtype_to_str + + +class WrapperHypers(TypedDict): + """Hypers to initialize the model. + + These are only use on a first initialization. + When loading a checkpoint, the hypers are ignored + and instead the components of the model are loaded. + """ + + # Models passed directly + model: NotRequired[dict] + additive_models: NotRequired[list[dict]] + scaler: NotRequired[dict] + + +class MetatrainWrapper(ModelInterface[WrapperHypers]): + __checkpoint_version__ = 1 + __supported_devices__ = ["cuda", "cpu"] + __supported_dtypes__ = [torch.float32, torch.float64] + __default_metadata__ = ModelMetadata() + component_labels: Dict[str, List[List[Labels]]] + NUM_FEATURE_TYPES: int = 2 # node + edge features + + def __init__( + self, + hypers: WrapperHypers, + dataset_info: DatasetInfo, + ): + super().__init__(hypers=hypers, dataset_info=dataset_info, metadata=self.__default_metadata__) + + if "model" in hypers: + self.model = hypers["model"] + self.additive_models = torch.nn.ModuleList(hypers["additive_models"]) + self.scaler = hypers["scaler"] + + def forward( + self, + systems: List[System], + outputs: Dict[str, ModelOutput], + selected_atoms: Optional[Labels] = None, + ) -> Dict[str, TensorMap]: + return_dict = self.model( + systems, outputs, selected_atoms=selected_atoms + ) + + with torch.profiler.record_function("MTT_WRAPPER::post-processing"): + if not self.training: + # at evaluation, we also introduce the scaler and additive contributions + return_dict = self.scaler.apply_scales( + systems, + return_dict, + selected_atoms=selected_atoms, + use_per_target_scales=True, + use_per_property_scales=True, + ) + + # For atomic basis targets, sparsify to create blocks with "atom_type" + # in the key dimensions, and ensure properties are unpadded. This is + # done before adding the additive contributions, which are also + # sparsified (by the additive models themselves, in eval mode). + # for k in atomic_predictions_dict.keys(): + # if self.model.dataset_info.targets[k].is_atomic_basis: + # return_dict[k] = sparsify_atomic_basis_target( + # systems, + # return_dict[k], + # self.dataset_info.targets[k].layout, + # species, + # ) + + for additive_model in self.additive_models: + outputs_for_additive_model: Dict[str, ModelOutput] = {} + for name, output in outputs.items(): + if name in additive_model.outputs: + outputs_for_additive_model[name] = output + additive_contributions = additive_model( + systems, + outputs_for_additive_model, + selected_atoms, + ) + for name in additive_contributions: + # TODO: uncomment this after metatensor.torch.add + # is updated to handle sparse sums + # return_dict[name] = metatensor.torch.add( + # return_dict[name], + # additive_contributions[name].to( + # device=return_dict[name].device, + # dtype=return_dict[name].dtype + # ), + # ) + # TODO: "manual" sparse sum: update to metatensor.torch.add + # after sparse sum is implemented in metatensor.operations + output_blocks: List[TensorBlock] = [] + for k, b in return_dict[name].items(): + if k in additive_contributions[name].keys: + output_blocks.append( + _add_block_block( + b, + additive_contributions[name] + .block(k) + .to(device=b.device, dtype=b.dtype), + ) + ) + else: + output_blocks.append( + TensorBlock( + values=b.values, + samples=b.samples, + components=b.components, + properties=b.properties, + ) + ) + return_dict[name] = TensorMap( + return_dict[name].keys, output_blocks + ) + return return_dict + + def requested_inputs(self) -> Dict[str, ModelOutput]: + requested_inputs = {} + + def _add_model_requested_inputs(model: ModelInterface): + if not hasattr(model, "requested_inputs"): + return + for name, output in model.requested_inputs().items(): + if name not in requested_inputs: + requested_inputs[name] = output + + for additive_model in self.additive_models: + _add_model_requested_inputs(additive_model) + + _add_model_requested_inputs(self.scaler) + _add_model_requested_inputs(self.model) + + return requested_inputs + + def requested_neighbor_lists(self) -> list[NeighborListOptions]: + requested_neighbor_lists = [] + + def _add_model_requested_neighbor_lists(model: ModelInterface): + if not hasattr(model, "requested_neighbor_lists"): + return + for nl in model.requested_neighbor_lists(): + if nl not in requested_neighbor_lists: + requested_neighbor_lists.append(nl) + + for additive_model in self.additive_models: + _add_model_requested_neighbor_lists(additive_model) + + _add_model_requested_neighbor_lists(self.scaler) + _add_model_requested_neighbor_lists(self.model) + + return requested_neighbor_lists + + def supported_outputs(self) -> Dict[str, ModelOutput]: + return self.model.supported_outputs() + + def export( + self, + metadata: Optional[ModelMetadata] = None, + ) -> AtomisticModel: + """ + Turn this model into an instance of + :py:class:`metatomic.torch.MetatensorAtomisticModel`, containing the model + itself, a definition of the model capabilities and some metadata about the + model. + + :param metadata: additional metadata to add in the model as specified by the + user. + + :return: An instance of :py:class:`metatomic.torch.MetatensorAtomisticModel` + """ + # Export all models, making sure that the additive models and scaler have + # the same dtype as the main model. + model = self.model.export(metadata) + dtype = getattr(torch, model.capabilities().dtype) + additive_models = [model.to(dtype).export(metadata) for model in self.additive_models] + scaler = self.scaler.to(dtype).export(metadata) + + # Get a list of the capabilities of each model + all_capabilities = [model.capabilities()] + [ + model.capabilities() for model in additive_models + ] + [scaler.capabilities()] + + # The interaction range of the model is the maximum interaction range + # of all the models involved. + all_interaction_ranges = [cap.interaction_range for cap in all_capabilities if cap.interaction_range is not None] + interaction_range = max(all_interaction_ranges) if all_interaction_ranges else None + + all_supported_devices = [cap.supported_devices for cap in all_capabilities] + # Get the intersection of all supported devices + supported_devices = set(all_supported_devices[0]) + for devices in all_supported_devices[1:]: + supported_devices.intersection_update(devices) + + # Build the wrapper model again with the exported modules. + to_export = self.__class__( + hypers=dict( + model=model.module, + additive_models=[model.module for model in additive_models], + scaler=scaler.module, + ), + dataset_info=self.dataset_info, + ) + + if metadata is None: + metadata = ModelMetadata() + + capabilities = ModelCapabilities( + outputs=self.supported_outputs(), + atomic_types=self.dataset_info.atomic_types, + interaction_range=interaction_range, + length_unit=self.dataset_info.length_unit, + supported_devices=self.model.__supported_devices__, + dtype=dtype_to_str(dtype), + ) + + return AtomisticModel(to_export.eval(), metadata, capabilities) + + @classmethod + def load_checkpoint( + cls, + checkpoint: Dict[str, Any], + context: Literal["restart", "finetune", "export"], + ) -> "ModelInterface": + hypers = {} + for k in ["model", "scaler"]: + subcheckpoint = checkpoint[k] + + architecture_name = subcheckpoint["architecture_name"] + model_cls = import_architecture(architecture_name).__model__ + hypers[k] = model_cls.load_checkpoint(subcheckpoint, context) + + hypers["additive_models"] = [] + for additive_model_checkpoint in checkpoint["additive_models"]: + architecture_name = additive_model_checkpoint["architecture_name"] + model_cls = import_architecture(architecture_name).__model__ + hypers["additive_models"].append( + model_cls.load_checkpoint(additive_model_checkpoint, context) + ) + + return cls(hypers=hypers, dataset_info=checkpoint["dataset_info"]) + + @classmethod + def upgrade_checkpoint(cls, checkpoint: Dict["str", Any]) -> Dict["str", Any]: + """ + Upgrade the checkpoint to the current version of the model. + + :param checkpoint: Checkpoint's state dictionary. + + :raises RuntimeError: if the checkpoint cannot be upgraded to the current + version of the model. + + :return: The upgraded checkpoint. + """ + for k in ["model", "scaler"]: + subcheckpoint = checkpoint[k] + + architecture_name = subcheckpoint["architecture_name"] + model_cls = import_architecture(architecture_name).__model__ + checkpoint[k] = model_cls.upgrade_checkpoint(subcheckpoint) + + for i, additive_model_checkpoint in enumerate(checkpoint["additive_models"]): + architecture_name = additive_model_checkpoint["architecture_name"] + model_cls = import_architecture(architecture_name).__model__ + checkpoint["additive_models"][i] = model_cls.upgrade_checkpoint( + additive_model_checkpoint + ) + + return checkpoint + + def get_checkpoint(self) -> Dict[str, Any]: + """ + Get the checkpoint of the model. This should contain all the information + needed by `load_checkpoint` to recreate the same model instance. + + :return: The model's checkpoint. + """ + checkpoint = { + "architecture_name": "metatrain_wrapper", + "model_ckpt_version": self.__checkpoint_version__, + "metadata": self.metadata, + "dataset_info": self.dataset_info, + "model": self.model.get_checkpoint(), + "additive_models": [m.get_checkpoint() for m in self.additive_models], + "scaler": self.scaler.get_checkpoint(), + } + return checkpoint + + def restart(self, dataset_info, model_hypers = None): + """ + Restart the model with new dataset_info and model_hypers. + This is used when the model is loaded from a checkpoint and the dataset_info + has changed (e.g. when finetuning on a new dataset). + + :param dataset_info: The new dataset_info to use. + :param model_hypers: The new model_hypers to use. If None, the current + hypers will be used. + """ + self.model.restart(dataset_info, model_hypers) + + composition_model = self.additive_models[0] + self.additive_models[0] = composition_model.restart( + dataset_info=DatasetInfo( + length_unit=dataset_info.length_unit, + atomic_types=dataset_info.atomic_types, + targets={ + target_name: target_info + for target_name, target_info in dataset_info.targets.items() + if composition_model.is_valid_target(target_name, target_info) + }, + ), + ) + + self.scaler = self.scaler.restart(dataset_info) + + self.dataset_info = dataset_info + + def remove_output(self, target_name: str) -> None: + """ + Remove a previously registered output target. + + Used to drop targets whose heads are no longer meaningful after a + backbone-altering fine-tuning run (``full``/``lora``). + + :param target_name: Name of the target to remove. + """ + self.model.remove_output(target_name) + + for additive_model in self.additive_models: + if target_name in additive_model.supported_outputs(): + additive_model.remove_output(target_name) + + if target_name in self.scaler.supported_outputs(): + self.scaler.remove_output(target_name) + + self.dataset_info.targets.pop(target_name, None) From 1f4231b4ff1c21b900648b7e30eaec1725201bb8 Mon Sep 17 00:00:00 2001 From: Pol Febrer Date: Fri, 18 Sep 2026 08:39:41 +0200 Subject: [PATCH 2/3] remove prints --- src/metatrain/scaler/model.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/src/metatrain/scaler/model.py b/src/metatrain/scaler/model.py index 8073adf100..6fdd898ef4 100644 --- a/src/metatrain/scaler/model.py +++ b/src/metatrain/scaler/model.py @@ -441,9 +441,6 @@ def remove_output(self, target_name: str) -> None: if hasattr(self, buffer_name): delattr(self, buffer_name) - print(list(self.dataset_info.targets.keys())) - print(list(self.target_infos.keys())) - def scales_to(self, device: torch.device, dtype: torch.dtype) -> None: if len(self.model.scales) != 0: if self.model.scales[list(self.model.scales.keys())[0]].device != device: From 08fbad5b15ad1154ffb4747a09e0bcb1a1c9c505 Mon Sep 17 00:00:00 2001 From: Pol Febrer Calabozo Date: Tue, 22 Sep 2026 18:26:45 +0200 Subject: [PATCH 3/3] MetatrainModel is not ModelInterface, loading old checkpoints --- src/metatrain/composition/trainer.py | 4 +- src/metatrain/pet/checkpoints.py | 19 +++ src/metatrain/pet/model.py | 2 +- src/metatrain/pet/trainer.py | 22 ++- src/metatrain/scaler/trainer.py | 4 +- src/metatrain/utils/abc.py | 24 +++- src/metatrain/utils/io.py | 26 ++-- src/metatrain/utils/wrapper.py | 207 ++++++++++++++++----------- 8 files changed, 190 insertions(+), 118 deletions(-) diff --git a/src/metatrain/composition/trainer.py b/src/metatrain/composition/trainer.py index 2340e7884b..2acae4a2eb 100644 --- a/src/metatrain/composition/trainer.py +++ b/src/metatrain/composition/trainer.py @@ -29,10 +29,10 @@ from metatrain.utils.transfer import batch_to from . import checkpoints -from .documentation import TrainerHypers +from .documentation import TrainerHypers, ModelHypers -class Trainer(TrainerInterface[TrainerHypers]): +class Trainer(TrainerInterface[TrainerHypers, ModelHypers]): __checkpoint_version__ = 2 def __init__(self, hypers: TrainerHypers): diff --git a/src/metatrain/pet/checkpoints.py b/src/metatrain/pet/checkpoints.py index 21472b5ac7..0a40a40fd5 100644 --- a/src/metatrain/pet/checkpoints.py +++ b/src/metatrain/pet/checkpoints.py @@ -373,6 +373,25 @@ def model_update_v15_v16(checkpoint: dict) -> None: updated[name] = value checkpoint[key] = updated +def model_update_v16_v17(checkpoint: dict) -> None: + """ + Update a v16 checkpoint to v17. + + It removes the additive models and scaler from the model checkpoint, + as this is now handled by the MetatrainModel wrapper. + + :param checkpoint: The checkpoint to update. + """ + removed_prefixes = ( + "additive_models.", + "scaler.", + ) + for key in ["model_state_dict", "best_model_state_dict"]: + if (state_dict := checkpoint.get(key)) is not None: + for k in list(state_dict): + for prefix in removed_prefixes: + if k.startswith(prefix): + state_dict.pop(k) ########################### # TRAINER ################# diff --git a/src/metatrain/pet/model.py b/src/metatrain/pet/model.py index 30d2e18b9d..635c5e3a28 100644 --- a/src/metatrain/pet/model.py +++ b/src/metatrain/pet/model.py @@ -57,7 +57,7 @@ class PET(ModelInterface[ModelHypers]): targets. """ - __checkpoint_version__ = 16 + __checkpoint_version__ = 17 __supported_devices__ = ["cuda", "cpu"] __supported_dtypes__ = [torch.float32, torch.float64] __default_metadata__ = ModelMetadata( diff --git a/src/metatrain/pet/trainer.py b/src/metatrain/pet/trainer.py index 0001905a5b..4813d6c2b5 100644 --- a/src/metatrain/pet/trainer.py +++ b/src/metatrain/pet/trainer.py @@ -49,7 +49,7 @@ from metatrain.utils.scaler import get_remove_scale_transform from metatrain.utils.system_data import get_system_data_transform from metatrain.utils.transfer import batch_to -from metatrain.utils.wrapper import MetatrainWrapper +from metatrain.utils.wrapper import MetatrainModel from . import checkpoints from .documentation import ModelHypers, TrainerHypers @@ -90,7 +90,7 @@ def lr_lambda(current_step: int) -> float: return scheduler -class Trainer(TrainerInterface[TrainerHypers]): +class Trainer(TrainerInterface[TrainerHypers, ModelHypers]): __checkpoint_version__ = 15 has_new_targets: bool @@ -112,7 +112,7 @@ def __init__(self, hypers: TrainerHypers) -> None: def setup( self, model_hypers: ModelHypers, dataset_info: DatasetInfo - ) -> ModelInterface: + ) -> MetatrainModel: self.dataset_info = dataset_info model_dataset_info = densify_atomic_basis_dataset_info(dataset_info) @@ -147,21 +147,19 @@ def setup( scaler_hypers = get_default_hypers("scaler")["model"] scaler = Scaler(hypers=scaler_hypers, dataset_info=dataset_info) - return MetatrainWrapper( - hypers=dict( - model=model, - additive_models=additive_models, - scaler=scaler, - ), + return MetatrainModel( + model=model, + additive_models=additive_models, + scaler=scaler, dataset_info=dataset_info ) def restart( self, - model: MetatrainWrapper, + model: MetatrainModel, dataset_info: DatasetInfo, model_hypers: ModelHypers - ) -> MetatrainWrapper: + ) -> MetatrainModel: # merge old and new dataset info merged_info = model.dataset_info.union(dataset_info) @@ -194,7 +192,7 @@ def restart( def train( self, - model: MetatrainWrapper, + model: MetatrainModel, dtype: torch.dtype, devices: List[torch.device], train_datasets: List[Union[Dataset, torch.utils.data.Subset]], diff --git a/src/metatrain/scaler/trainer.py b/src/metatrain/scaler/trainer.py index 4d3731f500..dd5e9f4c45 100644 --- a/src/metatrain/scaler/trainer.py +++ b/src/metatrain/scaler/trainer.py @@ -31,10 +31,10 @@ from metatrain.utils.per_atom import average_by_num_atoms from metatrain.utils.transfer import batch_to -from .documentation import TrainerHypers +from .documentation import TrainerHypers, ModelHypers -class Trainer(TrainerInterface[TrainerHypers]): +class Trainer(TrainerInterface[TrainerHypers, ModelHypers]): __checkpoint_version__ = 1 def __init__(self, hypers: TrainerHypers): diff --git a/src/metatrain/utils/abc.py b/src/metatrain/utils/abc.py index c59974588e..c9cdade289 100644 --- a/src/metatrain/utils/abc.py +++ b/src/metatrain/utils/abc.py @@ -21,12 +21,13 @@ ) from metatrain.utils.data.dataset import Dataset, DatasetInfo +#from metatrain.utils.wrapper import MetatrainModel +ModelHypersType = TypeVar("ModelHypersType") +TrainerHypersType = TypeVar("TrainerHypersType") -HypersType = TypeVar("HypersType") - -class ModelInterface(torch.nn.Module, Generic[HypersType], metaclass=ABCMeta): +class ModelInterface(torch.nn.Module, Generic[ModelHypersType], metaclass=ABCMeta): """ Abstract base class for a machine learning model in metatrain. @@ -70,7 +71,7 @@ class ModelInterface(torch.nn.Module, Generic[HypersType], metaclass=ABCMeta): """ def __init__( - self, hypers: HypersType, dataset_info: DatasetInfo, metadata: ModelMetadata + self, hypers: ModelHypersType, dataset_info: DatasetInfo, metadata: ModelMetadata ) -> None: """""" super().__init__() @@ -233,7 +234,7 @@ def get_checkpoint(self) -> Dict[str, Any]: """ -class TrainerInterface(Generic[HypersType], metaclass=ABCMeta): +class TrainerInterface(Generic[TrainerHypersType, ModelHypersType], metaclass=ABCMeta): """ Abstract base class for a model trainer in metatrain. @@ -250,7 +251,7 @@ class TrainerInterface(Generic[HypersType], metaclass=ABCMeta): This is used to upgrade checkpoints produced with earlier versions of the code. See :ref:`ckpt_version` for more information.""" - def __init__(self, hypers: HypersType): + def __init__(self, hypers: TrainerHypersType): required_attributes = [ "__checkpoint_version__", ] @@ -273,6 +274,15 @@ def __setattr__(self, name: str, value: Any) -> None: ) super().__setattr__(name, value) + #@abstractmethod Uncomment when all architectures support it. + def setup(self, model_hypers: ModelHypersType, dataset_info: DatasetInfo) -> Any: #MetatrainModel: + """ + Setup the trainer with the model hyper-parameters and dataset information. + + :param model_hypers: The hyper-parameters of the model to be trained. + :param dataset_info: Information about the dataset to be used for training. + """ + @abstractmethod def train( self, @@ -325,7 +335,7 @@ def upgrade_checkpoint(cls, checkpoint: Dict) -> Dict: def load_checkpoint( cls, checkpoint: Dict[str, Any], - hypers: HypersType, + hypers: TrainerHypersType, context: Literal["restart", "finetune"], ) -> "TrainerInterface": """ diff --git a/src/metatrain/utils/io.py b/src/metatrain/utils/io.py index 955b379e30..be5d746454 100644 --- a/src/metatrain/utils/io.py +++ b/src/metatrain/utils/io.py @@ -11,6 +11,8 @@ from .. import __version__ from .architectures import find_all_architectures, import_architecture +from .abc import ModelInterface +from .wrapper import MetatrainModel hf_pattern = re.compile( @@ -186,11 +188,10 @@ def load_model( checkpoint = torch.load(path, weights_only=False, map_location="cpu") return model_from_checkpoint(checkpoint, context="export") - def model_from_checkpoint( checkpoint: Dict[str, Any], context: Literal["restart", "finetune", "export"], -) -> torch.nn.Module: +) -> MetatrainModel: """ Load the checkpoint at the given ``path``, and create the corresponding model instance. The model architecture is determined from information stored inside the @@ -207,11 +208,17 @@ def model_from_checkpoint( :return: the loaded model instance. """ - architecture_name = checkpoint["architecture_name"] - if architecture_name == "metatrain_wrapper": - from .wrapper import MetatrainWrapper - return MetatrainWrapper.load_checkpoint(checkpoint, context=context) + return MetatrainModel.load_checkpoint( + checkpoint, + context=context, + ) + +def arch_model_from_checkpoint( + checkpoint: Dict[str, Any], + context: Literal["restart", "finetune", "export"] +) -> ModelInterface: + architecture_name = checkpoint["architecture_name"] if architecture_name not in find_all_architectures(): raise ValueError( @@ -283,9 +290,10 @@ def trainer_from_checkpoint( :return: the loaded trainer instance. """ - architecture_name = checkpoint["architecture_name"] - - if architecture_name == "metatrain_wrapper": + if "architecture_name" in checkpoint: + # This is an old checkpoint + architecture_name = checkpoint["architecture_name"] + else: architecture_name = checkpoint["model"]["architecture_name"] if architecture_name not in find_all_architectures(): diff --git a/src/metatrain/utils/wrapper.py b/src/metatrain/utils/wrapper.py index 6d7efb375b..e466c24e90 100644 --- a/src/metatrain/utils/wrapper.py +++ b/src/metatrain/utils/wrapper.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, List, Literal, Optional, NotRequired, Union +from typing import Any, Dict, List, Literal, Optional, NotRequired, Union, Callable import torch from metatensor.torch import Labels, TensorBlock, TensorMap @@ -8,55 +8,32 @@ ModelMetadata, ModelOutput, ModelCapabilities, - ModelEvaluationOptions, System, NeighborListOptions, ) from typing_extensions import TypedDict -from metatrain.scaler import Scaler +#from metatrain.scaler import Scaler from metatrain.utils.abc import ModelInterface from metatrain.utils.architectures import import_architecture from metatrain.utils.data import DatasetInfo -from metatrain.utils.data.atomic_basis_helpers import ( - sparsify_atomic_basis_target, -) from metatrain.utils.dtype import dtype_to_str - -class WrapperHypers(TypedDict): - """Hypers to initialize the model. - - These are only use on a first initialization. - When loading a checkpoint, the hypers are ignored - and instead the components of the model are loaded. - """ - - # Models passed directly - model: NotRequired[dict] - additive_models: NotRequired[list[dict]] - scaler: NotRequired[dict] - - -class MetatrainWrapper(ModelInterface[WrapperHypers]): +class MetatrainModel(torch.nn.Module): __checkpoint_version__ = 1 - __supported_devices__ = ["cuda", "cpu"] - __supported_dtypes__ = [torch.float32, torch.float64] - __default_metadata__ = ModelMetadata() - component_labels: Dict[str, List[List[Labels]]] - NUM_FEATURE_TYPES: int = 2 # node + edge features def __init__( self, - hypers: WrapperHypers, + model: ModelInterface, + additive_models: list[ModelInterface], + scaler: Optional[ModelInterface], #Scaler, dataset_info: DatasetInfo, ): - super().__init__(hypers=hypers, dataset_info=dataset_info, metadata=self.__default_metadata__) - - if "model" in hypers: - self.model = hypers["model"] - self.additive_models = torch.nn.ModuleList(hypers["additive_models"]) - self.scaler = hypers["scaler"] + super().__init__() + self.model = model + self.additive_models = torch.nn.ModuleList(additive_models) + self.scaler = scaler + self.dataset_info = dataset_info def forward( self, @@ -70,14 +47,15 @@ def forward( with torch.profiler.record_function("MTT_WRAPPER::post-processing"): if not self.training: - # at evaluation, we also introduce the scaler and additive contributions - return_dict = self.scaler.apply_scales( - systems, - return_dict, - selected_atoms=selected_atoms, - use_per_target_scales=True, - use_per_property_scales=True, - ) + if self.scaler is not None: + # at evaluation, we also introduce the scaler and additive contributions + return_dict = self.scaler.apply_scales( + systems, + return_dict, + selected_atoms=selected_atoms, + use_per_target_scales=True, + use_per_property_scales=True, + ) # For atomic basis targets, sparsify to create blocks with "atom_type" # in the key dimensions, and ensure properties are unpadded. This is @@ -152,7 +130,8 @@ def _add_model_requested_inputs(model: ModelInterface): for additive_model in self.additive_models: _add_model_requested_inputs(additive_model) - _add_model_requested_inputs(self.scaler) + if self.scaler is not None: + _add_model_requested_inputs(self.scaler) _add_model_requested_inputs(self.model) return requested_inputs @@ -170,7 +149,8 @@ def _add_model_requested_neighbor_lists(model: ModelInterface): for additive_model in self.additive_models: _add_model_requested_neighbor_lists(additive_model) - _add_model_requested_neighbor_lists(self.scaler) + if self.scaler is not None: + _add_model_requested_neighbor_lists(self.scaler) _add_model_requested_neighbor_lists(self.model) return requested_neighbor_lists @@ -198,12 +178,15 @@ def export( model = self.model.export(metadata) dtype = getattr(torch, model.capabilities().dtype) additive_models = [model.to(dtype).export(metadata) for model in self.additive_models] - scaler = self.scaler.to(dtype).export(metadata) + if self.scaler is not None: + scaler = self.scaler.to(dtype).export(metadata) + else: + scaler = None # Get a list of the capabilities of each model all_capabilities = [model.capabilities()] + [ model.capabilities() for model in additive_models - ] + [scaler.capabilities()] + ] + [scaler.capabilities()] if scaler is not None else [] # The interaction range of the model is the maximum interaction range # of all the models involved. @@ -218,11 +201,9 @@ def export( # Build the wrapper model again with the exported modules. to_export = self.__class__( - hypers=dict( - model=model.module, - additive_models=[model.module for model in additive_models], - scaler=scaler.module, - ), + model=model.module, + additive_models=[model.module for model in additive_models], + scaler=scaler.module, dataset_info=self.dataset_info, ) @@ -240,29 +221,85 @@ def export( return AtomisticModel(to_export.eval(), metadata, capabilities) + @staticmethod + def _ckpt_from_arch_ckpt(checkpoint: dict) -> dict: + """ + Convert a checkpoint from an architecture to one that can be loaded + by MetatrainModel. + + This is specially useful for converting checkpoints that were generated + before the introduction of the MetatrainModel class. + """ + from metatrain.utils.io import trainer_from_checkpoint + # Use the trainer to setup a MetatrainModel + trainer = trainer_from_checkpoint(checkpoint, context="restart", hypers={}) + model_data = checkpoint["model_data"] + mtt_model = trainer.setup( + model_hypers=model_data.get("hypers", model_data.get("model_hypers")), + dataset_info=model_data["dataset_info"] + ) + + # Get the new checkpoint skeleton from the MetatrainModel + new_ckpt = mtt_model.get_checkpoint() + + # Copy all the trainer keys. + for k in list(checkpoint): + if k not in new_ckpt["model"]: + new_ckpt[k] = checkpoint.pop(k) + + # The checkpoint of the model now goes to the "model" key. + # (here we replace completely the model in new_ckpt, since that one + # is simply an untrained model) + new_ckpt["model"] = checkpoint + + # Fill the state dicts of the scaler and additive models by finding + # them in the state dict of the original checkpoint. + scaler_state_dict = {} + additive_models_state_dict = [{}] * len(new_ckpt["additive_models"]) + + state_dict = checkpoint["model_state_dict"] + for k, v in state_dict.items(): + if k.startswith("scaler."): + scaler_state_dict[k.replace("scaler.", "")] = v + for i, additive_model_state_dict in enumerate(additive_models_state_dict): + if k.startswith(f"additive_models.{i}."): + additive_model_state_dict[k.replace(f"additive_models.{i}.", "")] = v + + if new_ckpt["scaler"] is None: + if len(scaler_state_dict) > 0: + raise ValueError( + "The checkpoint contains a scaler state dict, but the model " + "does not have a scaler." + ) + else: + new_ckpt["scaler"]["model_state_dict"] = scaler_state_dict + new_ckpt["scaler"]["best_model_state_dict"] = scaler_state_dict + + for i, additive_model_state_dict in enumerate(additive_models_state_dict): + new_ckpt["additive_models"][i]["model_state_dict"] = additive_model_state_dict + new_ckpt["additive_models"][i]["best_model_state_dict"] = additive_model_state_dict + + return new_ckpt + @classmethod def load_checkpoint( cls, checkpoint: Dict[str, Any], context: Literal["restart", "finetune", "export"], ) -> "ModelInterface": - hypers = {} - for k in ["model", "scaler"]: - subcheckpoint = checkpoint[k] - - architecture_name = subcheckpoint["architecture_name"] - model_cls = import_architecture(architecture_name).__model__ - hypers[k] = model_cls.load_checkpoint(subcheckpoint, context) - - hypers["additive_models"] = [] - for additive_model_checkpoint in checkpoint["additive_models"]: - architecture_name = additive_model_checkpoint["architecture_name"] - model_cls = import_architecture(architecture_name).__model__ - hypers["additive_models"].append( - model_cls.load_checkpoint(additive_model_checkpoint, context) - ) - - return cls(hypers=hypers, dataset_info=checkpoint["dataset_info"]) + from .io import arch_model_from_checkpoint + if "architecture_name" in checkpoint: + checkpoint = cls._ckpt_from_arch_ckpt(checkpoint) + + return cls( + model=arch_model_from_checkpoint(checkpoint["model"], context=context), + additive_models=[ + arch_model_from_checkpoint(additive_model_checkpoint, context=context) + for additive_model_checkpoint in checkpoint["additive_models"] + ], + scaler=arch_model_from_checkpoint(checkpoint["scaler"], context=context) if checkpoint["scaler"] is not None else None, + dataset_info=checkpoint["dataset_info"] + ) @classmethod def upgrade_checkpoint(cls, checkpoint: Dict["str", Any]) -> Dict["str", Any]: @@ -300,13 +337,11 @@ def get_checkpoint(self) -> Dict[str, Any]: :return: The model's checkpoint. """ checkpoint = { - "architecture_name": "metatrain_wrapper", "model_ckpt_version": self.__checkpoint_version__, - "metadata": self.metadata, "dataset_info": self.dataset_info, "model": self.model.get_checkpoint(), "additive_models": [m.get_checkpoint() for m in self.additive_models], - "scaler": self.scaler.get_checkpoint(), + "scaler": self.scaler.get_checkpoint() if self.scaler is not None else None, } return checkpoint @@ -322,20 +357,22 @@ def restart(self, dataset_info, model_hypers = None): """ self.model.restart(dataset_info, model_hypers) - composition_model = self.additive_models[0] - self.additive_models[0] = composition_model.restart( - dataset_info=DatasetInfo( - length_unit=dataset_info.length_unit, - atomic_types=dataset_info.atomic_types, - targets={ - target_name: target_info - for target_name, target_info in dataset_info.targets.items() - if composition_model.is_valid_target(target_name, target_info) - }, - ), - ) + if len(self.additive_models) > 0: + composition_model = self.additive_models[0] + self.additive_models[0] = composition_model.restart( + dataset_info=DatasetInfo( + length_unit=dataset_info.length_unit, + atomic_types=dataset_info.atomic_types, + targets={ + target_name: target_info + for target_name, target_info in dataset_info.targets.items() + if composition_model.is_valid_target(target_name, target_info) + }, + ), + ) - self.scaler = self.scaler.restart(dataset_info) + if self.scaler is not None: + self.scaler = self.scaler.restart(dataset_info) self.dataset_info = dataset_info @@ -354,7 +391,7 @@ def remove_output(self, target_name: str) -> None: if target_name in additive_model.supported_outputs(): additive_model.remove_output(target_name) - if target_name in self.scaler.supported_outputs(): + if self.scaler is not None and target_name in self.scaler.supported_outputs(): self.scaler.remove_output(target_name) self.dataset_info.targets.pop(target_name, None)