From e8d24b4a43ebc0b5f017ae82e82a4002c4bd8b53 Mon Sep 17 00:00:00 2001 From: Stephan Breimann Date: Sat, 19 Sep 2026 00:23:13 +0200 Subject: [PATCH] fix(data): seed load_dataset's random sampling with random_state (#582) `aa.load_dataset(random=True)` had no seed, so the first line of most workflows drew a different sample on every call and `options['random_state']` never reached it: a benchmark, tutorial or bug report written with `random=True` could not be reproduced by its own author. `load_dataset` now takes `random_state` (appended after `verbose`, so every positional call keeps working), resolved through `ut.check_random_state` like AAclust, CPP and dPULearn, which also gives it the documented `options['random_state']` override. The per-class draws share one `numpy.random.RandomState` instance, so the classes continue one stream instead of repeating it, and the legacy generator's permanent stream guarantee makes a seed reproduce the same frame in another process or environment. Nothing else changes. `random=False` keeps the deterministic head-of-class selection, and an unseeded `random=True` keeps drawing from numpy's global state: over all 14 bundled datasets (83 loading variants x 5 `PYTHONHASHSEED` values, 415 comparisons) the sha256 of the returned frame is identical to a pristine export of the parent commit, and with the global numpy state pinned the unseeded random path matches byte for byte too. An unseeded `random=True` deliberately stays silent: no other stochastic entry point in the package warns on `random_state=None`, `random=True` is an explicit opt-in to randomness, and a warning here would fire in the example notebook and every docs build for a documented default. Closes #582 Co-Authored-By: Claude Fable 5.1 --- CHANGELOG.md | 9 + aaanalysis/data_handling/_load_dataset.py | 26 +- examples/data_handling/load_dataset.ipynb | 832 +++++++++--------- .../data_handling_tests/test_load_dataset.py | 160 ++++ 4 files changed, 610 insertions(+), 417 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index f4dcca68c..8091d7090 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -129,6 +129,15 @@ notes — with cross-references and examples — live in axes' tick labels and raised `ValueError: not enough values to unpack (expected 2, got 0)`. Rendering is unchanged: the three `pytest-mpl` baselines are byte-identical before and after, and a caller-supplied `ax` is still drawn into, so multi-panel composition is unaffected. +- `load_dataset(random=True)` is reproducible: the new `random_state` parameter seeds the + balanced per-class sampling, and `options['random_state']` overrides it like everywhere else + in the package. The function had no seed at all, so the first line of most workflows drew a + different sample on every call and no benchmark, tutorial or bug report using `random=True` + could be reproduced. The per-class draws share one `numpy.random.RandomState`, whose stream is + stable across numpy versions, so a given seed reproduces the same frame in another process. + Everything else is untouched: `random=False` (the deterministic head-of-class selection) and + an unseeded `random=True` are byte-identical to before, verified over all 14 bundled datasets + (#582). ### Changed - `CPP.run(n_batches=...)` (scale-axis batching) now returns output identical to the diff --git a/aaanalysis/data_handling/_load_dataset.py b/aaanalysis/data_handling/_load_dataset.py index 2db2895ac..965fddebc 100644 --- a/aaanalysis/data_handling/_load_dataset.py +++ b/aaanalysis/data_handling/_load_dataset.py @@ -66,6 +66,20 @@ def post_check_df_seq(df_seq: pd.DataFrame, n: Optional[int] = None, name: Optio # Helper functions +def _get_sampling_rng(random_state: Optional[int] = None) -> Optional[np.random.RandomState]: + """Get random number generator for the balanced class sampling + + A single generator is shared by the per-class draws, so each class continues the same + stream instead of repeating it. ``np.random.RandomState`` is used (not ``default_rng``) + because its stream is guaranteed stable across numpy versions, which is what makes a + given seed reproduce the same sample in another process or environment. ``None`` keeps + the truly random default (pandas then draws from the global numpy random state). + """ + if random_state is None: + return None + return np.random.RandomState(random_state) + + def _is_aa_level(name: str) -> bool: return name.split("_")[0] == "AA" @@ -147,6 +161,7 @@ def load_dataset(name: str = "Overview", max_len: Optional[int] = None, aa_window_size: Union[int, None] = 9, verbose: bool = False, + random_state: Optional[int] = None, ) -> DataFrame: """ Load protein benchmarking datasets. @@ -178,6 +193,7 @@ def load_dataset(name: str = "Overview", Number of proteins per class, selected by index. If ``None``, the whole dataset will be returned. random : bool, default=False If ``True``, ``n`` randomly selected proteins per class will be chosen. + Set ``random_state`` to make that selection reproducible. non_canonical_aa : {'remove', 'keep', 'gap'}, default='remove' Options for handling non-canonical amino acids: @@ -200,6 +216,12 @@ def load_dataset(name: str = "Overview", If ``True``, report how many entries each removal step (``min_len``, ``max_len``, and ``non_canonical_aa='remove'``) dropped. Does not change the returned data. + random_state : int, optional + The seed used by the random number generator. If a positive integer, results of stochastic processes are + consistent, enabling reproducibility. If ``None``, stochastic processes will be truly random. Only used + together with ``random=True``; overridden by ``options['random_state']`` when set. + + .. versionadded:: 1.2.0 Returns ------- @@ -266,6 +288,7 @@ def load_dataset(name: str = "Overview", is_cs_dataset = _is_cleavage_site_dataset(name=name) check_aa_window_size(aa_window_size=aa_window_size, is_cs_dataset=is_cs_dataset) verbose = ut.check_verbose(verbose) + random_state = ut.check_random_state(random_state=random_state) # Load overview table if name == "Overview": @@ -311,7 +334,8 @@ def load_dataset(name: str = "Overview", if n is not None: labels = set(df_seq[ut.COL_LABEL]) if random: - df_seq = pd.concat([df_seq[df_seq[ut.COL_LABEL] == l].sample(n) for l in labels]) + rng = _get_sampling_rng(random_state=random_state) + df_seq = pd.concat([df_seq[df_seq[ut.COL_LABEL] == l].sample(n=n, random_state=rng) for l in labels]) else: df_seq = pd.concat([df_seq[df_seq[ut.COL_LABEL] == l].head(n) for l in labels]) post_check_df_seq(df_seq=df_seq, n=n, name=name) diff --git a/examples/data_handling/load_dataset.ipynb b/examples/data_handling/load_dataset.ipynb index d564faa01..f12ba5500 100644 --- a/examples/data_handling/load_dataset.ipynb +++ b/examples/data_handling/load_dataset.ipynb @@ -21,10 +21,10 @@ }, "collapsed": false, "execution": { - "iopub.execute_input": "2026-09-14T21:28:36.989110Z", - "iopub.status.busy": "2026-09-14T21:28:36.988989Z", - "iopub.status.idle": "2026-09-14T21:28:38.529658Z", - "shell.execute_reply": "2026-09-14T21:28:38.529357Z" + "iopub.execute_input": "2026-09-18T22:09:31.956656Z", + "iopub.status.busy": "2026-09-18T22:09:31.956267Z", + "iopub.status.idle": "2026-09-18T22:09:33.365636Z", + "shell.execute_reply": "2026-09-18T22:09:33.365453Z" } }, "outputs": [ @@ -39,238 +39,238 @@ "data": { "text/html": [ "\n", - "\n", + "
\n", " \n", " \n", " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", " \n", " \n", " \n", " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", " \n", " \n", "
 LevelDataset# SequencesAvg length# Amino acids# Positives# NegativesPredictorDescriptionReferenceLabelLevelDataset# SequencesAvg length# Amino acids# Positives# NegativesPredictorDescriptionReferenceLabel
1Amino acidAA_CASPASE3233796.587983185605705184900PROSPERousPrediction of c...3 cleavage siteSong et al., 20181 (adjacent to ... cleavage site)
2Amino acidAA_FURIN71831.0281695900316358840PROSPERousPrediction of f...n cleavage siteSong et al., 20181 (adjacent to ... cleavage site)
3Amino acidAA_LDR342345.7543861182483546982779IDP-Seq2SeqPrediction of l...d regions (LDR)Tang et al., 20201 (disordered), 0 (ordered)
4Amino acidAA_MMP2573546.2059343129762416310560PROSPERousPrediction of M...) cleavage siteSong et al., 20181 (adjacent to ... cleavage site)
5Amino acidAA_RNABIND221248.87330355001649248509GMKSVM-RUPrediction of R...(RBP60 dataset)Yang et al., 20211 (binding), 0 (non-binding)
6Amino acidAA_SA233796.58798318560510108284523PROSPERousPrediction of s...PASE3 data set)Song et al., 20181 (exposed/acce...non-accessible)
7SequenceSEQ_AMYLO14146.0000008484511903ReRF-PredPrediction of a...ognenic regionsTeng et al. 20211 (amyloidogeni...-amyloidogenic)
8SequenceSEQ_CAPSID7935424.030246336468038644071VIRALproPrediction of capdsid proteinsGaliez et al., 20161 (capsid prote...capsid protein)
9SequenceSEQ_DISULFIDE2547241.2524546144708971650DiproPrediction of d...es in sequencesCheng et al., 20061 (sequence wit...ithout SS bond)
10SequenceSEQ_LOCATION1835399.1269757323981045790nanPrediction of s...lasma membrane)Shen et al., 20191 (protein in c...asma membrane)
11SequenceSEQ_SOLUBLE17408254.611041443226987048704SOLproPrediction of s...oluble proteinsMagnan et al., 20091 (soluble), 0 (insoluble)
12SequenceSEQ_TAIL6668400.673365267169025744094VIRALproPrediction of tail proteinsGaliez et al., 20161 (tail protein...n-tail protein)
13DomainDOM_GSEC126737.809524929646363nanPrediction of g...tase substratesBreimann et al, 2024c1 (substrate), ...(non-substrate)
14DomainDOM_GSEC_PU694712.570605494524630nanPrediction of g...es (PU dataset)Breimann et al, 2024c1 (substrate), ...bstrate status)1Amino acidAA_CASPASE3233796.587983185605705184900PROSPERousPrediction of c...3 cleavage siteSong et al., 20181 (adjacent to ... cleavage site)
2Amino acidAA_FURIN71831.0281695900316358840PROSPERousPrediction of f...n cleavage siteSong et al., 20181 (adjacent to ... cleavage site)
3Amino acidAA_LDR342345.7543861182483546982779IDP-Seq2SeqPrediction of l...d regions (LDR)Tang et al., 20201 (disordered), 0 (ordered)
4Amino acidAA_MMP2573546.2059343129762416310560PROSPERousPrediction of M...) cleavage siteSong et al., 20181 (adjacent to ... cleavage site)
5Amino acidAA_RNABIND221248.87330355001649248509GMKSVM-RUPrediction of R...(RBP60 dataset)Yang et al., 20211 (binding), 0 (non-binding)
6Amino acidAA_SA233796.58798318560510108284523PROSPERousPrediction of s...PASE3 data set)Song et al., 20181 (exposed/acce...non-accessible)
7SequenceSEQ_AMYLO14146.0000008484511903ReRF-PredPrediction of a...ognenic regionsTeng et al. 20211 (amyloidogeni...-amyloidogenic)
8SequenceSEQ_CAPSID7935424.030246336468038644071VIRALproPrediction of capdsid proteinsGaliez et al., 20161 (capsid prote...capsid protein)
9SequenceSEQ_DISULFIDE2547241.2524546144708971650DiproPrediction of d...es in sequencesCheng et al., 20061 (sequence wit...ithout SS bond)
10SequenceSEQ_LOCATION1835399.1269757323981045790nanPrediction of s...lasma membrane)Shen et al., 20191 (protein in c...asma membrane)
11SequenceSEQ_SOLUBLE17408254.611041443226987048704SOLproPrediction of s...oluble proteinsMagnan et al., 20091 (soluble), 0 (insoluble)
12SequenceSEQ_TAIL6668400.673365267169025744094VIRALproPrediction of tail proteinsGaliez et al., 20161 (tail protein...n-tail protein)
13DomainDOM_GSEC126737.809524929646363nanPrediction of g...tase substratesBreimann et al, 2024c1 (substrate), ...(non-substrate)
14DomainDOM_GSEC_PU694712.570605494524630nanPrediction of g...es (PU dataset)Breimann et al, 2024c1 (substrate), ...bstrate status)
\n" @@ -310,10 +310,10 @@ }, "collapsed": false, "execution": { - "iopub.execute_input": "2026-09-14T21:28:38.531289Z", - "iopub.status.busy": "2026-09-14T21:28:38.531122Z", - "iopub.status.idle": "2026-09-14T21:28:38.618552Z", - "shell.execute_reply": "2026-09-14T21:28:38.618269Z" + "iopub.execute_input": "2026-09-18T22:09:33.366642Z", + "iopub.status.busy": "2026-09-18T22:09:33.366559Z", + "iopub.status.idle": "2026-09-18T22:09:33.424844Z", + "shell.execute_reply": "2026-09-18T22:09:33.424620Z" } }, "outputs": [ @@ -321,63 +321,63 @@ "data": { "text/html": [ "\n", - "\n", + "
\n", " \n", " \n", " \n", - " \n", - " \n", - " \n", - " \n", + " \n", + " \n", + " \n", + " \n", " \n", " \n", " \n", " \n", - " \n", - " \n", - " \n", - " \n", - " \n", + " \n", + " \n", + " \n", + " \n", + " \n", " \n", " \n", - " \n", - " \n", - " \n", - " \n", - " \n", + " \n", + " \n", + " \n", + " \n", + " \n", " \n", " \n", - " \n", - " \n", - " \n", - " \n", - " \n", + " \n", + " \n", + " \n", + " \n", + " \n", " \n", " \n", - " \n", - " \n", - " \n", - " \n", - " \n", + " \n", + " \n", + " \n", + " \n", + " \n", " \n", " \n", "
 entrygenesequencelabelentrygenesequencelabel
