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 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
@@ -0,0 +1,8 @@
1
+ """Entry point for ``python -m chemsplit``."""
2
+
3
+ import sys
4
+
5
+ from chemsplit.cli import main
6
+
7
+ if __name__ == "__main__":
8
+ sys.exit(main())
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())