From d9a30cbcab7486e4c2c492c2cc64f0053b432b28 Mon Sep 17 00:00:00 2001 From: William Benoit Date: Wed, 30 Sep 2026 10:11:29 -0400 Subject: [PATCH 1/3] Simplify data rules with split wildcard, one branch checkpoint, and a single aggregate script --- pipeline/config/config.yaml | 3 +- projects/data/data.smk | 339 +++++++----------- .../scripts/aggregate_testing_waveforms.py | 38 -- .../scripts/aggregate_training_waveforms.py | 21 -- .../data/scripts/aggregate_val_waveforms.py | 29 -- projects/data/scripts/aggregate_waveforms.py | 48 +++ projects/infer/infer.smk | 5 +- projects/train/train.smk | 4 +- 8 files changed, 193 insertions(+), 294 deletions(-) delete mode 100644 projects/data/scripts/aggregate_testing_waveforms.py delete mode 100644 projects/data/scripts/aggregate_training_waveforms.py delete mode 100644 projects/data/scripts/aggregate_val_waveforms.py create mode 100644 projects/data/scripts/aggregate_waveforms.py diff --git a/pipeline/config/config.yaml b/pipeline/config/config.yaml index 5ab577ea8..c5c25360e 100644 --- a/pipeline/config/config.yaml +++ b/pipeline/config/config.yaml @@ -119,8 +119,7 @@ aggregate_rules_local: true # under slurm or condor. Rules and keys not listed here use the profile's # default-resources. resources: - fetch_train_background: {mem_mb: 6144, runtime: 480} - fetch_test_background: {mem_mb: 6144, runtime: 480} + fetch_background: {mem_mb: 6144, runtime: 480} testing_waveforms_branch: {mem_mb: 8192, runtime: 60} aggregate_testing_waveforms: {mem_mb: 2048, runtime: 10} val_waveforms_branch: {mem_mb: 6144, runtime: 60} diff --git a/projects/data/data.smk b/projects/data/data.smk index dcfef8d3e..75eb2a43c 100644 --- a/projects/data/data.smk +++ b/projects/data/data.smk @@ -1,11 +1,8 @@ """Snakemake rules for data acquisition and waveform generation. Rules: - generate_train_segments checkpoint: query DQSegDB for train segments - generate_test_segments checkpoint: query DQSegDB for test segments - fetch_train_background download one chunk of train strain data - fetch_test_background download one chunk of test strain data - compute_waveform_branches checkpoint: enumerate (segment, shifts) combos + generate_segments checkpoint: query DQSegDB for a split's segments + fetch_background download one chunk of a split's strain data compute_psd PSDs that waveform jobs rejection-sample against val_waveforms_branch validation waveforms for one branch aggregate_val_waveforms merge per-branch validation waveforms @@ -21,14 +18,14 @@ Directory layout: {waveforms_dir}/train/{val_waveforms,training_waveforms,psd}.hdf5 {waveforms_dir}/test/{waveforms,rejected_parameters,psd}.hdf5 -The fetch rules wildcard over {start} and {duration}. -Testing waveforms wildcard over {wbranch_id}, which are indices into the -compute_waveform_branches output. +Segments, fetching and PSDs wildcard over {split}, "train" or "test", and +fetching over each file's {start} and {duration}. +Testing waveforms wildcard over a test file's {start} and {duration} and +a {timeslide}, from testing_waveform_branches. Validation waveforms over {vbranch_id}, a fixed number of jobs that split num_validation_signals between them. """ -import json import math from pathlib import Path @@ -50,17 +47,16 @@ training_branch_ids = [str(i) for i in range(num_train_waveform_jobs)] localrules: - compute_waveform_branches, compute_psd, # DQSegDB seems to be unreachable from execute points - generate_train_segments, - generate_test_segments, + generate_segments, wildcard_constraints: + split="train|test", start=r"\d{10}", duration=r"\d+", - wbranch_id=r"\d+", + timeslide=r"\d+", vbranch_id=r"\d+", tbranch_id=r"\d+", @@ -75,6 +71,16 @@ def _is_analyzeable_segment(start, stop, shifts, psd_length): return (stop - start) - max(shifts) - psd_length > 0 +def timeslide_shifts(timeslide): + """Each detector's time shift, in seconds, at a timeslide. + + `shifts` in the config is the step each detector moves per timeslide, + so timeslide n shifts each detector by n of its steps. + Timeslide 0 is zero lag. + """ + return [int(timeslide) * step for step in config["shifts"]] + + def _read_segments(segments_file): """Parse a segments file into (start, duration) pairs. @@ -107,20 +113,29 @@ def _segment_chunks(segments_file): return chunks -def get_train_background_files(wildcards): - """All training background file paths.""" - seg_file = checkpoints.generate_train_segments.get(**wildcards).output[0] +def background_files(split): + """All background file paths of a split, "train" or "test".""" + seg_file = checkpoints.generate_segments.get(split=split).output[0] return [ - str(train_bg / f"background-{s}-{d}.hdf5") for s, d in _segment_chunks(seg_file) + str(bg_dir / split / f"background-{s}-{d}.hdf5") + for s, d in _segment_chunks(seg_file) ] -def get_test_background_files(wildcards): - """All testing background file paths.""" - seg_file = checkpoints.generate_test_segments.get(**wildcards).output[0] - return [ - str(test_bg / f"background-{s}-{d}.hdf5") for s, d in _segment_chunks(seg_file) - ] +def train_background_files(wildcards): + """Every training background file as an input for the training rules. + + The list of files isn't known until generate_segments has queried the + training segments, which happens partway through a run. Because a rule + can't list all the files directly, it names this function as its input, + list these files directly. Instead it names this function as its input, + `background=train_background_files`, without calling it, and snakemake + calls it later. + + Snakemake passes every such function the rule's wildcards. They aren't + used here, but the argument is required. + """ + return background_files("train") # Shortest background file to compute waveform PSDs from @@ -133,11 +148,7 @@ def _psd_background_file(wildcards): rejection-sampled against. We take the last file of the split that is at least PSD_MIN_DURATION long. """ - if wildcards.split == "train": - files = get_train_background_files(wildcards) - else: - files = get_test_background_files(wildcards) - for fname in reversed(files): + for fname in reversed(background_files(wildcards.split)): if int(Path(fname).stem.split("-")[-1]) >= PSD_MIN_DURATION: return fname raise WorkflowError( @@ -146,71 +157,81 @@ def _psd_background_file(wildcards): ) -def _branch_params(wildcards, input): - """Read one branch's params from the branch map file.""" - with open(input.branch_map) as f: - branch = json.load(f)[wildcards.wbranch_id] - return { - "start": branch["start"], - "end": branch["end"], - "shifts": _fmt_list(branch["shifts"]), - } +def testing_waveform_branches(): + """The (start, duration, timeslide) of each testing waveform branch. + A branch is one test background file, identified by its start and + duration, analyzed at one timeslide (see timeslide_shifts). These are + the same (file, timeslide) pairs that inference runs on, so every + injection falls inside a single inference branch. -def get_waveform_branch_files(wildcards): - """Per-branch testing waveform outputs.""" - bmap_file = checkpoints.compute_waveform_branches.get(**wildcards).output[0] - with open(bmap_file) as f: - branch_map = json.load(f) - waveforms = expand( - str(test_waveforms / "branches" / "{wbranch_id}" / "waveforms.hdf5"), - wbranch_id=branch_map.keys(), - ) - rejected = expand( - str(test_waveforms / "branches" / "{wbranch_id}" / "rejected_parameters.hdf5"), - wbranch_id=branch_map.keys(), + Branches are added until they have room for num_testing_signals injections. + + Branches are named by file and timeslide rather than numbered, so their + files stay valid in a waveforms_dir shared between runs. + """ + seg_file = checkpoints.generate_segments.get(split="test").output[0] + chunks = _segment_chunks(seg_file) + psd_length = config["psd_length"] + target = config["num_testing_signals"] + edge = config["buffer"] + config["waveform_duration"] // 2 + stride = config["spacing"] + config["waveform_duration"] + branches, total, timeslide = [], 0, 0 + while total < target: + timeslide += 1 + shifts = timeslide_shifts(timeslide) + added = False + for start, duration in chunks: + if total >= target: + break + if not _is_analyzeable_segment(start, start + duration, shifts, psd_length): + continue + span = duration - psd_length - max(shifts) + slots = math.ceil((span - 2 * edge) / stride) + if slots <= 0: + continue + branches.append((start, duration, timeslide)) + total += slots + added = True + # the analyzed part of each file shrinks as the shifts grow, + # stop if nothing is getting added + if not added: + break + return branches + + +testing_branch_dir = test_waveforms / "branches" / "{start}-{duration}-{timeslide}" + + +def testing_waveform_file(start, duration, timeslide, name="waveforms"): + """One testing waveform branch's output file.""" + return str(testing_branch_dir / f"{name}.hdf5").format( + start=start, duration=duration, timeslide=timeslide ) - return {"waveforms": waveforms, "rejected": rejected} -checkpoint generate_train_segments: - """Query DQSegDB for valid training data segments.""" - output: - str(train_bg / "segments.txt"), - log: - str(data_log_dir / "generate_train_segments.log"), - container: - DATA_CONTAINER - params: - flags=_fmt_list(config["flags"]), - start=config["train_start"], - end=config["train_end"], - min_duration=config["train_min_duration"], - segment_server=config["segment_server"], - shell: - "generate-segments" - " --flags '{params.flags}'" - " --start {params.start}" - " --end {params.end}" - " --min_duration {params.min_duration}" - " --segment_server {params.segment_server}" - " --output_file {output}" - " &> {log}" +def get_testing_waveform_files(wildcards): + """Every testing waveform branch's outputs, keyed by output name.""" + branches = testing_waveform_branches() + return { + name: [testing_waveform_file(*b, name) for b in branches] + for name in ["waveforms", "rejected_parameters"] + } -checkpoint generate_test_segments: - """Query DQSegDB for valid test data segments.""" +checkpoint generate_segments: + """Query DQSegDB for a split's valid data segments.""" output: - str(test_bg / "segments.txt"), + str(bg_dir / "{split}" / "segments.txt"), log: - str(data_log_dir / "generate_test_segments.log"), + str(data_log_dir / "generate_segments-{split}.log"), container: DATA_CONTAINER params: flags=_fmt_list(config["flags"]), - start=config["test_start"], - end=config["test_end"], - min_duration=config["test_min_duration"], + start=lambda wc: config[f"{wc.split}_start"], + end=lambda wc: config[f"{wc.split}_end"], + min_duration=lambda wc: config[f"{wc.split}_min_duration"], segment_server=config["segment_server"], shell: "generate-segments" @@ -223,48 +244,20 @@ checkpoint generate_test_segments: " &> {log}" -rule fetch_train_background: - """Download one chunk of training strain data.""" - input: - str(train_bg / "segments.txt"), - output: - str(train_bg / "background-{start}-{duration}.hdf5"), - log: - str(data_log_dir / "fetch_train_background-{start}-{duration}.log"), - container: - DATA_CONTAINER - # `fetch` downloads with nproc=3 - threads: 4 - resources: - **rule_resources("fetch_train_background"), - params: - channels=_fmt_list(config["channels"]), - sample_rate=config["sample_rate"], - end=lambda wc: int(wc.start) + int(wc.duration), - shell: - "fetch-data" - " --start {wildcards.start}" - " --end {params.end}" - " --channels '{params.channels}'" - " --sample_rate {params.sample_rate}" - " --output_file {output}" - " &> {log}" - - -rule fetch_test_background: - """Download one chunk of test strain data.""" +rule fetch_background: + """Download one chunk of a split's strain data.""" input: - str(test_bg / "segments.txt"), + str(bg_dir / "{split}" / "segments.txt"), output: - str(test_bg / "background-{start}-{duration}.hdf5"), + str(bg_dir / "{split}" / "background-{start}-{duration}.hdf5"), log: - str(data_log_dir / "fetch_test_background-{start}-{duration}.log"), + str(data_log_dir / "fetch_background-{split}-{start}-{duration}.log"), container: DATA_CONTAINER # `fetch` downloads with nproc=3 threads: 4 resources: - **rule_resources("fetch_test_background"), + **rule_resources("fetch_background"), params: channels=_fmt_list(config["channels"]), sample_rate=config["sample_rate"], @@ -279,70 +272,6 @@ rule fetch_test_background: " &> {log}" -checkpoint compute_waveform_branches: - """Create a file of (start, end, shifts) branches for testing waveforms. - -Each branch covers the analyzed part of one test background file at one -timeslide: after the PSD burn-in at the file's start, and before the -timeslide loss (max shift) at its end. Every injection, including its -full waveform, falls inside a single inference branch. - -Adds branches until enough data is present to generate as many waveforms -as requested. Loops over background files, adding an additional timeslide -if the target has not yet been met. - -Runs locally on the submit node. -""" - input: - lambda wildcards: checkpoints.generate_test_segments.get(**wildcards).output[0], - output: - str(test_waveforms / "waveform_branch_map.json"), - run: - files = [ - (start, start + duration) - for start, duration in _segment_chunks(input[0]) - ] - shifts = config["shifts"] - psd_length = config["psd_length"] - target = config["num_testing_signals"] - edge = config["buffer"] + config["waveform_duration"] // 2 - stride = config["spacing"] + config["waveform_duration"] - branch_map, branch_id, total = {}, 0, 0 - i = 0 - while total < target: - i += 1 - shift = [i * s for s in shifts] - added = False - for start, end in files: - if total >= target: - break - if not _is_analyzeable_segment(start, end, shift, psd_length): - continue - avail_start = start + psd_length - avail_end = end - max(shift) - slots = math.ceil((avail_end - avail_start - 2 * edge) / stride) - if slots <= 0: - continue - branch_map[str(branch_id)] = { - "background": str( - test_bg / f"background-{start}-{end - start}.hdf5" - ), - "start": avail_start, - "end": avail_end, - "shifts": shift, - } - branch_id += 1 - total += slots - added = True - # files shrink as the shift grows, - # stop if nothing is getting added - if not added: - break - Path(output[0]).parent.mkdir(parents=True, exist_ok=True) - with open(output[0], "w") as f: - json.dump(branch_map, f, indent=2) - - rule compute_psd: """Compute the PSDs that a split's waveform jobs rejection-sample against so that each job reads a small file instead of a full @@ -354,8 +283,6 @@ background file. str(waveform_dir / "{split}" / "psd.hdf5"), log: str(data_log_dir / "compute_psd-{split}.log"), - wildcard_constraints: - split="train|test", container: DATA_CONTAINER params: @@ -366,28 +293,35 @@ background file. rule testing_waveforms_branch: - """Generate testing waveforms for one (segment, shifts) branch. + """Generate testing waveforms for one branch: the analyzed part of +the test background file starting at {start}, lasting {duration}, at +timeslide {timeslide}. The PSD burn-in at the file's start and +the timeslide loss (max shift) at its end are accounted for. Rejection-samples waveforms against the PSDs of the last fetched test-background chunk and writes the accepted injection set and the rejected parameters for this branch. """ input: - branch_map=str(test_waveforms / "waveform_branch_map.json"), psd_file=str(test_waveforms / "psd.hdf5"), output: - waveforms=str(test_waveforms / "branches" / "{wbranch_id}" / "waveforms.hdf5"), - rejected=str( - test_waveforms / "branches" / "{wbranch_id}" / "rejected_parameters.hdf5" - ), + waveforms=str(testing_branch_dir / "waveforms.hdf5"), + rejected=str(testing_branch_dir / "rejected_parameters.hdf5"), log: - str(data_log_dir / "testing_waveforms_branch-{wbranch_id}.log"), + str( + data_log_dir + / "testing_waveforms_branch-{start}-{duration}-{timeslide}.log" + ), container: DATA_CONTAINER resources: **rule_resources("testing_waveforms_branch"), params: - branch=_branch_params, + start=lambda wc: int(wc.start) + config["psd_length"], + end=lambda wc: ( + int(wc.start) + int(wc.duration) - max(timeslide_shifts(wc.timeslide)) + ), + shifts=lambda wc: _fmt_list(timeslide_shifts(wc.timeslide)), ifos=_fmt_list(config["ifos"]), prior=config["prior"], minimum_frequency=config["minimum_frequency"], @@ -405,10 +339,10 @@ the rejected parameters for this branch. seed=config["seed"], shell: "generate-testing-waveforms" - " --start {params.branch[start]}" - " --end {params.branch[end]}" + " --start {params.start}" + " --end {params.end}" " --ifos '{params.ifos}'" - " --shifts '{params.branch[shifts]}'" + " --shifts '{params.shifts}'" " --spacing {params.spacing}" " --buffer {params.buffer}" " --prior {params.prior}" @@ -432,10 +366,10 @@ the rejected parameters for this branch. rule aggregate_testing_waveforms: """Merge per-branch testing waveforms into the final injection set.""" input: - unpack(get_waveform_branch_files), + unpack(get_testing_waveform_files), output: waveforms=str(test_waveforms / "waveforms.hdf5"), - rejected=str(test_waveforms / "rejected_parameters.hdf5"), + rejected_parameters=str(test_waveforms / "rejected_parameters.hdf5"), log: str(data_log_dir / "aggregate_testing_waveforms.log"), localrule: config["aggregate_rules_local"] @@ -445,8 +379,9 @@ rule aggregate_testing_waveforms: **rule_resources("aggregate_testing_waveforms"), params: ifos=config["ifos"], + classes={"waveforms": "responses", "rejected_parameters": "parameters"}, script: - "scripts/aggregate_testing_waveforms.py" + "scripts/aggregate_waveforms.py" rule val_waveforms_branch: @@ -502,12 +437,12 @@ The PSDs are those of the last fetched train-background chunk. rule aggregate_val_waveforms: """Merge per-branch validation waveforms into val_waveforms.hdf5.""" input: - expand( + waveforms=expand( str(train_waveforms / "validation_tmp" / "waveforms-{vbranch_id}.hdf5"), vbranch_id=validation_branch_ids, ), output: - str(train_waveforms / "val_waveforms.hdf5"), + waveforms=str(train_waveforms / "val_waveforms.hdf5"), log: str(data_log_dir / "aggregate_val_waveforms.log"), localrule: config["aggregate_rules_local"] @@ -517,8 +452,9 @@ rule aggregate_val_waveforms: **rule_resources("aggregate_val_waveforms"), params: ifos=config["ifos"], + classes={"waveforms": "waveforms"}, script: - "scripts/aggregate_val_waveforms.py" + "scripts/aggregate_waveforms.py" if config["pregenerate_training_waveforms"]: @@ -563,12 +499,12 @@ if config["pregenerate_training_waveforms"]: rule aggregate_training_waveforms: """Merge per-branch training waveforms into training_waveforms.hdf5.""" input: - expand( + waveforms=expand( str(train_waveforms / "training_tmp" / "{tbranch_id}.hdf5"), tbranch_id=training_branch_ids, ), output: - str(train_waveforms / "training_waveforms.hdf5"), + waveforms=str(train_waveforms / "training_waveforms.hdf5"), log: str(data_log_dir / "aggregate_training_waveforms.log"), localrule: config["aggregate_rules_local"] @@ -576,5 +512,8 @@ if config["pregenerate_training_waveforms"]: DATA_CONTAINER resources: **rule_resources("aggregate_training_waveforms"), + params: + ifos=config["ifos"], + classes={"waveforms": "polarizations"}, script: - "scripts/aggregate_training_waveforms.py" + "scripts/aggregate_waveforms.py" diff --git a/projects/data/scripts/aggregate_testing_waveforms.py b/projects/data/scripts/aggregate_testing_waveforms.py deleted file mode 100644 index 08d410cbc..000000000 --- a/projects/data/scripts/aggregate_testing_waveforms.py +++ /dev/null @@ -1,38 +0,0 @@ -# ruff: noqa: F821 -"""Merge per-branch testing waveform files into the final injection set. - -Executed via the snakemake `script:` directive. -The `snakemake` object is injected by snakemake. -""" - -import sys -from pathlib import Path - -from ledger.injections import ( - InjectionParameterSet, - InterferometerResponseSet, - waveform_class_factory, -) - -# `script:` directives are not auto-redirected to the rule's log, so -# send stdout/stderr there to capture tracebacks on failure. -sys.stdout = sys.stderr = open(snakemake.log[0], "w", buffering=1) - -cls = waveform_class_factory( - snakemake.params.ifos, - InterferometerResponseSet, - "ResponseSet", -) - -# Keep the per-branch waveform files so that they can be -# used as infer inputs. -cls.aggregate( - [Path(i) for i in snakemake.input.waveforms], - snakemake.output.waveforms, - clean=False, -) -InjectionParameterSet.aggregate( - [Path(i) for i in snakemake.input.rejected], - snakemake.output.rejected, - clean=True, -) diff --git a/projects/data/scripts/aggregate_training_waveforms.py b/projects/data/scripts/aggregate_training_waveforms.py deleted file mode 100644 index 84421152e..000000000 --- a/projects/data/scripts/aggregate_training_waveforms.py +++ /dev/null @@ -1,21 +0,0 @@ -# ruff: noqa: F821 -"""Merge per-branch training waveform files into training_waveforms.hdf5. - -Executed via the snakemake `script:` directive. -The `snakemake` object is injected by snakemake. -""" - -import sys -from pathlib import Path - -from data.cleanup import remove_empty_dirs -from ledger.injections import WaveformPolarizationSet - -sys.stdout = sys.stderr = open(snakemake.log[0], "w", buffering=1) - -WaveformPolarizationSet.aggregate( - [Path(i) for i in list(snakemake.input)], - snakemake.output[0], - clean=True, -) -remove_empty_dirs(snakemake.input, Path(snakemake.output[0]).parent) diff --git a/projects/data/scripts/aggregate_val_waveforms.py b/projects/data/scripts/aggregate_val_waveforms.py deleted file mode 100644 index 5d9d8b030..000000000 --- a/projects/data/scripts/aggregate_val_waveforms.py +++ /dev/null @@ -1,29 +0,0 @@ -# ruff: noqa: F821 -"""Merge per-branch validation waveform files into val_waveforms.hdf5. - -Executed via the snakemake `script:` directive. -The `snakemake` object is injected by snakemake. -""" - -import sys -from pathlib import Path - -from data.cleanup import remove_empty_dirs -from ledger.injections import WaveformSet, waveform_class_factory - -# `script:` directives are not auto-redirected to the rule's log, so -# send stdout/stderr there to capture tracebacks on failure. -sys.stdout = sys.stderr = open(snakemake.log[0], "w", buffering=1) - -cls = waveform_class_factory( - snakemake.params.ifos, - WaveformSet, - "WaveformSet", -) - -cls.aggregate( - [Path(i) for i in list(snakemake.input)], - snakemake.output[0], - clean=True, -) -remove_empty_dirs(snakemake.input, Path(snakemake.output[0]).parent) diff --git a/projects/data/scripts/aggregate_waveforms.py b/projects/data/scripts/aggregate_waveforms.py new file mode 100644 index 000000000..55641d541 --- /dev/null +++ b/projects/data/scripts/aggregate_waveforms.py @@ -0,0 +1,48 @@ +# ruff: noqa: F821 +"""Merge per-branch waveform files into one file per named output. + +Each output is merged from the rule's input of the same name. +Per-branch files are deleted once merged along with their empty directories, +except for responses, which are also used for inference. + +Executed via the snakemake `script:` directive. +The `snakemake` object is injected by snakemake. +""" + +import sys +from pathlib import Path + +from data.cleanup import remove_empty_dirs +from ledger.injections import ( + InjectionParameterSet, + InterferometerResponseSet, + WaveformPolarizationSet, + WaveformSet, + waveform_class_factory, +) + +# `script:` directives are not auto-redirected to the rule's log, so +# send stdout/stderr there to capture tracebacks on failure. +sys.stdout = sys.stderr = open(snakemake.log[0], "w", buffering=1) + +ifos = snakemake.params.ifos +CLASSES = { + "responses": waveform_class_factory( + ifos, InterferometerResponseSet, "ResponseSet" + ), + "waveforms": waveform_class_factory(ifos, WaveformSet, "WaveformSet"), + "polarizations": WaveformPolarizationSet, + "parameters": InjectionParameterSet, +} + +for name, output in snakemake.output.items(): + kind = snakemake.params.classes[name] + cls = CLASSES[kind] + files = snakemake.input[name] + # snakemake gives a named input with one file as a string + if isinstance(files, str): + files = [files] + files = [Path(f) for f in files] + # Don't clean up responses because inference reads the individual branches + cls.aggregate(files, output, clean=kind != "responses") + remove_empty_dirs(files, Path(output).parent) diff --git a/projects/infer/infer.smk b/projects/infer/infer.smk index 84407e32b..b68a8ba60 100644 --- a/projects/infer/infer.smk +++ b/projects/infer/infer.smk @@ -166,10 +166,11 @@ else: includes zero-lag branches. Each branch records the waveform file of the testing waveform branch with the same file and shifts, which holds all of its injections, or null if there is none (e.g. zero-lag). + + Needs only the test segments, so it doesn't wait for fetching. """ input: - background=get_test_background_files, - waveform_branch_map=str(test_waveforms / "waveform_branch_map.json"), + str(bg_dir / "test" / "segments.txt"), output: str(infer_dir / "branch_map.json"), run: diff --git a/projects/train/train.smk b/projects/train/train.smk index 8a3a96266..a6338c093 100644 --- a/projects/train/train.smk +++ b/projects/train/train.smk @@ -78,7 +78,7 @@ if config["remote_train"]: NOTE: not yet functional """ input: - background=get_train_background_files, + background=train_background_files, val_waveforms=str(train_waveforms / "val_waveforms.hdf5"), train_waveforms=_train_waveform_inputs, output: @@ -107,7 +107,7 @@ else: overrides trainer.devices in the train config. """ input: - background=get_train_background_files, + background=train_background_files, val_waveforms=str(train_waveforms / "val_waveforms.hdf5"), train_waveforms=_train_waveform_inputs, output: From 2eed18410cd168af00cc678f3484868a8ddde18a Mon Sep 17 00:00:00 2001 From: William Benoit Date: Wed, 30 Sep 2026 10:17:32 -0400 Subject: [PATCH 2/3] Add left-out infer.smk changes --- projects/infer/infer.smk | 101 +++++++++++++++++---------------------- 1 file changed, 45 insertions(+), 56 deletions(-) diff --git a/projects/infer/infer.smk b/projects/infer/infer.smk index b68a8ba60..0d5ced422 100644 --- a/projects/infer/infer.smk +++ b/projects/infer/infer.smk @@ -158,13 +158,13 @@ if ANALYSIS_TYPE == "rnp": else: checkpoint compute_branch_map: - """Enumerate (background file, shifts) inference branches. + """Enumerate (background file, timeslide) inference branches. - The number of shift multiples is the minimum needed to accumulate + The number of timeslides is the minimum needed to accumulate Tb seconds of background livetime. Branches that are too short to analyze after shifting and PSD burn-in are dropped. Optionally includes zero-lag branches. Each branch records the waveform file of - the testing waveform branch with the same file and shifts, which + the testing waveform branch with the same file and timeslide, which holds all of its injections, or null if there is none (e.g. zero-lag). Needs only the test segments, so it doesn't wait for fetching. @@ -174,72 +174,61 @@ else: output: str(infer_dir / "branch_map.json"), run: - def _get_num_shifts(segments, Tb, shift, psd_length): + def _get_num_timeslides(segments, Tb, shift_step, psd_length): if Tb == 0: return 0 - livetime, num_shifts = 0, 0 + livetime, num_timeslides = 0, 0 durations = [stop - start - psd_length for start, stop in segments] while livetime < Tb: - num_shifts += 1 + num_timeslides += 1 for dur in durations: - dur -= shift * num_shifts + dur -= shift_step * num_timeslides if dur > 0: livetime += dur - return num_shifts + return num_timeslides - shifts = config["shifts"] psd_length = config["psd_length"] - segments = [] - for fname in input.background: - start, duration = map(float, Path(fname).stem.split("-")[-2:]) - segments.append((start, start + duration)) - num_shifts = _get_num_shifts( - segments, config["Tb"], max(shifts), psd_length + chunks = _segment_chunks(input[0]) + num_timeslides = _get_num_timeslides( + [(start, start + duration) for start, duration in chunks], + config["Tb"], + max(config["shifts"]), + psd_length, ) - with open(input.waveform_branch_map) as f: - wbmap = json.load(f) - wbranch_ids = { - (w["background"], tuple(w["shifts"])): i for i, w in wbmap.items() - } - max_waveform_shift = max(max(b["shifts"]) for b in wbmap.values()) - num_waveform_shifts = max_waveform_shift / max(shifts) - if num_waveform_shifts > num_shifts: + waveform_branches = set(testing_waveform_branches()) + num_waveform_timeslides = max( + (t for _, _, t in waveform_branches), default=0 + ) + if num_waveform_timeslides > num_timeslides: raise WorkflowError( - f"num_testing_signals requires {num_waveform_shifts} shift " - f"multiples but Tb={config['Tb']} only covers {num_shifts}. " + f"num_testing_signals requires {num_waveform_timeslides} " + f"timeslides but Tb={config['Tb']} only covers " + f"{num_timeslides}. " f"Reduce num_testing_signals or increase Tb." ) - branch_shifts = [ - [(j + 1) * s for s in shifts] for j in range(num_shifts) - ] - if zero_lag: - branch_shifts.insert(0, [0] * len(shifts)) - branch_map, i, matched = {}, 0, set() - for fname, (start, stop) in zip(input.background, segments): - for shift in branch_shifts: - if _is_analyzeable_segment(start, stop, shift, psd_length): - branch_map[str(i)] = { - "fname": str(fname), - "shifts": shift, - "waveforms": None, - } - wbranch_id = wbranch_ids.get((str(fname), tuple(shift))) - if wbranch_id is not None: - branch_map[str(i)]["waveforms"] = str( - test_waveforms - / "branches" - / wbranch_id - / "waveforms.hdf5" - ) - matched.add(wbranch_id) - i += 1 - # every testing waveform must be analyzed by some branch - unmatched = set(wbmap) - matched - if unmatched: - raise WorkflowError( - f"Testing waveform branches {sorted(unmatched, key=int)} " - "match no inference branch" - ) + # zero lag is timeslide 0, which has no testing waveforms + timeslides = range(0 if zero_lag else 1, num_timeslides + 1) + branch_map, i = {}, 0 + for start, duration in chunks: + for timeslide in timeslides: + shifts = timeslide_shifts(timeslide) + if not _is_analyzeable_segment( + start, start + duration, shifts, psd_length + ): + continue + branch = (start, duration, timeslide) + branch_map[str(i)] = { + "fname": str( + test_bg / f"background-{start}-{duration}.hdf5" + ), + "shifts": shifts, + "waveforms": ( + testing_waveform_file(*branch) + if branch in waveform_branches + else None + ), + } + i += 1 _assign_groups(branch_map, split_at_file_changes=True) Path(output[0]).parent.mkdir(parents=True, exist_ok=True) with open(output[0], "w") as f: From 04378229d708be89c73434376e23733c9a89e3f0 Mon Sep 17 00:00:00 2001 From: William Benoit Date: Wed, 30 Sep 2026 10:18:57 -0400 Subject: [PATCH 3/3] Link preprocessor args to export's and pass lowpass --- projects/export/export.smk | 9 +++------ projects/export/export.yaml | 5 ----- projects/export/export/cli.py | 13 +++++++++++++ projects/export/export_bns.yaml | 5 ----- 4 files changed, 16 insertions(+), 16 deletions(-) diff --git a/projects/export/export.smk b/projects/export/export.smk index e2487f10d..b705ae984 100644 --- a/projects/export/export.smk +++ b/projects/export/export.smk @@ -49,6 +49,7 @@ snakemake is invoked. fftlength=config["fftlength"] or "null", psd_length=config["psd_length"], highpass=config["highpass"], + lowpass=config["lowpass"] or "null", streams_per_gpu=config["streams_per_gpu"], gpu_env=gpu_env(config["inference_gpus"], 1), weights=( @@ -75,12 +76,8 @@ snakemake is invoked. " --fduration {params.fduration}" " --psd_length {params.psd_length}" " --streams_per_gpu {params.streams_per_gpu}" - " --preprocessor.init_args.kernel_length {params.kernel_length}" - " --preprocessor.init_args.sample_rate {params.sample_rate}" - " --preprocessor.init_args.inference_sampling_rate" - " {params.inference_sampling_rate}" - " --preprocessor.init_args.batch_size {params.batch_size}" - " --preprocessor.init_args.fduration {params.fduration}" + # the rest of the preprocessor's arguments are linked from the above " --preprocessor.init_args.fftlength {params.fftlength}" " --preprocessor.init_args.highpass {params.highpass}" + " --preprocessor.init_args.lowpass {params.lowpass}" " &> {log}" diff --git a/projects/export/export.yaml b/projects/export/export.yaml index 0915d8733..43035a385 100644 --- a/projects/export/export.yaml +++ b/projects/export/export.yaml @@ -20,11 +20,6 @@ psd_length: 64 preprocessor: class_path: utils.preprocessing.BatchWhitener init_args: - kernel_length: 1.5 - sample_rate: 2048 - inference_sampling_rate: 4 - batch_size: 128 - fduration: 1 fftlength: 2 highpass: 32 streams_per_gpu: 6 diff --git a/projects/export/export/cli.py b/projects/export/export/cli.py index 8300ef669..431dd037f 100644 --- a/projects/export/export/cli.py +++ b/projects/export/export/cli.py @@ -5,12 +5,25 @@ from export.main import export +# Arguments that the preprocessor shares with the model export +PREPROCESSOR_ARGUMENTS = [ + "kernel_length", + "sample_rate", + "inference_sampling_rate", + "batch_size", + "fduration", +] + def build_parser(): parser = jsonargparse.ArgumentParser() parser.add_argument("--config", action=jsonargparse.ActionConfigFile) parser.add_argument("--logfile", type=str, default=None) parser.add_function_arguments(export) + for arg in PREPROCESSOR_ARGUMENTS: + parser.link_arguments( + arg, f"preprocessor.init_args.{arg}", apply_on="parse" + ) return parser diff --git a/projects/export/export_bns.yaml b/projects/export/export_bns.yaml index 32d63b140..dd6a41f8b 100644 --- a/projects/export/export_bns.yaml +++ b/projects/export/export_bns.yaml @@ -18,11 +18,6 @@ psd_length: 64 preprocessor: class_path: utils.preprocessing.TimeSpectrogramPreprocessor init_args: - kernel_length: 20.0 - sample_rate: 2048 - inference_sampling_rate: 16 - batch_size: 128 - fduration: 2 fftlength: 2 schedule: [[0, 16, 512], [16, 20, 2048]] split: True