chemsplit 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- chemsplit/__init__.py +250 -0
- chemsplit/__main__.py +8 -0
- chemsplit/_devtools.py +420 -0
- chemsplit/_fp_similarity.py +154 -0
- chemsplit/_optimize.py +781 -0
- chemsplit/_pair_assign.py +143 -0
- chemsplit/_unionfind.py +94 -0
- chemsplit/audit.py +577 -0
- chemsplit/base.py +977 -0
- chemsplit/cli.py +135 -0
- chemsplit/clustering.py +288 -0
- chemsplit/datasets.py +656 -0
- chemsplit/determinism.py +160 -0
- chemsplit/exceptions.py +252 -0
- chemsplit/featurizers/__init__.py +101 -0
- chemsplit/featurizers/descriptors.py +78 -0
- chemsplit/featurizers/fingerprints.py +205 -0
- chemsplit/featurizers/precomputed.py +28 -0
- chemsplit/metrics.py +284 -0
- chemsplit/preprocess.py +418 -0
- chemsplit/registry.py +306 -0
- chemsplit/scaffolds.py +331 -0
- chemsplit/splitters/__init__.py +4 -0
- chemsplit/splitters/baseline.py +816 -0
- chemsplit/splitters/biomolecular.py +579 -0
- chemsplit/splitters/embedding.py +645 -0
- chemsplit/splitters/lineage.py +967 -0
- chemsplit/splitters/property_.py +983 -0
- chemsplit/splitters/protocol.py +772 -0
- chemsplit/splitters/scaffold.py +1220 -0
- chemsplit/splitters/similarity.py +1613 -0
- chemsplit/splitters/task.py +1692 -0
- chemsplit/types.py +41 -0
- chemsplit-0.1.0.dist-info/METADATA +186 -0
- chemsplit-0.1.0.dist-info/RECORD +39 -0
- chemsplit-0.1.0.dist-info/WHEEL +5 -0
- chemsplit-0.1.0.dist-info/entry_points.txt +2 -0
- chemsplit-0.1.0.dist-info/licenses/LICENSE +21 -0
- chemsplit-0.1.0.dist-info/top_level.txt +1 -0
chemsplit/__init__.py
ADDED
|
@@ -0,0 +1,250 @@
|
|
|
1
|
+
"""chemsplit: dataset-splitting strategies for cheminformatics machine learning.
|
|
2
|
+
"""
|
|
3
|
+
|
|
4
|
+
from __future__ import annotations
|
|
5
|
+
|
|
6
|
+
import importlib
|
|
7
|
+
from typing import Any
|
|
8
|
+
|
|
9
|
+
__version__ = "0.1.0"
|
|
10
|
+
|
|
11
|
+
__all__ = [
|
|
12
|
+
# -- 52 splitters across 9 families, in declared order: baseline (5), scaffold (6),
|
|
13
|
+
# similarity (10), embedding (3), property (5), lineage (4), task (8), biomolecular (5),
|
|
14
|
+
# protocol (6) --
|
|
15
|
+
"RandomSplitter",
|
|
16
|
+
"StratifiedRandomSplitter",
|
|
17
|
+
"KFoldSplitter",
|
|
18
|
+
"MonteCarloSplitter",
|
|
19
|
+
"PredefinedSplitter",
|
|
20
|
+
"MurckoScaffoldSplitter",
|
|
21
|
+
"GenericScaffoldSplitter",
|
|
22
|
+
"ScaffoldTreeSplitter",
|
|
23
|
+
"RingSystemSplitter",
|
|
24
|
+
"MatchedMolecularSeriesSplitter",
|
|
25
|
+
"ActivityCliffSplitter",
|
|
26
|
+
"SimilarityThresholdSplitter",
|
|
27
|
+
"ButinaSplitter",
|
|
28
|
+
"KMeansClusterSplitter",
|
|
29
|
+
"DensityClusterSplitter",
|
|
30
|
+
"SpectralSplitter",
|
|
31
|
+
"MaxMinSplitter",
|
|
32
|
+
"MaxDissimilaritySplitter",
|
|
33
|
+
"PerimeterSplitter",
|
|
34
|
+
"LeaveOneClusterOutSplitter",
|
|
35
|
+
"BalancedMultiTaskSplitter",
|
|
36
|
+
"UMAPClusterSplitter",
|
|
37
|
+
"ProjectionSplitter",
|
|
38
|
+
"LatentSpaceSplitter",
|
|
39
|
+
"PropertySplitter",
|
|
40
|
+
"LabelExtrapolationSplitter",
|
|
41
|
+
"StratifiedDistributionSplitter",
|
|
42
|
+
"MOODSplitter",
|
|
43
|
+
"AdversarialSplitter",
|
|
44
|
+
"TemporalSplitter",
|
|
45
|
+
"SIMPDSplitter",
|
|
46
|
+
"SourceSplitter",
|
|
47
|
+
"PartySplitter",
|
|
48
|
+
"HiSplitter",
|
|
49
|
+
"LoSplitter",
|
|
50
|
+
"ScaffoldHopSplitter",
|
|
51
|
+
"ColdDrugSplitter",
|
|
52
|
+
"ColdTargetSplitter",
|
|
53
|
+
"ColdPairSplitter",
|
|
54
|
+
"AVESplitter",
|
|
55
|
+
"DecoyBenchmarkSplitter",
|
|
56
|
+
"SequenceIdentitySplitter",
|
|
57
|
+
"ProteinFamilySplitter",
|
|
58
|
+
"BindingSiteSplitter",
|
|
59
|
+
"DepositionDateSplitter",
|
|
60
|
+
"ComplexJointSplitter",
|
|
61
|
+
"GroupKFoldSplitter",
|
|
62
|
+
"ThreeWaySplitter",
|
|
63
|
+
"RepeatedSplitter",
|
|
64
|
+
"NestedCVSplitter",
|
|
65
|
+
"ExternalHoldoutSplitter",
|
|
66
|
+
"ApplicabilityDomainSplitter",
|
|
67
|
+
# -- core types --
|
|
68
|
+
"BaseSplitter",
|
|
69
|
+
"GroupSplitter",
|
|
70
|
+
"SplitResult",
|
|
71
|
+
"Strictness",
|
|
72
|
+
# -- registry --
|
|
73
|
+
"get_splitter",
|
|
74
|
+
"list_splitters",
|
|
75
|
+
"SPLITTER_REGISTRY",
|
|
76
|
+
# -- audit --
|
|
77
|
+
"audit_split",
|
|
78
|
+
"LeakageReport",
|
|
79
|
+
"adversarial_validation",
|
|
80
|
+
"nn_similarity_profile",
|
|
81
|
+
"y_scramble_control",
|
|
82
|
+
# -- featurizers --
|
|
83
|
+
"get_featurizer",
|
|
84
|
+
"Featurizer",
|
|
85
|
+
# -- exceptions --
|
|
86
|
+
"ChemSplitError",
|
|
87
|
+
"ParameterError",
|
|
88
|
+
"ConfigurationError",
|
|
89
|
+
"UnknownSplitterError",
|
|
90
|
+
"UnknownFeaturizerError",
|
|
91
|
+
"UnknownMetricError",
|
|
92
|
+
"InputError",
|
|
93
|
+
"InputKindError",
|
|
94
|
+
"ColumnError",
|
|
95
|
+
"EmptyInputError",
|
|
96
|
+
"MoleculeParseError",
|
|
97
|
+
"DuplicateRecordError",
|
|
98
|
+
"LabelError",
|
|
99
|
+
"InfeasibleSplitError",
|
|
100
|
+
"DegenerateGroupingError",
|
|
101
|
+
"ConstraintUnsatisfiableError",
|
|
102
|
+
"EmptyPartitionError",
|
|
103
|
+
"ScalabilityError",
|
|
104
|
+
"MissingDependencyError",
|
|
105
|
+
"InvariantError",
|
|
106
|
+
# -- warnings --
|
|
107
|
+
"ChemSplitWarning",
|
|
108
|
+
"SizeToleranceWarning",
|
|
109
|
+
"DuplicateWarning",
|
|
110
|
+
"ParseWarning",
|
|
111
|
+
"StandardizationWarning",
|
|
112
|
+
"DeterminismWarning",
|
|
113
|
+
"DegenerateClusterWarning",
|
|
114
|
+
"SmallPartitionWarning",
|
|
115
|
+
"CircularityWarning",
|
|
116
|
+
"HomologyLeakWarning",
|
|
117
|
+
# -- version --
|
|
118
|
+
"__version__",
|
|
119
|
+
]
|
|
120
|
+
|
|
121
|
+
# name -> the submodule that actually defines it. Every entry not listed here (there are none
|
|
122
|
+
# left over once this dict is complete) would fall through to AttributeError in __getattr__.
|
|
123
|
+
_LAZY_SOURCE: dict[str, str] = {
|
|
124
|
+
# splitters/baseline.py
|
|
125
|
+
"RandomSplitter": "chemsplit.splitters.baseline",
|
|
126
|
+
"StratifiedRandomSplitter": "chemsplit.splitters.baseline",
|
|
127
|
+
"KFoldSplitter": "chemsplit.splitters.baseline",
|
|
128
|
+
"MonteCarloSplitter": "chemsplit.splitters.baseline",
|
|
129
|
+
"PredefinedSplitter": "chemsplit.splitters.baseline",
|
|
130
|
+
# splitters/scaffold.py
|
|
131
|
+
"MurckoScaffoldSplitter": "chemsplit.splitters.scaffold",
|
|
132
|
+
"GenericScaffoldSplitter": "chemsplit.splitters.scaffold",
|
|
133
|
+
"ScaffoldTreeSplitter": "chemsplit.splitters.scaffold",
|
|
134
|
+
"RingSystemSplitter": "chemsplit.splitters.scaffold",
|
|
135
|
+
"MatchedMolecularSeriesSplitter": "chemsplit.splitters.scaffold",
|
|
136
|
+
"ActivityCliffSplitter": "chemsplit.splitters.scaffold",
|
|
137
|
+
# splitters/similarity.py
|
|
138
|
+
"SimilarityThresholdSplitter": "chemsplit.splitters.similarity",
|
|
139
|
+
"ButinaSplitter": "chemsplit.splitters.similarity",
|
|
140
|
+
"KMeansClusterSplitter": "chemsplit.splitters.similarity",
|
|
141
|
+
"DensityClusterSplitter": "chemsplit.splitters.similarity",
|
|
142
|
+
"SpectralSplitter": "chemsplit.splitters.similarity",
|
|
143
|
+
"MaxMinSplitter": "chemsplit.splitters.similarity",
|
|
144
|
+
"MaxDissimilaritySplitter": "chemsplit.splitters.similarity",
|
|
145
|
+
"PerimeterSplitter": "chemsplit.splitters.similarity",
|
|
146
|
+
"LeaveOneClusterOutSplitter": "chemsplit.splitters.similarity",
|
|
147
|
+
"BalancedMultiTaskSplitter": "chemsplit.splitters.similarity",
|
|
148
|
+
# splitters/embedding.py
|
|
149
|
+
"UMAPClusterSplitter": "chemsplit.splitters.embedding",
|
|
150
|
+
"ProjectionSplitter": "chemsplit.splitters.embedding",
|
|
151
|
+
"LatentSpaceSplitter": "chemsplit.splitters.embedding",
|
|
152
|
+
# splitters/property_.py
|
|
153
|
+
"PropertySplitter": "chemsplit.splitters.property_",
|
|
154
|
+
"LabelExtrapolationSplitter": "chemsplit.splitters.property_",
|
|
155
|
+
"StratifiedDistributionSplitter": "chemsplit.splitters.property_",
|
|
156
|
+
"MOODSplitter": "chemsplit.splitters.property_",
|
|
157
|
+
"AdversarialSplitter": "chemsplit.splitters.property_",
|
|
158
|
+
# splitters/lineage.py
|
|
159
|
+
"TemporalSplitter": "chemsplit.splitters.lineage",
|
|
160
|
+
"SIMPDSplitter": "chemsplit.splitters.lineage",
|
|
161
|
+
"SourceSplitter": "chemsplit.splitters.lineage",
|
|
162
|
+
"PartySplitter": "chemsplit.splitters.lineage",
|
|
163
|
+
# splitters/task.py
|
|
164
|
+
"HiSplitter": "chemsplit.splitters.task",
|
|
165
|
+
"LoSplitter": "chemsplit.splitters.task",
|
|
166
|
+
"ScaffoldHopSplitter": "chemsplit.splitters.task",
|
|
167
|
+
"ColdDrugSplitter": "chemsplit.splitters.task",
|
|
168
|
+
"ColdTargetSplitter": "chemsplit.splitters.task",
|
|
169
|
+
"ColdPairSplitter": "chemsplit.splitters.task",
|
|
170
|
+
"AVESplitter": "chemsplit.splitters.task",
|
|
171
|
+
"DecoyBenchmarkSplitter": "chemsplit.splitters.task",
|
|
172
|
+
# splitters/biomolecular.py
|
|
173
|
+
"SequenceIdentitySplitter": "chemsplit.splitters.biomolecular",
|
|
174
|
+
"ProteinFamilySplitter": "chemsplit.splitters.biomolecular",
|
|
175
|
+
"BindingSiteSplitter": "chemsplit.splitters.biomolecular",
|
|
176
|
+
"DepositionDateSplitter": "chemsplit.splitters.biomolecular",
|
|
177
|
+
"ComplexJointSplitter": "chemsplit.splitters.biomolecular",
|
|
178
|
+
# splitters/protocol.py
|
|
179
|
+
"GroupKFoldSplitter": "chemsplit.splitters.protocol",
|
|
180
|
+
"ThreeWaySplitter": "chemsplit.splitters.protocol",
|
|
181
|
+
"RepeatedSplitter": "chemsplit.splitters.protocol",
|
|
182
|
+
"NestedCVSplitter": "chemsplit.splitters.protocol",
|
|
183
|
+
"ExternalHoldoutSplitter": "chemsplit.splitters.protocol",
|
|
184
|
+
"ApplicabilityDomainSplitter": "chemsplit.splitters.protocol",
|
|
185
|
+
# base.py
|
|
186
|
+
"BaseSplitter": "chemsplit.base",
|
|
187
|
+
"GroupSplitter": "chemsplit.base",
|
|
188
|
+
"SplitResult": "chemsplit.base",
|
|
189
|
+
"Strictness": "chemsplit.base",
|
|
190
|
+
# registry.py
|
|
191
|
+
"get_splitter": "chemsplit.registry",
|
|
192
|
+
"list_splitters": "chemsplit.registry",
|
|
193
|
+
"SPLITTER_REGISTRY": "chemsplit.registry",
|
|
194
|
+
# audit.py
|
|
195
|
+
"audit_split": "chemsplit.audit",
|
|
196
|
+
"LeakageReport": "chemsplit.audit",
|
|
197
|
+
"adversarial_validation": "chemsplit.audit",
|
|
198
|
+
"nn_similarity_profile": "chemsplit.audit",
|
|
199
|
+
"y_scramble_control": "chemsplit.audit",
|
|
200
|
+
# featurizers/__init__.py (cheap: no rdkit/sklearn at its own module scope)
|
|
201
|
+
"get_featurizer": "chemsplit.featurizers",
|
|
202
|
+
"Featurizer": "chemsplit.featurizers",
|
|
203
|
+
# exceptions.py (cheap: no rdkit/sklearn at its own module scope)
|
|
204
|
+
"ChemSplitError": "chemsplit.exceptions",
|
|
205
|
+
"ParameterError": "chemsplit.exceptions",
|
|
206
|
+
"ConfigurationError": "chemsplit.exceptions",
|
|
207
|
+
"UnknownSplitterError": "chemsplit.exceptions",
|
|
208
|
+
"UnknownFeaturizerError": "chemsplit.exceptions",
|
|
209
|
+
"UnknownMetricError": "chemsplit.exceptions",
|
|
210
|
+
"InputError": "chemsplit.exceptions",
|
|
211
|
+
"InputKindError": "chemsplit.exceptions",
|
|
212
|
+
"ColumnError": "chemsplit.exceptions",
|
|
213
|
+
"EmptyInputError": "chemsplit.exceptions",
|
|
214
|
+
"MoleculeParseError": "chemsplit.exceptions",
|
|
215
|
+
"DuplicateRecordError": "chemsplit.exceptions",
|
|
216
|
+
"LabelError": "chemsplit.exceptions",
|
|
217
|
+
"InfeasibleSplitError": "chemsplit.exceptions",
|
|
218
|
+
"DegenerateGroupingError": "chemsplit.exceptions",
|
|
219
|
+
"ConstraintUnsatisfiableError": "chemsplit.exceptions",
|
|
220
|
+
"EmptyPartitionError": "chemsplit.exceptions",
|
|
221
|
+
"ScalabilityError": "chemsplit.exceptions",
|
|
222
|
+
"MissingDependencyError": "chemsplit.exceptions",
|
|
223
|
+
"InvariantError": "chemsplit.exceptions",
|
|
224
|
+
"ChemSplitWarning": "chemsplit.exceptions",
|
|
225
|
+
"SizeToleranceWarning": "chemsplit.exceptions",
|
|
226
|
+
"DuplicateWarning": "chemsplit.exceptions",
|
|
227
|
+
"ParseWarning": "chemsplit.exceptions",
|
|
228
|
+
"StandardizationWarning": "chemsplit.exceptions",
|
|
229
|
+
"DeterminismWarning": "chemsplit.exceptions",
|
|
230
|
+
"DegenerateClusterWarning": "chemsplit.exceptions",
|
|
231
|
+
"SmallPartitionWarning": "chemsplit.exceptions",
|
|
232
|
+
"CircularityWarning": "chemsplit.exceptions",
|
|
233
|
+
"HomologyLeakWarning": "chemsplit.exceptions",
|
|
234
|
+
}
|
|
235
|
+
|
|
236
|
+
|
|
237
|
+
def __getattr__(name: str) -> Any:
|
|
238
|
+
if name == "__version__":
|
|
239
|
+
return __version__
|
|
240
|
+
source = _LAZY_SOURCE.get(name)
|
|
241
|
+
if source is None:
|
|
242
|
+
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
|
243
|
+
module = importlib.import_module(source)
|
|
244
|
+
value = getattr(module, name)
|
|
245
|
+
globals()[name] = value # cache on the package module so repeat access is free
|
|
246
|
+
return value
|
|
247
|
+
|
|
248
|
+
|
|
249
|
+
def __dir__() -> list[str]:
|
|
250
|
+
return sorted(__all__)
|
chemsplit/__main__.py
ADDED
chemsplit/_devtools.py
ADDED
|
@@ -0,0 +1,420 @@
|
|
|
1
|
+
"""Developer tooling: golden-file regeneration."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import argparse
|
|
6
|
+
import json
|
|
7
|
+
import os
|
|
8
|
+
import sys
|
|
9
|
+
from collections.abc import Callable
|
|
10
|
+
from pathlib import Path
|
|
11
|
+
from typing import Any
|
|
12
|
+
|
|
13
|
+
import numpy as np
|
|
14
|
+
|
|
15
|
+
_GOLDEN_DIR = Path(__file__).resolve().parent.parent.parent / "tests" / "golden"
|
|
16
|
+
|
|
17
|
+
#: A zero-arg callable returning (X, y, split_kwargs, ctor_kwargs) for one splitter's golden case.
|
|
18
|
+
PlanBuilder = Callable[[], tuple[Any, Any, dict[str, Any], dict[str, Any]]]
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def _fixture_cache() -> dict[str, Any]:
|
|
22
|
+
from chemsplit import datasets as ds
|
|
23
|
+
|
|
24
|
+
return {
|
|
25
|
+
"linear_series": ds.make_linear_series(n=80, seed=0),
|
|
26
|
+
"scaffold_families": ds.make_scaffold_families(n_scaffolds=8, per_scaffold=10, seed=0),
|
|
27
|
+
"two_clusters": ds.make_two_clusters(n=80, seed=0),
|
|
28
|
+
"activity_cliffs": ds.make_activity_cliffs(n_pairs=20, seed=0),
|
|
29
|
+
"dated_series": ds.make_dated_series(n=100, seed=0),
|
|
30
|
+
"multitask_sparse": ds.make_multitask_sparse(n=100, n_tasks=4, seed=0),
|
|
31
|
+
"interactions": ds.make_interactions(n_compounds=40, n_targets=10, seed=0),
|
|
32
|
+
"sequences": ds.make_sequences(n=20, families=4, seed=0),
|
|
33
|
+
"pathological": ds.make_pathological(),
|
|
34
|
+
"all_identical": ds.make_all_identical(n=30),
|
|
35
|
+
"singletons": ds.make_singletons(n=20, seed=0),
|
|
36
|
+
"label_extremes": ds.make_label_extremes(n=100, seed=0),
|
|
37
|
+
# SIMPDSplitter (simpd) needs n >= 200 to meaningfully run its GA.
|
|
38
|
+
"simpd_series": ds.make_scaffold_families(n_scaffolds=10, per_scaffold=20, seed=1),
|
|
39
|
+
# A parallel sequences array for ComplexJointSplitter (complex_joint), which wants
|
|
40
|
+
# ctx.sequences to derive its default sequence grouper even though its own `accepts`
|
|
41
|
+
# doesn't include "sequences" (it's a second axis alongside the ligand SMILES, not the
|
|
42
|
+
# primary X).
|
|
43
|
+
"aux_sequences_80": ds.make_sequences(n=80, families=4, seed=1),
|
|
44
|
+
}
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def _interactions_X_y(fixture: Any) -> tuple[list[tuple[str, str]], list[float]]:
|
|
48
|
+
X = [(fixture.smiles[ci], fixture.targets[ti]) for ci, ti, _y in fixture.interactions]
|
|
49
|
+
y = [float(_y) for _ci, _ti, _y in fixture.interactions]
|
|
50
|
+
return X, y
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
def _binary_y(y: np.ndarray) -> np.ndarray:
|
|
54
|
+
med = float(np.median(y))
|
|
55
|
+
return (np.asarray(y) > med).astype(np.int64)
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
# Each entry: splitter_id -> a zero-arg callable returning (X, y, split_kwargs, ctor_kwargs).
|
|
59
|
+
# split_kwargs are passed to split_result(X, y, **split_kwargs) (e.g. dates=, X_kind=,
|
|
60
|
+
# sequences=). ctor_kwargs are passed to the splitter's constructor alongside random_state=0.
|
|
61
|
+
def _build_plan(fixtures: dict[str, Any]) -> tuple[dict[str, Any], dict[str, str]]:
|
|
62
|
+
F = fixtures
|
|
63
|
+
plan: dict[str, Any] = {}
|
|
64
|
+
fixture_name_of: dict[str, str] = {}
|
|
65
|
+
|
|
66
|
+
def smiles_case(fid: str, *, y: np.ndarray | None = None, **ctor: Any) -> PlanBuilder:
|
|
67
|
+
return lambda: (F[fid].smiles, y, {}, ctor)
|
|
68
|
+
|
|
69
|
+
# -- baseline --
|
|
70
|
+
plan["random"] = smiles_case("linear_series")
|
|
71
|
+
plan["stratified_random"] = lambda: (F["linear_series"].smiles, F["linear_series"].y, {}, {})
|
|
72
|
+
plan["k_fold"] = lambda: (F["linear_series"].smiles, None, {}, {"n_splits": 3})
|
|
73
|
+
plan["monte_carlo"] = smiles_case("linear_series")
|
|
74
|
+
plan["predefined"] = lambda: (
|
|
75
|
+
F["linear_series"].smiles,
|
|
76
|
+
None,
|
|
77
|
+
{},
|
|
78
|
+
{
|
|
79
|
+
"assignment": {
|
|
80
|
+
"train": list(range(0, 60)),
|
|
81
|
+
"test": list(range(60, 80)),
|
|
82
|
+
}
|
|
83
|
+
},
|
|
84
|
+
)
|
|
85
|
+
|
|
86
|
+
# -- scaffold --
|
|
87
|
+
plan["murcko_scaffold"] = smiles_case("scaffold_families")
|
|
88
|
+
plan["generic_scaffold"] = smiles_case("scaffold_families")
|
|
89
|
+
plan["scaffold_tree"] = smiles_case("scaffold_families")
|
|
90
|
+
plan["ring_system"] = smiles_case("scaffold_families")
|
|
91
|
+
plan["matched_molecular_series"] = smiles_case("scaffold_families")
|
|
92
|
+
plan["activity_cliff"] = lambda: (
|
|
93
|
+
F["activity_cliffs"].smiles,
|
|
94
|
+
F["activity_cliffs"].y,
|
|
95
|
+
{},
|
|
96
|
+
{},
|
|
97
|
+
)
|
|
98
|
+
|
|
99
|
+
# -- similarity --
|
|
100
|
+
for sid in ["similarity_threshold", "butina", "k_means_cluster", "density_cluster", "spectral",
|
|
101
|
+
"max_min", "max_dissimilarity", "perimeter", "leave_one_cluster_out"]:
|
|
102
|
+
plan[sid] = smiles_case("two_clusters")
|
|
103
|
+
plan["balanced_multi_task"] = lambda: (
|
|
104
|
+
F["multitask_sparse"].smiles,
|
|
105
|
+
F["multitask_sparse"].y,
|
|
106
|
+
{},
|
|
107
|
+
{},
|
|
108
|
+
)
|
|
109
|
+
|
|
110
|
+
# -- embedding --
|
|
111
|
+
plan["umap_cluster"] = smiles_case("two_clusters")
|
|
112
|
+
plan["projection"] = smiles_case("two_clusters")
|
|
113
|
+
plan["latent_space"] = lambda: (
|
|
114
|
+
np.random.default_rng(0).standard_normal((60, 12)),
|
|
115
|
+
None,
|
|
116
|
+
{},
|
|
117
|
+
{},
|
|
118
|
+
)
|
|
119
|
+
|
|
120
|
+
# -- property --
|
|
121
|
+
plan["property"] = smiles_case("linear_series")
|
|
122
|
+
plan["label_extrapolation"] = lambda: (F["linear_series"].smiles, F["linear_series"].y, {}, {})
|
|
123
|
+
plan["stratified_distribution"] = lambda: (F["linear_series"].smiles, F["linear_series"].y, {}, {})
|
|
124
|
+
def _moodsplitter_case() -> tuple[list[str], None, dict[str, Any], dict[str, Any]]:
|
|
125
|
+
from chemsplit.registry import get_splitter
|
|
126
|
+
|
|
127
|
+
candidates = [
|
|
128
|
+
get_splitter("random", random_state=0),
|
|
129
|
+
get_splitter("butina", random_state=0),
|
|
130
|
+
]
|
|
131
|
+
return (
|
|
132
|
+
F["two_clusters"].smiles,
|
|
133
|
+
None,
|
|
134
|
+
{},
|
|
135
|
+
{"candidates": candidates, "deployment_set": F["singletons"].smiles},
|
|
136
|
+
)
|
|
137
|
+
|
|
138
|
+
plan["mood"] = _moodsplitter_case
|
|
139
|
+
plan["adversarial"] = smiles_case("two_clusters")
|
|
140
|
+
|
|
141
|
+
# -- lineage --
|
|
142
|
+
plan["temporal"] = lambda: (
|
|
143
|
+
F["dated_series"].smiles,
|
|
144
|
+
None,
|
|
145
|
+
{"dates": F["dated_series"].dates},
|
|
146
|
+
{},
|
|
147
|
+
)
|
|
148
|
+
plan["simpd"] = lambda: (
|
|
149
|
+
F["simpd_series"].smiles,
|
|
150
|
+
_binary_y(np.arange(len(F["simpd_series"].smiles)) % 2),
|
|
151
|
+
{},
|
|
152
|
+
{},
|
|
153
|
+
)
|
|
154
|
+
plan["source"] = lambda: (
|
|
155
|
+
F["scaffold_families"].smiles,
|
|
156
|
+
None,
|
|
157
|
+
{},
|
|
158
|
+
{"source": F["scaffold_families"].groups_true.tolist()},
|
|
159
|
+
)
|
|
160
|
+
plan["party"] = lambda: (
|
|
161
|
+
F["scaffold_families"].smiles,
|
|
162
|
+
None,
|
|
163
|
+
{},
|
|
164
|
+
{"party": F["scaffold_families"].groups_true.tolist(), "synthesis": "given"},
|
|
165
|
+
)
|
|
166
|
+
|
|
167
|
+
# -- task --
|
|
168
|
+
plan["hi"] = smiles_case("two_clusters")
|
|
169
|
+
plan["lo"] = lambda: (F["linear_series"].smiles, F["linear_series"].y, {}, {})
|
|
170
|
+
plan["scaffold_hop"] = lambda: (
|
|
171
|
+
F["scaffold_families"].smiles,
|
|
172
|
+
_binary_y(np.arange(len(F["scaffold_families"].smiles)) % 3 == 0),
|
|
173
|
+
{},
|
|
174
|
+
{"pharmacophore_similarity": "none"},
|
|
175
|
+
)
|
|
176
|
+
for sid in ["cold_drug", "cold_target", "cold_pair"]:
|
|
177
|
+
def _make(sid: str = sid) -> tuple[list[tuple[str, str]], list[float], dict[str, Any], dict[str, Any]]:
|
|
178
|
+
X, y = _interactions_X_y(F["interactions"])
|
|
179
|
+
return (X, y, {}, {})
|
|
180
|
+
plan[sid] = _make
|
|
181
|
+
plan["ave"] = lambda: (
|
|
182
|
+
F["two_clusters"].smiles,
|
|
183
|
+
_binary_y(np.arange(len(F["two_clusters"].smiles)) % 2),
|
|
184
|
+
{},
|
|
185
|
+
{},
|
|
186
|
+
)
|
|
187
|
+
plan["decoy_benchmark"] = lambda: (
|
|
188
|
+
F["scaffold_families"].smiles,
|
|
189
|
+
_binary_y(np.arange(len(F["scaffold_families"].smiles)) % 5 == 0),
|
|
190
|
+
{},
|
|
191
|
+
{"scheme": "spatial_random"},
|
|
192
|
+
)
|
|
193
|
+
|
|
194
|
+
# -- biomolecular --
|
|
195
|
+
plan["sequence_identity"] = lambda: (
|
|
196
|
+
F["sequences"].sequences,
|
|
197
|
+
None,
|
|
198
|
+
{"X_kind": "sequences", "sequences": F["sequences"].sequences},
|
|
199
|
+
{},
|
|
200
|
+
)
|
|
201
|
+
plan["protein_family"] = lambda: (
|
|
202
|
+
F["sequences"].sequences,
|
|
203
|
+
None,
|
|
204
|
+
{"X_kind": "sequences", "sequences": F["sequences"].sequences},
|
|
205
|
+
{"family_labels": F["sequences"].groups_true.tolist()},
|
|
206
|
+
)
|
|
207
|
+
plan["binding_site"] = lambda: (
|
|
208
|
+
F["sequences"].sequences,
|
|
209
|
+
None,
|
|
210
|
+
{"X_kind": "sequences", "sequences": F["sequences"].sequences},
|
|
211
|
+
{"representation": "pocket_sequence"},
|
|
212
|
+
)
|
|
213
|
+
plan["deposition_date"] = lambda: (
|
|
214
|
+
F["dated_series"].smiles,
|
|
215
|
+
None,
|
|
216
|
+
{"dates": F["dated_series"].dates},
|
|
217
|
+
{"cut_date": str(np.median(F["dated_series"].dates.astype("datetime64[D]").astype("int64")).astype("datetime64[D]"))},
|
|
218
|
+
)
|
|
219
|
+
plan["complex_joint"] = lambda: (
|
|
220
|
+
F["two_clusters"].smiles,
|
|
221
|
+
None,
|
|
222
|
+
{"sequences": F["aux_sequences_80"].sequences},
|
|
223
|
+
{},
|
|
224
|
+
)
|
|
225
|
+
|
|
226
|
+
# -- protocol --
|
|
227
|
+
plan["group_k_fold"] = smiles_case("scaffold_families", **{"grouper": "murcko_scaffold"})
|
|
228
|
+
plan["three_way"] = lambda: (
|
|
229
|
+
F["two_clusters"].smiles,
|
|
230
|
+
None,
|
|
231
|
+
{},
|
|
232
|
+
{"base_splitter": "random", "train_size": 0.6, "valid_size": 0.2, "test_size": 0.2},
|
|
233
|
+
)
|
|
234
|
+
plan["repeated"] = lambda: (
|
|
235
|
+
F["linear_series"].smiles,
|
|
236
|
+
None,
|
|
237
|
+
{},
|
|
238
|
+
{"base_splitter": "random", "n_repeats": 3},
|
|
239
|
+
)
|
|
240
|
+
plan["nested_cv"] = lambda: (
|
|
241
|
+
F["linear_series"].smiles,
|
|
242
|
+
None,
|
|
243
|
+
{},
|
|
244
|
+
{"outer_splitter": "random", "inner_splitter": "random"},
|
|
245
|
+
)
|
|
246
|
+
plan["external_holdout"] = lambda: (
|
|
247
|
+
F["two_clusters"].smiles[:60],
|
|
248
|
+
None,
|
|
249
|
+
{},
|
|
250
|
+
{"X_external": F["two_clusters"].smiles[60:70]},
|
|
251
|
+
)
|
|
252
|
+
plan["applicability_domain"] = lambda: (F["two_clusters"].smiles, None, {}, {"base_splitter": "random"})
|
|
253
|
+
|
|
254
|
+
fixture_name_of.update(
|
|
255
|
+
{
|
|
256
|
+
"random": "linear_series",
|
|
257
|
+
"stratified_random": "linear_series",
|
|
258
|
+
"k_fold": "linear_series",
|
|
259
|
+
"monte_carlo": "linear_series",
|
|
260
|
+
"predefined": "linear_series",
|
|
261
|
+
"murcko_scaffold": "scaffold_families",
|
|
262
|
+
"generic_scaffold": "scaffold_families",
|
|
263
|
+
"scaffold_tree": "scaffold_families",
|
|
264
|
+
"ring_system": "scaffold_families",
|
|
265
|
+
"matched_molecular_series": "scaffold_families",
|
|
266
|
+
"activity_cliff": "activity_cliffs",
|
|
267
|
+
"similarity_threshold": "two_clusters",
|
|
268
|
+
"butina": "two_clusters",
|
|
269
|
+
"k_means_cluster": "two_clusters",
|
|
270
|
+
"density_cluster": "two_clusters",
|
|
271
|
+
"spectral": "two_clusters",
|
|
272
|
+
"max_min": "two_clusters",
|
|
273
|
+
"max_dissimilarity": "two_clusters",
|
|
274
|
+
"perimeter": "two_clusters",
|
|
275
|
+
"leave_one_cluster_out": "two_clusters",
|
|
276
|
+
"balanced_multi_task": "multitask_sparse",
|
|
277
|
+
"umap_cluster": "two_clusters",
|
|
278
|
+
"projection": "two_clusters",
|
|
279
|
+
"latent_space": "synthetic_matrix",
|
|
280
|
+
"property": "linear_series",
|
|
281
|
+
"label_extrapolation": "linear_series",
|
|
282
|
+
"stratified_distribution": "linear_series",
|
|
283
|
+
"mood": "two_clusters",
|
|
284
|
+
"adversarial": "two_clusters",
|
|
285
|
+
"temporal": "dated_series",
|
|
286
|
+
"simpd": "scaffold_families",
|
|
287
|
+
"source": "scaffold_families",
|
|
288
|
+
"party": "scaffold_families",
|
|
289
|
+
"hi": "two_clusters",
|
|
290
|
+
"lo": "linear_series",
|
|
291
|
+
"scaffold_hop": "scaffold_families",
|
|
292
|
+
"cold_drug": "interactions",
|
|
293
|
+
"cold_target": "interactions",
|
|
294
|
+
"cold_pair": "interactions",
|
|
295
|
+
"ave": "two_clusters",
|
|
296
|
+
"decoy_benchmark": "scaffold_families",
|
|
297
|
+
"sequence_identity": "sequences",
|
|
298
|
+
"protein_family": "sequences",
|
|
299
|
+
"binding_site": "sequences",
|
|
300
|
+
"deposition_date": "dated_series",
|
|
301
|
+
"complex_joint": "two_clusters",
|
|
302
|
+
"group_k_fold": "scaffold_families",
|
|
303
|
+
"three_way": "two_clusters",
|
|
304
|
+
"repeated": "linear_series",
|
|
305
|
+
"nested_cv": "linear_series",
|
|
306
|
+
"external_holdout": "two_clusters",
|
|
307
|
+
"applicability_domain": "two_clusters",
|
|
308
|
+
}
|
|
309
|
+
)
|
|
310
|
+
|
|
311
|
+
return plan, fixture_name_of
|
|
312
|
+
|
|
313
|
+
|
|
314
|
+
def _stabilize_for_golden(obj: Any, *, key: str | None = None) -> Any:
|
|
315
|
+
"""Round floats (BLAS/build ULP noise) and mask ``umap_versions`` before golden comparison."""
|
|
316
|
+
if isinstance(obj, float):
|
|
317
|
+
return float(f"{obj:.9g}")
|
|
318
|
+
if isinstance(obj, dict):
|
|
319
|
+
if key == "umap_versions":
|
|
320
|
+
return dict.fromkeys(obj, "<version>")
|
|
321
|
+
return {k: _stabilize_for_golden(v, key=k) for k, v in obj.items()}
|
|
322
|
+
if isinstance(obj, list):
|
|
323
|
+
return [_stabilize_for_golden(v, key=key) for v in obj]
|
|
324
|
+
return obj
|
|
325
|
+
|
|
326
|
+
|
|
327
|
+
def _to_golden_payload(result: Any, ctx_extra: dict[str, Any] | None = None) -> dict[str, Any]:
|
|
328
|
+
"""``result`` is a SplitResult. Chooses the exact or tolerance tier based on
|
|
329
|
+
``metadata.get("nondeterministic_method")`` (per-instance; a couple of splitters, e.g. the
|
|
330
|
+
embedding family's t-SNE/MDS projection modes, can only know this after actually running)."""
|
|
331
|
+
nondeterministic = bool(result.metadata.get("nondeterministic_method", False))
|
|
332
|
+
if not nondeterministic:
|
|
333
|
+
stabilized = _stabilize_for_golden(json.loads(result.to_json()))
|
|
334
|
+
split_result_json = json.dumps(
|
|
335
|
+
stabilized, sort_keys=True, separators=(",", ":"), ensure_ascii=True
|
|
336
|
+
)
|
|
337
|
+
return {"tier": "exact", "split_result_json": split_result_json}
|
|
338
|
+
|
|
339
|
+
# Tolerance tier: sizes exact, group-size histogram as a multiset.
|
|
340
|
+
if result.groups is not None:
|
|
341
|
+
group_sizes = np.bincount(result.groups).tolist()
|
|
342
|
+
else:
|
|
343
|
+
group_sizes = None
|
|
344
|
+
return {
|
|
345
|
+
"tier": "tolerance",
|
|
346
|
+
"n_train": int(len(result.train)),
|
|
347
|
+
"n_valid": int(len(result.valid)),
|
|
348
|
+
"n_test": int(len(result.test)),
|
|
349
|
+
"n_discard": int(len(result.discard)),
|
|
350
|
+
"group_size_histogram": group_sizes,
|
|
351
|
+
}
|
|
352
|
+
|
|
353
|
+
|
|
354
|
+
def regenerate_goldens(
|
|
355
|
+
*, confirm: bool = False, splitter_ids: list[str] | None = None
|
|
356
|
+
) -> dict[str, str]:
|
|
357
|
+
if not confirm or os.environ.get("CHEMSPLIT_ALLOW_GOLDEN_REGEN") != "1":
|
|
358
|
+
raise RuntimeError(
|
|
359
|
+
"regenerate_goldens() refuses to run: pass confirm=True AND set "
|
|
360
|
+
"CHEMSPLIT_ALLOW_GOLDEN_REGEN=1 in the environment, so goldens are "
|
|
361
|
+
"never silently rewritten by an unrelated test run."
|
|
362
|
+
)
|
|
363
|
+
|
|
364
|
+
from chemsplit.registry import SPLITTER_REGISTRY, _ensure_built, get_splitter
|
|
365
|
+
|
|
366
|
+
_ensure_built()
|
|
367
|
+
_GOLDEN_DIR.mkdir(parents=True, exist_ok=True)
|
|
368
|
+
fixtures = _fixture_cache()
|
|
369
|
+
plan, fixture_name_of = _build_plan(fixtures)
|
|
370
|
+
|
|
371
|
+
written: dict[str, str] = {}
|
|
372
|
+
failures: dict[str, str] = {}
|
|
373
|
+
ids = splitter_ids if splitter_ids is not None else [
|
|
374
|
+
sid for sid in SPLITTER_REGISTRY if sid in plan
|
|
375
|
+
]
|
|
376
|
+
for sid in ids:
|
|
377
|
+
if sid not in plan:
|
|
378
|
+
failures[sid] = "no plan entry"
|
|
379
|
+
continue
|
|
380
|
+
builder = plan[sid]
|
|
381
|
+
try:
|
|
382
|
+
X, y, split_kwargs, ctor_kwargs = builder()
|
|
383
|
+
splitter = get_splitter(sid, random_state=0, **ctor_kwargs)
|
|
384
|
+
results = splitter.split_result(X, y, **split_kwargs)
|
|
385
|
+
result = results[0]
|
|
386
|
+
payload = _to_golden_payload(result)
|
|
387
|
+
fixture_name = fixture_name_of.get(sid, "custom")
|
|
388
|
+
out_path = _GOLDEN_DIR / f"{sid}__{fixture_name}__seed0.json"
|
|
389
|
+
with open(out_path, "w", encoding="utf-8") as fh:
|
|
390
|
+
json.dump(payload, fh, indent=2, sort_keys=True)
|
|
391
|
+
written[sid] = str(out_path)
|
|
392
|
+
print(f"OK {sid:16s} -> {out_path.name}")
|
|
393
|
+
except Exception as exc: # noqa: BLE001 - devtool, report and continue
|
|
394
|
+
failures[sid] = f"{type(exc).__name__}: {exc}"
|
|
395
|
+
print(f"FAIL {sid:16s} -> {type(exc).__name__}: {exc}", file=sys.stderr)
|
|
396
|
+
|
|
397
|
+
print(f"\n{len(written)}/{len(ids)} golden files written; {len(failures)} failures.")
|
|
398
|
+
if failures:
|
|
399
|
+
print("Failures:", file=sys.stderr)
|
|
400
|
+
for sid, msg in failures.items():
|
|
401
|
+
print(f" {sid}: {msg}", file=sys.stderr)
|
|
402
|
+
return written
|
|
403
|
+
|
|
404
|
+
|
|
405
|
+
def main(argv: list[str] | None = None) -> int:
|
|
406
|
+
parser = argparse.ArgumentParser(prog="python -m chemsplit._devtools")
|
|
407
|
+
sub = parser.add_subparsers(dest="command", required=True)
|
|
408
|
+
p = sub.add_parser("regenerate_goldens")
|
|
409
|
+
p.add_argument("--confirm", action="store_true")
|
|
410
|
+
p.add_argument("--splitter", action="append", dest="splitter_ids", default=None)
|
|
411
|
+
args = parser.parse_args(argv)
|
|
412
|
+
|
|
413
|
+
if args.command == "regenerate_goldens":
|
|
414
|
+
written = regenerate_goldens(confirm=args.confirm, splitter_ids=args.splitter_ids)
|
|
415
|
+
return 0 if written else 1
|
|
416
|
+
return 1
|
|
417
|
+
|
|
418
|
+
|
|
419
|
+
if __name__ == "__main__":
|
|
420
|
+
raise SystemExit(main())
|