diff --git a/tyff/_serialization.py b/tyff/_serialization.py new file mode 100644 index 0000000..9a08c12 --- /dev/null +++ b/tyff/_serialization.py @@ -0,0 +1,87 @@ +from typing import Any + +import openff.interchange +import openff.units +import torch + +from tyff._models import TensorForceField, TensorPotential, TensorVSites + + +def _dump_tensor(t: torch.Tensor) -> dict[str, Any]: + return {"dtype": str(t.dtype).removeprefix("torch."), "shape": list(t.shape), "data": t.flatten().tolist()} + + +def _load_tensor(v: Any) -> torch.Tensor: + return torch.tensor(v["data"], dtype=getattr(torch, v["dtype"])).reshape(v["shape"]) + + +def _dump_exceptions(d: dict[tuple[int, int], int] | None) -> list[tuple[int, int, int]] | None: + return None if d is None else [(i, j, v) for (i, j), v in d.items()] + + +def _load_exceptions(v: Any) -> dict[tuple[int, int], int] | None: + if v is None or isinstance(v, dict): + return v + return {(i, j): val for i, j, val in v} + + +def dump_tensor_potential(p: TensorPotential) -> dict: + return { + "type": p.type, + "fn": p.fn, + "parameters": _dump_tensor(p.parameters), + "parameter_keys": [k.model_dump() for k in p.parameter_keys], + "parameter_cols": list(p.parameter_cols), + "parameter_units": [str(u) for u in p.parameter_units], + "attributes": None if p.attributes is None else _dump_tensor(p.attributes), + "attribute_cols": p.attribute_cols, + "attribute_units": None if p.attribute_units is None else [str(u) for u in p.attribute_units], + "exceptions": _dump_exceptions(p.exceptions), + } + + +def load_tensor_potential(d: dict) -> TensorPotential: + return TensorPotential( + type=d["type"], + fn=d["fn"], + parameters=_load_tensor(d["parameters"]), + parameter_keys=[openff.interchange.models.PotentialKey.model_validate(k) for k in d["parameter_keys"]], + parameter_cols=tuple(d["parameter_cols"]), + parameter_units=tuple(openff.units.unit.Unit(u) for u in d["parameter_units"]), + attributes=None if d["attributes"] is None else _load_tensor(d["attributes"]), + attribute_cols=tuple(d["attribute_cols"]) if d["attribute_cols"] is not None else None, + attribute_units=( + None if d["attribute_units"] is None else tuple(openff.units.unit.Unit(u) for u in d["attribute_units"]) + ), + exceptions=_load_exceptions(d["exceptions"]), + ) + + +def dump_tensor_vsites(v: TensorVSites) -> dict: + return { + "keys": [k.model_dump() for k in v.keys], + "weights": [_dump_tensor(w) for w in v.weights], + "parameters": _dump_tensor(v.parameters), + } + + +def load_tensor_vsites(d: dict) -> TensorVSites: + return TensorVSites( + keys=[openff.interchange.models.PotentialKey.model_validate(k) for k in d["keys"]], + weights=[_load_tensor(w) for w in d["weights"]], + parameters=_load_tensor(d["parameters"]), + ) + + +def dump_tensor_force_field(tff: TensorForceField) -> dict: + return { + "potentials": [dump_tensor_potential(p) for p in tff.potentials], + "v_sites": None if tff.v_sites is None else dump_tensor_vsites(tff.v_sites), + } + + +def load_tensor_force_field(d: dict) -> TensorForceField: + return TensorForceField( + potentials=[load_tensor_potential(p) for p in d["potentials"]], + v_sites=None if d["v_sites"] is None else load_tensor_vsites(d["v_sites"]), + ) diff --git a/tyff/_tests/test_serialization.py b/tyff/_tests/test_serialization.py new file mode 100644 index 0000000..d314e13 --- /dev/null +++ b/tyff/_tests/test_serialization.py @@ -0,0 +1,111 @@ +import json + +import pytest +import torch + +from tyff._models import TensorForceField +from tyff._serialization import dump_tensor_force_field, load_tensor_force_field +from tyff._tests.utils import system_from_smiles + + +def compare_tensor_force_fields( + force_field1: TensorForceField, + force_field2: TensorForceField, +) -> bool: + if len(force_field1.potentials) != len(force_field2.potentials): + pytest.fail("Failed: Number of potentials do not match") + + for potential1, potential2 in zip(force_field1.potentials, force_field2.potentials): + if potential1.type != potential2.type: + pytest.fail(f"Failed: potential type mismatch ({potential1.type} != {potential2.type})") + if potential1.fn != potential2.fn: + pytest.fail(f"Failed: potential fn mismatch ({potential1.fn} != {potential2.fn})") + if potential1.parameter_units != potential2.parameter_units: + pytest.fail( + f"Failed: potential parameter_units mismatch " + f"({potential1.parameter_units} != {potential2.parameter_units})" + ) + if potential1.parameter_keys != potential2.parameter_keys: + pytest.fail( + f"Failed: potential parameter_keys mismatch " + f"({potential1.parameter_keys} != {potential2.parameter_keys})" + ) + if potential1.parameter_cols != potential2.parameter_cols: + pytest.fail( + f"Failed: potential parameter_cols mismatch " + f"({potential1.parameter_cols} != {potential2.parameter_cols})" + ) + if not torch.equal(potential1.parameters, potential2.parameters): + pytest.fail("Failed: potential parameters are not equal") + + # potential.attributes can be None + if potential1.attributes is None: + if potential2.attributes is not None: + pytest.fail("Failed: potential1.attributes is None but potential2.attributes is not None") + else: + if not torch.equal(potential1.attributes, potential2.attributes): + pytest.fail("Failed: potential1.attributes and potential2.attributes are not equal") + + if potential1.attribute_units != potential2.attribute_units: + pytest.fail( + f"Failed: potential attribute_units mismatch " + f"({potential1.attribute_units} != {potential2.attribute_units})" + ) + if potential1.attribute_cols != potential2.attribute_cols: + pytest.fail( + f"Failed: potential attribute_cols mismatch " + f"({potential1.attribute_cols} != {potential2.attribute_cols})" + ) + + if force_field1.v_sites is None: + assert force_field2.v_sites is None + + return True + else: + v_sites1 = force_field1.v_sites + v_sites2 = force_field2.v_sites + + if v_sites1.keys != v_sites2.keys: + pytest.fail("Failed: v_site keys mismatch") + + if not torch.equal(v_sites1.parameters, v_sites2.parameters): + pytest.fail("Failed: v_site parameters are not equal") + + for w1, w2 in zip(v_sites1.weights, v_sites2.weights): + if not torch.equal(w1, w2): + pytest.fail("Failed: v_site weights mismatch") + + return True + + +def test_basic_serialization(default_force_field, tmp_path): + + _, tensor_force_field = system_from_smiles( + ["O"], + [1], + default_force_field, + ) + + with open(tmp_path / "tensor_force_field.json", "w") as f: + json.dump(dump_tensor_force_field(tensor_force_field), f) + + with open(tmp_path / "tensor_force_field.json") as f: + loaded = load_tensor_force_field(json.load(f)) + + assert compare_tensor_force_fields(tensor_force_field, loaded) + + +def test_virtual_sites(v_site_force_field, tmp_path): + _, tensor_force_field = system_from_smiles( + ["O"], + [1], + v_site_force_field, + ) + + with open(tmp_path / "tensor_force_field.json", "w") as f: + json.dump(dump_tensor_force_field(tensor_force_field), f) + + with open(tmp_path / "tensor_force_field.json") as f: + loaded = load_tensor_force_field(json.load(f)) + + assert compare_tensor_force_fields(tensor_force_field, loaded)