Skip to content
Open
Show file tree
Hide file tree
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
17 changes: 16 additions & 1 deletion captum/_utils/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 "
Expand All @@ -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"
Expand Down
3 changes: 3 additions & 0 deletions captum/attr/_core/feature_ablation.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
)
Expand Down Expand Up @@ -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
)
Expand Down
3 changes: 3 additions & 0 deletions captum/attr/_core/shapley_value.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
)
Expand Down Expand Up @@ -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
)
Expand Down
4 changes: 4 additions & 0 deletions captum/attr/_utils/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
6 changes: 6 additions & 0 deletions tests/attr/test_common.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
36 changes: 36 additions & 0 deletions tests/attr/test_feature_ablation.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]])
Expand Down
35 changes: 35 additions & 0 deletions tests/attr/test_shapley.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Loading