diff --git a/src/metatrain/composition/model.py b/src/metatrain/composition/model.py index ba932fe286..4b039a9cdd 100644 --- a/src/metatrain/composition/model.py +++ b/src/metatrain/composition/model.py @@ -343,7 +343,6 @@ def supported_outputs(self) -> Dict[str, ModelOutput]: def _add_output(self, target_name: str, target_info: TargetInfo) -> None: self.outputs[target_name] = ModelOutput( - quantity=target_info.quantity, unit=target_info.unit, sample_kind="atom", description=target_info.description, diff --git a/src/metatrain/composition/tests/test_errors.py b/src/metatrain/composition/tests/test_errors.py index 56f5360277..d3f18520aa 100644 --- a/src/metatrain/composition/tests/test_errors.py +++ b/src/metatrain/composition/tests/test_errors.py @@ -108,4 +108,4 @@ def test_forward_unknown_output_raises(): systems = read_systems(DATASET_PATH) with pytest.raises(ValueError, match="not supported"): - model(systems[:1], {"nonexistent": ModelOutput(quantity="energy")}) + model(systems[:1], {"nonexistent": ModelOutput()}) diff --git a/src/metatrain/experimental/classifier/model.py b/src/metatrain/experimental/classifier/model.py index 8486806ed9..77188139f1 100644 --- a/src/metatrain/experimental/classifier/model.py +++ b/src/metatrain/experimental/classifier/model.py @@ -91,7 +91,7 @@ def set_wrapped_model(self, model: ModelInterface) -> None: # Store capabilities outputs = {name: ModelOutput() for name in self.dataset_info.targets.keys()} - outputs["feature"] = ModelOutput(quantity="", unit="", sample_kind="system") + outputs["feature"] = ModelOutput(unit="", sample_kind="system") self.capabilities = ModelCapabilities( outputs=outputs, atomic_types=old_capabilities.atomic_types, diff --git a/src/metatrain/experimental/classifier/trainer.py b/src/metatrain/experimental/classifier/trainer.py index 9a4e648836..834ad699f3 100644 --- a/src/metatrain/experimental/classifier/trainer.py +++ b/src/metatrain/experimental/classifier/trainer.py @@ -202,11 +202,7 @@ def get_lr_schedule(step): # Forward pass - request logits for training outputs = model( systems, - { - target_name_logits: ModelOutput( - quantity="", unit="", sample_kind="system" - ) - }, + {target_name_logits: ModelOutput(unit="", sample_kind="system")}, None, ) @@ -252,7 +248,7 @@ def get_lr_schedule(step): systems, { target_name_logits: ModelOutput( - quantity="", unit="", sample_kind="system" + unit="", sample_kind="system" ) }, None, diff --git a/src/metatrain/experimental/dpa3/model.py b/src/metatrain/experimental/dpa3/model.py index 8d3fcfb8aa..62fd33425c 100644 --- a/src/metatrain/experimental/dpa3/model.py +++ b/src/metatrain/experimental/dpa3/model.py @@ -214,7 +214,6 @@ def _add_output(self, target_name: str, target: TargetInfo) -> None: block.properties for block in target.layout.blocks() ] self.outputs[target_name] = ModelOutput( - quantity=target.quantity, unit=target.unit, sample_kind="atom", ) diff --git a/src/metatrain/experimental/dpa3/tests/test_compile.py b/src/metatrain/experimental/dpa3/tests/test_compile.py index dd7ff7c263..a03cd6ae48 100644 --- a/src/metatrain/experimental/dpa3/tests/test_compile.py +++ b/src/metatrain/experimental/dpa3/tests/test_compile.py @@ -55,7 +55,7 @@ def test_forward_from_batch_matches_forward(): # Standard forward out = model( systems, - {"mtt::U0": ModelOutput(quantity="energy", unit="", sample_kind="system")}, + {"mtt::U0": ModelOutput(unit="", sample_kind="system")}, ) # Verify the standard forward works assert out["mtt::U0"].block().values.numel() > 0 diff --git a/src/metatrain/experimental/dpa3/tests/test_pretrained.py b/src/metatrain/experimental/dpa3/tests/test_pretrained.py index 9bea57bbc3..eb30a3aef3 100644 --- a/src/metatrain/experimental/dpa3/tests/test_pretrained.py +++ b/src/metatrain/experimental/dpa3/tests/test_pretrained.py @@ -141,9 +141,7 @@ def test_pretrained_forward_matches(): for s in systems: get_system_with_neighbor_lists(s, base.requested_neighbor_lists()) - output_request = { - "mtt::U0": ModelOutput(quantity="energy", unit="", sample_kind="system") - } + output_request = {"mtt::U0": ModelOutput(unit="", sample_kind="system")} base_out = base(systems, output_request) pretrained_out = pretrained(systems, output_request) @@ -179,9 +177,7 @@ def test_pretrained_checkpoint_roundtrip(): for s in systems: get_system_with_neighbor_lists(s, pretrained.requested_neighbor_lists()) - output_request = { - "mtt::U0": ModelOutput(quantity="energy", unit="", sample_kind="system") - } + output_request = {"mtt::U0": ModelOutput(unit="", sample_kind="system")} orig_out = pretrained(systems, output_request) reloaded_out = reloaded(systems, output_request) diff --git a/src/metatrain/experimental/dpa3/tests/test_regression.py b/src/metatrain/experimental/dpa3/tests/test_regression.py index 975358bd5b..b4737c7f3b 100644 --- a/src/metatrain/experimental/dpa3/tests/test_regression.py +++ b/src/metatrain/experimental/dpa3/tests/test_regression.py @@ -44,7 +44,7 @@ def test_regression_init(): output = model( systems, - {"mtt::U0": ModelOutput(quantity="energy", unit="", sample_kind="system")}, + {"mtt::U0": ModelOutput(unit="", sample_kind="system")}, ) expected_output = torch.tensor( diff --git a/src/metatrain/experimental/flashmd/model.py b/src/metatrain/experimental/flashmd/model.py index 427410df8e..fc37f5269e 100644 --- a/src/metatrain/experimental/flashmd/model.py +++ b/src/metatrain/experimental/flashmd/model.py @@ -306,10 +306,8 @@ def requested_neighbor_lists(self) -> List[NeighborListOptions]: def requested_inputs(self) -> Dict[str, ModelOutput]: return { - "momentum": ModelOutput( - quantity="momentum", unit="(eV*u)^(1/2)", sample_kind="atom" - ), - "mass": ModelOutput(quantity="mass", unit="u", sample_kind="atom"), + "momentum": ModelOutput(unit="(eV*u)^(1/2)", sample_kind="atom"), + "mass": ModelOutput(unit="u", sample_kind="atom"), } def forward( @@ -1256,7 +1254,6 @@ def _add_output(self, target_name: str, target_info: TargetInfo) -> None: ] + [len(block.properties.values)] self.outputs[target_name] = ModelOutput( - quantity=target_info.quantity, unit=target_info.unit, sample_kind="atom", description=target_info.description, diff --git a/src/metatrain/experimental/flashmd/modules/additive.py b/src/metatrain/experimental/flashmd/modules/additive.py index 98df04a888..e453ca8000 100644 --- a/src/metatrain/experimental/flashmd/modules/additive.py +++ b/src/metatrain/experimental/flashmd/modules/additive.py @@ -41,7 +41,6 @@ def __init__(self, hypers: PositionAdditiveHypers, dataset_info: DatasetInfo): # skip momenta targets unless `also_momenta` is True continue self.outputs[key] = ModelOutput( - quantity=value.quantity, unit=value.unit, sample_kind="atom", description=value.description, @@ -63,7 +62,6 @@ def restart(self, dataset_info: DatasetInfo) -> "PositionAdditive": # skip momenta targets unless `also_momenta` is True continue self.outputs[key] = ModelOutput( - quantity=value.quantity, unit=value.unit, sample_kind="atom", description=value.description, diff --git a/src/metatrain/experimental/flashmd/tests/test_functionality.py b/src/metatrain/experimental/flashmd/tests/test_functionality.py index f15bf7353d..c81a7a9d40 100644 --- a/src/metatrain/experimental/flashmd/tests/test_functionality.py +++ b/src/metatrain/experimental/flashmd/tests/test_functionality.py @@ -100,8 +100,8 @@ def test_forward(): system.add_data("momentum", tmap) outputs = { - "position": ModelOutput(quantity="length", unit="angstrom", sample_kind="atom"), - "momentum": ModelOutput(quantity="length", unit="angstrom", sample_kind="atom"), + "position": ModelOutput(unit="angstrom", sample_kind="atom"), + "momentum": ModelOutput(unit="angstrom", sample_kind="atom"), } result_dict = model(systems, outputs) diff --git a/src/metatrain/experimental/flashmd_symplectic/model.py b/src/metatrain/experimental/flashmd_symplectic/model.py index 06185270e5..1b0c499c0c 100644 --- a/src/metatrain/experimental/flashmd_symplectic/model.py +++ b/src/metatrain/experimental/flashmd_symplectic/model.py @@ -288,10 +288,8 @@ def requested_neighbor_lists(self) -> List[NeighborListOptions]: def requested_inputs(self) -> Dict[str, ModelOutput]: return { - "momentum": ModelOutput( - quantity="momentum", unit="(eV*u)^(1/2)", sample_kind="atom" - ), - "mass": ModelOutput(quantity="mass", unit="u", sample_kind="atom"), + "momentum": ModelOutput(unit="(eV*u)^(1/2)", sample_kind="atom"), + "mass": ModelOutput(unit="u", sample_kind="atom"), } def forward( diff --git a/src/metatrain/experimental/flashmd_symplectic/tests/test_functionality.py b/src/metatrain/experimental/flashmd_symplectic/tests/test_functionality.py index 3b051a7ad9..f21df551cd 100644 --- a/src/metatrain/experimental/flashmd_symplectic/tests/test_functionality.py +++ b/src/metatrain/experimental/flashmd_symplectic/tests/test_functionality.py @@ -100,8 +100,8 @@ def test_forward(): system.add_data("momentum", tmap) outputs = { - "position": ModelOutput(quantity="length", unit="angstrom", sample_kind="atom"), - "momentum": ModelOutput(quantity="length", unit="angstrom", sample_kind="atom"), + "position": ModelOutput(unit="angstrom", sample_kind="atom"), + "momentum": ModelOutput(unit="angstrom", sample_kind="atom"), } result_dict = model(systems, outputs) diff --git a/src/metatrain/experimental/mace/model.py b/src/metatrain/experimental/mace/model.py index 7d11eb39aa..b81de8ad0f 100644 --- a/src/metatrain/experimental/mace/model.py +++ b/src/metatrain/experimental/mace/model.py @@ -284,7 +284,6 @@ def __init__(self, hypers: ModelHypers, dataset_info: DatasetInfo) -> None: targets = dataset_info.targets self.outputs = { k: ModelOutput( - quantity=targets[k].quantity if k in targets else "", unit=targets[k].unit if k in targets else "", sample_kind="atom", ) diff --git a/src/metatrain/experimental/space/model.py b/src/metatrain/experimental/space/model.py index 11f3833e5f..deb194cc50 100644 --- a/src/metatrain/experimental/space/model.py +++ b/src/metatrain/experimental/space/model.py @@ -644,7 +644,6 @@ def _train_dataset_info(self, dataset_info: DatasetInfo) -> DatasetInfo: def _add_output(self, target_name: str, target_info: TargetInfo) -> None: self.outputs[target_name] = ModelOutput( - quantity=target_info.quantity, unit=target_info.unit, sample_kind="atom", ) diff --git a/src/metatrain/experimental/space/tests/test_functionality.py b/src/metatrain/experimental/space/tests/test_functionality.py index 97551adb4e..19b7d0ca23 100644 --- a/src/metatrain/experimental/space/tests/test_functionality.py +++ b/src/metatrain/experimental/space/tests/test_functionality.py @@ -159,7 +159,7 @@ def test_multiple_targets(): model = SPACE(hypers, dataset_info) system = _make_system(model) outputs = { - "energy": ModelOutput(quantity="energy", unit="eV", sample_kind="system"), + "energy": ModelOutput(unit="eV", sample_kind="system"), "dipole": ModelOutput(sample_kind="system"), "non_conservative_stress": ModelOutput(sample_kind="system"), } diff --git a/src/metatrain/experimental/space/tests/test_regression.py b/src/metatrain/experimental/space/tests/test_regression.py index a7e7b5ae0c..ddb89cf51f 100644 --- a/src/metatrain/experimental/space/tests/test_regression.py +++ b/src/metatrain/experimental/space/tests/test_regression.py @@ -61,7 +61,7 @@ def test_regression_init(): output = model( systems, - {"mtt::U0": ModelOutput(quantity="energy", unit="", sample_kind="system")}, + {"mtt::U0": ModelOutput(unit="", sample_kind="system")}, ) expected_output = torch.tensor( @@ -158,7 +158,7 @@ def test_regression_train(): model = torch.jit.script(model) output = model( systems[:5], - {"mtt::U0": ModelOutput(quantity="energy", unit="", sample_kind="system")}, + {"mtt::U0": ModelOutput(unit="", sample_kind="system")}, ) expected_output = torch.tensor( @@ -258,11 +258,7 @@ def test_regression_train_spherical(device): ] output = model( systems, - { - "mtt::electron_density_basis": ModelOutput( - quantity="", unit="", sample_kind="atom" - ) - }, + {"mtt::electron_density_basis": ModelOutput(unit="", sample_kind="atom")}, ) expected_output = torch.tensor( diff --git a/src/metatrain/experimental/space/trainer.py b/src/metatrain/experimental/space/trainer.py index 9ee27d210b..7f39af9e04 100644 --- a/src/metatrain/experimental/space/trainer.py +++ b/src/metatrain/experimental/space/trainer.py @@ -58,7 +58,6 @@ def _get_requested_outputs(targets, target_info_dict): requested_outputs = {} for name, target in targets.items(): requested_outputs[name] = ModelOutput( - quantity=target_info_dict[name].quantity, unit=target_info_dict[name].unit, sample_kind=target_info_dict[name].sample_kind, explicit_gradients=target.block(0).gradients_list(), diff --git a/src/metatrain/gap/model.py b/src/metatrain/gap/model.py index c830073820..0a7d009a5c 100644 --- a/src/metatrain/gap/model.py +++ b/src/metatrain/gap/model.py @@ -73,7 +73,6 @@ def __init__(self, hypers: ModelHypers, dataset_info: DatasetInfo) -> None: self.outputs = { key: ModelOutput( - quantity=value.quantity, unit=value.unit, sample_kind="system", description=value.description, diff --git a/src/metatrain/llpr/model.py b/src/metatrain/llpr/model.py index 27f52b5313..2a61861806 100644 --- a/src/metatrain/llpr/model.py +++ b/src/metatrain/llpr/model.py @@ -155,7 +155,6 @@ def set_wrapped_model(self, model: ModelInterface) -> None: self.outputs_list.append(name) uncertainty_name = _get_uncertainty_name(name) additional_capabilities[uncertainty_name] = ModelOutput( - quantity=output.quantity, unit=output.unit, sample_kind=output.sample_kind, description=output.description, @@ -216,13 +215,10 @@ def set_wrapped_model(self, model: ModelInterface) -> None: ) if ensemble_output_name == "mtt::aux::energy_ensemble": ensemble_output_name = "energy_ensemble" - explicit_gradients = _ensemble_explicit_gradients( - old_capabilities.outputs[name] - ) + explicit_gradients = self._ensemble_explicit_gradients(name) if len(explicit_gradients) > 0: self.ensemble_gradient_outputs.append(ensemble_output_name) ensemble_outputs[ensemble_output_name] = ModelOutput( - quantity=old_capabilities.outputs[name].quantity, unit=old_capabilities.outputs[name].unit, sample_kind=old_capabilities.outputs[name].sample_kind, explicit_gradients=explicit_gradients, @@ -1146,10 +1142,9 @@ def generate_ensemble(self) -> None: if ensemble_name == "mtt::aux::energy_ensemble": ensemble_name = "energy_ensemble" new_outputs[ensemble_name] = ModelOutput( - quantity=old_outputs[name].quantity, unit=old_outputs[name].unit, sample_kind=old_outputs[name].sample_kind, - explicit_gradients=_ensemble_explicit_gradients(old_outputs[name]), + explicit_gradients=self._ensemble_explicit_gradients(name), description=f"ensemble of {name}", ) self.capabilities = ModelCapabilities( @@ -1314,29 +1309,18 @@ def upgrade_checkpoint(cls, checkpoint: Dict) -> Dict: def supported_outputs(self) -> Dict[str, ModelOutput]: return self.capabilities.outputs + def _ensemble_explicit_gradients(self, name: str) -> List[str]: + """Explicit gradients the ensemble of the ``name`` target is able to produce. -def _ensemble_explicit_gradients(output: ModelOutput) -> List[str]: - """Explicit gradients the ensemble of ``output`` is able to produce. - - ``_add_energy_ensemble_gradients`` differentiates one per-system scalar per - ensemble member, so what matters is that the output *is* a per-system energy, - not what it is called: an energy target may carry any name (``mtt::my_energy``) - and any variant (``energy/pbesol``). Selecting on the quantity rather than on the - name keeps all of those working. - - Other quantities get nothing: positions/strain gradients of them are not what - this computes. + Only energy targets are currently supported. - Note that ``sample_kind`` is deliberately *not* consulted here. In capabilities - it describes what the wrapped model is able to produce (PET reports ``"atom"`` - for its energy even when the target is per-system), not what a given call asks - for. Whether the *requested* sample kind is supported is checked in ``forward``, - where the request is actually known. - - :param output: the wrapped model's output the ensemble is built from. - :return: gradient names for the corresponding ensemble output. - """ - return ["positions", "strain"] if output.quantity == "energy" else [] + :param name: name of the target the ensemble is built from. + :return: gradient names for the corresponding ensemble output. + """ + target = self.dataset_info.targets.get(name) + if target is None or target.quantity != "energy": + return [] + return ["positions", "strain"] def _get_uncertainty_name(name: str) -> str: diff --git a/src/metatrain/pet/model.py b/src/metatrain/pet/model.py index 935df4adfa..63860c325f 100644 --- a/src/metatrain/pet/model.py +++ b/src/metatrain/pet/model.py @@ -270,7 +270,7 @@ def requested_neighbor_lists(self) -> List[NeighborListOptions]: def requested_inputs(self) -> Dict[str, ModelOutput]: if self.system_conditioning is not None: return { - key: ModelOutput(quantity="", unit="", sample_kind="system") + key: ModelOutput(unit="", sample_kind="system") for key in self.system_conditioning.required_data_keys } return {} @@ -1051,7 +1051,6 @@ def _add_output(self, target_name: str, target_info: TargetInfo) -> None: ] + [len(block.properties.values)] self.outputs[target_name] = ModelOutput( - quantity=target_info.quantity, unit=target_info.unit, sample_kind="atom", description=target_info.description, diff --git a/src/metatrain/pet/tests/test_regression.py b/src/metatrain/pet/tests/test_regression.py index a2761486ba..90ae3acc10 100644 --- a/src/metatrain/pet/tests/test_regression.py +++ b/src/metatrain/pet/tests/test_regression.py @@ -60,7 +60,7 @@ def test_regression_init(): output = model( systems, - {"mtt::U0": ModelOutput(quantity="energy", unit="", sample_kind="system")}, + {"mtt::U0": ModelOutput(unit="", sample_kind="system")}, ) expected_output = torch.tensor( @@ -307,7 +307,7 @@ def test_regression_energy_non_conservative_stress(batch_size): outputs = model( eval_systems, { - "energy": ModelOutput(quantity="energy", unit="", sample_kind="system"), + "energy": ModelOutput(unit="", sample_kind="system"), "non_conservative_stress": ModelOutput(sample_kind="system"), }, ) @@ -426,11 +426,7 @@ def test_regression_train_spherical(device): ] output = model( systems, - { - "mtt::electron_density_basis": ModelOutput( - quantity="", unit="", sample_kind="atom" - ) - }, + {"mtt::electron_density_basis": ModelOutput(unit="", sample_kind="atom")}, ) expected_output = torch.tensor( diff --git a/src/metatrain/scaler/model.py b/src/metatrain/scaler/model.py index 737b3c10a8..01c7c28d0b 100644 --- a/src/metatrain/scaler/model.py +++ b/src/metatrain/scaler/model.py @@ -304,7 +304,6 @@ def supported_outputs(self) -> Dict[str, ModelOutput]: def _add_output(self, target_name: str, target_info: TargetInfo) -> None: self.outputs[target_name] = ModelOutput( - quantity=target_info.quantity, unit=target_info.unit, sample_kind="atom", description=target_info.description, diff --git a/src/metatrain/scaler/tests/test_scaler.py b/src/metatrain/scaler/tests/test_scaler.py index fc2a99d484..b8f67ca1cf 100644 --- a/src/metatrain/scaler/tests/test_scaler.py +++ b/src/metatrain/scaler/tests/test_scaler.py @@ -2640,7 +2640,6 @@ def test_scaler_with_neighbor_list_additive_model(): outputs = { "energy": ModelOutput( - quantity="", unit="", sample_kind="system", ) diff --git a/src/metatrain/soap_bpnn/model.py b/src/metatrain/soap_bpnn/model.py index 72d932749e..e4ebb6622d 100644 --- a/src/metatrain/soap_bpnn/model.py +++ b/src/metatrain/soap_bpnn/model.py @@ -1232,7 +1232,6 @@ def _add_output(self, target_name: str, target: TargetInfo) -> None: ] self.outputs[target_name] = ModelOutput( - quantity=target.quantity, unit=target.unit, sample_kind="atom", description=target.description, diff --git a/src/metatrain/soap_bpnn/tests/test_functionality.py b/src/metatrain/soap_bpnn/tests/test_functionality.py index 8ab43be67e..191c606630 100644 --- a/src/metatrain/soap_bpnn/tests/test_functionality.py +++ b/src/metatrain/soap_bpnn/tests/test_functionality.py @@ -49,7 +49,7 @@ def test_scalar_output(legacy): system = _make_system(model) output = model( [system], - {"energy": ModelOutput(quantity="energy", unit="eV", sample_kind="system")}, + {"energy": ModelOutput(unit="eV", sample_kind="system")}, ) values = output["energy"].block().values assert values.shape == (1, 1) @@ -163,7 +163,7 @@ def test_mlp_head(add_lambda_basis): system = _make_system(model) output = model( [system], - {"energy": ModelOutput(quantity="energy", unit="eV", sample_kind="system")}, + {"energy": ModelOutput(unit="eV", sample_kind="system")}, ) values = output["energy"].block().values assert values.shape == (1, 1) @@ -195,7 +195,7 @@ def test_multiple_targets(legacy, add_lambda_basis): model = SoapBpnn(hypers, dataset_info) system = _make_system(model) outputs = { - "energy": ModelOutput(quantity="energy", unit="eV", sample_kind="system"), + "energy": ModelOutput(unit="eV", sample_kind="system"), "non_conservative_stress": ModelOutput(sample_kind="system"), } result = model([system], outputs) diff --git a/src/metatrain/soap_bpnn/tests/test_regression.py b/src/metatrain/soap_bpnn/tests/test_regression.py index 83226cbfd4..12a4cb3296 100644 --- a/src/metatrain/soap_bpnn/tests/test_regression.py +++ b/src/metatrain/soap_bpnn/tests/test_regression.py @@ -46,7 +46,7 @@ def test_regression_init(): output = model( systems, - {"mtt::U0": ModelOutput(quantity="energy", unit="", sample_kind="system")}, + {"mtt::U0": ModelOutput(unit="", sample_kind="system")}, ) expected_output = torch.tensor( @@ -128,7 +128,7 @@ def test_regression_train(device): ] output = model( systems[:5], - {"mtt::U0": ModelOutput(quantity="energy", unit="", sample_kind="system")}, + {"mtt::U0": ModelOutput(unit="", sample_kind="system")}, ) expected_output = torch.tensor( @@ -219,11 +219,7 @@ def test_regression_train_spherical(device): ] output = model( systems, - { - "mtt::electron_density_basis": ModelOutput( - quantity="", unit="", sample_kind="atom" - ) - }, + {"mtt::electron_density_basis": ModelOutput(unit="", sample_kind="atom")}, ) expected_output = torch.tensor( diff --git a/src/metatrain/utils/additive/zbl.py b/src/metatrain/utils/additive/zbl.py index 5b414d54e3..6f3851ecfb 100644 --- a/src/metatrain/utils/additive/zbl.py +++ b/src/metatrain/utils/additive/zbl.py @@ -57,7 +57,6 @@ def __init__(self, hypers: Dict, dataset_info: DatasetInfo): self.outputs = { key: ModelOutput( - quantity=value.quantity, unit=value.unit, sample_kind="atom", description=value.description, diff --git a/src/metatrain/utils/data/dataset.py b/src/metatrain/utils/data/dataset.py index 5feab7356b..9c9f785c4e 100644 --- a/src/metatrain/utils/data/dataset.py +++ b/src/metatrain/utils/data/dataset.py @@ -95,7 +95,7 @@ def __init__( ): # verify that `length_unit` and `atomic_types` are valid for metatomic _ = ModelCapabilities( - outputs={"energy": ModelOutput()}, + outputs={"energy": ModelOutput(unit="eV")}, length_unit=length_unit, atomic_types=atomic_types, ) diff --git a/src/metatrain/utils/data/target_info.py b/src/metatrain/utils/data/target_info.py index 16adc00a29..3368ce9dfd 100644 --- a/src/metatrain/utils/data/target_info.py +++ b/src/metatrain/utils/data/target_info.py @@ -52,8 +52,8 @@ def __init__( self._check_layout(layout) self.layout = layout - # verify that `quantity`, `unit` and `description` are valid for metatomic - _ = ModelOutput(quantity=quantity, unit=unit, description=description) + # verify that `unit` and `description` are valid for metatomic + _ = ModelOutput(unit=unit, description=description) self.quantity = quantity self.unit = unit diff --git a/src/metatrain/utils/evaluate_model.py b/src/metatrain/utils/evaluate_model.py index 69db73a06f..76042ff8f2 100644 --- a/src/metatrain/utils/evaluate_model.py +++ b/src/metatrain/utils/evaluate_model.py @@ -269,7 +269,6 @@ def _get_model_outputs( length_unit="", # this is only needed for unit conversions in MD engines outputs={ key: ModelOutput( - quantity=value.quantity, unit=value.unit, sample_kind=value.sample_kind, description=value.description, @@ -283,7 +282,6 @@ def _get_model_outputs( systems, { key: ModelOutput( - quantity=value.quantity, unit=value.unit, sample_kind=value.sample_kind, description=value.description, diff --git a/src/metatrain/utils/logging.py b/src/metatrain/utils/logging.py index 14b4be2340..4254a9cf5c 100644 --- a/src/metatrain/utils/logging.py +++ b/src/metatrain/utils/logging.py @@ -398,14 +398,29 @@ def setup_logging( formatter = logging.Formatter(format, datefmt="%Y-%m-%d %H:%M:%S", style="{") handlers: List[Union[logging.StreamHandler, logging.FileHandler]] = [] + class HideWarnings(logging.Filter): + messages_to_hide = [ + "is multi-threaded, use of fork() may lead to deadlocks in the child", + ] + + def filter(self, record: logging.LogRecord) -> bool: + for message in self.messages_to_hide: + if message in record.getMessage(): + return False + return True + + warnings_filter = HideWarnings() + stream_handler = logging.StreamHandler(sys.stdout) stream_handler.setFormatter(formatter) + stream_handler.addFilter(warnings_filter) handlers.append(stream_handler) if log_file and is_main_process(): log_file = check_file_extension(filename=log_file, extension=".log") file_handler = logging.FileHandler(filename=str(log_file), encoding="utf-8") file_handler.setFormatter(formatter) + file_handler.addFilter(warnings_filter) handlers.append(file_handler) csv_file = Path(log_file).with_suffix(".csv") diff --git a/src/metatrain/utils/testing/output.py b/src/metatrain/utils/testing/output.py index a95c6e3378..802ef9306d 100644 --- a/src/metatrain/utils/testing/output.py +++ b/src/metatrain/utils/testing/output.py @@ -645,7 +645,6 @@ def test_output_features( ) features_output_options = ModelOutput( - quantity="", unit="", sample_kind=sample_kind, ) @@ -730,7 +729,6 @@ def test_output_last_layer_features( # last-layer features per atom: ll_output_options = ModelOutput( - quantity="", unit="", sample_kind=sample_kind, ) diff --git a/tests/cli/test_train_model.py b/tests/cli/test_train_model.py index a976494ee1..26cf31372e 100644 --- a/tests/cli/test_train_model.py +++ b/tests/cli/test_train_model.py @@ -22,11 +22,13 @@ import metatrain.soap_bpnn from metatrain import RANDOM_SEED -from metatrain.cli.train import _process_restart_from, train_model +from metatrain.cli.train import _process_restart_from +from metatrain.cli.train import train_model as raw_train_model from metatrain.utils.data import build_train_dataloaders, build_val_dataloaders from metatrain.utils.data.readers.ase import read from metatrain.utils.data.writers import DiskDatasetWriter from metatrain.utils.errors import ArchitectureError +from metatrain.utils.logging import ROOT_LOGGER, setup_logging from metatrain.utils.neighbor_lists import get_system_with_neighbor_lists from metatrain.utils.pydantic import MetatrainValidationError from metatrain.utils.testing._utils import WANDB_AVAILABLE @@ -67,6 +69,11 @@ def options_spherical(): return OmegaConf.load(OPTIONS_SPHERICAL_PATH) +def train_model(*args, **kwargs): + with setup_logging(ROOT_LOGGER, log_file=None, level=logging.INFO): + return raw_train_model(*args, **kwargs) + + @pytest.mark.parametrize("output", [None, "mymodel.pt"]) def test_train(capfd, monkeypatch, tmp_path, output): """Test that training via the training cli runs without an error raise."""