1CAPSID_1name_1MVTHNVKINKHVTRR...DTPRIPATKLDEENV01CAPSID_1name_1MVTHNVKINKHVTRR...DTPRIPATKLDEENV0
2CAPSID_2name_2MKKRQKKMTLSNFTD...AMLEAVINARHFGEE02CAPSID_2name_2MKKRQKKMTLSNFTD...AMLEAVINARHFGEE0
3CAPSID_4072name_4072MALTTNDVITEDFVR...AWKAIFPEAAVKVDA13CAPSID_4072name_4072MALTTNDVITEDFVR...AWKAIFPEAAVKVDA1
4CAPSID_4073name_4073MGELTDNGVQLAKAQ...TCTNPAAHAKIRDLK14CAPSID_4073name_4073MGELTDNGVQLAKAQ...TCTNPAAHAKIRDLK1
\n" @@ -402,7 +402,7 @@ "collapsed": false }, "source": [ - "The sampling can be performed randomly by setting ``random=True``: " + "The sampling can be performed randomly by setting ``random=True``. Provide a fixed seed via ``random_state`` to make the random selection reproducible (the global ``aa.options['random_state']`` setting overrides it):" ] }, { @@ -416,10 +416,10 @@ }, "collapsed": false, "execution": { - "iopub.execute_input": "2026-09-14T21:28:38.619732Z", - "iopub.status.busy": "2026-09-14T21:28:38.619656Z", - "iopub.status.idle": "2026-09-14T21:28:38.662218Z", - "shell.execute_reply": "2026-09-14T21:28:38.661947Z" + "iopub.execute_input": "2026-09-18T22:09:33.425747Z", + "iopub.status.busy": "2026-09-18T22:09:33.425685Z", + "iopub.status.idle": "2026-09-18T22:09:33.469856Z", + "shell.execute_reply": "2026-09-18T22:09:33.469644Z" } }, "outputs": [ @@ -427,63 +427,63 @@ "data": { "text/html": [ "\n", - "\n", + "
\n", " \n", " \n", " \n", - " \n", - " \n", - " \n", - " \n", + " \n", + " \n", + " \n", + " \n", " \n", " \n", " \n", " \n", - " \n", - " \n", - " \n", - " \n", - " \n", + " \n", + " \n", + " \n", + " \n", + " \n", " \n", " \n", - " \n", - " \n", - " \n", - " \n", - " \n", + " \n", + " \n", + " \n", + " \n", + " \n", " \n", " \n", - " \n", - " \n", - " \n", - " \n", - " \n", + " \n", + " \n", + " \n", + " \n", + " \n", " \n", " \n", - " \n", - " \n", - " \n", - " \n", - " \n", + " \n", + " \n", + " \n", + " \n", + " \n", " \n", " \n", "
 entrygenesequencelabelentrygenesequencelabel
