diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 00000000..5acda731 --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,52 @@ +name: ci + +on: + pull_request: + push: + branches: [main] + workflow_dispatch: + +concurrency: + group: ci-${{ github.workflow }}-${{ github.ref }} + cancel-in-progress: true + +jobs: + test: + name: pytest (py${{ matrix.python }} / ${{ matrix.os }}) + runs-on: ${{ matrix.os }} + strategy: + fail-fast: false + matrix: + os: [ubuntu-latest] + python: ["3.10"] + + steps: + - uses: actions/checkout@v4 + + - uses: actions/setup-python@v5 + with: + python-version: ${{ matrix.python }} + + - name: Install system libraries + run: | + sudo apt-get update + sudo apt-get install -y libcairo2 libxrender1 libglib2.0-0 + + - name: Cache pip downloads + uses: actions/cache@v4 + with: + path: ~/.cache/pip + key: pip-${{ runner.os }}-py${{ matrix.python }}-${{ hashFiles('requirements-ci.txt', 'pyproject.toml') }} + restore-keys: | + pip-${{ runner.os }}-py${{ matrix.python }}- + + - name: Install CI dependencies + run: | + python -m pip install --upgrade pip wheel + pip install -r requirements-ci.txt + + - name: Install ms_pred (editable, no extra deps) + run: pip install -e . --no-deps + + - name: Run pytest + run: pytest -ra --durations=10 --maxfail=5 diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 00000000..d7c443a5 --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,3 @@ +[build-system] +requires = ["setuptools>=61", "wheel", "cython>=3", "numpy<2"] +build-backend = "setuptools.build_meta" diff --git a/pytest.ini b/pytest.ini new file mode 100644 index 00000000..33772184 --- /dev/null +++ b/pytest.ini @@ -0,0 +1,12 @@ +[pytest] +testpaths = tests +addopts = -ra --strict-markers +markers = + slow: takes more than 5s; excluded from default CI run + gpu: requires CUDA; excluded from CI + integration: cross-module / end-to-end +filterwarnings = + ignore::DeprecationWarning:rdkit.* + ignore::DeprecationWarning:pkg_resources + ignore::UserWarning:dgl.* + ignore::FutureWarning:torch.* diff --git a/requirements-ci.txt b/requirements-ci.txt new file mode 100644 index 00000000..30122896 --- /dev/null +++ b/requirements-ci.txt @@ -0,0 +1,48 @@ +# CPU-only dependency set for GitHub Actions CI. +# Source of truth for the CI install; runtime users should still follow README. + +# ---- PyTorch (CPU) ---- +--extra-index-url https://download.pytorch.org/whl/cpu +torch==2.4.0 +torchvision==0.19.0 +torchaudio==2.4.0 + +# ---- PyG extension wheels (CPU build) ---- +--find-links https://data.pyg.org/whl/torch-2.4.0+cpu.html +torch-scatter==2.1.2 +torch-sparse==0.6.18 +torch-geometric + +# ---- DGL (CPU build, py3.10 manylinux) ---- +https://data.dgl.ai/wheels/torch-2.4/dgl-2.4.0-cp310-cp310-manylinux1_x86_64.whl + +# ---- Pinned for py3.10 compatibility ---- +numpy<2 +scikit-learn>=1.4,<1.8 + +# ---- Core runtime deps mirrored from requirements.txt + import surface ---- +CairoSVG +cython>=3 +einops +h5py +ipdb +lightgbm +matplotlib +msbuddy +omegaconf +packaging +pandas +pathos +pubchempy +pytorch-lightning>=2.0 +rdkit +ray[tune] +seaborn +tqdm + +# pygmtools is a git source; pin to a commit so CI is reproducible +pygmtools @ git+https://github.com/Thinklab-SJTU/pygmtools.git@main + +# ---- Test-only ---- +pytest>=8 +pytest-xdist diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 00000000..0797c722 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,65 @@ +"""Shared pytest fixtures and global setup for the ms-pred test suite.""" + +import random + +import numpy as np +import pytest + + +def _seed_all(seed: int = 0) -> None: + random.seed(seed) + np.random.seed(seed) + try: + import torch + except ImportError: + return + torch.manual_seed(seed) + + +_seed_all(0) + +try: + from rdkit import RDLogger + + RDLogger.DisableLog("rdApp.*") +except ImportError: + pass + + +@pytest.fixture(autouse=True) +def _reseed_each_test(): + _seed_all(0) + yield + + +@pytest.fixture +def known_molecules(): + """SMILES, formula, exact monoisotopic mass for a small reference set. + + Masses are RDKit ExactMolWt values rounded to 4 decimals; tests should + compare with a tolerance, not by exact equality. + """ + return { + "water": ("O", "H2O", 18.0106), + "methane": ("C", "CH4", 16.0313), + "ethanol": ("CCO", "C2H6O", 46.0419), + "glucose": ("OC[C@H]1OC(O)[C@H](O)[C@@H](O)[C@@H]1O", "C6H12O6", 180.0634), + "caffeine": ("CN1C=NC2=C1C(=O)N(C)C(=O)N2C", "C8H10N4O2", 194.0804), + "aspirin": ("CC(=O)OC1=CC=CC=C1C(=O)O", "C9H8O4", 180.0423), + "fluorobenzene": ("Fc1ccccc1", "C6H5F", 96.0375), + } + + +@pytest.fixture +def tiny_spectrum(): + """A 5-peak spectrum: 2-D array with columns [m/z, intensity].""" + return np.array( + [ + [50.0, 0.10], + [100.0, 0.80], + [150.0, 0.45], + [200.0, 1.00], + [250.0, 0.20], + ], + dtype=np.float64, + ) diff --git a/tests/test_smoke_imports.py b/tests/test_smoke_imports.py new file mode 100644 index 00000000..b783424c --- /dev/null +++ b/tests/test_smoke_imports.py @@ -0,0 +1,40 @@ +"""Smoke tests: confirm every subpackage's anchor module imports cleanly. + +A failure here usually means a dependency rename, a top-level NameError, or +that one branch's refactor broke the import graph. +""" + +import importlib + +import pytest + +ANCHOR_MODULES = [ + "ms_pred", + "ms_pred.common", + "ms_pred.common.chem_utils", + "ms_pred.common.misc_utils", + "ms_pred.common.fingerprint", + "ms_pred.common.splitter", + "ms_pred.common.parallel_utils", + "ms_pred.common.plot_utils", + "ms_pred.common.denoising_utils", + "ms_pred.nn_utils", + "ms_pred.nn_utils.form_embedder", + "ms_pred.massformer_pred", + "ms_pred.dag_pred.dag_data", + "ms_pred.ffn_pred.ffn_data", + "ms_pred.gnn_pred.gnn_data", + "ms_pred.scarf_pred.scarf_data", + "ms_pred.marason.dag_data", + "ms_pred.magma.fragmentation", + "ms_pred.retrieval.bootstrap_metrics", + "ms_pred.autoregr_gen.autoregr_data", + "ms_pred.graff_ms.graff_ms_data", + "ms_pred.molnetms.molnetms_data", +] + + +@pytest.mark.parametrize("module_name", ANCHOR_MODULES) +def test_import_anchor_module(module_name: str) -> None: + module = importlib.import_module(module_name) + assert module is not None