pino-source-code / tests /test_upload_data_split.py
Matthew Ford
Add Fraterworks ingestion and upload data split support
2671b56
Raw
History Blame Contribute Delete
4.06 kB
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"])