diff --git a/captum/_utils/common.py b/captum/_utils/common.py index cab42bf2cd..703e7a5c28 100644 --- a/captum/_utils/common.py +++ b/captum/_utils/common.py @@ -117,6 +117,7 @@ def _validate_input( inputs: Tuple[Tensor, ...], baselines: Tuple[Union[Tensor, int, float], ...], draw_baseline_from_distrib: bool = False, + allow_broadcastable_baselines: bool = False, ) -> None: assert len(inputs) == len(baselines), ( "Input and baseline must have the same " @@ -137,10 +138,24 @@ def _validate_input( " Found baseline: {} and input: {} ".format(baseline, input) ) else: + baseline_is_broadcastable = False + if allow_broadcastable_baselines and isinstance(baseline, Tensor): + try: + baseline_is_broadcastable = ( + torch.broadcast_shapes(input.shape, baseline.shape) + == input.shape + ) + except RuntimeError: + pass assert ( isinstance(baseline, (int, float)) or input.shape == baseline.shape - or baseline.shape[0] == 1 + or ( + baseline.dim() > 0 + and baseline.shape[0] == 1 + and input.shape[1:] == baseline.shape[1:] + ) + or baseline_is_broadcastable ), ( "Baseline can be provided as a tensor for just one input and" " broadcasted to the batch or input and baseline must have the" diff --git a/captum/attr/_core/feature_ablation.py b/captum/attr/_core/feature_ablation.py index 505d4cb260..3111a2ab72 100644 --- a/captum/attr/_core/feature_ablation.py +++ b/captum/attr/_core/feature_ablation.py @@ -22,6 +22,7 @@ _is_tuple, _maybe_expand_parameters, _run_forward, + _validate_input, ) from captum._utils.exceptions import FeatureAblationFutureError from captum._utils.progress import NullProgress, progress, Progress @@ -501,6 +502,7 @@ def attribute( is_inputs_tuple = _is_tuple(inputs) formatted_inputs, baselines = _format_input_baseline(inputs, baselines) + _validate_input(formatted_inputs, baselines, allow_broadcastable_baselines=True) formatted_additional_forward_args = _format_additional_forward_args( additional_forward_args ) @@ -834,6 +836,7 @@ def attribute_future( # converting it into a tuple. is_inputs_tuple = _is_tuple(inputs) formatted_inputs, baselines = _format_input_baseline(inputs, baselines) + _validate_input(formatted_inputs, baselines, allow_broadcastable_baselines=True) formatted_additional_forward_args = _format_additional_forward_args( additional_forward_args ) diff --git a/captum/attr/_core/shapley_value.py b/captum/attr/_core/shapley_value.py index 32f8c95c2a..7b6a9ded9c 100644 --- a/captum/attr/_core/shapley_value.py +++ b/captum/attr/_core/shapley_value.py @@ -37,6 +37,7 @@ _is_mask_valid, _is_tuple, _run_forward, + _validate_input, ) from captum._utils.exceptions import ShapleyValueFutureError from captum._utils.progress import progress @@ -320,6 +321,7 @@ def attribute( # converting it into a tuple. is_inputs_tuple = _is_tuple(inputs) inputs_tuple, baselines = _format_input_baseline(inputs, baselines) + _validate_input(inputs_tuple, baselines) additional_forward_args = _format_additional_forward_args( additional_forward_args ) @@ -487,6 +489,7 @@ def attribute_future( ) -> Future[TensorOrTupleOfTensorsGeneric]: is_inputs_tuple = _is_tuple(inputs) inputs_tuple, baselines = _format_input_baseline(inputs, baselines) + _validate_input(inputs_tuple, baselines) additional_forward_args = _format_additional_forward_args( additional_forward_args ) diff --git a/captum/attr/_utils/common.py b/captum/attr/_utils/common.py index cab1a31050..5ffd3b51ea 100644 --- a/captum/attr/_utils/common.py +++ b/captum/attr/_utils/common.py @@ -292,6 +292,10 @@ def _tensorize_single_baseline(baseline, input): type(baselines), type(inputs) ) ) + assert len(inputs) == len(baselines), ( + "Input and baseline must have the same dimensions, baseline has " + f"{len(baselines)} features whereas input has {len(inputs)}." + ) return tuple( _tensorize_single_baseline(baseline, input) for baseline, input in zip(baselines, inputs) diff --git a/tests/attr/test_common.py b/tests/attr/test_common.py index a3e95a3a01..2052383c14 100644 --- a/tests/attr/test_common.py +++ b/tests/attr/test_common.py @@ -56,6 +56,12 @@ def test_validate_input(self) -> None: (torch.tensor([-1.0]),), (torch.tensor([-2.0]),), method="gausslegendre" ) + with self.assertRaisesRegex(AssertionError, "Baseline can be provided"): + _validate_input( + (torch.zeros((2, 3)),), + (torch.zeros((1, 4)),), + ) + def test_validate_nt_type(self) -> None: with self.assertRaises( AssertionError, diff --git a/tests/attr/test_feature_ablation.py b/tests/attr/test_feature_ablation.py index db293217c9..2917d9bd74 100644 --- a/tests/attr/test_feature_ablation.py +++ b/tests/attr/test_feature_ablation.py @@ -131,6 +131,42 @@ def test_simple_ablation_with_baselines(self) -> None: perturbations_per_eval=(1, 2, 3), ) + def test_perturbation_methods_reject_baseline_tuple_arity_mismatch(self) -> None: + inputs = (torch.tensor([[1.0]]), torch.tensor([[2.0]])) + attribution = FeatureAblation(lambda first, second: first[:, 0] + second[:, 0]) + + for baselines in ( + (torch.tensor([[0.0]]),), + (torch.tensor([[0.0]]),) * 3, + ): + with self.subTest(baseline_count=len(baselines)): + with self.assertRaisesRegex( + AssertionError, "Input and baseline must have the same" + ): + attribution.attribute(inputs, baselines=baselines) + + def test_perturbation_methods_reject_invalid_baseline_batch_size(self) -> None: + inputs = torch.tensor([[1.0], [2.0], [3.0]]) + + attribution = FeatureAblation(lambda values: values[:, 0]) + for baseline in ( + torch.tensor([[10.0], [20.0]]), + torch.tensor([[10.0, 20.0]]), + ): + with self.subTest(baseline_shape=baseline.shape): + with self.assertRaisesRegex(AssertionError, "Baseline can be provided"): + attribution.attribute(inputs, baselines=baseline) + + def test_accepts_broadcastable_tensor_baselines(self) -> None: + inputs = torch.tensor([[1.0, 2.0], [3.0, 4.0]]) + attribution = FeatureAblation(lambda values: values.sum(dim=1)) + + for baseline in (torch.tensor([0.5, 1.5]), torch.tensor(0.5)): + with self.subTest(baseline_shape=baseline.shape): + result = attribution.attribute(inputs, baselines=baseline) + expected = inputs - baseline + torch.testing.assert_close(result, expected) + def test_simple_ablation_boolean(self) -> None: ablation_algo = FeatureAblation(BasicModelBoolInput()) inp = torch.tensor([[True, False, True]]) diff --git a/tests/attr/test_shapley.py b/tests/attr/test_shapley.py index d7d9db56d1..1784143787 100644 --- a/tests/attr/test_shapley.py +++ b/tests/attr/test_shapley.py @@ -29,6 +29,41 @@ class Test(BaseTest): + def test_rejects_baseline_tuple_arity_mismatch(self) -> None: + inputs = (torch.tensor([[1.0]]), torch.tensor([[2.0]])) + attribution = ShapleyValueSampling( + lambda first, second: first[:, 0] + second[:, 0] + ) + + for baselines in ( + (torch.tensor([[0.0]]),), + (torch.tensor([[0.0]]),) * 3, + ): + with self.subTest(baseline_count=len(baselines)): + with self.assertRaisesRegex( + AssertionError, "Input and baseline must have the same" + ): + attribution.attribute(inputs, baselines=baselines) + with self.assertRaisesRegex( + AssertionError, "Input and baseline must have the same" + ): + attribution.attribute_future(inputs, baselines=baselines) + + def test_rejects_invalid_baseline_shape(self) -> None: + inputs = torch.tensor([[1.0], [2.0], [3.0]]) + attribution = ShapleyValueSampling(lambda values: values[:, 0]) + + for baseline in ( + torch.tensor([[10.0], [20.0]]), + torch.tensor([[10.0, 20.0]]), + torch.tensor(10.0), + ): + with self.subTest(baseline_shape=baseline.shape): + with self.assertRaisesRegex(AssertionError, "Baseline can be provided"): + attribution.attribute(inputs, baselines=baseline) + with self.assertRaisesRegex(AssertionError, "Baseline can be provided"): + attribution.attribute_future(inputs, baselines=baseline) + @parameterized.expand([True, False]) def test_simple_shapley_sampling(self, use_future: bool) -> None: inp = torch.tensor([[20.0, 50.0, 30.0]], requires_grad=True)