1CAPSID_2024name_2024MSKAGEIYYDIKIRE...NNRIIRTNRAQGVIL01CAPSID_2839name_2839MDTADIVWVEESVSA...KFVYPFDDKMSFLFA0
2CAPSID_2227name_2227MADIPPDPPALNTTP...PAACPYPCGHTFLRP02CAPSID_151name_151MPDTTTDPRPAPPAP...ADKRTHLLNRAPGHS0
3CAPSID_6439name_6439MLELEFEDVPNNIGS...SSAAREFLSKFGIRM13CAPSID_6241name_6241MRLFLRDFFMSMYTT...MLLADPDEFVSVQLA1
4CAPSID_5062name_5062AEGDDPAKAAFNSLQ...TIGIKLFKKFTSKAS14CAPSID_6362name_6362MSRLGPALDEFRSFK...PDSQSTGGDLSVIDS1
\n" @@ -497,7 +497,7 @@ } ], "source": [ - "df_seq = aa.load_dataset(name=\"SEQ_CAPSID\", n=2, random=True)\n", + "df_seq = aa.load_dataset(name=\"SEQ_CAPSID\", n=2, random=True, random_state=42)\n", "aa.display_df(df=df_seq)" ] }, @@ -522,10 +522,10 @@ }, "collapsed": false, "execution": { - "iopub.execute_input": "2026-09-14T21:28:38.663641Z", - "iopub.status.busy": "2026-09-14T21:28:38.663557Z", - "iopub.status.idle": "2026-09-14T21:28:38.679470Z", - "shell.execute_reply": "2026-09-14T21:28:38.679231Z" + "iopub.execute_input": "2026-09-18T22:09:33.470766Z", + "iopub.status.busy": "2026-09-18T22:09:33.470705Z", + "iopub.status.idle": "2026-09-18T22:09:33.484700Z", + "shell.execute_reply": "2026-09-18T22:09:33.484488Z" } }, "outputs": [ @@ -564,10 +564,10 @@ }, "collapsed": false, "execution": { - "iopub.execute_input": "2026-09-14T21:28:38.680525Z", - "iopub.status.busy": "2026-09-14T21:28:38.680438Z", - "iopub.status.idle": "2026-09-14T21:28:38.686082Z", - "shell.execute_reply": "2026-09-14T21:28:38.685824Z" + "iopub.execute_input": "2026-09-18T22:09:33.485632Z", + "iopub.status.busy": "2026-09-18T22:09:33.485568Z", + "iopub.status.idle": "2026-09-18T22:09:33.490371Z", + "shell.execute_reply": "2026-09-18T22:09:33.490186Z" } }, "outputs": [ @@ -598,10 +598,10 @@ "id": "85ac404f", "metadata": { "execution": { - "iopub.execute_input": "2026-09-14T21:28:38.687503Z", - "iopub.status.busy": "2026-09-14T21:28:38.687414Z", - "iopub.status.idle": "2026-09-14T21:28:38.694645Z", - "shell.execute_reply": "2026-09-14T21:28:38.694339Z" + "iopub.execute_input": "2026-09-18T22:09:33.491200Z", + "iopub.status.busy": "2026-09-18T22:09:33.491145Z", + "iopub.status.idle": "2026-09-18T22:09:33.496338Z", + "shell.execute_reply": "2026-09-18T22:09:33.496148Z" } }, "outputs": [ @@ -652,10 +652,10 @@ }, "collapsed": false, "execution": { - "iopub.execute_input": "2026-09-14T21:28:38.696133Z", - "iopub.status.busy": "2026-09-14T21:28:38.696013Z", - "iopub.status.idle": "2026-09-14T21:28:38.803865Z", - "shell.execute_reply": "2026-09-14T21:28:38.803572Z" + "iopub.execute_input": "2026-09-18T22:09:33.497148Z", + "iopub.status.busy": "2026-09-18T22:09:33.497097Z", + "iopub.status.idle": "2026-09-18T22:09:33.599518Z", + "shell.execute_reply": "2026-09-18T22:09:33.599303Z" } }, "outputs": [ @@ -663,58 +663,58 @@ "data": { "text/html": [ "\n", - "\n", + "
\n", " \n", " \n", " \n", - " \n", - " \n", - " \n", + " \n", + " \n", + " \n", " \n", " \n", " \n", " \n", - " \n", - " \n", - " \n", - " \n", + " \n", + " \n", + " \n", + " \n", " \n", " \n", - " \n", - " \n", - " \n", - " \n", + " \n", + " \n", + " \n", + " \n", " \n", " \n", - " \n", - " \n", - " \n", - " \n", + " \n", + " \n", + " \n", + " \n", " \n", " \n", - " \n", - " \n", - " \n", - " \n", + " \n", + " \n", + " \n", + " \n", " \n", " \n", "
 entrysequencelabelentrysequencelabel
1CASPASE3_1_pos2MSLFD01CASPASE3_1_pos2MSLFD0
2CASPASE3_1_pos3SLFDL02CASPASE3_1_pos3SLFDL0
3CASPASE3_1_pos126LRDSM13CASPASE3_1_pos126LRDSM1
4CASPASE3_1_pos127RDSML14CASPASE3_1_pos127RDSML1
\n" @@ -753,10 +753,10 @@ }, "collapsed": false, "execution": { - "iopub.execute_input": "2026-09-14T21:28:38.805071Z", - "iopub.status.busy": "2026-09-14T21:28:38.804978Z", - "iopub.status.idle": "2026-09-14T21:28:38.819941Z", - "shell.execute_reply": "2026-09-14T21:28:38.819703Z" + "iopub.execute_input": "2026-09-18T22:09:33.600439Z", + "iopub.status.busy": "2026-09-18T22:09:33.600377Z", + "iopub.status.idle": "2026-09-18T22:09:33.612642Z", + "shell.execute_reply": "2026-09-18T22:09:33.612448Z" } }, "outputs": [ @@ -764,112 +764,112 @@ "data": { "text/html": [ "\n", - "\n", + "
\n", " \n", " \n", " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", " \n", " \n", " \n", " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", - " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", " \n", " \n", "
 entrygenesequencelabeltmd_starttmd_stopjmd_ntmdjmd_centrygenesequencelabeltmd_starttmd_stopjmd_ntmdjmd_c
1P05067APPMLPGLALLLLAAWTA...GYENPTYKFFEQMQN1701723FAEDVGSNKGAIIGLMVGGVVIATVIVITLVMLKKKQYTSIHH
2P14925PamMAGRARSGLLLLLLG...EEEYSAPLPKPAPSS1868890KLSTEPGSGVSVVLITTLLVIPVLVLLAIVMFIRWKKSRAFGD
3P70180Npr3MRSLLLFTFSACVLL...RELREDSIRSHFSVA1477499PCKSSGGLEESAVTGIVVGALLGAGLLMAFYFFRKKYRITIER
4P12821ACEMGAASGRRGPGLLLP...SHGPQFGSEVELRHS212571276GLDLDAQQARVGQWLLLFLGIALLVATLGLSQRLFSIRHR
5P36896ACVR1BMAESAGASSFFPLVV...KKTLSQLSVQEDVKI2127149EHPSMWGPVELVGIIAGPVFLLFLIIIIVFLVINYHQRVYHNR
6Q8NER5ACVR1CMTRALCSALRQALLL...KKTISQLCVKEDCKA2114136PNAPKLGPMELAIIITVPVCLLSIAAMLTVWACQGRQCSYRKK1P05067APPMLPGLALLLLAAWTA...GYENPTYKFFEQMQN1701723FAEDVGSNKGAIIGLMVGGVVIATVIVITLVMLKKKQYTSIHH
2P14925PamMAGRARSGLLLLLLG...EEEYSAPLPKPAPSS1868890KLSTEPGSGVSVVLITTLLVIPVLVLLAIVMFIRWKKSRAFGD
3P70180Npr3MRSLLLFTFSACVLL...RELREDSIRSHFSVA1477499PCKSSGGLEESAVTGIVVGALLGAGLLMAFYFFRKKYRITIER
4P12821ACEMGAASGRRGPGLLLP...SHGPQFGSEVELRHS212571276GLDLDAQQARVGQWLLLFLGIALLVATLGLSQRLFSIRHR
5P36896ACVR1BMAESAGASSFFPLVV...KKTLSQLSVQEDVKI2127149EHPSMWGPVELVGIIAGPVFLLFLIIIIVFLVINYHQRVYHNR
6Q8NER5ACVR1CMTRALCSALRQALLL...KKTISQLCVKEDCKA2114136PNAPKLGPMELAIIITVPVCLLSIAAMLTVWACQGRQCSYRKK
\n" @@ -904,7 +904,7 @@ "name": "python", "nbconvert_exporter": "python", "pygments_lexer": "ipython3", - "version": "3.14.0" + "version": "3.13.11" } }, "nbformat": 4, diff --git a/tests/unit/data_handling_tests/test_load_dataset.py b/tests/unit/data_handling_tests/test_load_dataset.py index 641e5df8b..d920d9fe2 100644 --- a/tests/unit/data_handling_tests/test_load_dataset.py +++ b/tests/unit/data_handling_tests/test_load_dataset.py @@ -1,6 +1,12 @@ """ This is a script for testing the aa.load_dataset function. """ +import hashlib +import os +import subprocess +import sys +from pathlib import Path + from hypothesis import given, example, settings import hypothesis.strategies as some import numpy as np @@ -382,3 +388,157 @@ def test_resolves_sample_by_gene_end_to_end(self): sf = aa.SequenceFeature(verbose=False) df_parts = sf.get_df_parts(df_seq=df) assert resolve_sample_entry(df_seq=df, df_parts=df_parts, sample="APP") == "P05067" + + +# Int-labelled datasets, whose class blocks are concatenated in a process-independent +# order, so a seeded frame can be compared byte-for-byte across processes. +LIST_SEEDED_NAMES = ["DOM_GSEC", "SEQ_CAPSID", "SEQ_LOCATION"] + + +def _digest(df_seq): + """Byte-level digest of a frame (values, index, shape, and dtypes).""" + meta = f"{df_seq.shape}|{[f'{c}:{t}' for c, t in df_seq.dtypes.astype(str).items()]}" + return hashlib.sha256(df_seq.to_csv(index=True).encode("utf-8") + meta.encode("utf-8")).hexdigest() + + +class TestLoadDatasetRandomState: + """Test the 'random_state' parameter of load_dataset (issue #582).""" + + # Reproducibility under an explicit seed + def test_same_seed_returns_identical_frame(self): + """The same seed reproduces the sample exactly.""" + df_1 = aa.load_dataset(name="DOM_GSEC", n=5, random=True, random_state=42) + df_2 = aa.load_dataset(name="DOM_GSEC", n=5, random=True, random_state=42) + pd.testing.assert_frame_equal(df_1, df_2) + + @pytest.mark.parametrize("name", LIST_SEEDED_NAMES) + def test_same_seed_identical_for_every_level(self, name): + """A fixed seed is reproducible on the sequence and domain levels alike.""" + df_1 = aa.load_dataset(name=name, n=3, random=True, random_state=0) + df_2 = aa.load_dataset(name=name, n=3, random=True, random_state=0) + assert _digest(df_1) == _digest(df_2) + + @settings(max_examples=10, deadline=None) + @given(random_state=some.integers(min_value=0, max_value=10000)) + def test_any_seed_is_reproducible(self, random_state): + """Any valid seed reproduces its own sample (property-based).""" + kwargs = dict(name="SEQ_CAPSID", n=3, random=True, random_state=random_state) + pd.testing.assert_frame_equal(aa.load_dataset(**kwargs), aa.load_dataset(**kwargs)) + + def test_seed_zero_is_honoured(self): + """random_state=0 is a valid seed, not treated as 'unset'.""" + df_1 = aa.load_dataset(name="DOM_GSEC", n=5, random=True, random_state=0) + df_2 = aa.load_dataset(name="DOM_GSEC", n=5, random=True, random_state=0) + pd.testing.assert_frame_equal(df_1, df_2) + + def test_different_seeds_give_different_samples(self): + """Different seeds draw different samples (not a constant selection).""" + digests = {_digest(aa.load_dataset(name="SEQ_LOCATION", n=5, random=True, random_state=rs)) + for rs in [0, 1, 2, 3, 42]} + assert len(digests) > 1 + + def test_seeded_sample_is_balanced(self): + """A seeded random selection still returns n entries per class.""" + df_seq = aa.load_dataset(name="SEQ_LOCATION", n=5, random=True, random_state=42) + assert len(df_seq) == 5 * 2 + assert set(df_seq[ut.COL_LABEL]) == {0, 1} + + def test_columns_and_dtypes_unchanged_by_seed(self): + """Seeding changes which rows are drawn, never the columns or their dtypes.""" + df_unseeded = aa.load_dataset(name="DOM_GSEC", n=5, random=True) + df_seeded = aa.load_dataset(name="DOM_GSEC", n=5, random=True, random_state=42) + assert list(df_unseeded.columns) == list(df_seeded.columns) + assert df_unseeded.dtypes.astype(str).to_dict() == df_seeded.dtypes.astype(str).to_dict() + + def test_reproducible_across_processes(self): + """A seeded frame is identical in a fresh interpreter (KPI: across processes).""" + df_seq = aa.load_dataset(name="DOM_GSEC", n=5, random=True, random_state=42) + code = ( + "import hashlib, warnings; warnings.simplefilter('ignore');" + "import aaanalysis as aa;" + "df = aa.load_dataset(name='DOM_GSEC', n=5, random=True, random_state=42);" + "meta = f\"{df.shape}|{[f'{c}:{t}' for c, t in df.dtypes.astype(str).items()]}\";" + "print(hashlib.sha256(df.to_csv(index=True).encode('utf-8') + meta.encode('utf-8')).hexdigest())" + ) + env = dict(os.environ, MPLBACKEND="Agg", + PYTHONPATH=str(Path(aa.__file__).resolve().parents[1])) + out = subprocess.run([sys.executable, "-c", code], capture_output=True, text=True, env=env) + assert out.returncode == 0, out.stderr + assert out.stdout.strip().splitlines()[-1] == _digest(df_seq) + + # Global option override + def test_option_random_state_is_honoured(self): + """options['random_state'] seeds the sampling with no explicit argument.""" + aa.options["random_state"] = 42 + df_1 = aa.load_dataset(name="DOM_GSEC", n=5, random=True) + df_2 = aa.load_dataset(name="DOM_GSEC", n=5, random=True) + pd.testing.assert_frame_equal(df_1, df_2) + + def test_option_random_state_matches_explicit_seed(self): + """The option resolves to the same sample as passing that seed directly.""" + df_explicit = aa.load_dataset(name="DOM_GSEC", n=5, random=True, random_state=42) + aa.options["random_state"] = 42 + df_option = aa.load_dataset(name="DOM_GSEC", n=5, random=True) + pd.testing.assert_frame_equal(df_explicit, df_option) + + def test_option_random_state_overrides_argument(self): + """options['random_state'] wins over the per-call argument (documented contract).""" + df_explicit = aa.load_dataset(name="DOM_GSEC", n=5, random=True, random_state=42) + aa.options["random_state"] = 42 + df_option = aa.load_dataset(name="DOM_GSEC", n=5, random=True, random_state=7) + pd.testing.assert_frame_equal(df_explicit, df_option) + + # random=False is unaffected + @pytest.mark.parametrize("name", LIST_SEEDED_NAMES + ["AA_CASPASE3", "DOM_GSEC_PU"]) + def test_random_false_ignores_seed(self, name): + """random=False returns the head-of-class selection, seed or no seed.""" + df_plain = aa.load_dataset(name=name, n=3) + assert _digest(aa.load_dataset(name=name, n=3, random_state=42)) == _digest(df_plain) + assert _digest(aa.load_dataset(name=name, n=3, random_state=7)) == _digest(df_plain) + + def test_random_false_ignores_option(self): + """A global random_state does not touch the deterministic selection either.""" + df_plain = aa.load_dataset(name="DOM_GSEC", n=5) + aa.options["random_state"] = 42 + pd.testing.assert_frame_equal(aa.load_dataset(name="DOM_GSEC", n=5), df_plain) + + def test_full_dataset_ignores_seed(self): + """Without 'n' there is no sampling, so the seed cannot change the output.""" + df_plain = aa.load_dataset(name="SEQ_CAPSID") + assert _digest(aa.load_dataset(name="SEQ_CAPSID", random_state=42)) == _digest(df_plain) + + # Invalid input + def test_invalid_random_state_negative(self): + """A negative seed is rejected.""" + with pytest.raises(ValueError): + aa.load_dataset(name="DOM_GSEC", n=5, random=True, random_state=-1) + + @settings(max_examples=10, deadline=None) + @given(random_state=some.integers(max_value=-1)) + def test_invalid_random_state_negative_property(self, random_state): + """Every negative seed is rejected (property-based).""" + with pytest.raises(ValueError): + aa.load_dataset(name="DOM_GSEC", n=5, random=True, random_state=random_state) + + def test_invalid_random_state_float(self): + """A non-integer seed is rejected.""" + with pytest.raises(ValueError): + aa.load_dataset(name="DOM_GSEC", n=5, random=True, random_state=1.5) + + def test_invalid_random_state_str(self): + """A string seed is rejected with a message naming the parameter.""" + with pytest.raises(ValueError, match="random_state"): + aa.load_dataset(name="DOM_GSEC", n=5, random=True, random_state="42") + + def test_invalid_random_state_rejected_without_random(self): + """The seed is validated even when random=False (fail fast, not silently).""" + with pytest.raises(ValueError): + aa.load_dataset(name="DOM_GSEC", n=5, random_state=-5) + + def test_option_random_state_read_per_call(self): + """The option is resolved on every call, so changing it changes the next sample.""" + aa.options["random_state"] = 42 + df_1 = aa.load_dataset(name="SEQ_LOCATION", n=5, random=True) + aa.options["random_state"] = 7 + df_2 = aa.load_dataset(name="SEQ_LOCATION", n=5, random=True) + assert _digest(df_1) != _digest(df_2)