diff --git a/docs/changelog.rst b/docs/changelog.rst index 83a5728c..e2da9dda 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -5,6 +5,11 @@ Changelog Unreleased ---------- - Loss aggregation uses the configured floating-point dtype in eager and compiled execution, including empty and zero-weight aggregates. +- 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. +- 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/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 fd283106..3be67e2f 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 @@ -38,6 +39,34 @@ 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` | +| 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 +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/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/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..282c5527 --- /dev/null +++ b/tests/fixtures/README.md @@ -0,0 +1,46 @@ +# 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 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 | +| `numerical.py` | Synthetic tensors and factories | Imported only by `unit/conftest.py`; function | +| `functional.py` | Read-only `shared_model_ft` | Imported only by `functional/conftest.py`; module | + +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. + +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 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 +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. + +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..390ea678 --- /dev/null +++ b/tests/fixtures/files.py @@ -0,0 +1,89 @@ +"""Select sample paths and matching model/reflection pairs without loading them.""" + +from pathlib import Path + +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: + """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") diff --git a/tests/fixtures/functional.py b/tests/fixtures/functional.py new file mode 100644 index 00000000..5ee204a4 --- /dev/null +++ b/tests/fixtures/functional.py @@ -0,0 +1,25 @@ +"""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.model import ModelFT + + +@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)) 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..3eb7628e --- /dev/null +++ b/tests/fixtures/objects.py @@ -0,0 +1,152 @@ +"""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 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.""" + 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 _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..ba6caf3b 100644 --- a/tests/functional/conftest.py +++ b/tests/functional/conftest.py @@ -1,6 +1,3 @@ -""" -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 shared_model_ft # noqa: F401 diff --git a/tests/functional/test_io_functional.py b/tests/functional/test_io_functional.py deleted file mode 100644 index 9cdef611..00000000 --- a/tests/functional/test_io_functional.py +++ /dev/null @@ -1,311 +0,0 @@ -""" -Functional tests for I/O operations. - -Tests file loading and data processing with real crystallographic data. -""" - -import pytest -import torch -import numpy as np - - -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 - - @pytest.mark.integration - def test_cif_atom_properties(self, sample_cif_file): - """Test that atom properties are correctly loaded from CIF.""" - from torchref.model.model import Model - - model = Model() - model.load_cif(str(sample_cif_file)) - - 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, sample_cif_file): - """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() - - # 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.""" - - @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 - - @pytest.mark.integration - def test_mtz_data_properties(self, sample_mtz_file): - """Test that MTZ data properties are correctly loaded.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - # 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, sample_mtz_file): - """Test that resolution range is computed correctly.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - # 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.""" - - @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 - - -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.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"])) - - 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.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 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, sample_mtz_file): - """Test that French-Wilson conversion is applied.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - # 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, 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 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, sample_mtz_file): - """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: - # 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, sample_mtz_file): - """Test masks() method returns valid mask.""" - from torchref.io import ReflectionData - - data = ReflectionData() - data.load_mtz(str(sample_mtz_file)) - - 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_loss_weighting_functional.py b/tests/functional/test_loss_weighting_functional.py deleted file mode 100644 index 04e4f115..00000000 --- a/tests/functional/test_loss_weighting_functional.py +++ /dev/null @@ -1,196 +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 total.item() == pytest.approx(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 total.item() == pytest.approx(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 - assert total.item() == pytest.approx(2.5) - - 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/functional/test_model_ft_functional.py b/tests/functional/test_model_ft_functional.py index b8919254..ce9b5e71 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 @@ -29,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 @@ -53,45 +39,20 @@ 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) + 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 -class TestModelFTParametrization: - """Test ModelFT parametrization with real structures.""" - - def test_parametrization_built(self, sample_cif_file): - """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)) - - # Parametrization should be set - assert model.parametrization is not None - - def test_scattering_factors_available(self, sample_cif_file): - """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)) - - # 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 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,35 +68,33 @@ 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)) + + 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 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 @@ -145,155 +104,25 @@ def test_map_symmetry_available(self, sample_cif_file): 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, sample_cif_file, 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)) - - 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, sample_cif_file): - """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)) - - # 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.""" - - 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 TestModelFTCaching: - """Test ModelFT caching mechanism.""" - - def test_cache_initialization(self, sample_cif_file): - """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)) - - # Mixin cache should start empty (lazily initialized) - assert getattr(model, "_fwd_cached_output", None) is None - - def test_cache_usage(self, sample_cif_file): - """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)) - - # 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, sample_cif_file): - """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)) - - 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, 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 # 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) @@ -301,32 +130,35 @@ def test_fractional_to_cartesian(self, sample_cif_file): @pytest.mark.integration -class TestModelFTAnisoHandling: - """Test ModelFT handling of anisotropic parameters.""" - - def test_access_aniso_atoms(self, sample_cif_file): - """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)) - - # Check if aniso is available - if hasattr(model, 'get_aniso') or hasattr(model, 'aniso'): - # Structure has aniso - pass - - def test_isotropic_b_factors(self, sample_cif_file): - """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)) - - # 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_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/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/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/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/integration/test_io_cif.py b/tests/integration/test_io_cif.py index ff288cc5..2bd47fa1 100644 --- a/tests/integration/test_io_cif.py +++ b/tests/integration/test_io_cif.py @@ -6,116 +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 - - -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 +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 TestCIFSaving: @@ -125,18 +46,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 +66,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 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/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/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/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/refinement/test_loss_state.py b/tests/unit/refinement/test_loss_state.py index 553ff43e..dba1b464 100644 --- a/tests/unit/refinement/test_loss_state.py +++ b/tests/unit/refinement/test_loss_state.py @@ -104,7 +104,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) @@ -196,6 +198,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 @@ -244,12 +247,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): @@ -267,6 +269,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.""" @@ -320,15 +338,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 ec581ec3..d314e624 100644 --- a/tests/unit/refinement/test_loss_weighting.py +++ b/tests/unit/refinement/test_loss_weighting.py @@ -1,54 +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 - assert total.item() == pytest.approx(2.5) class TestDefaultGroupWeights: 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])) 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 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 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(()) + )