diff --git a/src/pyrecest/_backend/numpy/linalg.py b/src/pyrecest/_backend/numpy/linalg.py index 5cc1e7302..d0c396c9f 100644 --- a/src/pyrecest/_backend/numpy/linalg.py +++ b/src/pyrecest/_backend/numpy/linalg.py @@ -20,17 +20,19 @@ expm, ) +from .._shared_numpy.linalg import fractional_matrix_power as _fractional_matrix_power from .._shared_numpy.linalg import ( - fractional_matrix_power as _fractional_matrix_power, is_single_matrix_pd, - logm as _logm, +) +from .._shared_numpy.linalg import logm as _logm +from .._shared_numpy.linalg import ( polar, qr, quadratic_assignment, solve, solve_sylvester, - sqrtm as _sqrtm, ) +from .._shared_numpy.linalg import sqrtm as _sqrtm def _empty_zero_by_zero_matrix_result(value): diff --git a/src/pyrecest/_backend/pytorch/random.py b/src/pyrecest/_backend/pytorch/random.py index 54340db9d..eaf5728a5 100644 --- a/src/pyrecest/_backend/pytorch/random.py +++ b/src/pyrecest/_backend/pytorch/random.py @@ -267,9 +267,7 @@ def _validate_randint_array_dtype_bounds(low, high, dtype): # representable by the output dtype, as in randint(255, 256, dtype=uint8). # For int64, input tensors cannot represent max + 1, so every accepted high # value is already within the valid endpoint range. - if dtype != _torch.int64 and bool( - _torch.any(high_int64 > dtype_info.max + 1) - ): + if dtype != _torch.int64 and bool(_torch.any(high_int64 > dtype_info.max + 1)): raise ValueError(f"high is out of bounds for {dtype_name}") diff --git a/src/pyrecest/distributions/abstract_custom_distribution.py b/src/pyrecest/distributions/abstract_custom_distribution.py index 7fa21201e..d3a596f8b 100644 --- a/src/pyrecest/distributions/abstract_custom_distribution.py +++ b/src/pyrecest/distributions/abstract_custom_distribution.py @@ -10,7 +10,6 @@ from .abstract_distribution_type import AbstractDistributionType - _INVALID_INTEGRAL_TYPES = ( bool, np.bool_, diff --git a/src/pyrecest/distributions/cart_prod/hypercylindrical_state_space_subdivision_distribution.py b/src/pyrecest/distributions/cart_prod/hypercylindrical_state_space_subdivision_distribution.py index 82cd2a09e..3c63bbe58 100644 --- a/src/pyrecest/distributions/cart_prod/hypercylindrical_state_space_subdivision_distribution.py +++ b/src/pyrecest/distributions/cart_prod/hypercylindrical_state_space_subdivision_distribution.py @@ -38,7 +38,9 @@ from pyrecest.distributions.nonperiodic.custom_linear_distribution import ( CustomLinearDistribution, ) -from pyrecest.distributions.nonperiodic.gaussian_distribution import GaussianDistribution +from pyrecest.distributions.nonperiodic.gaussian_distribution import ( + GaussianDistribution, +) from pyrecest.distributions.nonperiodic.linear_mixture import LinearMixture diff --git a/src/pyrecest/distributions/hypersphere_subset/von_mises_fisher_distribution.py b/src/pyrecest/distributions/hypersphere_subset/von_mises_fisher_distribution.py index e5b96943b..fabc78a7b 100644 --- a/src/pyrecest/distributions/hypersphere_subset/von_mises_fisher_distribution.py +++ b/src/pyrecest/distributions/hypersphere_subset/von_mises_fisher_distribution.py @@ -33,7 +33,6 @@ from .abstract_hyperspherical_distribution import AbstractHypersphericalDistribution - _INVALID_REAL_SCALAR_TYPES = ( bool, np.bool_, diff --git a/src/pyrecest/distributions/hypersphere_subset/watson_distribution.py b/src/pyrecest/distributions/hypersphere_subset/watson_distribution.py index f4dcebb91..36e288612 100644 --- a/src/pyrecest/distributions/hypersphere_subset/watson_distribution.py +++ b/src/pyrecest/distributions/hypersphere_subset/watson_distribution.py @@ -29,7 +29,6 @@ from .abstract_hyperspherical_distribution import AbstractHypersphericalDistribution from .bingham_distribution import BinghamDistribution - _INVALID_REAL_SCALAR_TYPES = ( bool, np.bool_, diff --git a/src/pyrecest/evaluation/check_and_fix_config.py b/src/pyrecest/evaluation/check_and_fix_config.py index d90cb058a..b4429a4c0 100644 --- a/src/pyrecest/evaluation/check_and_fix_config.py +++ b/src/pyrecest/evaluation/check_and_fix_config.py @@ -1,7 +1,6 @@ from numbers import Integral, Real import numpy as np - from pyrecest.distributions import AbstractManifoldSpecificDistribution @@ -12,9 +11,9 @@ def _is_integer_count(value): def _validate_probability(value, name): - if isinstance(value, (bool, np.bool_, np.datetime64, np.timedelta64)) or not isinstance( - value, Real - ): + if isinstance( + value, (bool, np.bool_, np.datetime64, np.timedelta64) + ) or not isinstance(value, Real): raise TypeError(f"{name} must be a real scalar") value = float(value) if not np.isfinite(value) or not 0.0 <= value <= 1.0: diff --git a/src/pyrecest/evaluation/get_distance_function.py b/src/pyrecest/evaluation/get_distance_function.py index 44a2cb340..1e0b32515 100644 --- a/src/pyrecest/evaluation/get_distance_function.py +++ b/src/pyrecest/evaluation/get_distance_function.py @@ -164,9 +164,7 @@ def distance_function(xest, xtrue): return distance_function -def _target_matrix_candidates( - value, name: str -) -> list[tuple[numpy.ndarray, int]]: +def _target_matrix_candidates(value, name: str) -> list[tuple[numpy.ndarray, int]]: value = _as_real_numeric_array(value, name) if value.ndim not in (1, 2): raise ValueError(f"{name} must be a one- or two-dimensional target set") diff --git a/src/pyrecest/filters/global_nearest_neighbor.py b/src/pyrecest/filters/global_nearest_neighbor.py index 82f7b0f9e..7fd108c49 100644 --- a/src/pyrecest/filters/global_nearest_neighbor.py +++ b/src/pyrecest/filters/global_nearest_neighbor.py @@ -22,7 +22,6 @@ from .abstract_nearest_neighbor_tracker import AbstractNearestNeighborTracker - _INVALID_PAIRWISE_COST_SCALAR_TYPES = ( type(None), bool, diff --git a/src/pyrecest/filters/mode_rbpf_manifold_ukf_tracker.py b/src/pyrecest/filters/mode_rbpf_manifold_ukf_tracker.py index 18d94faa0..dc61fa613 100644 --- a/src/pyrecest/filters/mode_rbpf_manifold_ukf_tracker.py +++ b/src/pyrecest/filters/mode_rbpf_manifold_ukf_tracker.py @@ -876,9 +876,7 @@ def _normalize_probs(probs): ) total = float(np.sum(probs)) if total <= 0.0: - raise ValueError( - "initial_mode_probs must have positive total probability" - ) + raise ValueError("initial_mode_probs must have positive total probability") return probs / total @staticmethod diff --git a/src/pyrecest/filters/sequence_association.py b/src/pyrecest/filters/sequence_association.py index 6f62962c7..dc55ec467 100644 --- a/src/pyrecest/filters/sequence_association.py +++ b/src/pyrecest/filters/sequence_association.py @@ -360,9 +360,8 @@ def _validate_integer(value: object, name: str) -> int: raise ValueError(message) from exc if value_array.ndim != 0 or value_array.dtype == np.bool_: raise ValueError(message) - if ( - value_array.dtype.kind in {"S", "U", "c"} - or _is_temporal_scalar_array(value_array) + if value_array.dtype.kind in {"S", "U", "c"} or _is_temporal_scalar_array( + value_array ): raise ValueError(message) diff --git a/src/pyrecest/models/linear_gaussian.py b/src/pyrecest/models/linear_gaussian.py index 7ed22d4ae..c784cc2bc 100644 --- a/src/pyrecest/models/linear_gaussian.py +++ b/src/pyrecest/models/linear_gaussian.py @@ -2,8 +2,8 @@ from numbers import Complex, Integral, Real +from pyrecest.backend import all as backend_all from pyrecest.backend import ( - all as backend_all, asarray, ) from pyrecest.backend import copy as backend_copy diff --git a/src/pyrecest/smoothers/abstract_smoother.py b/src/pyrecest/smoothers/abstract_smoother.py index 867352fa3..64978d4d4 100644 --- a/src/pyrecest/smoothers/abstract_smoother.py +++ b/src/pyrecest/smoothers/abstract_smoother.py @@ -161,9 +161,7 @@ def _normalize_vector_sequence( # pylint: disable=too-many-return-statements expected_shape = (vector_dim,) shape_error = f"{name} must contain vectors with shape {expected_shape}." - if isinstance(values, (list, tuple)) and any( - value is None for value in values - ): + if isinstance(values, (list, tuple)) and any(value is None for value in values): values_arr = None else: try: diff --git a/src/pyrecest/utils/_point_set_registration_common.py b/src/pyrecest/utils/_point_set_registration_common.py index 46fdafacb..84405b9a2 100644 --- a/src/pyrecest/utils/_point_set_registration_common.py +++ b/src/pyrecest/utils/_point_set_registration_common.py @@ -25,7 +25,6 @@ from scipy.optimize import linear_sum_assignment from scipy.spatial.distance import cdist - _INVALID_REAL_SCALAR_TYPES = ( type(None), bool, diff --git a/src/pyrecest/utils/metrics.py b/src/pyrecest/utils/metrics.py index 0153d8e2a..9b869436a 100644 --- a/src/pyrecest/utils/metrics.py +++ b/src/pyrecest/utils/metrics.py @@ -579,11 +579,7 @@ def _as_covariance_stack( def _as_positive_int(value: Any, name: str) -> int: array = np.asarray(value) - if ( - array.ndim != 0 - or array.dtype == np.bool_ - or array.dtype.kind in {"M", "m"} - ): + if array.ndim != 0 or array.dtype == np.bool_ or array.dtype.kind in {"M", "m"}: raise ValueError(f"{name} must be a positive integer") scalar = array.item() if isinstance(scalar, (int, np.integer)) and not isinstance(scalar, bool): diff --git a/src/pyrecest/utils/point_set_registration.py b/src/pyrecest/utils/point_set_registration.py index 1a7d60b72..f291b5471 100644 --- a/src/pyrecest/utils/point_set_registration.py +++ b/src/pyrecest/utils/point_set_registration.py @@ -214,8 +214,10 @@ def _validate_positive_integer(value, name: str, *, minimum: int = 1) -> int: except TypeError: pass - if isinstance(value, np.ndarray) and value.shape == () and isinstance( - value.item(), temporal_types + if ( + isinstance(value, np.ndarray) + and value.shape == () + and isinstance(value.item(), temporal_types) ): raise ValueError(error_message) diff --git a/tests/backend/test_numpy_fractional_matrix_power_exponent_validation.py b/tests/backend/test_numpy_fractional_matrix_power_exponent_validation.py index d0a055781..299a69558 100644 --- a/tests/backend/test_numpy_fractional_matrix_power_exponent_validation.py +++ b/tests/backend/test_numpy_fractional_matrix_power_exponent_validation.py @@ -2,7 +2,6 @@ import pytest from pyrecest._backend.numpy import linalg - _MATRIX = np.diag([4.0, 9.0]) @@ -16,9 +15,7 @@ np.timedelta64(2, "ns"), np.datetime64("1970-01-01T00:00:00.000000002"), np.array(np.timedelta64(2, "ns"), dtype=object), - np.array( - np.datetime64("1970-01-01T00:00:00.000000002"), dtype=object - ), + np.array(np.datetime64("1970-01-01T00:00:00.000000002"), dtype=object), "0.5", 0.5 + 0.0j, ], diff --git a/tests/backend/test_numpy_linalg_zero_by_zero_matrix_functions.py b/tests/backend/test_numpy_linalg_zero_by_zero_matrix_functions.py index 92668c944..c96767d25 100644 --- a/tests/backend/test_numpy_linalg_zero_by_zero_matrix_functions.py +++ b/tests/backend/test_numpy_linalg_zero_by_zero_matrix_functions.py @@ -1,6 +1,5 @@ import numpy as np import pytest - from pyrecest._backend.numpy import linalg diff --git a/tests/backend/test_pytorch_randint_dtype_bounds.py b/tests/backend/test_pytorch_randint_dtype_bounds.py index d633ecfa4..3de40b1f3 100644 --- a/tests/backend/test_pytorch_randint_dtype_bounds.py +++ b/tests/backend/test_pytorch_randint_dtype_bounds.py @@ -1,7 +1,6 @@ import numpy as np import pytest - torch = pytest.importorskip("torch") from pyrecest._backend.pytorch import random # noqa: E402 @@ -16,9 +15,7 @@ ([0], [129], np.int8, "high is out of bounds for int8"), ], ) -def test_array_randint_rejects_bounds_outside_output_dtype( - low, high, dtype, message -): +def test_array_randint_rejects_bounds_outside_output_dtype(low, high, dtype, message): with pytest.raises(ValueError, match=message): random.randint(low, high, dtype=dtype) diff --git a/tests/backend/test_pytorch_random_multivariate_normal_keywords.py b/tests/backend/test_pytorch_random_multivariate_normal_keywords.py index f4d1982ee..4f4fd0e77 100644 --- a/tests/backend/test_pytorch_random_multivariate_normal_keywords.py +++ b/tests/backend/test_pytorch_random_multivariate_normal_keywords.py @@ -1,7 +1,6 @@ import numpy as np import pytest - torch = pytest.importorskip("torch") from pyrecest._backend.pytorch import random # noqa: E402 diff --git a/tests/distributions/test_abstract_mixture_temporal_weights.py b/tests/distributions/test_abstract_mixture_temporal_weights.py index e18ba4cb9..a6a3ccaf7 100644 --- a/tests/distributions/test_abstract_mixture_temporal_weights.py +++ b/tests/distributions/test_abstract_mixture_temporal_weights.py @@ -1,7 +1,6 @@ import unittest import numpy as np - from pyrecest.distributions.abstract_mixture import _validate_mixture_weight_values diff --git a/tests/distributions/test_complex_acg_sample_count_precision.py b/tests/distributions/test_complex_acg_sample_count_precision.py index be98466cb..53b99159a 100644 --- a/tests/distributions/test_complex_acg_sample_count_precision.py +++ b/tests/distributions/test_complex_acg_sample_count_precision.py @@ -2,7 +2,6 @@ from fractions import Fraction import numpy as np - from pyrecest.distributions.hypersphere_subset.complex_angular_central_gaussian_distribution import ( _validate_positive_sample_count, ) diff --git a/tests/distributions/test_complex_watson_temporal_sample_count.py b/tests/distributions/test_complex_watson_temporal_sample_count.py index 12dd3e549..ef8688604 100644 --- a/tests/distributions/test_complex_watson_temporal_sample_count.py +++ b/tests/distributions/test_complex_watson_temporal_sample_count.py @@ -1,11 +1,9 @@ import numpy as np -import pytest - import pyrecest.backend +import pytest from pyrecest.backend import array, complex128 from pyrecest.distributions import ComplexWatsonDistribution - pytestmark = pytest.mark.skipif( pyrecest.backend.__backend_name__ == "jax", reason="Complex Watson sampling is not supported on the JAX backend", diff --git a/tests/distributions/test_ellipsoidal_ball_temporal_sample_count.py b/tests/distributions/test_ellipsoidal_ball_temporal_sample_count.py index 0e5cd24c2..8a4eda059 100644 --- a/tests/distributions/test_ellipsoidal_ball_temporal_sample_count.py +++ b/tests/distributions/test_ellipsoidal_ball_temporal_sample_count.py @@ -1,6 +1,5 @@ import numpy as np import pytest - from pyrecest.backend import array, diag from pyrecest.distributions import EllipsoidalBallUniformDistribution diff --git a/tests/distributions/test_gaussian_marginalize_out_validation.py b/tests/distributions/test_gaussian_marginalize_out_validation.py index 9da6cbc21..510256396 100644 --- a/tests/distributions/test_gaussian_marginalize_out_validation.py +++ b/tests/distributions/test_gaussian_marginalize_out_validation.py @@ -1,7 +1,6 @@ import unittest import numpy as np - from pyrecest.backend import array from pyrecest.distributions import GaussianDistribution diff --git a/tests/distributions/test_gaussian_mixture_complex_mean.py b/tests/distributions/test_gaussian_mixture_complex_mean.py index 4604d27a1..48c3ba18d 100644 --- a/tests/distributions/test_gaussian_mixture_complex_mean.py +++ b/tests/distributions/test_gaussian_mixture_complex_mean.py @@ -11,12 +11,8 @@ class GaussianMixtureComplexMeanTest(unittest.TestCase): def test_set_mean_rejects_complex_target_without_mutating_components(self): - component_1 = GaussianDistribution( - array([0.0, 1.0]), diag(array([1.0, 2.0])) - ) - component_2 = GaussianDistribution( - array([2.0, 3.0]), diag(array([3.0, 4.0])) - ) + component_1 = GaussianDistribution(array([0.0, 1.0]), diag(array([1.0, 2.0]))) + component_2 = GaussianDistribution(array([2.0, 3.0]), diag(array([3.0, 4.0]))) mixture = GaussianMixture([component_1, component_2], array([0.25, 0.75])) original_means = [to_numpy(dist.mu).copy() for dist in mixture.dists] diff --git a/tests/distributions/test_hypercylindrical_zero_marginals.py b/tests/distributions/test_hypercylindrical_zero_marginals.py index d782d86cc..673982ab3 100644 --- a/tests/distributions/test_hypercylindrical_zero_marginals.py +++ b/tests/distributions/test_hypercylindrical_zero_marginals.py @@ -2,7 +2,6 @@ from math import pi import numpy as np - from pyrecest.backend import __backend_name__ as backend_name from pyrecest.backend import array from pyrecest.distributions.cart_prod.hypercylindrical_state_space_subdivision_distribution import ( diff --git a/tests/distributions/test_hypertoroidal_mixed_boolean_input_validation.py b/tests/distributions/test_hypertoroidal_mixed_boolean_input_validation.py index decf71c83..206131911 100644 --- a/tests/distributions/test_hypertoroidal_mixed_boolean_input_validation.py +++ b/tests/distributions/test_hypertoroidal_mixed_boolean_input_validation.py @@ -23,9 +23,7 @@ def test_rejects_mixed_boolean_evaluation_points(self): def test_numeric_python_sequences_remain_valid(self): npt.assert_allclose(as_shift_vector([0.0, 1.0], 2), [0.0, 1.0]) - npt.assert_allclose( - as_hypertoroidal_points([[0.0, 1.0]], 2), [[0.0, 1.0]] - ) + npt.assert_allclose(as_hypertoroidal_points([[0.0, 1.0]], 2), [[0.0, 1.0]]) if __name__ == "__main__": diff --git a/tests/distributions/test_hypertoroidal_python_scalar_integrand.py b/tests/distributions/test_hypertoroidal_python_scalar_integrand.py index a94be22f9..1c2077135 100644 --- a/tests/distributions/test_hypertoroidal_python_scalar_integrand.py +++ b/tests/distributions/test_hypertoroidal_python_scalar_integrand.py @@ -1,7 +1,6 @@ import numpy as np import numpy.testing as npt import pytest - from pyrecest import backend from pyrecest.distributions import AbstractHypertoroidalDistribution diff --git a/tests/distributions/test_linear_box_particle_temporal_count.py b/tests/distributions/test_linear_box_particle_temporal_count.py index e75c61e84..31e9ef685 100644 --- a/tests/distributions/test_linear_box_particle_temporal_count.py +++ b/tests/distributions/test_linear_box_particle_temporal_count.py @@ -1,6 +1,5 @@ import numpy as np import pytest - from pyrecest.backend import array from pyrecest.distributions.nonperiodic.linear_box_particle_distribution import ( LinearBoxParticleDistribution, diff --git a/tests/distributions/test_partially_wrapped_normal_hybrid_moment_order.py b/tests/distributions/test_partially_wrapped_normal_hybrid_moment_order.py index d1d045a31..15c01efee 100644 --- a/tests/distributions/test_partially_wrapped_normal_hybrid_moment_order.py +++ b/tests/distributions/test_partially_wrapped_normal_hybrid_moment_order.py @@ -2,7 +2,6 @@ import numpy as np import numpy.testing as npt - from pyrecest.backend import array from pyrecest.distributions.cart_prod.partially_wrapped_normal_distribution import ( PartiallyWrappedNormalDistribution, diff --git a/tests/distributions/test_piecewise_constant_interval_index_validation.py b/tests/distributions/test_piecewise_constant_interval_index_validation.py index 521852682..675aa5902 100644 --- a/tests/distributions/test_piecewise_constant_interval_index_validation.py +++ b/tests/distributions/test_piecewise_constant_interval_index_validation.py @@ -1,6 +1,5 @@ import numpy as np import pytest - from pyrecest.distributions.circle.piecewise_constant_distribution import ( PiecewiseConstantDistribution, ) diff --git a/tests/distributions/test_se2_dirac_temporal_particle_count.py b/tests/distributions/test_se2_dirac_temporal_particle_count.py index 20704ce01..8be4cda23 100644 --- a/tests/distributions/test_se2_dirac_temporal_particle_count.py +++ b/tests/distributions/test_se2_dirac_temporal_particle_count.py @@ -1,6 +1,5 @@ import numpy as np import pytest - from pyrecest.distributions import SE2DiracDistribution from pyrecest.distributions.cart_prod.abstract_hypercylindrical_distribution import ( AbstractHypercylindricalDistribution, @@ -18,10 +17,14 @@ def sample(self, n): raise AssertionError("invalid temporal counts must not reach sampling") def marginalize_linear(self): - raise AssertionError("marginalization must not be evaluated for count validation") + raise AssertionError( + "marginalization must not be evaluated for count validation" + ) def marginalize_periodic(self): - raise AssertionError("marginalization must not be evaluated for count validation") + raise AssertionError( + "marginalization must not be evaluated for count validation" + ) @pytest.mark.parametrize( diff --git a/tests/distributions/test_so3_conversion_validation.py b/tests/distributions/test_so3_conversion_validation.py index 09ebb67f1..d3b6f9e69 100644 --- a/tests/distributions/test_so3_conversion_validation.py +++ b/tests/distributions/test_so3_conversion_validation.py @@ -16,7 +16,6 @@ SO3TangentGaussianDistribution, ) - _TEMPORAL_VALUES = ( np.timedelta64(3, "ns"), np.timedelta64(3, "us"), diff --git a/tests/distributions/test_von_mises_temporal_sample_count.py b/tests/distributions/test_von_mises_temporal_sample_count.py index 467b42c50..66d58a2ad 100644 --- a/tests/distributions/test_von_mises_temporal_sample_count.py +++ b/tests/distributions/test_von_mises_temporal_sample_count.py @@ -1,6 +1,5 @@ import numpy as np import pytest - from pyrecest.distributions import VonMisesDistribution diff --git a/tests/distributions/test_wrapped_normal_temporal_sample_counts.py b/tests/distributions/test_wrapped_normal_temporal_sample_counts.py index 1a697b548..fc39951a1 100644 --- a/tests/distributions/test_wrapped_normal_temporal_sample_counts.py +++ b/tests/distributions/test_wrapped_normal_temporal_sample_counts.py @@ -1,6 +1,5 @@ import numpy as np import pytest - from pyrecest.distributions import WrappedNormalDistribution diff --git a/tests/evaluation/test_integer_estimate_history.py b/tests/evaluation/test_integer_estimate_history.py index 82f48dde1..385f8379e 100644 --- a/tests/evaluation/test_integer_estimate_history.py +++ b/tests/evaluation/test_integer_estimate_history.py @@ -1,7 +1,6 @@ import numpy as np import numpy.testing as npt import pytest - from pyrecest.backend import array, get_backend_name from pyrecest.evaluation import perform_predict_update_cycles from pyrecest.evaluation.configure_for_filter import register_filter_factory diff --git a/tests/filters/test_euclidean_boxed_particle_filter_zero_support_likelihood.py b/tests/filters/test_euclidean_boxed_particle_filter_zero_support_likelihood.py index f1de6cd21..b4144dcfe 100644 --- a/tests/filters/test_euclidean_boxed_particle_filter_zero_support_likelihood.py +++ b/tests/filters/test_euclidean_boxed_particle_filter_zero_support_likelihood.py @@ -2,7 +2,6 @@ import numpy as np import numpy.testing as npt - from pyrecest.backend import array, to_numpy from pyrecest.distributions.nonperiodic.linear_dirac_distribution import ( LinearDiracDistribution, diff --git a/tests/filters/test_gnn_pairwise_cost_weight_validation.py b/tests/filters/test_gnn_pairwise_cost_weight_validation.py index a3be13f89..d2f3f2e6f 100644 --- a/tests/filters/test_gnn_pairwise_cost_weight_validation.py +++ b/tests/filters/test_gnn_pairwise_cost_weight_validation.py @@ -36,15 +36,13 @@ def test_rejects_invalid_pairwise_cost_weights(self): def test_accepts_finite_nonnegative_scalar_weights(self): for valid_weight in (0, 0.5, np.float64(2.0)): with self.subTest(pairwise_cost_weight=valid_weight): - validated_weight = ( - GlobalNearestNeighbor._validate_pairwise_cost_weight(valid_weight) + validated_weight = GlobalNearestNeighbor._validate_pairwise_cost_weight( + valid_weight ) self.assertEqual(validated_weight, float(valid_weight)) def test_zero_weight_ignores_positive_infinite_pairwise_gate(self): - tracker = GlobalNearestNeighbor( - association_param={"pairwise_cost_weight": 0.0} - ) + tracker = GlobalNearestNeighbor(association_param={"pairwise_cost_weight": 0.0}) geometric_costs = np.array([[1.25]]) combined_costs = tracker._apply_pairwise_cost_matrix( @@ -54,9 +52,7 @@ def test_zero_weight_ignores_positive_infinite_pairwise_gate(self): npt.assert_array_equal(combined_costs, geometric_costs) def test_positive_weight_scales_pairwise_costs(self): - tracker = GlobalNearestNeighbor( - association_param={"pairwise_cost_weight": 2.0} - ) + tracker = GlobalNearestNeighbor(association_param={"pairwise_cost_weight": 2.0}) combined_costs = tracker._apply_pairwise_cost_matrix( np.array([[1.0, 2.0]]), np.array([[3.0, 4.0]]) diff --git a/tests/filters/test_gnn_pairwise_object_cost_validation.py b/tests/filters/test_gnn_pairwise_object_cost_validation.py index 1956a965a..53222d0fb 100644 --- a/tests/filters/test_gnn_pairwise_object_cost_validation.py +++ b/tests/filters/test_gnn_pairwise_object_cost_validation.py @@ -78,12 +78,10 @@ def test_rejects_native_temporal_dtypes(self): ) def test_accepts_real_numeric_values(self): - pairwise_cost_matrix = ( - GlobalNearestNeighbor._validate_pairwise_cost_matrix( - [[1, 2.5]], - 1, - 2, - ) + pairwise_cost_matrix = GlobalNearestNeighbor._validate_pairwise_cost_matrix( + [[1, 2.5]], + 1, + 2, ) np.testing.assert_allclose(pairwise_cost_matrix, [[1.0, 2.5]]) diff --git a/tests/filters/test_measurement_reliability_temporal_counts.py b/tests/filters/test_measurement_reliability_temporal_counts.py index 2abed5d82..4ff98c2a3 100644 --- a/tests/filters/test_measurement_reliability_temporal_counts.py +++ b/tests/filters/test_measurement_reliability_temporal_counts.py @@ -1,7 +1,6 @@ import unittest import numpy as np - from pyrecest.filters import ( normalize_active_measurement_mask, normalize_measurement_noise_covariances, diff --git a/tests/filters/test_mode_rbpf_manifold_ukf_tracker.py b/tests/filters/test_mode_rbpf_manifold_ukf_tracker.py index bd6e3c87a..ca3e7521e 100644 --- a/tests/filters/test_mode_rbpf_manifold_ukf_tracker.py +++ b/tests/filters/test_mode_rbpf_manifold_ukf_tracker.py @@ -161,7 +161,6 @@ def test_update_without_measurements_is_noop(self): npt.assert_allclose(posterior, prior) - def test_preserves_impossible_mode_probabilities(self): transition_matrix = np.eye(3) tracker = self.make_tracker( diff --git a/tests/filters/test_nonadditive_arraylike_samples.py b/tests/filters/test_nonadditive_arraylike_samples.py index ae6d5c316..67907c49e 100644 --- a/tests/filters/test_nonadditive_arraylike_samples.py +++ b/tests/filters/test_nonadditive_arraylike_samples.py @@ -1,5 +1,4 @@ import numpy.testing as npt - from pyrecest.backend import array from pyrecest.distributions import LinearDiracDistribution from pyrecest.filters.euclidean_particle_filter import EuclideanParticleFilter @@ -7,9 +6,7 @@ def test_predict_nonlinear_nonadditive_accepts_array_like_samples_and_weights(): particle_filter = EuclideanParticleFilter(n_particles=3, dim=1) - particle_filter.filter_state = LinearDiracDistribution( - array([[0.0], [1.0], [2.0]]) - ) + particle_filter.filter_state = LinearDiracDistribution(array([[0.0], [1.0], [2.0]])) particle_filter.predict_nonlinear_nonadditive( lambda particle, noise: particle + noise, diff --git a/tests/filters/test_particle_filter_count_precision.py b/tests/filters/test_particle_filter_count_precision.py index def36a006..9dee320f7 100644 --- a/tests/filters/test_particle_filter_count_precision.py +++ b/tests/filters/test_particle_filter_count_precision.py @@ -35,9 +35,7 @@ def test_particle_filter_count_validation_is_exact(validator): np.timedelta64(3, "ns"), np.datetime64("1970-01-01T00:00:00.000000003"), np.array(np.timedelta64(3, "ns"), dtype=object), - np.array( - np.datetime64("1970-01-01T00:00:00.000000003"), dtype=object - ), + np.array(np.datetime64("1970-01-01T00:00:00.000000003"), dtype=object), ], ids=["timedelta", "datetime", "object-timedelta", "object-datetime"], ) diff --git a/tests/filters/test_relaxed_s3f_process_noise_validation.py b/tests/filters/test_relaxed_s3f_process_noise_validation.py index b8d1cdca0..a02a22037 100644 --- a/tests/filters/test_relaxed_s3f_process_noise_validation.py +++ b/tests/filters/test_relaxed_s3f_process_noise_validation.py @@ -28,7 +28,9 @@ def test_rejects_nonsymmetric_process_noise_without_mutating_state(self): self.assertTrue( bool( - (filter_.filter_state.linear_distributions[0].C == covariance_before).all() + ( + filter_.filter_state.linear_distributions[0].C == covariance_before + ).all() ) ) @@ -45,7 +47,9 @@ def test_rejects_indefinite_process_noise_without_mutating_state(self): self.assertTrue( bool( - (filter_.filter_state.linear_distributions[0].C == covariance_before).all() + ( + filter_.filter_state.linear_distributions[0].C == covariance_before + ).all() ) ) diff --git a/tests/filters/test_sequence_association_temporal_validation.py b/tests/filters/test_sequence_association_temporal_validation.py index 895b4fc60..a3fb56a45 100644 --- a/tests/filters/test_sequence_association_temporal_validation.py +++ b/tests/filters/test_sequence_association_temporal_validation.py @@ -6,7 +6,6 @@ solve_viterbi_sequence_association, ) - _TEMPORAL_VALUES = ( pytest.param(np.timedelta64(1, "ns"), id="timedelta-ns"), pytest.param(np.timedelta64(1, "us"), id="timedelta-us"), diff --git a/tests/models/test_linear_gaussian_finite_inputs.py b/tests/models/test_linear_gaussian_finite_inputs.py index 07c87cc39..0227adce6 100644 --- a/tests/models/test_linear_gaussian_finite_inputs.py +++ b/tests/models/test_linear_gaussian_finite_inputs.py @@ -15,27 +15,19 @@ def test_models_reject_nonfinite_system_and_measurement_matrices(self): for value in (np.nan, np.inf, -np.inf): with self.subTest(model="transition", value=value): with self.assertRaisesRegex(ValueError, "matrix.*finite"): - LinearGaussianTransitionModel( - array([[value]]), array([[1.0]]) - ) + LinearGaussianTransitionModel(array([[value]]), array([[1.0]])) with self.subTest(model="measurement", value=value): with self.assertRaisesRegex(ValueError, "matrix.*finite"): - LinearGaussianMeasurementModel( - array([[value]]), array([[1.0]]) - ) + LinearGaussianMeasurementModel(array([[value]]), array([[1.0]])) def test_models_reject_nonfinite_noise_covariances(self): for value in (np.nan, np.inf, -np.inf): with self.subTest(model="transition", value=value): with self.assertRaisesRegex(ValueError, "noise_cov.*finite"): - LinearGaussianTransitionModel( - array([[1.0]]), array([[value]]) - ) + LinearGaussianTransitionModel(array([[1.0]]), array([[value]])) with self.subTest(model="measurement", value=value): with self.assertRaisesRegex(ValueError, "noise_cov.*finite"): - LinearGaussianMeasurementModel( - array([[1.0]]), array([[value]]) - ) + LinearGaussianMeasurementModel(array([[1.0]]), array([[value]])) def test_transition_model_rejects_nonfinite_offset(self): for value in (np.nan, np.inf, -np.inf): @@ -55,12 +47,8 @@ def test_identity_models_reject_nonfinite_scalar_noise(self): IdentityGaussianMeasurementModel(1, value) def test_prediction_rejects_nonfinite_state_inputs(self): - transition = LinearGaussianTransitionModel( - array([[1.0]]), array([[1.0]]) - ) - measurement = LinearGaussianMeasurementModel( - array([[1.0]]), array([[1.0]]) - ) + transition = LinearGaussianTransitionModel(array([[1.0]]), array([[1.0]])) + measurement = LinearGaussianMeasurementModel(array([[1.0]]), array([[1.0]])) for value in (np.nan, np.inf, -np.inf): with self.subTest(method="transition mean", value=value): diff --git a/tests/test_deprecation_helper.py b/tests/test_deprecation_helper.py index ef851b7ed..7251d5a62 100644 --- a/tests/test_deprecation_helper.py +++ b/tests/test_deprecation_helper.py @@ -25,9 +25,9 @@ def test_deprecated_decorator_supports_partial_callables(): def add(left, right): return left + right - legacy_add_one = deprecated( - since="2.3.0", remove_in="3.0.0", replacement="add" - )(functools.partial(add, 1)) + legacy_add_one = deprecated(since="2.3.0", remove_in="3.0.0", replacement="add")( + functools.partial(add, 1) + ) with warnings.catch_warnings(record=True) as caught: warnings.simplefilter("always") diff --git a/tests/test_evidence_terminal_posterior_validation.py b/tests/test_evidence_terminal_posterior_validation.py index 0f0cf09aa..33f52d94b 100644 --- a/tests/test_evidence_terminal_posterior_validation.py +++ b/tests/test_evidence_terminal_posterior_validation.py @@ -1,5 +1,4 @@ import pytest - from pyrecest.evidence import EvidenceComputationMode diff --git a/tests/test_gaussian_sampler_zero_samples.py b/tests/test_gaussian_sampler_zero_samples.py index 22397d2ca..bd4d6b108 100644 --- a/tests/test_gaussian_sampler_zero_samples.py +++ b/tests/test_gaussian_sampler_zero_samples.py @@ -1,5 +1,4 @@ import numpy as np - from pyrecest.sampling.euclidean_sampler import GaussianSampler diff --git a/tests/test_leopardi_small_symmetric_partitions.py b/tests/test_leopardi_small_symmetric_partitions.py index 90aa4b35c..5a2a8376d 100644 --- a/tests/test_leopardi_small_symmetric_partitions.py +++ b/tests/test_leopardi_small_symmetric_partitions.py @@ -3,7 +3,6 @@ import pytest from pyrecest.sampling.leopardi_sampler import get_partition_points_cartesian - pytestmark = pytest.mark.skipif( pyrecest.backend.__backend_name__ == "jax", reason="Leopardi sampling uses SciPy root finding.", diff --git a/tests/test_metrics_temporal_counts.py b/tests/test_metrics_temporal_counts.py index 84650c8ae..99e084275 100644 --- a/tests/test_metrics_temporal_counts.py +++ b/tests/test_metrics_temporal_counts.py @@ -1,9 +1,7 @@ import numpy as np import pytest - from pyrecest.utils.metrics import chi_square_confidence_bounds - _TEMPORAL_COUNTS = ( np.timedelta64(2, "ns"), np.timedelta64(2, "us"), diff --git a/tests/test_model_comparison_comparable_flags.py b/tests/test_model_comparison_comparable_flags.py index bef9fa05e..8edaa14e9 100644 --- a/tests/test_model_comparison_comparable_flags.py +++ b/tests/test_model_comparison_comparable_flags.py @@ -1,5 +1,4 @@ import pandas as pd - from pyrecest.evaluation.model_comparison import ( evidence_margin_table, paired_model_margin_decisions, diff --git a/tests/test_pytorch_split_index_contract.py b/tests/test_pytorch_split_index_contract.py index e4b94f39c..f066514b0 100644 --- a/tests/test_pytorch_split_index_contract.py +++ b/tests/test_pytorch_split_index_contract.py @@ -2,7 +2,6 @@ from fractions import Fraction import numpy as np - from pyrecest.backend_support._pytorch_split_index_contract import ( _normalize_split_section_count, ) diff --git a/tests/test_sigma_points_temporal_parameters.py b/tests/test_sigma_points_temporal_parameters.py index bcfaacd40..b7185b381 100644 --- a/tests/test_sigma_points_temporal_parameters.py +++ b/tests/test_sigma_points_temporal_parameters.py @@ -1,9 +1,7 @@ import numpy as np import pytest - from pyrecest.sampling import JulierSigmaPoints, MerweScaledSigmaPoints - _TEMPORAL_VALUES = ( np.timedelta64(2, "ns"), np.timedelta64(2, "us"), diff --git a/tests/tracking/test_hypothesis_replay_temporal_validation.py b/tests/tracking/test_hypothesis_replay_temporal_validation.py index 5a581260e..4dfa2c21c 100644 --- a/tests/tracking/test_hypothesis_replay_temporal_validation.py +++ b/tests/tracking/test_hypothesis_replay_temporal_validation.py @@ -8,7 +8,6 @@ rank_hypothesis_replays, ) - _TEMPORAL_VALUES = ( np.timedelta64(2, "ns"), np.datetime64("1970-01-01T00:00:00.000000002"), @@ -45,9 +44,7 @@ def test_temporal_record_statistics_are_ignored() -> None: records=[ { "nis": np.timedelta64(4, "ns"), - "residual_norm_m": np.datetime64( - "1970-01-01T00:00:00.000000005" - ), + "residual_norm_m": np.datetime64("1970-01-01T00:00:00.000000005"), }, { "nis": np.asarray(np.timedelta64(6, "ns"), dtype=object), diff --git a/tests/utils/test_association_model_failed_refit.py b/tests/utils/test_association_model_failed_refit.py index ee8abe1ed..24e6e53be 100644 --- a/tests/utils/test_association_model_failed_refit.py +++ b/tests/utils/test_association_model_failed_refit.py @@ -1,6 +1,5 @@ import numpy.testing as npt import pytest - from pyrecest.backend import array, zeros from pyrecest.utils import LogisticPairwiseAssociationModel @@ -17,9 +16,7 @@ def test_failed_refit_preserves_previous_fitted_state(): converged_before = model.converged_ class_weights_before = model.class_weights_ - replacement_features = array( - [[-2.0, 0.0], [-1.0, 0.0], [1.0, 0.0], [2.0, 0.0]] - ) + replacement_features = array([[-2.0, 0.0], [-1.0, 0.0], [1.0, 0.0], [2.0, 0.0]]) with pytest.raises( ValueError, match="At least one example must receive positive weight" ): diff --git a/tests/utils/test_history_recorder_empty_steps.py b/tests/utils/test_history_recorder_empty_steps.py index 75dd94f1e..d8073885d 100644 --- a/tests/utils/test_history_recorder_empty_steps.py +++ b/tests/utils/test_history_recorder_empty_steps.py @@ -14,9 +14,7 @@ def test_empty_padded_record_preserves_time_axis(): recorder.record("estimate", backend.array([1.0, 2.0]), pad_with_nan=True) recorder.record("estimate", backend.array([]), pad_with_nan=True) - history = recorder.record( - "estimate", backend.array([3.0]), pad_with_nan=True - ) + history = recorder.record("estimate", backend.array([3.0]), pad_with_nan=True) npt.assert_allclose( _to_numpy(history),