Skip to content

Commit 5f71bde

Browse files
committed
Refactor frame methods to improve waterfall handling and integrate observational context preservation in derived frames
1 parent dde82ba commit 5f71bde

5 files changed

Lines changed: 100 additions & 39 deletions

File tree

setigen/frame.py

Lines changed: 8 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -29,9 +29,10 @@
2929
_generate_sampled_noise,
3030
)
3131
from ._frame.io import (
32-
_decode_bytestrings,
33-
_encode_bytestrings,
34-
_update_waterfall,
32+
_check_waterfall,
33+
_get_waterfall,
34+
_save_fil,
35+
_save_hdf5,
3536
)
3637
from ._frame.signal import (
3738
_finalize_signal,
@@ -767,8 +768,7 @@ def get_waterfall(self) -> Any:
767768
Returns:
768769
Waterfall representation of the frame.
769770
"""
770-
_update_waterfall(self)
771-
return self.waterfall
771+
return _get_waterfall(self)
772772

773773
def check_waterfall(self) -> Any:
774774
"""Return the updated attached waterfall when one exists.
@@ -777,9 +777,7 @@ def check_waterfall(self) -> Any:
777777
Updated waterfall object or `None` when the frame has no attached
778778
waterfall.
779779
"""
780-
if self.waterfall is None:
781-
return None
782-
return self.get_waterfall()
780+
return _check_waterfall(self)
783781

784782
def save_fil(self, filename: PathLike, max_load: int = 1) -> None:
785783
"""Save frame data as a SIGPROC filterbank file.
@@ -788,10 +786,7 @@ def save_fil(self, filename: PathLike, max_load: int = 1) -> None:
788786
filename: Output `.fil` path.
789787
max_load: Maximum load parameter for a lazily created waterfall.
790788
"""
791-
_update_waterfall(self, filename=filename, max_load=max_load)
792-
_encode_bytestrings(self)
793-
self.waterfall.write_to_fil(filename)
794-
_decode_bytestrings(self)
789+
_save_fil(self, filename, max_load=max_load)
795790

796791
def save_hdf5(self, filename: PathLike, max_load: int = 1) -> None:
797792
"""Save frame data as an HDF5 waterfall file.
@@ -800,10 +795,7 @@ def save_hdf5(self, filename: PathLike, max_load: int = 1) -> None:
800795
filename: Output `.h5` path.
801796
max_load: Maximum load parameter for a lazily created waterfall.
802797
"""
803-
_update_waterfall(self, filename=filename, max_load=max_load)
804-
_encode_bytestrings(self)
805-
self.waterfall.write_to_hdf5(filename)
806-
_decode_bytestrings(self)
798+
_save_hdf5(self, filename, max_load=max_load)
807799

808800
def save_h5(self, filename: PathLike, max_load: int = 1) -> None:
809801
"""Save frame data as an HDF5 waterfall file.

setigen/integrate.py

Lines changed: 29 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
from __future__ import annotations
22

3+
import copy
34
from enum import Enum
45
from typing import Any
56

@@ -54,6 +55,27 @@ def _resolve_integration_mode(mode: IntegrationMode | str) -> IntegrationMode:
5455
return IntegrationMode.MEAN
5556

5657

58+
def _copy_frame_context(source: Any, target: Any) -> None:
59+
"""Copy non-shape observational context between frame-like objects.
60+
61+
Args:
62+
source: Original frame-like object.
63+
target: Derived frame-like object to update.
64+
"""
65+
metadata = copy.deepcopy(getattr(source, "metadata", {}))
66+
try:
67+
for key, value in source.get_params().items():
68+
if metadata.get(key) == value:
69+
metadata.pop(key)
70+
except AttributeError:
71+
pass
72+
if metadata:
73+
target.add_metadata(metadata)
74+
75+
if hasattr(source, "header"):
76+
target.header = copy.deepcopy(source.header)
77+
78+
5779
def integrate(
5880
fr: Any,
5981
axis: IntegrationAxis | str | int = 't',
@@ -101,15 +123,20 @@ def integrate(
101123
fch1=fr.fmid,
102124
ascending=fr.ascending,
103125
data=data,
104-
seed=fr.rng)
126+
seed=fr.rng,
127+
t_start=fr.t_start,
128+
source_name=fr.source_name)
105129
else:
106130
# Spectrum
107131
new_fr = Spectrum(df=fr.df,
108132
dt=fr.dt * fr.tchans,
109133
fch1=fr.fch1,
110134
ascending=fr.ascending,
111135
data=data,
112-
seed=fr.rng)
136+
seed=fr.rng,
137+
t_start=fr.t_start,
138+
source_name=fr.source_name)
139+
_copy_frame_context(fr, new_fr)
113140
return new_fr
114141
else:
115142
return data.flatten()

setigen/split_utils.py

Lines changed: 12 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@
88
from typing import Iterator
99

1010
from ._typing import PathLike
11+
from ._frame.io import _close_waterfall_handles
1112

1213

