Skip to content

Commit 52d8477

Browse files
committed
test(collective): consolidate related polarization tests to reduce redundancy
1 parent b0f48d4 commit 52d8477

1 file changed

Lines changed: 51 additions & 41 deletions

File tree

tests/test_unit/test_kinematics/test_collective.py

Lines changed: 51 additions & 41 deletions
Original file line numberDiff line numberDiff line change
@@ -553,8 +553,18 @@ def position_data_non_uniform_time():
553553
class TestComputePolarization:
554554
"""Test suite for the compute_polarization function."""
555555

556-
def test_polarization_aligned(self, position_data_aligned_individuals):
557-
"""Test polarization is 1.0 when all move same direction."""
556+
def test_polarization_aligned(
557+
self,
558+
position_data_aligned_individuals,
559+
position_data_diagonal_movement,
560+
):
561+
"""Test polarization is 1.0 when all move same direction.
562+
563+
Tests both horizontal and diagonal movement to verify that
564+
polarization is rotation-invariant (direction angle doesn't matter,
565+
only alignment between individuals).
566+
"""
567+
# Test horizontal alignment
558568
polarization = kinematics.compute_polarization(
559569
position_data_aligned_individuals
560570
)
@@ -569,6 +579,13 @@ def test_polarization_aligned(self, position_data_aligned_individuals):
569579
# (Skip first time point since velocity is computed via diff)
570580
assert np.allclose(polarization.values[1:], 1.0, atol=1e-10)
571581

582+
# Test diagonal alignment (rotation invariance)
583+
# Both individuals moving at 45 degrees should also yield pol=1.0
584+
polarization_diag = kinematics.compute_polarization(
585+
position_data_diagonal_movement
586+
)
587+
assert np.allclose(polarization_diag.values[1:], 1.0, atol=1e-10)
588+
572589
def test_polarization_opposite(self, position_data_opposite_individuals):
573590
"""Test polarization is 0.0 when individuals move opposite."""
574591
polarization = kinematics.compute_polarization(
@@ -604,17 +621,6 @@ def test_polarization_handles_nan(self, position_data_with_nan):
604621
# The frame with NaN should exclude that individual from calculation
605622
assert not np.all(np.isnan(polarization.values))
606623

607-
def test_polarization_range(self, position_data_aligned_individuals):
608-
"""Test that polarization values are in [0, 1] range."""
609-
polarization = kinematics.compute_polarization(
610-
position_data_aligned_individuals
611-
)
612-
613-
# Exclude NaN values from range check
614-
valid_values = polarization.values[~np.isnan(polarization.values)]
615-
assert np.all(valid_values >= 0.0)
616-
assert np.all(valid_values <= 1.0)
617-
618624
def test_invalid_input_type(self, position_data_aligned_individuals):
619625
"""Test that non-DataArray input raises TypeError."""
620626
with pytest.raises(TypeError, match="must be an xarray.DataArray"):
@@ -680,18 +686,6 @@ def test_polarization_partial_alignment(
680686
# Compare frames 1: avoid boundary differencing dependence at t=0.
681687
assert np.allclose(polarization.values[1:], expected, atol=1e-10)
682688

683-
def test_polarization_diagonal_movement(
684-
self, position_data_diagonal_movement
685-
):
686-
"""Test polarization with diagonal movement remains 1.0."""
687-
polarization = kinematics.compute_polarization(
688-
position_data_diagonal_movement
689-
)
690-
691-
# Both moving in same diagonal direction -> polarization = 1.0
692-
# Compare frames 1: avoid boundary differencing dependence at t=0.
693-
assert np.allclose(polarization.values[1:], 1.0, atol=1e-10)
694-
695689
# ==================== Edge Cases ====================
696690

697691
def test_polarization_single_individual(
@@ -809,25 +803,23 @@ def test_polarization_non_uniform_time(
809803

810804
# ==================== Output Properties ====================
811805

812-
def test_polarization_output_shape(
806+
def test_polarization_output_structure(
813807
self, position_data_aligned_individuals
814808
):
815-
"""Test that output has correct shape (time only)."""
809+
"""Test that output has correct structure (time dimension only).
810+
811+
Verifies both positive assertion (dims == time) and explicit
812+
absence of input dimensions that should be reduced over.
813+
"""
816814
polarization = kinematics.compute_polarization(
817815
position_data_aligned_individuals
818816
)
819817

818+
# Positive assertion: output has exactly time dimension
820819
assert polarization.dims == ("time",)
821820
assert len(polarization) == len(position_data_aligned_individuals.time)
822821

823-
def test_polarization_output_no_extra_dims(
824-
self, position_data_aligned_individuals
825-
):
826-
"""Test that output doesn't have keypoints or space dims."""
827-
polarization = kinematics.compute_polarization(
828-
position_data_aligned_individuals
829-
)
830-
822+
# Explicit absence checks (documents which dims are reduced)
831823
assert "keypoints" not in polarization.dims
832824
assert "space" not in polarization.dims
833825
assert "individuals" not in polarization.dims
@@ -912,8 +904,24 @@ def test_polarization_symmetry(self):
912904

913905
np.testing.assert_array_almost_equal(pol1.values, pol2.values)
914906

915-
def test_polarization_bounds_random_directions(self):
916-
"""Test polarization stays in [0, 1] with random-ish directions."""
907+
def test_polarization_bounds(self, position_data_aligned_individuals):
908+
"""Test polarization values are always in [0, 1] range.
909+
910+
Verifies bounds with both:
911+
1. Simple aligned data (deterministic, yields boundary value 1.0)
912+
2. Random directions (stochastic, yields distribution across range)
913+
"""
914+
# Test with simple aligned data (boundary case: all 1.0)
915+
polarization_simple = kinematics.compute_polarization(
916+
position_data_aligned_individuals
917+
)
918+
valid_simple = polarization_simple.values[
919+
~np.isnan(polarization_simple.values)
920+
]
921+
assert np.all(valid_simple >= 0.0)
922+
assert np.all(valid_simple <= 1.0)
923+
924+
# Test with random directions (interior values)
917925
time = [0, 1, 2, 3, 4]
918926
individuals = [f"id_{i}" for i in range(10)]
919927
keypoints = ["centroid"]
@@ -948,11 +956,13 @@ def test_polarization_bounds_random_directions(self):
948956
},
949957
)
950958

951-
polarization = kinematics.compute_polarization(da)
959+
polarization_random = kinematics.compute_polarization(da)
952960

953-
valid_values = polarization.values[~np.isnan(polarization.values)]
954-
assert np.all(valid_values >= 0.0)
955-
assert np.all(valid_values <= 1.0)
961+
valid_random = polarization_random.values[
962+
~np.isnan(polarization_random.values)
963+
]
964+
assert np.all(valid_random >= 0.0)
965+
assert np.all(valid_random <= 1.0)
956966

957967
# ==================== NaN Handling Edge Cases ====================
958968

0 commit comments

Comments
 (0)