IMP logo
IMP Reference Guide  develop.21ec66f3d8,2026/10/03
The Integrative Modeling Platform
jax.py
1 """@namespace IMP.jax
2  @brief Support for the JAX Python library.
3 
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.
8 """
9 
10 import numpy as np
11 import jax.numpy as jnp
12 import jax.random
13 import jax.tree_util
14 import IMP
15 
16 
17 class Space:
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."""
21 
22  def distance(dr):
23  """If given an array of particle-particle vectors, return an array
24  of distances. If given a single particle-particle vector, return
25  a single distance."""
26  pass
27 
28  def shift(r, dr):
29  """Shift r by dr and return new r"""
30  pass
31 
32  def shift_indexes(r, indexes, dr):
33  """Modify r[indexes] in place by adding dr"""
34  pass
35 
36 
37 class FreeSpace(Space):
38  """An unbounded space with no periodic boundary conditions."""
39 
40  @staticmethod
41  def distance(dr):
42  return jnp.linalg.norm(dr, axis=-1)
43 
44  @staticmethod
45  def shift(r, dr):
46  return r + dr
47 
48  @staticmethod
49  def shift_indexes(r, indexes, dr):
50  return r.at[indexes].add(dr)
51 
52 
54  """A space with periodic boundary conditions.
55 
56  @param side A 3D vector of the periodic boundary box dimensions.
57  """
58 
59  def __init__(self, side):
60  self.side = jnp.asarray(side)
61 
62  def distance(self, dr):
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)
65 
66  def shift(self, r, dr):
67  return jnp.mod(r + dr, self.side)
68 
69  def shift_indexes(self, r, indexes, dr):
70  newr = jnp.mod(r[indexes] + dr, self.side)
71  return r.at[indexes].set(newr)
72 
73 
75  """Get a new JAX random key seeded from IMP's random number generator"""
76  return jax.random.key(IMP.random_number_generator())
77 
78 
79 class _Remap:
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.
83 
84  @param full_view A NumPy view of the entire IMP Model.
85  @param indexes A NumPy array of indexes of the particles that
86  have the attribute.
87  @param mapping A NumPy array that maps original Particle indexes
88  to indexes into the `data` array.
89  """
90  def __init__(self, full_view, indexes, mapping):
91  self.full_view, self.indexes = full_view, indexes
92  self.mapping = mapping
93 
94 
95 @jax.tree_util.register_pytree_node_class
96 class _CompactArray:
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
103  to and from the GPU.
104 
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.
108  """
109 
110  def __init__(self, data, remap):
111  self.data, self.remap = data, remap
112 
113  @classmethod
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)
118  # IMP uses infinity to represent particles without the attribute
119  indexes = np.nonzero(full_view != np.inf)[0]
120  # Any particle not in `indexes` is mapped to the last element,
121  # which is inf (just as in the original full view)
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]))),
126  remap)
127 
128  def tree_flatten(self):
129  # Convert to JAX. Only `data` is sent to the device; everything
130  # else is static (aux data)
131  return (self.data,), self.remap
132 
133  @classmethod
134  def tree_unflatten(cls, aux_data, children):
135  return cls(data=children[0], remap=aux_data)
136 
137  def _map(self, idx):
138  return self.remap.mapping[idx]
139 
140  def __getitem__(self, idx):
141  """Lookup by particle index"""
142  return self.data[self._map(idx)]
143 
144  @property
145  def at(self):
146  """Like jax.Array.at; allow in-place array modification"""
147  return _AtCompactArray(self)
148 
149  def sync(self):
150  """Copy the JAX data back to the IMP Model"""
151  self.remap.full_view[self.remap.indexes] = self.data[:-1]
152 
153 
154 class _AtCompactArray:
155  """Helper class for _CompactArray.at"""
156  def __init__(self, arr):
157  self.arr = arr
158 
159  def __getitem__(self, idx):
160  return _AtIdxCompactArray(self.arr, self.arr._map(idx))
161 
162 
163 class _AtIdxCompactArray:
164  """Helper class for _CompactArray.at[idx]"""
165  def __init__(self, arr, rows):
166  self.arr, self.rows = arr, rows
167 
168  def _new(self, data):
169  return _CompactArray(data, self.arr.remap)
170 
171  def set(self, v):
172  return self._new(self.arr.data.at[self.rows].set(v))
173 
174  def add(self, v):
175  return self._new(self.arr.data.at[self.rows].add(v))
def shift_indexes
Modify r[indexes] in place by adding dr.
Definition: jax.py:32
def get_random_key
Get a new JAX random key seeded from IMP's random number generator.
Definition: jax.py:74
A space with periodic boundary conditions.
Definition: jax.py:53
The space in which restraints are evaluated.
Definition: jax.py:17
def shift
Shift r by dr and return new r.
Definition: jax.py:28
def distance
If given an array of particle-particle vectors, return an array of distances.
Definition: jax.py:22
RandomNumberGenerator random_number_generator
A shared non-GPU random number generator.