| from __future__ import annotations |
|
|
| import pytest |
|
|
| from pino.upload_data import ( |
| analyze_molecule_graph, |
| create_hub_excluded_component_split_from_records, |
| create_molecule_disjoint_split_from_records, |
| ) |
|
|
|
|
| def _record(cases: list[str], *, is_control: bool = False) -> dict: |
| return { |
| "is_control": is_control, |
| "metadata": { |
| "initial_components": [ |
| {"cas": "64-17-5", "weight_fraction": 0.85}, |
| *[ |
| {"cas": cas, "weight_fraction": 0.15 / len(cases)} |
| for cas in cases |
| ], |
| ] |
| }, |
| } |
|
|
|
|
| def _record_with_components(components: list[dict], *, formula_id: str = "F") -> dict: |
| return { |
| "formula_id": formula_id, |
| "metadata": {"initial_components": components}, |
| } |
|
|
|
|
| def test_molecule_disjoint_split_has_no_active_cas_overlap() -> None: |
| records = [ |
| _record(["100-00-1"]), |
| _record(["100-00-2"]), |
| _record(["100-00-3"], is_control=True), |
| _record(["100-00-4"], is_control=True), |
| _record(["100-00-1", "100-00-3"]), |
| _record(["100-00-2", "100-00-4"]), |
| ] |
|
|
| split = create_molecule_disjoint_split_from_records(records, train_ratio=0.5, seed=1) |
|
|
| assert set(split["train_compounds"]).isdisjoint(split["validation_compounds"]) |
| assert split["train"] |
| assert split["validation"] |
| assert len(split["excluded_indices"]) == 2 |
|
|
|
|
| def test_molecule_disjoint_split_rejects_invalid_ratio() -> None: |
| with pytest.raises(ValueError, match="train_ratio"): |
| create_molecule_disjoint_split_from_records([_record(["100-00-1"])], train_ratio=1.0) |
|
|
|
|
| def test_canonical_split_merges_smiles_and_cas_identity() -> None: |
| records = [ |
| _record_with_components( |
| [{"cas": "80-56-8", "smiles": "CC1=CCC2CC1C2(C)C", "weight_fraction": 1.0}], |
| formula_id="pinene_cas", |
| ), |
| _record_with_components( |
| [{"cas": "SMILES:CC1=CCC2CC1C2(C)C", "weight_fraction": 1.0}], |
| formula_id="pinene_smiles", |
| ), |
| _record_with_components( |
| [{"cas": "78-70-6", "smiles": "CC(=CCCC(C)(C=C)O)C", "weight_fraction": 1.0}], |
| formula_id="linalool", |
| ), |
| ] |
|
|
| report = analyze_molecule_graph(records) |
| duplicates = report["duplicate_identities_after_canonicalization"] |
|
|
| assert any( |
| "ALPHA-PINENE" in " ".join(item["aliases"]).upper() |
| or "80-56-8" in " ".join(item["aliases"]) |
| for item in duplicates |
| ) |
| assert report["n_molecules"] == 2 |
|
|
|
|
| def test_ethanol_only_control_is_excluded_from_molecule_split() -> None: |
| records = [ |
| _record_with_components( |
| [{"cas": "64-17-5", "smiles": "CCO", "weight_fraction": 1.0}], |
| formula_id="CONTROL_64-17-5", |
| ), |
| _record(["100-00-1"]), |
| _record(["100-00-2"]), |
| ] |
|
|
| split = create_molecule_disjoint_split_from_records(records, train_ratio=0.5, seed=1) |
|
|
| assert 0 in split["excluded_indices"] |
| assert records[0] not in split["train"] |
| assert records[0] not in split["validation"] |
|
|
|
|
| def test_hub_excluded_component_split_allows_hub_overlap_without_meaningful_leakage() -> None: |
| hub = {"cas": "25265-71-8", "smiles": "CC(CO)OC(C)CO", "weight_fraction": 0.8} |
| records = [ |
| _record_with_components([hub, {"cas": "100-00-1", "weight_fraction": 0.2}], formula_id="a1"), |
| _record_with_components([hub, {"cas": "100-00-2", "weight_fraction": 0.2}], formula_id="b1"), |
| _record_with_components([hub, {"cas": "100-00-2", "weight_fraction": 0.2}], formula_id="b2"), |
| _record_with_components([hub, {"cas": "100-00-3", "weight_fraction": 0.2}], formula_id="c1"), |
| ] |
|
|
| split = create_hub_excluded_component_split_from_records( |
| records, |
| train_ratio=0.5, |
| seed=3, |
| hub_threshold=0.5, |
| ) |
|
|
| assert split["validation"] |
| assert split["train"] |
| assert split["hub_compounds"] |
| assert set(split["train_compounds"]).isdisjoint(split["validation_compounds"]) |
|
|