|
10 | 10 | SEQUENCE_ROOT = Path("tests/sequence_data") |
11 | 11 |
|
12 | 12 |
|
13 | | -@contextmanager |
14 | | -def create_sequence_data( |
| 13 | +def _generate_sequence_data( |
| 14 | + sequence_root, |
15 | 15 | n_signals=10, |
16 | 16 | shifts_per_signal=False, |
17 | 17 | use_mem_mapped=False, |
| 18 | + start_time=0.0, |
18 | 19 | t_end=10.0, |
19 | 20 | sampling_rate=10.0, |
20 | 21 | contain_nans=False, |
21 | 22 | ): |
22 | | - try: |
23 | | - SEQUENCE_ROOT.mkdir(parents=True, exist_ok=True) |
24 | | - (SEQUENCE_ROOT / "meta").mkdir(parents=True, exist_ok=True) |
25 | | - |
26 | | - meta = { |
27 | | - "start_time": 0, |
28 | | - "end_time": t_end, |
29 | | - "modality": "sequence", |
30 | | - "sampling_rate": sampling_rate, |
31 | | - "phase_shift_per_signal": shifts_per_signal, |
32 | | - "is_mem_mapped": use_mem_mapped, |
33 | | - "n_signals": n_signals, |
34 | | - } |
35 | | - |
36 | | - timestamps = np.linspace( |
37 | | - meta["start_time"], |
38 | | - meta["end_time"], |
39 | | - int((meta["end_time"] - meta["start_time"]) * meta["sampling_rate"]) + 1, |
| 23 | + """Generates synthetic sequence data folders for testing interpolator logic.""" |
| 24 | + |
| 25 | + sequence_root = Path(sequence_root) |
| 26 | + sequence_root.mkdir(parents=True, exist_ok=True) |
| 27 | + (sequence_root / "meta").mkdir(parents=True, exist_ok=True) |
| 28 | + |
| 29 | + meta = { |
| 30 | + "start_time": start_time, |
| 31 | + "end_time": t_end, |
| 32 | + "modality": "sequence", |
| 33 | + "sampling_rate": sampling_rate, |
| 34 | + "phase_shift_per_signal": shifts_per_signal, |
| 35 | + "is_mem_mapped": use_mem_mapped, |
| 36 | + "n_signals": n_signals, |
| 37 | + } |
| 38 | + |
| 39 | + # Determine number of samples based on duration and rate |
| 40 | + duration = meta["end_time"] - meta["start_time"] |
| 41 | + n_samples = int(round(duration * meta["sampling_rate"])) + 1 |
| 42 | + |
| 43 | + timestamps = np.linspace(meta["start_time"], meta["end_time"], n_samples) |
| 44 | + |
| 45 | + data = np.random.rand(len(timestamps), n_signals) |
| 46 | + |
| 47 | + if contain_nans: |
| 48 | + nan_indices = np.random.choice( |
| 49 | + data.size, size=int(0.1 * data.size), replace=False |
40 | 50 | ) |
41 | | - np.save(SEQUENCE_ROOT / "timestamps.npy", timestamps) |
42 | | - meta["n_timestamps"] = len(timestamps) |
| 51 | + data.flat[nan_indices] = np.nan |
| 52 | + # ensure each row has at least one non-NaN |
| 53 | + if n_signals > 0: |
| 54 | + row_all_nan = np.isnan(data).all(axis=1) |
| 55 | + data[row_all_nan, 0] = 0.0 |
43 | 56 |
|
44 | | - data = np.random.rand(len(timestamps), n_signals) |
| 57 | + if not use_mem_mapped: |
| 58 | + np.save(sequence_root / "data.npy", data) |
| 59 | + else: |
| 60 | + filename = sequence_root / "data.mem" |
| 61 | + fp = np.memmap(filename, dtype=data.dtype, mode="w+", shape=data.shape) |
| 62 | + fp[:] = data[:] |
| 63 | + fp.flush() |
| 64 | + del fp |
45 | 65 |
|
46 | | - if contain_nans: |
47 | | - nan_indices = np.random.choice( |
48 | | - data.size, size=int(0.1 * data.size), replace=False |
49 | | - ) |
50 | | - data.flat[nan_indices] = np.nan |
| 66 | + np.save(sequence_root / "timestamps.npy", timestamps) |
| 67 | + meta["n_timestamps"] = len(timestamps) |
| 68 | + meta["dtype"] = str(data.dtype) |
51 | 69 |
|
52 | | - if not use_mem_mapped: |
53 | | - np.save(SEQUENCE_ROOT / "data.npy", data) |
54 | | - else: |
55 | | - filename = SEQUENCE_ROOT / "data.mem" |
| 70 | + # Handle per-signal phase shifts if required by the test case |
| 71 | + shifts = None |
| 72 | + if shifts_per_signal: |
| 73 | + shifts = np.random.rand(n_signals) / meta["sampling_rate"] * 0.9 |
| 74 | + np.save(sequence_root / "meta" / "phase_shifts.npy", shifts) |
56 | 75 |
|
57 | | - fp = np.memmap(filename, dtype=data.dtype, mode="w+", shape=data.shape) |
58 | | - fp[:] = data[:] |
59 | | - fp.flush() # Ensure data is written to disk |
60 | | - del fp |
61 | | - meta["dtype"] = str(data.dtype) |
| 76 | + with open(sequence_root / "meta.yml", "w") as f: |
| 77 | + yaml.safe_dump(meta, f) |
62 | 78 |
|
63 | | - if shifts_per_signal: |
64 | | - shifts = np.random.rand(n_signals) / meta["sampling_rate"] * 0.9 |
65 | | - np.save(SEQUENCE_ROOT / "meta" / "phase_shifts.npy", shifts) |
| 79 | + return timestamps, data, shifts |
66 | 80 |
|
67 | | - with open(SEQUENCE_ROOT / "meta.yml", "w") as f: |
68 | | - yaml.safe_dump(meta, f) |
69 | 81 |
|
70 | | - yield timestamps, data, shifts if shifts_per_signal else None |
| 82 | +@contextmanager |
| 83 | +def create_sequence_data( |
| 84 | + n_signals=10, |
| 85 | + shifts_per_signal=False, |
| 86 | + use_mem_mapped=False, |
| 87 | + t_end=10.0, |
| 88 | + sampling_rate=10.0, |
| 89 | + contain_nans=False, |
| 90 | + start_time=0.0, |
| 91 | +): |
| 92 | + """Context manager for temporary sequence data creation and cleanup.""" |
| 93 | + try: |
| 94 | + yield _generate_sequence_data( |
| 95 | + sequence_root=SEQUENCE_ROOT, |
| 96 | + n_signals=n_signals, |
| 97 | + shifts_per_signal=shifts_per_signal, |
| 98 | + use_mem_mapped=use_mem_mapped, |
| 99 | + t_end=t_end, |
| 100 | + sampling_rate=sampling_rate, |
| 101 | + contain_nans=contain_nans, |
| 102 | + start_time=start_time, |
| 103 | + ) |
71 | 104 | finally: |
72 | | - shutil.rmtree(SEQUENCE_ROOT) |
| 105 | + if SEQUENCE_ROOT.exists(): |
| 106 | + shutil.rmtree(SEQUENCE_ROOT) |
73 | 107 |
|
74 | 108 |
|
75 | 109 | @contextmanager |
76 | 110 | def sequence_data_and_interpolator(data_kwargs=None, interp_kwargs=None): |
77 | 111 | data_kwargs = data_kwargs or {} |
78 | 112 | interp_kwargs = interp_kwargs or {} |
79 | 113 | with create_sequence_data(**data_kwargs) as (timestamps, data, shifts): |
| 114 | + # Restore the helper expected by the rest of the test suite |
| 115 | + |
80 | 116 | with closing( |
81 | | - Interpolator.create("tests/sequence_data", **interp_kwargs) |
| 117 | + Interpolator.create(str(SEQUENCE_ROOT), **interp_kwargs) |
82 | 118 | ) as seq_interp: |
83 | 119 | yield timestamps, data, shifts, seq_interp |
0 commit comments