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.
Files changed (60) hide show
  1. pytesprocess/__init__.py +9 -0
  2. pytesprocess/_version.py +2 -0
  3. pytesprocess/cli/__init__.py +1 -0
  4. pytesprocess/cli/commands/__init__.py +5 -0
  5. pytesprocess/cli/commands/event.py +66 -0
  6. pytesprocess/cli/commands/filter.py +17 -0
  7. pytesprocess/cli/commands/ivsweep.py +29 -0
  8. pytesprocess/cli/common.py +86 -0
  9. pytesprocess/cli/main.py +81 -0
  10. pytesprocess/config/__init__.py +4 -0
  11. pytesprocess/config/loader.py +94 -0
  12. pytesprocess/config/manager.py +297 -0
  13. pytesprocess/config/resolvers/__init__.py +5 -0
  14. pytesprocess/config/resolvers/common.py +56 -0
  15. pytesprocess/config/resolvers/feature.py +293 -0
  16. pytesprocess/config/resolvers/salting.py +86 -0
  17. pytesprocess/config/resolvers/trigger.py +84 -0
  18. pytesprocess/config/selectors.py +108 -0
  19. pytesprocess/config/validation.py +314 -0
  20. pytesprocess/config/warnings.py +2 -0
  21. pytesprocess/core/__init__.py +10 -0
  22. pytesprocess/core/algorithms.py +1455 -0
  23. pytesprocess/core/didv.py +1648 -0
  24. pytesprocess/core/eventbuilder.py +495 -0
  25. pytesprocess/core/filterbuilder.py +81 -0
  26. pytesprocess/core/filterdata.py +1849 -0
  27. pytesprocess/core/ivsweep.py +2072 -0
  28. pytesprocess/core/noise.py +923 -0
  29. pytesprocess/core/noisemodel.py +1408 -0
  30. pytesprocess/core/oftrigger.py +1035 -0
  31. pytesprocess/core/template.py +450 -0
  32. pytesprocess/process/__init__.py +6 -0
  33. pytesprocess/process/data_source.py +185 -0
  34. pytesprocess/process/event_context.py +35 -0
  35. pytesprocess/process/feature_plan.py +186 -0
  36. pytesprocess/process/feature_resources.py +267 -0
  37. pytesprocess/process/features.py +1024 -0
  38. pytesprocess/process/filterprocess.py +1176 -0
  39. pytesprocess/process/ivprocess.py +1380 -0
  40. pytesprocess/process/processing_data.py +967 -0
  41. pytesprocess/process/randoms.py +921 -0
  42. pytesprocess/process/triggers.py +1011 -0
  43. pytesprocess/salting/__init__.py +7 -0
  44. pytesprocess/salting/generator.py +364 -0
  45. pytesprocess/salting/injector.py +329 -0
  46. pytesprocess/salting/sampling.py +84 -0
  47. pytesprocess/utils/__init__.py +5 -0
  48. pytesprocess/utils/arg_utils.py +122 -0
  49. pytesprocess/utils/dataframe_output.py +120 -0
  50. pytesprocess/utils/filter_hdf5.py +594 -0
  51. pytesprocess/utils/utils.py +701 -0
  52. pytesprocess/workflows/__init__.py +3 -0
  53. pytesprocess/workflows/processing.py +317 -0
  54. pytesprocess/workflows/salting.py +133 -0
  55. pytesprocess-0.1.1.dist-info/METADATA +211 -0
  56. pytesprocess-0.1.1.dist-info/RECORD +60 -0
  57. pytesprocess-0.1.1.dist-info/WHEEL +5 -0
  58. pytesprocess-0.1.1.dist-info/entry_points.txt +2 -0
  59. pytesprocess-0.1.1.dist-info/licenses/LICENSE +21 -0
  60. 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