From 8f8f28022e974129088caafa388138c4da01f40a Mon Sep 17 00:00:00 2001 From: Jonathan Roy Date: Sun, 20 Sep 2026 22:51:18 -0400 Subject: [PATCH 1/4] Restore PredSpecDB-aware add_dag_intens.py (from 02e6ade) Restores RuiXiWangTW's fix "fix magma inten" (02e6ade) for data_scripts/dag/add_dag_intens.py. The content of that fix was overwritten in 7b59014 ("add run scripts"), which returned this file to its earlier version; the other files touched by 02e6ade are unaffected. Without it, `add_dag_intens.py --magma-output` (as called by run_scripts/glacier/add_inten.sh) cannot read the PredSpecDB magma_tree.hdf5 written by the current run_magma.py: names are matched by Path(name).stem, so no entries match and the script stops with "ValueError: Empty list to process!". The mode without --magma-output, used by the ICEBERG and MARASON run scripts, produces identical output with either version of the file. Co-Authored-By: Claude Fable 5.1 --- data_scripts/dag/add_dag_intens.py | 244 +++++++++++++++++++++++------ 1 file changed, 196 insertions(+), 48 deletions(-) diff --git a/data_scripts/dag/add_dag_intens.py b/data_scripts/dag/add_dag_intens.py index 460283cb..72a4fe2d 100644 --- a/data_scripts/dag/add_dag_intens.py +++ b/data_scripts/dag/add_dag_intens.py @@ -16,6 +16,67 @@ import ms_pred.common as common +# Per-process handle caches so parallel workers reuse open HDF5 files instead +# of reopening them for every entry. +_PRED_DB_CACHE: dict = {} +_TRUE_H5_CACHE: dict = {} +_PRED_H5_CACHE: dict = {} + + +def _get_pred_db(path) -> common.PredSpecDB: + path = str(path) + db = _PRED_DB_CACHE.get(path) + if db is None: + db = common.PredSpecDB(path) + _PRED_DB_CACHE[path] = db + return db + + +def _get_true_h5(path) -> common.HDF5Dataset: + path = str(path) + h5 = _TRUE_H5_CACHE.get(path) + if h5 is None: + h5 = common.HDF5Dataset(path) + _TRUE_H5_CACHE[path] = h5 + return h5 + + +def _get_pred_h5(path) -> common.HDF5Dataset: + path = str(path) + h5 = _PRED_H5_CACHE.get(path) + if h5 is None: + h5 = common.HDF5Dataset(path) + _PRED_H5_CACHE[path] = h5 + return h5 + + +def _parse_legacy_magma_name(name: str): + """Parse legacy MAGMA tree keys like `spec_collision 30 eV.json`.""" + match = re.match( + r"^(.*?)_collision\s+([0-9]+\.?[0-9]*|nan)(?:\s*eV)?(?:\.json)?$", + name, + ) + if match is None: + return None + spec_id = match.group(1) + ce_raw = match.group(2) + ce_label = "nan" if ce_raw == "nan" else f"{float(ce_raw):.0f}" + return spec_id, ce_label + + +def _is_legacy_magma_h5(path: Path) -> bool: + """Detect legacy MAGMA HDF5 where top-level entries are JSON datasets.""" + pred_h5 = common.HDF5Dataset(path) + try: + for name in pred_h5.get_all_names(): + if name == "__predspec_manifest__": + continue + return _parse_legacy_magma_name(name) is not None and not hasattr(pred_h5[name], "keys") + return False + finally: + pred_h5.close() + + def get_args(): """get_args. """ @@ -33,52 +94,86 @@ def get_args(): "--magma-output", action="store_true", default=False, - help="If set, treat pred-dag-path as a MAGMA output HDF5 and add intensities to it." + help="If set, treat pred-dag-path as a MAGMA output PredSpecDB and add " + "gold intensities to it, writing JSON trees consumable by GLACIER." ) return parser.parse_args() +def _extract_raw_spec(true_dag_h5: common.HDF5Dataset, true_dag_name: str): + """Pull the gold (mass, intensity) peak list from a true DAG entry.""" + if true_dag_name not in true_dag_h5: + return None + true_dag = json.loads(true_dag_h5.read_str(true_dag_name)) + true_tbl = true_dag.get("output_tbl") + if not true_tbl or "mono_mass" not in true_tbl or "rel_inten" not in true_tbl: + return None + return [list(pair) for pair in zip(true_tbl["mono_mass"], true_tbl["rel_inten"])] + + def relabel_tree( - pred_dag_db: Path|common.PredSpecDB, + pred_dag_db: Path | common.PredSpecDB, true_dag_h5: Path, pred_dag_name: str, true_dag_name: str, out_dag_name: str, collision_energy: str, + remark=None, magma_output: bool = False, -) -> Tuple[str, str]: - """relabel_tree.""" - true_dag_h5 = common.HDF5Dataset(true_dag_h5) + legacy_magma_json: bool = False, +) -> Tuple[str, object]: + """relabel_tree. - if not true_dag_name in true_dag_h5: + Attach the gold spectrum (``raw_spec``) from the true DAG to a predicted / + MAGMA-annotated DAG. + + When ``magma_output`` is set, ``pred_dag_db`` points at a MAGMA ``PredSpecDB`` + (binary ``MassSpec`` arrays). We read the ``MassSpec``, expand its integer + fragments, and emit a JSON tree with the keys GLACIER's ``featurize_tree`` + expects (``root_canonical_smiles``, ``adduct``, ``collision_energy``, + ``frags`` as a dict of ``{"frag": int}``, and ``raw_spec``). + """ + true_h5 = _get_true_h5(true_dag_h5) + raw_spec = _extract_raw_spec(true_h5, true_dag_name) + if raw_spec is None: return None + if magma_output: - pred_dag_h5 = common.HDF5Dataset(pred_dag_db) - pred_dag = json.loads(pred_dag_h5.read_str(pred_dag_name)) - assert 'root_canonical_smiles' in pred_dag - assert 'frags' in pred_dag - assert 'collision_energy' in pred_dag - assert 'adduct' in pred_dag + if legacy_magma_json: + pred_h5 = _get_pred_h5(pred_dag_db) + if pred_dag_name not in pred_h5: + return None + pred_dag = json.loads(pred_h5.read_str(pred_dag_name)) + if not isinstance(pred_dag, dict): + return None + pred_dag["raw_spec"] = raw_spec + if collision_energy is not None: + pred_dag["collision_energy"] = float(collision_energy) + return out_dag_name, json.dumps(pred_dag, indent=2) + + pred_db = _get_pred_db(pred_dag_db) + spec = pred_db.read(pred_dag_name, collision_energy, remark) + if spec.root_canonical_smiles is None or not spec.has_frags: + return None + int_frags = spec.int_frags or [] + frags = {str(i): {"frag": int(frag)} for i, frag in enumerate(int_frags)} + out_dict = { + "root_canonical_smiles": spec.root_canonical_smiles, + "adduct": spec.adduct, + "collision_energy": float(spec.collision_energy), + "frags": frags, + "raw_spec": raw_spec, + } + return out_dag_name, json.dumps(out_dict, indent=2) else: - if not isinstance(pred_dag_db, common.PredSpecDB): - pred_dag_db = common.PredSpecDB(pred_dag_db) - pred_dag = pred_dag_db.read(pred_dag_name, collision_energy) + pred_db = pred_dag_db if isinstance(pred_dag_db, common.PredSpecDB) else _get_pred_db(pred_dag_db) + pred_dag = pred_db.read(pred_dag_name, collision_energy) assert pred_dag.root_canonical_smiles is not None assert pred_dag.frags is not None assert pred_dag.collision_energy is not None assert pred_dag.adduct is not None - true_dag = json.loads(true_dag_h5.read_str(true_dag_name)) - true_tbl = true_dag["output_tbl"] - if true_tbl is None: - return None - raw_spec = list(zip(true_tbl["mono_mass"], true_tbl["rel_inten"])) - - if not magma_output: pred_dag.meta["raw_spec"] = raw_spec return out_dag_name, pred_dag - else: - pred_dag["raw_spec"] = raw_spec - return out_dag_name, json.dumps(pred_dag, indent=2) def main(): @@ -92,31 +187,68 @@ def main(): out_dag_path.parent.mkdir(exist_ok=True) if args.magma_output: - # Treat pred_dag_path as a MAGMA output HDF5, add intensities from true DAGs - pred_dag_h5 = common.HDF5Dataset(pred_dag_path) - pred_dag_names = pred_dag_h5.get_all_names() - # Do not close pred_dag_h5 here + # Treat pred_dag_path as a MAGMA output PredSpecDB. Each spec is stored + # once per collision energy, while the true DAGs are keyed by + # "{spec}_collision {ce}". Expand the PredSpecDB per collision energy and + # match each entry to its gold spectrum. true_dag_h5 = common.HDF5Dataset(true_dag_path) - true_dag_names = true_dag_h5.get_all_names() - # Do not close true_dag_h5 here - # Match by stem (remove .json if present) - pred_to_true = {Path(n).stem: n for n in pred_dag_names} - true_by_stem = {Path(n).stem: n for n in true_dag_names} - matched = [(pred_to_true[k], true_by_stem[k], pred_to_true[k]) for k in pred_to_true if k in true_by_stem] - arg_dicts = [ - { - "pred_dag_db": pred_dag_path, - "true_dag_h5": true_dag_path, - "pred_dag_name": pred_name, - "true_dag_name": true_name, - "out_dag_name": out_name, - "collision_energy": common.get_collision_energy(true_name), - "magma_output": True, - } - for pred_name, true_name, out_name in matched - ] - pred_dag_h5.close() + true_lookup = {} + for true_name in true_dag_h5.get_all_names(): + spec_id = common.rm_collision_str(true_name) + ce_label = common.get_collision_energy(true_name) + true_lookup[(spec_id, ce_label)] = true_name true_dag_h5.close() + + arg_dicts = [] + is_legacy_magma = _is_legacy_magma_h5(pred_dag_path) + if is_legacy_magma: + pred_h5 = common.HDF5Dataset(pred_dag_path) + for pred_name in tqdm(pred_h5.get_all_names()): + parsed = _parse_legacy_magma_name(pred_name) + if parsed is None: + continue + spec_id, ce_label = parsed + true_name = true_lookup.get((spec_id, ce_label)) + if true_name is None: + continue + arg_dicts.append( + { + "pred_dag_db": pred_dag_path, + "true_dag_h5": true_dag_path, + "pred_dag_name": pred_name, + "true_dag_name": true_name, + "out_dag_name": f"{spec_id}_collision {ce_label}", + "collision_energy": ce_label, + "remark": None, + "magma_output": True, + "legacy_magma_json": True, + } + ) + pred_h5.close() + else: + pred_db = common.PredSpecDB(pred_dag_path, h5_persistent=True) + for pred_name in tqdm(pred_db.get_all_names()): + spec_id = pred_name[5:] if pred_name.startswith("pred_") else pred_name + ces, remarks = pred_db.get_entries(pred_name) + for ce, remark in zip(ces, remarks): + ce_label = f"{float(ce):.0f}" + true_name = true_lookup.get((spec_id, ce_label)) + if true_name is None: + continue + arg_dicts.append( + { + "pred_dag_db": pred_dag_path, + "true_dag_h5": true_dag_path, + "pred_dag_name": pred_name, + "true_dag_name": true_name, + "out_dag_name": f"{spec_id}_collision {ce_label}", + "collision_energy": ce, + "remark": remark, + "magma_output": True, + "legacy_magma_json": False, + } + ) + pred_db.close() else: pred_dag_h5 = common.HDF5Dataset(pred_dag_path) pred_dag_name_set = set(pred_dag_h5.get_all_names()) @@ -152,9 +284,25 @@ def main(): def write_func(outs): out_db = common.PredSpecDB(out_dag_path, mode='w') for out in outs: + if out is None: + continue out_db.write(*out) out_db.close() + if len(arg_dicts) == 0: + print( + f"No matching entries found for pred DAGs ({pred_dag_path}) and true DAGs ({true_dag_path}). " + f"Creating empty output at {out_dag_path}." + ) + if args.magma_output: + out_h5 = common.HDF5Dataset(out_dag_path, mode='w') + out_h5.close() + else: + out_db = common.PredSpecDB(out_dag_path, mode='w') + out_db.close() + print("success!") + return + # Run wrapper_fn = lambda arg_dict: relabel_tree(**arg_dict) num_workers = args.num_workers From 79f29b589162cf61c89657845ece1d5a95b8601a Mon Sep 17 00:00:00 2001 From: Jonathan Roy Date: Sun, 20 Sep 2026 22:51:18 -0400 Subject: [PATCH 2/4] Make GLACIER lr_scheduler_step compatible with pytorch-lightning 2.x Lightning 2.x calls the hook as lr_scheduler_step(scheduler, metric), so the three-positional-argument signature raised "TypeError: JointModel.lr_scheduler_step() missing 1 required positional argument: 'metric'" after the first optimizer step of glacier/train_joint.py. Use the metric=None default already used by the other models in this repository. Co-Authored-By: Claude Fable 5.1 --- src/ms_pred/glacier/joint_model.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/ms_pred/glacier/joint_model.py b/src/ms_pred/glacier/joint_model.py index 419a72ff..833047a7 100644 --- a/src/ms_pred/glacier/joint_model.py +++ b/src/ms_pred/glacier/joint_model.py @@ -826,7 +826,7 @@ def _is_no_decay_param(name: str, param: torch.nn.Parameter) -> bool: lr_decay_rate=self.lr_decay_rate, warmup=self.warmup) return {"optimizer": optimizer, "lr_scheduler": {"scheduler": scheduler, "frequency": 1, "interval": "step"}} - def lr_scheduler_step(self, scheduler, optimizer_idx, metric): + def lr_scheduler_step(self, scheduler, optimizer_idx, metric=None): # fix lightning API mismatch for torch>=2.0 # For LambdaLR, just call step() without arguments scheduler.step() From f9ba718958c56f1a97d322614dee5bf32e31f658 Mon Sep 17 00:00:00 2001 From: Jonathan Roy Date: Sun, 20 Sep 2026 22:55:29 -0400 Subject: [PATCH 3/4] Add smoke tests for the GLACIER training path test_add_dag_intens_reads_magma_predspecdb runs run_magma.py, 01_assign_subformulae.py and add_dag_intens.py --magma-output on a four-molecule dataset in a temporary directory and checks that the PredSpecDB magma_tree.hdf5 yields one `_collision ` JSON tree per spectrum that the GLACIER IntenDataset can load. test_glacier_fast_dev_run_cpu runs one CPU training step of the GLACIER JointModel with pl.Trainer(fast_dev_run=True), which exercises the lr_scheduler_step hook under the installed pytorch-lightning. Both run on CPU in a few seconds. Co-Authored-By: Claude Fable 5.1 --- tests/test_glacier_training.py | 229 +++++++++++++++++++++++++++++++++ 1 file changed, 229 insertions(+) create mode 100644 tests/test_glacier_training.py diff --git a/tests/test_glacier_training.py b/tests/test_glacier_training.py new file mode 100644 index 00000000..4623ab9c --- /dev/null +++ b/tests/test_glacier_training.py @@ -0,0 +1,229 @@ +"""Smoke tests for the GLACIER training path. + +Two regressions are covered: + +1. ``data_scripts/dag/add_dag_intens.py --magma-output`` must be able to read + the ``PredSpecDB`` ``magma_tree.hdf5`` written by ``run_magma.py`` and emit + the ``_collision `` JSON trees consumed by the GLACIER dataset. +2. ``glacier.joint_model.JointModel`` must complete a training step, including + the ``lr_scheduler_step`` hook, under the installed pytorch-lightning. + +Everything runs on CPU in a temporary directory with a four-molecule dataset. +""" +import json +import runpy +import sys +from pathlib import Path + +import pandas as pd +import pytest +import pytorch_lightning as pl +from torch.utils.data import DataLoader + +import ms_pred.common as common +from ms_pred.glacier.dataset import IntenDataset, TreeProcessor +from ms_pred.glacier.joint_model import JointModel +from ms_pred.magma.fragmentation import FRAGMENT_ENGINE_PARAMS, FragmentEngine + +REPO_ROOT = Path(__file__).resolve().parents[1] +COLLISION_ENERGY = 20 + +# (spec, smiles, formula, [M+H]+ m/z, fragment m/z values with a known subformula) +TINY_SPECS = [ + ("tiny_caffeine", "Cn1cnc2c1c(=O)n(C)c(=O)n2C", "C8H10N4O2", 195.0877, [138.0662, 110.0713]), + ("tiny_paracetamol", "CC(=O)Nc1ccc(O)cc1", "C8H9NO2", 152.0706, [110.0600, 93.0335]), + ("tiny_nicotine", "CN1CCCC1c1cccnc1", "C10H14N2", 163.1230, [132.0808, 84.0808]), + ("tiny_aspirin", "CC(=O)Oc1ccccc1C(=O)O", "C9H8O4", 181.0495, [163.0390, 121.0284]), +] + + +def _peaks(precursor_mz, frag_mzs): + mzs = [precursor_mz] + list(frag_mzs) + intens = [1.0, 0.6, 0.3][: len(mzs)] + return list(zip(mzs, intens)) + + +def _labels_df() -> pd.DataFrame: + return pd.DataFrame( + { + "dataset": "tiny", + "spec": [s[0] for s in TINY_SPECS], + "ionization": "[M+H]+", + "formula": [s[2] for s in TINY_SPECS], + "smiles": [s[1] for s in TINY_SPECS], + "inchikey": [f"TINY{i:010d}" for i in range(len(TINY_SPECS))], + "instrument": "Orbitrap", + "collision_energies": f"['{COLLISION_ENERGY}']", + "precursor": [s[3] for s in TINY_SPECS], + } + ) + + +def _write_tiny_dataset(data_dir: Path) -> None: + """labels.tsv + spec_files.hdf5 in the layout the data scripts expect.""" + data_dir.mkdir(parents=True, exist_ok=True) + _labels_df().to_csv(data_dir / "labels.tsv", sep="\t", index=False) + spec_h5 = common.HDF5Dataset(data_dir / "spec_files.hdf5", mode="w") + for spec, smiles, formula, precursor_mz, frag_mzs in TINY_SPECS: + peak_str = "\n".join(f"{mz} {inten}" for mz, inten in _peaks(precursor_mz, frag_mzs)) + spec_h5.write_str( + f"{spec}.ms", + f">compound {spec}\n>formula {formula}\n>parentmass {precursor_mz}\n" + f">ionization [M+H]+\n>smiles {smiles}\n\n" + f">collision {COLLISION_ENERGY}\n{peak_str}\n", + ) + spec_h5.close() + + +def _run_script(monkeypatch, script: str, *args) -> None: + """Run a repository script in-process as ``__main__`` (avoids one interpreter start-up per stage).""" + monkeypatch.setattr(sys, "argv", [script, *[str(a) for a in args]]) + runpy.run_path(str(REPO_ROOT / script), run_name="__main__") + + +def _tree_processor() -> TreeProcessor: + return TreeProcessor( + pe_embed_k=0, + root_encode="graphormer", + embed_elem_group=True, + multi_hop_max_dist=3, + ) + + +def _inten_dataset(magma_h5_path: Path, tree_processor: TreeProcessor) -> IntenDataset: + """Same construction as glacier/train_joint.py.""" + magma_h5 = common.HDF5Dataset(magma_h5_path) + name_to_json = {Path(i).stem: i for i in magma_h5.get_all_names()} + magma_h5.close() + return IntenDataset( + _labels_df(), + magma_h5=magma_h5_path, + magma_map=name_to_json, + num_workers=0, + root_encode="graphormer", + embed_elem_group=True, + tree_processor=tree_processor, + datatype="HDF5", + ) + + +@pytest.mark.integration +def test_add_dag_intens_reads_magma_predspecdb(tmp_path, monkeypatch): + """run_magma.py -> 01_assign_subformulae.py -> add_dag_intens.py --magma-output.""" + data_dir = tmp_path / "tiny" + _write_tiny_dataset(data_dir) + magma_dir = data_dir / "magma_outputs" + out_h5_path = magma_dir / "magma_tree_with_inten.hdf5" + + # --debug selects the serial code path of each script (no worker pool) + _run_script( + monkeypatch, + "src/ms_pred/magma/run_magma.py", + "--spectra-dir", data_dir / "spec_files.hdf5", + "--output-dir", magma_dir, + "--spec-labels", data_dir / "labels.tsv", + "--max-peaks", 50, + "--ppm-diff", 20, + "--workers", 1, + "--debug", + ) + magma_db = common.PredSpecDB(magma_dir / "magma_tree.hdf5") + assert sorted(magma_db.get_all_names()) == sorted(s[0] for s in TINY_SPECS) + magma_db.close() + + _run_script( + monkeypatch, + "data_scripts/forms/01_assign_subformulae.py", + "--data-dir", data_dir, + "--labels-file", data_dir / "labels.tsv", + "--use-all", + "--output-dir-name", "no_subform.hdf5", + "--num-workers", 1, + "--debug", + ) + _run_script( + monkeypatch, + "data_scripts/dag/add_dag_intens.py", + "--pred-dag-path", magma_dir / "magma_tree.hdf5", + "--true-dag-path", data_dir / "subformulae" / "no_subform.hdf5", + "--out-dag-path", out_h5_path, + "--num-workers", 0, + "--magma-output", + ) + + out_h5 = common.HDF5Dataset(out_h5_path) + names = sorted(out_h5.get_all_names()) + assert names == sorted(f"{s[0]}_collision {COLLISION_ENERGY}" for s in TINY_SPECS) + for name in names: + tree = json.loads(out_h5.read_str(name)) + for key in ["root_canonical_smiles", "adduct", "collision_energy", "frags", "raw_spec"]: + assert key in tree, f"{name} is missing {key}" + assert tree["collision_energy"] == COLLISION_ENERGY + assert len(tree["frags"]) > 0 + assert all("frag" in frag for frag in tree["frags"].values()) + assert len(tree["raw_spec"]) > 0 + out_h5.close() + + # The output must be consumable by the GLACIER dataset as train_joint.py builds it + dataset = _inten_dataset(out_h5_path, _tree_processor()) + assert len(dataset) == len(TINY_SPECS) + item = dataset[0] + assert item["frag_targs"].shape[0] > 0 + assert len(item["inten_targs"]) > 0 + + +def test_glacier_fast_dev_run_cpu(tmp_path): + """One CPU training step, including the lr scheduler hook, under the installed Lightning.""" + # Hand-built trees in the format written by add_dag_intens.py --magma-output; + # the only fragment is the intact molecule, which keeps this test independent of MAGMa. + magma_h5_path = tmp_path / "magma_tree_with_inten.hdf5" + magma_h5 = common.HDF5Dataset(magma_h5_path, mode="w") + for spec, smiles, _, precursor_mz, frag_mzs in TINY_SPECS: + engine = FragmentEngine(mol_str=smiles, **FRAGMENT_ENGINE_PARAMS) + tree = { + "root_canonical_smiles": engine.smiles, + "adduct": "[M+H]+", + "collision_energy": float(COLLISION_ENERGY), + "frags": {"0": {"frag": (1 << engine.natoms) - 1}}, + "raw_spec": [list(p) for p in _peaks(precursor_mz, frag_mzs)], + } + magma_h5.write_str(f"{spec}_collision {COLLISION_ENERGY}", json.dumps(tree)) + magma_h5.close() + + tree_processor = _tree_processor() + dataset = _inten_dataset(magma_h5_path, tree_processor) + assert len(dataset) == len(TINY_SPECS) + loader = DataLoader(dataset, batch_size=2, shuffle=False, collate_fn=dataset.get_collate_fn()) + + model = JointModel( + hidden_size=32, + graphormer_layers=1, + frag_decoder_layers=1, + frag_encoder_layers=0, + inten_decoder_layers=1, + inten_encoder_layers=1, + node_feats=dataset.get_node_feats(), + edge_feats=tree_processor.get_edge_feats(), + multi_hop_max_dist=3, + max_breakpoints=8, + embed_adduct=True, + embed_collision=True, + embed_elem_group=True, + embed_instrument=True, + encode_forms=True, + enable_aux_loss=True, + warmup=2, + magma_warmup_steps=2, + magma_decay_steps=2, + ) + trainer = pl.Trainer( + accelerator="cpu", + devices=1, + fast_dev_run=True, + logger=False, + enable_checkpointing=False, + enable_progress_bar=False, + enable_model_summary=False, + ) + trainer.fit(model, loader, loader) + assert trainer.global_step == 1 From e43e4f04f68ce84d46a1dd58a104867c815da097 Mon Sep 17 00:00:00 2001 From: Runzhong Wang <18309862+rogerwwww@users.noreply.github.com> Date: Mon, 21 Sep 2026 09:55:27 -0400 Subject: [PATCH 4/4] Delete tests/test_glacier_training.py --- tests/test_glacier_training.py | 229 --------------------------------- 1 file changed, 229 deletions(-) delete mode 100644 tests/test_glacier_training.py diff --git a/tests/test_glacier_training.py b/tests/test_glacier_training.py deleted file mode 100644 index 4623ab9c..00000000 --- a/tests/test_glacier_training.py +++ /dev/null @@ -1,229 +0,0 @@ -"""Smoke tests for the GLACIER training path. - -Two regressions are covered: - -1. ``data_scripts/dag/add_dag_intens.py --magma-output`` must be able to read - the ``PredSpecDB`` ``magma_tree.hdf5`` written by ``run_magma.py`` and emit - the ``_collision `` JSON trees consumed by the GLACIER dataset. -2. ``glacier.joint_model.JointModel`` must complete a training step, including - the ``lr_scheduler_step`` hook, under the installed pytorch-lightning. - -Everything runs on CPU in a temporary directory with a four-molecule dataset. -""" -import json -import runpy -import sys -from pathlib import Path - -import pandas as pd -import pytest -import pytorch_lightning as pl -from torch.utils.data import DataLoader - -import ms_pred.common as common -from ms_pred.glacier.dataset import IntenDataset, TreeProcessor -from ms_pred.glacier.joint_model import JointModel -from ms_pred.magma.fragmentation import FRAGMENT_ENGINE_PARAMS, FragmentEngine - -REPO_ROOT = Path(__file__).resolve().parents[1] -COLLISION_ENERGY = 20 - -# (spec, smiles, formula, [M+H]+ m/z, fragment m/z values with a known subformula) -TINY_SPECS = [ - ("tiny_caffeine", "Cn1cnc2c1c(=O)n(C)c(=O)n2C", "C8H10N4O2", 195.0877, [138.0662, 110.0713]), - ("tiny_paracetamol", "CC(=O)Nc1ccc(O)cc1", "C8H9NO2", 152.0706, [110.0600, 93.0335]), - ("tiny_nicotine", "CN1CCCC1c1cccnc1", "C10H14N2", 163.1230, [132.0808, 84.0808]), - ("tiny_aspirin", "CC(=O)Oc1ccccc1C(=O)O", "C9H8O4", 181.0495, [163.0390, 121.0284]), -] - - -def _peaks(precursor_mz, frag_mzs): - mzs = [precursor_mz] + list(frag_mzs) - intens = [1.0, 0.6, 0.3][: len(mzs)] - return list(zip(mzs, intens)) - - -def _labels_df() -> pd.DataFrame: - return pd.DataFrame( - { - "dataset": "tiny", - "spec": [s[0] for s in TINY_SPECS], - "ionization": "[M+H]+", - "formula": [s[2] for s in TINY_SPECS], - "smiles": [s[1] for s in TINY_SPECS], - "inchikey": [f"TINY{i:010d}" for i in range(len(TINY_SPECS))], - "instrument": "Orbitrap", - "collision_energies": f"['{COLLISION_ENERGY}']", - "precursor": [s[3] for s in TINY_SPECS], - } - ) - - -def _write_tiny_dataset(data_dir: Path) -> None: - """labels.tsv + spec_files.hdf5 in the layout the data scripts expect.""" - data_dir.mkdir(parents=True, exist_ok=True) - _labels_df().to_csv(data_dir / "labels.tsv", sep="\t", index=False) - spec_h5 = common.HDF5Dataset(data_dir / "spec_files.hdf5", mode="w") - for spec, smiles, formula, precursor_mz, frag_mzs in TINY_SPECS: - peak_str = "\n".join(f"{mz} {inten}" for mz, inten in _peaks(precursor_mz, frag_mzs)) - spec_h5.write_str( - f"{spec}.ms", - f">compound {spec}\n>formula {formula}\n>parentmass {precursor_mz}\n" - f">ionization [M+H]+\n>smiles {smiles}\n\n" - f">collision {COLLISION_ENERGY}\n{peak_str}\n", - ) - spec_h5.close() - - -def _run_script(monkeypatch, script: str, *args) -> None: - """Run a repository script in-process as ``__main__`` (avoids one interpreter start-up per stage).""" - monkeypatch.setattr(sys, "argv", [script, *[str(a) for a in args]]) - runpy.run_path(str(REPO_ROOT / script), run_name="__main__") - - -def _tree_processor() -> TreeProcessor: - return TreeProcessor( - pe_embed_k=0, - root_encode="graphormer", - embed_elem_group=True, - multi_hop_max_dist=3, - ) - - -def _inten_dataset(magma_h5_path: Path, tree_processor: TreeProcessor) -> IntenDataset: - """Same construction as glacier/train_joint.py.""" - magma_h5 = common.HDF5Dataset(magma_h5_path) - name_to_json = {Path(i).stem: i for i in magma_h5.get_all_names()} - magma_h5.close() - return IntenDataset( - _labels_df(), - magma_h5=magma_h5_path, - magma_map=name_to_json, - num_workers=0, - root_encode="graphormer", - embed_elem_group=True, - tree_processor=tree_processor, - datatype="HDF5", - ) - - -@pytest.mark.integration -def test_add_dag_intens_reads_magma_predspecdb(tmp_path, monkeypatch): - """run_magma.py -> 01_assign_subformulae.py -> add_dag_intens.py --magma-output.""" - data_dir = tmp_path / "tiny" - _write_tiny_dataset(data_dir) - magma_dir = data_dir / "magma_outputs" - out_h5_path = magma_dir / "magma_tree_with_inten.hdf5" - - # --debug selects the serial code path of each script (no worker pool) - _run_script( - monkeypatch, - "src/ms_pred/magma/run_magma.py", - "--spectra-dir", data_dir / "spec_files.hdf5", - "--output-dir", magma_dir, - "--spec-labels", data_dir / "labels.tsv", - "--max-peaks", 50, - "--ppm-diff", 20, - "--workers", 1, - "--debug", - ) - magma_db = common.PredSpecDB(magma_dir / "magma_tree.hdf5") - assert sorted(magma_db.get_all_names()) == sorted(s[0] for s in TINY_SPECS) - magma_db.close() - - _run_script( - monkeypatch, - "data_scripts/forms/01_assign_subformulae.py", - "--data-dir", data_dir, - "--labels-file", data_dir / "labels.tsv", - "--use-all", - "--output-dir-name", "no_subform.hdf5", - "--num-workers", 1, - "--debug", - ) - _run_script( - monkeypatch, - "data_scripts/dag/add_dag_intens.py", - "--pred-dag-path", magma_dir / "magma_tree.hdf5", - "--true-dag-path", data_dir / "subformulae" / "no_subform.hdf5", - "--out-dag-path", out_h5_path, - "--num-workers", 0, - "--magma-output", - ) - - out_h5 = common.HDF5Dataset(out_h5_path) - names = sorted(out_h5.get_all_names()) - assert names == sorted(f"{s[0]}_collision {COLLISION_ENERGY}" for s in TINY_SPECS) - for name in names: - tree = json.loads(out_h5.read_str(name)) - for key in ["root_canonical_smiles", "adduct", "collision_energy", "frags", "raw_spec"]: - assert key in tree, f"{name} is missing {key}" - assert tree["collision_energy"] == COLLISION_ENERGY - assert len(tree["frags"]) > 0 - assert all("frag" in frag for frag in tree["frags"].values()) - assert len(tree["raw_spec"]) > 0 - out_h5.close() - - # The output must be consumable by the GLACIER dataset as train_joint.py builds it - dataset = _inten_dataset(out_h5_path, _tree_processor()) - assert len(dataset) == len(TINY_SPECS) - item = dataset[0] - assert item["frag_targs"].shape[0] > 0 - assert len(item["inten_targs"]) > 0 - - -def test_glacier_fast_dev_run_cpu(tmp_path): - """One CPU training step, including the lr scheduler hook, under the installed Lightning.""" - # Hand-built trees in the format written by add_dag_intens.py --magma-output; - # the only fragment is the intact molecule, which keeps this test independent of MAGMa. - magma_h5_path = tmp_path / "magma_tree_with_inten.hdf5" - magma_h5 = common.HDF5Dataset(magma_h5_path, mode="w") - for spec, smiles, _, precursor_mz, frag_mzs in TINY_SPECS: - engine = FragmentEngine(mol_str=smiles, **FRAGMENT_ENGINE_PARAMS) - tree = { - "root_canonical_smiles": engine.smiles, - "adduct": "[M+H]+", - "collision_energy": float(COLLISION_ENERGY), - "frags": {"0": {"frag": (1 << engine.natoms) - 1}}, - "raw_spec": [list(p) for p in _peaks(precursor_mz, frag_mzs)], - } - magma_h5.write_str(f"{spec}_collision {COLLISION_ENERGY}", json.dumps(tree)) - magma_h5.close() - - tree_processor = _tree_processor() - dataset = _inten_dataset(magma_h5_path, tree_processor) - assert len(dataset) == len(TINY_SPECS) - loader = DataLoader(dataset, batch_size=2, shuffle=False, collate_fn=dataset.get_collate_fn()) - - model = JointModel( - hidden_size=32, - graphormer_layers=1, - frag_decoder_layers=1, - frag_encoder_layers=0, - inten_decoder_layers=1, - inten_encoder_layers=1, - node_feats=dataset.get_node_feats(), - edge_feats=tree_processor.get_edge_feats(), - multi_hop_max_dist=3, - max_breakpoints=8, - embed_adduct=True, - embed_collision=True, - embed_elem_group=True, - embed_instrument=True, - encode_forms=True, - enable_aux_loss=True, - warmup=2, - magma_warmup_steps=2, - magma_decay_steps=2, - ) - trainer = pl.Trainer( - accelerator="cpu", - devices=1, - fast_dev_run=True, - logger=False, - enable_checkpointing=False, - enable_progress_bar=False, - enable_model_summary=False, - ) - trainer.fit(model, loader, loader) - assert trainer.global_step == 1