@@ -553,8 +553,18 @@ def position_data_non_uniform_time():
553553class 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