Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
232 changes: 180 additions & 52 deletions tests/test_mparray.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,70 +47,198 @@ def test_mparray_self_join(T_A, T_B, k):
zone = int(np.ceil(m / 4))

ref_mp = naive.stump(T_B, m, exclusion_zone=zone, k=k)
comp_mp = stump(T_B, m, ignore_trivial=True, k=k)
cmp_mp = stump(T_B, m, ignore_trivial=True, k=k)
naive.replace_inf(ref_mp)
naive.replace_inf(comp_mp)
npt.assert_almost_equal(np.squeeze(ref_mp[:, :k]), comp_mp.P_)
npt.assert_almost_equal(np.squeeze(ref_mp[:, k : 2 * k]), comp_mp.I_)
npt.assert_almost_equal(np.squeeze(ref_mp[:, 2 * k]), comp_mp.left_I_)
npt.assert_almost_equal(np.squeeze(ref_mp[:, 2 * k + 1]), comp_mp.right_I_)

comp_mp = stump(pd.Series(T_B), m, ignore_trivial=True, k=k)
naive.replace_inf(comp_mp)
npt.assert_almost_equal(np.squeeze(ref_mp[:, :k]), comp_mp.P_)
npt.assert_almost_equal(np.squeeze(ref_mp[:, k : 2 * k]), comp_mp.I_)
npt.assert_almost_equal(np.squeeze(ref_mp[:, 2 * k]), comp_mp.left_I_)
npt.assert_almost_equal(np.squeeze(ref_mp[:, 2 * k + 1]), comp_mp.right_I_)
naive.replace_inf(cmp_mp)
npt.assert_allclose(
cmp_mp.P_.astype(np.float64),
np.squeeze(ref_mp[:, :k]).astype(np.float64),
atol=1.5e-07,
)
npt.assert_allclose(
cmp_mp.I_.astype(np.float64),
np.squeeze(ref_mp[:, k : 2 * k]).astype(np.float64),
atol=1.5e-07,
)
npt.assert_allclose(
cmp_mp.left_I_.astype(np.float64),
np.squeeze(ref_mp[:, 2 * k]).astype(np.float64),
atol=1.5e-07,
)
npt.assert_allclose(
cmp_mp.right_I_.astype(np.float64),
np.squeeze(ref_mp[:, 2 * k + 1]).astype(np.float64),
atol=1.5e-07,
)

cmp_mp = stump(pd.Series(T_B), m, ignore_trivial=True, k=k)
naive.replace_inf(cmp_mp)
npt.assert_allclose(
cmp_mp.P_.astype(np.float64),
np.squeeze(ref_mp[:, :k]).astype(np.float64),
atol=1.5e-07,
)
npt.assert_allclose(
cmp_mp.I_.astype(np.float64),
np.squeeze(ref_mp[:, k : 2 * k]).astype(np.float64),
atol=1.5e-07,
)
npt.assert_allclose(
cmp_mp.left_I_.astype(np.float64),
np.squeeze(ref_mp[:, 2 * k]).astype(np.float64),
atol=1.5e-07,
)
npt.assert_allclose(
cmp_mp.right_I_.astype(np.float64),
np.squeeze(ref_mp[:, 2 * k + 1]).astype(np.float64),
atol=1.5e-07,
)

ref_mp = naive.aamp(T_B, m, exclusion_zone=zone, k=k)
comp_mp = aamp(T_B, m, ignore_trivial=True, k=k)
cmp_mp = aamp(T_B, m, ignore_trivial=True, k=k)
naive.replace_inf(ref_mp)
naive.replace_inf(comp_mp)
npt.assert_almost_equal(np.squeeze(ref_mp[:, :k]), comp_mp.P_)
npt.assert_almost_equal(np.squeeze(ref_mp[:, k : 2 * k]), comp_mp.I_)
npt.assert_almost_equal(np.squeeze(ref_mp[:, 2 * k]), comp_mp.left_I_)
npt.assert_almost_equal(np.squeeze(ref_mp[:, 2 * k + 1]), comp_mp.right_I_)

