Download src/remedi/splits.py from ChatterjeeLab/ReMEDi: direct link, hf CLI and curl.
- Browser
- Download file 2.03 kB
-
https://huggingface.co/ChatterjeeLab/ReMEDi/resolve/main/src/remedi/splits.py
- Command line
-
hf download hf://ChatterjeeLab/ReMEDi/src/remedi/splits.py
-
curl -L -o splits.py https://huggingface.co/ChatterjeeLab/ReMEDi/resolve/main/src/remedi/splits.py
2.03 kB
| """Molecule-disjoint or scaffold-disjoint splits before feature fitting.""" | |
| import numpy as np | |
| import pandas as pd | |
| from .chemistry import canonical_smiles, molecule_id | |
| def make_splits(frame, seed=0, fractions=(.6, .15, .1, .15), scaffold=False): | |
| if len(fractions) != 4 or not np.isclose(sum(fractions), 1) or min(fractions) <= 0: | |
| raise ValueError("Four positive split fractions must sum to one") | |
| frame = frame[["smiles"]].drop_duplicates().copy() | |
| frame["smiles"] = frame.smiles.map(canonical_smiles) | |
| frame = frame.drop_duplicates("smiles") | |
| frame["molecule_id"] = frame.smiles.map(molecule_id) | |
| frame["group"] = frame.molecule_id | |
| if scaffold: | |
| from rdkit import Chem | |
| from rdkit.Chem.Scaffolds import MurckoScaffold | |
| def key(s): | |
| scaffold = MurckoScaffold.MurckoScaffoldSmiles(mol=Chem.MolFromSmiles(s)) | |
| # All acyclic structures remain one group in this conservative implementation. | |
| return scaffold or "ACYCLIC" | |
| frame["group"] = frame.smiles.map(key) | |
| groups = frame.groupby("group").size() | |
| if len(groups) < 8: | |
| raise ValueError("At least eight independent molecule/scaffold groups are required") | |
| rng = np.random.default_rng(seed) | |
| order = groups.index.to_numpy()[rng.permutation(len(groups))] | |
| # Allocate whole groups. Each split receives at least one group. | |
| counts = np.maximum(1, np.floor(np.asarray(fractions) * len(order)).astype(int)) | |
| while counts.sum() < len(order): | |
| counts[np.argmax(np.asarray(fractions) * len(order) - counts)] += 1 | |
| while counts.sum() > len(order): | |
| eligible = np.where(counts > 1, counts - np.asarray(fractions)*len(order), -np.inf) | |
| counts[np.argmax(eligible)] -= 1 | |
| assignments = {} | |
| pos = 0 | |
| for name, n in zip(["train", "tune", "calibration", "test"], counts): | |
| assignments.update({g: name for g in order[pos:pos+n]}) | |
| pos += n | |
| frame["split"] = frame.group.map(assignments) | |
| return frame.reset_index(drop=True) | |