-
Notifications
You must be signed in to change notification settings - Fork 0
TensorForceField serialization
#178
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Merged
Merged
Changes from all commits
Commits
Show all changes
5 commits
Select commit
Hold shift + click to select a range
1d8202c
Serialize and de-serialize `TensorForceField`
mattwthompson 402aae6
Remove Jacobian from merge
mattwthompson 3d38e67
Lint
mattwthompson 5b147f3
Add basic virtual site test
mattwthompson 93909c6
Remove some dead code
mattwthompson File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| 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} | ||
|
|
||
|
|
||
| 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"]), | ||
| ) | ||
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| 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) |
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
Uh oh!
There was an error while loading. Please reload this page.