comp_mp = aamp(pd.Series(T_B), m, ignore_trivial=True, k=k)
naive.replace_inf(comp_mp)
npt.assert_almost_equal(np.squeeze(ref_mp[:, :k]), comp_mp.P_)
npt.assert_almost_equal(np.squeeze(ref_mp[:, k : 2 * k]), comp_mp.I_)
npt.assert_almost_equal(np.squeeze(ref_mp[:, 2 * k]), comp_mp.left_I_)
npt.assert_almost_equal(np.squeeze(ref_mp[:, 2 * k + 1]), comp_mp.right_I_)
naive.replace_inf(cmp_mp)
npt.assert_allclose(
cmp_mp.P_.astype(np.float64),
np.squeeze(ref_mp[:, :k]).astype(np.float64),
atol=1.5e-07,
)
npt.assert_allclose(
cmp_mp.I_.astype(np.float64),
np.squeeze(ref_mp[:, k : 2 * k]).astype(np.float64),
atol=1.5e-07,
)
npt.assert_allclose(
cmp_mp.left_I_.astype(np.float64),
np.squeeze(ref_mp[:, 2 * k]).astype(np.float64),
atol=1.5e-07,
)
npt.assert_allclose(
cmp_mp.right_I_.astype(np.float64),
np.squeeze(ref_mp[:, 2 * k + 1]).astype(np.float64),
atol=1.5e-07,
)

cmp_mp = aamp(pd.Series(T_B), m, ignore_trivial=True, k=k)
naive.replace_inf(cmp_mp)
npt.assert_allclose(
cmp_mp.P_.astype(np.float64),
np.squeeze(ref_mp[:, :k]).astype(np.float64),
atol=1.5e-07,
)
npt.assert_allclose(
cmp_mp.I_.astype(np.float64),
np.squeeze(ref_mp[:, k : 2 * k]).astype(np.float64),
atol=1.5e-07,
)
npt.assert_allclose(
cmp_mp.left_I_.astype(np.float64),
np.squeeze(ref_mp[:, 2 * k]).astype(np.float64),
atol=1.5e-07,
)
npt.assert_allclose(
cmp_mp.right_I_.astype(np.float64),
np.squeeze(ref_mp[:, 2 * k + 1]).astype(np.float64),
atol=1.5e-07,
)


@pytest.mark.parametrize("T_A, T_B", test_data)
@pytest.mark.parametrize("k", kNN)
def test_mparray_A_B_join(T_A, T_B, k):
m = 3
ref_mp = naive.stump(T_A, m, T_B=T_B, k=k)
comp_mp = stump(T_A, m, T_B, ignore_trivial=False, k=k)
cmp_mp = stump(T_A, m, T_B, ignore_trivial=False, k=k)
naive.replace_inf(ref_mp)
naive.replace_inf(comp_mp)
npt.assert_almost_equal(np.squeeze(ref_mp[:, :k]), comp_mp.P_)
npt.assert_almost_equal(np.squeeze(ref_mp[:, k : 2 * k]), comp_mp.I_)
npt.assert_almost_equal(np.squeeze(ref_mp[:, 2 * k]), comp_mp.left_I_)
npt.assert_almost_equal(np.squeeze(ref_mp[:, 2 * k + 1]), comp_mp.right_I_)

comp_mp = stump(pd.Series(T_A), m, pd.Series(T_B), ignore_trivial=False, k=k)
naive.replace_inf(comp_mp)
npt.assert_almost_equal(np.squeeze(ref_mp[:, :k]), comp_mp.P_)
npt.assert_almost_equal(np.squeeze(ref_mp[:, k : 2 * k]), comp_mp.I_)
npt.assert_almost_equal(np.squeeze(ref_mp[:, 2 * k]), comp_mp.left_I_)
npt.assert_almost_equal(np.squeeze(ref_mp[:, 2 * k + 1]), comp_mp.right_I_)
naive.replace_inf(cmp_mp)
npt.assert_allclose(
cmp_mp.P_.astype(np.float64),
np.squeeze(ref_mp[:, :k]).astype(np.float64),
atol=1.5e-07,
)
npt.assert_allclose(
cmp_mp.I_.astype(np.float64),
np.squeeze(ref_mp[:, k : 2 * k]).astype(np.float64),
atol=1.5e-07,
)
npt.assert_allclose(
cmp_mp.left_I_.astype(np.float64),
np.squeeze(ref_mp[:, 2 * k]).astype(np.float64),
atol=1.5e-07,
)
npt.assert_allclose(
cmp_mp.right_I_.astype(np.float64),
np.squeeze(ref_mp[:, 2 * k + 1]).astype(np.float64),
atol=1.5e-07,
)

