pytesprocess 0.1.1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- pytesprocess/__init__.py +9 -0
- pytesprocess/_version.py +2 -0
- pytesprocess/cli/__init__.py +1 -0
- pytesprocess/cli/commands/__init__.py +5 -0
- pytesprocess/cli/commands/event.py +66 -0
- pytesprocess/cli/commands/filter.py +17 -0
- pytesprocess/cli/commands/ivsweep.py +29 -0
- pytesprocess/cli/common.py +86 -0
- pytesprocess/cli/main.py +81 -0
- pytesprocess/config/__init__.py +4 -0
- pytesprocess/config/loader.py +94 -0
- pytesprocess/config/manager.py +297 -0
- pytesprocess/config/resolvers/__init__.py +5 -0
- pytesprocess/config/resolvers/common.py +56 -0
- pytesprocess/config/resolvers/feature.py +293 -0
- pytesprocess/config/resolvers/salting.py +86 -0
- pytesprocess/config/resolvers/trigger.py +84 -0
- pytesprocess/config/selectors.py +108 -0
- pytesprocess/config/validation.py +314 -0
- pytesprocess/config/warnings.py +2 -0
- pytesprocess/core/__init__.py +10 -0
- pytesprocess/core/algorithms.py +1455 -0
- pytesprocess/core/didv.py +1648 -0
- pytesprocess/core/eventbuilder.py +495 -0
- pytesprocess/core/filterbuilder.py +81 -0
- pytesprocess/core/filterdata.py +1849 -0
- pytesprocess/core/ivsweep.py +2072 -0
- pytesprocess/core/noise.py +923 -0
- pytesprocess/core/noisemodel.py +1408 -0
- pytesprocess/core/oftrigger.py +1035 -0
- pytesprocess/core/template.py +450 -0
- pytesprocess/process/__init__.py +6 -0
- pytesprocess/process/data_source.py +185 -0
- pytesprocess/process/event_context.py +35 -0
- pytesprocess/process/feature_plan.py +186 -0
- pytesprocess/process/feature_resources.py +267 -0
- pytesprocess/process/features.py +1024 -0
- pytesprocess/process/filterprocess.py +1176 -0
- pytesprocess/process/ivprocess.py +1380 -0
- pytesprocess/process/processing_data.py +967 -0
- pytesprocess/process/randoms.py +921 -0
- pytesprocess/process/triggers.py +1011 -0
- pytesprocess/salting/__init__.py +7 -0
- pytesprocess/salting/generator.py +364 -0
- pytesprocess/salting/injector.py +329 -0
- pytesprocess/salting/sampling.py +84 -0
- pytesprocess/utils/__init__.py +5 -0
- pytesprocess/utils/arg_utils.py +122 -0
- pytesprocess/utils/dataframe_output.py +120 -0
- pytesprocess/utils/filter_hdf5.py +594 -0
- pytesprocess/utils/utils.py +701 -0
- pytesprocess/workflows/__init__.py +3 -0
- pytesprocess/workflows/processing.py +317 -0
- pytesprocess/workflows/salting.py +133 -0
- pytesprocess-0.1.1.dist-info/METADATA +211 -0
- pytesprocess-0.1.1.dist-info/RECORD +60 -0
- pytesprocess-0.1.1.dist-info/WHEEL +5 -0
- pytesprocess-0.1.1.dist-info/entry_points.txt +2 -0
- pytesprocess-0.1.1.dist-info/licenses/LICENSE +21 -0
- pytesprocess-0.1.1.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,7 @@
|
|
|
1
|
+
"""Salting generation and injection helpers."""
|
|
2
|
+
|
|
3
|
+
from .generator import SaltGenerator
|
|
4
|
+
from .injector import SaltInjector
|
|
5
|
+
from .sampling import sample_pdf, sample_dm_distributions
|
|
6
|
+
|
|
7
|
+
__all__ = ["SaltGenerator", "SaltInjector", "sample_pdf", "sample_dm_distributions"]
|
|
@@ -0,0 +1,364 @@
|
|
|
1
|
+
"""Salt-event metadata generation.
|
|
2
|
+
|
|
3
|
+
Generation selects raw-data coordinates and records enough metadata to
|
|
4
|
+
reconstruct each injected pulse later. It intentionally does not require raw
|
|
5
|
+
waveform injection and therefore remains independent of trigger/feature reads.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
import numpy as np
|
|
12
|
+
import qetpy as qp
|
|
13
|
+
|
|
14
|
+
from pytesdaqx.io import AcquisitionCatalog
|
|
15
|
+
from qetpy.utils import convert_channel_name_to_list, convert_channel_list_to_name
|
|
16
|
+
|
|
17
|
+
from pytesprocess.utils import extract_stream_id
|
|
18
|
+
from .sampling import sample_dm_distributions
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class SaltGenerator:
|
|
22
|
+
"""Generate salt metadata from an acquisition and FilterData object."""
|
|
23
|
+
|
|
24
|
+
def __init__(self, filter_data, verbose=True):
|
|
25
|
+
self._filter_data = filter_data
|
|
26
|
+
self._verbose = bool(verbose)
|
|
27
|
+
self._data_path = None
|
|
28
|
+
self._streams = None
|
|
29
|
+
self._restricted = False
|
|
30
|
+
self._catalog = None
|
|
31
|
+
self._sample_rate_hz = None
|
|
32
|
+
self._positions = None
|
|
33
|
+
|
|
34
|
+
@property
|
|
35
|
+
def sample_rate_hz(self):
|
|
36
|
+
return self._sample_rate_hz
|
|
37
|
+
|
|
38
|
+
@property
|
|
39
|
+
def catalog(self):
|
|
40
|
+
return self._catalog
|
|
41
|
+
|
|
42
|
+
def set_raw_data(self, data_path, streams=None, restricted=False):
|
|
43
|
+
self._data_path = data_path
|
|
44
|
+
self._streams = streams
|
|
45
|
+
self._restricted = bool(restricted)
|
|
46
|
+
|
|
47
|
+
catalog_full = AcquisitionCatalog(data_path, verbose=self._verbose)
|
|
48
|
+
self._catalog = catalog_full.filter(
|
|
49
|
+
measurement_types="background",
|
|
50
|
+
streams=streams,
|
|
51
|
+
restricted=self._restricted,
|
|
52
|
+
)
|
|
53
|
+
selected = self._catalog.select_files_by_stream(stream_key="stream_id")
|
|
54
|
+
if not selected:
|
|
55
|
+
raise ValueError("No background streams were found for salting.")
|
|
56
|
+
self._sample_rate_hz = self._catalog.sample_rate_hz
|
|
57
|
+
return self._catalog
|
|
58
|
+
|
|
59
|
+
def get_positions(self):
|
|
60
|
+
return self._positions
|
|
61
|
+
|
|
62
|
+
def set_positions(self, dataframe):
|
|
63
|
+
self._positions = self._ensure_stream_id(dataframe)
|
|
64
|
+
|
|
65
|
+
def clear_positions(self):
|
|
66
|
+
self._positions = None
|
|
67
|
+
|
|
68
|
+
def generate_positions(self, nrandoms, *, min_separation_msec=0,
|
|
69
|
+
edge_exclusion_msec=0, random_seed=None):
|
|
70
|
+
if self._data_path is None:
|
|
71
|
+
raise ValueError("Raw data must be set before generating salt positions.")
|
|
72
|
+
# Local import avoids a package-initialization cycle.
|
|
73
|
+
from pytesprocess.process.randoms import Randoms
|
|
74
|
+
randoms = Randoms(
|
|
75
|
+
self._data_path,
|
|
76
|
+
streams=self._streams,
|
|
77
|
+
data_type="background",
|
|
78
|
+
restricted=self._restricted,
|
|
79
|
+
verbose=False,
|
|
80
|
+
)
|
|
81
|
+
dataframe = randoms.process(
|
|
82
|
+
nrandoms=int(nrandoms),
|
|
83
|
+
min_separation_msec=float(min_separation_msec),
|
|
84
|
+
edge_exclusion_msec=float(edge_exclusion_msec),
|
|
85
|
+
partition_target_duration_s=None,
|
|
86
|
+
random_seed=random_seed,
|
|
87
|
+
lgc_save=False,
|
|
88
|
+
lgc_output=True,
|
|
89
|
+
)
|
|
90
|
+
self._positions = self._ensure_stream_id(dataframe)
|
|
91
|
+
if self._verbose:
|
|
92
|
+
print(f"INFO: {len(self._positions)} salting events randomly selected!")
|
|
93
|
+
return self._positions
|
|
94
|
+
|
|
95
|
+
def generate_metadata(
|
|
96
|
+
self,
|
|
97
|
+
channels,
|
|
98
|
+
*,
|
|
99
|
+
template_tag="default",
|
|
100
|
+
dpdi_tag=None,
|
|
101
|
+
dpdi_poles=None,
|
|
102
|
+
collection_efficiency=1.0,
|
|
103
|
+
energy_eV=None,
|
|
104
|
+
dm_pdf_file=None,
|
|
105
|
+
nsalt=100,
|
|
106
|
+
positions=None,
|
|
107
|
+
min_separation_msec=None,
|
|
108
|
+
edge_exclusion_msec=0,
|
|
109
|
+
random_seed=None,
|
|
110
|
+
salting_livetime_s=None,
|
|
111
|
+
):
|
|
112
|
+
"""Generate one salt metadata dataframe.
|
|
113
|
+
|
|
114
|
+
Exactly one of ``energy_eV`` or ``dm_pdf_file`` must be provided.
|
|
115
|
+
Fixed-energy generation is intentionally one energy at a time; an
|
|
116
|
+
outer workflow may call this repeatedly for an energy scan.
|
|
117
|
+
"""
|
|
118
|
+
if (energy_eV is None) == (dm_pdf_file is None):
|
|
119
|
+
raise ValueError(
|
|
120
|
+
'Specify exactly one of "energy_eV" or "dm_pdf_file".'
|
|
121
|
+
)
|
|
122
|
+
if (dpdi_tag is None) != (dpdi_poles is None):
|
|
123
|
+
raise ValueError(
|
|
124
|
+
'Both "dpdi_tag" and "dpdi_poles" must be set or both omitted.'
|
|
125
|
+
)
|
|
126
|
+
|
|
127
|
+
channel_list = convert_channel_name_to_list(channels)
|
|
128
|
+
channel_name = convert_channel_list_to_name(channel_list)
|
|
129
|
+
efficiencies = self._normalize_efficiency(
|
|
130
|
+
collection_efficiency, len(channel_list)
|
|
131
|
+
)
|
|
132
|
+
dpdi_tags = self._normalize_per_channel(
|
|
133
|
+
dpdi_tag, len(channel_list), label="dpdi_tag"
|
|
134
|
+
)
|
|
135
|
+
dpdi_poles_list = self._normalize_per_channel(
|
|
136
|
+
dpdi_poles, len(channel_list), label="dpdi_poles"
|
|
137
|
+
)
|
|
138
|
+
|
|
139
|
+
template, time_array, template_metadata = self._filter_data.get_template(
|
|
140
|
+
channel_name, tag=template_tag, return_metadata=True
|
|
141
|
+
)
|
|
142
|
+
template_length = int(np.asarray(template).shape[-1])
|
|
143
|
+
template_pretrigger = int(
|
|
144
|
+
template_metadata.get("nb_pretrigger_samples", template_length // 2)
|
|
145
|
+
)
|
|
146
|
+
|
|
147
|
+
if energy_eV is not None:
|
|
148
|
+
if isinstance(energy_eV, (list, tuple, np.ndarray)):
|
|
149
|
+
if len(energy_eV) != 1:
|
|
150
|
+
raise ValueError(
|
|
151
|
+
"generate_metadata accepts one fixed energy at a time."
|
|
152
|
+
)
|
|
153
|
+
energy_eV = energy_eV[0]
|
|
154
|
+
energies = np.full(int(nsalt), float(energy_eV), dtype=float)
|
|
155
|
+
masses = None
|
|
156
|
+
salting_type = np.asarray(
|
|
157
|
+
[f"energy_{float(energy_eV):g}_eV"] * len(energies)
|
|
158
|
+
)
|
|
159
|
+
else:
|
|
160
|
+
energies, masses = sample_dm_distributions(
|
|
161
|
+
dm_pdf_file, int(nsalt), random_seed=random_seed
|
|
162
|
+
)
|
|
163
|
+
salting_type = np.asarray(["dm_pdf"] * len(energies))
|
|
164
|
+
|
|
165
|
+
event_count = len(energies)
|
|
166
|
+
if positions is not None:
|
|
167
|
+
position_df = self._ensure_stream_id(positions)
|
|
168
|
+
self._positions = position_df
|
|
169
|
+
else:
|
|
170
|
+
if min_separation_msec is None:
|
|
171
|
+
min_separation_msec = 1000.0 * template_length / float(
|
|
172
|
+
self._sample_rate_hz
|
|
173
|
+
)
|
|
174
|
+
position_df = self.generate_positions(
|
|
175
|
+
event_count,
|
|
176
|
+
min_separation_msec=min_separation_msec,
|
|
177
|
+
edge_exclusion_msec=edge_exclusion_msec,
|
|
178
|
+
random_seed=random_seed,
|
|
179
|
+
)
|
|
180
|
+
|
|
181
|
+
if len(position_df) != event_count:
|
|
182
|
+
raise ValueError(
|
|
183
|
+
f"Salt position count ({len(position_df)}) does not match "
|
|
184
|
+
f"generated event count ({event_count})."
|
|
185
|
+
)
|
|
186
|
+
self._positions = position_df
|
|
187
|
+
|
|
188
|
+
scales, peaks, normalizations = self._calculate_template_scales(
|
|
189
|
+
channel_list=channel_list,
|
|
190
|
+
template=template,
|
|
191
|
+
time_array=np.asarray(time_array),
|
|
192
|
+
dpdi_tags=dpdi_tags,
|
|
193
|
+
dpdi_poles_list=dpdi_poles_list,
|
|
194
|
+
energies_eV=energies,
|
|
195
|
+
efficiencies=efficiencies,
|
|
196
|
+
)
|
|
197
|
+
|
|
198
|
+
salt_dict = {
|
|
199
|
+
"salt_template_tag": np.asarray([str(template_tag)] * event_count),
|
|
200
|
+
"salt_recoil_energy_eV": energies,
|
|
201
|
+
"saltchanname": np.asarray([channel_name] * event_count),
|
|
202
|
+
"salt_channel_expression": np.asarray([channel_name] * event_count),
|
|
203
|
+
"salting_type": salting_type,
|
|
204
|
+
"salt_template_length_samples": np.full(
|
|
205
|
+
event_count, template_length, dtype=np.int64
|
|
206
|
+
),
|
|
207
|
+
"salt_template_pretrigger_samples": np.full(
|
|
208
|
+
event_count, template_pretrigger, dtype=np.int64
|
|
209
|
+
),
|
|
210
|
+
}
|
|
211
|
+
if masses is not None:
|
|
212
|
+
salt_dict["salt_dm_mass_MeV"] = masses
|
|
213
|
+
if salting_livetime_s is None:
|
|
214
|
+
position_columns = set(position_df.get_column_names())
|
|
215
|
+
if "randoms_livetime_s" in position_columns and len(position_df):
|
|
216
|
+
values = position_df.evaluate(
|
|
217
|
+
"randoms_livetime_s", array_type="numpy"
|
|
218
|
+
)
|
|
219
|
+
if len(values):
|
|
220
|
+
salting_livetime_s = float(values[0])
|
|
221
|
+
if salting_livetime_s is not None:
|
|
222
|
+
salt_dict["salting_livetime"] = np.full(
|
|
223
|
+
event_count, float(salting_livetime_s)
|
|
224
|
+
)
|
|
225
|
+
if dpdi_tag is not None:
|
|
226
|
+
if len(set(str(value) for value in dpdi_tags)) == 1 and len(set(int(value) for value in dpdi_poles_list)) == 1:
|
|
227
|
+
salt_dict["salt_dpdi_tag"] = np.asarray(
|
|
228
|
+
[str(dpdi_tags[0])] * event_count
|
|
229
|
+
)
|
|
230
|
+
salt_dict["salt_dpdi_poles"] = np.full(
|
|
231
|
+
event_count, int(dpdi_poles_list[0]), dtype=np.int16
|
|
232
|
+
)
|
|
233
|
+
else:
|
|
234
|
+
for index, channel in enumerate(channel_list):
|
|
235
|
+
salt_dict[f"salt_dpdi_tag_{channel}"] = np.asarray(
|
|
236
|
+
[str(dpdi_tags[index])] * event_count
|
|
237
|
+
)
|
|
238
|
+
salt_dict[f"salt_dpdi_poles_{channel}"] = np.full(
|
|
239
|
+
event_count, int(dpdi_poles_list[index]), dtype=np.int16
|
|
240
|
+
)
|
|
241
|
+
|
|
242
|
+
for index, channel in enumerate(channel_list):
|
|
243
|
+
scale = scales[channel]
|
|
244
|
+
# ``salt_amplitude`` is retained as a compatibility alias. Its
|
|
245
|
+
# historical injector treated it as the multiplier applied to the
|
|
246
|
+
# stored template, which is now the explicit definition.
|
|
247
|
+
salt_dict[f"salt_template_scale_{channel}"] = scale
|
|
248
|
+
salt_dict[f"salt_amplitude_{channel}"] = scale
|
|
249
|
+
salt_dict[f"salt_peak_amplitude_amps_{channel}"] = peaks[channel]
|
|
250
|
+
# Keep historical meaning of salt_energy_eV_<chan> (recoil energy)
|
|
251
|
+
# and add the collected/channel energy explicitly.
|
|
252
|
+
salt_dict[f"salt_energy_eV_{channel}"] = energies
|
|
253
|
+
salt_dict[f"salt_collected_energy_eV_{channel}"] = (
|
|
254
|
+
energies * efficiencies[index]
|
|
255
|
+
)
|
|
256
|
+
salt_dict[f"salt_collection_efficiency_{channel}"] = np.full(
|
|
257
|
+
event_count, efficiencies[index], dtype=float
|
|
258
|
+
)
|
|
259
|
+
if normalizations[channel] is not None:
|
|
260
|
+
salt_dict[f"salt_energy_normalization_eV_per_amp_{channel}"] = (
|
|
261
|
+
np.full(event_count, normalizations[channel], dtype=float)
|
|
262
|
+
)
|
|
263
|
+
|
|
264
|
+
output = position_df.copy()
|
|
265
|
+
for key, values in salt_dict.items():
|
|
266
|
+
output[key] = values
|
|
267
|
+
return output
|
|
268
|
+
|
|
269
|
+
def _calculate_template_scales(
|
|
270
|
+
self, *, channel_list, template, time_array, dpdi_tags,
|
|
271
|
+
dpdi_poles_list, energies_eV, efficiencies,
|
|
272
|
+
):
|
|
273
|
+
scales = {}
|
|
274
|
+
peaks = {}
|
|
275
|
+
normalizations = {}
|
|
276
|
+
for index, channel in enumerate(channel_list):
|
|
277
|
+
temp = self._select_template_channel(template, index, len(channel_list))
|
|
278
|
+
normalization = None
|
|
279
|
+
if dpdi_tags is not None:
|
|
280
|
+
dpdi, _ = self._filter_data.get_dpdi(
|
|
281
|
+
channel,
|
|
282
|
+
poles=int(dpdi_poles_list[index]),
|
|
283
|
+
tag=dpdi_tags[index],
|
|
284
|
+
)
|
|
285
|
+
dpdi = np.asarray(dpdi)
|
|
286
|
+
if dpdi.ndim > 1:
|
|
287
|
+
dpdi = dpdi[0]
|
|
288
|
+
normalization = float(
|
|
289
|
+
qp.get_energy_normalization(
|
|
290
|
+
time_array, np.asarray(temp), dpdi=dpdi, lgc_ev=True
|
|
291
|
+
)
|
|
292
|
+
)
|
|
293
|
+
if not np.isfinite(normalization) or normalization == 0:
|
|
294
|
+
raise ValueError(
|
|
295
|
+
f"Invalid energy normalization for salting channel {channel}."
|
|
296
|
+
)
|
|
297
|
+
scale = energies_eV * efficiencies[index] / normalization
|
|
298
|
+
else:
|
|
299
|
+
# Without dPdI, retain the historical convention that the
|
|
300
|
+
# stored template is already normalized for multiplication by
|
|
301
|
+
# the requested energy-like amplitude.
|
|
302
|
+
scale = energies_eV * efficiencies[index]
|
|
303
|
+
|
|
304
|
+
scale = np.asarray(scale, dtype=float)
|
|
305
|
+
scales[channel] = scale
|
|
306
|
+
peaks[channel] = np.max(np.abs(np.asarray(temp))) * np.abs(scale)
|
|
307
|
+
normalizations[channel] = normalization
|
|
308
|
+
return scales, peaks, normalizations
|
|
309
|
+
|
|
310
|
+
@staticmethod
|
|
311
|
+
def _select_template_channel(template, index, nb_channels):
|
|
312
|
+
template = np.asarray(template)
|
|
313
|
+
if nb_channels == 1:
|
|
314
|
+
if template.ndim != 1:
|
|
315
|
+
return np.asarray(template).reshape(-1)
|
|
316
|
+
return template
|
|
317
|
+
# FilterData stores multi-channel templates as [nchan, ntemplate, nsamp].
|
|
318
|
+
selected = np.asarray(template[index])
|
|
319
|
+
if selected.ndim > 1:
|
|
320
|
+
selected = selected[0]
|
|
321
|
+
return selected
|
|
322
|
+
|
|
323
|
+
@staticmethod
|
|
324
|
+
def _normalize_efficiency(value, nb_channels):
|
|
325
|
+
if np.isscalar(value):
|
|
326
|
+
return np.full(nb_channels, float(value), dtype=float)
|
|
327
|
+
values = np.asarray(value, dtype=float).reshape(-1)
|
|
328
|
+
if len(values) != nb_channels:
|
|
329
|
+
raise ValueError(
|
|
330
|
+
"collection_efficiency must be scalar or have one value per channel."
|
|
331
|
+
)
|
|
332
|
+
return values
|
|
333
|
+
|
|
334
|
+
@staticmethod
|
|
335
|
+
def _normalize_per_channel(value, nb_channels, label):
|
|
336
|
+
if value is None:
|
|
337
|
+
return None
|
|
338
|
+
if isinstance(value, (list, tuple, np.ndarray)):
|
|
339
|
+
values = list(value)
|
|
340
|
+
if len(values) != nb_channels:
|
|
341
|
+
raise ValueError(
|
|
342
|
+
f"{label} must be scalar or have one value per channel."
|
|
343
|
+
)
|
|
344
|
+
return values
|
|
345
|
+
return [value] * nb_channels
|
|
346
|
+
|
|
347
|
+
@staticmethod
|
|
348
|
+
def _ensure_stream_id(dataframe):
|
|
349
|
+
if dataframe is None:
|
|
350
|
+
raise ValueError("Salt position dataframe is None.")
|
|
351
|
+
columns = set(dataframe.get_column_names())
|
|
352
|
+
if "stream_id" in columns:
|
|
353
|
+
return dataframe
|
|
354
|
+
if "stream_number" not in columns:
|
|
355
|
+
raise ValueError(
|
|
356
|
+
"Salt position dataframe needs stream_id or stream_number."
|
|
357
|
+
)
|
|
358
|
+
source = "stream_number"
|
|
359
|
+
numbers = dataframe.evaluate(source, array_type="numpy")
|
|
360
|
+
dataframe = dataframe.copy()
|
|
361
|
+
dataframe["stream_id"] = np.asarray(
|
|
362
|
+
[extract_stream_id(int(number)) for number in numbers]
|
|
363
|
+
)
|
|
364
|
+
return dataframe
|
|
@@ -0,0 +1,329 @@
|
|
|
1
|
+
"""Raw-waveform salt injection."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
import vaex as vx
|
|
7
|
+
|
|
8
|
+
from qetpy.utils import convert_channel_name_to_list
|
|
9
|
+
from pytesprocess.utils import extract_stream_id
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class SaltInjector:
|
|
13
|
+
"""Inject saved salt metadata into arbitrary raw-data windows."""
|
|
14
|
+
|
|
15
|
+
def __init__(self, filter_data, dataframe=None, verbose=True):
|
|
16
|
+
self._filter_data = filter_data
|
|
17
|
+
self._dataframe = None
|
|
18
|
+
self._verbose = bool(verbose)
|
|
19
|
+
self._template_cache = {}
|
|
20
|
+
if dataframe is not None:
|
|
21
|
+
self.set_dataframe(dataframe)
|
|
22
|
+
|
|
23
|
+
@property
|
|
24
|
+
def dataframe(self):
|
|
25
|
+
return self._dataframe
|
|
26
|
+
|
|
27
|
+
@property
|
|
28
|
+
def is_modern(self):
|
|
29
|
+
if self._dataframe is None:
|
|
30
|
+
return False
|
|
31
|
+
columns = set(self._dataframe.get_column_names())
|
|
32
|
+
return {"stream_id", "stream_trigger_index"}.issubset(columns)
|
|
33
|
+
|
|
34
|
+
def set_dataframe(self, dataframe):
|
|
35
|
+
if isinstance(dataframe, str):
|
|
36
|
+
dataframe = vx.open(dataframe)
|
|
37
|
+
if dataframe is None or not hasattr(dataframe, "get_column_names"):
|
|
38
|
+
raise ValueError("Unrecognized salting dataframe argument.")
|
|
39
|
+
if len(dataframe) < 1:
|
|
40
|
+
raise ValueError("No salting events found in dataframe.")
|
|
41
|
+
self._dataframe = dataframe
|
|
42
|
+
|
|
43
|
+
def inject_window(self, channels, trace, *, stream_id,
|
|
44
|
+
trace_start_index, include_metadata=False):
|
|
45
|
+
"""Inject every salt pulse overlapping one stream-global raw window."""
|
|
46
|
+
if self._dataframe is None:
|
|
47
|
+
return (trace, {}) if include_metadata else trace
|
|
48
|
+
if not self.is_modern:
|
|
49
|
+
raise ValueError(
|
|
50
|
+
"Window-based injection requires modern salting coordinates "
|
|
51
|
+
"(stream_id and stream_trigger_index)."
|
|
52
|
+
)
|
|
53
|
+
|
|
54
|
+
sid = extract_stream_id(stream_id)
|
|
55
|
+
filtered = self._dataframe[self._dataframe["stream_id"] == sid]
|
|
56
|
+
if len(filtered) == 0:
|
|
57
|
+
return (trace, {}) if include_metadata else trace
|
|
58
|
+
|
|
59
|
+
output, metadata = self._inject_rows(
|
|
60
|
+
channels, trace, filtered,
|
|
61
|
+
trigger_column="stream_trigger_index",
|
|
62
|
+
trace_start_index=int(trace_start_index),
|
|
63
|
+
stream_id=sid,
|
|
64
|
+
)
|
|
65
|
+
if include_metadata:
|
|
66
|
+
return output, metadata
|
|
67
|
+
return output
|
|
68
|
+
|
|
69
|
+
def inject_record(self, channels, trace, *, stream_id=None,
|
|
70
|
+
stream_number=None, global_segment_number=None,
|
|
71
|
+
include_metadata=False):
|
|
72
|
+
"""Inject salts into a full HDF5/native record.
|
|
73
|
+
|
|
74
|
+
HDF5 salt dataframes use ``global_segment_number`` and
|
|
75
|
+
``segment_trigger_index``.
|
|
76
|
+
"""
|
|
77
|
+
if self._dataframe is None:
|
|
78
|
+
return (trace, {}) if include_metadata else trace
|
|
79
|
+
columns = set(self._dataframe.get_column_names())
|
|
80
|
+
filtered = self._dataframe
|
|
81
|
+
|
|
82
|
+
required = {"global_segment_number", "segment_trigger_index"}
|
|
83
|
+
if not required.issubset(columns):
|
|
84
|
+
raise ValueError(
|
|
85
|
+
"Salting dataframe needs global_segment_number and "
|
|
86
|
+
"segment_trigger_index for HDF5 record injection."
|
|
87
|
+
)
|
|
88
|
+
if global_segment_number is None:
|
|
89
|
+
raise ValueError(
|
|
90
|
+
"global_segment_number is required for HDF5 salt injection."
|
|
91
|
+
)
|
|
92
|
+
|
|
93
|
+
filtered = filtered[
|
|
94
|
+
filtered["global_segment_number"] == int(global_segment_number)
|
|
95
|
+
]
|
|
96
|
+
if stream_id is not None and "stream_id" in columns:
|
|
97
|
+
filtered = filtered[
|
|
98
|
+
filtered["stream_id"] == extract_stream_id(stream_id)
|
|
99
|
+
]
|
|
100
|
+
elif stream_number is not None and "stream_number" in columns:
|
|
101
|
+
filtered = filtered[
|
|
102
|
+
filtered["stream_number"] == int(stream_number)
|
|
103
|
+
]
|
|
104
|
+
trigger_column = "segment_trigger_index"
|
|
105
|
+
sid = extract_stream_id(
|
|
106
|
+
stream_id if stream_id is not None else stream_number
|
|
107
|
+
)
|
|
108
|
+
|
|
109
|
+
if len(filtered) == 0:
|
|
110
|
+
return (trace, {}) if include_metadata else trace
|
|
111
|
+
output, metadata = self._inject_rows(
|
|
112
|
+
channels, trace, filtered,
|
|
113
|
+
trigger_column=trigger_column,
|
|
114
|
+
trace_start_index=0,
|
|
115
|
+
stream_id=sid,
|
|
116
|
+
)
|
|
117
|
+
metadata["global_segment_number"] = int(global_segment_number)
|
|
118
|
+
if include_metadata:
|
|
119
|
+
return output, metadata
|
|
120
|
+
return output
|
|
121
|
+
|
|
122
|
+
def _inject_rows(self, channels, trace, dataframe, *, trigger_column,
|
|
123
|
+
trace_start_index, stream_id=None):
|
|
124
|
+
channel_list = self._normalize_channels(channels)
|
|
125
|
+
trace_array = np.asarray(trace).copy()
|
|
126
|
+
input_was_1d = trace_array.ndim == 1
|
|
127
|
+
if input_was_1d:
|
|
128
|
+
trace_array = trace_array.reshape(1, -1)
|
|
129
|
+
if trace_array.ndim != 2 or len(channel_list) != trace_array.shape[0]:
|
|
130
|
+
raise ValueError(
|
|
131
|
+
"Number of channels is incompatible with salting trace shape."
|
|
132
|
+
)
|
|
133
|
+
|
|
134
|
+
columns = set(dataframe.get_column_names())
|
|
135
|
+
common = {}
|
|
136
|
+
for name in (
|
|
137
|
+
trigger_column, "salt_template_tag", "salt_channel_expression",
|
|
138
|
+
"saltchanname", "salting_type", "salt_recoil_energy_eV",
|
|
139
|
+
"salt_template_pretrigger_samples", "salt_template_length_samples",
|
|
140
|
+
"dataframe_group_name", "dataframe_group_id",
|
|
141
|
+
"dataframe_group_number", "dataframe_file_index",
|
|
142
|
+
):
|
|
143
|
+
if name in columns:
|
|
144
|
+
common[name] = self._evaluate(dataframe, name)
|
|
145
|
+
|
|
146
|
+
scale_arrays = {}
|
|
147
|
+
for channel in channel_list:
|
|
148
|
+
scale_col = f"salt_template_scale_{channel}"
|
|
149
|
+
if scale_col not in columns:
|
|
150
|
+
scale_col = f"salt_amplitude_{channel}"
|
|
151
|
+
if scale_col in columns:
|
|
152
|
+
scale_arrays[channel] = self._evaluate(dataframe, scale_col)
|
|
153
|
+
|
|
154
|
+
injected_rows = 0
|
|
155
|
+
injected_types = []
|
|
156
|
+
injected_energies = []
|
|
157
|
+
injected_group_names = []
|
|
158
|
+
injected_group_ids = []
|
|
159
|
+
injected_group_numbers = []
|
|
160
|
+
window_start = int(trace_start_index)
|
|
161
|
+
window_stop = window_start + trace_array.shape[-1]
|
|
162
|
+
|
|
163
|
+
for row_index in range(len(dataframe)):
|
|
164
|
+
expression = self._row_value(
|
|
165
|
+
common, "salt_channel_expression", row_index,
|
|
166
|
+
fallback=self._row_value(common, "saltchanname", row_index)
|
|
167
|
+
)
|
|
168
|
+
if expression is None:
|
|
169
|
+
continue
|
|
170
|
+
expression = str(expression)
|
|
171
|
+
salt_channels = convert_channel_name_to_list(expression)
|
|
172
|
+
template_tag = str(
|
|
173
|
+
self._row_value(common, "salt_template_tag", row_index,
|
|
174
|
+
fallback="default")
|
|
175
|
+
)
|
|
176
|
+
salt_trigger = int(common[trigger_column][row_index])
|
|
177
|
+
|
|
178
|
+
# Modern salt metadata carries template support explicitly, which
|
|
179
|
+
# lets us reject non-overlapping rows without loading templates.
|
|
180
|
+
template_length = self._row_value(
|
|
181
|
+
common, "salt_template_length_samples", row_index
|
|
182
|
+
)
|
|
183
|
+
pretrigger = self._row_value(
|
|
184
|
+
common, "salt_template_pretrigger_samples", row_index
|
|
185
|
+
)
|
|
186
|
+
if (template_length is not None and not self._is_missing(template_length)
|
|
187
|
+
and pretrigger is not None and not self._is_missing(pretrigger)):
|
|
188
|
+
template_length = int(template_length)
|
|
189
|
+
pretrigger = int(pretrigger)
|
|
190
|
+
pulse_start = salt_trigger - pretrigger
|
|
191
|
+
pulse_stop = pulse_start + template_length
|
|
192
|
+
if pulse_stop <= window_start or pulse_start >= window_stop:
|
|
193
|
+
continue
|
|
194
|
+
template, template_meta = self._get_template(
|
|
195
|
+
expression, template_tag
|
|
196
|
+
)
|
|
197
|
+
else:
|
|
198
|
+
template, template_meta = self._get_template(
|
|
199
|
+
expression, template_tag
|
|
200
|
+
)
|
|
201
|
+
template_length = int(np.asarray(template).shape[-1])
|
|
202
|
+
if pretrigger is None or self._is_missing(pretrigger):
|
|
203
|
+
pretrigger = template_meta.get(
|
|
204
|
+
"nb_pretrigger_samples", template_length // 2
|
|
205
|
+
)
|
|
206
|
+
pretrigger = int(pretrigger)
|
|
207
|
+
pulse_start = salt_trigger - pretrigger
|
|
208
|
+
pulse_stop = pulse_start + template_length
|
|
209
|
+
|
|
210
|
+
overlap_start = max(window_start, pulse_start)
|
|
211
|
+
overlap_stop = min(window_stop, pulse_stop)
|
|
212
|
+
if overlap_start >= overlap_stop:
|
|
213
|
+
continue
|
|
214
|
+
|
|
215
|
+
row_injected = False
|
|
216
|
+
for trace_chan_index, channel in enumerate(channel_list):
|
|
217
|
+
if channel not in salt_channels:
|
|
218
|
+
continue
|
|
219
|
+
salt_chan_index = salt_channels.index(channel)
|
|
220
|
+
if channel not in scale_arrays:
|
|
221
|
+
continue
|
|
222
|
+
scale = scale_arrays[channel][row_index]
|
|
223
|
+
if self._is_missing(scale):
|
|
224
|
+
continue
|
|
225
|
+
temp = self._select_template_channel(
|
|
226
|
+
template, salt_chan_index, len(salt_channels)
|
|
227
|
+
)
|
|
228
|
+
pulse = np.asarray(temp, dtype=float) * float(scale)
|
|
229
|
+
|
|
230
|
+
trace_lo = overlap_start - window_start
|
|
231
|
+
trace_hi = overlap_stop - window_start
|
|
232
|
+
pulse_lo = overlap_start - pulse_start
|
|
233
|
+
pulse_hi = overlap_stop - pulse_start
|
|
234
|
+
trace_array[trace_chan_index, trace_lo:trace_hi] += pulse[pulse_lo:pulse_hi]
|
|
235
|
+
row_injected = True
|
|
236
|
+
|
|
237
|
+
if row_injected:
|
|
238
|
+
injected_rows += 1
|
|
239
|
+
value = self._row_value(common, "salting_type", row_index)
|
|
240
|
+
if value is not None and not self._is_missing(value):
|
|
241
|
+
injected_types.append(str(value))
|
|
242
|
+
energy = self._row_value(common, "salt_recoil_energy_eV", row_index)
|
|
243
|
+
if energy is not None and not self._is_missing(energy):
|
|
244
|
+
injected_energies.append(float(energy))
|
|
245
|
+
group_name = self._row_value(common, "dataframe_group_name", row_index)
|
|
246
|
+
group_id = self._row_value(common, "dataframe_group_id", row_index)
|
|
247
|
+
group_number = self._row_value(common, "dataframe_group_number", row_index)
|
|
248
|
+
if group_name is not None and not self._is_missing(group_name):
|
|
249
|
+
injected_group_names.append(str(group_name))
|
|
250
|
+
if group_id is not None and not self._is_missing(group_id):
|
|
251
|
+
injected_group_ids.append(str(group_id))
|
|
252
|
+
if group_number is not None and not self._is_missing(group_number):
|
|
253
|
+
injected_group_numbers.append(int(group_number))
|
|
254
|
+
|
|
255
|
+
unique_types = list(dict.fromkeys(injected_types))
|
|
256
|
+
metadata = {}
|
|
257
|
+
if injected_rows:
|
|
258
|
+
metadata = {
|
|
259
|
+
"salting_type": (unique_types[0] if len(unique_types) == 1 else "multiple"),
|
|
260
|
+
"salting_types": unique_types,
|
|
261
|
+
"nb_salts": injected_rows,
|
|
262
|
+
}
|
|
263
|
+
if stream_id is not None:
|
|
264
|
+
metadata["stream_id"] = stream_id
|
|
265
|
+
names = list(dict.fromkeys(injected_group_names))
|
|
266
|
+
ids = list(dict.fromkeys(injected_group_ids))
|
|
267
|
+
numbers = list(dict.fromkeys(injected_group_numbers))
|
|
268
|
+
if len(names) == 1:
|
|
269
|
+
metadata["salting_dataframe_group_name"] = names[0]
|
|
270
|
+
if len(ids) == 1:
|
|
271
|
+
metadata["salting_dataframe_group_id"] = ids[0]
|
|
272
|
+
if len(numbers) == 1:
|
|
273
|
+
metadata["salting_dataframe_group_number"] = numbers[0]
|
|
274
|
+
if injected_energies:
|
|
275
|
+
metadata["salt_recoil_energy_eV"] = injected_energies
|
|
276
|
+
|
|
277
|
+
output = trace_array[0] if input_was_1d else trace_array
|
|
278
|
+
return output, metadata
|
|
279
|
+
|
|
280
|
+
def _get_template(self, expression, tag):
|
|
281
|
+
key = (str(expression), str(tag))
|
|
282
|
+
cached = self._template_cache.get(key)
|
|
283
|
+
if cached is not None:
|
|
284
|
+
return cached
|
|
285
|
+
template, _, metadata = self._filter_data.get_template(
|
|
286
|
+
expression, tag=tag, return_metadata=True
|
|
287
|
+
)
|
|
288
|
+
value = (template, metadata)
|
|
289
|
+
self._template_cache[key] = value
|
|
290
|
+
return value
|
|
291
|
+
|
|
292
|
+
@staticmethod
|
|
293
|
+
def _normalize_channels(channels):
|
|
294
|
+
if isinstance(channels, str):
|
|
295
|
+
return convert_channel_name_to_list(channels)
|
|
296
|
+
return list(channels)
|
|
297
|
+
|
|
298
|
+
@staticmethod
|
|
299
|
+
def _evaluate(dataframe, column):
|
|
300
|
+
values = dataframe.evaluate(column, array_type="numpy")
|
|
301
|
+
if np.ma.isMaskedArray(values):
|
|
302
|
+
values = values.filled(np.nan)
|
|
303
|
+
return np.asarray(values)
|
|
304
|
+
|
|
305
|
+
@staticmethod
|
|
306
|
+
def _row_value(common, name, index, fallback=None):
|
|
307
|
+
values = common.get(name)
|
|
308
|
+
if values is None:
|
|
309
|
+
return fallback
|
|
310
|
+
return values[index]
|
|
311
|
+
|
|
312
|
+
@staticmethod
|
|
313
|
+
def _is_missing(value):
|
|
314
|
+
if value is None:
|
|
315
|
+
return True
|
|
316
|
+
try:
|
|
317
|
+
return bool(np.isnan(value))
|
|
318
|
+
except (TypeError, ValueError):
|
|
319
|
+
return False
|
|
320
|
+
|
|
321
|
+
@staticmethod
|
|
322
|
+
def _select_template_channel(template, index, nb_channels):
|
|
323
|
+
template = np.asarray(template)
|
|
324
|
+
if nb_channels == 1:
|
|
325
|
+
return template.reshape(-1) if template.ndim != 1 else template
|
|
326
|
+
selected = np.asarray(template[index])
|
|
327
|
+
if selected.ndim > 1:
|
|
328
|
+
selected = selected[0]
|
|
329
|
+
return selected
|