Skip to content

Commit bea760a

Browse files
committed
fix: linear-interpolation
1 parent 05cf038 commit bea760a

3 files changed

Lines changed: 14 additions & 12 deletions

File tree

diff_output.txt

40.3 KB
Binary file not shown.

experanto/interpolators.py

Lines changed: 6 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -315,8 +315,8 @@ def interpolate(
315315
idx_upper = idx_upper[~overflow_mask]
316316
idx_lower = idx_lower[~overflow_mask]
317317

318-
times_lower = idx_lower * self.time_delta
319-
times_upper = idx_upper * self.time_delta
318+
times_lower = self.start_time + (idx_lower * self.time_delta)
319+
times_upper = self.start_time + (idx_upper * self.time_delta)
320320
denom = times_upper - times_lower
321321

322322
times_valid = valid_times[~overflow_mask]
@@ -449,14 +449,12 @@ def interpolate(
449449

450450
valid = valid[~overflow_mask.any(axis=1)]
451451

452-
times_lower = idx_lower * self.time_delta
453-
times_upper = idx_upper * self.time_delta
452+
times_lower = self.start_time + (idx_lower * self.time_delta) + self._phase_shifts[np.newaxis, :]
453+
times_upper = self.start_time + (idx_upper * self.time_delta) + self._phase_shifts[np.newaxis, :]
454454
denom = times_upper - times_lower
455455

456-
time_dim = valid_times[:, np.newaxis] - self._phase_shifts[np.newaxis, :]
457-
458-
lower_numerator = times_upper - time_dim
459-
upper_numerator = time_dim - times_lower
456+
lower_numerator = times_upper - valid_times[:, np.newaxis]
457+
upper_numerator = valid_times[:, np.newaxis] - times_lower
460458

461459
lower_signal_ratio = lower_numerator / denom
462460
upper_signal_ratio = upper_numerator / denom

tests/test_sequence_interpolator.py

Lines changed: 8 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -216,14 +216,16 @@ def test_nearest_neighbor_interpolation_with_phase_shifts_handles_nans(
216216
@pytest.mark.parametrize("use_mem_mapped", [False, True])
217217
@pytest.mark.parametrize("contain_nans", [False, True])
218218
@pytest.mark.parametrize("keep_nans", [False, True])
219+
@pytest.mark.parametrize("start_time", [0.0, 15.0])
219220
def test_linear_interpolation(
220-
n_signals, sampling_rate, use_mem_mapped, contain_nans, keep_nans
221+
n_signals, sampling_rate, use_mem_mapped, contain_nans, keep_nans, start_time
221222
):
222223
with sequence_data_and_interpolator(
223224
data_kwargs={
224225
"n_signals": n_signals,
225226
"use_mem_mapped": use_mem_mapped,
226-
"t_end": 5.0,
227+
"start_time": start_time,
228+
"t_end": start_time + 5.0,
227229
"sampling_rate": sampling_rate,
228230
"contain_nans": contain_nans,
229231
},
@@ -269,14 +271,16 @@ def test_linear_interpolation(
269271
@pytest.mark.parametrize("sampling_rate", [3.0, 10.0, 100.0])
270272
@pytest.mark.parametrize("use_mem_mapped", [False, True])
271273
@pytest.mark.parametrize("keep_nans", [False, True])
274+
@pytest.mark.parametrize("start_time", [0.0, 15.0])
272275
def test_linear_interpolation_with_phase_shifts(
273-
n_signals, sampling_rate, use_mem_mapped, keep_nans
276+
n_signals, sampling_rate, use_mem_mapped, keep_nans, start_time
274277
):
275278
with sequence_data_and_interpolator(
276279
data_kwargs={
277280
"n_signals": n_signals,
278281
"use_mem_mapped": use_mem_mapped,
279-
"t_end": 5.0,
282+
"start_time": start_time,
283+
"t_end": start_time + 5.0,
280284
"sampling_rate": sampling_rate,
281285
"shifts_per_signal": True,
282286
},

0 commit comments

Comments
 (0)