2 @brief Support for the JAX Python library.
4 IMP currently has rudimentary support for running on a graphics
5 processing unit (GPU) or similar systems such as
6 Tensor Processing Units (TPUs). This support uses the
7 [JAX](https://docs.jax.dev/) Python library.
11 import jax.numpy
as jnp
18 """The space in which restraints are evaluated. See FreeSpace for the
19 default unbounded space, or PeriodicSpace for a space that implements
20 periodic boundary conditions."""
23 """If given an array of particle-particle vectors, return an array
24 of distances. If given a single particle-particle vector, return
29 """Shift r by dr and return new r"""
33 """Modify r[indexes] in place by adding dr"""
37 class FreeSpace(Space):
38 """An unbounded space with no periodic boundary conditions."""
42 return jnp.linalg.norm(dr, axis=-1)
50 return r.at[indexes].add(dr)
54 """A space with periodic boundary conditions.
56 @param side A 3D vector of the periodic boundary box dimensions.
59 def __init__(self, side):
60 self.side = jnp.asarray(side)
63 p_dr = jnp.mod(dr + self.side * 0.5, self.side) - 0.5 * self.side
64 return jnp.linalg.norm(p_dr, axis=-1)
66 def shift(self, r, dr):
67 return jnp.mod(r + dr, self.side)
70 newr = jnp.mod(r[indexes] + dr, self.side)
71 return r.at[indexes].set(newr)
75 """Get a new JAX random key seeded from IMP's random number generator"""
80 """A mapping from full particle indexes to and from the compact array.
81 This is used in _CompactArray, below. It is hashed by identity so
82 that it (unlike the arrays it contains) can be used as pytree aux data.
84 @param full_view A NumPy view of the entire IMP Model.
85 @param indexes A NumPy array of indexes of the particles that
87 @param mapping A NumPy array that maps original Particle indexes
88 to indexes into the `data` array.
90 def __init__(self, full_view, indexes, mapping):
91 self.full_view, self.indexes = full_view, indexes
92 self.mapping = mapping
95 @jax.tree_util.register_pytree_node_class
97 """Access an IMP Model Float attribute as a compacted or sparse array.
98 An IMP attribute array (returned by Model.get_numpy()) can be sparsely
99 populated. Compact it down to a flat array of only the particles that
100 have the attribute. Original Particle indexes are mapped to indexes
101 into the compacted array at JAX trace time. This should make JAX code
102 more performant since only the compacted array needs to be transferred
105 @param data Compact array of only the used particles.
106 @param remap A _Remap object that maps full particle indexes to and
107 from the compact array.
110 def __init__(self, data, remap):
111 self.data, self.remap = data, remap
114 def from_model(cls, m, fk):
115 """Create a new CompactArray for the given FloatKey `fk` in the
116 given IMP Model `m`."""
117 full_view = m.get_numpy(fk)
119 indexes = np.nonzero(full_view != np.inf)[0]
122 mapping = np.full(len(full_view), len(full_view) + 1, dtype=np.int32)
123 mapping[indexes] = np.arange(len(indexes), dtype=np.int32)
124 remap = _Remap(full_view, indexes, mapping)
125 return cls(np.concatenate((full_view[indexes], np.array([np.inf]))),
128 def tree_flatten(self):
131 return (self.data,), self.remap
134 def tree_unflatten(cls, aux_data, children):
135 return cls(data=children[0], remap=aux_data)
138 return self.remap.mapping[idx]
140 def __getitem__(self, idx):
141 """Lookup by particle index"""
142 return self.data[self._map(idx)]
146 """Like jax.Array.at; allow in-place array modification"""
147 return _AtCompactArray(self)
150 """Copy the JAX data back to the IMP Model"""
151 self.remap.full_view[self.remap.indexes] = self.data[:-1]
154 class _AtCompactArray:
155 """Helper class for _CompactArray.at"""
156 def __init__(self, arr):
159 def __getitem__(self, idx):
160 return _AtIdxCompactArray(self.arr, self.arr._map(idx))
163 class _AtIdxCompactArray:
164 """Helper class for _CompactArray.at[idx]"""
165 def __init__(self, arr, rows):
166 self.arr, self.rows = arr, rows
168 def _new(self, data):
169 return _CompactArray(data, self.arr.remap)
172 return self._new(self.arr.data.at[self.rows].set(v))
175 return self._new(self.arr.data.at[self.rows].add(v))
def shift_indexes
Modify r[indexes] in place by adding dr.
def get_random_key
Get a new JAX random key seeded from IMP's random number generator.
A space with periodic boundary conditions.
The space in which restraints are evaluated.
def shift
Shift r by dr and return new r.
def distance
If given an array of particle-particle vectors, return an array of distances.
RandomNumberGenerator random_number_generator
A shared non-GPU random number generator.