File size: 22,566 Bytes
e16c7ac | 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 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 | """Adversarial edge-case tests for the magnet public API (magnet/api.py).
Pure-logic tests (single-vs-list detection, validation, solvent checking, solute selection) run
whenever the torch stack is installed. Tests marked @needs_model load real checkpoints and are
skipped unless the weights are present; they stay few and fast (n_passes=1, symmetrize=False, tiny
molecules). Run the whole suite with `pytest`; to include the real-model tests, check out the
Hugging Face copy so the weights land in model_checkpoints/.
"""
import os
import warnings
import numpy as np
import pytest
warnings.filterwarnings("ignore")
import torch
import magnet.api as api
from magnet import scaling
from magnet.eqV2.edge_rot_mat import init_edge_rot_mat
_HERE = os.path.dirname(os.path.abspath(__file__))
_REPO = os.path.dirname(_HERE) # magnet/ -> repo root
_CKPT_REL = os.path.join("model_checkpoints", "MagNET-Zero", "MagNET-Zero_1H.ckpt")
_HAVE_MODELS = os.path.exists(os.path.join(_REPO, _CKPT_REL))
needs_model = pytest.mark.skipif(not _HAVE_MODELS, reason="checkpoints not available")
# ---- tiny geometries (validity irrelevant for API-contract tests) ----
CH4_Z = np.array([6, 1, 1, 1, 1])
CH4_XYZ = np.array([[0., 0., 0.],
[0.629, 0.629, 0.629],
[-0.629, -0.629, 0.629],
[-0.629, 0.629, -0.629],
[0.629, -0.629, -0.629]])
# formaldehyde CH2O: exercises H, C, and an O (should be NaN in shifts)
CH2O_Z = np.array([6, 8, 1, 1])
CH2O_XYZ = np.array([[0., 0., 0.], [1.2, 0., 0.], [-0.5, 0.94, 0.], [-0.5, -0.94, 0.]])
# one chloroform molecule, atoms in molecule order C H Cl Cl Cl
CHCL3_Z = np.array([6, 1, 17, 17, 17])
CHCL3_XYZ = np.array([[0., 0., 0.], [0., 0., 1.1], [1.7, 0., -0.4],
[-0.85, 1.47, -0.4], [-0.85, -1.47, -0.4]])
# =====================================================================
# _as_batch: single-vs-list detection and validation
# =====================================================================
def test_single_ndarray_is_single():
an, xyz, single = api._as_batch(CH4_Z, CH4_XYZ)
assert single is True and len(an) == 1 and an[0].shape == (5,) and xyz[0].shape == (5, 3)
def test_single_pythonlist_is_single():
an, xyz, single = api._as_batch([6, 1, 1, 1, 1], CH4_XYZ.tolist())
assert single is True and len(xyz) == 1
def test_one_atom_molecule_is_single():
an, xyz, single = api._as_batch(np.array([1]), np.array([[0., 0., 0.]]))
assert single is True and xyz[0].shape == (1, 3)
def test_list_of_molecules_is_batch():
an, xyz, single = api._as_batch([CH4_Z, CH2O_Z], [CH4_XYZ, CH2O_XYZ])
assert single is False and len(an) == 2
def test_batch_size_one_list_is_batch():
an, xyz, single = api._as_batch([CH4_Z], [CH4_XYZ])
assert single is False and len(an) == 1
def test_3d_ndarray_is_batch():
Z = np.stack([CH4_Z, CH4_Z]); X = np.stack([CH4_XYZ, CH4_XYZ]) # (2,5) and (2,5,3)
an, xyz, single = api._as_batch(Z, X)
assert single is False and len(xyz) == 2 and xyz[0].shape == (5, 3)
def test_ragged_batch_ok():
an, xyz, single = api._as_batch([CH4_Z, CH2O_Z], [CH4_XYZ, CH2O_XYZ])
assert xyz[0].shape == (5, 3) and xyz[1].shape == (4, 3)
def test_empty_list_raises():
with pytest.raises(ValueError, match="empty input"):
api._as_batch([], [])
def test_empty_ndarray_raises():
with pytest.raises(ValueError, match="empty input"):
api._as_batch(np.array([]), np.zeros((0, 3)))
def test_mismatched_atom_counts_single_raises():
with pytest.raises(ValueError):
api._as_batch(np.array([6, 1, 1]), np.zeros((4, 3)))
def test_wrong_last_dim_raises():
with pytest.raises(ValueError):
api._as_batch(np.array([6, 1]), np.zeros((2, 2)))
def test_mismatched_list_lengths_raises():
with pytest.raises(ValueError):
api._as_batch([CH4_Z, CH4_Z], [CH4_XYZ, CH4_XYZ, CH4_XYZ])
def test_coords_1d_single_atom_flat_is_misread():
# coordinates=(3,) flat for one atom: coordinates[0] is a scalar (ndim 0) -> not single ->
# treated as a 3-molecule batch of 0-d coords -> raises. Documents the (N,3) requirement.
with pytest.raises(ValueError):
api._as_batch(np.array([6]), np.array([0., 0., 0.]))
# =====================================================================
# _check_passes
# =====================================================================
@pytest.mark.parametrize("bad", [0, -1, 1.5, "3", None, 2.0])
def test_check_passes_rejects(bad):
with pytest.raises(ValueError, match="n_passes"):
api._check_passes(bad)
@pytest.mark.parametrize("good", [1, 10, np.int64(5)])
def test_check_passes_accepts(good):
api._check_passes(good) # no raise
def test_check_passes_bool_quirk():
# bool is a subclass of int in Python; True passes as "1 pass", False is rejected.
api._check_passes(True) # accepted (quirk, not a real bug)
with pytest.raises(ValueError):
api._check_passes(False)
# =====================================================================
# predict_shieldings argument validation (pure logic, before model load)
# =====================================================================
def test_predict_shieldings_bad_model():
with pytest.raises(ValueError, match="model must be one of"):
api.predict_shieldings(CH4_Z, CH4_XYZ, model="bogus")
@pytest.mark.parametrize("m,fn", [("MagNET-PCM", "implicit_solvent_correction"),
("MagNET-x", "explicit_solvent_correction")])
def test_predict_shieldings_correction_redirect(m, fn):
with pytest.raises(ValueError, match=fn):
api.predict_shieldings(CH4_Z, CH4_XYZ, model=m)
def test_predict_shieldings_bad_passes():
with pytest.raises(ValueError, match="n_passes"):
api.predict_shieldings(CH4_Z, CH4_XYZ, n_passes=0)
# =====================================================================
# predict_shifts solvent validation (pure logic, before model load)
# =====================================================================
def test_predict_shifts_unknown_solvent():
with pytest.raises(ValueError, match="unknown solvent"):
api.predict_shifts(CH4_Z, CH4_XYZ, solvent="ethanol")
def test_predict_shifts_solvent_case_sensitive():
with pytest.raises(ValueError, match="unknown solvent"):
api.predict_shifts(CH4_Z, CH4_XYZ, solvent="Chloroform")
def test_predict_shifts_water_not_rejected_by_validation():
# "water" must map to TIP4P and NOT trip the unknown-solvent guard.
tables = scaling.published_scaling_tables()
assert "TIP4P" in tables["C"] and "TIP4P" in tables["H"]
# the error message offers "water" (not "TIP4P") as an option:
options = ["water" if s == "TIP4P" else s for s in tables["C"]]
assert "water" in options and "TIP4P" not in options
# =====================================================================
# _validate_solvent: atom-type-count check against a named solvent
# =====================================================================
def test_validate_chloroform():
api._validate_solvent("chloroform", CHCL3_Z) # no raise
def test_validate_two_chloroforms():
api._validate_solvent("chloroform", np.concatenate([CHCL3_Z, CHCL3_Z]))
def test_validate_water():
api._validate_solvent("water", np.array([8, 1, 1]))
def test_validate_benzene():
api._validate_solvent("benzene", np.array([6] * 6 + [1] * 6))
def test_validate_methanol():
api._validate_solvent("methanol", np.array([6, 8, 1, 1, 1, 1]))
def test_validate_order_within_block_independent():
# order within one molecule's block does not matter
api._validate_solvent("water", np.array([1, 8, 1]))
def test_validate_empty_raises():
with pytest.raises(ValueError, match="no solvent atoms"):
api._validate_solvent("chloroform", np.array([], dtype=int))
def test_validate_wrong_count_raises():
# 7 atoms is not a whole number of chloroforms (5 atoms each)
with pytest.raises(ValueError, match="not a whole number"):
api._validate_solvent("chloroform", np.array([6, 6, 1, 1, 17, 17, 17]))
def test_validate_wrong_composition_raises():
# 5 atoms (one chloroform-sized block) but the wrong elements
with pytest.raises(ValueError, match="not whole chloroform"):
api._validate_solvent("chloroform", np.array([6, 6, 6, 6, 6]))
def test_validate_grouped_by_element_raises():
# two chloroforms' worth of atoms (correct TOTALS: 2 C, 2 H, 6 Cl) but grouped by element, so the
# contiguous 5-atom blocks are not whole molecules. Total-count check would miss this; block does not.
grouped = np.array([6, 6, 1, 1, 17, 17, 17, 17, 17, 17])
with pytest.raises(ValueError, match="not whole chloroform"):
api._validate_solvent("chloroform", grouped)
def test_validate_two_chloroforms_contiguous_ok():
# same atoms as above but laid out as two whole contiguous molecules: fine
api._validate_solvent("chloroform", np.concatenate([CHCL3_Z, CHCL3_Z]))
def test_validate_mismatched_solvent_raises():
# a benzene block passed as methanol (both divisible cases would still fail composition)
with pytest.raises(ValueError, match="not whole methanol"):
api._validate_solvent("methanol", np.array([6, 6, 6, 6, 6, 6]))
# =====================================================================
# supported-solvent tables stay in sync with the code (validation)
# =====================================================================
EXPLICIT_SOLVENTS = ["chloroform", "benzene", "methanol", "water"]
PREDICT_SHIFTS_SOLVENTS = [
"tetrahydrofuran", "dichloromethane", "chloroform", "toluene", "benzene", "chlorobenzene",
"acetone", "dimethylsulfoxide", "acetonitrile", "trifluoroethanol", "methanol", "water",
]
def _one_molecule(solvent):
"""Build one molecule of `solvent` as an atomic-number array from the composition table."""
return np.array([element for element, count in api._SOLVENT_COMPOSITION[solvent].items()
for _ in range(count)])
def test_explicit_solvent_set_matches_run_magnet():
# composition table, the model's block-size table, and the documented list all agree
from magnet import run_magnet
assert sorted(api._SOLVENT_COMPOSITION) == sorted(run_magnet.N_ATOMS_PER_SOLVENT)
assert sorted(api._SOLVENT_COMPOSITION) == sorted(EXPLICIT_SOLVENTS)
def test_solvent_composition_sums_to_block_size():
from magnet import run_magnet
for solvent, comp in api._SOLVENT_COMPOSITION.items():
assert sum(comp.values()) == run_magnet.N_ATOMS_PER_SOLVENT[solvent]
@pytest.mark.parametrize("solvent", EXPLICIT_SOLVENTS)
def test_validate_accepts_one_molecule(solvent):
api._validate_solvent(solvent, _one_molecule(solvent)) # must not raise
@pytest.mark.parametrize("solvent", EXPLICIT_SOLVENTS)
def test_validate_accepts_two_contiguous_molecules(solvent):
one = _one_molecule(solvent)
api._validate_solvent(solvent, np.concatenate([one, one])) # must not raise
@pytest.mark.parametrize("solvent", EXPLICIT_SOLVENTS)
def test_validate_rejects_dropping_one_atom(solvent):
# one atom short of a whole molecule -> wrong count
short = _one_molecule(solvent)[:-1]
with pytest.raises(ValueError, match="not a whole number"):
api._validate_solvent(solvent, short)
def test_predict_shifts_supports_the_12_documented_solvents():
tables = scaling.published_scaling_tables()
for nucleus in ("H", "C"):
keys = {"water" if s == "TIP4P" else s for s in tables[nucleus]}
assert keys == set(PREDICT_SHIFTS_SOLVENTS)
def test_explicit_rejects_predict_only_solvent():
# acetone is valid for predict_shifts but not for explicit corrections
atoms = np.concatenate([np.array([6, 1, 1, 1]), CHCL3_Z])
xyz = np.zeros((len(atoms), 3))
with pytest.raises(ValueError, match="solvent must be one of"):
api.explicit_solvent_correction(atoms, xyz, solute_atoms=[0, 1, 2, 3], solvent="acetone",
n_passes=1, symmetrize=False)
# =====================================================================
# _explicit_one: solute selection, reordering, negative/dup/range
# =====================================================================
def _system(): # methane solute (0..3) + one chloroform (4..8)
atoms = np.concatenate([np.array([6, 1, 1, 1]), CHCL3_Z])
xyz = np.concatenate([CH4_XYZ[:4] + 10.0, CHCL3_XYZ]) # solute far, solvent near origin
return atoms, xyz
def test_explicit_one_basic():
atoms, xyz = _system()
solute_an, full_an, full_xyz = api._explicit_one(atoms, xyz, [0, 1, 2, 3], "chloroform")
assert np.array_equal(solute_an, np.array([6, 1, 1, 1]))
assert np.array_equal(full_an, atoms) # already solute-first
def test_explicit_one_reorders_solvent_first_input():
# input has chloroform first, methane second; solute indices point at the methane
atoms = np.concatenate([CHCL3_Z, np.array([6, 1, 1, 1])])
xyz = np.concatenate([CHCL3_XYZ, CH4_XYZ[:4] + 10.0])
solute_an, full_an, full_xyz = api._explicit_one(atoms, xyz, [5, 6, 7, 8], "chloroform")
assert np.array_equal(solute_an, np.array([6, 1, 1, 1]))
# reordered: solute first, then the solvent block
assert np.array_equal(full_an[:4], np.array([6, 1, 1, 1]))
assert np.array_equal(np.sort(full_an[4:]), np.sort(CHCL3_Z))
def test_explicit_one_output_order_follows_solute_atoms():
atoms, xyz = _system()
solute_an, full_an, _ = api._explicit_one(atoms, xyz, [3, 2, 1, 0], "chloroform") # reversed
assert np.array_equal(solute_an, np.array([1, 1, 1, 6])) # matches requested order
assert np.array_equal(full_an[:4], np.array([1, 1, 1, 6]))
def test_explicit_one_negative_indices():
atoms, xyz = _system() # n=9
solute_an, _, _ = api._explicit_one(atoms, xyz, [-9, -8, -7, -6], "chloroform") # -> 0,1,2,3
assert np.array_equal(solute_an, np.array([6, 1, 1, 1]))
def test_explicit_one_wrong_solvent_raises():
atoms, xyz = _system() # solvent is chloroform
with pytest.raises(ValueError, match="not .* benzene"):
api._explicit_one(atoms, xyz, [0, 1, 2, 3], "benzene")
def test_explicit_one_duplicate_raises():
atoms, xyz = _system()
with pytest.raises(ValueError, match="duplicate"):
api._explicit_one(atoms, xyz, [0, 0, 1, 2], "chloroform")
def test_explicit_one_negative_normalizes_to_duplicate_raises():
atoms, xyz = _system() # n=9; index 2 and -7 both -> 2
with pytest.raises(ValueError, match="duplicate"):
api._explicit_one(atoms, xyz, [2, -7], "chloroform")
def test_explicit_one_out_of_range_raises():
atoms, xyz = _system()
with pytest.raises(ValueError, match="out of range"):
api._explicit_one(atoms, xyz, [0, 1, 2, 9], "chloroform")
def test_explicit_one_very_negative_raises():
atoms, xyz = _system()
with pytest.raises(ValueError, match="out of range"):
api._explicit_one(atoms, xyz, [-100], "chloroform")
def test_explicit_one_all_atoms_solute_raises():
atoms, xyz = _system()
with pytest.raises(ValueError, match="no solvent atoms"):
api._explicit_one(atoms, xyz, list(range(len(atoms))), "chloroform")
def test_explicit_one_empty_solute_raises():
atoms, xyz = _system()
with pytest.raises(ValueError, match="empty"):
api._explicit_one(atoms, xyz, [], "chloroform")
# =====================================================================
# explicit_solvent_correction: pure-logic validation before any model call
# =====================================================================
def test_explicit_unknown_solvent_raises():
atoms, xyz = _system()
with pytest.raises(ValueError, match="solvent must be one of"):
api.explicit_solvent_correction(atoms, xyz, solute_atoms=[0, 1, 2, 3], solvent="ethanol",
n_passes=1, symmetrize=False)
def test_explicit_solvent_mismatch_raises():
# declare benzene but the snapshot has chloroform -> count check fails before any model load
atoms, xyz = _system()
with pytest.raises(ValueError, match="not .* benzene"):
api.explicit_solvent_correction(atoms, xyz, solute_atoms=[0, 1, 2, 3], solvent="benzene",
n_passes=1, symmetrize=False)
# =====================================================================
# REAL MODEL TESTS (few, tiny, n_passes=1, symmetrize=False)
# =====================================================================
@needs_model
def test_real_predict_shieldings_single_returns_array():
out = api.predict_shieldings(CH4_Z, CH4_XYZ, n_passes=1, symmetrize=False)
assert isinstance(out, np.ndarray) and out.shape == (5,)
assert np.all(np.isfinite(out)) # H and C all predicted, none NaN
@needs_model
def test_real_predict_shieldings_list_returns_list():
out = api.predict_shieldings([CH4_Z], [CH4_XYZ], n_passes=1, symmetrize=False)
assert isinstance(out, list) and len(out) == 1 and out[0].shape == (5,)
@needs_model
def test_real_predict_shieldings_batch_shapes():
out = api.predict_shieldings([CH4_Z, CH2O_Z], [CH4_XYZ, CH2O_XYZ],
n_passes=1, symmetrize=False)
assert isinstance(out, list) and out[0].shape == (5,) and out[1].shape == (4,)
@needs_model
def test_real_unsupported_element_rejected():
# PH3-like: phosphorus (15) is out of vocab
Z = np.array([15, 1, 1, 1])
X = np.array([[0., 0., 0.], [1., 0., 0.], [0., 1., 0.], [0., 0., 1.]])
with pytest.raises(ValueError, match="unsupported atomic numbers"):
api.predict_shieldings(Z, X, n_passes=1, symmetrize=False)
@needs_model
def test_real_predict_shifts_components_reproduce():
d = api.predict_shifts(CH2O_Z, CH2O_XYZ, solvent="chloroform", n_passes=1,
symmetrize=False, return_components=True)
assert set(d) == {"shifts", "zero_shielding", "pcm_correction", "coefficients"}
shifts, sigma, delta = d["shifts"], d["zero_shielding"], d["pcm_correction"]
coef = d["coefficients"]
# coefficients must match chloroform rows
tbl = scaling.published_scaling_tables()
assert coef["H"] == tbl["H"]["chloroform"] and coef["C"] == tbl["C"]["chloroform"]
# reproduce shift = intercept + stationary*zero + pcm*pcm for H (idx 2,3) and C (idx 0)
for idx, nuc in [(0, "C"), (2, "H"), (3, "H")]:
expect = coef[nuc]["intercept"] + coef[nuc]["stationary"] * sigma[idx] \
+ coef[nuc]["pcm"] * delta[idx]
assert np.isclose(shifts[idx], expect, atol=1e-6), (idx, nuc, shifts[idx], expect)
# O atom (idx 1) must be NaN, H/C finite
assert np.isnan(shifts[1])
assert np.all(np.isfinite([shifts[0], shifts[2], shifts[3]]))
@needs_model
def test_real_predict_shifts_acetone_matches_readme():
# the README "Your First Prediction" worked example: acetone in chloroform on an AIMNet2
# geometry, default n_passes/symmetrize. Guards the documented numbers against model drift.
Z = np.array([6, 6, 8, 6, 1, 1, 1, 1, 1, 1])
X = np.array([
[1.2913, -0.5947, -0.0016], [0.0029, 0.1931, -0.0010], [-0.0174, 1.3994, -0.0003],
[-1.2743, -0.6189, 0.0004], [-1.0822, -1.6899, 0.0014], [-1.8597, -0.3513, 0.8791],
[-1.8605, -0.3530, -0.8783], [1.3415, -1.2208, 0.8916], [1.3170, -1.2669, -0.8614],
[2.1415, 0.0788, -0.0298],
])
shifts = api.predict_shifts(Z, X, solvent="chloroform")
carbonyl = shifts[[1]].mean()
methyl_C = shifts[[0, 3]].mean()
methyl_H = shifts[[4, 5, 6, 7, 8, 9]].mean()
assert np.isclose(carbonyl, 207.6, atol=1.0), carbonyl
assert np.isclose(methyl_C, 30.7, atol=0.5), methyl_C
assert np.isclose(methyl_H, 2.20, atol=0.15), methyl_H
@needs_model
def test_real_predict_shifts_water_maps_to_tip4p():
d = api.predict_shifts(CH4_Z, CH4_XYZ, solvent="water", n_passes=1,
symmetrize=False, return_components=True)
tbl = scaling.published_scaling_tables()
assert d["coefficients"]["H"] == tbl["H"]["TIP4P"]
assert d["coefficients"]["C"] == tbl["C"]["TIP4P"]
@needs_model
def test_real_single_vs_batch_same_shapes():
single = api.predict_shieldings(CH4_Z, CH4_XYZ, n_passes=1, symmetrize=False)
batch = api.predict_shieldings([CH4_Z, CH4_Z], [CH4_XYZ, CH4_XYZ],
n_passes=1, symmetrize=False)
assert single.shape == batch[0].shape == batch[1].shape
# values are close (stochastic frames -> not identical); documents non-determinism
assert np.allclose(single, batch[0], atol=0.5)
@needs_model
@pytest.mark.xfail(reason="known: zero-edge graph (lone atom) gives a cryptic torch error, not a "
"clean ValueError; a lone atom is not a valid NMR input",
raises=RuntimeError)
def test_real_single_atom_solute():
# isolated single atom: no neighbors -> empty edge tensor -> deep torch.min() crash.
api.predict_shieldings(np.array([6]), np.array([[0., 0., 0.]]),
n_passes=1, symmetrize=False)
@needs_model
@pytest.mark.xfail(reason="known: atoms all beyond the model cutoff -> zero edges -> torch "
"RuntimeError, not a clean ValueError",
raises=RuntimeError)
def test_real_atoms_beyond_cutoff_crash():
# two atoms 50 A apart: no edges within cutoff -> same zero-edge crash as one atom.
api.predict_shieldings(np.array([6, 1]), np.array([[0., 0., 0.], [0., 0., 50.]]),
n_passes=1, symmetrize=False)
# ---- degenerate-geometry guard in the local-frame construction (no weights needed) ----
def test_edge_rot_mat_rejects_coincident_atoms():
# an edge between two atoms at the same position has ~zero length; the old code printed a
# warning and then divided by ~zero, silently producing NaN frames that poison the whole
# forward pass. It must raise instead.
edge_vecs = torch.tensor([[1.0, 0.0, 0.0], [0.0, 0.0, 0.0]])
with pytest.raises(ValueError, match="overlapping atoms"):
init_edge_rot_mat(edge_vecs)
def test_edge_rot_mat_accepts_separated_atoms():
edge_vecs = torch.tensor([[1.0, 0.0, 0.0], [0.0, 1.5, 0.0], [0.0, 0.0, 2.0]])
rot = init_edge_rot_mat(edge_vecs)
assert rot.shape == (3, 3, 3)
# each frame is a proper rotation: R @ R^T == I
identity = torch.eye(3).expand(3, 3, 3)
assert torch.allclose(torch.bmm(rot, rot.transpose(1, 2)), identity, atol=1e-4)
|