1314
def split_waterfall_generator(
@@ -32,10 +33,13 @@ def split_waterfall_generator(
3233
"""
3334

3435
info_wf = Waterfall(waterfall_fn, load_data=False)
35-
fch1 = info_wf.header['fch1']
36-
nchans = info_wf.header['nchans']
37-
df = info_wf.header['foff']
38-
tchans_tot = info_wf.container.selection_shape[0]
36+
try:
37+
fch1 = info_wf.header['fch1']
38+
nchans = info_wf.header['nchans']
39+
df = info_wf.header['foff']
40+
tchans_tot = info_wf.container.selection_shape[0]
41+
finally:
42+
_close_waterfall_handles(info_wf)
3943

4044
if f_shift is None:
4145
f_shift = fchans
@@ -58,7 +62,10 @@ def split_waterfall_generator(
5862
t_start=0,
5963
t_stop=tchans)
6064

61-
yield waterfall
65+
try:
66+
yield waterfall
67+
finally:
68+
_close_waterfall_handles(waterfall)
6269

6370
f_start += f_shift * df
6471
f_stop += f_shift * df

setigen/waterfall_utils.py

Lines changed: 30 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55
from blimpy import Waterfall
66

77
from ._typing import PathLike
8+
from ._frame.io import _close_waterfall_handles
89

910

1011
def max_freq(waterfall: PathLike | Waterfall) -> float:
@@ -44,15 +45,20 @@ def get_data(waterfall: PathLike | Waterfall, db: bool = False) -> np.ndarray:
4445
Raises:
4546
ValueError: If the waterfall input type is unsupported.
4647
"""
47-
if isinstance(waterfall, (str, PurePath)):
48+
owns_waterfall = isinstance(waterfall, (str, PurePath))
49+
if owns_waterfall:
4850
waterfall = Waterfall(waterfall)
4951
elif not isinstance(waterfall, Waterfall):
5052
raise ValueError('Invalid data file!')
5153

52-
if db:
53-
return 10 * np.log10(waterfall.data[:, 0, :])
54-
55-
return waterfall.data[:, 0, :]
54+
try:
55+
data = waterfall.data[:, 0, :]
56+
if db:
57+
data = 10 * np.log10(data)
58+
return np.array(data, copy=owns_waterfall)
59+
finally:
60+
if owns_waterfall:
61+
_close_waterfall_handles(waterfall)
5662

5763

5864
def get_fs(waterfall: PathLike | Waterfall) -> np.ndarray:
@@ -67,16 +73,20 @@ def get_fs(waterfall: PathLike | Waterfall) -> np.ndarray:
6773
Raises:
6874
ValueError: If the waterfall input type is unsupported.
6975
"""
70-
if isinstance(waterfall, (str, PurePath)):
76+
owns_waterfall = isinstance(waterfall, (str, PurePath))
77+
if owns_waterfall:
7178
waterfall = Waterfall(waterfall, load_data=False)
7279
elif not isinstance(waterfall, Waterfall):
7380
raise ValueError('Invalid data file!')
7481

75-
fch1 = waterfall.header['fch1']
76-
df = waterfall.header['foff']
77-
fchans = waterfall.header['nchans']
78-
79-
return fch1 + np.arange(fchans) * df
82+
try:
83+
fch1 = waterfall.header['fch1']
84+
df = waterfall.header['foff']
85+
fchans = waterfall.header['nchans']
86+
return fch1 + np.arange(fchans) * df
87+
finally:
88+
if owns_waterfall:
89+
_close_waterfall_handles(waterfall)
8090

8191

8292
def get_ts(waterfall: PathLike | Waterfall) -> np.ndarray:
@@ -91,12 +101,16 @@ def get_ts(waterfall: PathLike | Waterfall) -> np.ndarray:
91101
Raises:
92102
ValueError: If the waterfall input type is unsupported.
93103
"""
94-
if isinstance(waterfall, (str, PurePath)):
104+
owns_waterfall = isinstance(waterfall, (str, PurePath))
105+
if owns_waterfall:
95106
waterfall = Waterfall(waterfall, load_data=False)
96107
elif not isinstance(waterfall, Waterfall):
97108
raise ValueError('Invalid data file!')
98109

99-
tsamp = waterfall.header['tsamp']
100-
tchans = waterfall.container.selection_shape[0]
101-
102-
return np.arange(tchans) * tsamp
110+
try:
111+
tsamp = waterfall.header['tsamp']
112+
tchans = waterfall.container.selection_shape[0]
113+
return np.arange(tchans) * tsamp
114+
finally:
115+
if owns_waterfall:
116+
_close_waterfall_handles(waterfall)

tests/test_frame_creation.py

Lines changed: 21 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -129,6 +129,27 @@ def test_axis_semantics():
129129
assert frame.get_drift_rate(0, 3, reference="centers") == pytest.approx(7.5)
130130

131131

132+
def test_integrated_frames_preserve_observational_context():
133+
frame = stg.Frame(shape=(4, 8),
134+
t_start=123456,
135+
source_name="Target",
136+
data=np.ones((4, 8)))
137+
frame.header = {"rawdatafile": "original.raw", "source_name": "Target"}
138+
frame.add_metadata({"drift_rate": 1.25})
139+
140+
spectrum = stg.spectrum(frame, mode="sum")
141+
timeseries = stg.timeseries(frame, mode="sum")
142+
143+
for derived in [spectrum, timeseries]:
144+
assert derived.t_start == frame.t_start
145+
assert derived.source_name == frame.source_name
146+
assert derived.header == frame.header
147+
assert derived.header is not frame.header
148+
assert derived.metadata["drift_rate"] == frame.metadata["drift_rate"]
149+
assert derived.metadata["fchans"] == derived.fchans
150+
assert derived.metadata["tchans"] == derived.tchans
151+
152+
132153
def test_h5_ingestion_does_not_retain_live_waterfall(tmp_path):
133154
frame = stg.Frame(shape=(4, 8), seed=0)
134155
h5_path = tmp_path / "frame.h5"

0 commit comments

Comments
 (0)