Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
83 changes: 83 additions & 0 deletions src/modernmolbert/eval/contributed_datasets.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand All @@ -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
21 changes: 21 additions & 0 deletions tests/test_eval_clintox.py
Original file line number Diff line number Diff line change
@@ -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
21 changes: 21 additions & 0 deletions tests/test_eval_esol.py
Original file line number Diff line number Diff line change
@@ -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
Loading