Skip to content

Commit 9600f84

Browse files
committed
feat: Implement deterministic seed management for reproducibility (#6)
Add comprehensive reproducibility support for all random operations: - Add SeedManager class in core/reproducibility.py for global seed control - Add seed parameter to GridFIA API constructor and set_seed() method - Fix parallel processing workers (_bootstrap_worker, _permutation_worker) to accept and use deterministic seeds when global seed is set - Fix statistical analysis (permutation test, bootstrap test) to use reproducible random states when global seed is set - Add comprehensive test suite with 28 tests covering all reproducibility scenarios including API, parallel workers, and statistical analysis Examples: # Set seed at API initialization api = GridFIA(seed=42) # Or set seed after initialization api.set_seed(42) # Or use SeedManager directly from gridfia.core.reproducibility import set_seed set_seed(42) Fixes #6
1 parent 9cb3340 commit 9600f84

5 files changed

Lines changed: 827 additions & 30 deletions

File tree

gridfia/api.py

Lines changed: 46 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -85,7 +85,11 @@ class GridFIA:
8585
>>> api.create_maps("data/nc_forest.zarr", map_type="diversity")
8686
"""
8787

88-
def __init__(self, config: Optional[Union[str, Path, GridFIASettings]] = None):
88+
def __init__(
89+
self,
90+
config: Optional[Union[str, Path, GridFIASettings]] = None,
91+
seed: Optional[int] = None
92+
):
8993
"""
9094
Initialize GridFIA API.
9195
@@ -94,6 +98,18 @@ def __init__(self, config: Optional[Union[str, Path, GridFIASettings]] = None):
9498
config : str, Path, or GridFIASettings, optional
9599
Configuration file path or settings object.
96100
If None, uses default settings.
101+
seed : int, optional
102+
Random seed for reproducibility. If provided, all random
103+
operations (bootstrap, permutation tests, etc.) will be
104+
deterministic.
105+
106+
Examples
107+
--------
108+
>>> api = GridFIA(seed=42) # Reproducible results
109+
>>> result1 = api.calculate_metrics(zarr_path)
110+
>>>
111+
>>> api2 = GridFIA(seed=42) # Same seed = same results
112+
>>> result2 = api2.calculate_metrics(zarr_path)
97113
"""
98114
if config is None:
99115
self.settings = GridFIASettings()
@@ -102,10 +118,39 @@ def __init__(self, config: Optional[Union[str, Path, GridFIASettings]] = None):
102118
else:
103119
self.settings = config
104120

121+
# Set seed for reproducibility
122+
self._seed = seed
123+
if seed is not None:
124+
from .core.reproducibility import SeedManager
125+
SeedManager.set_global_seed(seed)
126+
105127
# Lock for thread-safe lazy initialization of components
106128
self._init_lock = threading.Lock()
107129
self._rest_client = None
108130
self._processor = None
131+
132+
def set_seed(self, seed: int) -> None:
133+
"""
134+
Set random seed for reproducibility.
135+
136+
Parameters
137+
----------
138+
seed : int
139+
Random seed value.
140+
141+
Examples
142+
--------
143+
>>> api = GridFIA()
144+
>>> api.set_seed(42)
145+
"""
146+
from .core.reproducibility import SeedManager
147+
self._seed = seed
148+
SeedManager.set_global_seed(seed)
149+
150+
@property
151+
def seed(self) -> Optional[int]:
152+
"""Get current random seed."""
153+
return self._seed
109154

110155
@property
111156
def rest_client(self) -> BigMapRestClient:

gridfia/core/analysis/statistical_analysis.py

Lines changed: 44 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -391,21 +391,39 @@ def _permutation_test(
391391
# Sequential permutation test implementation
392392
# Observed difference
393393
observed_diff = group1.mean() - group2.mean()
394-
394+
395395
# Combine all data
396396
combined = np.concatenate([group1.values, group2.values])
397397
n1, n2 = len(group1), len(group2)
398-
398+
399+
# Check for global seed from SeedManager for reproducibility
400+
try:
401+
from ..reproducibility import SeedManager
402+
global_seed = SeedManager.get_seed()
403+
except ImportError:
404+
global_seed = None
405+
406+
# Create random state for reproducibility if seed is set
407+
if global_seed is not None:
408+
rng = np.random.RandomState(global_seed)
409+
else:
410+
rng = np.random
411+
399412
# Permutation distribution
400413
perm_diffs = []
401-
for _ in range(n_permutations):
402-
# Randomly permute the combined data
403-
np.random.shuffle(combined)
404-
414+
for i in range(n_permutations):
415+
# Randomly permute the combined data (using reproducible RNG if seed set)
416+
if global_seed is not None:
417+
# Create iteration-specific seed for reproducibility
418+
iter_rng = np.random.RandomState(global_seed + i)
419+
shuffled = iter_rng.permutation(combined)
420+
else:
421+
shuffled = rng.permutation(combined)
422+
405423
# Split into two groups of original sizes
406-
perm_group1 = combined[:n1]
407-
perm_group2 = combined[n1:n1+n2]
408-
424+
perm_group1 = shuffled[:n1]
425+
perm_group2 = shuffled[n1:n1+n2]
426+
409427
# Calculate difference
410428
perm_diff = perm_group1.mean() - perm_group2.mean()
411429
perm_diffs.append(perm_diff)
@@ -463,15 +481,27 @@ def _bootstrap_test(
463481
logger.debug("Parallel processing not available for bootstrap")
464482

465483
# Sequential bootstrap implementation
484+
# Check for global seed from SeedManager for reproducibility
485+
try:
486+
from ..reproducibility import SeedManager
487+
global_seed = SeedManager.get_seed()
488+
except ImportError:
489+
global_seed = None
490+
466491
# Bootstrap distributions
467492
group1_boots = []
468493
group2_boots = []
469494
diff_boots = []
470-
471-
for _ in range(n_bootstrap):
472-
# Bootstrap samples
473-
boot1 = resample(group1.values, n_samples=len(group1))
474-
boot2 = resample(group2.values, n_samples=len(group2))
495+
496+
for i in range(n_bootstrap):
497+
# Bootstrap samples with reproducible random state if seed is set
498+
if global_seed is not None:
499+
iter_seed = global_seed + i
500+
boot1 = resample(group1.values, n_samples=len(group1), random_state=iter_seed)
501+
boot2 = resample(group2.values, n_samples=len(group2), random_state=iter_seed + n_bootstrap)
502+
else:
503+
boot1 = resample(group1.values, n_samples=len(group1))
504+
boot2 = resample(group2.values, n_samples=len(group2))
475505

476506
mean1 = np.mean(boot1)
477507
mean2 = np.mean(boot2)

gridfia/core/reproducibility.py

Lines changed: 212 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,212 @@
1+
"""
2+
Reproducibility utilities for deterministic random operations.
3+
4+
This module provides seed management to ensure reproducible results
5+
across runs, including support for parallel processing.
6+
"""
7+
8+
import random
9+
from contextlib import contextmanager
10+
from typing import Optional, Generator
11+
import logging
12+
13+
import numpy as np
14+
15+
logger = logging.getLogger(__name__)
16+
17+
18+
class SeedManager:
19+
"""
20+
Manage random seeds for reproducibility across the package.
21+
22+
This class provides global seed management that affects all random
23+
operations in GridFIA, including parallel processing workers.
24+
25+
Examples
26+
--------
27+
>>> from gridfia.core.reproducibility import SeedManager
28+
>>>
29+
>>> # Set global seed
30+
>>> SeedManager.set_global_seed(42)
31+
>>>
32+
>>> # Use temporary seed for specific operation
33+
>>> with SeedManager.temporary_seed(123):
34+
... result = some_random_operation()
35+
"""
36+
37+
_global_seed: Optional[int] = None
38+
_random_state: Optional[np.random.RandomState] = None
39+
40+
@classmethod
41+
def set_global_seed(cls, seed: int) -> None:
42+
"""
43+
Set global seed for all random operations.
44+
45+
Parameters
46+
----------
47+
seed : int
48+
Seed value for random number generators.
49+
50+
Notes
51+
-----
52+
This affects:
53+
- Python's random module
54+
- NumPy's random module
55+
- All GridFIA calculations using random operations
56+
"""
57+
cls._global_seed = seed
58+
cls._random_state = np.random.RandomState(seed)
59+
60+
# Set seeds for standard libraries
61+
random.seed(seed)
62+
np.random.seed(seed)
63+
64+
logger.info(f"Global seed set to {seed}")
65+
66+
@classmethod
67+
def get_seed(cls) -> Optional[int]:
68+
"""
69+
Get current global seed.
70+
71+
Returns
72+
-------
73+
int or None
74+
Current global seed, or None if not set.
75+
"""
76+
return cls._global_seed
77+
78+
@classmethod
79+
def get_random_state(cls) -> Optional[np.random.RandomState]:
80+
"""
81+
Get the global RandomState object.
82+
83+
Returns
84+
-------
85+
np.random.RandomState or None
86+
RandomState object for reproducible random operations.
87+
"""
88+
return cls._random_state
89+
90+
@classmethod
91+
def derive_seed(cls, offset: int = 0) -> int:
92+
"""
93+
Derive a deterministic seed from the global seed.
94+
95+
Useful for creating unique but reproducible seeds for
96+
parallel workers or nested operations.
97+
98+
Parameters
99+
----------
100+
offset : int
101+
Offset to add to the global seed.
102+
103+
Returns
104+
-------
105+
int
106+
Derived seed value.
107+
108+
Raises
109+
------
110+
ValueError
111+
If no global seed has been set.
112+
"""
113+
if cls._global_seed is None:
114+
raise ValueError("No global seed set. Call set_global_seed() first.")
115+
return cls._global_seed + offset
116+
117+
@classmethod
118+
def get_worker_seed(cls, worker_id: int) -> Optional[int]:
119+
"""
120+
Get a deterministic seed for a parallel worker.
121+
122+
Parameters
123+
----------
124+
worker_id : int
125+
Unique identifier for the worker (0, 1, 2, ...).
126+
127+
Returns
128+
-------
129+
int or None
130+
Seed for the worker, or None if no global seed is set.
131+
"""
132+
if cls._global_seed is None:
133+
return None
134+
# Use a prime multiplier to spread seeds and avoid collisions
135+
return cls._global_seed + (worker_id * 997)
136+
137+
@classmethod
138+
@contextmanager
139+
def temporary_seed(cls, seed: int) -> Generator[None, None, None]:
140+
"""
141+
Context manager for temporary seed override.
142+
143+
Parameters
144+
----------
145+
seed : int
146+
Temporary seed to use within the context.
147+
148+
Yields
149+
------
150+
None
151+
152+
Examples
153+
--------
154+
>>> with SeedManager.temporary_seed(999):
155+
... # Operations here use seed 999
156+
... result = np.random.rand()
157+
>>> # Original seed is restored here
158+
"""
159+
# Save current state
160+
old_random_state = random.getstate()
161+
old_numpy_state = np.random.get_state()
162+
old_global_seed = cls._global_seed
163+
164+
try:
165+
# Set temporary seed
166+
random.seed(seed)
167+
np.random.seed(seed)
168+
cls._global_seed = seed
169+
yield
170+
finally:
171+
# Restore original state
172+
random.setstate(old_random_state)
173+
np.random.set_state(old_numpy_state)
174+
cls._global_seed = old_global_seed
175+
176+
@classmethod
177+
def reset(cls) -> None:
178+
"""
179+
Reset seed manager to initial state (no global seed).
180+
"""
181+
cls._global_seed = None
182+
cls._random_state = None
183+
logger.info("Seed manager reset")
184+
185+
186+
def set_seed(seed: int) -> None:
187+
"""
188+
Convenience function to set global seed.
189+
190+
Parameters
191+
----------
192+
seed : int
193+
Seed value for random number generators.
194+
195+
Examples
196+
--------
197+
>>> from gridfia.core.reproducibility import set_seed
198+
>>> set_seed(42)
199+
"""
200+
SeedManager.set_global_seed(seed)
201+
202+
203+
def get_seed() -> Optional[int]:
204+
"""
205+
Convenience function to get current global seed.
206+
207+
Returns
208+
-------
209+
int or None
210+
Current global seed.
211+
"""
212+
return SeedManager.get_seed()

0 commit comments

Comments
 (0)