Skip to content

Commit 478a2e1

Browse files
authored
Migrate tests/test_motifs.py from npt.assert_almost_equal to npt.assert_allclose (#1202)
1 parent aa897ee commit 478a2e1

1 file changed

Lines changed: 58 additions & 56 deletions

File tree

‎tests/test_motifs.py‎

Lines changed: 58 additions & 56 deletions
Original file line numberDiff line numberDiff line change
@@ -166,20 +166,20 @@ def test_motifs_one_motif():
166166
m = 3
167167
max_motifs = 1
168168

169-
left_indices = [[0, 5, 9]]
170-
left_profile_values = [[0.0, 0.0, 0.0]]
169+
ref_indices = [[0, 5, 9]]
170+
ref_profile_values = [[0.0, 0.0, 0.0]]
171171

172172
mp = naive.stump(T, m)
173-
right_distance_values, right_indices = motifs(
173+
cmp_distance_values, cmp_indices = motifs(
174174
T,
175175
mp[:, 0],
176176
max_distance=lambda D: 0.001, # Also test lambda functionality
177177
max_motifs=max_motifs,
178178
cutoff=np.inf,
179179
)
180180

181-
npt.assert_array_equal(left_indices, right_indices)
182-
npt.assert_almost_equal(left_profile_values, right_distance_values)
181+
npt.assert_array_equal(cmp_indices, ref_indices)
182+
npt.assert_allclose(cmp_distance_values, ref_profile_values, atol=1.5e-07)
183183

184184

185185
def test_motifs_two_motifs():
@@ -217,8 +217,8 @@ def test_motifs_two_motifs():
217217

218218
mp = naive.stump(T, m)
219219

220-
# left_indices = [[70, 170, -1], [10, 210, 110]]
221-
left_profile_values = [
220+
# ref_indices = [[70, 170, -1], [10, 210, 110]]
221+
ref_profile_values = [
222222
[0.0, 0.0, np.nan],
223223
[
224224
0.0,
@@ -227,7 +227,7 @@ def test_motifs_two_motifs():
227227
],
228228
]
229229

230-
right_distance_values, right_indices = motifs(
230+
cmp_distance_values, cmp_indices = motifs(
231231
T,
232232
mp[:, 0],
233233
max_motifs=max_motifs,
@@ -237,7 +237,7 @@ def test_motifs_two_motifs():
237237

238238
# We ignore indices because of sorting ambiguities for equal distances.
239239
# As long as the distances are correct, the indices will be too.
240-
npt.assert_almost_equal(left_profile_values, right_distance_values)
240+
npt.assert_allclose(cmp_distance_values, ref_profile_values, atol=1.5e-07)
241241

242242

243243
def test_motifs_max_matches():
@@ -277,20 +277,20 @@ def test_motifs_max_matches():
277277
max_motifs = 2
278278
max_matches = 3
279279

280-
left_indices = [[0, 7], [4, 11]]
281-
left_profile_values = [
280+
ref_indices = [[0, 7], [4, 11]]
281+
ref_profile_values = [
282282
[0.0, 0.0],
283283
[
284284
0.0,
285285
naive.distance(
286-
core.z_norm(T[left_indices[1][0] : left_indices[1][0] + m]),
287-
core.z_norm(T[left_indices[1][1] : left_indices[1][1] + m]),
286+
core.z_norm(T[ref_indices[1][0] : ref_indices[1][0] + m]),
287+
core.z_norm(T[ref_indices[1][1] : ref_indices[1][1] + m]),
288288
),
289289
],
290290
]
291291

292292
mp = naive.stump(T, m)
293-
right_distance_values, right_indices = motifs(
293+
cmp_distance_values, cmp_indices = motifs(
294294
T,
295295
mp[:, 0],
296296
max_motifs=max_motifs,
@@ -301,7 +301,7 @@ def test_motifs_max_matches():
301301

302302
# We ignore indices because of sorting ambiguities for equal distances.
303303
# As long as the distances are correct, the indices will be too.
304-
npt.assert_almost_equal(left_profile_values, right_distance_values)
304+
npt.assert_allclose(cmp_distance_values, ref_profile_values, atol=1.5e-07)
305305

306306

307307
def test_motifs_max_matches_max_distances_inf():
@@ -342,21 +342,21 @@ def test_motifs_max_matches_max_distances_inf():
342342
max_matches = 2
343343
max_distance = np.inf
344344

345-
left_indices = [[0, 7], [4, 11]]
346-
left_profile_values = [
345+
ref_indices = [[0, 7], [4, 11]]
346+
ref_profile_values = [
347347
[0.0, 0.0],
348348
[
349349
0.0,
350350
naive.distance(
351-
core.z_norm(T[left_indices[1][0] : left_indices[1][0] + m]),
352-
core.z_norm(T[left_indices[1][1] : left_indices[1][1] + m]),
351+
core.z_norm(T[ref_indices[1][0] : ref_indices[1][0] + m]),
352+
core.z_norm(T[ref_indices[1][1] : ref_indices[1][1] + m]),
353353
),
354354
],
355355
]
356356

357357
# set `row_wise` to True so that we can compare the indices of motifs as well
358358
mp = naive.stump(T, m, row_wise=True)
359-
right_distance_values, right_indices = motifs(
359+
cmp_distance_values, cmp_indices = motifs(
360360
T,
361361
mp[:, 0],
362362
max_motifs=max_motifs,
@@ -365,8 +365,8 @@ def test_motifs_max_matches_max_distances_inf():
365365
max_matches=max_matches,
366366
)
367367

368-
npt.assert_almost_equal(left_indices, right_indices)
369-
npt.assert_almost_equal(left_profile_values, right_distance_values)
368+
npt.assert_allclose(cmp_indices, ref_indices, atol=1.5e-07)
369+
npt.assert_allclose(cmp_distance_values, ref_profile_values, atol=1.5e-07)
370370

371371

372372
def test_naive_match_exclusion_zone():
@@ -378,12 +378,12 @@ def test_naive_match_exclusion_zone():
378378
m = Q.shape[0]
379379
excl_zone = int(np.ceil(m / 4))
380380

381-
left = [
381+
ref = [
382382
[0, 1],
383383
[naive.distance(core.z_norm(Q), core.z_norm(T[5 : 5 + m])), 5],
384384
[naive.distance(core.z_norm(Q), core.z_norm(T[9 : 9 + m])), 9],
385385
]
386-
right = list(
386+
cmp = list(
387387
naive_match(
388388
Q,
389389
T,
@@ -392,9 +392,11 @@ def test_naive_match_exclusion_zone():
392392
)
393393
)
394394
# To avoid sorting errors we first sort based on distance and then based on indices
395-
right.sort(key=lambda x: (x[1], x[0]))
395+
cmp.sort(key=lambda x: (x[1], x[0]))
396396

397-
npt.assert_almost_equal(left, right)
397+
npt.assert_allclose(
398+
np.array(cmp).astype(np.float64), np.array(ref).astype(np.float64), atol=1.5e-07
399+
)
398400

399401

400402
@pytest.mark.parametrize("Q, T", test_data)
@@ -403,21 +405,21 @@ def test_match(Q, T):
403405
excl_zone = int(np.ceil(m / 4))
404406
max_distance = 0.3
405407

406-
left = naive_match(
408+
ref = naive_match(
407409
Q,
408410
T,
409411
excl_zone,
410412
max_distance=max_distance,
411413
)
412414

413-
right = match(
415+
cmp = match(
414416
Q,
415417
T,
416418
max_matches=None,
417419
max_distance=lambda D: max_distance, # also test lambda functionality
418420
)
419421

420-
npt.assert_almost_equal(left, right)
422+
npt.assert_allclose(cmp.astype(np.float64), ref.astype(np.float64), atol=1.5e-07)
421423

422424

423425
@pytest.mark.parametrize("Q, T", test_data)
@@ -426,7 +428,7 @@ def test_match_mean_stddev(Q, T):
426428
excl_zone = int(np.ceil(m / 4))
427429
max_distance = 0.3
428430

429-
left = naive_match(
431+
ref = naive_match(
430432
Q,
431433
T,
432434
excl_zone,
@@ -435,7 +437,7 @@ def test_match_mean_stddev(Q, T):
435437

436438
M_T, Σ_T = naive.compute_mean_std(T, len(Q))
437439

438-
right = match(
440+
cmp = match(
439441
Q,
440442
T,
441443
M_T,
@@ -444,7 +446,7 @@ def test_match_mean_stddev(Q, T):
444446
max_distance=lambda D: max_distance, # also test lambda functionality
445447
)
446448

447-
npt.assert_almost_equal(left, right)
449+
npt.assert_allclose(cmp.astype(np.float64), ref.astype(np.float64), atol=1.5e-07)
448450

449451

450452
@pytest.mark.parametrize("Q, T", test_data)
@@ -457,28 +459,28 @@ def test_match_isconstant(Q, T):
457459
naive.isconstant_func_stddev_threshold, quantile_threshold=0.05
458460
)
459461

460-
left = naive_match(
462+
ref = naive_match(
461463
Q,
462464
T,
463465
excl_zone,
464466
max_distance=max_distance,
465467
T_subseq_isconstant=T_subseq_isconstant,
466468
)
467469

468-
right = match(
470+
cmp = match(
469471
Q,
470472
T,
471473
max_matches=None,
472474
max_distance=lambda D: max_distance, # also test lambda functionality
473475
T_subseq_isconstant=T_subseq_isconstant,
474476
)
475477

476-
npt.assert_almost_equal(left, right)
478+
npt.assert_allclose(cmp.astype(np.float64), ref.astype(np.float64), atol=1.5e-07)
477479

478480
# Test for when Q is constant
479481
Q_subseq_isconstant = np.array([True])
480482

481-
left = naive_match(
483+
ref = naive_match(
482484
Q,
483485
T,
484486
excl_zone,
@@ -487,7 +489,7 @@ def test_match_isconstant(Q, T):
487489
Q_subseq_isconstant=Q_subseq_isconstant,
488490
)
489491

490-
right = match(
492+
cmp = match(
491493
Q,
492494
T,
493495
max_matches=None,
@@ -496,7 +498,7 @@ def test_match_isconstant(Q, T):
496498
Q_subseq_isconstant=Q_subseq_isconstant,
497499
)
498500

499-
npt.assert_almost_equal(left, right)
501+
npt.assert_allclose(cmp.astype(np.float64), ref.astype(np.float64), atol=1.5e-07)
500502

501503

502504
@pytest.mark.parametrize("Q, T", test_data)
@@ -505,7 +507,7 @@ def test_match_mean_stddev_isconstant(Q, T):
505507
excl_zone = int(np.ceil(m / 4))
506508
max_distance = 0.3
507509

508-
left = naive_match(
510+
ref = naive_match(
509511
Q,
510512
T,
511513
excl_zone,
@@ -515,7 +517,7 @@ def test_match_mean_stddev_isconstant(Q, T):
515517
T_subseq_isconstant = naive.rolling_isconstant(T, m)
516518
M_T, Σ_T = naive.compute_mean_std(T, len(Q))
517519

518-
right = match(
520+
cmp = match(
519521
Q,
520522
T,
521523
M_T,
@@ -525,7 +527,7 @@ def test_match_mean_stddev_isconstant(Q, T):
525527
T_subseq_isconstant=T_subseq_isconstant,
526528
)
527529

528-
npt.assert_almost_equal(left, right)
530+
npt.assert_allclose(cmp.astype(np.float64), ref.astype(np.float64), atol=1.5e-07)
529531

530532

531533
def test_multi_match():
@@ -536,21 +538,21 @@ def test_multi_match():
536538
excl_zone = int(np.ceil(m / 4))
537539
max_distance = 0.3
538540

539-
left = naive_multi_match(
541+
ref = naive_multi_match(
540542
Q,
541543
T,
542544
excl_zone,
543545
max_distance=max_distance,
544546
)
545547

546-
right = match(
548+
cmp = match(
547549
Q,
548550
T,
549551
max_matches=None,
550552
max_distance=lambda D: max_distance, # also test lambda functionality
551553
)
552554

553-
npt.assert_almost_equal(left, right)
555+
npt.assert_allclose(cmp.astype(np.float64), ref.astype(np.float64), atol=1.5e-07)
554556

555557

556558
def test_multi_match_isconstant():
@@ -572,7 +574,7 @@ def test_multi_match_isconstant():
572574
]
573575
)
574576

575-
left = naive_multi_match(
577+
ref = naive_multi_match(
576578
Q,
577579
T,
578580
excl_zone,
@@ -581,7 +583,7 @@ def test_multi_match_isconstant():
581583
Q_subseq_isconstant=Q_subseq_isconstant,
582584
)
583585

584-
right = match(
586+
cmp = match(
585587
Q,
586588
T,
587589
max_matches=None,
@@ -590,7 +592,7 @@ def test_multi_match_isconstant():
590592
Q_subseq_isconstant=Q_subseq_isconstant,
591593
)
592594

593-
npt.assert_almost_equal(left, right)
595+
npt.assert_allclose(cmp.astype(np.float64), ref.astype(np.float64), atol=1.5e-07)
594596

595597

596598
def test_motifs():
@@ -608,7 +610,7 @@ def test_motifs():
608610

609611
# performant
610612
mp = naive.stump(T, m, row_wise=True)
611-
comp_distance, comp_indices = motifs(
613+
cmp_distance, cmp_indices = motifs(
612614
T,
613615
mp[:, 0].astype(np.float64),
614616
min_neighbors=1,
@@ -618,8 +620,8 @@ def test_motifs():
618620
max_motifs=max_motifs,
619621
)
620622

621-
npt.assert_almost_equal(ref_indices, comp_indices)
622-
npt.assert_almost_equal(ref_distances, comp_distance)
623+
npt.assert_allclose(cmp_indices, ref_indices, atol=1.5e-07)
624+
npt.assert_allclose(cmp_distance, ref_distances, atol=1.5e-07)
623625

624626

625627
def test_motifs_with_isconstant():
@@ -643,7 +645,7 @@ def test_motifs_with_isconstant():
643645

644646
# performant
645647
mp = naive.stump(T, m, row_wise=True, T_A_subseq_isconstant=isconstant_custom_func)
646-
comp_distance, comp_indices = motifs(
648+
cmp_distance, cmp_indices = motifs(
647649
T,
648650
mp[:, 0].astype(np.float64),
649651
min_neighbors=1,
@@ -654,8 +656,8 @@ def test_motifs_with_isconstant():
654656
T_subseq_isconstant=isconstant_custom_func,
655657
)
656658

657-
npt.assert_almost_equal(ref_distances, comp_distance)
658-
npt.assert_almost_equal(ref_indices, comp_indices)
659+
npt.assert_allclose(cmp_distance, ref_distances, atol=1.5e-07)
660+
npt.assert_allclose(cmp_indices, ref_indices, atol=1.5e-07)
659661

660662

661663
def test_motifs_with_max_matches_none():
@@ -669,7 +671,7 @@ def test_motifs_with_max_matches_none():
669671

670672
# performant
671673
mp = naive.stump(T, m, row_wise=True)
672-
comp_distance, comp_indices = motifs(
674+
cmp_distance, cmp_indices = motifs(
673675
T,
674676
mp[:, 0].astype(np.float64),
675677
min_neighbors=1,
@@ -681,5 +683,5 @@ def test_motifs_with_max_matches_none():
681683

682684
ref_len = len(T) - m + 1
683685

684-
npt.assert_(ref_len >= comp_distance.shape[1])
685-
npt.assert_(ref_len >= comp_indices.shape[1])
686+
npt.assert_(ref_len >= cmp_distance.shape[1])
687+
npt.assert_(ref_len >= cmp_indices.shape[1])

0 commit comments

Comments
 (0)