Mean attribution distance over random input perturbations within a small radius. Summarises explanation stability on average rather than worst-case. ↓ better.

sensitivity_avg

sensitivity_avg(
    explainer: Explainer | Attribution,
    inputs: TensorOrTupleOfTensorsGeneric,
    perturb_func: Callable = default_perturb_func,
    perturb_radius: float = 0.02,
    n_perturb_samples: int = 10,
    norm_ord: str = "fro",
    max_examples_per_batch: int | None = None,
    multi_target: bool = False,
    **kwargs: Any,
) -> Tensor | list[Tensor]

Average sensitivity — mean attribution distance under input perturbations. ↓ better.

Wraps sensitivity_max_and_avg and returns only the avg component. See sensitivity_max_and_avg for full argument documentation.

Source code in torchxai/metrics/robustness/sensitivity.py
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
def sensitivity_avg(
    explainer: Explainer | Attribution,
    inputs: TensorOrTupleOfTensorsGeneric,
    perturb_func: Callable = default_perturb_func,
    perturb_radius: float = 0.02,
    n_perturb_samples: int = 10,
    norm_ord: str = "fro",
    max_examples_per_batch: int | None = None,
    multi_target: bool = False,
    **kwargs: Any,
) -> Tensor | list[Tensor]:
    """Average sensitivity — mean attribution distance under input perturbations. ↓ better.

    Wraps `sensitivity_max_and_avg` and returns only the avg component.
    See `sensitivity_max_and_avg` for full argument documentation.
    """
    outputs = sensitivity_max_and_avg(
        explainer,
        inputs,
        perturb_func=perturb_func,
        perturb_radius=perturb_radius,
        n_perturb_samples=n_perturb_samples,
        norm_ord=norm_ord,
        max_examples_per_batch=max_examples_per_batch,
        multi_target=multi_target,
        **kwargs,
    )
    assert isinstance(outputs, tuple), (
        "Expected outputs to be a tuple of (sensitivity_max, sensitivity_avg)."
    )
    _, sensitivity_avg = outputs
    return sensitivity_avg