A toolkit for quantitative evaluation of data attribution methods.
See the code
Interpretability toolkit for quantitative evaluation of data attribution methods in PyTorch.
quanda quanda is under active development. Note the release version to ensure reproducibility of your work. Contributions, bug reports, and feature requests are welcome.
Training data attribution (TDA) methods attribute model output on a specific test sample to the training dataset that it was trained on. They reveal the training datapoints responsible for the model's decisions. Existing methods achieve this by estimating the counterfactual effect of removing datapoints from the training set (Koh and Liang, 2017; Park et al., 2023; Bae et al., 2024) tracking the contributions of training points to the loss reduction throughout training (Pruthi et al., 2020), using interpretable surrogate models (Yeh et al., 2018) or finding training samples that are deemed similar to the test sample by the model (Caruana et. al, 1999; Hanawa et. al, 2021). In addition to model understanding, TDA has been used in a variety of applications such as debugging model behavior (Koh and Liang, 2017; Yeh et al., 2018; K and Søgaard, 2021; Guo et al., 2021), data summarization (Khanna et al., 2019; Marion et al., 2023; Yang et al., 2023), dataset selection (Engstrom et al., 2024; Chhabra et al., 2024), fact tracing (Akyurek et al., 2022) and machine unlearning (Warnecke et al., 2023).
Although there are various demonstrations of TDA’s potential for interpretability and practical applications, the critical question of how TDA methods should be effectively evaluated remains open. Several approaches have been proposed by the community, which can be categorized into three groups:
| Library | Reference |
|---|---|
| Captum (Similarity Influence, Arnoldi Influence Function, TracIn) | Caruana et al., 1999; Schioppa et al., 2022; Koh and Liang, 2017; Pruthi et al., 2020 |
| TRAK (TRAK) | Park et al., 2023 |
| Representer Point Selection (Representer Point Selection) | Yeh et al., 2018 |
| Kronfluence (Kronfluence) | Grosse et al., 2023 |
| Dattri (Influence Functions: Explicit / CG / LiSSA / DataInf, Arnoldi, EK-FAC, TracInCP, Grad-Dot, Grad-Cos, TRAK) | Deng et al., 2024 |
Linear Datamodeling Score (Park et al., 2023): Measures the correlation between the (grouped) attribution scores and the actual output of models trained on different subsets of the training set. For each subset, the linear datamodeling score compares the actual model output to the sum of attribution scores from the subset using Spearman rank correlation.
Class Detection / Subclass Detection (Hanawa et al., 2021): Measures the proportion of identical classes or subclasses in the top-1 training samples over the test dataset. If the attributions are based on similarity, they are expected to be predictive of the class of the test datapoint, as well as different subclasses under a single label.
Model Randomization (Hanawa et al., 2021): Measures the correlation between the original TDA and the TDA of a model with randomized weights. Since the attributions are expected to depend on model parameters, the correlation between original and randomized attributions should be low.
Top-K Cardinality (Barshan et al., 2020): Measures the cardinality of the union of the top-K training samples over the test set. A low value uncovers a specific failure mode: a limited pool of top attributions shared across many test samples. Note that only a very low score is suspicious: a low score need not indicate poor attribution quality (a model may legitimately rely on few training samples), and a medium score is not necessarily better than a high one.
Mislabeled Data Detection (Koh and Liang, 2017): Computes the proportion of noisy training labels detected as a function of the percentage of inspected training samples. The samples are inspected in order according to their global TDA ranking, which is computed using local attributions. This produces a cumulative mislabeling detection curve. We expect to see a curve that rapidly increases as we check more of the training data, thus we compute the area under this curve
Shortcut Detection (Yolcu et al., 2025): Assuming a known shortcut, or Clever-Hans effect has been identified in the model, this metric evaluates how effectively a TDA method can identify shortcut samples as the most influential in predicting cases with the shortcut artifact. This process is referred to as Domain Mismatch Debugging in the original paper.
Mixed Datasets (Hammoudeh and Lowd, 2022): In a setting where a model has been trained on two datasets: a clean dataset (e.g. CIFAR-10) and an adversarial (e.g. zeros from MNIST), this metric evaluates how well the model ranks the importance (attribution) of adversarial samples compared to clean samples when making predictions on an adversarial example.
Mean Reciprocal Rank (MRR) (Akyurek et al., 2022): For fact-tracing settings, measures the mean reciprocal rank of the highest-ranked entailing proponent across fact queries.
Recall@k (Akyurek et al., 2022): For fact-tracing settings, measures the proportion of facts for which an entailing proponent appears in the top-k retrievals.
Tail Patch (Chang et al., 2024): For fact-tracing settings, measures the incremental change in target-sequence probability after taking a single training step on retrieved proponents.
| Benchmark | Output range | Better |
|---|---|---|
| ClassDetection | [0, 1] | higher |
| SubclassDetection | [0, 1] | higher |
| MislabelingDetection | [0, 1] | higher |
| ShortcutDetection | [0, 1] | higher |
| MixedDatasets | [0, 1] | higher |
| TopKCardinality | [0, 1] | only very low is suspicious |
| ModelRandomization | [-1, 1] | closer to 0 |
| LinearDatamodelingScore | [-1, 1] | higher |
| MRR | [0, 1] | higher |
| RecallAtK | [0, 1] | higher |
| TailPatch | [-1, 1] | higher |
quanda comes with a few pre-computed benchmarks that can be conveniently used for evaluation in a plug-and-play manner. We are planning to significantly expand the number of benchmarks in the future. The benchmark IDs listed below are to be passed to load_pretrained. The following benchmarks are currently available:
Some settings additionally come with hyperparameter variants, such as LDS at different retraining subset sizes alpha. See Available Benchmarks for the full list of benchmark IDs.
| Metric | Type | Modality | Benchmark IDs (Dataset / Model) |
|---|---|---|---|
| TopKCardinalityMetric | Heuristic | Vision | mnist_top_k_cardinality (MNIST / LeNet) cifar_top_k_cardinality (CIFAR-10 / ResNet-9) awa2_top_k_cardinality (AWA2 / ResNet-50) |
| Text | qnli_top_k_cardinality (QNLI / BERT) | ||
| ModelRandomizationMetric | Heuristic | Vision | mnist_model_randomization (MNIST / LeNet) cifar_model_randomization (CIFAR-10 / ResNet-9) awa2_model_randomization (AWA2 / ResNet-50) |
| Text | qnli_model_randomization (QNLI / BERT) | ||
| MixedDatasetsMetric | Heuristic | Vision | mnist_mixed_datasets (MNIST / LeNet) cifar_mixed_datasets (CIFAR-10 / ResNet-9) awa2_mixed_datasets (AWA2 / ResNet-50) |
| Text | qnli_mixed_datasets (QNLI / BERT) | ||
| ClassDetectionMetric | Downstream-Task-Evaluator | Vision | mnist_class_detection (MNIST / LeNet) cifar_class_detection (CIFAR-10 / ResNet-9) awa2_class_detection (AWA2 / ResNet-50) |
| Text | qnli_class_detection (QNLI / BERT) | ||
| SubclassDetectionMetric | Downstream-Task-Evaluator | Vision | mnist_subclass_detection (MNIST / LeNet) cifar_subclass_detection (CIFAR-10 / ResNet-9) awa2_subclass_detection (AWA2 / ResNet-50) |
| MislabelingDetectionMetric | Downstream-Task-Evaluator | Vision | mnist_mislabeling_detection (MNIST / LeNet) cifar_mislabeling_detection (CIFAR-10 / ResNet-9) awa2_mislabeling_detection (AWA2 / ResNet-50) |
| Text | qnli_mislabeling_detection (QNLI / BERT) | ||
| ShortcutDetectionMetric | Downstream-Task-Evaluator | Vision | mnist_shortcut_detection (MNIST / LeNet) cifar_shortcut_detection (CIFAR-10 / ResNet-9) awa2_shortcut_detection (AWA2 / ResNet-50) |
| MRRMetric | Downstream-Task-Evaluator | Causal LM | gpt2_trex_openwebtext_ft_mrr (T-REx / GPT-2 fine-tuned on OpenWebText) |
| RecallAtKMetric | Downstream-Task-Evaluator | Causal LM | gpt2_trex_openwebtext_ft_recall_at_k (T-REx / GPT-2 fine-tuned on OpenWebText) |
| TailPatchMetric | Downstream-Task-Evaluator | Causal LM | gpt2_trex_openwebtext_ft_tail_patch (T-REx / GPT-2 fine-tuned on OpenWebText) |
| LinearDatamodelingMetric | Ground Truth | Vision | mnist_linear_datamodeling (MNIST / LeNet) cifar_linear_datamodeling (CIFAR-10 / ResNet-9) awa2_linear_datamodeling (AWA2 / ResNet-50) |
| Text | qnli_linear_datamodeling (QNLI / BERT) |
To install quanda from a local clone of this repository, run:
pip install -e .
quanda requires Python 3.10, 3.11 or 3.12. It is recommended to use a virtual environment to install the package.
In the following usage examples, we will be using the SimilarityInfluence data attribution from Captum.
To begin using quanda metrics, you need the following components:
model): A PyTorch model that has already been trained on a relevant dataset. As a placeholder, we used the layer name "avgpool" below. Please replace it with the name of one of the layers in your model.train_set): The dataset used during the training of the model.eval_set): The dataset to be used as test inputs for generating explanations. Explanations are generated with respect to an output neuron corresponding to a certain class. This class can be selected to be the ground truth label of the test points, or the classes predicted by the model. In the following we will use the predicted labels to generate explanations.
Next, we demonstrate how to evaluate explanations using the Model Randomization metric.from torch.utils.data import DataLoader
from tqdm import tqdm
from quanda.explainers.wrappers import CaptumSimilarity
from quanda.metrics.heuristics import ModelRandomizationMetric
We now create our explainer. The device to be used by the explainer and metrics is inherited from the model, thus we set the model device explicitly.
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
model.to(DEVICE)
explainer_kwargs = {
"layers": "fc_2",
"model_id": "default_model_id",
"cache_dir": cache_dir,
}
explainer = CaptumSimilarity(
model=model, train_dataset=dataset, **explainer_kwargs
)
The ModelRandomizationMetric needs to instantiate a new explainer to generate explanations for a randomized model. These will be compared with the explanations of the original model. Therefore, explainer_cls is passed directly to the metric along with initialization parameters of the explainer for the randomized model.
explainer_kwargs = {
"layers": "fc_2",
"model_id": "randomized_model_id",
"cache_dir": cache_dir,
}
ckpt_path = os.path.join(cache_dir, "model_rand_ckpt.pth")
torch.save(model.state_dict(), ckpt_path)
model_rand = ModelRandomizationMetric(
model=model,
model_id="randomized_model_id",
cache_dir=cache_dir,
train_dataset=dataset,
checkpoints=ckpt_path,
explainer_cls=CaptumSimilarity,
expl_kwargs=explainer_kwargs,
correlation_fn="spearman",
seed=42,
)
We now start producing explanations with our TDA method. We go through the test set batch-by-batch. For each batch, we first generate the attributions using the predicted labels, and we then update the metric with the produced explanations to showcase how to concurrently handle the explanation and evaluation processes.
test_loader = DataLoader(eval_set, batch_size=batch_size, shuffle=False)
for test_data, _ in tqdm(test_loader):
test_data = test_data.to(DEVICE)
target = model(test_data).argmax(dim=-1)
tda = explainer.explain(test_data=test_data, targets=target)
model_rand.update(
explanations=tda, test_data=test_data, test_targets=target
)
print("Randomization metric output:", model_rand.compute())
quanda benchmarks allow us to streamline the evaluation process by downloading the necessary data and models, and running the evaluation in a single command. The following code demonstrates how to use the mnist_subclass_detection benchmark:
from quanda.explainers.wrappers import CaptumSimilarity
from quanda.benchmarks.downstream_eval import SubclassDetection
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
model.to(DEVICE)
explainer_kwargs = {
"layers": "fc_2",
"model_id": "default_model_id",
"cache_dir": cache_dir,
}
subclass_detect = SubclassDetection.load_pretrained(
bench_id="mnist_subclass_detection",
cache_dir=cache_dir,
)
score = subclass_detect.evaluate(
explainer_cls=CaptumSimilarity,
expl_kwargs=explainer_kwargs,
batch_size=batch_size,
max_eval_n=16,
)["score"]
print(f"Subclass Detection Score: {score}")
While we provide a number of benchmarks with pre-computed assets, quanda Benchmark objects also expose a train interface for preparing benchmarks from scratch. To train a benchmark, specify its components in a single YAML file (see quanda/benchmarks/resources/configs).
import torch
from quanda.explainers.wrappers import CaptumSimilarity
from quanda.benchmarks.downstream_eval import MislabelingDetection
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
model.to(DEVICE)
explainer_kwargs = {
"layers": "fc_2",
"model_id": "top_k_model",
"cache_dir": cache_dir,
}
For mislabeling detection, we will train a model from scratch using a dataset with a portion of labels flipped.
with open(
"tests/assets/mnist_local_bench/83edb41-default_MislabelingDetection.yaml",
"r",
) as f:
mislabel_config = yaml.safe_load(f)
mislabel_config["bench_save_dir"] = os.path.join(
cache_dir, "mislabeling_detection_bench"
)
mislabeling_detection = MislabelingDetection.train(
mislabel_config,
device=DEVICE,
)
We can now call the evaluate method to directly start the evaluation process on the benchmark.
score = mislabeling_detection.evaluate(
explainer_cls=CaptumSimilarity,
expl_kwargs=explainer_kwargs,
batch_size=batch_size,
max_eval_n=16,
)["score"]
print(f"Mislabeling Detection Score: {score}")
More detailed examples can be found in the tutorials folder.
You can also use Hydra for benchmark training configuration, as shown in scripts/train.py.
In addition to the built-in explainers, quanda supports the evaluation of custom explainer methods. This section provides a guide on how to create a wrapper for a custom explainer that matches our interface.
Your custom explainer should inherit from the base Explainer class provided by quanda. The first step is to initialize your custom explainer within the __init__ method.
from quanda.explainers.base import Explainer
class CustomExplainer(Explainer):
def __init__(self, model, train_dataset, **kwargs):
super().__init__(model, train_dataset, **kwargs)
# Initialize your explainer here
The core of your wrapper is the explain method. This function should take test samples and their corresponding target values as input and return a 2D tensor containing the influence scores.
test: The test batch for which explanations are generated.targets: The target values for the explanations.Ensure that the output tensor has the shape (test_samples, train_samples), where the entries in the train samples dimension are ordered in the same order as in the train_dataset that is being attributed.
def explain(
self,
test_data: torch.Tensor,
targets: Union[List[int], torch.Tensor]
) -> torch.Tensor:
# Compute your influence scores here
return influence_scores
By default, quanda includes a built-in method for calculating self-influence scores. This base implementation computes all attributions over the training dataset, and collects the diagonal values in the attribution matrix. However, you can override this method to provide a more efficient implementation. This method should calculate how much each training sample influences itself and return a tensor of the computed self-influence scores.
def self_influence(self, batch_size: int = 1) -> torch.Tensor:
# Compute your self-influence scores here
return self_influence_scores
For detailed examples, we refer to the existing explainer wrappers in quanda.
Controlled Setting Evaluation: Many metrics require access to ground truth labels for datasets, such as the indices of the "shortcut samples" in the Shortcut Detection metric, or the mislabeling (noisy) label indices for the Mislabeling Detection Metric. However, users often may not have access to these labels. To address this, we recommend either using one of our pre-built benchmark suites (see Benchmarks section) or generating (train method) a custom benchmark for comparing explainers. Benchmarks provide a controlled environment for systematic evaluation.
Explainer Caching: Many explainers in our library generate re-usable cache. The cache_dir and model_id parameters passed to various class instances are used to store these intermediary results. Ensure each experiment is assigned a unique combination of these arguments. Failing to do so could lead to incorrect reuse of cached results. If you wish to avoid re-using cached results, you can set the load_from_disk parameter to False.
Benchmark Dataset Caching: Benchmark initialization methods involve caching a HuggingFace dataset locally to the HF_HOME cache path. We recommend ensuring that the environment variable is set as needed and caching the dataset into the directory in advance of loading the benchmark.
Explainers Are Expensive To Calculate: Certain explainers, such as CaptumTracInCPFastRandProj, may lead to OutOfMemory (OOM) issues when applied to large models or datasets. In such cases, we recommend adjusting memory usage by either reducing the dataset size or using smaller models to avoid these issues.
We have included a few tutorials to demonstrate the usage of quanda:
To install the library with tutorial dependencies, run:
pip install -e '.[tutorials]'
We welcome contributions to quanda! You could contribute by:
A detailed guide on how to contribute to quanda can be found here.
If you have any questions regarding the codebase, please open an issue or contact us via email at dilyabareeva@gmail.com or galip.uemit.yolcu@hhi.fraunhofer.de.
@misc{bareeva2024quandainterpretabilitytoolkittraining,
title={Quanda: An Interpretability Toolkit for Training Data Attribution Evaluation and Beyond},
author={Dilyara Bareeva and Galip Ümit Yolcu and Anna Hedström and Niklas Schmolenski and Thomas Wiegand and Wojciech Samek and Sebastian Lapuschkin},
year={2024},
eprint={2410.07158},
archivePrefix={arXiv},
primaryClass={cs.LG},
url={https://arxiv.org/abs/2410.07158},
}
A toolkit for quantitative evaluation of data attribution methods.
See the code
Interpretability toolkit for quantitative evaluation of data attribution methods in PyTorch.
quanda quanda is under active development. Note the release version to ensure reproducibility of your work. Contributions, bug reports, and feature requests are welcome.
Training data attribution (TDA) methods attribute model output on a specific test sample to the training dataset that it was trained on. They reveal the training datapoints responsible for the model's decisions. Existing methods achieve this by estimating the counterfactual effect of removing datapoints from the training set (Koh and Liang, 2017; Park et al., 2023; Bae et al., 2024) tracking the contributions of training points to the loss reduction throughout training (Pruthi et al., 2020), using interpretable surrogate models (Yeh et al., 2018) or finding training samples that are deemed similar to the test sample by the model (Caruana et. al, 1999; Hanawa et. al, 2021). In addition to model understanding, TDA has been used in a variety of applications such as debugging model behavior (Koh and Liang, 2017; Yeh et al., 2018; K and Søgaard, 2021; Guo et al., 2021), data summarization (Khanna et al., 2019; Marion et al., 2023; Yang et al., 2023), dataset selection (Engstrom et al., 2024; Chhabra et al., 2024), fact tracing (Akyurek et al., 2022) and machine unlearning (Warnecke et al., 2023).
Although there are various demonstrations of TDA’s potential for interpretability and practical applications, the critical question of how TDA methods should be effectively evaluated remains open. Several approaches have been proposed by the community, which can be categorized into three groups:
| Library | Reference |
|---|---|
| Captum (Similarity Influence, Arnoldi Influence Function, TracIn) | Caruana et al., 1999; Schioppa et al., 2022; Koh and Liang, 2017; Pruthi et al., 2020 |
| TRAK (TRAK) | Park et al., 2023 |
| Representer Point Selection (Representer Point Selection) | Yeh et al., 2018 |
| Kronfluence (Kronfluence) | Grosse et al., 2023 |
| Dattri (Influence Functions: Explicit / CG / LiSSA / DataInf, Arnoldi, EK-FAC, TracInCP, Grad-Dot, Grad-Cos, TRAK) | Deng et al., 2024 |
Linear Datamodeling Score (Park et al., 2023): Measures the correlation between the (grouped) attribution scores and the actual output of models trained on different subsets of the training set. For each subset, the linear datamodeling score compares the actual model output to the sum of attribution scores from the subset using Spearman rank correlation.
Class Detection / Subclass Detection (Hanawa et al., 2021): Measures the proportion of identical classes or subclasses in the top-1 training samples over the test dataset. If the attributions are based on similarity, they are expected to be predictive of the class of the test datapoint, as well as different subclasses under a single label.
Model Randomization (Hanawa et al., 2021): Measures the correlation between the original TDA and the TDA of a model with randomized weights. Since the attributions are expected to depend on model parameters, the correlation between original and randomized attributions should be low.
Top-K Cardinality (Barshan et al., 2020): Measures the cardinality of the union of the top-K training samples over the test set. A low value uncovers a specific failure mode: a limited pool of top attributions shared across many test samples. Note that only a very low score is suspicious: a low score need not indicate poor attribution quality (a model may legitimately rely on few training samples), and a medium score is not necessarily better than a high one.
Mislabeled Data Detection (Koh and Liang, 2017): Computes the proportion of noisy training labels detected as a function of the percentage of inspected training samples. The samples are inspected in order according to their global TDA ranking, which is computed using local attributions. This produces a cumulative mislabeling detection curve. We expect to see a curve that rapidly increases as we check more of the training data, thus we compute the area under this curve
Shortcut Detection (Yolcu et al., 2025): Assuming a known shortcut, or Clever-Hans effect has been identified in the model, this metric evaluates how effectively a TDA method can identify shortcut samples as the most influential in predicting cases with the shortcut artifact. This process is referred to as Domain Mismatch Debugging in the original paper.
Mixed Datasets (Hammoudeh and Lowd, 2022): In a setting where a model has been trained on two datasets: a clean dataset (e.g. CIFAR-10) and an adversarial (e.g. zeros from MNIST), this metric evaluates how well the model ranks the importance (attribution) of adversarial samples compared to clean samples when making predictions on an adversarial example.
Mean Reciprocal Rank (MRR) (Akyurek et al., 2022): For fact-tracing settings, measures the mean reciprocal rank of the highest-ranked entailing proponent across fact queries.
Recall@k (Akyurek et al., 2022): For fact-tracing settings, measures the proportion of facts for which an entailing proponent appears in the top-k retrievals.
Tail Patch (Chang et al., 2024): For fact-tracing settings, measures the incremental change in target-sequence probability after taking a single training step on retrieved proponents.
| Benchmark | Output range | Better |
|---|---|---|
| ClassDetection | [0, 1] | higher |
| SubclassDetection | [0, 1] | higher |
| MislabelingDetection | [0, 1] | higher |
| ShortcutDetection | [0, 1] | higher |
| MixedDatasets | [0, 1] | higher |
| TopKCardinality | [0, 1] | only very low is suspicious |
| ModelRandomization | [-1, 1] | closer to 0 |
| LinearDatamodelingScore | [-1, 1] | higher |
| MRR | [0, 1] | higher |
| RecallAtK | [0, 1] | higher |
| TailPatch | [-1, 1] | higher |
quanda comes with a few pre-computed benchmarks that can be conveniently used for evaluation in a plug-and-play manner. We are planning to significantly expand the number of benchmarks in the future. The benchmark IDs listed below are to be passed to load_pretrained. The following benchmarks are currently available:
Some settings additionally come with hyperparameter variants, such as LDS at different retraining subset sizes alpha. See Available Benchmarks for the full list of benchmark IDs.
| Metric | Type | Modality | Benchmark IDs (Dataset / Model) |
|---|---|---|---|
| TopKCardinalityMetric | Heuristic | Vision | mnist_top_k_cardinality (MNIST / LeNet) cifar_top_k_cardinality (CIFAR-10 / ResNet-9) awa2_top_k_cardinality (AWA2 / ResNet-50) |
| Text | qnli_top_k_cardinality (QNLI / BERT) | ||
| ModelRandomizationMetric | Heuristic | Vision | mnist_model_randomization (MNIST / LeNet) cifar_model_randomization (CIFAR-10 / ResNet-9) awa2_model_randomization (AWA2 / ResNet-50) |
| Text | qnli_model_randomization (QNLI / BERT) | ||
| MixedDatasetsMetric | Heuristic | Vision | mnist_mixed_datasets (MNIST / LeNet) cifar_mixed_datasets (CIFAR-10 / ResNet-9) awa2_mixed_datasets (AWA2 / ResNet-50) |
| Text | qnli_mixed_datasets (QNLI / BERT) | ||
| ClassDetectionMetric | Downstream-Task-Evaluator | Vision | mnist_class_detection (MNIST / LeNet) cifar_class_detection (CIFAR-10 / ResNet-9) awa2_class_detection (AWA2 / ResNet-50) |
| Text | qnli_class_detection (QNLI / BERT) | ||
| SubclassDetectionMetric | Downstream-Task-Evaluator | Vision | mnist_subclass_detection (MNIST / LeNet) cifar_subclass_detection (CIFAR-10 / ResNet-9) awa2_subclass_detection (AWA2 / ResNet-50) |
| MislabelingDetectionMetric | Downstream-Task-Evaluator | Vision | mnist_mislabeling_detection (MNIST / LeNet) cifar_mislabeling_detection (CIFAR-10 / ResNet-9) awa2_mislabeling_detection (AWA2 / ResNet-50) |
| Text | qnli_mislabeling_detection (QNLI / BERT) | ||
| ShortcutDetectionMetric | Downstream-Task-Evaluator | Vision | mnist_shortcut_detection (MNIST / LeNet) cifar_shortcut_detection (CIFAR-10 / ResNet-9) awa2_shortcut_detection (AWA2 / ResNet-50) |
| MRRMetric | Downstream-Task-Evaluator | Causal LM | gpt2_trex_openwebtext_ft_mrr (T-REx / GPT-2 fine-tuned on OpenWebText) |
| RecallAtKMetric | Downstream-Task-Evaluator | Causal LM | gpt2_trex_openwebtext_ft_recall_at_k (T-REx / GPT-2 fine-tuned on OpenWebText) |
| TailPatchMetric | Downstream-Task-Evaluator | Causal LM | gpt2_trex_openwebtext_ft_tail_patch (T-REx / GPT-2 fine-tuned on OpenWebText) |
| LinearDatamodelingMetric | Ground Truth | Vision | mnist_linear_datamodeling (MNIST / LeNet) cifar_linear_datamodeling (CIFAR-10 / ResNet-9) awa2_linear_datamodeling (AWA2 / ResNet-50) |
| Text | qnli_linear_datamodeling (QNLI / BERT) |
To install quanda from a local clone of this repository, run:
pip install -e .
quanda requires Python 3.10, 3.11 or 3.12. It is recommended to use a virtual environment to install the package.
In the following usage examples, we will be using the SimilarityInfluence data attribution from Captum.
To begin using quanda metrics, you need the following components:
model): A PyTorch model that has already been trained on a relevant dataset. As a placeholder, we used the layer name "avgpool" below. Please replace it with the name of one of the layers in your model.train_set): The dataset used during the training of the model.eval_set): The dataset to be used as test inputs for generating explanations. Explanations are generated with respect to an output neuron corresponding to a certain class. This class can be selected to be the ground truth label of the test points, or the classes predicted by the model. In the following we will use the predicted labels to generate explanations.
Next, we demonstrate how to evaluate explanations using the Model Randomization metric.from torch.utils.data import DataLoader
from tqdm import tqdm
from quanda.explainers.wrappers import CaptumSimilarity
from quanda.metrics.heuristics import ModelRandomizationMetric
We now create our explainer. The device to be used by the explainer and metrics is inherited from the model, thus we set the model device explicitly.
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
model.to(DEVICE)
explainer_kwargs = {
"layers": "fc_2",
"model_id": "default_model_id",
"cache_dir": cache_dir,
}
explainer = CaptumSimilarity(
model=model, train_dataset=dataset, **explainer_kwargs
)
The ModelRandomizationMetric needs to instantiate a new explainer to generate explanations for a randomized model. These will be compared with the explanations of the original model. Therefore, explainer_cls is passed directly to the metric along with initialization parameters of the explainer for the randomized model.
explainer_kwargs = {
"layers": "fc_2",
"model_id": "randomized_model_id",
"cache_dir": cache_dir,
}
ckpt_path = os.path.join(cache_dir, "model_rand_ckpt.pth")
torch.save(model.state_dict(), ckpt_path)
model_rand = ModelRandomizationMetric(
model=model,
model_id="randomized_model_id",
cache_dir=cache_dir,
train_dataset=dataset,
checkpoints=ckpt_path,
explainer_cls=CaptumSimilarity,
expl_kwargs=explainer_kwargs,
correlation_fn="spearman",
seed=42,
)
We now start producing explanations with our TDA method. We go through the test set batch-by-batch. For each batch, we first generate the attributions using the predicted labels, and we then update the metric with the produced explanations to showcase how to concurrently handle the explanation and evaluation processes.
test_loader = DataLoader(eval_set, batch_size=batch_size, shuffle=False)
for test_data, _ in tqdm(test_loader):
test_data = test_data.to(DEVICE)
target = model(test_data).argmax(dim=-1)
tda = explainer.explain(test_data=test_data, targets=target)
model_rand.update(
explanations=tda, test_data=test_data, test_targets=target
)
print("Randomization metric output:", model_rand.compute())
quanda benchmarks allow us to streamline the evaluation process by downloading the necessary data and models, and running the evaluation in a single command. The following code demonstrates how to use the mnist_subclass_detection benchmark:
from quanda.explainers.wrappers import CaptumSimilarity
from quanda.benchmarks.downstream_eval import SubclassDetection
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
model.to(DEVICE)
explainer_kwargs = {
"layers": "fc_2",
"model_id": "default_model_id",
"cache_dir": cache_dir,
}
subclass_detect = SubclassDetection.load_pretrained(
bench_id="mnist_subclass_detection",
cache_dir=cache_dir,
)
score = subclass_detect.evaluate(
explainer_cls=CaptumSimilarity,
expl_kwargs=explainer_kwargs,
batch_size=batch_size,
max_eval_n=16,
)["score"]
print(f"Subclass Detection Score: {score}")
While we provide a number of benchmarks with pre-computed assets, quanda Benchmark objects also expose a train interface for preparing benchmarks from scratch. To train a benchmark, specify its components in a single YAML file (see quanda/benchmarks/resources/configs).
import torch
from quanda.explainers.wrappers import CaptumSimilarity
from quanda.benchmarks.downstream_eval import MislabelingDetection
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
model.to(DEVICE)
explainer_kwargs = {
"layers": "fc_2",
"model_id": "top_k_model",
"cache_dir": cache_dir,
}
For mislabeling detection, we will train a model from scratch using a dataset with a portion of labels flipped.
with open(
"tests/assets/mnist_local_bench/83edb41-default_MislabelingDetection.yaml",
"r",
) as f:
mislabel_config = yaml.safe_load(f)
mislabel_config["bench_save_dir"] = os.path.join(
cache_dir, "mislabeling_detection_bench"
)
mislabeling_detection = MislabelingDetection.train(
mislabel_config,
device=DEVICE,
)
We can now call the evaluate method to directly start the evaluation process on the benchmark.
score = mislabeling_detection.evaluate(
explainer_cls=CaptumSimilarity,
expl_kwargs=explainer_kwargs,
batch_size=batch_size,
max_eval_n=16,
)["score"]
print(f"Mislabeling Detection Score: {score}")
More detailed examples can be found in the tutorials folder.
You can also use Hydra for benchmark training configuration, as shown in scripts/train.py.
In addition to the built-in explainers, quanda supports the evaluation of custom explainer methods. This section provides a guide on how to create a wrapper for a custom explainer that matches our interface.
Your custom explainer should inherit from the base Explainer class provided by quanda. The first step is to initialize your custom explainer within the __init__ method.
from quanda.explainers.base import Explainer
class CustomExplainer(Explainer):
def __init__(self, model, train_dataset, **kwargs):
super().__init__(model, train_dataset, **kwargs)
# Initialize your explainer here
The core of your wrapper is the explain method. This function should take test samples and their corresponding target values as input and return a 2D tensor containing the influence scores.
test: The test batch for which explanations are generated.targets: The target values for the explanations.Ensure that the output tensor has the shape (test_samples, train_samples), where the entries in the train samples dimension are ordered in the same order as in the train_dataset that is being attributed.
def explain(
self,
test_data: torch.Tensor,
targets: Union[List[int], torch.Tensor]
) -> torch.Tensor:
# Compute your influence scores here
return influence_scores
By default, quanda includes a built-in method for calculating self-influence scores. This base implementation computes all attributions over the training dataset, and collects the diagonal values in the attribution matrix. However, you can override this method to provide a more efficient implementation. This method should calculate how much each training sample influences itself and return a tensor of the computed self-influence scores.
def self_influence(self, batch_size: int = 1) -> torch.Tensor:
# Compute your self-influence scores here
return self_influence_scores
For detailed examples, we refer to the existing explainer wrappers in quanda.
Controlled Setting Evaluation: Many metrics require access to ground truth labels for datasets, such as the indices of the "shortcut samples" in the Shortcut Detection metric, or the mislabeling (noisy) label indices for the Mislabeling Detection Metric. However, users often may not have access to these labels. To address this, we recommend either using one of our pre-built benchmark suites (see Benchmarks section) or generating (train method) a custom benchmark for comparing explainers. Benchmarks provide a controlled environment for systematic evaluation.
Explainer Caching: Many explainers in our library generate re-usable cache. The cache_dir and model_id parameters passed to various class instances are used to store these intermediary results. Ensure each experiment is assigned a unique combination of these arguments. Failing to do so could lead to incorrect reuse of cached results. If you wish to avoid re-using cached results, you can set the load_from_disk parameter to False.
Benchmark Dataset Caching: Benchmark initialization methods involve caching a HuggingFace dataset locally to the HF_HOME cache path. We recommend ensuring that the environment variable is set as needed and caching the dataset into the directory in advance of loading the benchmark.
Explainers Are Expensive To Calculate: Certain explainers, such as CaptumTracInCPFastRandProj, may lead to OutOfMemory (OOM) issues when applied to large models or datasets. In such cases, we recommend adjusting memory usage by either reducing the dataset size or using smaller models to avoid these issues.
We have included a few tutorials to demonstrate the usage of quanda:
To install the library with tutorial dependencies, run:
pip install -e '.[tutorials]'
We welcome contributions to quanda! You could contribute by:
A detailed guide on how to contribute to quanda can be found here.
If you have any questions regarding the codebase, please open an issue or contact us via email at dilyabareeva@gmail.com or galip.uemit.yolcu@hhi.fraunhofer.de.
@misc{bareeva2024quandainterpretabilitytoolkittraining,
title={Quanda: An Interpretability Toolkit for Training Data Attribution Evaluation and Beyond},
author={Dilyara Bareeva and Galip Ümit Yolcu and Anna Hedström and Niklas Schmolenski and Thomas Wiegand and Wojciech Samek and Sebastian Lapuschkin},
year={2024},
eprint={2410.07158},
archivePrefix={arXiv},
primaryClass={cs.LG},
url={https://arxiv.org/abs/2410.07158},
}