ekwan16 commited on
Commit
fe222d1
·
verified ·
1 Parent(s): 53b4926

sync code from github@2603923ed5412e943ea2128d491e9c54f907332f

Browse files
magnet/inference.py CHANGED
@@ -138,29 +138,35 @@ def _predict_once(model_H, model_C, solute_atomic_numbers, geometry, atomic_numb
138
  return y_pred_combined
139
 
140
 
141
- def _predict_batch(model_H, model_C, solute_atomic_numbers, geometries, atomic_numbers,
142
- N_atoms_per_solvent, solvent_distance_threshold, device):
143
- """Run every geometry in `geometries` through both heads, one forward pass per head.
144
-
145
- Each geometry becomes one graph, and the graphs are concatenated into a single
146
- torch_geometric Batch, so a list of geometries costs two forward passes rather than two per
147
- geometry. The model recomputes `natoms` from `batch.batch`, so the per-graph `batch` and
148
- `natoms` that yield_data writes are dropped before concatenating and PyG assigns its own.
149
-
150
- Returns one (n_solute,) array per geometry, in the order given.
 
 
 
 
 
 
151
  """
152
- n = len(geometries)
153
  combined = [None] * n
154
  for atom_type in ['H', 'C']:
155
  data_list = []
156
- for geom in geometries:
157
  data = yield_data(
158
  solute_atomic_numbers = solute_atomic_numbers,
159
- geometries = geom,
160
  atomic_numbers = atomic_numbers,
161
  shieldings = None,
162
- N_atoms_per_solvent = N_atoms_per_solvent if N_atoms_per_solvent is not None else 3,
163
- solvent_distance_threshold = solvent_distance_threshold,
164
  atom_type = atom_type,
165
  )
166
  # PyG owns `batch` on a Batch, and the model rebuilds `natoms` from it; keeping the
@@ -191,8 +197,37 @@ def _predict_batch(model_H, model_C, solute_atomic_numbers, geometries, atomic_n
191
  return combined
192
 
193
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
194
  def predict_shieldings(model_H, model_C, solute_atomic_numbers, geometry, atomic_numbers = None, N_atoms_per_solvent = None, solvent_distance_threshold = None, device = 'cpu', n_passes = 1, mirror_average = False, symmetrize = None, max_batch_graphs = MAX_BATCH_GRAPHS):
195
- """Predict 1H/13C shieldings for the solute atoms.
196
 
197
  A single forward pass is NOT deterministic: each edge picks a random local
198
  reference frame (eqV2/edge_rot_mat.py), so with a finite spherical-harmonic grid
@@ -206,47 +241,89 @@ def predict_shieldings(model_H, model_C, solute_atomic_numbers, geometry, atomic
206
  spurious error.
207
 
208
  The passes are independent of one another, so they are run as one batch per head rather than
209
- one forward pass each: n_passes with mirror_average costs 2*n_passes graphs and two forward
210
- passes, not 4*n_passes. `max_batch_graphs` caps how many graphs go through at once, so a large
211
- solvated system at a high n_passes does not have to fit in memory all at once. The answer does
212
- not depend on it.
213
 
214
  `symmetrize` is the old name for `mirror_average` and still works, with a DeprecationWarning.
215
  """
216
  mirror_average = resolve_mirror_average(mirror_average, symmetrize, "predict_shieldings")
 
 
 
 
 
 
 
217
 
218
- solute_atomic_numbers = np.asarray(solute_atomic_numbers)
219
- if atomic_numbers is None:
220
- atomic_numbers = solute_atomic_numbers
221
- else:
222
- # coerce so the element comparisons in yield_data (atomic_numbers == 1) stay array-wise;
223
- # a bare Python list would compare to a scalar False and silently mis-mask
224
- atomic_numbers = np.asarray(atomic_numbers)
225
 
226
- # MagNET was only ever trained on these elements. On anything else it would emit a confident but
227
- # meaningless prediction, so refuse it rather than let bad numbers through (validate the full
228
- # system, which for MagNET-x includes the solvent atoms).
229
- unsupported = sorted(set(np.unique(atomic_numbers).tolist()) - SUPPORTED_ELEMENTS)
230
- if unsupported:
231
- raise ValueError(
232
- f"MagNET supports only the elements {sorted(SUPPORTED_ELEMENTS)} "
233
- f"(H, C, N, O, F, S, Cl); the input contains unsupported atomic numbers {unsupported}. "
234
- f"Exclude molecules with these elements before predicting.")
235
 
236
- geometry = np.asarray(geometry, dtype=float)
237
- geometries = [geometry]
238
- if mirror_average:
239
- reflected = geometry.copy()
240
- reflected[..., 0] = -reflected[..., 0] # mirror across the yz-plane (an improper rotation)
241
- geometries.append(reflected)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
242
 
243
- # every pass of every geometry, run together rather than one at a time
244
- repeated = [geom for geom in geometries for _ in range(n_passes)]
 
 
 
 
 
 
245
 
246
- preds = []
247
- for start in range(0, len(repeated), max_batch_graphs):
248
- preds.extend(_predict_batch(model_H, model_C, solute_atomic_numbers,
249
- repeated[start:start + max_batch_graphs], atomic_numbers,
250
- N_atoms_per_solvent, solvent_distance_threshold, device))
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
251
  # atleast_1d keeps a one-atom solute a (1,) array instead of a 0-d scalar after the per-pass squeeze
252
- return np.atleast_1d(np.mean(preds, axis=0))
 
138
  return y_pred_combined
139
 
140
 
141
+ def _predict_batch(model_H, model_C, graphs, device):
142
+ """Run every graph in `graphs` through both heads, one forward pass per head.
143
+
144
+ Each entry of `graphs` is one geometry to predict for, as
145
+ `(solute_atomic_numbers, geometry, atomic_numbers, N_atoms_per_solvent,
146
+ solvent_distance_threshold)`. They need not be the same molecule: the graphs are concatenated
147
+ into a single torch_geometric Batch, and the answers are split apart again by which graph each
148
+ atom came from, so molecules of different sizes batch together as readily as repeated passes
149
+ of one.
150
+
151
+ The model recomputes `natoms` from `batch.batch`, so the per-graph `batch` and `natoms` that
152
+ yield_data writes are dropped before concatenating and PyG assigns its own.
153
+ `unique_molecule_batch` rides along unread: yield_data consumes it for solvent filtering
154
+ before this point, and no model looks at it, so nothing has to be offset across graphs.
155
+
156
+ Returns one (n_solute,) array per entry, in the order given.
157
  """
158
+ n = len(graphs)
159
  combined = [None] * n
160
  for atom_type in ['H', 'C']:
161
  data_list = []
162
+ for solute_atomic_numbers, geometry, atomic_numbers, per_solvent, threshold in graphs:
163
  data = yield_data(
164
  solute_atomic_numbers = solute_atomic_numbers,
165
+ geometries = geometry,
166
  atomic_numbers = atomic_numbers,
167
  shieldings = None,
168
+ N_atoms_per_solvent = per_solvent if per_solvent is not None else 3,
169
+ solvent_distance_threshold = threshold,
170
  atom_type = atom_type,
171
  )
172
  # PyG owns `batch` on a Batch, and the model rebuilds `natoms` from it; keeping the
 
197
  return combined
198
 
199
 
200
+ def _check_elements(atomic_numbers):
201
+ """Refuse a system holding an element MagNET was never trained on.
202
+
203
+ MagNET was only ever trained on SUPPORTED_ELEMENTS. On anything else it would emit a confident
204
+ but meaningless prediction, so refuse it rather than let bad numbers through (validate the full
205
+ system, which for MagNET-x includes the solvent atoms).
206
+
207
+ Raises:
208
+ ValueError: naming every atomic number that is not supported.
209
+ """
210
+ unsupported = sorted(set(np.unique(atomic_numbers).tolist()) - SUPPORTED_ELEMENTS)
211
+ if unsupported:
212
+ raise ValueError(
213
+ f"MagNET supports only the elements {sorted(SUPPORTED_ELEMENTS)} "
214
+ f"(H, C, N, O, F, S, Cl); the input contains unsupported atomic numbers {unsupported}. "
215
+ f"Exclude molecules with these elements before predicting.")
216
+
217
+
218
+ def _geometries_of(geometry, mirror_average, n_passes):
219
+ """Say which geometries one prediction runs over, the mirror image and the passes included."""
220
+ geometry = np.asarray(geometry, dtype=float)
221
+ geometries = [geometry]
222
+ if mirror_average:
223
+ reflected = geometry.copy()
224
+ reflected[..., 0] = -reflected[..., 0] # mirror across the yz-plane (an improper rotation)
225
+ geometries.append(reflected)
226
+ return [geom for geom in geometries for _ in range(n_passes)]
227
+
228
+
229
  def predict_shieldings(model_H, model_C, solute_atomic_numbers, geometry, atomic_numbers = None, N_atoms_per_solvent = None, solvent_distance_threshold = None, device = 'cpu', n_passes = 1, mirror_average = False, symmetrize = None, max_batch_graphs = MAX_BATCH_GRAPHS):
230
+ """Predict 1H/13C shieldings for the solute atoms of one molecule.
231
 
232
  A single forward pass is NOT deterministic: each edge picks a random local
233
  reference frame (eqV2/edge_rot_mat.py), so with a finite spherical-harmonic grid
 
241
  spurious error.
242
 
243
  The passes are independent of one another, so they are run as one batch per head rather than
244
+ one forward pass each. `predict_shieldings_batch` batches whole molecules together as well,
245
+ and is what to call for more than one.
 
 
246
 
247
  `symmetrize` is the old name for `mirror_average` and still works, with a DeprecationWarning.
248
  """
249
  mirror_average = resolve_mirror_average(mirror_average, symmetrize, "predict_shieldings")
250
+ return predict_shieldings_batch(
251
+ model_H, model_C, [solute_atomic_numbers], [geometry],
252
+ atomic_numbers_list = None if atomic_numbers is None else [atomic_numbers],
253
+ N_atoms_per_solvent = N_atoms_per_solvent,
254
+ solvent_distance_threshold = solvent_distance_threshold, device = device,
255
+ n_passes = n_passes, mirror_average = mirror_average,
256
+ max_batch_graphs = max_batch_graphs)[0]
257
 
 
 
 
 
 
 
 
258
 
259
+ def predict_shieldings_batch(model_H, model_C, solute_atomic_numbers_list, geometries_list, atomic_numbers_list = None, N_atoms_per_solvent = None, solvent_distance_threshold = None, device = 'cpu', n_passes = 1, mirror_average = False, symmetrize = None, max_batch_graphs = MAX_BATCH_GRAPHS):
260
+ """Predict 1H/13C shieldings for many molecules at once.
 
 
 
 
 
 
 
261
 
262
+ Every molecule's passes, and its mirror image where `mirror_average` is set, go through as one
263
+ batch rather than one forward pass each, and molecules share those batches with one another.
264
+ A 31-atom graph occupies very little of a GPU, so what costs the time is the number of forward
265
+ passes and not the size of any one of them: batching across molecules is what fills the card.
266
+
267
+ The molecules need not be the same size or the same shape. Each answer is split out by which
268
+ graph its atoms came from, so a list of molecules comes back as a list of per-atom arrays in
269
+ the order given, exactly as calling `predict_shieldings` on each would have.
270
+
271
+ Args:
272
+ model_H, model_C: the 1H and 13C models to run.
273
+ solute_atomic_numbers_list: one array of solute atomic numbers per molecule.
274
+ geometries_list: one (n, 3) coordinate array per molecule, in the same order.
275
+ atomic_numbers_list: the whole system per molecule where it differs from the solute (the
276
+ explicit-solvent path), or None where every system is its own solute.
277
+ N_atoms_per_solvent, solvent_distance_threshold: the explicit-solvent options, applied to
278
+ every molecule.
279
+ device: where to run.
280
+ n_passes: how many forward passes to average per geometry.
281
+ mirror_average: whether to average over the mirror image as well.
282
+ symmetrize: the old name for `mirror_average`, which still works and warns.
283
+ max_batch_graphs: how many graphs go through one forward pass at most. It bounds memory
284
+ and changes no answer.
285
 
286
+ Returns:
287
+ One (n_solute,) array per molecule, in the order given.
288
+
289
+ Raises:
290
+ ValueError: a molecule holds an element MagNET was not trained on, or the lists differ in
291
+ length.
292
+ """
293
+ mirror_average = resolve_mirror_average(mirror_average, symmetrize, "predict_shieldings_batch")
294
 
295
+ if len(solute_atomic_numbers_list) != len(geometries_list):
296
+ raise ValueError(
297
+ f"got {len(solute_atomic_numbers_list)} solutes and {len(geometries_list)} geometries; "
298
+ f"they name the same molecules and must be the same length")
299
+ if atomic_numbers_list is not None and len(atomic_numbers_list) != len(geometries_list):
300
+ raise ValueError(
301
+ f"got {len(atomic_numbers_list)} systems and {len(geometries_list)} geometries; "
302
+ f"they name the same molecules and must be the same length")
303
+
304
+ # every graph to run, and which molecule each one belongs to
305
+ graphs, molecule_of_graph = [], []
306
+ for i, (solute, geometry) in enumerate(zip(solute_atomic_numbers_list, geometries_list)):
307
+ solute = np.asarray(solute)
308
+ if atomic_numbers_list is None:
309
+ whole = solute
310
+ else:
311
+ # coerce so the element comparisons in yield_data (atomic_numbers == 1) stay array-wise;
312
+ # a bare Python list would compare to a scalar False and silently mis-mask
313
+ whole = np.asarray(atomic_numbers_list[i])
314
+ _check_elements(whole)
315
+ for geom in _geometries_of(geometry, mirror_average, n_passes):
316
+ graphs.append((solute, geom, whole, N_atoms_per_solvent, solvent_distance_threshold))
317
+ molecule_of_graph.append(i)
318
+
319
+ per_graph = []
320
+ for start in range(0, len(graphs), max_batch_graphs):
321
+ per_graph.extend(_predict_batch(model_H, model_C, graphs[start:start + max_batch_graphs],
322
+ device))
323
+
324
+ # average each molecule's own passes, and nothing else's
325
+ gathered = [[] for _ in geometries_list]
326
+ for prediction, molecule in zip(per_graph, molecule_of_graph):
327
+ gathered[molecule].append(prediction)
328
  # atleast_1d keeps a one-atom solute a (1,) array instead of a 0-d scalar after the per-pass squeeze
329
+ return [np.atleast_1d(np.mean(passes, axis=0)) for passes in gathered]
magnet/pyproject.toml CHANGED
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
4
 
5
  [project]
6
  name = "magnet-nmr"
7
- version = "0.2.0"
8
  description = "MagNET: equivariant neural networks for NMR chemical-shift (shielding) prediction"
9
  readme = {text = "MagNET: equivariant neural networks for NMR chemical-shift (shielding) prediction. Models, datasets, install instructions, and usage examples are at https://github.com/ekwan/MagNET", content-type = "text/markdown"}
10
  requires-python = ">=3.9"
 
4
 
5
  [project]
6
  name = "magnet-nmr"
7
+ version = "0.2.1"
8
  description = "MagNET: equivariant neural networks for NMR chemical-shift (shielding) prediction"
9
  readme = {text = "MagNET: equivariant neural networks for NMR chemical-shift (shielding) prediction. Models, datasets, install instructions, and usage examples are at https://github.com/ekwan/MagNET", content-type = "text/markdown"}
10
  requires-python = ">=3.9"
magnet/run_magnet.py CHANGED
@@ -19,7 +19,8 @@ import numpy as np
19
  import torch
20
 
21
  from magnet.model import MagNET_Lightning
22
- from magnet.inference import predict_shieldings, resolve_mirror_average
 
23
 
24
  MODEL_CHECKPOINTS = {
25
  # foundation model (predicts the gas-phase shielding the rovibrational/QCD analysis builds on)
@@ -98,13 +99,10 @@ def _predict_with(key_H, key_C, atomic_numbers_list, geometries_list,
98
  device = _default_device()
99
  model_H = load_model_to_device(MODEL_CHECKPOINTS[key_H], device, checkpoints_dir=checkpoints_dir)
100
  model_C = load_model_to_device(MODEL_CHECKPOINTS[key_C], device, checkpoints_dir=checkpoints_dir)
101
- out = []
102
- for entry in zip(atomic_numbers_list, geometries_list):
103
- atomic_numbers, geometry = entry[0], entry[1]
104
- out.append(np.atleast_1d(predict_shieldings(
105
- model_H, model_C, solute_atomic_numbers=atomic_numbers, geometry=geometry,
106
- device=device, n_passes=n_passes, mirror_average=mirror_average, **predict_kwargs).squeeze()))
107
- return out
108
 
109
 
110
  def compute_MagNET_foundation_shieldings(atomic_numbers_list, geometries_list,
 
19
  import torch
20
 
21
  from magnet.model import MagNET_Lightning
22
+ from magnet.inference import (predict_shieldings, predict_shieldings_batch,
23
+ resolve_mirror_average)
24
 
25
  MODEL_CHECKPOINTS = {
26
  # foundation model (predicts the gas-phase shielding the rovibrational/QCD analysis builds on)
 
99
  device = _default_device()
100
  model_H = load_model_to_device(MODEL_CHECKPOINTS[key_H], device, checkpoints_dir=checkpoints_dir)
101
  model_C = load_model_to_device(MODEL_CHECKPOINTS[key_C], device, checkpoints_dir=checkpoints_dir)
102
+ # one batch across the molecules as well as their passes, rather than a call per molecule
103
+ return [np.atleast_1d(one.squeeze()) for one in predict_shieldings_batch(
104
+ model_H, model_C, list(atomic_numbers_list), list(geometries_list),
105
+ device=device, n_passes=n_passes, mirror_average=mirror_average, **predict_kwargs)]
 
 
 
106
 
107
 
108
  def compute_MagNET_foundation_shieldings(atomic_numbers_list, geometries_list,
tests/test_magnet.py CHANGED
@@ -345,3 +345,75 @@ def test_one_atom_solute_stays_one_dimensional():
345
  [-0.629, 0.629, -0.629], [0.629, -0.629, -0.629]])
346
  out = predict_shieldings(mH, mC, solute_atomic_numbers=z, geometry=xyz, device=dev, n_passes=2)
347
  assert out.ndim == 1 and out.shape == (5,)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
345
  [-0.629, 0.629, -0.629], [0.629, -0.629, -0.629]])
346
  out = predict_shieldings(mH, mC, solute_atomic_numbers=z, geometry=xyz, device=dev, n_passes=2)
347
  assert out.ndim == 1 and out.shape == (5,)
348
+
349
+
350
+ # =====================================================================
351
+ # batching across molecules
352
+ # =====================================================================
353
+
354
+ def test_batch_refuses_mismatched_lists():
355
+ """Two lists naming the same molecules have to be the same length."""
356
+ from magnet.inference import predict_shieldings_batch
357
+ z = np.array([6, 1, 1, 1, 1])
358
+ xyz = np.zeros((5, 3))
359
+ with pytest.raises(ValueError, match="same length"):
360
+ predict_shieldings_batch(None, None, [z, z], [xyz])
361
+ with pytest.raises(ValueError, match="same length"):
362
+ predict_shieldings_batch(None, None, [z], [xyz], atomic_numbers_list=[z, z])
363
+
364
+
365
+ def test_batch_refuses_unsupported_element():
366
+ """A molecule out of vocabulary is refused before any model runs, as the single path does."""
367
+ from magnet.inference import predict_shieldings_batch
368
+ good = np.array([6, 1, 1, 1, 1])
369
+ bad = np.array([15, 1, 1, 1]) # phosphorus
370
+ with pytest.raises(ValueError, match="unsupported atomic numbers"):
371
+ predict_shieldings_batch(None, None, [good, bad], [np.zeros((5, 3)), np.zeros((4, 3))])
372
+
373
+
374
+ @pytest.mark.skipif(not os.path.exists(CKPT), reason="released MagNET-Zero checkpoints not present")
375
+ def test_batch_matches_one_at_a_time():
376
+ """Molecules batched together answer as they do apart, whatever their sizes.
377
+
378
+ The molecules are deliberately of two different atom counts, since what splits the answers
379
+ apart is which graph each atom came from and not any assumption that they match.
380
+ """
381
+ from magnet.inference import predict_shieldings_batch
382
+ dev = "cpu"
383
+ mH = MagNET_Lightning.load_from_checkpoint(os.path.join(CKPT, "MagNET-Zero_1H.ckpt"), map_location=dev).eval()
384
+ mC = MagNET_Lightning.load_from_checkpoint(os.path.join(CKPT, "MagNET-Zero_13C.ckpt"), map_location=dev).eval()
385
+ methane_z = np.array([6, 1, 1, 1, 1])
386
+ methane_xyz = np.array([[0., 0., 0.], [0.629, 0.629, 0.629], [-0.629, -0.629, 0.629],
387
+ [-0.629, 0.629, -0.629], [0.629, -0.629, -0.629]])
388
+ ethanol_z = np.array([6, 6, 8, 1, 1, 1, 1, 1, 1])
389
+ ethanol_xyz = np.array([[1.17, -0.24, 0.], [0., 0.55, 0.], [-1.16, -0.25, 0.],
390
+ [2.08, 0.37, 0.], [1.17, -0.87, 0.89], [1.17, -0.87, -0.89],
391
+ [0.01, 1.19, 0.88], [0.01, 1.19, -0.88], [-1.90, 0.35, 0.]])
392
+ zs, xyzs = [methane_z, ethanol_z], [methane_xyz, ethanol_xyz]
393
+ kw = dict(device=dev, n_passes=30, mirror_average=True)
394
+ apart = [predict_shieldings(mH, mC, solute_atomic_numbers=z, geometry=x, **kw)
395
+ for z, x in zip(zs, xyzs)]
396
+ together = predict_shieldings_batch(mH, mC, zs, xyzs, **kw)
397
+ assert [one.shape for one in together] == [one.shape for one in apart]
398
+ for a, b in zip(apart, together):
399
+ # both average the same frame noise away, so they agree to well inside it
400
+ assert np.max(np.abs(a - b)) < 0.05
401
+
402
+
403
+ @pytest.mark.skipif(not os.path.exists(CKPT), reason="released MagNET-Zero checkpoints not present")
404
+ def test_batch_chunking_does_not_change_the_answer():
405
+ """max_batch_graphs splits the work and nothing else, across molecules as within one."""
406
+ from magnet.inference import predict_shieldings_batch
407
+ dev = "cpu"
408
+ mH = MagNET_Lightning.load_from_checkpoint(os.path.join(CKPT, "MagNET-Zero_1H.ckpt"), map_location=dev).eval()
409
+ mC = MagNET_Lightning.load_from_checkpoint(os.path.join(CKPT, "MagNET-Zero_13C.ckpt"), map_location=dev).eval()
410
+ z = np.array([6, 1, 1, 1, 1])
411
+ xyz = np.array([[0., 0., 0.], [0.629, 0.629, 0.629], [-0.629, -0.629, 0.629],
412
+ [-0.629, 0.629, -0.629], [0.629, -0.629, -0.629]])
413
+ kw = dict(device=dev, n_passes=20, mirror_average=True)
414
+ whole = predict_shieldings_batch(mH, mC, [z, z, z], [xyz, xyz + 0.01, xyz - 0.01],
415
+ max_batch_graphs=1000, **kw)
416
+ chunked = predict_shieldings_batch(mH, mC, [z, z, z], [xyz, xyz + 0.01, xyz - 0.01],
417
+ max_batch_graphs=3, **kw)
418
+ for a, b in zip(whole, chunked):
419
+ assert np.max(np.abs(a - b)) < 0.05