diff --git a/src/modernmolbert/eval/contributed_datasets.py b/src/modernmolbert/eval/contributed_datasets.py index 2969c02..b484cd0 100644 --- a/src/modernmolbert/eval/contributed_datasets.py +++ b/src/modernmolbert/eval/contributed_datasets.py @@ -103,6 +103,63 @@ def load_example_activity_dataset(*, root: str | Path) -> EvalDataset: ) +def load_esol(*, root: str | Path) -> EvalDataset: + root = Path(root) + task_names = ["target"] + + train = pd.read_csv(root / "train.csv") + valid_path = root / "valid.csv" + valid = pd.read_csv(valid_path) if valid_path.exists() else None + test = pd.read_csv(root / "test.csv") + + train = _drop_missing_labels(train, task_names=task_names) + test = _drop_missing_labels(test, task_names=task_names) + if valid is not None: + valid = _drop_missing_labels(valid, task_names=task_names) + + return make_eval_dataset_from_splits( + name="esol", + task_type="regression", + task_names=task_names, + train=train, + valid=valid, + test=test, + smiles_column="smiles", + metadata={"source": "MoleculeNet ESOL", "missing_label_policy": "Dropped"}, + ) + + +def load_clintox(*, root: str | Path) -> EvalDataset: + root = Path(root) + task_names = ["FDA_APPROVED"] + + train = pd.read_csv(root / "train.csv") + valid_path = root / "valid.csv" + valid = pd.read_csv(valid_path) if valid_path.exists() else None + test = pd.read_csv(root / "test.csv") + + train = _drop_missing_labels(train, task_names=task_names) + test = _drop_missing_labels(test, task_names=task_names) + if valid is not None: + valid = _drop_missing_labels(valid, task_names=task_names) + + _validate_binary_labels(train, task_name="FDA_APPROVED", split_name="train") + _validate_binary_labels(test, task_name="FDA_APPROVED", split_name="test") + if valid is not None: + _validate_binary_labels(valid, task_name="FDA_APPROVED", split_name="valid") + + return make_eval_dataset_from_splits( + name="clintox", + task_type="classification", + task_names=task_names, + train=train, + valid=valid, + test=test, + smiles_column="smiles", + metadata={"source": "MoleculeNet ClinTox", "missing_label_policy": "Dropped"}, + ) + + def register_contributed_datasets() -> None: """Register project-maintained contributed datasets. @@ -126,4 +183,30 @@ def register_contributed_datasets() -> None: # ) # ) + register_dataset( + DatasetSpec( + name="esol", + task_type="regression", + task_names=("target",), + loader=load_esol, + description="ESOL water solubility regression from MoleculeNet.", + source="https://moleculenet.org/datasets-1", + citation="Wu et al. MoleculeNet: a benchmark for molecular machine learning. Chemical Science (2018).", + license="MIT", + ) + ) + + register_dataset( + DatasetSpec( + name="clintox", + task_type="classification", + task_names=("FDA_APPROVED",), + loader=load_clintox, + description="ClinTox FDA approval classification from MoleculeNet.", + source="https://moleculenet.org/datasets-1", + citation="Wu et al. MoleculeNet: a benchmark for molecular machine learning. Chemical Science (2018).", + license="MIT", + ) + ) + return None diff --git a/tests/test_eval_clintox.py b/tests/test_eval_clintox.py new file mode 100644 index 0000000..f4d2505 --- /dev/null +++ b/tests/test_eval_clintox.py @@ -0,0 +1,21 @@ +from pathlib import Path +import pandas as pd +from modernmolbert.eval.contributed_datasets import load_clintox + + +def test_load_clintox(tmp_path: Path) -> None: + root = tmp_path / "clintox" + root.mkdir() + + pd.DataFrame({"smiles": ["CCO", "CCN", "CCC"], "FDA_APPROVED": [0, 1, None]}).to_csv( + root / "train.csv", index=False + ) + pd.DataFrame({"smiles": ["CCCl", "CCBr"], "FDA_APPROVED": [0, 1]}).to_csv( + root / "test.csv", index=False + ) + + dataset = load_clintox(root=root) + dataset.check() + + assert dataset.name == "clintox" + assert len(dataset.train) == 2 diff --git a/tests/test_eval_esol.py b/tests/test_eval_esol.py new file mode 100644 index 0000000..15e2f12 --- /dev/null +++ b/tests/test_eval_esol.py @@ -0,0 +1,21 @@ +from pathlib import Path +import pandas as pd +from modernmolbert.eval.contributed_datasets import load_esol + + +def test_load_esol(tmp_path: Path) -> None: + root = tmp_path / "esol" + root.mkdir() + + pd.DataFrame({"smiles": ["CCO", "CCN", "CCC"], "target": [1.2, -0.5, None]}).to_csv( + root / "train.csv", index=False + ) + pd.DataFrame({"smiles": ["CCCl", "CCBr"], "target": [0.1, 0.2]}).to_csv( + root / "test.csv", index=False + ) + + dataset = load_esol(root=root) + dataset.check() + + assert dataset.name == "esol" + assert len(dataset.train) == 2