diff --git a/projects/data/data.smk b/projects/data/data.smk index f43c58b62..431900bb7 100644 --- a/projects/data/data.smk +++ b/projects/data/data.smk @@ -413,6 +413,8 @@ The PSDs are those of the last fetched train-background chunk. lowpass=config["lowpass"] or "null", snr_threshold=config["snr_threshold"], max_num_samples=config["max_num_samples"], + # offset by branch so each branch samples different parameters + seed=lambda wc: config["seed"] + int(wc.vbranch_id), shell: "generate-validation-waveforms" " --num_signals {params.num_signals}" @@ -429,6 +431,7 @@ The PSDs are those of the last fetched train-background chunk. " --snr_threshold {params.snr_threshold}" " --psd {input.psd_file}" " --max_num_samples {params.max_num_samples}" + " --seed {params.seed}" " --output_file {output}" " &> {log}" diff --git a/projects/data/data/waveforms/rejection.py b/projects/data/data/waveforms/rejection.py index 92cef3f76..e7a739829 100644 --- a/projects/data/data/waveforms/rejection.py +++ b/projects/data/data/waveforms/rejection.py @@ -4,6 +4,7 @@ import numpy as np import torch +from bilby.core.utils import random as bilby_random from data.waveforms.utils import convert_to_detector_frame, load_psds from ledger.injections import ( BilbyParameterSet, @@ -30,7 +31,12 @@ def rejection_sample( snr_threshold: float, psd: Path | torch.Tensor, max_num_samples: int, + seed: int | None = None, ) -> tuple[ResponseSetFields, InjectionParameterSet]: + # bilby priors sample from bilby's own generator + if seed is not None: + bilby_random.seed(seed) + # get the detector tensors and vertices # for projecting our waveforms tensors, vertices = get_ifo_geometry(*ifos) diff --git a/projects/data/data/waveforms/utils.py b/projects/data/data/waveforms/utils.py index 8c155f300..60bd9a8a7 100644 --- a/projects/data/data/waveforms/utils.py +++ b/projects/data/data/waveforms/utils.py @@ -1,12 +1,13 @@ +import hashlib import logging import random import time from pathlib import Path -from zlib import adler32 import h5py import numpy as np import torch +from bilby.core.utils import random as bilby_random from gwpy.timeseries import TimeSeriesDict @@ -22,13 +23,16 @@ def seed_worker( start: float, stop: float, shifts: list[float], seed: int ) -> np.random.Generator: fingerprint = str((start, stop) + tuple(shifts)) - worker_hash = adler32(fingerprint.encode()) + digest = hashlib.sha256(fingerprint.encode()).digest() + worker_hash = int.from_bytes(digest[:8], "big") combined = seed + worker_hash logging.info( f"Seeding data generation with seed {seed}, " f"augmented by worker seed {worker_hash}" ) random.seed(combined) + # bilby priors sample from bilby's own generator + bilby_random.seed(combined) return np.random.default_rng(combined)