cmp_mp = stump(pd.Series(T_A), m, pd.Series(T_B), ignore_trivial=False, k=k)
naive.replace_inf(cmp_mp)
npt.assert_allclose(
cmp_mp.P_.astype(np.float64),
np.squeeze(ref_mp[:, :k]).astype(np.float64),
atol=1.5e-07,
)
npt.assert_allclose(
cmp_mp.I_.astype(np.float64),
np.squeeze(ref_mp[:, k : 2 * k]).astype(np.float64),
atol=1.5e-07,
)
npt.assert_allclose(
cmp_mp.left_I_.astype(np.float64),
np.squeeze(ref_mp[:, 2 * k]).astype(np.float64),
atol=1.5e-07,
)
npt.assert_allclose(
cmp_mp.right_I_.astype(np.float64),
np.squeeze(ref_mp[:, 2 * k + 1]).astype(np.float64),
atol=1.5e-07,
)

ref_mp = naive.aamp(T_A, m, T_B=T_B, k=k)
comp_mp = aamp(T_A, m, T_B, ignore_trivial=False, k=k)
cmp_mp = aamp(T_A, m, T_B, ignore_trivial=False, k=k)
naive.replace_inf(ref_mp)
naive.replace_inf(comp_mp)
npt.assert_almost_equal(np.squeeze(ref_mp[:, :k]), comp_mp.P_)
npt.assert_almost_equal(np.squeeze(ref_mp[:, k : 2 * k]), comp_mp.I_)
npt.assert_almost_equal(np.squeeze(ref_mp[:, 2 * k]), comp_mp.left_I_)
npt.assert_almost_equal(np.squeeze(ref_mp[:, 2 * k + 1]), comp_mp.right_I_)

comp_mp = aamp(pd.Series(T_A), m, pd.Series(T_B), ignore_trivial=False, k=k)
naive.replace_inf(comp_mp)
npt.assert_almost_equal(np.squeeze(ref_mp[:, :k]), comp_mp.P_)
npt.assert_almost_equal(np.squeeze(ref_mp[:, k : 2 * k]), comp_mp.I_)
npt.assert_almost_equal(np.squeeze(ref_mp[:, 2 * k]), comp_mp.left_I_)
npt.assert_almost_equal(np.squeeze(ref_mp[:, 2 * k + 1]), comp_mp.right_I_)
naive.replace_inf(cmp_mp)
npt.assert_allclose(
cmp_mp.P_.astype(np.float64),
np.squeeze(ref_mp[:, :k]).astype(np.float64),
atol=1.5e-07,
)
npt.assert_allclose(
cmp_mp.I_.astype(np.float64),
np.squeeze(ref_mp[:, k : 2 * k]).astype(np.float64),
atol=1.5e-07,
)
npt.assert_allclose(
cmp_mp.left_I_.astype(np.float64),
np.squeeze(ref_mp[:, 2 * k]).astype(np.float64),
atol=1.5e-07,
)
npt.assert_allclose(
cmp_mp.right_I_.astype(np.float64),
np.squeeze(ref_mp[:, 2 * k + 1]).astype(np.float64),
atol=1.5e-07,
)

cmp_mp = aamp(pd.Series(T_A), m, pd.Series(T_B), ignore_trivial=False, k=k)
naive.replace_inf(cmp_mp)
npt.assert_allclose(
cmp_mp.P_.astype(np.float64),
np.squeeze(ref_mp[:, :k]).astype(np.float64),
atol=1.5e-07,
)
npt.assert_allclose(
cmp_mp.I_.astype(np.float64),
np.squeeze(ref_mp[:, k : 2 * k]).astype(np.float64),
atol=1.5e-07,
)
npt.assert_allclose(
cmp_mp.left_I_.astype(np.float64),
np.squeeze(ref_mp[:, 2 * k]).astype(np.float64),
atol=1.5e-07,
)
npt.assert_allclose(
cmp_mp.right_I_.astype(np.float64),
np.squeeze(ref_mp[:, 2 * k + 1]).astype(np.float64),
atol=1.5e-07,
)
Loading