Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions projects/online/config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion projects/online/online/dataloading/__init__.py
Original file line number Diff line number Diff line change
@@ -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
23 changes: 18 additions & 5 deletions projects/online/online/dataloading/arrakis.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,16 +23,28 @@ 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
data, which is determined by the least common multiple
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
Expand All @@ -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
Expand Down Expand Up @@ -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 (
Expand Down
19 changes: 18 additions & 1 deletion projects/online/online/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@
data_iterator,
offline_data_iterator,
arrakis_data_iterator,
check_replay_id,
get_block_duration,
stream_channels,
)
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -731,14 +746,16 @@ 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(
strain_channels=channels,
ifos=ifos,
sample_rate=sample_rate,
state_channels=state_channels,
replay_id=replay_id,
)
elif data_source == "frames":
update_size = 1
Expand Down
Loading