IMP logo
IMP Reference Guide  develop.345e71cb9a,2026/08/09
The Integrative Modeling Platform
samplers.py
1 """@namespace IMP.pmi.samplers
2  Sampling of the system.
3 """
4 
5 import IMP
6 import IMP.core
7 from IMP.pmi.tools import get_restraint_set
8 
9 
10 class _SerialReplicaExchange:
11  """Dummy replica exchange class used in non-MPI builds.
12  It should act similarly to IMP.mpi.ReplicaExchange
13  on a single processor.
14  """
15  def __init__(self):
16  self.__params = {}
17 
18  def get_number_of_replicas(self):
19  return 1
20 
21  def create_temperatures(self, tmin, tmax, nrep):
22  return [tmin]
23 
24  def get_my_index(self):
25  return 0
26 
27  def set_my_parameter(self, key, val):
28  self.__params[key] = val
29 
30  def get_my_parameter(self, key):
31  return self.__params[key]
32 
33  def get_friend_index(self, step):
34  return 0
35 
36  def get_friend_parameter(self, key, findex):
37  return self.get_my_parameter(key)
38 
39  def do_exchange(self, myscore, fscore, findex):
40  return False
41 
42  def set_was_used(self, was_used):
43  self.was_used = was_used
44 
45 
46 class _SamplerBase:
47  def __init__(self, model, start_frame):
48  self.model = model
49  # that is -1 because mc/md has not yet run
50  self.nframe = start_frame - 1
51  self.simulated_annealing = False
52 
53  def set_simulated_annealing(self, min_temp, max_temp, min_temp_time,
54  max_temp_time):
55  self.simulated_annealing = True
56  self.tempmin = min_temp
57  self.tempmax = max_temp
58  self.timemin = min_temp_time
59  self.timemax = max_temp_time
60 
61  def temp_simulated_annealing(self):
62  if self.nframe % (self.timemin + self.timemax) < self.timemin:
63  value = 0.0
64  else:
65  value = 1.0
66  temp = self.tempmin + (self.tempmax - self.tempmin) * value
67  return temp
68 
69 
70 class MonteCarlo(_SamplerBase):
71  """Sample using Monte Carlo"""
72 
73  # check that isd is installed
74  try:
75  import IMP.isd
76  isd_available = True
77  except ImportError:
78  isd_available = False
79 
80  def __init__(self, model, objects=None, temp=1.0, filterbyname=None,
81  score_moved=False, start_frame=0):
82  """Setup Monte Carlo sampling
83  @param model The IMP Model
84  @param objects What to sample (a list of Movers)
85  @param temp The MC temperature
86  @param filterbyname Not used
87  @param score_moved If True, attempt to speed up sampling by
88  caching scoring function terms on particles that didn't move
89  @param start_frame The starting frame number
90  """
91  super().__init__(model, start_frame=start_frame)
92  self.losp = [
93  "Rigid_Bodies",
94  "Floppy_Bodies",
95  "Nuisances",
96  "X_coord",
97  "Weights"
98  "Surfaces"]
99  self.selfadaptive = False
100  self.temp = temp
101  self.mvs = []
102  self.mvslabels = []
103  self.label = "None"
104  self.movers_data = {}
105  self.use_jax = False
106  self._jax_optimizer = None
107  self._jax_state = None
108 
109  self.mvs = objects
110 
111  # SerialMover
112  self.smv = IMP.core.SerialMover(self.mvs)
113 
114  self.mc = IMP.core.MonteCarlo(self.model)
115  self.mc.set_scoring_function(get_restraint_set(self.model))
116  self.mc.set_return_best(False)
117  self.mc.set_score_moved(score_moved)
118  self.mc.set_kt(self.temp)
119  self.mc.add_mover(self.smv)
120 
121  def set_use_jax(self, nstep):
122  """Request that sampling of the scoring function is done using
123  JAX instead of IMP's internal C++ implementation (requires
124  that all PMI restraints used have a JAX implementation)."""
125  self.use_jax = True
126  self._jax_optimizer = self.mc._get_jax_optimizer(
127  nstep * self.get_number_of_movers())
128  self._jax_state = self._jax_optimizer.get_initial_state()
129 
130  def get_jax_model(self):
131  """Get the current JAX Model used by the sampler."""
132  return self._jax_state.jm
133 
134  def set_kt(self, temp):
135  self.temp = temp
136  self.mc.set_kt(temp)
137  if self._jax_state is not None:
138  self._jax_state.temperature = temp
139 
140  def get_mc(self):
141  return self.mc
142 
143  def set_scoring_function(self, objectlist):
144  rs = IMP.RestraintSet(self.model, 1.0, 'sfo')
145  for ob in objectlist:
146  rs.add_restraint(ob.get_restraint())
148  self.mc.set_scoring_function(sf)
149 
150  def set_self_adaptive(self, isselfadaptive=True):
151  self.selfadaptive = isselfadaptive
152 
153  def get_number_of_movers(self):
154  return len(self.smv.get_movers())
155 
156  def get_particle_types(self):
157  return self.losp
158 
159  def optimize(self, nstep):
160  self.nframe += 1
161  if self.use_jax:
162  score, self._jax_state = self._jax_optimizer.optimize(
163  self._jax_state)
164  else:
165  score = self.mc.optimize(nstep * self.get_number_of_movers())
166 
167  # apply simulated annealing protocol
168  if self.simulated_annealing:
169  self.set_kt(self.temp_simulated_annealing())
170 
171  # apply self adaptive protocol
172  if self.selfadaptive:
173  self.apply_self_adaptive()
174  return score
175 
177  """Modify parameters of individual movers to try to keep acceptance
178  rate around 50%"""
179  if self.use_jax:
180  raise NotImplementedError(
181  "Adaptive protocol is not yet implemented for JAX")
182  for i, mv in enumerate(self.mvs):
183 
184  mvacc = mv.get_number_of_accepted()
185  mvprp = mv.get_number_of_proposed()
186  if mv not in self.movers_data:
187  accept = float(mvacc) / float(mvprp)
188  self.movers_data[mv] = (mvacc, mvprp)
189  else:
190  oldmvacc, oldmvprp = self.movers_data[mv]
191  accept = float(mvacc-oldmvacc) / float(mvprp-oldmvprp)
192  self.movers_data[mv] = (mvacc, mvprp)
193  if accept < 0.05:
194  accept = 0.05
195  if accept > 1.0:
196  accept = 1.0
197 
198  if isinstance(mv, IMP.core.NormalMover):
199  stepsize = mv.get_sigma()
200  if 0.4 > accept or accept > 0.6:
201  mv.set_sigma(stepsize * 2 * accept)
202 
203  if isinstance(mv, IMP.isd.WeightMover):
204  stepsize = mv.get_radius()
205  if 0.4 > accept or accept > 0.6:
206  mv.set_radius(stepsize * 2 * accept)
207 
208  if isinstance(mv, IMP.core.RigidBodyMover):
209  mr = mv.get_maximum_rotation()
210  mt = mv.get_maximum_translation()
211  if 0.4 > accept or accept > 0.6:
212  mv.set_maximum_rotation(mr * 2 * accept)
213  mv.set_maximum_translation(mt * 2 * accept)
214 
215  if isinstance(mv, IMP.pmi.TransformMover):
216  mr = mv.get_maximum_rotation()
217  mt = mv.get_maximum_translation()
218  if 0.4 > accept or accept > 0.6:
219  mv.set_maximum_rotation(mr * 2 * accept)
220  mv.set_maximum_translation(mt * 2 * accept)
221 
222  if isinstance(mv, IMP.core.BallMover):
223  mr = mv.get_radius()
224  if 0.4 > accept or accept > 0.6:
225  mv.set_radius(mr * 2 * accept)
226 
227  def set_label(self, label):
228  self.label = label
229 
230  def get_frame_number(self):
231  return self.nframe
232 
233  def get_output(self):
234  output = {}
235  for i, mv in enumerate(self.smv.get_movers()):
236  mvname = mv.get_name()
237  mvacc = mv.get_number_of_accepted()
238  mvprp = mv.get_number_of_proposed()
239  try:
240  mvacr = float(mvacc) / float(mvprp)
241  except: # noqa: E722
242  mvacr = 0.0
243  output["MonteCarlo_Acceptance_" +
244  mvname + "_" + str(i)] = str(mvacr)
245  if "Nuisances" in mvname:
246  output["MonteCarlo_StepSize_" + mvname + "_" + str(i)] = \
247  str(IMP.core.NormalMover.get_from(mv).get_sigma())
248  if "Weights" in mvname:
249  output["MonteCarlo_StepSize_" + mvname + "_" + str(i)] = \
250  str(IMP.isd.WeightMover.get_from(mv).get_radius())
251  output["MonteCarlo_Temperature"] = str(self.mc.get_kt())
252  output["MonteCarlo_Nframe"] = str(self.nframe)
253  return output
254 
255 
256 class MolecularDynamics(_SamplerBase):
257  """Sample using molecular dynamics"""
258 
259  def __init__(self, model, objects, kt, gamma=0.01, maximum_time_step=1.0,
260  sf=None, use_jax=False, start_frame=0):
261  """Setup MD
262  @param model The IMP Model
263  @param objects What to sample. Use flat list of particles
264  @param kt Temperature
265  @param gamma Viscosity parameter
266  @param maximum_time_step MD max time step
267  @param start_frame The starting frame number
268  """
269  super().__init__(model, start_frame=start_frame)
270 
271  # check if using PMI1 objects dictionary, or just list of particles
272  try:
273  for obj in objects:
274  psamp = obj.get_particles_to_sample()
275  to_sample = psamp['Floppy_Bodies_SimplifiedModel'][0]
276  except: # noqa: E722
277  to_sample = objects
278 
280  self.model, to_sample, kt/0.0019872041, gamma)
281  self.md = IMP.atom.MolecularDynamics(self.model)
282  self.md.set_maximum_time_step(maximum_time_step)
283  if sf:
284  self.md.set_scoring_function(sf)
285  else:
286  self.md.set_scoring_function(get_restraint_set(self.model))
287  self.md.add_optimizer_state(self.ltstate)
288 
289  def set_use_jax(self, nstep):
290  """Request that sampling of the scoring function is done using
291  JAX instead of IMP's internal C++ implementation (requires
292  that all PMI restraints used have a JAX implementation)."""
293  raise NotImplementedError("JAX currently only supported for MC")
294 
295  def set_kt(self, kt):
296  temp = kt/0.0019872041
297  self.ltstate.set_temperature(temp)
298  self.md.assign_velocities(temp)
299 
300  def set_gamma(self, gamma):
301  self.ltstate.set_gamma(gamma)
302 
303  def optimize(self, nsteps):
304  # apply simulated annealing protocol
305  self.nframe += 1
306  if self.simulated_annealing:
307  self.set_kt(self.temp_simulated_annealing())
308  return self.md.optimize(nsteps)
309 
310  def get_output(self):
311  output = {}
312  output["MolecularDynamics_KineticEnergy"] = \
313  str(self.md.get_kinetic_energy())
314  return output
315 
316 
318  """Sample using conjugate gradients"""
319 
320  def __init__(self, model, objects):
321  self.model = model
322  self.nframe = -1
323  self.cg = IMP.core.ConjugateGradients(self.model)
324  self.cg.set_scoring_function(get_restraint_set(self.model))
325 
326  def set_label(self, label):
327  self.label = label
328 
329  def get_frame_number(self):
330  return self.nframe
331 
332  def optimize(self, nstep):
333  self.nframe += 1
334  self.cg.optimize(nstep)
335 
336  def set_scoring_function(self, objectlist):
337  rs = IMP.RestraintSet(self.model, 1.0, 'sfo')
338  for ob in objectlist:
339  rs.add_restraint(ob.get_restraint())
341  self.cg.set_scoring_function(sf)
342 
343  def get_output(self):
344  output = {}
345  output["ConjugatedGradients_Nframe"] = str(self.nframe)
346  return output
347 
348 
349 class _ReplicaExchangeStats:
350  """Statistics for replica exchange.
351  This is in a separate class so that we can pickle it easily for
352  restarts"""
353  def __init__(self):
354  self.nattempts = 0
355  self.nmintemp = 0
356  self.nmaxtemp = 0
357  self.nsuccess = 0
358 
359  def get_output(self):
360  output = {}
361  if self.nattempts != 0:
362  output["ReplicaExchange_SwapSuccessRatio"] = str(
363  float(self.nsuccess) / self.nattempts)
364  output["ReplicaExchange_MinTempFrequency"] = str(
365  float(self.nmintemp) / self.nattempts)
366  output["ReplicaExchange_MaxTempFrequency"] = str(
367  float(self.nmaxtemp) / self.nattempts)
368  else:
369  output["ReplicaExchange_SwapSuccessRatio"] = str(0)
370  output["ReplicaExchange_MinTempFrequency"] = str(0)
371  output["ReplicaExchange_MaxTempFrequency"] = str(0)
372  return output
373 
374 
376  """Sample using replica exchange"""
377 
378  def __init__(self, model, tempmin, tempmax, samplerobjects, test=True,
379  replica_exchange_object=None):
380  '''
381  samplerobjects can be a list of MonteCarlo or MolecularDynamics
382  '''
383 
384  self.model = model
385  self.samplerobjects = samplerobjects
386  # min and max temperature
387  self.TEMPMIN_ = tempmin
388  self.TEMPMAX_ = tempmax
389 
390  if replica_exchange_object is None:
391  # initialize Replica Exchange class
392  try:
393  import IMP.mpi
394  print('ReplicaExchange: MPI was found. '
395  'Using Parallel Replica Exchange')
396  self.rem = IMP.mpi.ReplicaExchange()
397  except ImportError:
398  print('ReplicaExchange: Could not find MPI. '
399  'Using Serial Replica Exchange')
400  self.rem = _SerialReplicaExchange()
401 
402  else:
403  # get the replica exchange class instance from elsewhere
404  print('got existing rex object')
405  self.rem = replica_exchange_object
406 
407  # get number of replicas
408  nproc = self.rem.get_number_of_replicas()
409 
410  if nproc % 2 != 0 and not test:
411  raise Exception(
412  "number of replicas has to be even. "
413  "set test=True to run with odd number of replicas.")
414  # create array of temperatures, in geometric progression
415  temp = self.rem.create_temperatures(
416  self.TEMPMIN_,
417  self.TEMPMAX_,
418  nproc)
419  # get replica index
420  self.temperatures = temp
421 
422  myindex = self.rem.get_my_index()
423  # set initial value of the parameter (temperature) to exchange
424  self.rem.set_my_parameter("temp", [self.temperatures[myindex]])
425  for so in self.samplerobjects:
426  so.set_kt(self.temperatures[myindex])
427  # Acceptance, etc. statistics
428  self.stats = _ReplicaExchangeStats()
429 
430  def get_temperatures(self):
431  return self.temperatures
432 
433  def get_my_temp(self):
434  return self.rem.get_my_parameter("temp")[0]
435 
436  def get_my_index(self):
437  return self.rem.get_my_index()
438 
439  def swap_temp(self, nframe, score=None):
440  if score is None:
441  score = self.model.evaluate(False)
442  # get my replica index and temperature
443  _ = self.rem.get_my_index()
444  mytemp = self.rem.get_my_parameter("temp")[0]
445 
446  if mytemp == self.TEMPMIN_:
447  self.stats.nmintemp += 1
448 
449  if mytemp == self.TEMPMAX_:
450  self.stats.nmaxtemp += 1
451 
452  # score divided by kbt
453  myscore = score / mytemp
454 
455  # get my friend index and temperature
456  findex = self.rem.get_friend_index(nframe)
457  ftemp = self.rem.get_friend_parameter("temp", findex)[0]
458  # score divided by kbt
459  fscore = score / ftemp
460 
461  # try exchange
462  flag = self.rem.do_exchange(myscore, fscore, findex)
463 
464  self.stats.nattempts += 1
465  # if accepted, change temperature
466  if (flag):
467  for so in self.samplerobjects:
468  so.set_kt(ftemp)
469  self.stats.nsuccess += 1
470 
471  def get_output(self):
472  output = self.stats.get_output()
473  output["ReplicaExchange_CurrentTemp"] = str(self.get_my_temp())
474  return output
475 
476 
477 class MPI_values:
478  def __init__(self, replica_exchange_object=None):
479  """Query values (ie score, and others)
480  from a set of parallel jobs"""
481 
482  if replica_exchange_object is None:
483  # initialize Replica Exchange class
484  try:
485  import IMP.mpi
486  print('MPI_values: MPI was found. '
487  'Using Parallel Replica Exchange')
488  self.rem = IMP.mpi.ReplicaExchange()
489  except ImportError:
490  print('MPI_values: Could not find MPI. '
491  'Using Serial Replica Exchange')
492  self.rem = _SerialReplicaExchange()
493 
494  else:
495  # get the replica exchange class instance from elsewhere
496  print('got existing rex object')
497  self.rem = replica_exchange_object
498 
499  def set_value(self, name, value):
500  self.rem.set_my_parameter(name, [value])
501 
502  def get_values(self, name):
503  values = []
504  for i in range(self.rem.get_number_of_replicas()):
505  v = self.rem.get_friend_parameter(name, i)[0]
506  values.append(v)
507  return values
508 
509  def get_percentile(self, name):
510  value = self.rem.get_my_parameter(name)[0]
511  values = sorted(self.get_values(name))
512  ind = values.index(value)
513  percentile = float(ind)/len(values)
514  return percentile
def __init__
samplerobjects can be a list of MonteCarlo or MolecularDynamics
Definition: samplers.py:378
A Monte Carlo optimizer.
Definition: MonteCarlo.h:46
A class to implement Hamiltonian Replica Exchange.
Maintains temperature during molecular dynamics.
Sample using molecular dynamics.
Definition: samplers.py:256
Miscellaneous utilities.
Definition: pmi/tools.py:1
Modify the transformation of a rigid body.
Simple conjugate gradients optimizer.
def set_use_jax
Request that sampling of the scoring function is done using JAX instead of IMP's internal C++ impleme...
Definition: samplers.py:121
Sample using conjugate gradients.
Definition: samplers.py:317
Create a scoring function on a list of restraints.
Move continuous particle variables by perturbing them within a ball.
Object used to hold a set of restraints.
Definition: RestraintSet.h:41
Simple molecular dynamics simulator.
Code that uses the MPI parallel library.
def set_use_jax
Request that sampling of the scoring function is done using JAX instead of IMP's internal C++ impleme...
Definition: samplers.py:289
A mover that perturbs a Weight particle.
Definition: WeightMover.h:20
Modify the transformation of a rigid body.
def __init__
Setup Monte Carlo sampling.
Definition: samplers.py:80
Modify a set of continuous variables using a normal distribution.
Definition: NormalMover.h:23
Basic functionality that is expected to be used by a wide variety of IMP users.
Sample using Monte Carlo.
Definition: samplers.py:70
The general base class for IMP exceptions.
Definition: exception.h:48
def apply_self_adaptive
Modify parameters of individual movers to try to keep acceptance rate around 50%. ...
Definition: samplers.py:176
Applies a list of movers one at a time.
Definition: SerialMover.h:26
def get_jax_model
Get the current JAX Model used by the sampler.
Definition: samplers.py:130
Sample using replica exchange.
Definition: samplers.py:375
Inferential scoring building on methods developed as part of the Inferential Structure Determination ...