-
Notifications
You must be signed in to change notification settings - Fork 34
Expand file tree
/
Copy pathcore.py
More file actions
77 lines (62 loc) · 2.52 KB
/
Copy pathcore.py
File metadata and controls
77 lines (62 loc) · 2.52 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
# SPDX-License-Identifier: MIT
# See LICENSE.md and CONTRIBUTORS.md at https://github.com/SSAGESLabs/PySAGES
from importlib import import_module
from pysages.backends.contexts import JaxMDContext
from pysages.typing import Callable, Optional
class SamplingContext:
"""
PySAGES simulation context. Manages access to the backend-dependent simulation context.
"""
def __init__(
self,
sampling_method,
context_generator: Callable,
callback: Optional[Callable] = None,
context_args: dict = {},
**kwargs,
):
"""
Automatically identifies the backend and binds the sampling method to
the simulation context.
"""
self._backend_name = None
context = context_generator(**context_args)
module_name = type(context).__module__
if module_name.startswith("ase.md"):
self._backend_name = "ase"
elif module_name.startswith("hoomd"):
self._backend_name = "hoomd"
elif isinstance(context, JaxMDContext):
self._backend_name = "jax-md"
elif module_name.startswith("lammps"):
self._backend_name = "lammps"
elif module_name.startswith("simtk.openmm") or module_name.startswith("openmm"):
self._backend_name = "openmm"
if self._backend_name is None:
backends = ", ".join(supported_backends())
raise ValueError(f"Invalid backend {module_name}: supported options are ({backends})")
self.context = context
self.method = sampling_method
self.run = None
backend = import_module("." + self._backend_name, package="pysages.backends")
self.sampler = backend.bind(self, callback, **kwargs)
# `self.run` *must* be set by the backend bind function.
assert self.run is not None
@property
def backend_name(self):
return self._backend_name
def __enter__(self):
"""
Trampoline 'with statements' to the wrapped context when the backend supports it.
"""
if hasattr(self.context, "__enter__"):
return self.context.__enter__()
return self.context
def __exit__(self, exc_type, exc_value, exc_traceback):
"""
Trampoline 'with statements' to the wrapped context when the backend supports it.
"""
if hasattr(self.context, "__exit__"):
self.context.__exit__(exc_type, exc_value, exc_traceback)
def supported_backends():
return ("ase", "hoomd", "jax-md", "lammps", "openmm")