Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
87 changes: 87 additions & 0 deletions tyff/_serialization.py
Original file line number Diff line number Diff line change
@@ -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}
Comment thread
mattwthompson marked this conversation as resolved.


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"]),
)
111 changes: 111 additions & 0 deletions tyff/_tests/test_serialization.py
Original file line number Diff line number Diff line change
@@ -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)
Loading