From 1d8202cf18d49c5040404a23755cd6698bee01d0 Mon Sep 17 00:00:00 2001 From: "Matthew W. Thompson" Date: Tue, 29 Sep 2026 15:05:02 -0500 Subject: [PATCH 1/5] Serialize and de-serialize `TensorForceField` --- tyff/_serialization.py | 115 ++++++++++++++++++++++++++++++ tyff/_tests/test_serialization.py | 69 ++++++++++++++++++ tyff/compute/_jacobian.py | 107 +++++++++++++++++++++++++++ 3 files changed, 291 insertions(+) create mode 100644 tyff/_serialization.py create mode 100644 tyff/_tests/test_serialization.py create mode 100644 tyff/compute/_jacobian.py diff --git a/tyff/_serialization.py b/tyff/_serialization.py new file mode 100644 index 0000000..f7914c3 --- /dev/null +++ b/tyff/_serialization.py @@ -0,0 +1,115 @@ +from typing import Annotated, Any + +import openff.interchange +import openff.units +import torch +from pydantic import BeforeValidator, PlainSerializer + +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: + if isinstance(v, torch.Tensor): + return v + return torch.tensor(v["data"], dtype=getattr(torch, v["dtype"])).reshape(v["shape"]) + + +_Tensor = Annotated[torch.Tensor, BeforeValidator(_load_tensor), PlainSerializer(_dump_tensor)] + + +def _dump_sparse(t: torch.Tensor) -> dict[str, Any]: + t = t.coalesce() + return { + "dtype": str(t.dtype).removeprefix("torch."), + "shape": list(t.shape), + "indices": t.indices().tolist(), + "values": t.values().tolist(), + } + + +def _load_sparse(v: Any) -> torch.Tensor: + if isinstance(v, torch.Tensor): + return v + return torch.sparse_coo_tensor( + v["indices"], v["values"], size=v["shape"], dtype=getattr(torch, v["dtype"]) + ).coalesce() + + +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} + + +# tyff/_serialization.py +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"]), + ) + + +# tyff/_serialization.py (continued) + + +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.VirtualSiteKey.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..6f0340a --- /dev/null +++ b/tyff/_tests/test_serialization.py @@ -0,0 +1,69 @@ +import torch +import json + +from tyff._serialization import dump_tensor_force_field, load_tensor_force_field +from tyff._tests.utils import system_from_smiles +from tyff._models import TensorForceField + +def compare_tensor_force_fields( + force_field1: TensorForceField, + force_field2: TensorForceField, +) -> bool: + if len(force_field1.potentials) != len(force_field2.potentials): + print("Failed: Number of potentials do not match") + return False + + for potential1, potential2 in zip(force_field1.potentials, force_field2.potentials): + if potential1.type != potential2.type: + print(f"Failed: potential type mismatch ({potential1.type} != {potential2.type})") + return False + if potential1.fn != potential2.fn: + print(f"Failed: potential fn mismatch ({potential1.fn} != {potential2.fn})") + return False + if potential1.parameter_units != potential2.parameter_units: + print(f"Failed: potential parameter_units mismatch ({potential1.parameter_units} != {potential2.parameter_units})") + return False + if potential1.parameter_keys != potential2.parameter_keys: + print(f"Failed: potential parameter_keys mismatch ({potential1.parameter_keys} != {potential2.parameter_keys})") + return False + if potential1.parameter_cols != potential2.parameter_cols: + print(f"Failed: potential parameter_cols mismatch ({potential1.parameter_cols} != {potential2.parameter_cols})") + return False + if not torch.equal(potential1.parameters, potential2.parameters): + print("Failed: potential parameters are not equal") + return False + + # potential.attributes can be None + if potential1.attributes is None: + if potential2.attributes is not None: + print("Failed: potential1.attributes is None but potential2.attributes is not None") + return False + else: + if not torch.equal(potential1.attributes, potential2.attributes): + print("Failed: potential1.attributes and potential2.attributes are not equal") + return False + + if potential1.attribute_units != potential2.attribute_units: + print(f"Failed: potential attribute_units mismatch ({potential1.attribute_units} != {potential2.attribute_units})") + return False + if potential1.attribute_cols != potential2.attribute_cols: + print(f"Failed: potential attribute_cols mismatch ({potential1.attribute_cols} != {potential2.attribute_cols})") + return False + + 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) diff --git a/tyff/compute/_jacobian.py b/tyff/compute/_jacobian.py new file mode 100644 index 0000000..2181f35 --- /dev/null +++ b/tyff/compute/_jacobian.py @@ -0,0 +1,107 @@ +"""Compute the ensemble average(s) and Jacobian matrix associated with a job.""" + +import glob +import json +import pathlib + +import torch +from openff.interchange import Interchange + +from tyff.compute._files import ProductionFiles +from tyff.configs.liquid import BulkLiquid + + +def _get_ensemble_average_and_jacobian( + production_future: dict[str, ProductionFiles], + job_dir: str, +) -> tuple[dict[str, torch.Tensor], torch.Tensor]: + import openmm.unit + import torch + + import tyff.mm + from tyff.mm._ops import _pack_force_field, _unpack_force_field + + # try to safeguard against the previous step failing + assert pathlib.Path(production_future["simulation_files"]["msgpack_trajectory"].filepath).exists() + + # the arguments we really care about are: + # system: tyff.TensorSystem, + # frames_path: pathlib.Path, + # temperature: float, + # pressure: float | None, + # force_field: tyff.TensorForceField, + # + # so grab them from scattered files we expect to be in this job directory. + + # TODO: Handle case of cas simulations + compute_config = BulkLiquid(**json.load(open(f"{job_dir}/compute_config.json"))) # type: ignore[typeddict-item] + + temperature = compute_config["temperature"] # kelvin, float + pressure = compute_config.get("pressure") # atmosphere, float | None + + interchanges = [] + for path_index, interchange_path in enumerate(glob.glob(f"{job_dir}/single_molecule_interchange_*.json")): + unique_molecule_index = int(pathlib.Path(interchange_path).stem.split("single_molecule_interchange_")[-1]) + + # hope we're loading up the single-molecule interchanges in the same order as we have unique molecules + assert unique_molecule_index == path_index, (unique_molecule_index, path_index) + + with open(interchange_path) as f: + interchanges.append(Interchange.model_validate_json(f.read())) + + assert len(interchanges) > 0, "Did not find single-molecule `Interchange`s as expected" + + tensor_force_field, tensor_topologies = tyff.converters.convert_interchange(interchanges) + + # must sync this up with tyff/compute/_pack.py if ever either change + n_molecules = compute_config["n_molecules"] + # also hope ordering lines up + n_copies = [int(n_molecules * x) for x in compute_config["x"]] + + system = tyff.TensorSystem( + topologies=tensor_topologies, + n_copies=n_copies, + is_periodic=True, # need to handle the case of gas simulations + ) + + frames_path = pathlib.Path(f"{job_dir}/production_trajectory.msgpack") + + # Use existing tyff packing order, including attributes and optional v-sites. + tensors, parameter_lookup, attribute_lookup, has_v_sites = _pack_force_field(tensor_force_field) + + assert tensors is not None and len(tensors) > 0 + + parameters = torch.cat([t.detach().reshape(-1) for t in tensors if t is not None]).requires_grad_(True) + pieces = iter(parameters.split([t.numel() for t in tensors if t is not None])) + tensors = tuple(None if t is None else next(pieces).reshape(t.shape) for t in tensors) + worker_ff = _unpack_force_field( + tensors, + parameter_lookup, + attribute_lookup, + has_v_sites, + tensor_force_field, + ) + means, _ = tyff.mm.compute_ensemble_averages( + system, + worker_ff, + frames_path, + temperature * openmm.unit.kelvin, + None if pressure is None else pressure * openmm.unit.atmosphere, + ) + + # originally this was sorted(means), + # https://github.com/openforcefield/tyff/pull/173#discussion_r4126720329 + # + # if that's changed back then we need to make BOTH + # the jacobian and ensemble averages sorted in serialization AND return values + # https://github.com/openforcefield/tyff/pull/173#discussion_r4126720329 + jacobian = torch.stack([torch.autograd.grad(means[name], parameters, retain_graph=True)[0] for name in means]) + + # save ensemble averages and jacobian in each job dir, consider revisiting this decision in the future + with open(f"{job_dir}/ensemble_averages.json", "w") as f: + json.dump({name: value.detach().tolist() for name, value in means.items()}, f) + + with open(f"{job_dir}/jacobian.pt", "wb") as f: + torch.save(jacobian.detach(), f) + + return {name: value.detach() for name, value in means.items()}, jacobian.detach() From 402aae6f5a5c8a2d06fc67b743a124f751409205 Mon Sep 17 00:00:00 2001 From: "Matthew W. Thompson" Date: Wed, 30 Sep 2026 11:32:59 -0500 Subject: [PATCH 2/5] Remove Jacobian from merge --- tyff/compute/_jacobian.py | 107 -------------------------------------- 1 file changed, 107 deletions(-) delete mode 100644 tyff/compute/_jacobian.py diff --git a/tyff/compute/_jacobian.py b/tyff/compute/_jacobian.py deleted file mode 100644 index 2181f35..0000000 --- a/tyff/compute/_jacobian.py +++ /dev/null @@ -1,107 +0,0 @@ -"""Compute the ensemble average(s) and Jacobian matrix associated with a job.""" - -import glob -import json -import pathlib - -import torch -from openff.interchange import Interchange - -from tyff.compute._files import ProductionFiles -from tyff.configs.liquid import BulkLiquid - - -def _get_ensemble_average_and_jacobian( - production_future: dict[str, ProductionFiles], - job_dir: str, -) -> tuple[dict[str, torch.Tensor], torch.Tensor]: - import openmm.unit - import torch - - import tyff.mm - from tyff.mm._ops import _pack_force_field, _unpack_force_field - - # try to safeguard against the previous step failing - assert pathlib.Path(production_future["simulation_files"]["msgpack_trajectory"].filepath).exists() - - # the arguments we really care about are: - # system: tyff.TensorSystem, - # frames_path: pathlib.Path, - # temperature: float, - # pressure: float | None, - # force_field: tyff.TensorForceField, - # - # so grab them from scattered files we expect to be in this job directory. - - # TODO: Handle case of cas simulations - compute_config = BulkLiquid(**json.load(open(f"{job_dir}/compute_config.json"))) # type: ignore[typeddict-item] - - temperature = compute_config["temperature"] # kelvin, float - pressure = compute_config.get("pressure") # atmosphere, float | None - - interchanges = [] - for path_index, interchange_path in enumerate(glob.glob(f"{job_dir}/single_molecule_interchange_*.json")): - unique_molecule_index = int(pathlib.Path(interchange_path).stem.split("single_molecule_interchange_")[-1]) - - # hope we're loading up the single-molecule interchanges in the same order as we have unique molecules - assert unique_molecule_index == path_index, (unique_molecule_index, path_index) - - with open(interchange_path) as f: - interchanges.append(Interchange.model_validate_json(f.read())) - - assert len(interchanges) > 0, "Did not find single-molecule `Interchange`s as expected" - - tensor_force_field, tensor_topologies = tyff.converters.convert_interchange(interchanges) - - # must sync this up with tyff/compute/_pack.py if ever either change - n_molecules = compute_config["n_molecules"] - # also hope ordering lines up - n_copies = [int(n_molecules * x) for x in compute_config["x"]] - - system = tyff.TensorSystem( - topologies=tensor_topologies, - n_copies=n_copies, - is_periodic=True, # need to handle the case of gas simulations - ) - - frames_path = pathlib.Path(f"{job_dir}/production_trajectory.msgpack") - - # Use existing tyff packing order, including attributes and optional v-sites. - tensors, parameter_lookup, attribute_lookup, has_v_sites = _pack_force_field(tensor_force_field) - - assert tensors is not None and len(tensors) > 0 - - parameters = torch.cat([t.detach().reshape(-1) for t in tensors if t is not None]).requires_grad_(True) - pieces = iter(parameters.split([t.numel() for t in tensors if t is not None])) - tensors = tuple(None if t is None else next(pieces).reshape(t.shape) for t in tensors) - worker_ff = _unpack_force_field( - tensors, - parameter_lookup, - attribute_lookup, - has_v_sites, - tensor_force_field, - ) - means, _ = tyff.mm.compute_ensemble_averages( - system, - worker_ff, - frames_path, - temperature * openmm.unit.kelvin, - None if pressure is None else pressure * openmm.unit.atmosphere, - ) - - # originally this was sorted(means), - # https://github.com/openforcefield/tyff/pull/173#discussion_r4126720329 - # - # if that's changed back then we need to make BOTH - # the jacobian and ensemble averages sorted in serialization AND return values - # https://github.com/openforcefield/tyff/pull/173#discussion_r4126720329 - jacobian = torch.stack([torch.autograd.grad(means[name], parameters, retain_graph=True)[0] for name in means]) - - # save ensemble averages and jacobian in each job dir, consider revisiting this decision in the future - with open(f"{job_dir}/ensemble_averages.json", "w") as f: - json.dump({name: value.detach().tolist() for name, value in means.items()}, f) - - with open(f"{job_dir}/jacobian.pt", "wb") as f: - torch.save(jacobian.detach(), f) - - return {name: value.detach() for name, value in means.items()}, jacobian.detach() From 3d38e67f2551445c7130a6d85c0950412abad7f9 Mon Sep 17 00:00:00 2001 From: "Matthew W. Thompson" Date: Wed, 30 Sep 2026 11:36:50 -0500 Subject: [PATCH 3/5] Lint --- tyff/_tests/test_serialization.py | 49 +++++++++++++++++++++---------- 1 file changed, 34 insertions(+), 15 deletions(-) diff --git a/tyff/_tests/test_serialization.py b/tyff/_tests/test_serialization.py index 6f0340a..8317fb0 100644 --- a/tyff/_tests/test_serialization.py +++ b/tyff/_tests/test_serialization.py @@ -1,57 +1,76 @@ -import torch 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 -from tyff._models import TensorForceField + def compare_tensor_force_fields( - force_field1: TensorForceField, - force_field2: TensorForceField, + force_field1: TensorForceField, + force_field2: TensorForceField, ) -> bool: if len(force_field1.potentials) != len(force_field2.potentials): - print("Failed: Number of potentials do not match") + pytest.fail("Failed: Number of potentials do not match") return False for potential1, potential2 in zip(force_field1.potentials, force_field2.potentials): if potential1.type != potential2.type: - print(f"Failed: potential type mismatch ({potential1.type} != {potential2.type})") + pytest.fail(f"Failed: potential type mismatch ({potential1.type} != {potential2.type})") return False if potential1.fn != potential2.fn: - print(f"Failed: potential fn mismatch ({potential1.fn} != {potential2.fn})") + pytest.fail(f"Failed: potential fn mismatch ({potential1.fn} != {potential2.fn})") return False if potential1.parameter_units != potential2.parameter_units: - print(f"Failed: potential parameter_units mismatch ({potential1.parameter_units} != {potential2.parameter_units})") + pytest.fail( + f"Failed: potential parameter_units mismatch " + f"({potential1.parameter_units} != {potential2.parameter_units})" + ) return False if potential1.parameter_keys != potential2.parameter_keys: - print(f"Failed: potential parameter_keys mismatch ({potential1.parameter_keys} != {potential2.parameter_keys})") + pytest.fail( + f"Failed: potential parameter_keys mismatch " + f"({potential1.parameter_keys} != {potential2.parameter_keys})" + ) return False if potential1.parameter_cols != potential2.parameter_cols: - print(f"Failed: potential parameter_cols mismatch ({potential1.parameter_cols} != {potential2.parameter_cols})") + pytest.fail( + f"Failed: potential parameter_cols mismatch " + f"({potential1.parameter_cols} != {potential2.parameter_cols})" + ) return False if not torch.equal(potential1.parameters, potential2.parameters): - print("Failed: potential parameters are not equal") + pytest.fail("Failed: potential parameters are not equal") return False # potential.attributes can be None if potential1.attributes is None: if potential2.attributes is not None: - print("Failed: potential1.attributes is None but potential2.attributes is not None") + pytest.fail("Failed: potential1.attributes is None but potential2.attributes is not None") return False else: if not torch.equal(potential1.attributes, potential2.attributes): - print("Failed: potential1.attributes and potential2.attributes are not equal") + pytest.fail("Failed: potential1.attributes and potential2.attributes are not equal") return False if potential1.attribute_units != potential2.attribute_units: - print(f"Failed: potential attribute_units mismatch ({potential1.attribute_units} != {potential2.attribute_units})") + pytest.fail( + f"Failed: potential attribute_units mismatch " + f"({potential1.attribute_units} != {potential2.attribute_units})" + ) return False if potential1.attribute_cols != potential2.attribute_cols: - print(f"Failed: potential attribute_cols mismatch ({potential1.attribute_cols} != {potential2.attribute_cols})") + pytest.fail( + f"Failed: potential attribute_cols mismatch " + f"({potential1.attribute_cols} != {potential2.attribute_cols})" + ) return False return True + def test_basic_serialization(default_force_field, tmp_path): _, tensor_force_field = system_from_smiles( From 5b147f3335b7c12b1d77e8cdea9d14ba5170a395 Mon Sep 17 00:00:00 2001 From: "Matthew W. Thompson" Date: Thu, 1 Oct 2026 10:01:36 -0500 Subject: [PATCH 4/5] Add basic virtual site test --- tyff/_serialization.py | 6 +---- tyff/_tests/test_serialization.py | 45 +++++++++++++++++++++++-------- 2 files changed, 35 insertions(+), 16 deletions(-) diff --git a/tyff/_serialization.py b/tyff/_serialization.py index f7914c3..ced61fc 100644 --- a/tyff/_serialization.py +++ b/tyff/_serialization.py @@ -49,7 +49,6 @@ def _load_exceptions(v: Any) -> dict[tuple[int, int], int] | None: return {(i, j): val for i, j, val in v} -# tyff/_serialization.py def dump_tensor_potential(p: TensorPotential) -> dict: return { "type": p.type, @@ -82,9 +81,6 @@ def load_tensor_potential(d: dict) -> TensorPotential: ) -# tyff/_serialization.py (continued) - - def dump_tensor_vsites(v: TensorVSites) -> dict: return { "keys": [k.model_dump() for k in v.keys], @@ -95,7 +91,7 @@ def dump_tensor_vsites(v: TensorVSites) -> dict: def load_tensor_vsites(d: dict) -> TensorVSites: return TensorVSites( - keys=[openff.interchange.models.VirtualSiteKey.model_validate(k) for k in d["keys"]], + 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"]), ) diff --git a/tyff/_tests/test_serialization.py b/tyff/_tests/test_serialization.py index 8317fb0..d314e13 100644 --- a/tyff/_tests/test_serialization.py +++ b/tyff/_tests/test_serialization.py @@ -14,59 +14,66 @@ def compare_tensor_force_fields( ) -> bool: if len(force_field1.potentials) != len(force_field2.potentials): pytest.fail("Failed: Number of potentials do not match") - return False 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})") - return False if potential1.fn != potential2.fn: pytest.fail(f"Failed: potential fn mismatch ({potential1.fn} != {potential2.fn})") - return False if potential1.parameter_units != potential2.parameter_units: pytest.fail( f"Failed: potential parameter_units mismatch " f"({potential1.parameter_units} != {potential2.parameter_units})" ) - return False if potential1.parameter_keys != potential2.parameter_keys: pytest.fail( f"Failed: potential parameter_keys mismatch " f"({potential1.parameter_keys} != {potential2.parameter_keys})" ) - return False if potential1.parameter_cols != potential2.parameter_cols: pytest.fail( f"Failed: potential parameter_cols mismatch " f"({potential1.parameter_cols} != {potential2.parameter_cols})" ) - return False if not torch.equal(potential1.parameters, potential2.parameters): pytest.fail("Failed: potential parameters are not equal") - return False # 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") - return False else: if not torch.equal(potential1.attributes, potential2.attributes): pytest.fail("Failed: potential1.attributes and potential2.attributes are not equal") - return False if potential1.attribute_units != potential2.attribute_units: pytest.fail( f"Failed: potential attribute_units mismatch " f"({potential1.attribute_units} != {potential2.attribute_units})" ) - return False if potential1.attribute_cols != potential2.attribute_cols: pytest.fail( f"Failed: potential attribute_cols mismatch " f"({potential1.attribute_cols} != {potential2.attribute_cols})" ) - return False + + 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 @@ -86,3 +93,19 @@ def test_basic_serialization(default_force_field, tmp_path): 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) From 93909c69a8aac0935e733c285909b335de8fc295 Mon Sep 17 00:00:00 2001 From: "Matthew W. Thompson" Date: Thu, 1 Oct 2026 12:02:32 -0500 Subject: [PATCH 5/5] Remove some dead code --- tyff/_serialization.py | 26 +------------------------- 1 file changed, 1 insertion(+), 25 deletions(-) diff --git a/tyff/_serialization.py b/tyff/_serialization.py index ced61fc..9a08c12 100644 --- a/tyff/_serialization.py +++ b/tyff/_serialization.py @@ -1,9 +1,8 @@ -from typing import Annotated, Any +from typing import Any import openff.interchange import openff.units import torch -from pydantic import BeforeValidator, PlainSerializer from tyff._models import TensorForceField, TensorPotential, TensorVSites @@ -13,32 +12,9 @@ def _dump_tensor(t: torch.Tensor) -> dict[str, Any]: def _load_tensor(v: Any) -> torch.Tensor: - if isinstance(v, torch.Tensor): - return v return torch.tensor(v["data"], dtype=getattr(torch, v["dtype"])).reshape(v["shape"]) -_Tensor = Annotated[torch.Tensor, BeforeValidator(_load_tensor), PlainSerializer(_dump_tensor)] - - -def _dump_sparse(t: torch.Tensor) -> dict[str, Any]: - t = t.coalesce() - return { - "dtype": str(t.dtype).removeprefix("torch."), - "shape": list(t.shape), - "indices": t.indices().tolist(), - "values": t.values().tolist(), - } - - -def _load_sparse(v: Any) -> torch.Tensor: - if isinstance(v, torch.Tensor): - return v - return torch.sparse_coo_tensor( - v["indices"], v["values"], size=v["shape"], dtype=getattr(torch, v["dtype"]) - ).coalesce() - - 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()]