From a42861f33fc1813b2ea2bee21d7f334bf1340408 Mon Sep 17 00:00:00 2001 From: William Benoit Date: Mon, 5 Oct 2026 11:05:42 -0400 Subject: [PATCH] Add replay_id to online search --- projects/online/config.yaml | 1 + .../online/online/dataloading/__init__.py | 2 +- projects/online/online/dataloading/arrakis.py | 23 +++++++++++++++---- projects/online/online/main.py | 19 ++++++++++++++- 4 files changed, 38 insertions(+), 7 deletions(-) diff --git a/projects/online/config.yaml b/projects/online/config.yaml index 7c7054b02..c698ab20a 100644 --- a/projects/online/config.yaml +++ b/projects/online/config.yaml @@ -46,6 +46,7 @@ amplfi_parameter_sampler: ./prior.yaml channels: ["H1:GDS-CALIB_STRAIN_CLEAN_INJ1_O4Replay", "L1:GDS-CALIB_STRAIN_CLEAN_INJ1_O4Replay", "V1:Hrec_hoft_16384Hz_INJ1_O4Replay"] state_channels: ["H1:GDS-CALIB_STATE_VECTOR", "L1:GDS-CALIB_STATE_VECTOR", "V1:DQ_ANALYSIS_STATE_VECTOR"] data_source: "frames" +replay_id: null sample_rate: 2048 astro_event_rate: 31 kernel_length: 1.5 diff --git a/projects/online/online/dataloading/__init__.py b/projects/online/online/dataloading/__init__.py index 65e18a521..f1481182a 100644 --- a/projects/online/online/dataloading/__init__.py +++ b/projects/online/online/dataloading/__init__.py @@ -1,4 +1,4 @@ from .arrakis import data_iterator as arrakis_data_iterator -from .arrakis import get_block_duration, stream_channels +from .arrakis import check_replay_id, get_block_duration, stream_channels from .offline import offline_data_iterator from .online import data_iterator diff --git a/projects/online/online/dataloading/arrakis.py b/projects/online/online/dataloading/arrakis.py index 40202e739..eff094d39 100644 --- a/projects/online/online/dataloading/arrakis.py +++ b/projects/online/online/dataloading/arrakis.py @@ -23,8 +23,20 @@ def stream_channels( return channels +def check_replay_id(replay_id: str) -> None: + """Check that `replay_id` is one registered on the Arrakis server""" + replays = Client().replays() + if replay_id not in replays: + raise ValueError( + f"Unknown replay ID {replay_id}. " + f"Available replays: {sorted(replays)}" + ) + + def get_block_duration( - channels: list[str], metadata: dict | None = None + channels: list[str], + metadata: dict | None = None, + replay_id: str | None = None, ) -> float: """ The cadence at which the server will deliver blocks of @@ -32,7 +44,7 @@ def get_block_duration( of the individual stride of each channel. """ if not metadata: - metadata = Client().describe(channels) + metadata = Client().describe(channels, replay_id=replay_id) strides = [metadata[channel].stride for channel in channels] # Strides are returned in nanoseconds return lcm(*strides) / Time.SECONDS @@ -59,13 +71,14 @@ def data_iterator( sample_rate: float, state_channels: dict[str, str] | None = None, numtaps: int | None = 60, + replay_id: str | None = None, ) -> torch.Tensor: channels = stream_channels(strain_channels, ifos, state_channels) client = Client() - metadata = client.describe(channels) + metadata = client.describe(channels, replay_id=replay_id) strain_sample_rate = get_strain_sample_rate(strain_channels, metadata) - block_duration = get_block_duration(channels, metadata) + block_duration = get_block_duration(channels, metadata, replay_id) # build resampling filter factor = strain_sample_rate / sample_rate @@ -97,7 +110,7 @@ def data_iterator( # a discontinuous jump expected_t0 = None - blocks = client.stream(channels) + blocks = client.stream(channels, replay_id=replay_id) for block in blocks: # Check if the expected t0 differs by more than half a sample discontinuous = expected_t0 is not None and ( diff --git a/projects/online/online/main.py b/projects/online/online/main.py index 75cccdf72..8a9b7e4bb 100644 --- a/projects/online/online/main.py +++ b/projects/online/online/main.py @@ -21,6 +21,7 @@ data_iterator, offline_data_iterator, arrakis_data_iterator, + check_replay_id, get_block_duration, stream_channels, ) @@ -384,6 +385,7 @@ def main( integration_window_length: float, astro_event_rate: float, data_source: Literal["frames", "arrakis"] = "frames", + replay_id: str | None = None, state_channels: Optional[list[str]] = None, fftlength: Optional[float] = None, highpass: Optional[float] = None, @@ -469,6 +471,11 @@ def main( Length of output integration window in seconds astro_event_rate: Prior on rate of astrophysical events in units Gpc^-3 yr^-1 + replay_id: + Arrakis replay to stream from. Must be one of the + replays registered on the Arrakis server, and some + channels are only available within a replay. Only + valid when `data_source` is "arrakis" fftlength: FFT length in seconds (defaults to kernel_length + fduration) highpass: @@ -539,6 +546,14 @@ def main( # accounted for search_start = gps_now() + # check the replay before spawning any subprocesses + if replay_id is not None: + if data_source != "arrakis": + raise ValueError( + "replay_id should be set only when data_source='arrakis'" + ) + check_replay_id(replay_id) + # create various queues for message # passing between subprocesses error_queue = Queue() @@ -731,7 +746,8 @@ def main( if data_source == "arrakis": update_size = get_block_duration( - stream_channels(channels, ifos, state_channels) + stream_channels(channels, ifos, state_channels), + replay_id=replay_id, ) logging.info(f"Arrakis update size: {update_size} s") data_it = arrakis_data_iterator( @@ -739,6 +755,7 @@ def main( ifos=ifos, sample_rate=sample_rate, state_channels=state_channels, + replay_id=replay_id, ) elif data_source == "frames": update_size = 1