From bde818fbf17de0276652eac6b3f6670caefee400 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 7 Sep 2026 16:46:19 +0200 Subject: [PATCH 1/5] test: separate fixtures and share read-only functional setup --- docs/changelog.rst | 1 + tests/README.md | 5 +- tests/conftest.py | 530 ++----------------- tests/fixtures/README.md | 40 ++ tests/fixtures/__init__.py | 5 + tests/fixtures/devices.py | 119 +++++ tests/fixtures/files.py | 98 ++++ tests/fixtures/functional.py | 42 ++ tests/fixtures/numerical.py | 195 +++++++ tests/fixtures/objects.py | 163 ++++++ tests/fixtures/paths.py | 68 +++ tests/fixtures/precision.py | 51 ++ tests/functional/conftest.py | 11 +- tests/functional/test_io_functional.py | 181 +++---- tests/functional/test_model_ft_functional.py | 159 +++--- tests/integration/conftest.py | 9 +- tests/unit/conftest.py | 166 +----- tests/unit/structure_factor/conftest.py | 35 +- 18 files changed, 1004 insertions(+), 874 deletions(-) create mode 100644 tests/fixtures/README.md create mode 100644 tests/fixtures/__init__.py create mode 100644 tests/fixtures/devices.py create mode 100644 tests/fixtures/files.py create mode 100644 tests/fixtures/functional.py create mode 100644 tests/fixtures/numerical.py create mode 100644 tests/fixtures/objects.py create mode 100644 tests/fixtures/paths.py create mode 100644 tests/fixtures/precision.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 3284fb85..b5b00059 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Organized test fixtures into focused modules and reused module-scoped loaded objects for read-only functional checks while retaining fresh objects for mutation and loading tests. - Rigid-body refinement stores its Euler angles pre-multiplied by the chain's radius of gyration, so a unit step in an angle and a unit step in a translation displace atoms comparably. In radians against Angstroms the rotation block of the Hessian carried 190-530x the curvature of the translation block on 1DAW and 3E98 -- the geometric ``Rg**2``, 411 and 442/516 -- putting ``cond(H)`` at 1e3-5e3, which is why six parameters needed ~250 L-BFGS iterations to place. Dividing the scale out in ``forward()`` brings the ratio to 0.4-1.3 and ``cond(H)`` to 3-18. Over ten structures the step then converges rather than exhausting its iteration budget, on about half the gradient evaluations, with R-free no worse anywhere. ``RigidXYZTensor.rotation_radians`` returns the physical angle, and setting ``angle_scale`` to ones restores the unscaled parametrization. Not a fix for the one or two negative Hessian eigenvalues at the finer cutoffs -- scaling a saddle leaves it a saddle -- and those counts are unchanged - The rigid-body step no longer co-refines the scaler in the same L-BFGS as the rigid parameters. The body target centres on ``alpha*|F_calc|`` and ``alpha`` absorbs a rescaling of ``F_calc`` exactly, so the scale had a flat direction there; ``SCALE_TARGETS`` already excludes every alpha-centred row from the scale fit for this reason, and 0.6.2 fixed the same thing in the main driver. ``refine_scaler`` (objective ``ls``) owns the scale, between cutoffs - Fixed ``refine_rigid_body`` leaving the caller's reflection data truncated. ``cut_res`` masks in place and returns ``self``, so each cutoff stamped its resolution mask on the caller's own object and the restore had nothing to restore to -- it only looked correct because the default schedule ends at the native limit. With ``--rigid-body-cutoffs 6,4`` on a 2.05 A dataset, 20138 of 23352 reflections stayed masked out for the rest of the run, R-factors included diff --git a/tests/README.md b/tests/README.md index fd283106..ff1a3bc5 100644 --- a/tests/README.md +++ b/tests/README.md @@ -6,7 +6,8 @@ This directory contains the complete test suite for torchref. ``` tests/ -├── conftest.py # Root fixtures (paths, devices, skip decorators) +├── conftest.py # Fixture registration and test-selection hooks +├── fixtures/ # Shared setup, grouped by responsibility (see fixtures/README.md) ├── pytest.ini # Pytest configuration ├── __init__.py ├── files/ # Test data files (CIF, PDB, MTZ) @@ -15,7 +16,7 @@ tests/ │ ├── mtz/ # Reflection MTZ files │ └── cif_sf/ # Structure factor CIF files ├── unit/ # Unit tests (fast, no I/O) -│ ├── conftest.py # Unit test fixtures (mock data) +│ ├── conftest.py # Imports scoped numerical fixtures │ ├── math_functions/ # Math module tests │ ├── model/ # Model module tests │ ├── refinement/ # Refinement module tests diff --git a/tests/conftest.py b/tests/conftest.py index b8b7178e..833cd222 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,35 +1,26 @@ -""" -Root pytest configuration and shared fixtures for torchref tests. +"""Register shared fixture plugins and gate tests on host capabilities.""" -This module provides fixtures that are automatically available to all test files. -""" import importlib.util import shutil import warnings import pytest -import torchref -import torch -import numpy as np -from pathlib import Path +# Apply process-wide settings before the device fixtures import torch. +import torchref # noqa: F401 + +pytest_plugins = ( + "tests.fixtures.paths", + "tests.fixtures.files", + "tests.fixtures.devices", + "tests.fixtures.precision", + "tests.fixtures.objects", +) -# Optional Amber/ensemble stack: OpenMM (pip ``[amber]`` extra) and AmberTools -# (antechamber/tleap — conda-only, detected on PATH). Tests that need them are -# tagged ``@pytest.mark.openmm`` (OpenMM only) or ``@pytest.mark.amber`` (OpenMM -# + AmberTools) and auto-skipped below when the stack is absent. _HAS_OPENMM = importlib.util.find_spec("openmm") is not None _HAS_AMBERTOOLS = bool(shutil.which("antechamber") and shutil.which("tleap")) -def _cuda_available() -> bool: - return torch.cuda.is_available() - - -def _mps_available() -> bool: - return hasattr(torch.backends, "mps") and torch.backends.mps.is_available() - - def pytest_addoption(parser): """Add custom command line options.""" parser.addoption( @@ -58,24 +49,34 @@ def pytest_addoption(parser): help="Deprecated no-op: accelerator tests now run automatically.", ) parser.addoption( - "--run-slow", - action="store_true", - default=False, - help="Run slow tests" + "--run-slow", action="store_true", default=False, help="Run slow tests" ) def pytest_configure(config): """Configure pytest markers.""" config.addinivalue_line("markers", "unit: Unit tests (fast, no I/O)") - config.addinivalue_line("markers", "integration: Integration tests (slower, real I/O)") - config.addinivalue_line("markers", "gpu: Needs any accelerator (CUDA or MPS); auto-skipped if none") - config.addinivalue_line("markers", "cuda: Needs CUDA specifically (e.g. Triton); auto-skipped if absent") - config.addinivalue_line("markers", "mps: Needs MPS specifically (Metal kernels); auto-skipped if absent") + config.addinivalue_line( + "markers", "integration: Integration tests (slower, real I/O)" + ) + config.addinivalue_line( + "markers", "gpu: Needs any accelerator (CUDA or MPS); auto-skipped if none" + ) + config.addinivalue_line( + "markers", "cuda: Needs CUDA specifically (e.g. Triton); auto-skipped if absent" + ) + config.addinivalue_line( + "markers", "mps: Needs MPS specifically (Metal kernels); auto-skipped if absent" + ) config.addinivalue_line("markers", "cuda_only: Deprecated alias for 'cuda'") config.addinivalue_line("markers", "slow: Slow tests (skipped by default)") - config.addinivalue_line("markers", "openmm: Needs OpenMM (the [amber] extra); skipped if absent") - config.addinivalue_line("markers", "amber: Needs OpenMM + AmberTools (antechamber/tleap); skipped if absent") + config.addinivalue_line( + "markers", "openmm: Needs OpenMM (the [amber] extra); skipped if absent" + ) + config.addinivalue_line( + "markers", + "amber: Needs OpenMM + AmberTools (antechamber/tleap); skipped if absent", + ) if config.getoption("--run-gpu"): # UserWarning, not DeprecationWarning: pytest.ini filters the latter, @@ -109,6 +110,8 @@ def pytest_collection_modifyitems(config, items): mask a forgotten marker, and turns "this host cannot run it" into a silent pass instead of the visible skip or the real error. """ + from tests.fixtures.devices import _cuda_available, _mps_available + has_cuda = _cuda_available() has_mps = _mps_available() @@ -138,7 +141,9 @@ def pytest_collection_modifyitems(config, items): ) skip_slow = pytest.mark.skip(reason="Need --run-slow option to run") - skip_openmm = pytest.mark.skip(reason="OpenMM not installed (pip install '.[amber]')") + skip_openmm = pytest.mark.skip( + reason="OpenMM not installed (pip install '.[amber]')" + ) skip_amber = pytest.mark.skip( reason="AmberTools (antechamber/tleap) not on PATH (conda install ambertools)" ) @@ -174,466 +179,3 @@ def pytest_collection_modifyitems(config, items): item.add_marker(skip_amber) elif "openmm" in item.keywords and not _HAS_OPENMM: item.add_marker(skip_openmm) - - -# ============================================================================= -# Path Fixtures -# ============================================================================= - -@pytest.fixture(scope="session") -def tests_root() -> Path: - """Root of the tests directory.""" - return Path(__file__).parent - - -@pytest.fixture(scope="session") -def project_root() -> Path: - """Root of the project.""" - return Path(__file__).parent.parent - - -@pytest.fixture(scope="session") -def test_files_dir(tests_root) -> Path: - """Path to test files directory.""" - return tests_root / "files" - - -@pytest.fixture(scope="session") -def cif_dir(test_files_dir) -> Path: - """Path to CIF model files.""" - return test_files_dir / "cif" - - -@pytest.fixture(scope="session") -def cif_sf_dir(test_files_dir) -> Path: - """Path to CIF structure factor files.""" - return test_files_dir / "cif_sf" - - -@pytest.fixture(scope="session") -def mtz_dir(test_files_dir) -> Path: - """Path to MTZ reflection files.""" - return test_files_dir / "mtz" - - -@pytest.fixture(scope="session") -def pdb_dir(test_files_dir) -> Path: - """Path to PDB model files.""" - return test_files_dir / "pdb" - - -@pytest.fixture(scope="session") -def external_monomer_library(project_root) -> Path: - """Path to external monomer library.""" - return project_root / "external_monomer_library" - - -# ============================================================================= -# Device Fixtures -# ============================================================================= - -@pytest.fixture(scope="session") -def cpu_device() -> torch.device: - """CPU torch device.""" - return torch.device("cpu") - - -@pytest.fixture(scope="session") -def gpu_device() -> torch.device: - """GPU torch device (only use with @pytest.mark.gpu). - - Prefers CUDA, falls back to MPS; skips if neither is available. Prefer the - backend-specific ``cuda_device`` / ``mps_device`` below when a test needs one - particular backend -- this fixture's preference order means a - ``cuda``-marked test asking for it on a dual-backend host could be handed - MPS, which is why the MPS tests used to carry a ``type != 'mps'`` skip to - undo it. - """ - accel = _accelerator() - if accel is None: - pytest.skip("No accelerator (CUDA or MPS) on this host") - return accel - - -@pytest.fixture(scope="session") -def cuda_device() -> torch.device: - """Canonical CUDA device for ``cuda``-marked tests. - - Deliberately unguarded. What runs is decided by the ``cuda`` marker in - :func:`pytest_collection_modifyitems` and nowhere else, so this fixture does - not re-check availability: on a host without CUDA the test is *meant* to - error with the real backend error rather than be quietly skipped here. - """ - return torch.device("cuda", 0) - - -@pytest.fixture(scope="session") -def mps_device() -> torch.device: - """Canonical MPS device for ``mps``-marked tests. - - Unguarded for the same reason as :func:`cuda_device` -- the ``mps`` marker - owns the decision. - """ - return torch.device("mps", 0) - - -def _accelerator() -> "torch.device | None": - """The canonical accelerator this host can actually use, or ``None``. - - Indices are filled in (``cuda:0`` / ``mps:0``) so the value compares equal - to a device read back off a real tensor -- ``torch.device('mps')`` and - ``torch.device('mps:0')`` are *not* equal even though they name the same - physical device. - """ - if _cuda_available(): - return torch.device("cuda", torch.cuda.current_device()) - if _mps_available(): - return torch.device("mps", 0) - return None - - -# Built at import time so the ``gpu`` mark is attached during *collection*. -# Adding it later (e.g. via ``request.node.add_marker`` inside the fixture) is -# too late for ``pytest_collection_modifyitems`` to gate on. -_DEVICE_PARAMS = [pytest.param(torch.device("cpu"), id="cpu")] -_ACCELERATOR = _accelerator() -if _ACCELERATOR is not None: - _DEVICE_PARAMS.append( - pytest.param( - _ACCELERATOR, - id=_ACCELERATOR.type, - # Backend-specific mark, so a CUDA-less host skips the cuda leg and - # a non-Mac skips the mps leg, each with an accurate reason. - marks=getattr(pytest.mark, _ACCELERATOR.type), - ) - ) - - -@pytest.fixture(params=_DEVICE_PARAMS) -def any_device(request) -> torch.device: - """Every device this host can actually use, one test run per device. - - The CPU leg always runs. The accelerator leg is ``gpu``-marked, so a plain - ``pytest`` run skips it and ``pytest --run-gpu`` picks up CUDA on a CUDA - box or MPS on a Mac. On a CPU-only host the accelerator parameter does not - exist at all, so there is no skip noise. - """ - return request.param - - -@pytest.fixture(scope="session") -def _device_model_cache() -> dict: - """``{device_str: ModelFT}`` built at most once per device, per session.""" - return {} - - -@pytest.fixture -def device_model_bundle(_device_model_cache, pdb_dir, any_device): - """A loaded model on ``any_device``, for target conformance tests. - - The existing ``loaded_model`` / ``model_and_data`` fixtures are - function-scoped and construct on the process default, so a - device-parametrized sweep over them would reload the structure once per - test per device. This caches one model per device instead. - - Shared mutable state: callers must treat the bundle as read-only. A test - that moves a *target* will drag the borrowed model with it, poisoning every - later test on that device -- see ``test_target_device_round_trip``, which - deliberately builds its own. - """ - key = str(any_device) - if key not in _device_model_cache: - pdb = pdb_dir / "1DAW.pdb" - if not pdb.exists(): - pytest.skip("1DAW.pdb fixture not present") - from torchref.model import ModelFT - - _device_model_cache[key] = ModelFT(device=any_device, verbose=0).load_pdb( - str(pdb) - ) - return {"model": _device_model_cache[key]} - - -@pytest.fixture -def device(request) -> torch.device: - """Default test device. - - Uses the package-wide auto-detected default (``torchref.device.current``) - so tests run on whichever device the user's machine resolved to at - import time: cuda -> mps -> cpu. Tests marked ``@pytest.mark.cuda_only`` - are skipped when CUDA is not available. - """ - from torchref.config import get_default_device - - markers = {m.name for m in request.node.iter_markers()} - if "cuda_only" in markers and not torch.cuda.is_available(): - pytest.skip("Test requires CUDA") - if "gpu" in markers and not (_cuda_available() or _mps_available()): - pytest.skip("No GPU (CUDA or MPS) available") - return get_default_device() - - -# ============================================================================= -# Numerical Fixtures -# ============================================================================= - -@pytest.fixture -def rtol() -> float: - """Relative tolerance for floating point comparisons.""" - return 1e-5 - - -@pytest.fixture -def atol() -> float: - """Absolute tolerance for floating point comparisons.""" - return 1e-8 - - -# ============================================================================= -# Sample File Fixtures -# ============================================================================= - -@pytest.fixture(scope="session") -def sample_cif_file(cif_dir): - """Return a sample CIF file for testing.""" - cif_file = cif_dir / "1DAW.cif" - if cif_file.exists(): - return cif_file - # Try any available CIF file - cif_files = list(cif_dir.glob("*.cif")) - if cif_files: - return cif_files[0] - pytest.skip("No CIF files found in test data") - - -@pytest.fixture(scope="session") -def sample_mtz_file(mtz_dir): - """Return a sample MTZ file for testing.""" - mtz_file = mtz_dir / "1DAW.mtz" - if mtz_file.exists(): - return mtz_file - # Try any available MTZ file - mtz_files = list(mtz_dir.glob("*.mtz")) - if mtz_files: - return mtz_files[0] - pytest.skip("No MTZ files found in test data") - - -@pytest.fixture(scope="session") -def sample_pdb_file(pdb_dir): - """Return a sample PDB file for testing.""" - pdb_files = sorted(pdb_dir.glob("*.pdb")) - if not pdb_files: - pytest.skip("No PDB files found in test data directory") - return pdb_files[0] - - -@pytest.fixture(scope="session") -def sample_structure_factor_cif(cif_sf_dir): - """Return a sample structure factor CIF file.""" - sf_files = sorted(cif_sf_dir.glob("*.cif")) - if not sf_files: - pytest.skip("No structure factor CIF files found") - return sf_files[0] - - -@pytest.fixture(scope="session") -def sample_structure_pair(cif_dir, mtz_dir): - """Return a matching pair of CIF model and MTZ reflections.""" - # Try to find matching files - pdb_id = "1DAW" - cif_file = cif_dir / f"{pdb_id}.cif" - mtz_file = mtz_dir / f"{pdb_id}.mtz" - - if cif_file.exists() and mtz_file.exists(): - return {"model": cif_file, "reflections": mtz_file} - - # Try to find any matching pair - cif_files = {f.stem: f for f in cif_dir.glob("*.cif")} - mtz_files = {f.stem: f for f in mtz_dir.glob("*.mtz")} - - common_ids = set(cif_files.keys()) & set(mtz_files.keys()) - if common_ids: - pdb_id = sorted(common_ids)[0] - return {"model": cif_files[pdb_id], "reflections": mtz_files[pdb_id]} - - pytest.skip("No matching CIF/MTZ pairs found in test data") - - -@pytest.fixture(scope="session") -def all_structure_pairs(cif_dir, mtz_dir): - """Return all matching pairs of CIF models and MTZ reflections.""" - cif_files = {f.stem: f for f in cif_dir.glob("*.cif")} - mtz_files = {f.stem: f for f in mtz_dir.glob("*.mtz")} - - common_ids = set(cif_files.keys()) & set(mtz_files.keys()) - - if not common_ids: - pytest.skip("No matching CIF/MTZ pairs found in test data") - - return [ - {"pdb_id": pdb_id, "model": cif_files[pdb_id], "reflections": mtz_files[pdb_id]} - for pdb_id in sorted(common_ids) - ] - - -@pytest.fixture(scope="session") -def all_cif_files(cif_dir): - """Return all available CIF test structure files.""" - cif_files = sorted(cif_dir.glob("*.cif")) - if not cif_files: - pytest.skip("No CIF files found in test data directory") - return cif_files - - -@pytest.fixture(scope="session") -def all_test_structures(all_structure_pairs): - """Return all loaded model/data pairs for comprehensive testing.""" - from torchref.model.model import Model - from torchref.io import ReflectionData - - structures = [] - for pair in all_structure_pairs: - try: - model = Model() - model.load_cif(str(pair["model"])) - - data = ReflectionData() - data.load_mtz(str(pair["reflections"])) - - structures.append({ - "pdb_id": pair["pdb_id"], - "model": model, - "data": data, - "model_path": pair["model"], - "data_path": pair["reflections"] - }) - except Exception: - # Skip structures that fail to load - continue - - if not structures: - pytest.skip("No structures could be loaded") - - return structures - - -@pytest.fixture(scope="session") -def monomer_library_path(project_root): - """Get path to the monomer library as a string. - - Returns - ------- - str - Absolute path to the external_monomer_library directory. - """ - lib_path = project_root / "external_monomer_library" - if not lib_path.exists(): - pytest.skip("Monomer library not found") - return str(lib_path) - - -# ============================================================================= -# Real Object Fixtures -# ============================================================================= - -@pytest.fixture -def loaded_model(sample_cif_file): - """Fixture providing a fully loaded Model from a real CIF file.""" - from torchref.model.model import Model - - model = Model() - model.load_cif(str(sample_cif_file)) - return model - - -@pytest.fixture -def loaded_reflection_data(sample_mtz_file): - """Fixture providing fully loaded ReflectionData from a real MTZ file.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - return data - - -@pytest.fixture -def model_and_data(sample_structure_pair): - """Fixture providing matching model and reflection data.""" - from torchref.model.model import Model - from torchref.io import ReflectionData - - model = Model() - model.load_cif(str(sample_structure_pair["model"])) - - data = ReflectionData() - data.load_mtz(str(sample_structure_pair["reflections"])) - - return {"model": model, "data": data} - - -@pytest.fixture -def model_with_symmetry(loaded_model): - """Fixture providing model with initialized symmetry.""" - from torchref.symmetry import SpaceGroup - - sg = SpaceGroup(loaded_model.spacegroup) - return {"model": loaded_model, "symmetry": sg} - - -@pytest.fixture -def initialized_scaler(model_and_data): - """Fixture providing initialized Scaler with model and data.""" - from torchref.scaling.scaler import Scaler - - model = model_and_data["model"] - data = model_and_data["data"] - - scaler = Scaler(model=model, data=data, nbins=10, verbose=0) - return scaler - - -@pytest.fixture -def model_with_restraints(loaded_model): - """Fixture providing model with built restraints.""" - from torchref.topology.restraints import Restraints - - restraints = Restraints( - pdb=loaded_model.pdb, - xyz_fn=loaded_model.xyz, - vdw_radii_fn=loaded_model.get_vdw_radii, - verbose=0 - ) - restraints.build_restraints() - return {"model": loaded_model, "restraints": restraints} - -@pytest.fixture -def double_cpu(): - """float64/complex128 on CPU for the duration of a test; restore afterwards. - - Required rather than cosmetic for anything touching eager structure factors: - ``iso_structure_factor_torched`` casts ``hkl`` to the *global* ``dtypes.float`` - (``torchref/base/direct_summation/isotropic.py:121``), so under the default float32 - config a float64 leaf produces a dtype-mismatched matmul. - - Promoted here from three byte-similar copies in ``tests/unit/test_kernel_fixes.py``, - ``tests/unit/test_gradient_correctness.py`` and - ``tests/integration/test_dtype_config_float64.py``. This version also restores - ``sigma_cutoff_ed``, which none of those did -- so a test that changed the cutoff - leaked it into everything that ran afterwards. - """ - import torchref - from torchref.config import device as _device, dtypes as _dtypes - - f0, c0, d0 = _dtypes.float, _dtypes.complex, _device.current - s0 = torchref.sigma_cutoff_ed.value - _dtypes.float = torch.float64 - _dtypes.complex = torch.complex128 - _device.current = torch.device("cpu") - try: - yield - finally: - _dtypes.float = f0 - _dtypes.complex = c0 - _device.current = d0 - torchref.sigma_cutoff_ed.value = s0 diff --git a/tests/fixtures/README.md b/tests/fixtures/README.md new file mode 100644 index 00000000..da39ba78 --- /dev/null +++ b/tests/fixtures/README.md @@ -0,0 +1,40 @@ +# Fixture ownership + +The root `tests/conftest.py` owns pytest options, markers, capability gating, +and the `pytest_plugins` registry. Put reusable setup in the modules below. +Keep a fixture in its test module when only that module needs it. + +| Module | Responsibility | Visibility / lifetime | +|---|---|---| +| `paths.py` | Repository, bundled-data and optional library paths | All tests; session | +| `files.py` | Sample-file selection and matching structure pairs | All tests; session; no model loading | +| `devices.py` | Configured device, explicit backends, device parametrization | All tests; existing per-fixture scopes | +| `precision.py` | Comparison tolerances and CPU-double reference context | All tests; reference fixture restores state after each test | +| `objects.py` | Mutable models, data, scalers and restraints | All tests; fresh per test except explicitly shared bundles | +| `numerical.py` | Synthetic tensors and factories | Imported only by `unit/conftest.py`; function | +| `functional.py` | `shared_model`, `shared_model_ft`, `shared_reflection_data` | Imported only by `functional/conftest.py`; module | + +Existing fixture names remain available without imports in tests. Import reusable +helpers from their defining module, never from the root `conftest.py`. Subtree +conftests import fixture functions explicitly; register shared plugins only at +the root so pytest also works when invoked from a subdirectory. + +Use `shared_*` only for read-only checks. They capture the package configuration +at module setup and may populate derived caches. Do not move them, change their +parameters, tables, masks or grids, backpropagate through them, or use them in +tests that switch global configuration. A target or scaler can mutate a model it +borrows, so a shared model must not be passed to such an operation. + +Tests that verify loading must invoke the loader themselves. Tests of mutation, +device movement, or empty caches use fresh objects. `loaded_model`, +`loaded_model_ft`, `loaded_reflection_data`, and their composed fixtures in `objects.py` provide +fresh mutable objects per test. The explicitly shared session bundles in that +module retain their documented ownership contracts. + +Use `cpu_double_precision()` to scope an explicit numerical reference, or request +`double_cpu` for a single test. The structure-factor package uses the same context +at package scope; both usages restore dtype, device, and density cutoff on exit. + +This separation preserves the existing numerical-factory allocation policy and +test-selection policy. Those policies are independent of fixture registration and +scope, and can be revised in their respective modules. diff --git a/tests/fixtures/__init__.py b/tests/fixtures/__init__.py new file mode 100644 index 00000000..660fb98a --- /dev/null +++ b/tests/fixtures/__init__.py @@ -0,0 +1,5 @@ +"""Provide pytest fixtures grouped by responsibility. + +Root conftest registers shared plugins. Unit and functional conftests explicitly +import their scoped fixtures; this package deliberately re-exports none. +""" diff --git a/tests/fixtures/devices.py b/tests/fixtures/devices.py new file mode 100644 index 00000000..579daeff --- /dev/null +++ b/tests/fixtures/devices.py @@ -0,0 +1,119 @@ +"""Provide explicit backend fixtures and the configured default device. + +Capability probes are shared with collection hooks and structure-factor cases. +Device parametrization is constructed at import so collection can see its marks. +""" + +import pytest +import torch + + +def _cuda_available() -> bool: + return torch.cuda.is_available() + + +def _mps_available() -> bool: + return hasattr(torch.backends, "mps") and torch.backends.mps.is_available() + + +def _accelerator() -> "torch.device | None": + """The canonical accelerator this host can actually use, or ``None``. + + Indices are filled in (``cuda:0`` / ``mps:0``) so the value compares equal + to a device read back off a real tensor -- ``torch.device('mps')`` and + ``torch.device('mps:0')`` are *not* equal even though they name the same + physical device. + """ + if _cuda_available(): + return torch.device("cuda", torch.cuda.current_device()) + if _mps_available(): + return torch.device("mps", 0) + return None + + +@pytest.fixture(scope="session") +def cpu_device() -> torch.device: + """CPU torch device.""" + return torch.device("cpu") + + +@pytest.fixture(scope="session") +def gpu_device() -> torch.device: + """Select CUDA, then MPS, for tests marked ``gpu``. + + Skip if neither backend is available. Use ``cuda_device`` or ``mps_device`` + when the test exercises a backend-specific contract. + """ + accel = _accelerator() + if accel is None: + pytest.skip("No accelerator (CUDA or MPS) on this host") + return accel + + +@pytest.fixture(scope="session") +def cuda_device() -> torch.device: + """Canonical CUDA device for ``cuda``-marked tests. + + Deliberately unguarded. What runs is decided by the ``cuda`` marker in + :func:`pytest_collection_modifyitems` and nowhere else, so this fixture does + not re-check availability: on a host without CUDA the test is *meant* to + error with the real backend error rather than be quietly skipped here. + """ + return torch.device("cuda", 0) + + +@pytest.fixture(scope="session") +def mps_device() -> torch.device: + """Canonical MPS device for ``mps``-marked tests. + + Unguarded for the same reason as :func:`cuda_device` -- the ``mps`` marker + owns the decision. + """ + return torch.device("mps", 0) + + +# Built at import time so the ``gpu`` mark is attached during *collection*. +# Adding it later (e.g. via ``request.node.add_marker`` inside the fixture) is +# too late for ``pytest_collection_modifyitems`` to gate on. +_DEVICE_PARAMS = [pytest.param(torch.device("cpu"), id="cpu")] +_ACCELERATOR = _accelerator() +if _ACCELERATOR is not None: + _DEVICE_PARAMS.append( + pytest.param( + _ACCELERATOR, + id=_ACCELERATOR.type, + # Backend-specific mark, so a CUDA-less host skips the cuda leg and + # a non-Mac skips the mps leg, each with an accurate reason. + marks=getattr(pytest.mark, _ACCELERATOR.type), + ) + ) + + +@pytest.fixture(params=_DEVICE_PARAMS) +def any_device(request: pytest.FixtureRequest) -> torch.device: + """Every device this host can actually use, one test run per device. + + The CPU leg always runs. An available accelerator runs automatically and + carries its backend-specific marker. No accelerator leg is created on a + CPU-only host. + """ + return request.param + + +@pytest.fixture +def device(request: pytest.FixtureRequest) -> torch.device: + """Default test device. + + Uses the package-wide auto-detected default (``torchref.device.current``) + so tests run on whichever device the user's machine resolved to at + import time: cuda -> mps -> cpu. Tests marked ``@pytest.mark.cuda_only`` + are skipped when CUDA is not available. + """ + from torchref.config import get_default_device + + markers = {m.name for m in request.node.iter_markers()} + if "cuda_only" in markers and not torch.cuda.is_available(): + pytest.skip("Test requires CUDA") + if "gpu" in markers and not (_cuda_available() or _mps_available()): + pytest.skip("No GPU (CUDA or MPS) available") + return get_default_device() diff --git a/tests/fixtures/files.py b/tests/fixtures/files.py new file mode 100644 index 00000000..d248775e --- /dev/null +++ b/tests/fixtures/files.py @@ -0,0 +1,98 @@ +"""Select sample paths and matching model/reflection pairs without loading them.""" + +from pathlib import Path + +import pytest + + +@pytest.fixture(scope="session") +def sample_cif_file(cif_dir: Path) -> Path: + """Return a sample CIF file for testing.""" + cif_file = cif_dir / "1DAW.cif" + if cif_file.exists(): + return cif_file + # Try any available CIF file + cif_files = list(cif_dir.glob("*.cif")) + if cif_files: + return cif_files[0] + pytest.skip("No CIF files found in test data") + + +@pytest.fixture(scope="session") +def sample_mtz_file(mtz_dir: Path) -> Path: + """Return a sample MTZ file for testing.""" + mtz_file = mtz_dir / "1DAW.mtz" + if mtz_file.exists(): + return mtz_file + # Try any available MTZ file + mtz_files = list(mtz_dir.glob("*.mtz")) + if mtz_files: + return mtz_files[0] + pytest.skip("No MTZ files found in test data") + + +@pytest.fixture(scope="session") +def sample_pdb_file(pdb_dir: Path) -> Path: + """Return a sample PDB file for testing.""" + pdb_files = sorted(pdb_dir.glob("*.pdb")) + if not pdb_files: + pytest.skip("No PDB files found in test data directory") + return pdb_files[0] + + +@pytest.fixture(scope="session") +def sample_structure_factor_cif(cif_sf_dir: Path) -> Path: + """Return a sample structure factor CIF file.""" + sf_files = sorted(cif_sf_dir.glob("*.cif")) + if not sf_files: + pytest.skip("No structure factor CIF files found") + return sf_files[0] + + +@pytest.fixture(scope="session") +def sample_structure_pair(cif_dir: Path, mtz_dir: Path) -> dict[str, Path]: + """Return a matching pair of CIF model and MTZ reflections.""" + # Try to find matching files + pdb_id = "1DAW" + cif_file = cif_dir / f"{pdb_id}.cif" + mtz_file = mtz_dir / f"{pdb_id}.mtz" + + if cif_file.exists() and mtz_file.exists(): + return {"model": cif_file, "reflections": mtz_file} + + # Try to find any matching pair + cif_files = {f.stem: f for f in cif_dir.glob("*.cif")} + mtz_files = {f.stem: f for f in mtz_dir.glob("*.mtz")} + + common_ids = set(cif_files.keys()) & set(mtz_files.keys()) + if common_ids: + pdb_id = min(common_ids) + return {"model": cif_files[pdb_id], "reflections": mtz_files[pdb_id]} + + pytest.skip("No matching CIF/MTZ pairs found in test data") + + +@pytest.fixture(scope="session") +def all_structure_pairs(cif_dir: Path, mtz_dir: Path) -> list[dict[str, Path | str]]: + """Return all matching pairs of CIF models and MTZ reflections.""" + cif_files = {f.stem: f for f in cif_dir.glob("*.cif")} + mtz_files = {f.stem: f for f in mtz_dir.glob("*.mtz")} + + common_ids = set(cif_files.keys()) & set(mtz_files.keys()) + + if not common_ids: + pytest.skip("No matching CIF/MTZ pairs found in test data") + + return [ + {"pdb_id": pdb_id, "model": cif_files[pdb_id], "reflections": mtz_files[pdb_id]} + for pdb_id in sorted(common_ids) + ] + + +@pytest.fixture(scope="session") +def all_cif_files(cif_dir: Path) -> list[Path]: + """Return all available CIF test structure files.""" + cif_files = sorted(cif_dir.glob("*.cif")) + if not cif_files: + pytest.skip("No CIF files found in test data directory") + return cif_files diff --git a/tests/fixtures/functional.py b/tests/fixtures/functional.py new file mode 100644 index 00000000..18168242 --- /dev/null +++ b/tests/fixtures/functional.py @@ -0,0 +1,42 @@ +"""Share loaded objects within a functional module for read-only checks. + +These fixtures capture the configured dtype/device at module setup. Callers may +populate derived caches but must not change parameters, tables, grids, masks, +device, or configuration. Tests of loading, mutation, and empty caches construct +fresh objects instead. No loaded objects are shared across test modules. +""" + +from __future__ import annotations + +from pathlib import Path +from typing import TYPE_CHECKING + +import pytest + +if TYPE_CHECKING: + from torchref.io import ReflectionData + from torchref.model import Model, ModelFT + + +@pytest.fixture(scope="module") +def shared_model(sample_cif_file: Path) -> Model: + """Load the sample CIF once per module for read-only atomic-model checks.""" + from torchref.model import Model + + return Model(verbose=0).load_cif(str(sample_cif_file)) + + +@pytest.fixture(scope="module") +def shared_model_ft(sample_cif_file: Path) -> ModelFT: + """Load a read-only Fourier model with a 2 Å resolution limit per module.""" + from torchref.model import ModelFT + + return ModelFT(max_res=2.0, verbose=0).load_cif(str(sample_cif_file)) + + +@pytest.fixture(scope="module") +def shared_reflection_data(sample_mtz_file: Path) -> ReflectionData: + """Load the sample MTZ once per module for read-only reflection checks.""" + from torchref.io import ReflectionData + + return ReflectionData().load_mtz(str(sample_mtz_file)) diff --git a/tests/fixtures/numerical.py b/tests/fixtures/numerical.py new file mode 100644 index 00000000..c0d54fc3 --- /dev/null +++ b/tests/fixtures/numerical.py @@ -0,0 +1,195 @@ +"""Generate small synthetic numerical inputs for unit tests. + +Imported by the unit conftest only. Factories return fresh CPU tensors on every +call, using TorchRef's numeric dtypes. They reset the global NumPy random seed; +``random_seed`` also resets PyTorch's seed. Accelerator coverage requires an +explicit move by the caller under this allocation policy. +""" + +from collections.abc import Callable + +import numpy as np +import pytest +import torch + +from torchref.config import dtypes + + +@pytest.fixture +def random_seed() -> int: + """Set random seed for reproducibility.""" + seed = 42 + np.random.seed(seed) + torch.manual_seed(seed) + return seed + + +@pytest.fixture +def random_coordinates() -> Callable[..., torch.Tensor]: + """Return a factory for Cartesian coordinates of shape (n_atoms, 3) in Å.""" + + def _generate(n_atoms: int = 10, seed: int = 42) -> torch.Tensor: + np.random.seed(seed) + return torch.tensor(np.random.rand(n_atoms, 3) * 10, dtype=dtypes.float) + + return _generate + + +@pytest.fixture +def random_fractional_coordinates() -> Callable[..., torch.Tensor]: + """Return a factory for fractional coordinates (n_atoms, 3) in [0, 1).""" + + def _generate(n_atoms: int = 10, seed: int = 42) -> torch.Tensor: + np.random.seed(seed) + return torch.tensor(np.random.rand(n_atoms, 3), dtype=dtypes.float) + + return _generate + + +@pytest.fixture +def random_adp() -> Callable[..., torch.Tensor]: + """Return a factory for isotropic B-factors (n_atoms,) in [10, 60) Ų.""" + + def _generate(n_atoms: int = 10, seed: int = 42) -> torch.Tensor: + np.random.seed(seed) + return torch.tensor(np.random.rand(n_atoms) * 50 + 10, dtype=dtypes.float) + + return _generate + + +@pytest.fixture +def random_occupancies() -> Callable[..., torch.Tensor]: + """Return a factory for dimensionless occupancies (n_atoms,) in [0.5, 1).""" + + def _generate(n_atoms: int = 10, seed: int = 42) -> torch.Tensor: + np.random.seed(seed) + return torch.tensor(np.random.rand(n_atoms) * 0.5 + 0.5, dtype=dtypes.float) + + return _generate + + +@pytest.fixture +def mock_cell() -> torch.Tensor: + """Return an orthorhombic cell (6,), lengths in Å and angles in degrees.""" + return torch.tensor([50.0, 60.0, 70.0, 90.0, 90.0, 90.0], dtype=dtypes.float) + + +@pytest.fixture +def mock_cell_triclinic() -> torch.Tensor: + """Return a triclinic cell (6,), lengths in Å and angles in degrees.""" + return torch.tensor([40.0, 50.0, 60.0, 70.0, 80.0, 85.0], dtype=dtypes.float) + + +@pytest.fixture +def mock_hkl_indices() -> Callable[..., torch.Tensor]: + """Return a factory for floating HKL triples (n_kept, 3), excluding the origin. + + The output uses ``dtypes.float``; ``n_kept`` can be less than the requested + reflection count when the origin is sampled. + """ + + def _generate( + n_reflections: int = 100, max_index: int = 10, seed: int = 42 + ) -> torch.Tensor: + np.random.seed(seed) + h = np.random.randint(-max_index, max_index + 1, n_reflections) + k = np.random.randint(-max_index, max_index + 1, n_reflections) + l = np.random.randint(-max_index, max_index + 1, n_reflections) + # Exclude (0,0,0) + mask = ~((h == 0) & (k == 0) & (l == 0)) + h, k, l = h[mask], k[mask], l[mask] + return torch.tensor(np.stack([h, k, l], axis=1), dtype=dtypes.float) + + return _generate + + +@pytest.fixture +def mock_structure_factors() -> Callable[..., torch.Tensor]: + """Return a factory for complex structure factors (n_reflections,) in electrons.""" + + def _generate(n_reflections: int = 100, seed: int = 42) -> torch.Tensor: + np.random.seed(seed) + real = np.random.randn(n_reflections) * 100 + imag = np.random.randn(n_reflections) * 100 + return torch.tensor(real + 1j * imag, dtype=dtypes.complex) + + return _generate + + +@pytest.fixture +def mock_F_obs() -> Callable[..., torch.Tensor]: + """Return a factory for observed amplitudes (n_reflections,) in electrons.""" + + def _generate(n_reflections: int = 100, seed: int = 42) -> torch.Tensor: + np.random.seed(seed) + # Positive values with realistic distribution + return torch.tensor( + np.abs(np.random.randn(n_reflections) * 100) + 10, dtype=dtypes.float + ) + + return _generate + + +@pytest.fixture +def mock_F_sigma() -> Callable[..., torch.Tensor]: + """Return a factory for amplitude uncertainties (n_reflections,) in electrons.""" + + def _generate(n_reflections: int = 100, seed: int = 42) -> torch.Tensor: + np.random.seed(seed) + return torch.tensor( + np.abs(np.random.randn(n_reflections) * 5) + 1, dtype=dtypes.float + ) + + return _generate + + +@pytest.fixture +def mock_aniso_u() -> Callable[..., torch.Tensor]: + """Return a factory for Cartesian U tensors (n_atoms, 6) in Ų. + + Components are ordered U11, U22, U33, U12, U13, U23. + """ + + def _generate(n_atoms: int = 10, seed: int = 42) -> torch.Tensor: + np.random.seed(seed) + # Diagonal elements (positive) + u11 = np.random.rand(n_atoms) * 0.05 + 0.02 + u22 = np.random.rand(n_atoms) * 0.05 + 0.02 + u33 = np.random.rand(n_atoms) * 0.05 + 0.02 + # Off-diagonal elements (can be negative, smaller magnitude) + u12 = (np.random.rand(n_atoms) - 0.5) * 0.02 + u13 = (np.random.rand(n_atoms) - 0.5) * 0.02 + u23 = (np.random.rand(n_atoms) - 0.5) * 0.02 + return torch.tensor( + np.stack([u11, u22, u33, u12, u13, u23], axis=1), dtype=dtypes.float + ) + + return _generate + + +@pytest.fixture +def mock_scattering_factors() -> Callable[..., torch.Tensor]: + """Return a factory for scattering factors (n_reflections, n_atoms) in electrons.""" + + def _generate( + n_reflections: int = 100, n_atoms: int = 10, seed: int = 42 + ) -> torch.Tensor: + np.random.seed(seed) + # Decreasing with resolution (approximate) + return torch.tensor( + np.random.rand(n_reflections, n_atoms) * 5 + 1, dtype=dtypes.float + ) + + return _generate + + +@pytest.fixture +def mock_weights() -> Callable[..., torch.Tensor]: + """Return a factory for dimensionless weights (n_atoms, 1) summing to one.""" + + def _generate(n_atoms: int = 10, seed: int = 42) -> torch.Tensor: + np.random.seed(seed) + weights = np.random.rand(n_atoms) + return torch.tensor(weights / weights.sum(), dtype=dtypes.float).reshape(-1, 1) + + return _generate diff --git a/tests/fixtures/objects.py b/tests/fixtures/objects.py new file mode 100644 index 00000000..34fb00c8 --- /dev/null +++ b/tests/fixtures/objects.py @@ -0,0 +1,163 @@ +"""Load fresh mutable models, reflection data, scalers, and restraints. + +Function-scoped fixtures isolate test mutations. The explicitly shared device +bundle caches one model per device and must be treated as read-only by callers. +""" + +from __future__ import annotations + +from pathlib import Path +from typing import TYPE_CHECKING, Any + +import pytest +import torch + +if TYPE_CHECKING: + from torchref.io import ReflectionData + from torchref.model import Model, ModelFT + from torchref.scaling import Scaler + + +@pytest.fixture +def loaded_model(sample_cif_file: Path) -> Model: + """Load a fresh mutable Model from the sample CIF file.""" + from torchref.model.model import Model + + model = Model() + model.load_cif(str(sample_cif_file)) + return model + + +@pytest.fixture +def loaded_model_ft(sample_cif_file: Path) -> ModelFT: + """Load a fresh mutable Fourier model with a 2 Å resolution limit.""" + from torchref.model import ModelFT + + return ModelFT(max_res=2.0, verbose=0).load_cif(str(sample_cif_file)) + + +@pytest.fixture +def loaded_reflection_data(sample_mtz_file: Path) -> ReflectionData: + """Load fresh mutable reflection data from the sample MTZ file.""" + from torchref.io import ReflectionData + + data = ReflectionData() + data.load_mtz(str(sample_mtz_file)) + return data + + +@pytest.fixture +def model_and_data(sample_structure_pair: dict[str, Path]) -> dict[str, Any]: + """Load a fresh matching model and reflection dataset.""" + from torchref.io import ReflectionData + from torchref.model.model import Model + + model = Model() + model.load_cif(str(sample_structure_pair["model"])) + + data = ReflectionData() + data.load_mtz(str(sample_structure_pair["reflections"])) + + return {"model": model, "data": data} + + +@pytest.fixture +def model_with_symmetry(loaded_model: Model) -> dict[str, Any]: + """Pair a fresh model with initialized symmetry.""" + from torchref.symmetry import SpaceGroup + + sg = SpaceGroup(loaded_model.spacegroup) + return {"model": loaded_model, "symmetry": sg} + + +@pytest.fixture +def initialized_scaler(model_and_data: dict[str, Any]) -> Scaler: + """Build a scaler around a fresh matching model and dataset.""" + from torchref.scaling.scaler import Scaler + + model = model_and_data["model"] + data = model_and_data["data"] + + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) + return scaler + + +@pytest.fixture +def model_with_restraints(loaded_model: Model) -> dict[str, Any]: + """Build restraints around a fresh model.""" + from torchref.topology.restraints import Restraints + + restraints = Restraints( + pdb=loaded_model.pdb, + xyz_fn=loaded_model.xyz, + vdw_radii_fn=loaded_model.get_vdw_radii, + verbose=0, + ) + restraints.build_restraints() + return {"model": loaded_model, "restraints": restraints} + + +@pytest.fixture(scope="session") +def all_test_structures( + all_structure_pairs: list[dict[str, Any]], +) -> list[dict[str, Any]]: + """Return all loaded model/data pairs for comprehensive testing.""" + from torchref.io import ReflectionData + from torchref.model.model import Model + + structures = [] + for pair in all_structure_pairs: + try: + model = Model() + model.load_cif(str(pair["model"])) + + data = ReflectionData() + data.load_mtz(str(pair["reflections"])) + + structures.append( + { + "pdb_id": pair["pdb_id"], + "model": model, + "data": data, + "model_path": pair["model"], + "data_path": pair["reflections"], + } + ) + except Exception: + # Skip structures that fail to load + continue + + if not structures: + pytest.skip("No structures could be loaded") + + return structures + + +@pytest.fixture(scope="session") +def _device_model_cache() -> dict: + """``{device_str: ModelFT}`` built at most once per device, per session.""" + return {} + + +@pytest.fixture +def device_model_bundle( + _device_model_cache: dict[str, ModelFT], pdb_dir: Path, any_device: torch.device +) -> dict[str, ModelFT]: + """Borrow a session-shared model on the requested device. + + Notes + ----- + Treat the model as read-only, including when a target borrows it. Moving a + target can move its model too; tests of movement need a fresh model. + """ + key = str(any_device) + if key not in _device_model_cache: + pdb = pdb_dir / "1DAW.pdb" + if not pdb.exists(): + pytest.skip("1DAW.pdb fixture not present") + from torchref.model import ModelFT + + _device_model_cache[key] = ModelFT(device=any_device, verbose=0).load_pdb( + str(pdb) + ) + return {"model": _device_model_cache[key]} diff --git a/tests/fixtures/paths.py b/tests/fixtures/paths.py new file mode 100644 index 00000000..16ac98d4 --- /dev/null +++ b/tests/fixtures/paths.py @@ -0,0 +1,68 @@ +"""Locate bundled test data and optional monomer-library installations.""" + +from pathlib import Path + +import pytest + + +@pytest.fixture(scope="session") +def tests_root() -> Path: + """Return the root of the test tree.""" + return Path(__file__).resolve().parents[1] + + +@pytest.fixture(scope="session") +def project_root() -> Path: + """Return the project root.""" + return Path(__file__).resolve().parents[2] + + +@pytest.fixture(scope="session") +def test_files_dir(tests_root: Path) -> Path: + """Return the bundled test-data directory.""" + return tests_root / "files" + + +@pytest.fixture(scope="session") +def cif_dir(test_files_dir: Path) -> Path: + """Return the model CIF directory.""" + return test_files_dir / "cif" + + +@pytest.fixture(scope="session") +def cif_sf_dir(test_files_dir: Path) -> Path: + """Return the structure-factor CIF directory.""" + return test_files_dir / "cif_sf" + + +@pytest.fixture(scope="session") +def mtz_dir(test_files_dir: Path) -> Path: + """Return the MTZ reflection directory.""" + return test_files_dir / "mtz" + + +@pytest.fixture(scope="session") +def pdb_dir(test_files_dir: Path) -> Path: + """Return the model PDB directory.""" + return test_files_dir / "pdb" + + +@pytest.fixture(scope="session") +def external_monomer_library(project_root: Path) -> Path: + """Return the optional external monomer-library path without checking it.""" + return project_root / "external_monomer_library" + + +@pytest.fixture(scope="session") +def monomer_library_path(project_root: Path) -> str: + """Get path to the monomer library as a string. + + Returns + ------- + str + Absolute path to the external_monomer_library directory. + """ + lib_path = project_root / "external_monomer_library" + if not lib_path.exists(): + pytest.skip("Monomer library not found") + return str(lib_path) diff --git a/tests/fixtures/precision.py b/tests/fixtures/precision.py new file mode 100644 index 00000000..6769c4de --- /dev/null +++ b/tests/fixtures/precision.py @@ -0,0 +1,51 @@ +"""Scope numerical reference configuration and expose comparison tolerances.""" + +from collections.abc import Iterator +from contextlib import contextmanager + +import pytest +import torch + +import torchref +from torchref.config import device, dtypes + + +@contextmanager +def cpu_double_precision() -> Iterator[None]: + """Temporarily select CPU float64/complex128 for numerical references. + + Notes + ----- + Mutate process-wide TorchRef defaults, not PyTorch factory defaults. Restore + float/complex dtype, device, and density cutoff even when the body raises. + Objects allocated inside the context retain their own dtype and device. + """ + original = dtypes.float, dtypes.complex, device.current + cutoff = torchref.sigma_cutoff_ed.value + dtypes.float = torch.float64 + dtypes.complex = torch.complex128 + device.current = torch.device("cpu") + try: + yield + finally: + dtypes.float, dtypes.complex, device.current = original + torchref.sigma_cutoff_ed.value = cutoff + + +@pytest.fixture +def double_cpu() -> Iterator[None]: + """Use CPU double precision for one test and restore configuration afterward.""" + with cpu_double_precision(): + yield + + +@pytest.fixture +def rtol() -> float: + """Relative tolerance for floating point comparisons.""" + return 1e-5 + + +@pytest.fixture +def atol() -> float: + """Absolute tolerance for floating point comparisons.""" + return 1e-8 diff --git a/tests/functional/conftest.py b/tests/functional/conftest.py index 58e2352e..e7a5b0ed 100644 --- a/tests/functional/conftest.py +++ b/tests/functional/conftest.py @@ -1,6 +1,7 @@ -""" -Functional test fixtures. +"""Expose module-shared read-only fixtures to functional tests.""" -All shared fixtures (sample files, loaded models, scalers, restraints, etc.) -are defined in the root tests/conftest.py and are automatically available here. -""" +from tests.fixtures.functional import ( # noqa: F401 + shared_model, + shared_model_ft, + shared_reflection_data, +) diff --git a/tests/functional/test_io_functional.py b/tests/functional/test_io_functional.py index 9cdef611..8779ea39 100644 --- a/tests/functional/test_io_functional.py +++ b/tests/functional/test_io_functional.py @@ -6,7 +6,6 @@ import pytest import torch -import numpy as np class TestCIFReadingFunctional: @@ -16,49 +15,45 @@ class TestCIFReadingFunctional: def test_load_multiple_cif_files(self, cif_dir): """Test loading multiple CIF files successfully.""" from torchref.model.model import Model - + cif_files = list(cif_dir.glob("*.cif")) assert len(cif_files) > 0, "No CIF files found in test directory" - + for cif_file in cif_files: model = Model() model.load_cif(str(cif_file)) - + # Each file should load with atoms n_atoms = model.xyz().shape[0] assert n_atoms > 0, f"No atoms loaded from {cif_file}" - + # Should have cell parameters assert model.cell is not None assert len(model.cell) == 6 @pytest.mark.integration - def test_cif_atom_properties(self, sample_cif_file): + def test_cif_atom_properties(self, shared_model): """Test that atom properties are correctly loaded from CIF.""" - from torchref.model.model import Model - - model = Model() - model.load_cif(str(sample_cif_file)) - + model = shared_model + pdb = model.pdb - + # Check required columns exist - required_cols = ['x', 'y', 'z', 'element', 'resname', 'chainid', 'resseq'] + required_cols = ["x", "y", "z", "element", "resname", "chainid", "resseq"] for col in required_cols: - assert col in pdb.columns or col.upper() in pdb.columns, f"Missing column: {col}" + assert col in pdb.columns or col.upper() in pdb.columns, ( + f"Missing column: {col}" + ) @pytest.mark.integration - def test_cif_element_types(self, sample_cif_file): + def test_cif_element_types(self, shared_model): """Test that element types are properly assigned.""" - from torchref.model.model import Model - - model = Model() - model.load_cif(str(sample_cif_file)) - - elements = model.pdb['element'].unique() - + model = shared_model + + elements = model.pdb["element"].unique() + # Should have common protein elements - common_elements = ['C', 'N', 'O', 'S'] + common_elements = ["C", "N", "O", "S"] found_any = any(elem in elements for elem in common_elements) assert found_any, "No common elements found" @@ -70,58 +65,52 @@ class TestMTZReadingFunctional: def test_load_multiple_mtz_files(self, mtz_dir): """Test loading multiple MTZ files successfully.""" from torchref.io import ReflectionData - + mtz_files = list(mtz_dir.glob("*.mtz")) assert len(mtz_files) > 0, "No MTZ files found in test directory" - + for mtz_file in mtz_files: data = ReflectionData() data.load_mtz(str(mtz_file)) - + # Each file should load with reflections n_refl = data.hkl.shape[0] assert n_refl > 0, f"No reflections loaded from {mtz_file}" - + # Should have cell parameters assert data.cell is not None @pytest.mark.integration - def test_mtz_data_properties(self, sample_mtz_file): + def test_mtz_data_properties(self, shared_reflection_data): """Test that MTZ data properties are correctly loaded.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - + data = shared_reflection_data + # Check HKL indices are integers or can be converted hkl = data.hkl assert hkl.shape[1] == 3, "HKL should have 3 columns" - + # Check F values are loaded assert data.F is not None assert data.F.shape[0] == hkl.shape[0] - + # Check sigma values - if hasattr(data, 'F_sigma') and data.F_sigma is not None: + if hasattr(data, "F_sigma") and data.F_sigma is not None: assert data.F_sigma.shape[0] == hkl.shape[0] @pytest.mark.integration - def test_mtz_resolution_range(self, sample_mtz_file): + def test_mtz_resolution_range(self, shared_reflection_data): """Test that resolution range is computed correctly.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - + data = shared_reflection_data + # Check if resolution data is available - if hasattr(data, 'd') and data.d is not None: + if hasattr(data, "d") and data.d is not None: d_min = data.d.min().item() d_max = data.d.max().item() - + # Resolution should be positive assert d_min > 0 assert d_max > d_min - + # Typical protein data: 0.8 - 500 Å assert d_min > 0.5 assert d_max < 1000 @@ -134,16 +123,16 @@ class TestSFCIFReadingFunctional: def test_load_sf_cif(self, cif_sf_dir): """Test loading structure factor CIF files.""" from torchref.io import ReflectionData - + sf_files = list(cif_sf_dir.glob("*.cif")) if not sf_files: pytest.skip("No SF-CIF files found") - + for sf_file in sf_files: data = ReflectionData() try: data.load_cif(str(sf_file)) - + # Should have loaded reflections if data.hkl is not None: assert data.hkl.shape[0] > 0 @@ -158,40 +147,42 @@ class TestDataConsistencyFunctional: @pytest.mark.integration def test_cell_parameters_match(self, sample_structure_pair): """Test that cell parameters match between model and reflections.""" - from torchref.model.model import Model from torchref.io import ReflectionData - + from torchref.model.model import Model + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + model_cell = model.cell data_cell = data.cell - + if model_cell is not None and data_cell is not None: # Convert to tensors if needed if not isinstance(model_cell, torch.Tensor): model_cell = torch.tensor(model_cell) if not isinstance(data_cell, torch.Tensor): data_cell = torch.tensor(data_cell) - + # Cell parameters should be similar (1% tolerance) - assert torch.allclose(model_cell.float(), data_cell.float(), rtol=0.01, atol=0.1) + assert torch.allclose( + model_cell.float(), data_cell.float(), rtol=0.01, atol=0.1 + ) @pytest.mark.integration def test_spacegroup_consistency(self, sample_structure_pair): """Test that spacegroup is consistent.""" - from torchref.model.model import Model from torchref.io import ReflectionData - + from torchref.model.model import Model + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + # Both should have spacegroup defined assert model.spacegroup is not None @@ -203,13 +194,13 @@ class TestDataBinningFunctional: def test_get_bins(self, sample_mtz_file): """Test resolution binning of reflection data.""" from torchref.io import ReflectionData - + data = ReflectionData() data.load_mtz(str(sample_mtz_file)) - + # Get bins bins, n_bins = data.get_bins(n_bins=10) - + assert bins is not None assert bins.shape[0] == data.hkl.shape[0] assert bins.min() >= 0 @@ -219,20 +210,20 @@ def test_get_bins(self, sample_mtz_file): def test_mean_res_per_bin(self, sample_mtz_file): """Test mean resolution per bin calculation.""" from torchref.io import ReflectionData - + data = ReflectionData() data.load_mtz(str(sample_mtz_file)) - + # Get bins first bins, n_bins = data.get_bins(n_bins=10) - + # Get mean resolution per bin - if hasattr(data, 'mean_res_per_bin'): + if hasattr(data, "mean_res_per_bin"): mean_res = data.mean_res_per_bin() - + assert mean_res is not None assert len(mean_res) == n_bins - + # Mean resolution should decrease with bin index (low res to high res) # or increase (high res to low res) - depends on implementation assert torch.all(torch.isfinite(mean_res)) @@ -242,16 +233,13 @@ class TestFrenchWilsonFunctional: """Test French-Wilson conversion with real data.""" @pytest.mark.integration - def test_french_wilson_applied(self, sample_mtz_file): + def test_french_wilson_applied(self, shared_reflection_data): """Test that French-Wilson conversion is applied.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - + data = shared_reflection_data + # After French-Wilson, F values should be non-negative valid_F = data.F[~torch.isnan(data.F)] - + if len(valid_F) > 0: # All valid F values should be >= 0 assert torch.all(valid_F >= 0) @@ -261,32 +249,28 @@ class TestRfreeHandlingFunctional: """Test R-free flag handling.""" @pytest.mark.integration - def test_rfree_flags_loaded(self, sample_mtz_file): + def test_rfree_flags_loaded(self, shared_reflection_data): """Test that R-free flags are loaded or generated.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - + data = shared_reflection_data + # Should have rfree attribute - if hasattr(data, 'rfree') and data.rfree is not None: + if hasattr(data, "rfree") and data.rfree is not None: assert data.rfree.shape[0] == data.hkl.shape[0] - + # Should be boolean or can be converted to boolean - assert data.rfree.dtype == torch.bool or torch.all((data.rfree == 0) | (data.rfree == 1)) + assert data.rfree.dtype == torch.bool or torch.all( + (data.rfree == 0) | (data.rfree == 1) + ) @pytest.mark.integration - def test_rfree_fraction(self, sample_mtz_file): + def test_rfree_fraction(self, shared_reflection_data): """Test R-free set fraction is reasonable.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - if hasattr(data, 'rfree') and data.rfree is not None: + data = shared_reflection_data + + if hasattr(data, "rfree") and data.rfree is not None: # Work set mask (True for work, False for test) work_fraction = data.rfree.float().mean().item() - + # Typically 90-95% work set, 5-10% test set # So work_fraction should be 0.9-0.95 typically assert 0.7 < work_fraction <= 1.0 @@ -296,16 +280,13 @@ class TestMaskHandlingFunctional: """Test reflection mask handling.""" @pytest.mark.integration - def test_masks_method(self, sample_mtz_file): + def test_masks_method(self, shared_reflection_data): """Test masks() method returns valid mask.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - if hasattr(data, 'masks'): + data = shared_reflection_data + + if hasattr(data, "masks"): mask = data.masks() - + assert mask is not None assert mask.shape[0] == data.hkl.shape[0] assert mask.dtype == torch.bool diff --git a/tests/functional/test_model_ft_functional.py b/tests/functional/test_model_ft_functional.py index b8919254..a6abd808 100644 --- a/tests/functional/test_model_ft_functional.py +++ b/tests/functional/test_model_ft_functional.py @@ -4,10 +4,9 @@ These tests exercise the ModelFT class with real crystallographic data, testing the FFT-based structure factor calculation pipeline. """ + import pytest import torch -import numpy as np -from pathlib import Path @pytest.mark.integration @@ -17,7 +16,7 @@ class TestModelFTInitialization: def test_modelft_empty_initialization(self): """Test empty ModelFT initialization.""" from torchref.model.model_ft import ModelFT - + model = ModelFT() assert model is not None assert model.max_res == 1.0 # Default @@ -32,10 +31,10 @@ def test_modelft_with_custom_resolution(self): def test_modelft_load_cif(self, sample_cif_file): """Test loading a CIF file into ModelFT.""" from torchref.model.model_ft import ModelFT - + model = ModelFT(max_res=2.0, verbose=0) model.load_cif(str(sample_cif_file)) - + # Verify basic properties assert model.xyz() is not None assert model.xyz().shape[0] > 0 @@ -59,39 +58,33 @@ def test_modelft_has_gridsize(self, sample_cif_file): class TestModelFTParametrization: """Test ModelFT parametrization with real structures.""" - def test_parametrization_built(self, sample_cif_file): + def test_parametrization_built(self, shared_model_ft): """Test that parametrization is built after loading.""" - from torchref.model.model_ft import ModelFT - - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) - + + model = shared_model_ft + # Parametrization should be set assert model.parametrization is not None - def test_scattering_factors_available(self, sample_cif_file): + def test_scattering_factors_available(self, shared_model_ft): """Test that scattering factors can be computed.""" - from torchref.model.model_ft import ModelFT - - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) - + + model = shared_model_ft + # Should be able to access atom properties xyz = model.xyz() assert xyz is not None assert xyz.dtype == torch.float32 or xyz.dtype == torch.float64 -@pytest.mark.integration +@pytest.mark.integration class TestModelFTGridOperations: """Test ModelFT grid operations.""" - def test_setup_grid(self, sample_cif_file): + def test_setup_grid(self, loaded_model_ft): """An explicit grid size overrides the resolution-derived one.""" - from torchref.model.model_ft import ModelFT - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) + model = loaded_model_ft derived = model.grid_shape assert derived is not None and len(derived) == 3 @@ -107,16 +100,15 @@ def test_setup_grid(self, sample_cif_file): class TestModelFTRealSpaceMap: """Test ModelFT real space electron density map construction.""" - def test_get_real_space_grid(self, sample_cif_file): + def test_get_real_space_grid(self, loaded_model_ft): """Test getting real space grid.""" - from torchref.model.model_ft import ModelFT from torchref.base.math_torch import get_real_grid - - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) + + # The grid helper moves its Cell in place when targeting CPU. + model = loaded_model_ft assert model.gridsize is not None - grid = get_real_grid(model.cell, max_res=2.0, device='cpu') + grid = get_real_grid(model.cell, max_res=2.0, device="cpu") assert grid is not None assert len(grid.shape) == 4 # Should be 4D (nx, ny, nz, 3) @@ -126,16 +118,14 @@ def test_get_real_space_grid(self, sample_cif_file): class TestModelFTSymmetry: """Test ModelFT symmetry operations.""" - def test_map_symmetry_available(self, sample_cif_file): + def test_map_symmetry_available(self, shared_model_ft): """Test map symmetry is available after loading.""" - from torchref.model.model_ft import ModelFT - - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) - + + model = shared_model_ft + # Model should have spacegroup after loading assert model.spacegroup is not None - + # The map operator comes from the space group, keyed on the grid shape. gridsize = model.grid_shape assert gridsize is not None @@ -149,21 +139,20 @@ def test_map_symmetry_available(self, sample_cif_file): class TestModelFTStateDictFunctional: """Test ModelFT state dict operations with real data.""" - def test_save_and_load_state_dict(self, sample_cif_file, tmp_path): + def test_save_and_load_state_dict(self, loaded_model_ft, tmp_path): """Test saving and loading state dict.""" from torchref.model.model_ft import ModelFT - - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) - + + model = loaded_model_ft + original_xyz = model.xyz().clone() - + # Save state dict state_dict = model.state_dict() - + # Create new model and load state model2 = ModelFT(max_res=2.0, verbose=0) - + # We need to ensure proper initialization # For now just verify state_dict works assert state_dict is not None @@ -174,25 +163,23 @@ def test_save_and_load_state_dict(self, sample_cif_file, tmp_path): class TestModelFTForwardPass: """Test ModelFT forward pass (structure factor calculation).""" - def test_forward_method_exists(self, sample_cif_file): + def test_forward_method_exists(self, shared_model_ft): """Test that forward method is available.""" - from torchref.model.model_ft import ModelFT - - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) - + + model = shared_model_ft + # Check forward method exists - assert hasattr(model, 'forward') + assert hasattr(model, "forward") def test_build_map_method(self, sample_cif_file): """Test build_map method if available.""" from torchref.model.model_ft import ModelFT - + model = ModelFT(max_res=3.0, verbose=0) # Lower res for faster test model.load_cif(str(sample_cif_file)) - + # Check build_map method - if hasattr(model, 'build_map'): + if hasattr(model, "build_map"): # Try to build map try: model.build_map() @@ -209,22 +196,22 @@ class TestModelFTMultipleStructures: def test_modelft_multiple_structures(self, all_structure_pairs): """Test ModelFT works with different structures.""" from torchref.model.model_ft import ModelFT - + tested = 0 for pair in all_structure_pairs[:3]: # Test first 3 try: model = ModelFT(max_res=3.0, verbose=0) model.load_cif(str(pair["model"])) - + # Basic checks assert model.xyz() is not None assert model.xyz().shape[0] > 0 - + tested += 1 except Exception as e: # Some structures may fail to load continue - + assert tested >= 1, "At least one structure should load" @@ -232,27 +219,23 @@ def test_modelft_multiple_structures(self, all_structure_pairs): class TestModelFTCaching: """Test ModelFT caching mechanism.""" - def test_cache_initialization(self, sample_cif_file): + def test_cache_initialization(self, loaded_model_ft): """Test that CachedForwardMixin cache starts empty.""" - from torchref.model.model_ft import ModelFT - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) + model = loaded_model_ft # Mixin cache should start empty (lazily initialized) assert getattr(model, "_fwd_cached_output", None) is None - def test_cache_usage(self, sample_cif_file): + def test_cache_usage(self, shared_model_ft): """Test that cache can be used for computations.""" - from torchref.model.model_ft import ModelFT - - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) - + + model = shared_model_ft + # Access xyz twice - should use caching xyz1 = model.xyz() xyz2 = model.xyz() - + # Should return same tensor assert torch.allclose(xyz1, xyz2) @@ -261,14 +244,12 @@ def test_cache_usage(self, sample_cif_file): class TestModelFTCoordinateOperations: """Test ModelFT coordinate operations.""" - def test_cartesian_to_fractional(self, sample_cif_file): + def test_cartesian_to_fractional(self, shared_model_ft): """Test coordinate conversion.""" - from torchref.model.model_ft import ModelFT from torchref.base.math_torch import cartesian_to_fractional_torch - - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) - + + model = shared_model_ft + xyz = model.xyz() cell = model.cell @@ -278,17 +259,15 @@ def test_cartesian_to_fractional(self, sample_cif_file): # Fractional coords should be bounded (mostly between 0 and 1) assert frac.shape == xyz.shape - def test_fractional_to_cartesian(self, sample_cif_file): + def test_fractional_to_cartesian(self, shared_model_ft): """Test fractional to cartesian conversion.""" - from torchref.model.model_ft import ModelFT from torchref.base.math_torch import ( cartesian_to_fractional_torch, - fractional_to_cartesian_torch + fractional_to_cartesian_torch, ) - - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) - + + model = shared_model_ft + xyz = model.xyz() cell = model.cell @@ -304,24 +283,20 @@ def test_fractional_to_cartesian(self, sample_cif_file): class TestModelFTAnisoHandling: """Test ModelFT handling of anisotropic parameters.""" - def test_access_aniso_atoms(self, sample_cif_file): + def test_access_aniso_atoms(self, shared_model_ft): """Test accessing anisotropic atom information.""" - from torchref.model.model_ft import ModelFT - - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) - + + model = shared_model_ft + # Check if aniso is available - if hasattr(model, 'get_aniso') or hasattr(model, 'aniso'): + if hasattr(model, "get_aniso") or hasattr(model, "aniso"): # Structure has aniso pass - def test_isotropic_b_factors(self, sample_cif_file): + def test_isotropic_b_factors(self, shared_model_ft): """Test accessing isotropic B-factors.""" - from torchref.model.model_ft import ModelFT - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) + model = shared_model_ft # Get B-factors (now accessed via adp()) b_factors = model.adp() diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py index 8bdf5d61..424307e5 100644 --- a/tests/integration/conftest.py +++ b/tests/integration/conftest.py @@ -1,7 +1,6 @@ -""" -Integration test specific fixtures. -Integration tests use real file I/O and test the full pipeline. +"""Use shared fixtures registered by the root conftest for integration tests. -All shared fixtures (sample files, path fixtures, monomer library, etc.) -are defined in the root tests/conftest.py and are automatically available here. +Pipeline-specific fixtures belong in their consuming modules. Mutable loaded +objects from ``tests.fixtures.objects`` are function-scoped unless documented +as explicitly shared. """ diff --git a/tests/unit/conftest.py b/tests/unit/conftest.py index 96ce096c..ecc8cdde 100644 --- a/tests/unit/conftest.py +++ b/tests/unit/conftest.py @@ -1,148 +1,18 @@ -""" -Unit test specific fixtures. -Unit tests should NOT use real file I/O - use mocks or minimal in-memory data. -""" -import pytest -import torch -import numpy as np - -from torchref.config import dtypes - - -@pytest.fixture -def random_seed(): - """Set random seed for reproducibility.""" - seed = 42 - np.random.seed(seed) - torch.manual_seed(seed) - return seed - - -@pytest.fixture -def random_coordinates(): - """Generate random atomic coordinates.""" - def _generate(n_atoms: int = 10, seed: int = 42): - np.random.seed(seed) - return torch.tensor(np.random.rand(n_atoms, 3) * 10, dtype=dtypes.float) - return _generate - - -@pytest.fixture -def random_fractional_coordinates(): - """Generate random fractional coordinates (0-1 range).""" - def _generate(n_atoms: int = 10, seed: int = 42): - np.random.seed(seed) - return torch.tensor(np.random.rand(n_atoms, 3), dtype=dtypes.float) - return _generate - - -@pytest.fixture -def random_adp(): - """Generate random ADPs (atomic displacement parameters, reasonable range 10-60 Ų).""" - def _generate(n_atoms: int = 10, seed: int = 42): - np.random.seed(seed) - return torch.tensor(np.random.rand(n_atoms) * 50 + 10, dtype=dtypes.float) - return _generate - - -@pytest.fixture -def random_occupancies(): - """Generate random occupancies (0-1 range).""" - def _generate(n_atoms: int = 10, seed: int = 42): - np.random.seed(seed) - return torch.tensor(np.random.rand(n_atoms) * 0.5 + 0.5, dtype=dtypes.float) - return _generate - - -@pytest.fixture -def mock_cell(): - """Mock cell parameters [a, b, c, alpha, beta, gamma].""" - return torch.tensor([50.0, 60.0, 70.0, 90.0, 90.0, 90.0], dtype=dtypes.float) - - -@pytest.fixture -def mock_cell_triclinic(): - """Mock triclinic cell parameters.""" - return torch.tensor([40.0, 50.0, 60.0, 70.0, 80.0, 85.0], dtype=dtypes.float) - - -@pytest.fixture -def mock_hkl_indices(): - """Generate mock HKL indices.""" - def _generate(n_reflections: int = 100, max_index: int = 10, seed: int = 42): - np.random.seed(seed) - h = np.random.randint(-max_index, max_index + 1, n_reflections) - k = np.random.randint(-max_index, max_index + 1, n_reflections) - l = np.random.randint(-max_index, max_index + 1, n_reflections) - # Exclude (0,0,0) - mask = ~((h == 0) & (k == 0) & (l == 0)) - h, k, l = h[mask], k[mask], l[mask] - return torch.tensor(np.stack([h, k, l], axis=1), dtype=dtypes.float) - return _generate - - -@pytest.fixture -def mock_structure_factors(): - """Generate mock structure factors (complex).""" - def _generate(n_reflections: int = 100, seed: int = 42): - np.random.seed(seed) - real = np.random.randn(n_reflections) * 100 - imag = np.random.randn(n_reflections) * 100 - return torch.tensor(real + 1j * imag, dtype=dtypes.complex) - return _generate - - -@pytest.fixture -def mock_F_obs(): - """Generate mock observed structure factor amplitudes.""" - def _generate(n_reflections: int = 100, seed: int = 42): - np.random.seed(seed) - # Positive values with realistic distribution - return torch.tensor(np.abs(np.random.randn(n_reflections) * 100) + 10, dtype=dtypes.float) - return _generate - - -@pytest.fixture -def mock_F_sigma(): - """Generate mock sigma values for F_obs.""" - def _generate(n_reflections: int = 100, seed: int = 42): - np.random.seed(seed) - return torch.tensor(np.abs(np.random.randn(n_reflections) * 5) + 1, dtype=dtypes.float) - return _generate - - -@pytest.fixture -def mock_aniso_u(): - """Generate mock anisotropic U tensor components [U11, U22, U33, U12, U13, U23].""" - def _generate(n_atoms: int = 10, seed: int = 42): - np.random.seed(seed) - # Diagonal elements (positive) - u11 = np.random.rand(n_atoms) * 0.05 + 0.02 - u22 = np.random.rand(n_atoms) * 0.05 + 0.02 - u33 = np.random.rand(n_atoms) * 0.05 + 0.02 - # Off-diagonal elements (can be negative, smaller magnitude) - u12 = (np.random.rand(n_atoms) - 0.5) * 0.02 - u13 = (np.random.rand(n_atoms) - 0.5) * 0.02 - u23 = (np.random.rand(n_atoms) - 0.5) * 0.02 - return torch.tensor(np.stack([u11, u22, u33, u12, u13, u23], axis=1), dtype=dtypes.float) - return _generate - - -@pytest.fixture -def mock_scattering_factors(): - """Generate mock scattering factors.""" - def _generate(n_reflections: int = 100, n_atoms: int = 10, seed: int = 42): - np.random.seed(seed) - # Decreasing with resolution (approximate) - return torch.tensor(np.random.rand(n_reflections, n_atoms) * 5 + 1, dtype=dtypes.float) - return _generate - - -@pytest.fixture -def mock_weights(): - """Generate mock weights for atoms.""" - def _generate(n_atoms: int = 10, seed: int = 42): - np.random.seed(seed) - weights = np.random.rand(n_atoms) - return torch.tensor(weights / weights.sum(), dtype=dtypes.float).reshape(-1, 1) - return _generate +"""Expose synthetic fixtures only to the unit-test subtree.""" + +from tests.fixtures.numerical import ( # noqa: F401 + mock_aniso_u, + mock_cell, + mock_cell_triclinic, + mock_F_obs, + mock_F_sigma, + mock_hkl_indices, + mock_scattering_factors, + mock_structure_factors, + mock_weights, + random_adp, + random_coordinates, + random_fractional_coordinates, + random_occupancies, + random_seed, +) diff --git a/tests/unit/structure_factor/conftest.py b/tests/unit/structure_factor/conftest.py index bcbf9be4..7ab42565 100644 --- a/tests/unit/structure_factor/conftest.py +++ b/tests/unit/structure_factor/conftest.py @@ -13,13 +13,11 @@ import torch import torchref -from torchref.config import device as device_cfg, dtypes - -from tests.conftest import _accelerator +from tests.fixtures.devices import _accelerator +from tests.fixtures.precision import cpu_double_precision from . import helpers as H - # --------------------------------------------------------------------------- # Device axis # --------------------------------------------------------------------------- @@ -86,7 +84,9 @@ def ds_device_dtype_kernels(): for name in H.ds_kernels_for(device, dtype): out.append( pytest.param( - device, dtype, name, + device, + dtype, + name, id=f"{device.type}-{str(dtype).replace('torch.float', 'f')}-{name}", marks=dev_param.marks, ) @@ -103,30 +103,9 @@ def ds_device_dtype_kernels(): # --------------------------------------------------------------------------- @pytest.fixture(scope="package", autouse=True) def _float64_cpu(): - """float64/complex128 on CPU for this package; restore afterwards. - - Required, not cosmetic: ``iso_structure_factor_torched`` casts ``hkl`` to the - *global* ``dtypes.float`` (``torchref/base/direct_summation/isotropic.py:121``), so - under the default float32 config a float64 leaf produces a dtype-mismatched matmul. - That is why the pre-existing tests wrapped every eager-SF call in a ``double_cpu`` - fixture. - - ``sigma_cutoff_ed`` is restored here too -- the three copies of ``double_cpu`` this - replaces did not, so a test that changed the cutoff leaked it into everything that - ran after it. - """ - f0, c0, d0 = dtypes.float, dtypes.complex, device_cfg.current - s0 = torchref.sigma_cutoff_ed.value - dtypes.float = torch.float64 - dtypes.complex = torch.complex128 - device_cfg.current = torch.device("cpu") - try: + """Scope the package's CPU double-precision reference configuration.""" + with cpu_double_precision(): yield - finally: - dtypes.float = f0 - dtypes.complex = c0 - device_cfg.current = d0 - torchref.sigma_cutoff_ed.value = s0 @pytest.fixture From d7eebba0d5f9d02f01c866ceb34353134f4d8db8 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 7 Sep 2026 16:51:52 +0200 Subject: [PATCH 2/5] test: consolidate weighting contracts under their API owners LossState owns hierarchical multiplication, aggregation, zero weights and cached reads. Default group weights remain separate. NLL checks move to base metrics; gradnorm smoke checks become exact RMS expectations. Validation: 36 passed on default MPS and CPU float64. Follow-up: all-zero aggregate returns torch float32 under configured float64; the retained zero-weight contract checks the value, not output dtype. --- docs/changelog.rst | 1 + .../test_loss_weighting_functional.py | 198 ------------------ tests/unit/base/test_loss.py | 36 ++++ tests/unit/refinement/test_loss_state.py | 46 +++- tests/unit/refinement/test_loss_weighting.py | 51 +---- tests/unit/utils/test_gradnorm.py | 91 +++----- 6 files changed, 100 insertions(+), 323 deletions(-) delete mode 100644 tests/functional/test_loss_weighting_functional.py create mode 100644 tests/unit/base/test_loss.py diff --git a/docs/changelog.rst b/docs/changelog.rst index b5b00059..4e17fba9 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Consolidated weighting tests by API ownership and strengthened Gaussian-likelihood, gradient-norm, and cached-loss assertions. - Organized test fixtures into focused modules and reused module-scoped loaded objects for read-only functional checks while retaining fresh objects for mutation and loading tests. - Rigid-body refinement stores its Euler angles pre-multiplied by the chain's radius of gyration, so a unit step in an angle and a unit step in a translation displace atoms comparably. In radians against Angstroms the rotation block of the Hessian carried 190-530x the curvature of the translation block on 1DAW and 3E98 -- the geometric ``Rg**2``, 411 and 442/516 -- putting ``cond(H)`` at 1e3-5e3, which is why six parameters needed ~250 L-BFGS iterations to place. Dividing the scale out in ``forward()`` brings the ratio to 0.4-1.3 and ``cond(H)`` to 3-18. Over ten structures the step then converges rather than exhausting its iteration budget, on about half the gradient evaluations, with R-free no worse anywhere. ``RigidXYZTensor.rotation_radians`` returns the physical angle, and setting ``angle_scale`` to ones restores the unscaled parametrization. Not a fix for the one or two negative Hessian eigenvalues at the finer cutoffs -- scaling a saddle leaves it a saddle -- and those counts are unchanged - The rigid-body step no longer co-refines the scaler in the same L-BFGS as the rigid parameters. The body target centres on ``alpha*|F_calc|`` and ``alpha`` absorbs a rescaling of ``F_calc`` exactly, so the scale had a flat direction there; ``SCALE_TARGETS`` already excludes every alpha-centred row from the scale fit for this reason, and 0.6.2 fixed the same thing in the main driver. ``refine_scaler`` (objective ``ls``) owns the scale, between cutoffs diff --git a/tests/functional/test_loss_weighting_functional.py b/tests/functional/test_loss_weighting_functional.py deleted file mode 100644 index f1bcb602..00000000 --- a/tests/functional/test_loss_weighting_functional.py +++ /dev/null @@ -1,198 +0,0 @@ -""" -Functional tests for loss weighting module. - -These tests exercise the loss weighting strategies with realistic data. -Updated to use the new component_weighting and LossState architecture. -""" -import pytest -import torch -import numpy as np -from unittest.mock import Mock - - -@pytest.mark.integration -class TestLossStateWeightingFunctional: - """Test LossState weighting functionality.""" - - def test_loss_state_add_and_get_weights(self): - """Test adding and getting weights from LossState.""" - from torchref.refinement.loss_state import LossState - - state = LossState() - state.set_weight('xray', 1.5) - state.set_weight('geometry', 0.7) - - assert state.get_weight('xray') == 1.5 - assert state.get_weight('geometry') == 0.7 - - def test_loss_state_hierarchical_weights(self): - """Test hierarchical weights in LossState.""" - from torchref.refinement.loss_state import LossState - - state = LossState() - state.set_weight('geometry', 2.0) - state.set_weight('geometry/bond', 3.0) - - # Effective weight should be product: 2.0 * 3.0 = 6.0 - effective = state.get_effective_weight('geometry/bond') - assert effective == 6.0 - - -@pytest.mark.integration -class TestWeightingMathOperations: - """Test mathematical operations with weights.""" - - def test_total_weighted_loss_from_state(self): - """Test computing total weighted loss from LossState via aggregate.""" - from torchref.refinement.loss_state import LossState - - state = LossState() - state.register_target('xray', lambda: torch.tensor(10.0)) - state.register_target('geometry', lambda: torch.tensor(5.0)) - state.register_target('adp', lambda: torch.tensor(2.0)) - - state.set_weight('xray', 1.0) - state.set_weight('geometry', 0.5) - state.set_weight('adp', 0.25) - - total = state.aggregate() - - # Expected: 10*1.0 + 5*0.5 + 2*0.25 = 10 + 2.5 + 0.5 = 13.0 - assert torch.isclose(total, torch.tensor(13.0)) - - -@pytest.mark.integration -class TestNLLXrayFunction: - """Test the NLL X-ray function used in weighting.""" - - def test_nll_xray_basic(self): - """Test basic NLL X-ray calculation.""" - from torchref.base.math_torch import nll_xray - - fobs = torch.tensor([100.0, 200.0, 300.0], dtype=torch.float32) - fcalc = torch.tensor([105.0, 195.0, 305.0], dtype=torch.float32) - sigma = torch.tensor([10.0, 15.0, 20.0], dtype=torch.float32) - - nll = nll_xray(fobs, fcalc, sigma) - - # nll returns per-reflection values - assert torch.all(torch.isfinite(nll)) - - def test_nll_decreases_with_better_fit(self): - """Test that NLL decreases as fit improves.""" - from torchref.base.math_torch import nll_xray - - fobs = torch.tensor([100.0], dtype=torch.float32) - sigma = torch.tensor([10.0], dtype=torch.float32) - - # Good fit - fcalc_good = torch.tensor([100.0], dtype=torch.float32) - nll_good = nll_xray(fobs, fcalc_good, sigma) - - # Bad fit - fcalc_bad = torch.tensor([150.0], dtype=torch.float32) - nll_bad = nll_xray(fobs, fcalc_bad, sigma) - - # Good fit should have lower NLL - assert nll_good < nll_bad - - -@pytest.mark.integration -class TestGradnormUtility: - """Test the gradnorm utility function.""" - - def test_gradnorm_basic(self): - """Test basic gradnorm calculation.""" - from torchref.utils.gradnorm import gradnorm - - # Create simple parameter - param = torch.tensor([1.0, 2.0, 3.0], requires_grad=True) - - # Create loss - loss = param.sum() - - # Compute gradient norm - norm = gradnorm(loss, [param]) - - assert torch.isfinite(norm) - assert norm > 0 - - def test_gradnorm_with_multiple_params(self): - """Test gradnorm with multiple parameters.""" - from torchref.utils.gradnorm import gradnorm - - param1 = torch.tensor([1.0, 2.0], requires_grad=True) - param2 = torch.tensor([3.0, 4.0], requires_grad=True) - - loss = param1.sum() + param2.sum() - - norm = gradnorm(loss, [param1, param2]) - - assert torch.isfinite(norm) - assert norm > 0 - - -@pytest.mark.integration -class TestWeightingEdgeCases: - """Test edge cases in weighting.""" - - def test_zero_weight(self): - """Test zero weight (disabling a loss term).""" - from torchref.refinement.loss_state import LossState - - state = LossState() - state.register_target('adp', lambda: torch.tensor(100.0)) - state.set_weight('adp', 0.0) - - # Zero weight should effectively disable ADP term - total = state.aggregate() - assert torch.isclose(total, torch.tensor(0.0)) - - -@pytest.mark.integration -class TestLossAggregatorFunctional: - """Test LossAggregator functionality.""" - - def test_aggregator_basic(self): - """Test basic aggregator functionality (LossState.aggregate).""" - from torchref.refinement.loss_state import LossState - - state = LossState() - state.register_target('xray', lambda: torch.tensor(2.0)) - state.register_target('bond', lambda: torch.tensor(1.0)) - state.set_weight('xray', 1.0) - state.set_weight('bond', 0.5) - - total = state.aggregate() - - # Expected: 2.0 * 1.0 + 1.0 * 0.5 = 2.5 - expected = torch.tensor(2.5) - assert torch.isclose(total, expected) - - def test_loss_state_caches_losses(self): - """Test that LossState caches computed losses.""" - from torchref.refinement.loss_state import LossState - - call_count = [0] - def counting_target(): - call_count[0] += 1 - return torch.tensor(2.0) - - state = LossState() - state.register_target('xray', counting_target) - state.set_weight('xray', 1.0) - - # register_target probes the target once to walk the autograd graph; - # reset the counter so we measure only aggregate() invocations. - call_count[0] = 0 - - # First aggregation computes the loss - total1 = state.aggregate() - assert call_count[0] == 1 - - # Get cached loss doesn't recompute - cached = state.get_loss('xray') - assert cached is not None - assert torch.isclose(cached, torch.tensor(2.0)) - - diff --git a/tests/unit/base/test_loss.py b/tests/unit/base/test_loss.py new file mode 100644 index 00000000..eff9ffcb --- /dev/null +++ b/tests/unit/base/test_loss.py @@ -0,0 +1,36 @@ +"""Pin the amplitude-metric Gaussian likelihood's value and reduction contract.""" + +import math + +import pytest +import torch + +from torchref.base.metrics.loss import nll_xray, nll_xray_mean, nll_xray_sum +from torchref.config import get_default_device, get_float_dtype + +pytestmark = pytest.mark.unit + + +def test_gaussian_nll_value_and_reduction() -> None: + """The NLL includes its normalization and sums over reflections.""" + obs = torch.tensor( + [10.0, 20.0, 30.0], dtype=get_float_dtype(), device=get_default_device() + ) + sigma = obs.new_tensor([1.0, 2.0, 4.0]) + calc = obs + sigma + expected = obs.new_tensor(1.5 + math.log(8.0) + 1.5 * math.log(2.0 * math.pi)) + + torch.testing.assert_close(nll_xray(obs, calc, sigma), expected) + torch.testing.assert_close(nll_xray_sum(obs, calc, sigma), expected) + torch.testing.assert_close(nll_xray_mean(obs, calc, sigma), expected / obs.numel()) + + +def test_gaussian_nll_penalizes_amplitude_error() -> None: + """A one-sigma residual adds one half per reflection to the perfect-fit NLL.""" + obs = torch.tensor( + [10.0, 20.0, 30.0], dtype=get_float_dtype(), device=get_default_device() + ) + sigma = torch.ones_like(obs) + good = nll_xray(obs, obs, sigma) + bad = nll_xray(obs, obs + sigma, sigma) + torch.testing.assert_close(bad - good, obs.new_tensor(1.5)) diff --git a/tests/unit/refinement/test_loss_state.py b/tests/unit/refinement/test_loss_state.py index c6126464..837b0940 100644 --- a/tests/unit/refinement/test_loss_state.py +++ b/tests/unit/refinement/test_loss_state.py @@ -44,7 +44,9 @@ def test_register_target(self): from torchref.refinement.loss_state import LossState state = LossState() - target_fn = lambda: torch.tensor(1.0) + + def target_fn(): + return torch.tensor(1.0) result = state.register_target("geometry/bond", target_fn) @@ -136,6 +138,7 @@ def test_set_weight(self): result = state.set_weight("geometry", 0.5) assert state.weights["geometry"] == 0.5 + assert state.get_weight("geometry") == 0.5 assert result is state # Method chaining @pytest.mark.unit @@ -184,12 +187,11 @@ def test_get_effective_weight_hierarchical(self): from torchref.refinement.loss_state import LossState state = LossState() - state.set_weight("geometry", 0.5) - state.set_weight("geometry/bond", 2.0) + state.set_weight("geometry", 2.0) + state.set_weight("geometry/bond", 3.0) - # geometry/bond -> geometry (0.5) * geometry/bond (2.0) = 1.0 effective = state.get_effective_weight("geometry/bond") - assert effective == 1.0 + assert effective == 6.0 @pytest.mark.unit def test_get_effective_weight_missing_intermediate(self): @@ -207,6 +209,22 @@ def test_get_effective_weight_missing_intermediate(self): class TestAggregation: """Tests for loss aggregation.""" + @pytest.mark.unit + def test_zero_weight(self): + """A zero weight contributes zero to the aggregate.""" + from torchref.config import get_default_device, get_float_dtype + from torchref.refinement.loss_state import LossState + + value = torch.tensor( + 100.0, dtype=get_float_dtype(), device=get_default_device() + ) + state = LossState() + state.register_target("adp", lambda: value) + state.set_weight("adp", 0.0) + total = state.aggregate() + assert total.ndim == 0 + assert total.item() == 0.0 + @pytest.mark.unit def test_aggregate_simple(self): """Test simple aggregation.""" @@ -260,15 +278,25 @@ def test_aggregate_default_weights(self): @pytest.mark.unit def test_aggregate_caches_losses(self): """Test that aggregate caches computed losses.""" + from torchref.config import get_default_device, get_float_dtype from torchref.refinement.loss_state import LossState state = LossState() - state.register_target("xray", lambda: torch.tensor(2.0)) + value = torch.tensor(2.0, dtype=get_float_dtype(), device=get_default_device()) + calls = 0 - state.aggregate(log_values=False) + def target(): + nonlocal calls + calls += 1 + return value - loss = state.get_loss("xray") - assert torch.isclose(loss, torch.tensor(2.0)) + state.register_target("xray", target) + # Registration probes the autograd graph; count only subsequent evaluations. + calls = 0 + state.aggregate(log_values=False) + assert calls == 1 + torch.testing.assert_close(state.get_loss("xray"), value) + assert calls == 1 class TestHistoryLogging: diff --git a/tests/unit/refinement/test_loss_weighting.py b/tests/unit/refinement/test_loss_weighting.py index 08a7fe0c..d314e624 100644 --- a/tests/unit/refinement/test_loss_weighting.py +++ b/tests/unit/refinement/test_loss_weighting.py @@ -1,55 +1,6 @@ -""" -Unit tests for LossState weight handling. - -Covers the retained ``LossState`` weight API (``set_weight`` / -``get_effective_weight`` / ``aggregate``). The standalone weighting -schemes were removed; refinement now aggregates at uniform weight by -default, with explicit per-target/group multipliers set via the -``LossState`` weight dict. -""" +"""Pin refinement's default group weights; LossState owns weight arithmetic.""" import pytest -import torch - - -class TestLossStateWeights: - """Tests for the LossState weight dict (hierarchical multipliers).""" - - @pytest.mark.unit - def test_hierarchical_weights_multiply(self): - """Test that hierarchical weights multiply in get_effective_weight.""" - from torchref.refinement.loss_state import LossState - - state = LossState() - # Set group weight - state.set_weight('geometry', 2.0) - # Set component weight - state.set_weight('geometry/bond', 3.0) - - # Effective weight should multiply: 2.0 * 3.0 = 6.0 - effective = state.get_effective_weight('geometry/bond') - assert effective == 6.0 - - -class TestTotalLossFromState: - """Tests for computing total loss from LossState.""" - - @pytest.mark.unit - def test_total_weighted_loss(self): - """Test computing total weighted loss from state via aggregate.""" - from torchref.refinement.loss_state import LossState - - state = LossState() - state.register_target('xray', lambda: torch.tensor(2.0)) - state.register_target('bond', lambda: torch.tensor(1.0)) - state.set_weight('xray', 1.0) - state.set_weight('bond', 0.5) - - total = state.aggregate(log_values=False) - - # Expected: 2.0 * 1.0 + 1.0 * 0.5 = 2.5 - expected = torch.tensor(2.5) - assert torch.isclose(total, expected) class TestDefaultGroupWeights: diff --git a/tests/unit/utils/test_gradnorm.py b/tests/unit/utils/test_gradnorm.py index d8e1a621..62378f84 100644 --- a/tests/unit/utils/test_gradnorm.py +++ b/tests/unit/utils/test_gradnorm.py @@ -1,75 +1,34 @@ -""" -Unit tests for torchref.utils.gradnorm +"""Pin the RMS gradient norm across one or several parameter tensors.""" -Tests gradient norm calculation utilities. -""" +import math import pytest import torch -import torch.nn as nn +from torchref.config import get_default_device, get_float_dtype +from torchref.utils.gradnorm import gradnorm -class TestGradNorm: - """Tests for gradient norm calculation.""" +pytestmark = pytest.mark.unit - @pytest.mark.unit - def test_gradnorm_basic(self): - """Test basic gradient norm calculation.""" - from torchref.utils.gradnorm import gradnorm - - # Simple linear model - model = nn.Linear(10, 1, bias=False) - x = torch.randn(5, 10) - y = torch.randn(5, 1) - - # Forward pass - pred = model(x) - loss = ((pred - y) ** 2).mean() - - # Calculate gradient norm - grad_norm = gradnorm(loss, model.parameters()) - - assert isinstance(grad_norm, torch.Tensor) - assert grad_norm.ndim == 0 # Scalar - assert grad_norm >= 0 # Non-negative - @pytest.mark.unit - def test_gradnorm_zero_gradient(self): - """Gradient norm should handle zero gradients.""" - from torchref.utils.gradnorm import gradnorm - - model = nn.Linear(10, 1, bias=False) - - # Create a loss that depends on the model but has zero gradient - x = torch.randn(3, 10) - pred = model(x) - loss = (pred * 0.0).sum() # Zero gradient - # DON'T call backward before gradnorm - it calls backward internally - - grad_norm = gradnorm(loss, model.parameters()) - - # Should be 0 (zero gradients) - assert torch.isclose(grad_norm, torch.tensor(0.0, dtype=grad_norm.dtype), atol=1e-10) +@pytest.mark.parametrize("split", [False, True], ids=["single", "multiple"]) +def test_gradnorm_rms(split: bool) -> None: + """The norm weights individual gradient elements, not parameter tensors.""" + values = torch.tensor( + [1.0, 2.0, 3.0], dtype=get_float_dtype(), device=get_default_device() + ) + chunks = (values[:1], values[1:]) if split else (values,) + params = [chunk.clone().requires_grad_() for chunk in chunks] + loss = sum((param.square().sum() for param in params)) + expected = values.new_tensor(math.sqrt(56.0 / 3.0)) + torch.testing.assert_close(gradnorm(loss, iter(params)), expected) - @pytest.mark.unit - def test_gradnorm_multiple_params(self): - """Test gradient norm with multiple parameter groups.""" - from torchref.utils.gradnorm import gradnorm - - # Model with multiple layers - model = nn.Sequential( - nn.Linear(10, 5), - nn.ReLU(), - nn.Linear(5, 1) - ) - - x = torch.randn(3, 10) - y = torch.randn(3, 1) - - pred = model(x) - loss = ((pred - y) ** 2).mean() - - grad_norm = gradnorm(loss, model.parameters()) - - assert isinstance(grad_norm, torch.Tensor) - assert grad_norm >= 0 + +def test_gradnorm_zero_gradient() -> None: + """A connected loss with zero derivative has zero RMS gradient.""" + param = torch.ones( + 3, dtype=get_float_dtype(), device=get_default_device(), requires_grad=True + ) + torch.testing.assert_close( + gradnorm((param * 0).sum(), [param]), param.new_zeros(()) + ) From 7912ce1390b811142f42263abeef7d83fc7191d2 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 7 Sep 2026 16:58:27 +0200 Subject: [PATCH 3/5] test: replace target arithmetic demonstrations with kernel contracts Remove local-only arithmetic and duplicate Target initialization; nn.Module inheritance remains in the comprehensive Target contract. Keep anisotropic DELU gradient routing, SIGD references and R-factor calls. Add deposited-coordinate bond, angle, chiral, plane and SIMU references plus exact LS weighting/mask values. Retain kernel/gradient boundary tests for torsion and DELU. Validation: 48 passed/2 CUDA skips including gradient guards; 34 passed on CPU float64 after accounting for the squared-distance regularizer. Eight zero-return fault injections detected. --- docs/changelog.rst | 1 + tests/README.md | 16 + tests/RUNNING_TESTS.md | 21 +- tests/unit/base/test_target_values.py | 127 +++++++ tests/unit/refinement/test_targets.py | 200 ---------- .../refinement/test_targets_comprehensive.py | 341 ++---------------- 6 files changed, 181 insertions(+), 525 deletions(-) create mode 100644 tests/unit/base/test_target_values.py delete mode 100644 tests/unit/refinement/test_targets.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 4e17fba9..a8fe18e8 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Replaced local-arithmetic target tests with configured-device production-kernel checks on deposited coordinates and explicit least-squares expectations. - Consolidated weighting tests by API ownership and strengthened Gaussian-likelihood, gradient-norm, and cached-loss assertions. - Organized test fixtures into focused modules and reused module-scoped loaded objects for read-only functional checks while retaining fresh objects for mutation and loading tests. - Rigid-body refinement stores its Euler angles pre-multiplied by the chain's radius of gyration, so a unit step in an angle and a unit step in a translation displace atoms comparably. In radians against Angstroms the rotation block of the Hessian carried 190-530x the curvature of the translation block on 1DAW and 3E98 -- the geometric ``Rg**2``, 411 and 442/516 -- putting ``cond(H)`` at 1e3-5e3, which is why six parameters needed ~250 L-BFGS iterations to place. Dividing the scale out in ``forward()`` brings the ratio to 0.4-1.3 and ``cond(H)`` to 3-18. Over ten structures the step then converges rather than exhausting its iteration budget, on about half the gradient evaluations, with R-free no worse anywhere. ``RigidXYZTensor.rotation_radians`` returns the physical angle, and setting ``angle_scale`` to ones restores the unscaled parametrization. Not a fix for the one or two negative Hessian eigenvalues at the finer cutoffs -- scaling a saddle leaves it a saddle -- and those counts are unchanged diff --git a/tests/README.md b/tests/README.md index ff1a3bc5..9e0ae0d6 100644 --- a/tests/README.md +++ b/tests/README.md @@ -39,6 +39,22 @@ tests/ ## Running Tests +### Coverage ownership + +| Contract | Owner | +|---|---| +| Loss weights, aggregation, cached loss reads | `unit/refinement/test_loss_state.py` | +| Refinement's default group weights | `unit/refinement/test_loss_weighting.py` | +| Gaussian amplitude-metric values and reductions | `unit/base/test_loss.py` | +| Restraint kernel values on deposited coordinates | `unit/base/test_target_values.py` | +| Gradient RMS norm | `unit/utils/test_gradnorm.py` | +| Numerical derivatives and backend parity | `unit/test_gradient_correctness.py`, `unit/structure_factor/` | + +A production call must participate in the assertion: computing a formula only in +the test does not check its implementation. Kernel values, target registration, +device transitions, and default configuration are separate contracts even when +they exercise the same class. Keep mutation tests on fresh objects. + ### Quick Local Run (on login node, for small tests only) ```bash diff --git a/tests/RUNNING_TESTS.md b/tests/RUNNING_TESTS.md index f2c8c120..109079d7 100644 --- a/tests/RUNNING_TESTS.md +++ b/tests/RUNNING_TESTS.md @@ -92,7 +92,7 @@ pytest tests/unit/refinement/ -v pytest tests/unit/refinement/test_loss_weighting.py -v # Target/loss functions -pytest tests/unit/refinement/test_targets.py -v +pytest tests/unit/base/test_target_values.py tests/unit/base/test_loss.py -v ``` ### Scaling @@ -176,17 +176,17 @@ pytest tests/unit/model/test_parameter_wrappers.py::TestMixedTensorOperations -v #### Refinement Classes ```bash -# Fixed weighting -pytest tests/unit/refinement/test_loss_weighting.py::TestFixedWeighting -v +# Weight handling +pytest tests/unit/refinement/test_loss_state.py::TestWeightManagement -v -# Resolution-dependent weighting -pytest tests/unit/refinement/test_loss_weighting.py::TestResolutionDependentWeighting -v +# Default group weights +pytest tests/unit/refinement/test_loss_weighting.py::TestDefaultGroupWeights -v # Gaussian NLL loss -pytest tests/unit/refinement/test_targets.py::TestGaussianNLL -v +pytest tests/unit/base/test_loss.py -v # Least squares target -pytest tests/unit/refinement/test_targets.py::TestLeastSquaresTarget -v +pytest tests/unit/base/test_target_values.py -k least_squares -v ``` #### Symmetry Classes @@ -361,13 +361,14 @@ pytest tests/unit --lf -v | `math_functions/test_math_numpy.py` | `TestCoordinateTransformations`, `TestScatteringVectors`, `TestRFactorCalculations`, `TestRotation` | | `model/test_model.py` | `TestModelInitialization`, `TestModelDeviceHandling` | | `model/test_parameter_wrappers.py` | `TestMixedTensorInitialization`, `TestMixedTensorOperations`, `TestMixedTensorDeviceHandling`, `TestOccupancyTensor`, `TestPositiveMixedTensor` | -| `refinement/test_loss_weighting.py` | `TestFixedWeighting`, `TestResolutionDependentWeighting`, `TestLossWeightingModule` | -| `refinement/test_targets.py` | `TestTargetBase`, `TestGaussianNLL`, `TestLeastSquaresTarget`, `TestRiceNLL`, `TestTargetDeviceHandling`, `TestNumericStability` | +| `refinement/test_loss_weighting.py` | `TestDefaultGroupWeights` | +| `base/test_target_values.py` | Deposited-coordinate restraint values and least-squares weighting | +| `base/test_loss.py` | Gaussian NLL values and reductions | | `scaling/test_scaler.py` | `TestScalerInitialization`, `TestScalerDeviceHandling`, `TestScalingCalculations`, `TestBFactorScaling`, `TestAnisotropicScaling` | | `symmetrie/test_symmetrie.py` | `TestSymmetryInitialization`, `TestSymmetryMatrices`, `TestSymmetryApplication`, `TestSymmetryDeviceHandling`, `TestSpaceGroupMapping` | | `io/test_data.py` | `TestReflectionDataInitialization`, `TestReflectionDataDeviceMovement`, `TestReflectionDataAttributes`, `TestReflectionDataProperties`, `TestMockReflectionData` | | `restraints/test_restraints.py` | `TestRestraintsInitialization`, `TestBondRestraintCalculations`, `TestAngleRestraintCalculations`, `TestTorsionRestraintCalculations`, `TestRestraintDeviceHandling`, `TestRestraintNumericStability` | -| `utils/test_gradnorm.py` | `TestGradNorm` | +| `utils/test_gradnorm.py` | RMS norms for single/multiple parameters and zero gradients | | `utils/test_utils.py` | `TestModuleReference`, `TestCIFReader` | ### Integration Tests (`tests/integration/`) diff --git a/tests/unit/base/test_target_values.py b/tests/unit/base/test_target_values.py new file mode 100644 index 00000000..bf27f0a6 --- /dev/null +++ b/tests/unit/base/test_target_values.py @@ -0,0 +1,127 @@ +"""Compare restraint kernels with host references on deposited Cartesian coordinates.""" + +import math + +import numpy as np +import pytest +import torch + +from torchref.base.targets._common import EPS +from torchref.base.targets.adp import adp_simu_math +from torchref.base.targets.angle import angle_math +from torchref.base.targets.bond import bond_math +from torchref.base.targets.chiral import chiral_math +from torchref.base.targets.planarity import planarity_math +from torchref.base.targets.xray_ls import ls_xray_loss_math +from torchref.config import get_default_device, get_float_dtype, get_int_dtype + +pytestmark = pytest.mark.unit + + +@pytest.fixture(scope="module") +def deposited_atoms(sample_cif_file): + """Return detached Cartesian coordinates (Å) and isotropic B-factors (Ų).""" + from torchref.model import Model + + model = Model(verbose=0) + model.load_cif(str(sample_cif_file)) + return model.xyz().detach().clone(), model.adp().detach().clone() + + +def _indices(rows, device): + return torch.tensor(rows, dtype=get_int_dtype(), device=device) + + +def _gaussian_sum(residual, sigma): + return np.sum( + 0.5 * (residual / sigma) ** 2 + np.log(sigma) + 0.5 * math.log(2 * math.pi) + ) + + +def test_bond_value(deposited_atoms) -> None: + """Bond lengths enter a summed Gaussian NLL in Å.""" + xyz, _ = deposited_atoms + host = xyz[:4].cpu().numpy().astype(np.float64) + idx = _indices([[0, 1], [2, 3]], xyz.device) + refs = xyz.new_tensor([1.4, 1.5]) + sigma = xyz.new_tensor([0.1, 0.2]) + # The kernel regularizes squared distance to keep coincident-atom gradients finite. + distance = np.sqrt(np.sum((host[[0, 2]] - host[[1, 3]]) ** 2, axis=1) + EPS) + expected = _gaussian_sum(distance - refs.cpu().numpy(), sigma.cpu().numpy()) + torch.testing.assert_close( + bond_math(xyz, idx, refs, sigma), xyz.new_tensor(expected) + ) + + +def test_angle_value(deposited_atoms) -> None: + """Angles and their restraint sigmas enter the NLL in radians.""" + import gemmi + + xyz, _ = deposited_atoms + positions = [gemmi.Position(*row) for row in xyz[:4].cpu().tolist()] + angles = np.array( + [gemmi.calculate_angle(*positions[:3]), gemmi.calculate_angle(*positions[1:4])] + ) + idx = _indices([[0, 1, 2], [1, 2, 3]], xyz.device) + refs = xyz.new_tensor([1.8, 2.0]) + sigma = xyz.new_tensor([0.1, 0.2]) + expected = _gaussian_sum(angles - refs.cpu().numpy(), sigma.cpu().numpy()) + torch.testing.assert_close( + angle_math(xyz, idx, refs, sigma), xyz.new_tensor(expected) + ) + + +def test_chiral_value(deposited_atoms) -> None: + """The signed scalar triple product, without a 1/6 factor, sets chirality.""" + xyz, _ = deposited_atoms + host = xyz[:4].cpu().numpy().astype(np.float64) + volume = np.linalg.det(host[1:] - host[0]) + idx = _indices([[0, 1, 2, 3]], xyz.device) + refs = xyz.new_tensor([2.0]) + sigma = xyz.new_tensor([0.5]) + expected = _gaussian_sum(volume - 2.0, 0.5) + torch.testing.assert_close( + chiral_math(xyz, idx, refs, sigma), xyz.new_tensor(expected) + ) + + +def test_planarity_value(deposited_atoms) -> None: + """The plane penalty sums signed-distance Gaussian NLLs over its atoms.""" + xyz, _ = deposited_atoms + host = xyz[:5].cpu().numpy().astype(np.float64) + centered = host - host.mean(axis=0) + _, _, vh = np.linalg.svd(centered, full_matrices=False) + distances = centered @ vh[-1] + idx = _indices([[0, 1, 2, 3, 4]], xyz.device) + sigma = xyz.new_full((1, 5), 0.2) + expected = _gaussian_sum(distances, 0.2) + torch.testing.assert_close( + planarity_math(xyz, [(idx, sigma)]), xyz.new_tensor(expected) + ) + + +def test_simu_value(deposited_atoms) -> None: + """SIMU penalizes differences of deposited isotropic B-factors in Ų.""" + _, b = deposited_atoms + host = b[:4].cpu().numpy().astype(np.float64) + idx = _indices([[0, 1], [2, 3]], b.device) + expected = _gaussian_sum(host[[0, 2]] - host[[1, 3]], 2.0) + torch.testing.assert_close( + adp_simu_math(b, idx, b.new_tensor(2.0)), b.new_tensor(expected) + ) + + +@pytest.mark.parametrize("weighting, expected", [("sigma", 6.5), ("unit", 20.0)]) +def test_least_squares_value_and_mask(weighting: str, expected: float) -> None: + """Least squares sums half squared amplitude errors using the selected weights.""" + obs = torch.tensor( + [10.0, 20.0, 30.0], dtype=get_float_dtype(), device=get_default_device() + ) + calc = -obs - obs.new_tensor([2.0, 6.0, 50.0]) + sigma = obs.new_tensor([1.0, 2.0, 5.0]) + mask = torch.tensor([True, True, False], device=obs.device) + loss = ls_xray_loss_math(obs, calc, sigma, mask, weighting=weighting) + torch.testing.assert_close(loss, obs.new_tensor(expected)) + torch.testing.assert_close( + ls_xray_loss_math(obs, obs, sigma, weighting=weighting), obs.new_zeros(()) + ) diff --git a/tests/unit/refinement/test_targets.py b/tests/unit/refinement/test_targets.py deleted file mode 100644 index fe2728e7..00000000 --- a/tests/unit/refinement/test_targets.py +++ /dev/null @@ -1,200 +0,0 @@ -""" -Unit tests for torchref.refinement.targets - -Tests target (loss) functions for crystallographic refinement. -Note: These are unit tests so we test the functions in isolation with mock data. -""" - -import pytest -import torch -import torch.nn as nn -import numpy as np - - -class TestTargetBase: - """Tests for base Target class.""" - - @pytest.mark.unit - def test_target_empty_initialization(self): - """Test Target can be initialized without arguments.""" - from torchref.refinement.targets import Target - - target = Target() - - assert target.verbose == 0 - - @pytest.mark.unit - def test_target_is_nn_module(self): - """Target should be a nn.Module.""" - from torchref.refinement.targets import Target - - target = Target() - - assert isinstance(target, nn.Module) - - -class TestGaussianNLL: - """Tests for Gaussian NLL calculation logic.""" - - @pytest.mark.unit - def test_gaussian_nll_identical_gives_small_loss(self, mock_F_obs, mock_F_sigma): - """When Fobs = Fcalc, NLL should be small (just the log sigma term).""" - from torchref.base.math_torch import nll_xray - - fobs = mock_F_obs(n_reflections=100) - sigma = mock_F_sigma(n_reflections=100) - fcalc = fobs.clone().to(torch.complex64) # |Fcalc| = Fobs - - # Calculate manually what Gaussian NLL should be - # NLL = 0.5*((fobs - |fcalc|)/sigma)^2 + log(sigma) + 0.5*log(2pi) - diff = fobs - torch.abs(fcalc) - expected_data_term = 0.5 * ((diff / sigma) ** 2) - - # Data term should be ~0 when fobs = |fcalc| - assert torch.allclose(expected_data_term, torch.zeros_like(expected_data_term), atol=1e-5) - - @pytest.mark.unit - def test_gaussian_nll_positive(self, mock_F_obs, mock_F_sigma): - """NLL should generally be positive or close to zero.""" - fobs = mock_F_obs(n_reflections=100) - sigma = mock_F_sigma(n_reflections=100) - fcalc = mock_F_obs(n_reflections=100, seed=123).to(torch.complex64) # Different - - # Simple Gaussian NLL - diff = fobs - torch.abs(fcalc) - eps = torch.median(sigma) * 0.1 - sigma_safe = torch.clamp(sigma, min=eps) - log_2pi = torch.log(torch.tensor(2.0 * np.pi)) - nll = 0.5 * (diff ** 2) / (sigma_safe ** 2) + torch.log(sigma_safe) + 0.5 * log_2pi - - # Mean NLL should be finite - assert torch.isfinite(nll.mean()) - - -class TestLeastSquaresTarget: - """Tests for Least Squares target calculation.""" - - @pytest.mark.unit - def test_least_squares_identical_zero(self, mock_F_obs): - """LS loss should be 0 when Fobs = Fcalc.""" - fobs = mock_F_obs(n_reflections=100) - fcalc = fobs.clone() - - # Simple LS: sum((fobs - fcalc)^2) - loss = torch.sum((fobs - fcalc) ** 2) - - assert torch.isclose(loss, torch.tensor(0.0, dtype=loss.dtype), atol=1e-10) - - @pytest.mark.unit - def test_least_squares_scaled(self, mock_F_obs): - """Test LS loss with scaled Fcalc.""" - fobs = mock_F_obs(n_reflections=100) - fcalc = fobs * 1.1 # 10% scaled - - loss = torch.mean((fobs - fcalc) ** 2) - - # Should be (0.1 * fobs)^2 on average - expected_loss = torch.mean((0.1 * fobs) ** 2) - assert torch.isclose(loss, expected_loss, rtol=1e-5) - - @pytest.mark.unit - def test_least_squares_weighted(self, mock_F_obs, mock_F_sigma): - """Test weighted LS with sigma weights.""" - fobs = mock_F_obs(n_reflections=100) - sigma = mock_F_sigma(n_reflections=100) - fcalc = mock_F_obs(n_reflections=100, seed=123) - - # Weighted LS: sum(w * (fobs - fcalc)^2) where w = 1/sigma^2 - weights = 1.0 / (sigma ** 2) - diff = fobs - fcalc - weighted_loss = torch.sum(weights * (diff ** 2)) - - assert torch.isfinite(weighted_loss) - assert weighted_loss >= 0 - - -class TestRiceNLL: - """Tests for Rice distribution NLL (used for acentric reflections).""" - - @pytest.mark.unit - def test_rice_nll_components(self, mock_F_obs, mock_F_sigma): - """Test components of Rice NLL calculation.""" - from torch.special import i0 - - fobs = mock_F_obs(n_reflections=50) - sigma = mock_F_sigma(n_reflections=50) - fcalc_amp = mock_F_obs(n_reflections=50, seed=123) - - # Rice NLL components - # NLL = (Fo^2 + Fc^2)/(2σ^2) - log(I0(Fo*Fc/σ^2)) - log(Fo/σ^2) - - # Check I0 calculation - x = fobs * fcalc_amp / (sigma ** 2) - bessel_i0 = i0(x) - - # I0 should be >= 1 for x >= 0 - assert torch.all(bessel_i0 >= 1.0) - - -class TestTargetDeviceHandling: - """Tests for proper device handling in targets.""" - - @pytest.mark.unit - def test_target_cpu_tensors(self, mock_F_obs, mock_F_sigma): - """Test calculations work on CPU.""" - fobs = mock_F_obs(n_reflections=100) - sigma = mock_F_sigma(n_reflections=100) - - # Simple calculation on CPU - loss = torch.mean((fobs / sigma) ** 2) - - assert loss.device.type == 'cpu' - assert torch.isfinite(loss) - - @pytest.mark.unit - @pytest.mark.gpu - def test_target_gpu_tensors(self, mock_F_obs, mock_F_sigma, gpu_device): - """Test calculations work on GPU.""" - fobs = mock_F_obs(n_reflections=100).to(gpu_device) - sigma = mock_F_sigma(n_reflections=100).to(gpu_device) - - loss = torch.mean((fobs / sigma) ** 2) - - assert loss.device.type == gpu_device.type - assert torch.isfinite(loss) - - -class TestNumericStability: - """Tests for numeric stability in target calculations.""" - - @pytest.mark.unit - def test_small_sigma_handling(self, mock_F_obs): - """Test handling of very small sigma values.""" - fobs = mock_F_obs(n_reflections=100) - sigma = torch.ones_like(fobs) * 1e-10 # Very small sigma - fcalc = mock_F_obs(n_reflections=100, seed=123) - - # Clamped sigma approach - eps = torch.median(sigma) * 0.1 - sigma_safe = torch.clamp(sigma, min=max(eps, 1e-6)) - - diff = fobs - fcalc - loss = torch.mean((diff / sigma_safe) ** 2) - - assert torch.isfinite(loss) - - @pytest.mark.unit - def test_zero_fcalc_handling(self, mock_F_obs, mock_F_sigma): - """Test handling of zero Fcalc values.""" - fobs = mock_F_obs(n_reflections=100) - sigma = mock_F_sigma(n_reflections=100) - fcalc = torch.zeros_like(fobs, dtype=torch.complex64) # All zero - - fcalc_amp = torch.abs(fcalc) # Will be zero - diff = fobs - fcalc_amp - - loss = torch.mean(diff ** 2) - - # Should just be mean of fobs^2 - expected = torch.mean(fobs ** 2) - assert torch.isclose(loss, expected, rtol=1e-5) diff --git a/tests/unit/refinement/test_targets_comprehensive.py b/tests/unit/refinement/test_targets_comprehensive.py index 21aa6c9e..0593a129 100644 --- a/tests/unit/refinement/test_targets_comprehensive.py +++ b/tests/unit/refinement/test_targets_comprehensive.py @@ -4,16 +4,16 @@ These tests focus on individual target classes with mock/minimal data to achieve higher coverage of the targets module. """ + +import numpy as np import pytest import torch -import numpy as np -from unittest.mock import MagicMock, PropertyMock - # ============================================================================= # Base Target Tests # ============================================================================= + @pytest.mark.unit class TestBaseTarget: """Test base Target class functionality.""" @@ -24,18 +24,19 @@ def test_target_initialization_empty(self): target = Target() assert target.verbose == 0 + assert isinstance(target, torch.nn.Module) def test_target_initialization_with_verbose(self): """Test initialization with verbose setting.""" from torchref.refinement.targets import Target - + target = Target(verbose=2) assert target.verbose == 2 def test_target_forward_not_implemented(self): """Test that forward raises NotImplementedError.""" from torchref.refinement.targets import Target - + target = Target() with pytest.raises(NotImplementedError): target.forward() @@ -45,6 +46,7 @@ def test_target_forward_not_implemented(self): # X-ray Target Tests # ============================================================================= + @pytest.mark.unit class TestXrayTargetBase: """Test XrayTarget base class.""" @@ -71,39 +73,9 @@ def test_gaussian_target_initialization(self): assert target._model is None assert target._data is None - def test_gaussian_nll_computation(self): - """Test Gaussian NLL computation with mock data.""" - from torchref.base.math_torch import nll_xray - - # Test the underlying function - fobs = torch.tensor([1.0, 2.0, 3.0, 4.0], dtype=torch.float32) - fcalc = torch.tensor([1.1, 1.9, 3.2, 3.8], dtype=torch.float32) - sigma = torch.tensor([0.1, 0.1, 0.1, 0.1], dtype=torch.float32) - - loss = nll_xray(fobs, fcalc, sigma).mean() - assert torch.isfinite(loss) # NLL can be negative depending on normalization -@pytest.mark.unit -class TestLeastSquaresXrayTarget: - """Test LeastSquaresXrayTarget.""" - - def test_least_squares_computation(self): - """Test least squares computation with mock data.""" - # Least squares: sum of (fobs - fcalc)^2 / sigma^2 - fobs = torch.tensor([1.0, 2.0, 3.0, 4.0], dtype=torch.float32) - fcalc = torch.tensor([1.1, 1.9, 3.2, 3.8], dtype=torch.float32) - sigma = torch.tensor([0.1, 0.1, 0.1, 0.1], dtype=torch.float32) - - diff = fobs - fcalc - weights = 1.0 / (sigma ** 2) - loss = 0.5 * torch.sum(weights * (diff ** 2)) - - assert torch.isfinite(loss) - assert loss >= 0 - - @pytest.mark.unit class TestRiceXrayTarget: """Test RiceXrayTarget.""" @@ -121,6 +93,7 @@ def test_rice_target_initialization(self): # Geometry Target Tests # ============================================================================= + @pytest.mark.unit class TestGeometryTargetBase: """Test GeometryTarget base class.""" @@ -144,30 +117,6 @@ def test_bond_target_initialization(self): target = BondTarget() assert target._model is None - def test_bond_deviation_calculation(self): - """Test bond deviation calculation with mock data.""" - # Create mock coordinates for a simple bond - xyz = torch.tensor([ - [0.0, 0.0, 0.0], - [1.5, 0.0, 0.0], # 1.5 Å bond - ], dtype=torch.float32) - - # Bond indices - i_atoms = torch.tensor([0]) - j_atoms = torch.tensor([1]) - - # Expected distance and sigma - d_expected = torch.tensor([1.54]) # Expected C-C bond - sigma = torch.tensor([0.02]) - - # Calculate actual distances - d_actual = torch.norm(xyz[i_atoms] - xyz[j_atoms], dim=1) - - # Calculate deviation - deviation = (d_actual - d_expected) / sigma - - assert torch.isfinite(deviation).all() - @pytest.mark.unit class TestAngleTarget: @@ -180,27 +129,6 @@ def test_angle_target_initialization(self): target = AngleTarget() assert target._model is None - def test_angle_calculation(self): - """Test angle calculation with mock data.""" - # Create mock coordinates for a 90-degree angle - xyz = torch.tensor([ - [1.0, 0.0, 0.0], # Atom 1 - [0.0, 0.0, 0.0], # Atom 2 (vertex) - [0.0, 1.0, 0.0], # Atom 3 - ], dtype=torch.float32) - - # Vectors - v1 = xyz[0] - xyz[1] - v2 = xyz[2] - xyz[1] - - # Calculate angle - cos_angle = torch.dot(v1, v2) / (torch.norm(v1) * torch.norm(v2)) - angle = torch.acos(cos_angle) - angle_deg = torch.rad2deg(angle) - - # Should be approximately 90 degrees - assert torch.isclose(angle_deg, torch.tensor(90.0), atol=0.1) - @pytest.mark.unit class TestTorsionTarget: @@ -213,34 +141,6 @@ def test_torsion_target_initialization(self): target = TorsionTarget() assert target._model is None - def test_torsion_angle_calculation(self): - """Test torsion angle calculation.""" - # Create mock coordinates for a torsion - # Atoms in a plane should give ~0 or ~180 degree torsion - xyz = torch.tensor([ - [0.0, 0.0, 0.0], - [1.0, 0.0, 0.0], - [2.0, 0.0, 0.0], - [3.0, 0.0, 0.0], - ], dtype=torch.float32) - - # Calculate torsion using standard formula - b1 = xyz[1] - xyz[0] - b2 = xyz[2] - xyz[1] - b3 = xyz[3] - xyz[2] - - # Normal vectors - n1 = torch.linalg.cross(b1, b2) - n2 = torch.linalg.cross(b2, b3) - - # Torsion angle - if torch.norm(n1) > 1e-6 and torch.norm(n2) > 1e-6: - cos_torsion = torch.dot(n1, n2) / (torch.norm(n1) * torch.norm(n2)) - # Clamp to valid range - cos_torsion = torch.clamp(cos_torsion, -1.0, 1.0) - torsion = torch.acos(cos_torsion) - assert torch.isfinite(torsion) - @pytest.mark.unit class TestPlanarityTarget: @@ -253,29 +153,6 @@ def test_planarity_target_initialization(self): target = PlanarityTarget() assert target._model is None - def test_planarity_calculation(self): - """Test planarity calculation for coplanar atoms.""" - # Atoms in the XY plane - xyz = torch.tensor([ - [0.0, 0.0, 0.0], - [1.0, 0.0, 0.0], - [1.0, 1.0, 0.0], - [0.0, 1.0, 0.0], - ], dtype=torch.float32) - - # Calculate centroid - centroid = xyz.mean(dim=0) - - # Center coordinates - centered = xyz - centroid - - # SVD to find plane - U, S, Vh = torch.linalg.svd(centered) - - # The smallest singular value indicates planarity - # For perfectly coplanar points, it should be ~0 - assert S[-1] < 0.1 - @pytest.mark.unit class TestChiralTarget: @@ -288,26 +165,6 @@ def test_chiral_target_initialization(self): target = ChiralTarget() assert target._model is None - def test_chiral_volume_calculation(self): - """Test chiral volume calculation.""" - # Create a tetrahedron - xyz = torch.tensor([ - [1.0, 0.0, -1.0/np.sqrt(2)], # Center - [0.0, 0.0, 1.0/np.sqrt(2)], # Atom 1 - [1.0, 1.0, 0.0], # Atom 2 - [1.0, -1.0, 0.0], # Atom 3 - ], dtype=torch.float32) - - # Vectors from center to other atoms - v1 = xyz[1] - xyz[0] - v2 = xyz[2] - xyz[0] - v3 = xyz[3] - xyz[0] - - # Chiral volume (scalar triple product) - chiral_vol = torch.dot(v1, torch.linalg.cross(v2, v3)) - - assert torch.isfinite(chiral_vol) - @pytest.mark.unit class TestNonBondedTarget: @@ -337,6 +194,7 @@ def test_total_geometry_target_initialization(self): # ADP Target Tests # ============================================================================= + @pytest.mark.unit class TestADPTargetBase: """Test ADPTarget base class.""" @@ -349,61 +207,10 @@ def test_adp_target_initialization(self): assert target._model is None -@pytest.mark.unit -class TestADPSimilarityTarget: - """Test ADPSimilarityTarget (SIMU restraint).""" - - def test_simu_calculation(self): - """Test SIMU calculation with mock B-factors.""" - # Create mock B-factors for nearby atoms - b_factors = torch.tensor([20.0, 21.0, 22.0, 50.0], dtype=torch.float32) - - # Pairs of similar atoms (indices) - i_atoms = torch.tensor([0, 1]) - j_atoms = torch.tensor([1, 2]) - - # Calculate difference - diff = b_factors[i_atoms] - b_factors[j_atoms] - - # SIMU restraint loss - sigma = 1.0 # B-factor sigma - simu_loss = (diff / sigma).pow(2).mean() - - assert torch.isfinite(simu_loss) - assert simu_loss >= 0 - - @pytest.mark.unit class TestRigidBondTarget: """Test RigidBondTarget (DELU restraint).""" - def test_delu_calculation(self): - """Test DELU calculation with mock U matrices.""" - # Create mock anisotropic U matrices (6 parameters each) - # U11, U22, U33, U12, U13, U23 - u1 = torch.tensor([0.05, 0.06, 0.04, 0.01, 0.005, -0.01], dtype=torch.float32) - u2 = torch.tensor([0.05, 0.06, 0.04, 0.01, 0.005, -0.01], dtype=torch.float32) - - # Bond vector (normalized) - bond_vec = torch.tensor([1.0, 0.0, 0.0], dtype=torch.float32) - - # Calculate U components along bond direction - # For Uij, the component along direction v is v^T U v - def u_along_direction(u_params, direction): - """Calculate U component along a direction.""" - U11, U22, U33, U12, U13, U23 = u_params - vx, vy, vz = direction - return (U11 * vx * vx + U22 * vy * vy + U33 * vz * vz + - 2 * U12 * vx * vy + 2 * U13 * vx * vz + 2 * U23 * vy * vz) - - u1_bond = u_along_direction(u1, bond_vec) - u2_bond = u_along_direction(u2, bond_vec) - - # DELU restraint: difference should be small - diff = u1_bond - u2_bond - - assert torch.isfinite(diff) - def test_aniso_path_runs_and_routes_grad_to_u(self, pdb_dir): """The anisotropic DELU path actually executes and feeds gradient to the U tensors. Regression for the dead ``hasattr(model, "u_aniso")`` gate, @@ -467,27 +274,11 @@ def test_matches_inverse_gamma_nll(self): beta = float(b.mean()) * (alpha - 1.0) mode = beta / (alpha + 1.0) - expected = ( - -sps.invgamma.logpdf(b.numpy(), alpha, scale=beta).sum() - + sps.invgamma.logpdf(mode, alpha, scale=beta) * len(b) - ) + expected = -sps.invgamma.logpdf( + b.numpy(), alpha, scale=beta + ).sum() + sps.invgamma.logpdf(mode, alpha, scale=beta) * len(b) assert float(adp_sigd_math(b, a, s0)) == pytest.approx(expected, rel=1e-10) - def test_alpha_sets_log_width(self): - """std(log B) = sqrt(trigamma(alpha)), the bridge the design rests on. - - This is what lets alpha play the role the log-normal's sigma played, and - is the basis for reporting ``implied_std_log_adp``. - """ - from scipy import stats as sps - from scipy.special import polygamma - - for alpha in (3.5, 7.4): - draws = sps.invgamma.rvs(alpha, scale=100.0, size=400000, random_state=1) - assert np.log(draws).std() == pytest.approx( - np.sqrt(polygamma(1, alpha)), rel=2e-2 - ) - def test_monotonically_increasing_in_spread(self): """The loss must never reward spreading the B distribution out. @@ -577,6 +368,7 @@ def test_gradient_pushes_toward_the_mode(self): # R-factor Tests # ============================================================================= + @pytest.mark.unit class TestRfactorCalculations: """Test R-factor calculation functions.""" @@ -584,15 +376,15 @@ class TestRfactorCalculations: def test_get_rfactors_basic(self): """Test basic R-factor calculation.""" from torchref.base.math_torch import get_rfactors - + fobs = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0], dtype=torch.float32) fcalc = torch.tensor([1.1, 2.1, 3.1, 4.1, 5.1], dtype=torch.float32) - + # Create rfree mask (1 reflection in test set) rfree_mask = torch.tensor([True, True, True, True, False], dtype=torch.bool) - + r_work, r_free = get_rfactors(fobs, fcalc, rfree_mask) - + # Both should be small since fcalc is close to fobs assert r_work < 0.2 # r_free only has one reflection @@ -600,33 +392,33 @@ def test_get_rfactors_basic(self): def test_get_rfactors_perfect_fit(self): """Test R-factor with perfect fit.""" from torchref.base.math_torch import get_rfactors - + fobs = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0], dtype=torch.float32) fcalc = fobs.clone() # Perfect fit - + rfree_mask = torch.tensor([True, True, True, True, False], dtype=torch.bool) - + r_work, r_free = get_rfactors(fobs, fcalc, rfree_mask) - + assert r_work < 0.001 # Should be ~0 def test_bin_wise_rfactors(self): """Test bin-wise R-factor calculation.""" from torchref.base.math_torch import bin_wise_rfactors - + n_refl = 100 n_bins = 5 - + fobs = torch.rand(n_refl) + 1.0 fcalc = fobs * (1 + 0.1 * torch.randn(n_refl)) # Note: rfree=True means work set (not free set) rfree_mask = torch.rand(n_refl) > 0.1 - + # Ensure all bins are represented bins = torch.arange(n_refl) % n_bins - + r_work_bins, r_free_bins = bin_wise_rfactors(fobs, fcalc, rfree_mask, bins) - + # Should have results for each bin assert len(r_work_bins) == n_bins assert len(r_free_bins) == n_bins @@ -636,88 +428,7 @@ def test_bin_wise_rfactors(self): # Loss Function Tests # ============================================================================= -@pytest.mark.unit -class TestLossFunctions: - """Test individual loss functions from math_torch.""" - - def test_nll_xray(self): - """Test NLL X-ray loss function.""" - from torchref.base.math_torch import nll_xray - - fobs = torch.tensor([1.0, 2.0, 3.0], dtype=torch.float32) - fcalc = torch.tensor([1.1, 2.1, 3.1], dtype=torch.float32) - sigma = torch.tensor([0.1, 0.1, 0.1], dtype=torch.float32) - - loss = nll_xray(fobs, fcalc, sigma).mean() - - # NLL can be negative depending on normalization - assert torch.isfinite(loss) - - def test_least_squares_manual(self): - """Test least squares loss calculation.""" - # Manual least squares implementation - fobs = torch.tensor([1.0, 2.0, 3.0], dtype=torch.float32) - fcalc = torch.tensor([1.1, 2.1, 3.1], dtype=torch.float32) - sigma = torch.tensor([0.1, 0.1, 0.1], dtype=torch.float32) - - diff = fobs - fcalc - weights = 1.0 / (sigma ** 2) - loss = 0.5 * torch.sum(weights * (diff ** 2)) / len(fobs) - - assert torch.isfinite(loss) - assert loss >= 0 - - def test_nll_xray_with_mask(self): - """Test NLL X-ray with masking.""" - from torchref.base.math_torch import nll_xray - - fobs = torch.tensor([1.0, 2.0, 3.0, float('nan')], dtype=torch.float32) - fcalc = torch.tensor([1.1, 2.1, 3.1, 0.0], dtype=torch.float32) - sigma = torch.tensor([0.1, 0.1, 0.1, 0.1], dtype=torch.float32) - - # Only use finite values - valid = torch.isfinite(fobs) - loss = nll_xray(fobs[valid], fcalc[valid], sigma[valid]).mean() - - assert torch.isfinite(loss) - # ============================================================================= # Helper Function Tests # ============================================================================= - -@pytest.mark.unit -class TestTargetHelpers: - """Test helper functions used in targets.""" - - def test_distance_calculation(self): - """Test distance calculation between atom pairs.""" - xyz = torch.tensor([ - [0.0, 0.0, 0.0], - [3.0, 4.0, 0.0], # Distance = 5.0 - ], dtype=torch.float32) - - distance = torch.norm(xyz[1] - xyz[0]) - - assert torch.isclose(distance, torch.tensor(5.0)) - - def test_angle_from_vectors(self): - """Test angle calculation from vectors.""" - v1 = torch.tensor([1.0, 0.0, 0.0], dtype=torch.float32) - v2 = torch.tensor([0.0, 1.0, 0.0], dtype=torch.float32) - - cos_angle = torch.dot(v1, v2) / (torch.norm(v1) * torch.norm(v2)) - angle = torch.acos(cos_angle) - angle_deg = torch.rad2deg(angle) - - assert torch.isclose(angle_deg, torch.tensor(90.0)) - - def test_cross_product(self): - """Test cross product for normal vectors.""" - v1 = torch.tensor([1.0, 0.0, 0.0], dtype=torch.float32) - v2 = torch.tensor([0.0, 1.0, 0.0], dtype=torch.float32) - - normal = torch.linalg.cross(v1, v2) - - # Should be [0, 0, 1] - assert torch.allclose(normal, torch.tensor([0.0, 0.0, 1.0])) From b3f58713ddefe7b916ec1caa3dfa140f67550801 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 7 Sep 2026 17:03:59 +0200 Subject: [PATCH 4/5] test: consolidate reader contracts and exercise ModelFT caching CIF and MTZ integration tests own field shapes, crystal metadata, bin means and pair consistency; required assertions no longer depend on optional attributes. Existing CIF-to-PDB writing and device movement stay separate. Replace empty ModelFT state/forward/aniso checks with actual forward/cache checks, leaving restoration and anisotropic coverage in model unit tests. Remove unconsumed shared Model/ReflectionData fixtures. Validation: 13 default-MPS cases passed plus the corrected cache case; CPU float64 run 15 passed, 1 slow skip. Cache recomputation compares real magnitudes to allow backend reduction order. --- docs/changelog.rst | 1 + tests/README.md | 3 + tests/fixtures/README.md | 9 +- tests/fixtures/functional.py | 19 +- tests/functional/conftest.py | 6 +- tests/functional/test_io_functional.py | 213 ---------------- tests/functional/test_model_ft_functional.py | 197 +++------------ tests/functional/test_targets_functional.py | 62 ++--- tests/integration/test_io_cif.py | 118 +++------ tests/integration/test_io_reflections.py | 248 ++++++------------- 10 files changed, 180 insertions(+), 696 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index a8fe18e8..011122e3 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Consolidated CIF/MTZ loading contracts, checked configured tensor placement, and replaced ModelFT smoke checks with exercised forward-cache behavior. - Replaced local-arithmetic target tests with configured-device production-kernel checks on deposited coordinates and explicit least-squares expectations. - Consolidated weighting tests by API ownership and strengthened Gaussian-likelihood, gradient-norm, and cached-loss assertions. - Organized test fixtures into focused modules and reused module-scoped loaded objects for read-only functional checks while retaining fresh objects for mutation and loading tests. diff --git a/tests/README.md b/tests/README.md index 9e0ae0d6..9298352f 100644 --- a/tests/README.md +++ b/tests/README.md @@ -48,6 +48,9 @@ tests/ | Gaussian amplitude-metric values and reductions | `unit/base/test_loss.py` | | Restraint kernel values on deposited coordinates | `unit/base/test_target_values.py` | | Gradient RMS norm | `unit/utils/test_gradnorm.py` | +| CIF atomic fields and crystal metadata | `integration/test_io_cif.py` | +| MTZ fields, resolution bins and model/data crystal agreement | `integration/test_io_reflections.py` | +| ModelFT forward cache and grid integration | `functional/test_model_ft_functional.py` | | Numerical derivatives and backend parity | `unit/test_gradient_correctness.py`, `unit/structure_factor/` | A production call must participate in the assertion: computing a formula only in diff --git a/tests/fixtures/README.md b/tests/fixtures/README.md index da39ba78..7361c509 100644 --- a/tests/fixtures/README.md +++ b/tests/fixtures/README.md @@ -12,10 +12,10 @@ Keep a fixture in its test module when only that module needs it. | `precision.py` | Comparison tolerances and CPU-double reference context | All tests; reference fixture restores state after each test | | `objects.py` | Mutable models, data, scalers and restraints | All tests; fresh per test except explicitly shared bundles | | `numerical.py` | Synthetic tensors and factories | Imported only by `unit/conftest.py`; function | -| `functional.py` | `shared_model`, `shared_model_ft`, `shared_reflection_data` | Imported only by `functional/conftest.py`; module | +| `functional.py` | Read-only `shared_model_ft` | Imported only by `functional/conftest.py`; module | -Existing fixture names remain available without imports in tests. Import reusable -helpers from their defining module, never from the root `conftest.py`. Subtree +Fixtures are available without imports in tests. Import reusable helpers from +their defining module, never from the root `conftest.py`. Subtree conftests import fixture functions explicitly; register shared plugins only at the root so pytest also works when invoked from a subdirectory. @@ -25,7 +25,8 @@ parameters, tables, masks or grids, backpropagate through them, or use them in tests that switch global configuration. A target or scaler can mutate a model it borrows, so a shared model must not be passed to such an operation. -Tests that verify loading must invoke the loader themselves. Tests of mutation, +Tests that verify loading must execute a fresh loader, directly or through a +function-scoped fixture. Tests of mutation, device movement, or empty caches use fresh objects. `loaded_model`, `loaded_model_ft`, `loaded_reflection_data`, and their composed fixtures in `objects.py` provide fresh mutable objects per test. The explicitly shared session bundles in that diff --git a/tests/fixtures/functional.py b/tests/fixtures/functional.py index 18168242..5ee204a4 100644 --- a/tests/fixtures/functional.py +++ b/tests/fixtures/functional.py @@ -14,16 +14,7 @@ import pytest if TYPE_CHECKING: - from torchref.io import ReflectionData - from torchref.model import Model, ModelFT - - -@pytest.fixture(scope="module") -def shared_model(sample_cif_file: Path) -> Model: - """Load the sample CIF once per module for read-only atomic-model checks.""" - from torchref.model import Model - - return Model(verbose=0).load_cif(str(sample_cif_file)) + from torchref.model import ModelFT @pytest.fixture(scope="module") @@ -32,11 +23,3 @@ def shared_model_ft(sample_cif_file: Path) -> ModelFT: from torchref.model import ModelFT return ModelFT(max_res=2.0, verbose=0).load_cif(str(sample_cif_file)) - - -@pytest.fixture(scope="module") -def shared_reflection_data(sample_mtz_file: Path) -> ReflectionData: - """Load the sample MTZ once per module for read-only reflection checks.""" - from torchref.io import ReflectionData - - return ReflectionData().load_mtz(str(sample_mtz_file)) diff --git a/tests/functional/conftest.py b/tests/functional/conftest.py index e7a5b0ed..ba6caf3b 100644 --- a/tests/functional/conftest.py +++ b/tests/functional/conftest.py @@ -1,7 +1,3 @@ """Expose module-shared read-only fixtures to functional tests.""" -from tests.fixtures.functional import ( # noqa: F401 - shared_model, - shared_model_ft, - shared_reflection_data, -) +from tests.fixtures.functional import shared_model_ft # noqa: F401 diff --git a/tests/functional/test_io_functional.py b/tests/functional/test_io_functional.py index 8779ea39..49c82a99 100644 --- a/tests/functional/test_io_functional.py +++ b/tests/functional/test_io_functional.py @@ -5,7 +5,6 @@ """ import pytest -import torch class TestCIFReadingFunctional: @@ -31,32 +30,6 @@ def test_load_multiple_cif_files(self, cif_dir): assert model.cell is not None assert len(model.cell) == 6 - @pytest.mark.integration - def test_cif_atom_properties(self, shared_model): - """Test that atom properties are correctly loaded from CIF.""" - model = shared_model - - pdb = model.pdb - - # Check required columns exist - required_cols = ["x", "y", "z", "element", "resname", "chainid", "resseq"] - for col in required_cols: - assert col in pdb.columns or col.upper() in pdb.columns, ( - f"Missing column: {col}" - ) - - @pytest.mark.integration - def test_cif_element_types(self, shared_model): - """Test that element types are properly assigned.""" - model = shared_model - - elements = model.pdb["element"].unique() - - # Should have common protein elements - common_elements = ["C", "N", "O", "S"] - found_any = any(elem in elements for elem in common_elements) - assert found_any, "No common elements found" - class TestMTZReadingFunctional: """Functional tests for MTZ file reading.""" @@ -80,41 +53,6 @@ def test_load_multiple_mtz_files(self, mtz_dir): # Should have cell parameters assert data.cell is not None - @pytest.mark.integration - def test_mtz_data_properties(self, shared_reflection_data): - """Test that MTZ data properties are correctly loaded.""" - data = shared_reflection_data - - # Check HKL indices are integers or can be converted - hkl = data.hkl - assert hkl.shape[1] == 3, "HKL should have 3 columns" - - # Check F values are loaded - assert data.F is not None - assert data.F.shape[0] == hkl.shape[0] - - # Check sigma values - if hasattr(data, "F_sigma") and data.F_sigma is not None: - assert data.F_sigma.shape[0] == hkl.shape[0] - - @pytest.mark.integration - def test_mtz_resolution_range(self, shared_reflection_data): - """Test that resolution range is computed correctly.""" - data = shared_reflection_data - - # Check if resolution data is available - if hasattr(data, "d") and data.d is not None: - d_min = data.d.min().item() - d_max = data.d.max().item() - - # Resolution should be positive - assert d_min > 0 - assert d_max > d_min - - # Typical protein data: 0.8 - 500 Å - assert d_min > 0.5 - assert d_max < 1000 - class TestSFCIFReadingFunctional: """Functional tests for structure factor CIF reading.""" @@ -139,154 +77,3 @@ def test_load_sf_cif(self, cif_sf_dir): except Exception as e: # Some files may not be valid SF-CIF format pass - - -class TestDataConsistencyFunctional: - """Test consistency between model and data files.""" - - @pytest.mark.integration - def test_cell_parameters_match(self, sample_structure_pair): - """Test that cell parameters match between model and reflections.""" - from torchref.io import ReflectionData - from torchref.model.model import Model - - model = Model() - model.load_cif(str(sample_structure_pair["model"])) - - data = ReflectionData() - data.load_mtz(str(sample_structure_pair["reflections"])) - - model_cell = model.cell - data_cell = data.cell - - if model_cell is not None and data_cell is not None: - # Convert to tensors if needed - if not isinstance(model_cell, torch.Tensor): - model_cell = torch.tensor(model_cell) - if not isinstance(data_cell, torch.Tensor): - data_cell = torch.tensor(data_cell) - - # Cell parameters should be similar (1% tolerance) - assert torch.allclose( - model_cell.float(), data_cell.float(), rtol=0.01, atol=0.1 - ) - - @pytest.mark.integration - def test_spacegroup_consistency(self, sample_structure_pair): - """Test that spacegroup is consistent.""" - from torchref.io import ReflectionData - from torchref.model.model import Model - - model = Model() - model.load_cif(str(sample_structure_pair["model"])) - - data = ReflectionData() - data.load_mtz(str(sample_structure_pair["reflections"])) - - # Both should have spacegroup defined - assert model.spacegroup is not None - - -class TestDataBinningFunctional: - """Test data binning operations.""" - - @pytest.mark.integration - def test_get_bins(self, sample_mtz_file): - """Test resolution binning of reflection data.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - # Get bins - bins, n_bins = data.get_bins(n_bins=10) - - assert bins is not None - assert bins.shape[0] == data.hkl.shape[0] - assert bins.min() >= 0 - assert bins.max() < n_bins - - @pytest.mark.integration - def test_mean_res_per_bin(self, sample_mtz_file): - """Test mean resolution per bin calculation.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - # Get bins first - bins, n_bins = data.get_bins(n_bins=10) - - # Get mean resolution per bin - if hasattr(data, "mean_res_per_bin"): - mean_res = data.mean_res_per_bin() - - assert mean_res is not None - assert len(mean_res) == n_bins - - # Mean resolution should decrease with bin index (low res to high res) - # or increase (high res to low res) - depends on implementation - assert torch.all(torch.isfinite(mean_res)) - - -class TestFrenchWilsonFunctional: - """Test French-Wilson conversion with real data.""" - - @pytest.mark.integration - def test_french_wilson_applied(self, shared_reflection_data): - """Test that French-Wilson conversion is applied.""" - data = shared_reflection_data - - # After French-Wilson, F values should be non-negative - valid_F = data.F[~torch.isnan(data.F)] - - if len(valid_F) > 0: - # All valid F values should be >= 0 - assert torch.all(valid_F >= 0) - - -class TestRfreeHandlingFunctional: - """Test R-free flag handling.""" - - @pytest.mark.integration - def test_rfree_flags_loaded(self, shared_reflection_data): - """Test that R-free flags are loaded or generated.""" - data = shared_reflection_data - - # Should have rfree attribute - if hasattr(data, "rfree") and data.rfree is not None: - assert data.rfree.shape[0] == data.hkl.shape[0] - - # Should be boolean or can be converted to boolean - assert data.rfree.dtype == torch.bool or torch.all( - (data.rfree == 0) | (data.rfree == 1) - ) - - @pytest.mark.integration - def test_rfree_fraction(self, shared_reflection_data): - """Test R-free set fraction is reasonable.""" - data = shared_reflection_data - - if hasattr(data, "rfree") and data.rfree is not None: - # Work set mask (True for work, False for test) - work_fraction = data.rfree.float().mean().item() - - # Typically 90-95% work set, 5-10% test set - # So work_fraction should be 0.9-0.95 typically - assert 0.7 < work_fraction <= 1.0 - - -class TestMaskHandlingFunctional: - """Test reflection mask handling.""" - - @pytest.mark.integration - def test_masks_method(self, shared_reflection_data): - """Test masks() method returns valid mask.""" - data = shared_reflection_data - - if hasattr(data, "masks"): - mask = data.masks() - - assert mask is not None - assert mask.shape[0] == data.hkl.shape[0] - assert mask.dtype == torch.bool diff --git a/tests/functional/test_model_ft_functional.py b/tests/functional/test_model_ft_functional.py index a6abd808..2c6b9b4c 100644 --- a/tests/functional/test_model_ft_functional.py +++ b/tests/functional/test_model_ft_functional.py @@ -28,19 +28,6 @@ def test_modelft_with_custom_resolution(self): model = ModelFT(max_res=1.5) assert model.max_res == 1.5 - def test_modelft_load_cif(self, sample_cif_file): - """Test loading a CIF file into ModelFT.""" - from torchref.model.model_ft import ModelFT - - model = ModelFT(max_res=2.0, verbose=0) - model.load_cif(str(sample_cif_file)) - - # Verify basic properties - assert model.xyz() is not None - assert model.xyz().shape[0] > 0 - assert model.cell is not None - assert len(model.cell) == 6 - def test_modelft_has_gridsize(self, sample_cif_file): """The grid resolves from the loaded cell and space group on first read.""" from torchref.model.model_ft import ModelFT @@ -52,29 +39,10 @@ def test_modelft_has_gridsize(self, sample_cif_file): assert model.gridsize is not None assert len(model.gridsize) == 3 assert all(g > 0 for g in model.gridsize) - - -@pytest.mark.integration -class TestModelFTParametrization: - """Test ModelFT parametrization with real structures.""" - - def test_parametrization_built(self, shared_model_ft): - """Test that parametrization is built after loading.""" - - model = shared_model_ft - - # Parametrization should be set - assert model.parametrization is not None - - def test_scattering_factors_available(self, shared_model_ft): - """Test that scattering factors can be computed.""" - - model = shared_model_ft - - # Should be able to access atom properties - xyz = model.xyz() - assert xyz is not None - assert xyz.dtype == torch.float32 or xyz.dtype == torch.float64 + assert model.xyz().shape[0] > 0 + assert model.parametrization + assert model.adp().shape == (len(model.xyz()),) + assert torch.all(model.adp() >= 0) @pytest.mark.integration @@ -104,14 +72,15 @@ def test_get_real_space_grid(self, loaded_model_ft): """Test getting real space grid.""" from torchref.base.math_torch import get_real_grid - # The grid helper moves its Cell in place when targeting CPU. model = loaded_model_ft assert model.gridsize is not None - grid = get_real_grid(model.cell, max_res=2.0, device="cpu") + grid = get_real_grid(model.cell, max_res=2.0, device=model.device) assert grid is not None assert len(grid.shape) == 4 # Should be 4D (nx, ny, nz, 3) + assert grid.device == model.xyz().device + assert grid.dtype == model.xyz().dtype @pytest.mark.integration @@ -135,60 +104,6 @@ def test_map_symmetry_available(self, shared_model_ft): assert operator.map_shape == gridsize -@pytest.mark.integration -class TestModelFTStateDictFunctional: - """Test ModelFT state dict operations with real data.""" - - def test_save_and_load_state_dict(self, loaded_model_ft, tmp_path): - """Test saving and loading state dict.""" - from torchref.model.model_ft import ModelFT - - model = loaded_model_ft - - original_xyz = model.xyz().clone() - - # Save state dict - state_dict = model.state_dict() - - # Create new model and load state - model2 = ModelFT(max_res=2.0, verbose=0) - - # We need to ensure proper initialization - # For now just verify state_dict works - assert state_dict is not None - assert len(state_dict) > 0 - - -@pytest.mark.integration -class TestModelFTForwardPass: - """Test ModelFT forward pass (structure factor calculation).""" - - def test_forward_method_exists(self, shared_model_ft): - """Test that forward method is available.""" - - model = shared_model_ft - - # Check forward method exists - assert hasattr(model, "forward") - - def test_build_map_method(self, sample_cif_file): - """Test build_map method if available.""" - from torchref.model.model_ft import ModelFT - - model = ModelFT(max_res=3.0, verbose=0) # Lower res for faster test - model.load_cif(str(sample_cif_file)) - - # Check build_map method - if hasattr(model, "build_map"): - # Try to build map - try: - model.build_map() - assert model.map is not None - except Exception as e: - # May fail if missing dependencies - pytest.skip(f"build_map not available: {e}") - - @pytest.mark.integration class TestModelFTMultipleStructures: """Test ModelFT with multiple structures.""" @@ -215,50 +130,10 @@ def test_modelft_multiple_structures(self, all_structure_pairs): assert tested >= 1, "At least one structure should load" -@pytest.mark.integration -class TestModelFTCaching: - """Test ModelFT caching mechanism.""" - - def test_cache_initialization(self, loaded_model_ft): - """Test that CachedForwardMixin cache starts empty.""" - - model = loaded_model_ft - - # Mixin cache should start empty (lazily initialized) - assert getattr(model, "_fwd_cached_output", None) is None - - def test_cache_usage(self, shared_model_ft): - """Test that cache can be used for computations.""" - - model = shared_model_ft - - # Access xyz twice - should use caching - xyz1 = model.xyz() - xyz2 = model.xyz() - - # Should return same tensor - assert torch.allclose(xyz1, xyz2) - - @pytest.mark.integration class TestModelFTCoordinateOperations: """Test ModelFT coordinate operations.""" - def test_cartesian_to_fractional(self, shared_model_ft): - """Test coordinate conversion.""" - from torchref.base.math_torch import cartesian_to_fractional_torch - - model = shared_model_ft - - xyz = model.xyz() - cell = model.cell - - # Convert to fractional - frac = cartesian_to_fractional_torch(xyz, cell.data) - - # Fractional coords should be bounded (mostly between 0 and 1) - assert frac.shape == xyz.shape - def test_fractional_to_cartesian(self, shared_model_ft): """Test fractional to cartesian conversion.""" from torchref.base.math_torch import ( @@ -273,6 +148,7 @@ def test_fractional_to_cartesian(self, shared_model_ft): # Round trip conversion frac = cartesian_to_fractional_torch(xyz, cell.data) + assert frac.shape == xyz.shape xyz_back = fractional_to_cartesian_torch(frac, cell.data) # Should get back original coordinates (float32 roundtrip) @@ -280,28 +156,35 @@ def test_fractional_to_cartesian(self, shared_model_ft): @pytest.mark.integration -class TestModelFTAnisoHandling: - """Test ModelFT handling of anisotropic parameters.""" - - def test_access_aniso_atoms(self, shared_model_ft): - """Test accessing anisotropic atom information.""" - - model = shared_model_ft - - # Check if aniso is available - if hasattr(model, "get_aniso") or hasattr(model, "aniso"): - # Structure has aniso - pass - - def test_isotropic_b_factors(self, shared_model_ft): - """Test accessing isotropic B-factors.""" - - model = shared_model_ft - - # Get B-factors (now accessed via adp()) - b_factors = model.adp() - - assert b_factors is not None - assert b_factors.shape[0] == model.xyz().shape[0] - # B-factors should be positive - assert torch.all(b_factors > 0) or torch.all(b_factors >= 0) +def test_forward_cache_contract( + loaded_model_ft, loaded_reflection_data, monkeypatch +) -> None: + """A model computes complex structure factors and caches only until invalidation.""" + from unittest.mock import Mock + + from torchref.config import caching, get_complex_dtype + + model = loaded_model_ft + hkl = loaded_reflection_data.hkl[:32] + monkeypatch.setattr(caching, "value", True) + forward = Mock(wraps=model.forward) + monkeypatch.setattr(model, "forward", forward) + assert getattr(model, "_fwd_cached_output", None) is None + + first = model(hkl) + assert first.shape == (len(hkl),) + assert first.dtype == get_complex_dtype() + assert first.device == hkl.device + assert torch.isfinite(first).all() + assert first.abs().sum() > 0 + assert model(hkl) is first + assert forward.call_count == 1 + + refreshed = model(hkl, recalc=True) + assert forward.call_count == 2 + assert refreshed is not first + # Accelerator reductions need not repeat bit-for-bit after recomputation. + relative_error = torch.linalg.vector_norm( + (refreshed - first).abs() + ) / torch.linalg.vector_norm(first.abs()) + assert relative_error < 256 * torch.finfo(first.real.dtype).eps diff --git a/tests/functional/test_targets_functional.py b/tests/functional/test_targets_functional.py index 5eb505e1..c3f2f70d 100644 --- a/tests/functional/test_targets_functional.py +++ b/tests/functional/test_targets_functional.py @@ -4,9 +4,9 @@ Tests target functions with real model and data objects. """ +import numpy as np import pytest import torch -import numpy as np class TestXrayTargetsFunctional: @@ -15,9 +15,9 @@ class TestXrayTargetsFunctional: @pytest.mark.integration def test_gaussian_nll_with_real_data(self, sample_structure_pair): """Test Gaussian NLL calculation with real reflection data.""" - from torchref.model.model import Model - from torchref.io import ReflectionData from torchref.base.math_torch import nll_xray + from torchref.io import ReflectionData + from torchref.model.model import Model model = Model() model.load_cif(str(sample_structure_pair["model"])) @@ -43,8 +43,8 @@ def test_gaussian_nll_with_real_data(self, sample_structure_pair): @pytest.mark.integration def test_least_squares_with_real_data(self, sample_structure_pair): """Test least squares calculation with real data.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model model = Model() model.load_cif(str(sample_structure_pair["model"])) @@ -74,9 +74,9 @@ class TestRfactorCalculationsFunctional: @pytest.mark.integration def test_rfactor_with_real_data(self, sample_structure_pair): """Test R-factor calculation with real reflection data.""" - from torchref.model.model import Model - from torchref.io import ReflectionData from torchref.base.math_torch import get_rfactors + from torchref.io import ReflectionData + from torchref.model.model import Model model = Model() model.load_cif(str(sample_structure_pair["model"])) @@ -110,9 +110,9 @@ def test_rfactor_with_real_data(self, sample_structure_pair): @pytest.mark.integration def test_bin_wise_rfactors(self, sample_structure_pair): """Test bin-wise R-factor calculation.""" - from torchref.model.model import Model - from torchref.io import ReflectionData from torchref.base.math_torch import bin_wise_rfactors + from torchref.io import ReflectionData + from torchref.model.model import Model model = Model() model.load_cif(str(sample_structure_pair["model"])) @@ -234,27 +234,6 @@ def test_angle_target_with_real_structure(self, sample_cif_file, external_monome assert torch.isfinite(loss) -class TestStructureFactorCalculationFunctional: - """Functional tests for structure factor calculation.""" - - @pytest.mark.integration - def test_fcalc_shape_matches_data(self, sample_structure_pair): - """Test that calculated structure factors have correct shape.""" - from torchref.model.model import Model - from torchref.io import ReflectionData - - model = Model() - model.load_cif(str(sample_structure_pair["model"])) - - data = ReflectionData() - data.load_mtz(str(sample_structure_pair["reflections"])) - - # Check if model has fcalc calculation method - if hasattr(model, 'calc_fcalc'): - fcalc = model.calc_fcalc(data) - - # Fcalc should have same number of reflections as data - assert fcalc.shape[0] == data.hkl.shape[0] class TestScalingWithRealData: @@ -263,8 +242,8 @@ class TestScalingWithRealData: @pytest.mark.integration def test_scaler_initialization_with_real_data(self, sample_structure_pair): """Test scaler initialization with real model and data.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler model = Model() @@ -285,8 +264,8 @@ def test_scaler_initialization_with_real_data(self, sample_structure_pair): @pytest.mark.integration def test_anisotropy_correction_values(self, sample_structure_pair): """Test that anisotropy correction produces reasonable values.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler model = Model() @@ -315,8 +294,8 @@ class TestMathFunctionsFunctional: @pytest.mark.integration def test_scattering_vectors_from_real_data(self, sample_structure_pair): """Test scattering vector calculation with real HKL and cell.""" - from torchref.io import ReflectionData from torchref.base.math_torch import get_scattering_vectors + from torchref.io import ReflectionData data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) @@ -333,11 +312,11 @@ def test_scattering_vectors_from_real_data(self, sample_structure_pair): @pytest.mark.integration def test_coordinate_transformations_with_real_cell(self, sample_cif_file): """Test coordinate transformations with real unit cell.""" - from torchref.model.model import Model from torchref.base.math_torch import ( cartesian_to_fractional_torch, - fractional_to_cartesian_torch + fractional_to_cartesian_torch, ) + from torchref.model.model import Model model = Model() model.load_cif(str(sample_cif_file)) @@ -439,8 +418,8 @@ class TestNLLFunctionsFunctional: @pytest.mark.integration def test_nll_xray_with_identical_data(self, sample_structure_pair): """Test NLL is minimal when Fobs equals Fcalc.""" - from torchref.io import ReflectionData from torchref.base.math_torch import nll_xray + from torchref.io import ReflectionData data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) @@ -463,8 +442,8 @@ def test_nll_xray_with_identical_data(self, sample_structure_pair): @pytest.mark.integration def test_nll_xray_increases_with_error(self, sample_structure_pair): """Test NLL increases as Fcalc differs from Fobs.""" - from torchref.io import ReflectionData from torchref.base.math_torch import nll_xray + from torchref.io import ReflectionData data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) @@ -494,8 +473,8 @@ def test_nll_xray_increases_with_error(self, sample_structure_pair): @pytest.mark.integration def test_nll_xray_lognormal(self, sample_structure_pair): """Test lognormal NLL calculation.""" - from torchref.io import ReflectionData from torchref.base.math_torch import nll_xray_lognormal + from torchref.io import ReflectionData data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) @@ -520,8 +499,9 @@ class TestRiceDistributionFunctional: @pytest.mark.integration def test_rice_nll_acentric(self, sample_structure_pair): """Test Rice NLL for acentric reflections.""" - from torchref.io import ReflectionData from torch.special import i0 + + from torchref.io import ReflectionData data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) @@ -592,8 +572,8 @@ def test_sigma_weighting(self, sample_structure_pair): @pytest.mark.integration def test_resolution_weighting(self, sample_structure_pair): """Test resolution-based weighting.""" - from torchref.io import ReflectionData from torchref.base.math_torch import get_scattering_vectors + from torchref.io import ReflectionData data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) @@ -704,9 +684,9 @@ class TestCombinedLossFunctional: @pytest.mark.integration def test_xray_plus_geometry_loss(self, sample_structure_pair, external_monomer_library): """Test combining X-ray and geometry losses.""" - from torchref.model.model import Model - from torchref.io import ReflectionData from torchref.base.math_torch import nll_xray + from torchref.io import ReflectionData + from torchref.model.model import Model model = Model() model.load_cif(str(sample_structure_pair["model"])) diff --git a/tests/integration/test_io_cif.py b/tests/integration/test_io_cif.py index ff288cc5..7e8b23d7 100644 --- a/tests/integration/test_io_cif.py +++ b/tests/integration/test_io_cif.py @@ -6,85 +6,37 @@ import pytest import torch -from pathlib import Path - -class TestCIFLoading: - """Tests for loading CIF model files.""" - - @pytest.mark.integration - def test_load_model_cif(self, sample_cif_file): - """Test loading a real CIF model file.""" - from torchref.model.model import Model - - model = Model() - model.load_cif(str(sample_cif_file)) - - # Basic checks - use xyz().shape[0] for atom count - n_atoms = model.xyz().shape[0] - assert n_atoms > 0 - assert hasattr(model, 'xyz') - assert hasattr(model, 'adp') - assert hasattr(model, 'occupancy') - - @pytest.mark.integration - def test_model_atom_counts(self, sample_cif_file): - """Test that model has consistent atom counts.""" - from torchref.model.model import Model - - model = Model() - model.load_cif(str(sample_cif_file)) - - # All arrays should have same number of atoms - n_atoms = model.xyz().shape[0] - assert model.xyz().shape[0] == n_atoms - assert model.adp().shape[0] == n_atoms - assert model.occupancy().shape[0] == n_atoms - - @pytest.mark.integration - def test_model_cell_parameters(self, sample_cif_file): - """Test that model has valid cell parameters.""" - from torchref.model.model import Model - - model = Model() - model.load_cif(str(sample_cif_file)) - - # Cell should have 6 parameters - assert len(model.cell) == 6 - # All cell parameters should be positive - assert all(p > 0 for p in model.cell[:3].tolist()) # a, b, c - # Angles should be reasonable (0-180) - assert all(0 < p <= 180 for p in model.cell[3:].tolist()) # alpha, beta, gamma - - @pytest.mark.integration - def test_model_spacegroup(self, sample_cif_file): - """Test that model has a valid spacegroup.""" - from torchref.model.model import Model - - model = Model() - model.load_cif(str(sample_cif_file)) - - # Spacegroup should be set (can be string or gemmi.SpaceGroup) - assert model.spacegroup is not None - # Check it can be converted to string representation - assert len(str(model.spacegroup)) > 0 - - @pytest.mark.integration - def test_model_element_types(self, sample_cif_file): - """Test that element types are recognized.""" - from torchref.model.model import Model - - model = Model() - model.load_cif(str(sample_cif_file)) - - # Should have pdb DataFrame with element column - assert hasattr(model, 'pdb') - assert 'element' in model.pdb.columns - # Elements should be strings like 'C', 'N', 'O', etc. - elements = set(model.pdb['element'].unique()) - common_elements = {'C', 'N', 'O', 'S', 'H', 'CA', 'MG', 'ZN', 'FE'} - # At least some elements should be recognized - assert len(elements.intersection(common_elements)) > 0 or len(elements) > 0 +from torchref.config import canonical_device, get_default_device, get_float_dtype + + +@pytest.mark.integration +def test_cif_loading_contract(loaded_model, sample_cif_file) -> None: + """A deposited CIF supplies aligned atomic tensors and its crystal metadata.""" + import gemmi + + model = loaded_model + reference = gemmi.read_structure(str(sample_cif_file)) + xyz, adp, occupancy = model.xyz(), model.adp(), model.occupancy() + assert xyz.shape == (len(model.pdb), 3) + assert len(xyz) > 0 + assert adp.shape == occupancy.shape == (len(xyz),) + for tensor in (xyz, adp, occupancy, model.cell.data): + assert tensor.dtype == get_float_dtype() + assert canonical_device(tensor.device) == canonical_device(get_default_device()) + assert torch.isfinite(tensor).all() + assert torch.all(adp >= 0) + assert {"x", "y", "z", "element", "resname", "chainid", "resseq"} <= set( + model.pdb.columns + ) + assert {"C", "N", "O"} <= set(model.pdb.element) + torch.testing.assert_close( + model.cell.data, xyz.new_tensor(reference.cell.parameters) + ) + assert ( + model.spacegroup.number + == gemmi.find_spacegroup_by_name(reference.spacegroup_hm).number + ) class TestMultipleCIFFiles: @@ -125,18 +77,18 @@ class TestCIFSaving: def test_save_and_reload_cif(self, sample_cif_file, tmp_path): """Test saving a model to CIF and reloading it.""" from torchref.model.model import Model - + # Load original model1 = Model() model1.load_cif(str(sample_cif_file)) n_atoms1 = model1.xyz().shape[0] - + # Save to temp file using write_pdb (CIF saving may not exist) output_path = tmp_path / "test_output.pdb" model1.write_pdb(str(output_path)) - + assert output_path.exists() - + # add_hydrogens=False on reload: what is under test is whether the written # file round-trips, not whether generation reruns. Regenerating on reload can # legitimately differ, because ``write_pdb`` does not emit LINK records -- so a @@ -145,6 +97,6 @@ def test_save_and_reload_cif(self, sample_cif_file, tmp_path): model2 = Model(add_hydrogens=False) model2.load_pdb(str(output_path)) n_atoms2 = model2.xyz().shape[0] - + # Compare atom counts assert n_atoms2 == n_atoms1 diff --git a/tests/integration/test_io_reflections.py b/tests/integration/test_io_reflections.py index 79642253..cb9e9a1b 100644 --- a/tests/integration/test_io_reflections.py +++ b/tests/integration/test_io_reflections.py @@ -6,106 +6,72 @@ import pytest import torch -from pathlib import Path - -class TestMTZLoading: - """Tests for loading MTZ reflection files.""" - - @pytest.mark.integration - def test_load_mtz_file(self, sample_mtz_file): - """Test loading a real MTZ file.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - # Should have reflections loaded - assert hasattr(data, 'hkl') - assert hasattr(data, 'F') - assert data.hkl is not None - - @pytest.mark.integration - def test_mtz_reflection_counts(self, sample_mtz_file): - """Test that MTZ has consistent reflection counts.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - n_refl = data.hkl.shape[0] - assert n_refl > 0 - - # F should match hkl count - if data.F is not None: - assert data.F.shape[0] == n_refl - - @pytest.mark.integration - def test_mtz_hkl_indices(self, sample_mtz_file): - """Test HKL indices are valid integers.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - # HKL should have 3 columns - assert data.hkl.shape[1] == 3 - - # Should contain integer-like values - hkl_rounded = torch.round(data.hkl) - assert torch.allclose(data.hkl, hkl_rounded) - - @pytest.mark.integration - def test_mtz_cell_parameters(self, sample_mtz_file): - """Test that MTZ has valid cell parameters.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - if hasattr(data, 'cell') and data.cell is not None: - assert len(data.cell) == 6 - assert all(c > 0 for c in data.cell[:3].tolist()) - - @pytest.mark.integration - def test_mtz_spacegroup(self, sample_mtz_file): - """Test that MTZ has a valid spacegroup.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - # Spacegroup should be set (can be string or gemmi.SpaceGroup) - assert data.spacegroup is not None - # Check it can be converted to string representation - assert len(str(data.spacegroup)) > 0 - - @pytest.mark.integration - def test_mtz_sigma_values(self, sample_mtz_file): - """Test that sigma values are loaded.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - if hasattr(data, 'F_sigma') and data.F_sigma is not None: - assert data.F_sigma.shape[0] == data.F.shape[0] - # Check that non-NaN sigma values are positive - valid_sigma = data.F_sigma[~torch.isnan(data.F_sigma)] - if len(valid_sigma) > 0: - assert torch.all(valid_sigma > 0) - - @pytest.mark.integration - def test_mtz_rfree_flags(self, sample_mtz_file): - """Test that R-free flags are loaded or generated.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - # Should have rfree_flags (loaded or generated) - if hasattr(data, 'rfree_flags') and data.rfree_flags is not None: - assert data.rfree_flags.shape[0] == data.hkl.shape[0] +from torchref.config import ( + canonical_device, + get_default_device, + get_float_dtype, + get_int_dtype, +) + + +@pytest.mark.integration +def test_mtz_loading_contract(loaded_reflection_data, sample_mtz_file) -> None: + """MTZ loading supplies aligned observations, masks and crystal metadata.""" + import gemmi + + data = loaded_reflection_data + reference = gemmi.read_mtz_file(str(sample_mtz_file)) + n = len(data.hkl) + assert n > 0 + assert data.hkl.shape == (n, 3) + assert data.hkl.dtype == get_int_dtype() + assert canonical_device(data.hkl.device) == canonical_device(get_default_device()) + for tensor in (data.F, data.F_sigma, data.resolution): + assert tensor.shape == (n,) + assert tensor.dtype == get_float_dtype() + assert canonical_device(tensor.device) == canonical_device(get_default_device()) + mask = data.masks() + assert mask.shape == (n,) + assert mask.dtype == torch.bool + assert mask.any() + assert torch.isfinite(data.F[mask]).all() + assert torch.all(data.F[mask] >= 0) + assert torch.isfinite(data.F_sigma[mask]).all() + assert torch.all(data.F_sigma[mask] > 0) + assert torch.isfinite(data.resolution).all() + assert torch.all(data.resolution > 0) + assert data.rfree_flags.shape == (n,) + assert data.rfree_flags.dtype == torch.bool + assert data.rfree_flags.any() and (~data.rfree_flags).any() + assert 0.7 < data.rfree_flags.to(get_float_dtype()).mean().item() < 1.0 + torch.testing.assert_close( + data.cell.data, data.F.new_tensor(reference.cell.parameters) + ) + assert data.spacegroup.number == reference.spacegroup.number + + +@pytest.mark.integration +def test_resolution_bins(loaded_reflection_data) -> None: + """Every bin mean equals the mean d-spacing of its unmasked reflections.""" + data = loaded_reflection_data + bins, n_bins = data.get_bins(n_bins=10) + assert bins.shape == (len(data.hkl),) + assert n_bins > 0 + assert bins.min() >= 0 and bins.max() < n_bins + groups = [(bins == i) & data.masks() for i in range(n_bins)] + assert all(group.any() for group in groups) + expected = torch.stack([data.resolution[group].mean() for group in groups]) + torch.testing.assert_close(data.mean_res_per_bin(), expected) + + +@pytest.mark.integration +def test_structure_pair_consistency(model_and_data) -> None: + """Matching model and reflection files describe the same crystal.""" + model, data = model_and_data["model"], model_and_data["data"] + assert len(model.xyz()) > 0 and len(data.hkl) > 0 + torch.testing.assert_close(model.cell.data, data.cell.data, rtol=0.01, atol=0.1) + assert model.spacegroup.number == data.spacegroup.number class TestSFCIFLoading: @@ -115,95 +81,27 @@ class TestSFCIFLoading: def test_load_sf_cif(self, sample_structure_factor_cif): """Test loading a structure factor CIF file.""" from torchref.io import ReflectionData - + data = ReflectionData() data.load_cif(str(sample_structure_factor_cif)) - - assert data.hkl is not None + + assert data.hkl.shape[0] > 0 + assert data.hkl.shape[1] == 3 class TestReflectionDataProperties: """Tests for computed properties of reflection data.""" - @pytest.mark.integration - def test_resolution_calculation(self, sample_mtz_file): - """Test resolution can be calculated from loaded data.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - # Should have resolution attribute - if hasattr(data, 'resolution') and data.resolution is not None: - assert torch.all(data.resolution > 0) - assert torch.all(torch.isfinite(data.resolution)) - - @pytest.mark.integration - def test_wilson_b_factor(self, sample_mtz_file): - """Test Wilson B-factor is calculated.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - # Wilson B should be calculated during loading - if hasattr(data, 'wilson_b') and data.wilson_b is not None: - assert data.wilson_b > 0 - @pytest.mark.integration def test_data_device_movement(self, sample_mtz_file, cpu_device): """Test moving reflection data to different devices.""" from torchref.io import ReflectionData - + data = ReflectionData() data.load_mtz(str(sample_mtz_file)) - + # Move to device data = data.to(cpu_device) - - # Tensors should be on correct device - if data.hkl is not None: - assert data.hkl.device == cpu_device - if data.F is not None: - assert data.F.device == cpu_device - - -class TestMatchingDataPairs: - """Tests using matching model and reflection data.""" - @pytest.mark.integration - def test_load_structure_pair(self, sample_structure_pair): - """Test loading matching model and reflection data.""" - from torchref.model.model import Model - from torchref.io import ReflectionData - - model = Model() - model.load_cif(str(sample_structure_pair["model"])) - - data = ReflectionData() - data.load_mtz(str(sample_structure_pair["reflections"])) - - # Both should load successfully - n_atoms = model.xyz().shape[0] - assert n_atoms > 0 - assert data.hkl is not None - - @pytest.mark.integration - def test_cell_consistency(self, sample_structure_pair): - """Test that model and data have consistent cell parameters.""" - from torchref.model.model import Model - from torchref.io import ReflectionData - - model = Model() - model.load_cif(str(sample_structure_pair["model"])) - - data = ReflectionData() - data.load_mtz(str(sample_structure_pair["reflections"])) - - # Cell parameters should be similar (may have small differences) - if hasattr(data, 'cell') and data.cell is not None: - model_cell = torch.tensor(model.cell) - data_cell = torch.tensor(data.cell) - - # Allow 1% tolerance for cell parameters - assert torch.allclose(model_cell, data_cell, rtol=0.01, atol=0.1) + assert data.hkl.device == cpu_device + assert data.F.device == cpu_device From fb97d23438e7f8127e3c2419270e6d7936cf4bc8 Mon Sep 17 00:00:00 2001 From: HatPdotS Date: Mon, 7 Sep 2026 17:12:03 +0200 Subject: [PATCH 5/5] test: make extended structure coverage explicit and lazy Replace CIF/MTZ/SF-CIF loops with named per-file compatibility cases and an input inventory guard. Keep 1DAW reader contracts in the quick suite; require --run-slow for extras. Replace eager all_test_structures with one fresh named crystal per scaler/restraint case. Move the extra ModelFT loading cases and symmetry file sweep coverage into the compatibility panel. Merge space-group name cases under their unit owner without dropping parameter variants. Validation: 44 extended cases passed; final regression 477 passed/70 skipped, including 41 explicitly slow cases. Full collection: 2587 cases. No production files changed. --- docs/changelog.rst | 1 + docs/user_guide/testing.rst | 26 +- tests/README.md | 9 + tests/fixtures/README.md | 7 +- tests/fixtures/files.py | 43 ++-- tests/fixtures/objects.py | 61 ++--- tests/functional/test_io_functional.py | 79 ------ tests/functional/test_model_ft_functional.py | 26 -- .../functional/test_restraints_functional.py | 157 +++++------- tests/functional/test_scaler_functional.py | 232 ++++++++---------- tests/helpers/structure_cases.py | 27 ++ tests/integration/test_io_cif.py | 31 --- .../test_structure_compatibility.py | 101 ++++++++ .../integration/test_symmetry_integration.py | 41 +--- tests/unit/symmetry/test_symmetry.py | 29 ++- 15 files changed, 398 insertions(+), 472 deletions(-) delete mode 100644 tests/functional/test_io_functional.py create mode 100644 tests/helpers/structure_cases.py create mode 100644 tests/integration/test_structure_compatibility.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 011122e3..3cb46505 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,7 @@ Changelog Unreleased ---------- +- Named broad structure-compatibility cases explicitly, moved extra datasets to the slow tier, and removed eager all-model loading and swallowed reader failures. - Consolidated CIF/MTZ loading contracts, checked configured tensor placement, and replaced ModelFT smoke checks with exercised forward-cache behavior. - Replaced local-arithmetic target tests with configured-device production-kernel checks on deposited coordinates and explicit least-squares expectations. - Consolidated weighting tests by API ownership and strengthened Gaussian-likelihood, gradient-norm, and cached-loss assertions. diff --git a/docs/user_guide/testing.rst b/docs/user_guide/testing.rst index cbdc1c8c..66e60273 100644 --- a/docs/user_guide/testing.rst +++ b/docs/user_guide/testing.rst @@ -39,7 +39,13 @@ PDB ID d_min (Å) Space group ``tests/files/`` also holds partial sets — ``1AK5_with_H.pdb`` + ``1AK5.mtz`` (no CIF), ``7L84.pdb`` + ``7L84-sf.cif`` (no MTZ), ``test_ihm_ensemble.cif`` — so a test that globs one directory and assumes a matching file in another will -fail on those. Use ``sample_structure_pair`` / ``all_test_structures``. +fail on those. Use ``sample_structure_pair`` for the quick reference crystal, +or ``compatibility_structure_pair`` for named extended cases. The latter carries +the ``slow`` marker and selects paths without loading objects. + +``tests/helpers/structure_cases.py`` assigns bundled CIF, MTZ and SF-CIF files to +the quick or extended compatibility panel. Additional files require an explicit +assignment; directory growth does not silently expand numerical test work. Running Tests ------------- @@ -99,11 +105,11 @@ The Amber stack, if you want it: Fixtures -------- -Almost everything lives in the root ``tests/conftest.py`` and is therefore -available from every category — the ``integration/`` and ``functional/`` -conftests are docstrings only. Mock data is the exception: -``tests/unit/conftest.py``. Read those two files for the authoritative list; the -ones you will reach for most: +Reusable setup lives in ``tests/fixtures/``. The root ``tests/conftest.py`` +registers shared plugins and owns test-selection hooks. The unit conftest exposes +synthetic numerical factories; the functional conftest exposes the module-scoped, +read-only Fourier-model fixture. See ``tests/fixtures/README.md`` for ownership +and mutation rules. Common fixtures include: - Paths (session-scoped): ``tests_root``, ``project_root``, ``test_files_dir``, the per-format ``cif_dir``, ``mtz_dir``, ``pdb_dir``, ``cif_sf_dir``, and @@ -118,9 +124,11 @@ ones you will reach for most: ``mock_aniso_u``, ``mock_scattering_factors``, ``mock_weights``. - Real files: ``sample_cif_file``, ``sample_pdb_file``, ``sample_mtz_file``, ``sample_structure_factor_cif``, ``sample_structure_pair`` (matched model + - data), ``all_structure_pairs``, ``all_test_structures``. -- Loaded objects: ``loaded_model``, ``loaded_reflection_data``, - ``model_and_data``, ``initialized_scaler``. + data), ``compatibility_structure_pair`` (one named slow crystal). +- Loaded objects: ``loaded_model``, ``loaded_model_ft``, ``loaded_reflection_data``, + ``model_and_data``, ``initialized_scaler``. ``compatibility_model`` and + ``compatibility_model_and_data`` load only the current slow case and remain + function-scoped to isolate mutations. The mock-data fixtures yield a *factory* taking ``n_atoms`` / ``n_reflections`` and ``seed``; ``mock_cell`` and ``mock_cell_triclinic`` yield the tensor diff --git a/tests/README.md b/tests/README.md index 9298352f..3be67e2f 100644 --- a/tests/README.md +++ b/tests/README.md @@ -51,6 +51,7 @@ tests/ | CIF atomic fields and crystal metadata | `integration/test_io_cif.py` | | MTZ fields, resolution bins and model/data crystal agreement | `integration/test_io_reflections.py` | | ModelFT forward cache and grid integration | `functional/test_model_ft_functional.py` | +| Extra deposited files and input inventory | `integration/test_structure_compatibility.py`, `helpers/structure_cases.py` | | Numerical derivatives and backend parity | `unit/test_gradient_correctness.py`, `unit/structure_factor/` | A production call must participate in the assertion: computing a formula only in @@ -58,6 +59,14 @@ the test does not check its implementation. Kernel values, target registration, device transitions, and default configuration are separate contracts even when they exercise the same class. Keep mutation tests on fresh objects. +The quick reader contracts use 1DAW. Extended reader compatibility runs with +`pytest tests/integration/test_structure_compatibility.py --run-slow`; each file +is a separate case and must succeed. The manifest covers the bundled CIF, MTZ +and SF-CIF inputs, including the IHM fixture and reflection-only depositions. +Adding a data file requires an explicit coverage assignment in the manifest. +Extended scaler and restraint cases use 2DQ6 (trigonal) and 3A5V (body-centred +tetragonal), with fresh objects per case and `--run-slow` required. + ### Quick Local Run (on login node, for small tests only) ```bash diff --git a/tests/fixtures/README.md b/tests/fixtures/README.md index 7361c509..282c5527 100644 --- a/tests/fixtures/README.md +++ b/tests/fixtures/README.md @@ -7,7 +7,7 @@ Keep a fixture in its test module when only that module needs it. | Module | Responsibility | Visibility / lifetime | |---|---|---| | `paths.py` | Repository, bundled-data and optional library paths | All tests; session | -| `files.py` | Sample-file selection and matching structure pairs | All tests; session; no model loading | +| `files.py` | Sample paths and named compatibility pairs | All tests; sample paths session-scoped, extended pairs function-scoped; no loading | | `devices.py` | Configured device, explicit backends, device parametrization | All tests; existing per-fixture scopes | | `precision.py` | Comparison tolerances and CPU-double reference context | All tests; reference fixture restores state after each test | | `objects.py` | Mutable models, data, scalers and restraints | All tests; fresh per test except explicitly shared bundles | @@ -32,6 +32,11 @@ device movement, or empty caches use fresh objects. `loaded_model`, fresh mutable objects per test. The explicitly shared session bundles in that module retain their documented ownership contracts. +`compatibility_structure_pair` selects named slow cases from +`tests/helpers/structure_cases.py`. `compatibility_model` loads just that model; +`compatibility_model_and_data` adds observations only when needed. Skipped slow +cases do not load any structures. + Use `cpu_double_precision()` to scope an explicit numerical reference, or request `double_cpu` for a single test. The structure-factor package uses the same context at package scope; both usages restore dtype, device, and density cutoff on exit. diff --git a/tests/fixtures/files.py b/tests/fixtures/files.py index d248775e..390ea678 100644 --- a/tests/fixtures/files.py +++ b/tests/fixtures/files.py @@ -4,6 +4,23 @@ import pytest +from tests.helpers.structure_cases import EXTENDED_PAIR_CODES + + +@pytest.fixture( + params=[pytest.param(code, marks=pytest.mark.slow) for code in EXTENDED_PAIR_CODES] +) +def compatibility_structure_pair( + cif_dir: Path, mtz_dir: Path, request: pytest.FixtureRequest +) -> dict: + """Select one named extended crystal without loading its model or observations.""" + code = request.param + return { + "pdb_id": code, + "model": cif_dir / f"{code}.cif", + "reflections": mtz_dir / f"{code}.mtz", + } + @pytest.fixture(scope="session") def sample_cif_file(cif_dir: Path) -> Path: @@ -70,29 +87,3 @@ def sample_structure_pair(cif_dir: Path, mtz_dir: Path) -> dict[str, Path]: return {"model": cif_files[pdb_id], "reflections": mtz_files[pdb_id]} pytest.skip("No matching CIF/MTZ pairs found in test data") - - -@pytest.fixture(scope="session") -def all_structure_pairs(cif_dir: Path, mtz_dir: Path) -> list[dict[str, Path | str]]: - """Return all matching pairs of CIF models and MTZ reflections.""" - cif_files = {f.stem: f for f in cif_dir.glob("*.cif")} - mtz_files = {f.stem: f for f in mtz_dir.glob("*.mtz")} - - common_ids = set(cif_files.keys()) & set(mtz_files.keys()) - - if not common_ids: - pytest.skip("No matching CIF/MTZ pairs found in test data") - - return [ - {"pdb_id": pdb_id, "model": cif_files[pdb_id], "reflections": mtz_files[pdb_id]} - for pdb_id in sorted(common_ids) - ] - - -@pytest.fixture(scope="session") -def all_cif_files(cif_dir: Path) -> list[Path]: - """Return all available CIF test structure files.""" - cif_files = sorted(cif_dir.glob("*.cif")) - if not cif_files: - pytest.skip("No CIF files found in test data directory") - return cif_files diff --git a/tests/fixtures/objects.py b/tests/fixtures/objects.py index 34fb00c8..3eb7628e 100644 --- a/tests/fixtures/objects.py +++ b/tests/fixtures/objects.py @@ -18,6 +18,31 @@ from torchref.scaling import Scaler +@pytest.fixture +def compatibility_model(compatibility_structure_pair: dict) -> Model: + """Load a fresh model for one slow compatibility case.""" + from torchref.model import Model + + path = compatibility_structure_pair["model"] + assert path.is_file() + return Model(verbose=0).load_cif(str(path)) + + +@pytest.fixture +def compatibility_model_and_data( + compatibility_model: Model, compatibility_structure_pair: dict +) -> dict: + """Load observations only for the single crystal used by the current pipeline case.""" + from torchref.io import ReflectionData + + path = compatibility_structure_pair["reflections"] + assert path.is_file() + return { + "model": compatibility_model, + "data": ReflectionData(verbose=0).load_mtz(str(path)), + } + + @pytest.fixture def loaded_model(sample_cif_file: Path) -> Model: """Load a fresh mutable Model from the sample CIF file.""" @@ -97,42 +122,6 @@ def model_with_restraints(loaded_model: Model) -> dict[str, Any]: return {"model": loaded_model, "restraints": restraints} -@pytest.fixture(scope="session") -def all_test_structures( - all_structure_pairs: list[dict[str, Any]], -) -> list[dict[str, Any]]: - """Return all loaded model/data pairs for comprehensive testing.""" - from torchref.io import ReflectionData - from torchref.model.model import Model - - structures = [] - for pair in all_structure_pairs: - try: - model = Model() - model.load_cif(str(pair["model"])) - - data = ReflectionData() - data.load_mtz(str(pair["reflections"])) - - structures.append( - { - "pdb_id": pair["pdb_id"], - "model": model, - "data": data, - "model_path": pair["model"], - "data_path": pair["reflections"], - } - ) - except Exception: - # Skip structures that fail to load - continue - - if not structures: - pytest.skip("No structures could be loaded") - - return structures - - @pytest.fixture(scope="session") def _device_model_cache() -> dict: """``{device_str: ModelFT}`` built at most once per device, per session.""" diff --git a/tests/functional/test_io_functional.py b/tests/functional/test_io_functional.py deleted file mode 100644 index 49c82a99..00000000 --- a/tests/functional/test_io_functional.py +++ /dev/null @@ -1,79 +0,0 @@ -""" -Functional tests for I/O operations. - -Tests file loading and data processing with real crystallographic data. -""" - -import pytest - - -class TestCIFReadingFunctional: - """Functional tests for CIF file reading.""" - - @pytest.mark.integration - def test_load_multiple_cif_files(self, cif_dir): - """Test loading multiple CIF files successfully.""" - from torchref.model.model import Model - - cif_files = list(cif_dir.glob("*.cif")) - assert len(cif_files) > 0, "No CIF files found in test directory" - - for cif_file in cif_files: - model = Model() - model.load_cif(str(cif_file)) - - # Each file should load with atoms - n_atoms = model.xyz().shape[0] - assert n_atoms > 0, f"No atoms loaded from {cif_file}" - - # Should have cell parameters - assert model.cell is not None - assert len(model.cell) == 6 - - -class TestMTZReadingFunctional: - """Functional tests for MTZ file reading.""" - - @pytest.mark.integration - def test_load_multiple_mtz_files(self, mtz_dir): - """Test loading multiple MTZ files successfully.""" - from torchref.io import ReflectionData - - mtz_files = list(mtz_dir.glob("*.mtz")) - assert len(mtz_files) > 0, "No MTZ files found in test directory" - - for mtz_file in mtz_files: - data = ReflectionData() - data.load_mtz(str(mtz_file)) - - # Each file should load with reflections - n_refl = data.hkl.shape[0] - assert n_refl > 0, f"No reflections loaded from {mtz_file}" - - # Should have cell parameters - assert data.cell is not None - - -class TestSFCIFReadingFunctional: - """Functional tests for structure factor CIF reading.""" - - @pytest.mark.integration - def test_load_sf_cif(self, cif_sf_dir): - """Test loading structure factor CIF files.""" - from torchref.io import ReflectionData - - sf_files = list(cif_sf_dir.glob("*.cif")) - if not sf_files: - pytest.skip("No SF-CIF files found") - - for sf_file in sf_files: - data = ReflectionData() - try: - data.load_cif(str(sf_file)) - - # Should have loaded reflections - if data.hkl is not None: - assert data.hkl.shape[0] > 0 - except Exception as e: - # Some files may not be valid SF-CIF format - pass diff --git a/tests/functional/test_model_ft_functional.py b/tests/functional/test_model_ft_functional.py index 2c6b9b4c..ce9b5e71 100644 --- a/tests/functional/test_model_ft_functional.py +++ b/tests/functional/test_model_ft_functional.py @@ -104,32 +104,6 @@ def test_map_symmetry_available(self, shared_model_ft): assert operator.map_shape == gridsize -@pytest.mark.integration -class TestModelFTMultipleStructures: - """Test ModelFT with multiple structures.""" - - def test_modelft_multiple_structures(self, all_structure_pairs): - """Test ModelFT works with different structures.""" - from torchref.model.model_ft import ModelFT - - tested = 0 - for pair in all_structure_pairs[:3]: # Test first 3 - try: - model = ModelFT(max_res=3.0, verbose=0) - model.load_cif(str(pair["model"])) - - # Basic checks - assert model.xyz() is not None - assert model.xyz().shape[0] > 0 - - tested += 1 - except Exception as e: - # Some structures may fail to load - continue - - assert tested >= 1, "At least one structure should load" - - @pytest.mark.integration class TestModelFTCoordinateOperations: """Test ModelFT coordinate operations.""" diff --git a/tests/functional/test_restraints_functional.py b/tests/functional/test_restraints_functional.py index 76534421..841022cc 100644 --- a/tests/functional/test_restraints_functional.py +++ b/tests/functional/test_restraints_functional.py @@ -21,10 +21,7 @@ def test_build_restraints_from_cif(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, - xyz_fn=model.xyz, - vdw_radii_fn=model.get_vdw_radii, - verbose=0 + pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 ) restraints.build_restraints() @@ -42,34 +39,31 @@ def test_bond_restraints_built(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, - xyz_fn=model.xyz, - vdw_radii_fn=model.get_vdw_radii, - verbose=0 + pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 ) restraints.build_restraints() # Check bond restraints exist - assert 'bond' in restraints.restraints + assert "bond" in restraints.restraints # Check intra-residue bonds - if 'intra' in restraints.restraints['bond']: - bond_intra = restraints.restraints['bond']['intra'] - assert 'indices' in bond_intra - assert 'references' in bond_intra - assert 'sigmas' in bond_intra + if "intra" in restraints.restraints["bond"]: + bond_intra = restraints.restraints["bond"]["intra"] + assert "indices" in bond_intra + assert "references" in bond_intra + assert "sigmas" in bond_intra # Indices should be 2D with shape (N, 2) - indices = bond_intra['indices'] + indices = bond_intra["indices"] assert len(indices.shape) == 2 assert indices.shape[1] == 2 # References should match number of bonds - assert bond_intra['references'].shape[0] == indices.shape[0] - assert bond_intra['sigmas'].shape[0] == indices.shape[0] + assert bond_intra["references"].shape[0] == indices.shape[0] + assert bond_intra["sigmas"].shape[0] == indices.shape[0] # Bond lengths should be positive and reasonable (0.5-3.0 Å) - refs = bond_intra['references'] + refs = bond_intra["references"] assert torch.all(refs > 0.5) assert torch.all(refs < 3.0) @@ -83,29 +77,26 @@ def test_angle_restraints_built(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, - xyz_fn=model.xyz, - vdw_radii_fn=model.get_vdw_radii, - verbose=0 + pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 ) restraints.build_restraints() # Check angle restraints exist - assert 'angle' in restraints.restraints + assert "angle" in restraints.restraints - if 'intra' in restraints.restraints['angle']: - angle_intra = restraints.restraints['angle']['intra'] - assert 'indices' in angle_intra - assert 'references' in angle_intra - assert 'sigmas' in angle_intra + if "intra" in restraints.restraints["angle"]: + angle_intra = restraints.restraints["angle"]["intra"] + assert "indices" in angle_intra + assert "references" in angle_intra + assert "sigmas" in angle_intra # Indices should be 2D with shape (N, 3) - indices = angle_intra['indices'] + indices = angle_intra["indices"] assert len(indices.shape) == 2 assert indices.shape[1] == 3 # References should match number of angles - assert angle_intra['references'].shape[0] == indices.shape[0] + assert angle_intra["references"].shape[0] == indices.shape[0] @pytest.mark.integration def test_torsion_restraints_built(self, sample_cif_file): @@ -117,25 +108,22 @@ def test_torsion_restraints_built(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, - xyz_fn=model.xyz, - vdw_radii_fn=model.get_vdw_radii, - verbose=0 + pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 ) restraints.build_restraints() # Check torsion restraints exist - assert 'torsion' in restraints.restraints + assert "torsion" in restraints.restraints - if 'intra' in restraints.restraints['torsion']: - torsion_intra = restraints.restraints['torsion']['intra'] - assert 'indices' in torsion_intra - assert 'references' in torsion_intra - assert 'sigmas' in torsion_intra - assert 'periods' in torsion_intra + if "intra" in restraints.restraints["torsion"]: + torsion_intra = restraints.restraints["torsion"]["intra"] + assert "indices" in torsion_intra + assert "references" in torsion_intra + assert "sigmas" in torsion_intra + assert "periods" in torsion_intra # Indices should be 2D with shape (N, 4) - indices = torsion_intra['indices'] + indices = torsion_intra["indices"] assert len(indices.shape) == 2 assert indices.shape[1] == 4 @@ -149,23 +137,20 @@ def test_plane_restraints_built(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, - xyz_fn=model.xyz, - vdw_radii_fn=model.get_vdw_radii, - verbose=0 + pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 ) restraints.build_restraints() # Check plane restraints exist - assert 'plane' in restraints.restraints + assert "plane" in restraints.restraints # Planes are grouped by atom count (3_atoms, 4_atoms, etc.) - plane_restraints = restraints.restraints['plane'] + plane_restraints = restraints.restraints["plane"] if len(list(plane_restraints.keys())) > 0: # Check at least one plane group exists for key, plane_group in plane_restraints.items(): - if 'indices' in plane_group: - indices = plane_group['indices'] + if "indices" in plane_group: + indices = plane_group["indices"] # Planes need at least 3 atoms if len(indices.shape) == 2: assert indices.shape[1] >= 3 @@ -184,15 +169,12 @@ def test_bond_deviations(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, - xyz_fn=model.xyz, - vdw_radii_fn=model.get_vdw_radii, - verbose=0 + pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 ) restraints.build_restraints() # Compute bond deviations - if hasattr(restraints, 'bond_deviations'): + if hasattr(restraints, "bond_deviations"): deviations, sigmas = restraints.bond_deviations() assert torch.all(torch.isfinite(deviations)) @@ -211,15 +193,12 @@ def test_angle_deviations(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, - xyz_fn=model.xyz, - vdw_radii_fn=model.get_vdw_radii, - verbose=0 + pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 ) restraints.build_restraints() # Compute angle deviations - if hasattr(restraints, 'angle_deviations'): + if hasattr(restraints, "angle_deviations"): deviations, sigmas = restraints.angle_deviations() assert torch.all(torch.isfinite(deviations)) @@ -231,28 +210,20 @@ class TestRestraintsMultipleStructures: @pytest.mark.integration @pytest.mark.slow - def test_restraints_multiple_cif_files(self, cif_dir): - """Test building restraints for multiple CIF files.""" - from torchref.model.model import Model + def test_restraints_multiple_cif_files(self, compatibility_model): + """Each extended crystal supplies bond and angle restraints.""" from torchref.topology.restraints import Restraints - cif_files = list(cif_dir.glob("*.cif"))[:3] # First 3 structures - - for cif_file in cif_files: - model = Model() - model.load_cif(str(cif_file)) - - restraints = Restraints( - pdb=model.pdb, - xyz_fn=model.xyz, - vdw_radii_fn=model.get_vdw_radii, - verbose=0 - ) - restraints.build_restraints() - - # Should have built restraints for each structure - assert 'bond' in restraints.restraints - assert 'angle' in restraints.restraints + model = compatibility_model + restraints = Restraints( + pdb=model.pdb, + xyz_fn=model.xyz, + vdw_radii_fn=model.get_vdw_radii, + verbose=0, + ) + restraints.build_restraints() + assert "bond" in restraints.restraints + assert "angle" in restraints.restraints class TestRestraintsDeviceHandling: @@ -268,16 +239,13 @@ def test_restraints_device_movement(self, sample_cif_file, cpu_device): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, - xyz_fn=model.xyz, - vdw_radii_fn=model.get_vdw_radii, - verbose=0 + pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 ) restraints.build_restraints() # Check that tensors are on the correct device - if 'bond' in restraints.restraints and 'intra' in restraints.restraints['bond']: - bond_indices = restraints.restraints['bond']['intra']['indices'] + if "bond" in restraints.restraints and "intra" in restraints.restraints["bond"]: + bond_indices = restraints.restraints["bond"]["intra"]["indices"] assert bond_indices.device == cpu_device @@ -294,10 +262,7 @@ def test_cif_dict_loaded(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, - xyz_fn=model.xyz, - vdw_radii_fn=model.get_vdw_radii, - verbose=0 + pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 ) # CIF dict should be populated with residue restraints @@ -305,10 +270,13 @@ def test_cif_dict_loaded(self, sample_cif_file): assert len(restraints.cif_dict) > 0 # Should have standard amino acids - common_residues = ['ALA', 'GLY', 'VAL', 'LEU', 'ILE'] + common_residues = ["ALA", "GLY", "VAL", "LEU", "ILE"] for res in common_residues: if res in restraints.cif_dict: - assert 'bonds' in restraints.cif_dict[res] or 'angles' in restraints.cif_dict[res] + assert ( + "bonds" in restraints.cif_dict[res] + or "angles" in restraints.cif_dict[res] + ) @pytest.mark.integration def test_unique_residues_detected(self, sample_cif_file): @@ -320,10 +288,7 @@ def test_unique_residues_detected(self, sample_cif_file): model.load_cif(str(sample_cif_file)) restraints = Restraints( - pdb=model.pdb, - xyz_fn=model.xyz, - vdw_radii_fn=model.get_vdw_radii, - verbose=0 + pdb=model.pdb, xyz_fn=model.xyz, vdw_radii_fn=model.get_vdw_radii, verbose=0 ) # Should have detected unique residues diff --git a/tests/functional/test_scaler_functional.py b/tests/functional/test_scaler_functional.py index d1364665..e6c3b3a7 100644 --- a/tests/functional/test_scaler_functional.py +++ b/tests/functional/test_scaler_functional.py @@ -6,7 +6,23 @@ import pytest import torch -import numpy as np + + +@pytest.mark.integration +def test_scaler_crystal_compatibility(compatibility_model_and_data) -> None: + """Each extended crystal produces finite anisotropic scale corrections.""" + from torchref.scaling import Scaler + + model = compatibility_model_and_data["model"] + data = compatibility_model_and_data["data"] + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) + scaler.setup_anisotropy_correction() + assert scaler.s is not None + assert scaler.bins is not None + assert scaler.U is not None + correction = scaler.anisotropy_correction() + assert correction.shape == (len(data.hkl),) + assert torch.isfinite(correction).all() class TestScalerCreationFunctional: @@ -15,18 +31,18 @@ class TestScalerCreationFunctional: @pytest.mark.integration def test_scaler_full_initialization(self, sample_structure_pair): """Test full scaler initialization with model and data.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + scaler = Scaler(model=model, data=data, nbins=20, verbose=0) - + # Check all components are initialized assert scaler.model is not None assert scaler._data is not None @@ -39,8 +55,8 @@ def test_scaler_full_initialization(self, sample_structure_pair): @pytest.mark.parametrize("nbins", [5, 10, 15, 20]) def test_scaler_with_different_nbins(self, sample_structure_pair, nbins): """Test scaler with different bin counts.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler model = Model() @@ -62,18 +78,18 @@ class TestScatteringVectorsFunctional: @pytest.mark.integration def test_scattering_vectors_shape(self, sample_structure_pair): """Test scattering vectors have correct shape.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) - + # s should have shape (n_reflections, 3) n_refl = data.hkl.shape[0] assert scaler.s.shape == (n_refl, 3) @@ -81,21 +97,21 @@ def test_scattering_vectors_shape(self, sample_structure_pair): @pytest.mark.integration def test_scattering_vectors_magnitude(self, sample_structure_pair): """Test scattering vector magnitudes are reasonable.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) - + # Calculate |s| = sin(theta)/lambda = 1/(2d) s_mag = torch.norm(scaler.s, dim=1) - + # For typical protein data: # Low resolution (d=100Å): |s| ~ 0.005 # High resolution (d=1Å): |s| ~ 0.5 @@ -109,26 +125,26 @@ class TestAnisotropyCorrectionFunctional: @pytest.mark.integration def test_anisotropy_setup_and_compute(self, sample_structure_pair): """Test setting up and computing anisotropy correction.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) scaler.setup_anisotropy_correction() - + # U parameters should exist - assert hasattr(scaler, 'U') + assert hasattr(scaler, "U") assert scaler.U.shape == (6,) # U11, U22, U33, U12, U13, U23 - + # Compute correction correction = scaler.anisotropy_correction() - + # Correction should be positive (exponential) assert correction.shape[0] == data.hkl.shape[0] assert torch.all(correction > 0) @@ -137,22 +153,22 @@ def test_anisotropy_setup_and_compute(self, sample_structure_pair): @pytest.mark.integration def test_anisotropy_correction_near_unity(self, sample_structure_pair): """Test anisotropy correction starts near unity with small U.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) scaler.setup_anisotropy_correction() - + # With small random U values, correction should be close to 1 correction = scaler.anisotropy_correction() - + # Most values should be between 0.5 and 2.0 for small U mean_correction = correction.mean().item() assert 0.5 < mean_correction < 2.0 @@ -164,20 +180,20 @@ class TestBinwiseBfactorFunctional: @pytest.mark.integration def test_setup_binwise_bfactor(self, sample_structure_pair): """Test setting up bin-wise B-factor parameters.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) scaler.setup_bin_wise_bfactor() - assert hasattr(scaler, 'bin_wise_bfactor') + assert hasattr(scaler, "bin_wise_bfactor") assert scaler.bin_wise_bfactor.shape == (10,) # Initially should be zeros assert torch.allclose( @@ -187,24 +203,24 @@ def test_setup_binwise_bfactor(self, sample_structure_pair): @pytest.mark.integration def test_binwise_bfactor_correction(self, sample_structure_pair): """Test computing bin-wise B-factor correction.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) scaler.setup_bin_wise_bfactor() - + # Set some non-zero B-factors scaler.bin_wise_bfactor.data = torch.linspace(0, 20, 10, device=scaler.device) - + correction = scaler.bin_wise_bfactor_correction() - + # Correction should have same length as reflections assert correction.shape[0] == data.hkl.shape[0] # Should be positive (exponential) @@ -218,35 +234,35 @@ class TestScalerStateDictFunctional: @pytest.mark.integration def test_save_and_load_state_dict(self, sample_structure_pair, tmp_path): """Test saving and loading scaler state.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + # Create scaler with some setup scaler1 = Scaler(model=model, data=data, nbins=10, verbose=0) scaler1.setup_anisotropy_correction() scaler1.setup_bin_wise_bfactor() - + # Modify parameters scaler1.U.data = torch.randn(6, device=scaler1.device) scaler1.bin_wise_bfactor.data = torch.randn(10, device=scaler1.device) - + # Save state state_path = tmp_path / "scaler_state.pt" torch.save(scaler1.state_dict(), state_path) - + # Create new scaler and load state scaler2 = Scaler(model=model, data=data, nbins=10, verbose=0) scaler2.setup_anisotropy_correction() scaler2.setup_bin_wise_bfactor() scaler2.load_state_dict(torch.load(state_path, weights_only=False)) - + # Parameters should match assert torch.allclose(scaler1.U, scaler2.U) assert torch.allclose(scaler1.bin_wise_bfactor, scaler2.bin_wise_bfactor) @@ -258,18 +274,18 @@ class TestScalerHKLPropertyFunctional: @pytest.mark.integration def test_hkl_property(self, sample_structure_pair): """Test that HKL property returns correct indices.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) - + # HKL from scaler should match data hkl = scaler.hkl assert hkl is not None @@ -283,43 +299,45 @@ class TestScalerDeviceOperationsFunctional: @pytest.mark.integration def test_scaler_cpu_operation(self, sample_structure_pair): """Test scaler works on CPU.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - - scaler = Scaler(model=model, data=data, nbins=10, verbose=0, device=torch.device('cpu')) + + scaler = Scaler( + model=model, data=data, nbins=10, verbose=0, device=torch.device("cpu") + ) scaler.setup_anisotropy_correction() - - assert scaler.device.type == 'cpu' - assert scaler.s.device.type == 'cpu' - assert scaler.U.device.type == 'cpu' + + assert scaler.device.type == "cpu" + assert scaler.s.device.type == "cpu" + assert scaler.U.device.type == "cpu" @pytest.mark.integration def test_scaler_cpu_method(self, sample_structure_pair): """Test scaler.cpu() method.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) scaler.setup_anisotropy_correction() scaler.cpu() - + # All tensors should be on CPU for param in scaler.parameters(): - assert param.device.type == 'cpu' + assert param.device.type == "cpu" class TestScalerUMatrixFunctional: @@ -328,23 +346,23 @@ class TestScalerUMatrixFunctional: @pytest.mark.integration def test_u_to_matrix_conversion(self, sample_structure_pair): """Test conversion from U parameters to 3x3 matrix.""" - from torchref.model.model import Model + from torchref.base.math_torch import U_to_matrix from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - from torchref.base.math_torch import U_to_matrix - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) scaler.setup_anisotropy_correction() - + # Convert U vector to matrix U_matrix = U_to_matrix(scaler.U) - + # Should be 3x3 assert U_matrix.shape == (3, 3) # Should be symmetric @@ -357,24 +375,24 @@ class TestScalerGradientsFunctional: @pytest.mark.integration def test_anisotropy_gradients(self, sample_structure_pair): """Test gradients flow through anisotropy correction.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) scaler.setup_anisotropy_correction() - + # Compute correction and loss correction = scaler.anisotropy_correction() loss = correction.sum() loss.backward() - + # U should have gradients assert scaler.U.grad is not None assert torch.all(torch.isfinite(scaler.U.grad)) @@ -382,56 +400,24 @@ def test_anisotropy_gradients(self, sample_structure_pair): @pytest.mark.integration def test_binwise_bfactor_gradients(self, sample_structure_pair): """Test gradients flow through bin-wise B-factor correction.""" - from torchref.model.model import Model from torchref.io import ReflectionData + from torchref.model.model import Model from torchref.scaling.scaler import Scaler - + model = Model() model.load_cif(str(sample_structure_pair["model"])) - + data = ReflectionData() data.load_mtz(str(sample_structure_pair["reflections"])) - + scaler = Scaler(model=model, data=data, nbins=10, verbose=0) scaler.setup_bin_wise_bfactor() - + # Compute correction and loss correction = scaler.bin_wise_bfactor_correction() loss = correction.sum() loss.backward() - + # bin_wise_bfactor should have gradients assert scaler.bin_wise_bfactor.grad is not None assert torch.all(torch.isfinite(scaler.bin_wise_bfactor.grad)) - - -class TestScalerMultipleStructuresFunctional: - """Functional tests with multiple structures.""" - - @pytest.mark.integration - def test_scaler_with_different_structures(self, all_test_structures): - """Test scaler works with different crystal structures.""" - from torchref.scaling.scaler import Scaler - - tested = 0 - for struct in all_test_structures: - pdb_id = struct["pdb_id"] - model = struct["model"] - data = struct["data"] - - scaler = Scaler(model=model, data=data, nbins=10, verbose=0) - scaler.setup_anisotropy_correction() - - # Verify scaler is set up correctly - assert scaler.s is not None - assert scaler.bins is not None - assert scaler.U is not None - - correction = scaler.anisotropy_correction() - assert torch.all(torch.isfinite(correction)) - - tested += 1 - if tested >= 3: # Test first 3 structures - break - - assert tested >= 1, "No test structures with both CIF and MTZ found" diff --git a/tests/helpers/structure_cases.py b/tests/helpers/structure_cases.py new file mode 100644 index 00000000..74a84c28 --- /dev/null +++ b/tests/helpers/structure_cases.py @@ -0,0 +1,27 @@ +"""Name compatibility datasets explicitly so adding a file cannot grow test work silently. + +The quick reader contracts use 1DAW. The broader panel exercises deposited files +across crystal systems and file encodings in the slow tier. Pair-based pipeline +checks use trigonal 2DQ6 and body-centred tetragonal 3A5V in addition to their +separate 1DAW checks. +""" + +MODEL_CODES = ( + "1DAW", # C-centred monoclinic; quick reference structure. + "2DQ6", # Trigonal. + "3A5V", # Body-centred tetragonal. + "3E98", # Monoclinic screw axis. + "3GR5", # Hexagonal screw axis. + "3K7M", # Cubic. + "3VRJ", # Additional monoclinic deposition. + "4BX9", # Tetragonal screw axis. + "5BOV", # Triclinic P1. + "6G9X", # Orthorhombic. +) + +MTZ_CODES = MODEL_CODES + ("1AK5", "1BYW", "1VER", "6JZA", "6SXW", "6VHI") +SF_CIF_CODES = MODEL_CODES + ("7L84",) +EXTENDED_PAIR_CODES = ("2DQ6", "3A5V") +MODEL_CIF_FILES = tuple(f"{code}.cif" for code in MODEL_CODES) + ( + "test_ihm_ensemble.cif", +) diff --git a/tests/integration/test_io_cif.py b/tests/integration/test_io_cif.py index 7e8b23d7..2bd47fa1 100644 --- a/tests/integration/test_io_cif.py +++ b/tests/integration/test_io_cif.py @@ -39,37 +39,6 @@ def test_cif_loading_contract(loaded_model, sample_cif_file) -> None: ) -class TestMultipleCIFFiles: - """Tests that load multiple CIF files.""" - - @pytest.mark.integration - @pytest.mark.slow - def test_load_all_test_structures(self, all_cif_files): - """Test loading all available test structures.""" - from torchref.model.model import Model - - loaded = 0 - errors = [] - - for cif_file in all_cif_files: - try: - model = Model() - model.load_cif(str(cif_file)) - n_atoms = model.xyz().shape[0] - assert n_atoms > 0 - loaded += 1 - except Exception as e: - errors.append((cif_file.name, str(e))) - - # Report - print(f"\nLoaded {loaded}/{len(all_cif_files)} structures") - if errors: - print(f"Errors: {errors}") - - # Should load at least most structures - assert loaded > 0 - - class TestCIFSaving: """Tests for saving CIF files.""" diff --git a/tests/integration/test_structure_compatibility.py b/tests/integration/test_structure_compatibility.py new file mode 100644 index 00000000..2ad00c10 --- /dev/null +++ b/tests/integration/test_structure_compatibility.py @@ -0,0 +1,101 @@ +"""Exercise explicitly named extra structure files in the slow compatibility tier.""" + +import pytest +import torch + +from tests.helpers.structure_cases import ( + EXTENDED_PAIR_CODES, + MODEL_CIF_FILES, + MTZ_CODES, + SF_CIF_CODES, +) +from torchref.config import ( + canonical_device, + get_default_device, + get_float_dtype, + get_int_dtype, +) + +pytestmark = pytest.mark.integration + + +@pytest.mark.parametrize( + "directory, expected", + [ + ("cif", MODEL_CIF_FILES), + ("mtz", tuple(f"{code}.mtz" for code in MTZ_CODES)), + ("cif_sf", tuple(f"{code}-sf.cif" for code in SF_CIF_CODES)), + ], + ids=["models", "mtz", "sf-cif"], +) +def test_compatibility_inventory(test_files_dir, directory, expected) -> None: + """Every bundled input has an explicit quick or extended coverage assignment.""" + suffix = ".mtz" if directory == "mtz" else ".cif" + actual = {path.name for path in (test_files_dir / directory).glob(f"*{suffix}")} + assert actual == set(expected) + + +@pytest.mark.slow +@pytest.mark.parametrize( + "filename", [name for name in MODEL_CIF_FILES if name != "1DAW.cif"] +) +def test_model_cif_compatibility(cif_dir, filename) -> None: + """Each extra CIF loads atoms and finite symmetry operators on the default device.""" + from torchref.model import Model + + path = cif_dir / filename + assert path.is_file() + model = Model(verbose=0).load_cif(str(path)) + xyz = model.xyz() + assert xyz.shape == (len(model.pdb), 3) + assert len(xyz) > 0 + assert xyz.dtype == get_float_dtype() + assert canonical_device(xyz.device) == canonical_device(get_default_device()) + assert torch.isfinite(xyz).all() + assert model.cell.data.shape == (6,) + assert model.spacegroup.matrices.shape[0] > 0 + assert torch.isfinite(model.spacegroup.matrices).all() + + +@pytest.mark.slow +@pytest.mark.parametrize("code", EXTENDED_PAIR_CODES) +def test_modelft_cif_compatibility(cif_dir, code) -> None: + """Fourier models initialize their scattering parametrization in distinct crystals.""" + from torchref.model import ModelFT + + path = cif_dir / f"{code}.cif" + assert path.is_file() + model = ModelFT(max_res=3.0, verbose=0).load_cif(str(path)) + assert len(model.xyz()) > 0 + assert model.parametrization + assert all(size > 0 for size in model.grid_shape) + + +@pytest.mark.slow +@pytest.mark.parametrize( + "directory, filename, loader", + [("mtz", f"{code}.mtz", "load_mtz") for code in MTZ_CODES if code != "1DAW"] + + [ + ("cif_sf", f"{code}-sf.cif", "load_cif") + for code in SF_CIF_CODES + if code != "1DAW" + ], + ids=[f"mtz-{code}" for code in MTZ_CODES if code != "1DAW"] + + [f"sf-cif-{code}" for code in SF_CIF_CODES if code != "1DAW"], +) +def test_reflection_file_compatibility( + test_files_dir, directory, filename, loader +) -> None: + """Every named MTZ/SF-CIF must load reflections; one success cannot mask another failure.""" + from torchref.io import ReflectionData + + path = test_files_dir / directory / filename + assert path.is_file() + data = ReflectionData(verbose=0) + getattr(data, loader)(str(path)) + assert data.hkl.shape == (len(data.hkl), 3) + assert len(data.hkl) > 0 + assert data.hkl.dtype == get_int_dtype() + assert canonical_device(data.hkl.device) == canonical_device(get_default_device()) + assert data.cell.data.shape == (6,) + assert data.F.shape == (len(data.hkl),) diff --git a/tests/integration/test_symmetry_integration.py b/tests/integration/test_symmetry_integration.py index 8bb0a52d..a39ce804 100644 --- a/tests/integration/test_symmetry_integration.py +++ b/tests/integration/test_symmetry_integration.py @@ -6,7 +6,6 @@ import pytest import torch -from pathlib import Path class TestSpaceGroupInitialization: @@ -106,8 +105,8 @@ class TestSpaceGroupDevice: @pytest.mark.integration def test_spacegroup_default_device(self): """Test SpaceGroup matrices land on the configured default device.""" - from torchref.symmetry import SpaceGroup from torchref.config import get_default_device + from torchref.symmetry import SpaceGroup sg = SpaceGroup("P 21 21 21") @@ -140,7 +139,7 @@ def test_expand_coordinates(self, sample_cif_file): # The model should be able to generate symmetry mates # Check if there's an expand method - if hasattr(sg, 'expand') or hasattr(sg, 'expand_atoms'): + if hasattr(sg, "expand") or hasattr(sg, "expand_atoms"): expanded = sg.expand(xyz) assert expanded.shape[0] >= xyz.shape[0] @@ -148,23 +147,6 @@ def test_expand_coordinates(self, sample_cif_file): class TestSpacegroupVariants: """Tests for different spacegroup conventions.""" - @pytest.mark.integration - @pytest.mark.parametrize("sg_name", [ - "P 1", # Triclinic - "P 21", # Monoclinic - "P 21 21 21", # Orthorhombic - "P 43 21 2", # Tetragonal - "P 3 2 1", # Trigonal - "P 6 2 2", # Hexagonal - "P 2 3", # Cubic - ]) - def test_common_spacegroups(self, sg_name): - """Test loading common spacegroups.""" - from torchref.symmetry import SpaceGroup - - sg = SpaceGroup(sg_name) - assert sg.matrices is not None - @pytest.mark.integration def test_spacegroup_name_variations(self): """Test that different spacegroup name formats work.""" @@ -181,25 +163,6 @@ def test_spacegroup_name_variations(self): class TestSpaceGroupWithData: """Tests for SpaceGroup with real crystallographic data.""" - @pytest.mark.integration - def test_spacegroup_with_multiple_structures(self, cif_dir): - """Test SpaceGroup for multiple structures.""" - from torchref.model.model import Model - from torchref.symmetry import SpaceGroup - - cif_files = list(cif_dir.glob("*.cif"))[:3] - - for cif_file in cif_files: - model = Model() - model.load_cif(str(cif_file)) - - sg = SpaceGroup(model.spacegroup) - - # Should have valid matrices - assert sg.matrices is not None - assert sg.matrices.shape[0] >= 1 - assert torch.all(torch.isfinite(sg.matrices)) - @pytest.mark.integration def test_spacegroup_consistent_with_cell(self, sample_cif_file): """Test that SpaceGroup is consistent with unit cell.""" diff --git a/tests/unit/symmetry/test_symmetry.py b/tests/unit/symmetry/test_symmetry.py index aab7bcce..afca8c07 100644 --- a/tests/unit/symmetry/test_symmetry.py +++ b/tests/unit/symmetry/test_symmetry.py @@ -6,7 +6,6 @@ import pytest import torch -import torch.nn as nn class TestSpaceGroupInitialization: @@ -131,7 +130,9 @@ def test_rotation_matrices_determinant(self): for i in range(sg.matrices.shape[0]): det = torch.linalg.det(sg.matrices[i]) - assert torch.isclose(torch.abs(det), torch.tensor(1.0, dtype=det.dtype), atol=1e-5) + assert torch.isclose( + torch.abs(det), torch.tensor(1.0, dtype=det.dtype), atol=1e-5 + ) class TestSpaceGroupApplication: @@ -189,10 +190,10 @@ def test_spacegroup_cpu(self): """Test SpaceGroup on CPU.""" from torchref.symmetry import SpaceGroup - sg = SpaceGroup("P21", device=torch.device('cpu')) + sg = SpaceGroup("P21", device=torch.device("cpu")) - assert sg.matrices.device.type == 'cpu' - assert sg.translations.device.type == 'cpu' + assert sg.matrices.device.type == "cpu" + assert sg.translations.device.type == "cpu" @pytest.mark.unit @pytest.mark.gpu @@ -220,7 +221,23 @@ class TestSpaceGroupMapping: """Tests for space group name mapping.""" @pytest.mark.unit - @pytest.mark.parametrize("sg_name", ["P1", "P21", "P212121", "C2", "P21212"]) + @pytest.mark.parametrize( + "sg_name", + [ + "P1", + "P21", + "P212121", + "C2", + "P21212", + "P 1", + "P 21", + "P 21 21 21", + "P 43 21 2", + "P 3 2 1", + "P 6 2 2", + "P 2 3", + ], + ) def test_common_spacegroups(self, sg_name): """Test common crystallographic space groups.""" from torchref.symmetry import SpaceGroup