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/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 diff --git a/projects/infer/infer.smk b/projects/infer/infer.smk index 84407e32b..0d5ced422 100644 --- a/projects/infer/infer.smk +++ b/projects/infer/infer.smk @@ -158,87 +158,77 @@ 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. """ 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: - 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: 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: