Datamodels: Predicting Predictions from Training Data

Andrew Ilyas, Sung Min Park, Logan Engstrom, Guillaume Leclerc, Aleksander Madry

Introduction

What kinds of biases does my (machine learning) system exhibit? What correlations does it exploit? On what subpopulations does it perform well (or poorly)?

A recent body of work in machine learning suggests that the answers to these questions lie within both the learning algorithm and the training data used [CLK+19a, GDG17a, IST+19a, Hoo21a, JTM21a]. However, it is often difficult to understand how algorithms and data combine to yield model predictions. In this work, we present datamodeling—a framework for tackling this question by forming an explicit model for predictions in terms of the training data.

Consider a typical machine learning setup, starting with a training set SS comprising dd input-label pairs. The focal point of this setup is a learning algorithm A\mathcal{A} that takes in such a training set of input-label pairs, and outputs a trained model. (Note that this learning algorithm does not have to be deterministic—for example, A\mathcal{A} might encode the process of training a deep neural network from random initialization using stochastic gradient descent.)

Now, consider a fixed input xx (e.g., a photo from the test set of a computer vision benchmark) and define

where we leave “outcome” intentionally broad to capture a variety of use cases. For example, fA(x;S)f_{\mathcal{A}}({x};{S}) may be the cross-entropy loss of a classifier on xx, or the squared-error of a regression model on xx. The potentially stochastic nature of A\mathcal{A} means that fA(x;S)f_{\mathcal{A}}({x};{S}) is a random variable.

Broadly, we aim to understand how the training examples in SS combine through the learning algorithm A\mathcal{A} to yield fA(x;S)f_{\mathcal{A}}({x};{S}) (again, for the specifically chosen input xx). To this end, we leverage a classic technique for studying black-box functions: surrogate modeling [SWM+89a]. In surrogate modeling, one replaces a complex black-box function with an inexact but significantly easier-to-analyze approximation, then uses the latter to shed light on the behavior of the original function.

In our context, the complex black-box function is fA(x;⋅)f_{\mathcal{A}}({x};{\cdot}). We thus aim to find a simple surrogate function g(S′)g(S^{\prime}) whose output roughly matches fA(x;S′)f_{\mathcal{A}}({x};{S^{\prime}}) for a variety of training sets S′S^{\prime} (but again, for a fixed input xx). Achieving this goal would reduce the challenge of scrutinizing fA(x;⋅)f_{\mathcal{A}}({x};{\cdot})—and more generally, the map from training data to predictions through learning algorithm A\mathcal{A}—to the (hopefully easier) task of analyzing gg.

By parameterizing the surrogate function gg (e.g., as gθg_{\theta}, for a parameter vector θ\theta), we transform the challenge of constructing a surrogate function into a supervised learning problem. In this problem, the “training examples” are subsets S′⊂SS^{\prime}\subset S of the original task’s training set SS, and the corresponding “labels” are given by fA(x;S′)f_{\mathcal{A}}({x};{S^{\prime}}) (which we can compute by simply training a new model on S′S^{\prime} with algorithm A\mathcal{A}, and evaluating on xx). Our goal is then to fit a parametric function gθg_{\theta} mapping the former to the latter.

We now formalize this idea as datamodeling—a framework that forms the basis of our work. In this framework, we first fix a distribution over subsets that we will use to collect “training data” for gθg_{\theta},

and then use DS\mathcal{D}_{S} to collect a datamodel training set, or a collection of pairs

where Si∼DSS_{i}\sim\mathcal{D}_{S}, and again fA(x;Si)f_{\mathcal{A}}({x};{S_{i}}) is the result of training a model on SiS_{i} and evaluating on xx (cf. (1)).

We next focus on how to parameterize our surrogate function gθg_{\theta}. In theory, gθg_{\theta} can be any map that takes as input subsets of the training set, and returns estimates of fA(x;⋅)f_{\mathcal{A}}({x};{\cdot}). However, to simplify gθg_{\theta} we ignore the actual contents of the subsets SiS_{i}, and instead focus solely on the presence of each training example of SS within SiS_{i}. In particular, we consider the characteristic vector corresponding to each SiS_{i},

a vector that indicates which elements of the original training set SS belong to a given subset SiS_{i}. We then define a datamodel for a given input xx as a function

and L(⋅,⋅)\mathcal{L}(\cdot,\cdot) is a fixed loss function (e.g., squared-error). This setup (4) places datamodels squarely within the realm of supervised learning: e.g., we can easily test a given datamodel by sampling new subset-output pairs {(Si,fA(x;Si))}\{(S_{i},f_{\mathcal{A}}({x};{S_{i}}))\} and computing average loss. For completeness, we restate the entire datamodeling framework below: {defn}[Datamodeling] Consider a fixed training set SS, a learning algorithm A\mathcal{A}, a target example xx, and a distribution DS\mathcal{D}_{S} over subsets of SS. For any set S′⊂SS^{\prime}\subset S, let fA(x;S′)f_{\mathcal{A}}({x};{S^{\prime}}) be the (stochastic) output of training a model on S′S^{\prime} using A\mathcal{A}, and evaluating on xx. A datamodel for xx is a parametric function gθg_{\theta} optimized to predict fA(x;Si)f_{\mathcal{A}}({x};{S_{i}}) from training subsets Si∼DSS_{i}\sim\mathcal{D}_{S}, i.e.,

Datamodeling studies model classes, not specific models: Datamodeling focuses on the entire distribution of models induced by the algorithm A\mathcal{A}, rather than a specific model. Recent work suggests this distinction is particularly significant for modern learning algorithms (e.g., neural networks), where models can exhibit drastically different behavior depending on only the choice of random seed during training [NB20a, JNB+21a, DHM+20a, ZGK+21a]—we discuss this further in Section 6.

Datamodels are target example-specific: A datamodel gθg_{\theta} predicts model outputs on a specific but arbitrary target example xx. This xx might be an example from the test set, a synthetically generated example, or even (as we will see in Section 3.1) an example from the training set SS itself. We will often work with collections of datamodels corresponding to a set of target examples (e.g., we might consider a test set {x1,…xn}\{x_{1},\ldots x_{n}\} with corresponding datamodels {gθ1,…gθn}\{g_{\theta_{1}},\ldots g_{\theta_{n}}\}). In Section 3 we show that as long as the learning algorithm A\mathcal{A} and the training set SS are fixed, computing a collection of datamodels simultaneously is not much harder than computing a single one.

1 Roadmap and contributions

The key contribution of our work is the datamodeling framework described above, which allows us to analyze the behavior of a machine learning algorithm A\mathcal{A} in terms of the training data. In the remainder of this work, we show how to instantiate, implement, and apply this framework.

We begin in Section 2 by considering a concrete instantiation of datamodeling in which the map gθg_{\theta} is a linear function. Then, in Section 3 we develop the remaining machinery required to apply this instantiation to deep neural networks trained on standard image datasets. In the rest of the paper, we find that:

Datamodels successfully predict model outputs (§ 3.2, Figure 2): despite their simplicity, datamodels yield predictions that match expected model outputs on new sets SS drawn from the same distribution DS\mathcal{D}_{S}. (For example, the Pearson correlation between predicted and ground-truth outputs is r>0.99r>0.99.)

Datamodels successfully predict counterfactuals (§ 4.1, Figure 2): predictions correlate with model outputs even on out-of-distribution training subsets (Figures 8, 45 and Appendix F.1) allowing us to estimate the causal effect of removing training images on a given test prediction. Leveraging this ability, we find that for 50% of CIFAR-10 [Kri09a] test images, models can be made incorrect by removing less than 200 target-specific training points (i.e., 0.4% of the total training set size). If one mislabels the training examples instead of only removing them, 35 label-specific points suffice.

Datamodel weights encode similarity (§ 4.2, Figure 3): the most positive (resp., negative) datamodel weights tend to correspond to similar training images from the same (resp., different) class as the target example xx. We use this property to identify significant train-test leakage across both datasets we study (CIFAR-10 and Functional Map of the World [KSM+20a, CFW+18a]).

enable (qualitatively) high-quality clustering;

allow us to identify model-relevant subpopulations that we can causally verify in a natural sense;

have a number of advantages over representations derived from, e.g., the penultimate layer of a fixed pre-trained network, such as higher effective dimensionality and (a priori) human-meaningful coordinates.

More broadly, datamodels turn out to be a versatile tool for understanding how learning algorithms leverage their training data. In Section 6, we contextualize datamodeling with respect to several ongoing lines of work in machine learning and statistics. We conclude, in Section 7, by outlining a variety of directions for future work on both improving and applying datamodels.

Constructing (linear) datamodels

As described in Section 1, building datamodels comprises the following steps:

pick a parameterized class of functions gθg_{\theta};

sample a collection of subsets Si⊂SS_{i}\subset S from a fixed training set according to a distribution DS\mathcal{D}_{S};

for each subset SiS_{i}, train a model using algorithm A\mathcal{A}, evaluate the model on target input xx using the relevant metric (e.g., loss); collect the resulting pair (1Si,fA(x;Si))(\bm{1}_{S_{i}},f_{\mathcal{A}}({x};{S_{i}}));

split the collected dataset of subset-output pairs into a datamodel training set of size mm, a datamodel validation set of size mvalm_{val}, and a datamodel test set of size mtestm_{test};

estimate parameters θ\theta by fitting gθg_{\theta} on subset-output pairs, i.e., by minimizing

over the collected datamodel training set, and use the validation set to perform model selection.

We now explicitly instantiate this framework, with the goal of understanding the predictions of (deep) classification models. To this end, we revisit steps (a)-(e) above, and consider each relevant aspect—the sampling distribution DS\mathcal{D}_{S}, the output function fA(x;S)f_{\mathcal{A}}({x};{S}), the parameterized family gθg_{\theta}, and the loss function L(⋅,⋅)\mathcal{L}(\cdot,\cdot)—separately:

The first design choice to make is which family of parameterized surrogate functions gθg_{\theta} to optimize over. At first, one might be inclined to use a complex family of functions in the hope of reducing potential misspecification error. After all, gθg_{\theta} is meant to be a surrogate for the end-to-end training of a deep classifier. In this work, however, we will instantiate datamodeling by taking gθ(⋅)g_{\theta}(\cdot) to be a simple linear mapping

where we recall that 1Si\bm{1}_{S_{i}} is the size-dd characteristic vector of SiS_{i} within SS (cf. (3)).

While we will allow gθ(⋅)g_{\theta}(\cdot) to fit a bias term as above, for notational convenience we omit θ0\theta_{0} throughout this work and will simply write θ⊤1Si\theta^{\top}\bm{1}_{S_{i}} to represent a datamodel prediction for the set SiS_{i}.

In step (a) of the estimation process above, we collect a “datamodel training set” by sampling subsets Si⊂SS_{i}\subset S from a distribution DS\mathcal{D}_{S}. A simple first choice for DS\mathcal{D}_{S}—and indeed, the one we consider for the remainder of this work—is the distribution of random α\alpha-fraction subsets of the training set. Formally, we set

This design choice reduces the choice of DS\mathcal{D}_{S} to a choice of subsampling fraction α∈(0,1)\alpha\in(0,1), a decision whose impact we explore in Section 5. In practice, we estimate datamodels for several choices of α\alpha, as it turns out that the value of α\alpha corresponding to the most useful datamodels can vary by setting.

Recall that for any subset S′⊂SS^{\prime}\subset S of the training set SS, fA(x;S′)f_{\mathcal{A}}({x};{S^{\prime}}) is intended to be a specific (potentially stochastic) function representing the output of a model trained on S′S^{\prime} and evaluated on a target example xx. There are, however, several candidates for fA(x;S′)f_{\mathcal{A}}({x};{S^{\prime}}) based on which model output we opt to track.

In the context of understanding classifiers, perhaps the simplest such candidate is the correctness function (i.e., a stochastic function that is 11 if the model trained on S′S^{\prime} is correct on xx, and otherwise). However, while the correctness function may be a natural choice for fA(x;S′)f_{\mathcal{A}}({x};{S^{\prime}}), it turns out to be suboptimal in two ways. First, fitting to the correctness function ignores potentially valuable information about the model’s confidence in a given decision. Second, recall that our procedure fits model outputs using a least-squares linear model, which is not designed to properly handle discrete (binary) dependent variables.

A natural way to improve over our initial candidate would thus be to use continuous output function, such as cross-entropy loss or correct-label confidence. But which exact function should we choose? In Appendix C, we describe a heuristic that we use to guide our choice of the correct-class margin:

where we recall that dd is the size of the original training set SS. We can use cross-validation to select the regularization parameter λ\lambda for each specific target example xx.

Accurately predicting outputs with datamodels

We now demonstrate how datamodels can be applied in the context of deep neural networks—specifically, we consider deep image classifiers trained on two standard datasets: CIFAR-10 [Kri09a] and Functional Map of the World (FMoW) [KSM+20a] (see Appendix D.1 for more information on each dataset).

1 Implementation details

Before applying datamodels to our two tasks of interest, we address a few remaining technical aspects of datamodel estimation:

Rather than repeat the entire datamodel estimation process for each target example xx of interest separately, we can estimate datamodels for an entire set of target examples simultaneously through model reuse. Specifically, we train a large pool of models on subsets Si⊂SS_{i}\subset S sampled from the distribution DS\mathcal{D}_{S}, and use the same models to compute outputs fA(x;Si)f_{\mathcal{A}}({x};{S_{i}}) for each target example xx.

The cost of obtaining a single subset-output pair can be non-trivial—in our case, it involves training a ResNet from scratch on CIFAR-10. It turns out, however, that recent advances in fast neural network training [Pag18a, LIE+22a] allow us to train a wealth of models on different α\alpha-subsets of each dataset very efficiently. (For example, for α=50%\alpha=50\% we can use [LIE+22a] to train 40,000 models/day on an 8×A1008\times\text{A}100 GPU machine; see Appendix D.2 for details.) We train m=300,000m=300,000 CIFAR models and m=150,000m=150,000 FMoW models on α=50%\alpha=50\% subsets of each dataset. We also train mm models for each subsampling fraction α∈{10%,20%,75%}\alpha\in\{10\%,20\%,75\%\}, using α\alpha to scale mm. See Table 1 for a summary of the models trained.

Recall that the target example xx for which we estimate a datamodel can be arbitrary. In particular, xx could itself be a training example—indeed, as we mention above, our goal is to estimate a datamodel for every image in the FMoW and CIFAR-10 test and training sets. When xx is in the training set, however, we slightly alter the datamodel estimation objective (8) to exclude training sets SiS_{i} containing the target example:

However, most readily available LASSO solvers require too much memory or are prohibitively slow for our values of nn (the number of datamodels to estimate), mm (the number of models trained and thus the size of the datamodel training set of subset-output pairs), and dd (the size of the original tasks training set and thus the input dimensionality of the regression problem in (8)). We therefore built a custom solver leveraging the works of [WSM21a] and [LIE+22a]—details of our implementation are in Appendix E.1.

2 Results: linear datamodels can predict deep network training

For both datasets considered (CIFAR-10 and FMoW), we minimize objectives (8) (respectively, (9)) yielding a datamodel gθig_{\theta_{i}} for each example xix_{i} in the test set (respectively, training set). We now assess the quality of these datamodels in terms of how well they predict model outputs on unseen subsets (i.e., fresh samples from DS\mathcal{D}_{S}). We refer to this process as on-distribution evaluation because we are interested in subsets SiS_{i}, sampled from the same distribution DS\mathcal{D}_{S} as the datamodel training set, but not the exact ones used for estimation. (In fact, recall that we explicitly held out mtestm_{test} subset-output pairs for evaluation in Section 2.)

We next study the dependence of datamodel estimation on the size of the datamodel training set mm. Specifically, we can measure the on-distribution average mean-squared error (MSE) as

To evaluate (10), we replace the inner expectation with an empirical average, again using a heldout set of samples that was not used for estimation.

Note that OPT is independent of the estimator gθg_{\theta} and measures only the inherent variance in the prediction problem, i.e., loss that will necessarily be incurred due only to inherent noise in deep network training.

Finally, in Appendix E.2, we study the effect of the regularization parameter λ\lambda (cf. (8) and (9)) on datamodel performance. In particular, in Figure 19 we plot the variation in average MSE, on both on-distribution subsets (i.e., the exact subsets that we used to optimize (8)) and unseen subsets, as we vary the regularization parameter λ\lambda in (8). We find that—as predicted by classical learning theory—setting λ=0\lambda=0 leads to overfit datamodels, i.e., estimators gθg_{\theta} that perform well on the exact subsets that were used to estimate them, but are poor output predictors on new subsets SiS_{i} sampled from DS\mathcal{D}_{S}. (In fact, using m=300,000m=300,000 trained models with λ=0\lambda=0 results in higher MSE than using only m=10,000m=10,000 with optimal λ\lambda, i.e., the left-most datapoint in Figure 1).

Leveraging datamodels

Now that we have introduced (Section 1), instantiated (Section 2), and implemented (Section 3) the datamodeling framework, we turn to some of its applications. Specifically, we will now show how to apply datamodels within three different contexts:

We originally constructed datamodels to predict the outcome of training a model on random (α\alpha-)subsets of the training set. However, it turns out that we can also use datamodels to predict model outputs on arbitrary subsets (i.e., subsets that are “off-distribution” from the perspective of the datamodel prediction task). To illustrate the utility of this capability, we will use datamodels to (a) identify predictions that are brittle to removal of relatively few training points, and (b) estimate data counterfactuals, i.e., the causal effects of removing groups of training examples.

We demonstrate that datamodels can identify, for any given target example xx, a set of visually similar examples in the training data. Leveraging this ability, we will identify instances of train-test leakage, i.e., when test examples are duplicated (or nearly duplicated) within the training set.

1 Counterfactual prediction

So far, we have computed and evaluated datamodels entirely within a supervised learning framework. In particular, we constructed datamodels with the goal of predicting the outcome of training on random subsets of the training set (sampled from a distribution DS\mathcal{D}_{S} (6)) and evaluating on a fixed target example xx. Accordingly, for each target example xx, we evaluated its datamodel gθg_{\theta} by (a) sampling new random subsets SiS_{i} (from the same distribution); (b) training (a neural network) on each one of these subsets; (c) measuring correct-class margin on the target example xx; and (d) comparing the results to the datamodel’s predictions (namely, gθ(Si)g_{\theta}(S_{i})) in terms of expected mean-squared error (see (10)) over the distribution of subsets.

We will now go beyond this framework, and use datamodels to predict the outcome of training on arbitrary subsets of the training set. In particular, consider a fixed target example xx with corresponding datamodel gθg_{\theta}. For any subset S′S^{\prime} of the training set SS, we will use the datamodel-predicted outcome of training on S′S^{\prime} and evaluating on xx, i.e., gθ(1S′),g_{\theta}(\bm{1}_{S^{\prime}}), in place of the ground-truth outcome fA(x;S′).f_{\mathcal{A}}({x};{S^{\prime}}). Since S′S^{\prime} is an arbitrary subset of the training set, it is “out-of-distribution” with respect to the distribution of fixed-size subsets DS\mathcal{D}_{S} that we designed the datamodel to operate on. As such, using datamodel predictions in place of end-to-end-model training in this manner is not a priori guaranteed to work. Nevertheless, we will demonstrate through two applications that datamodels can in fact be effective proxies for end-to-end model training, even for such out-of-distribution subsets.

[Proxy for end-to-end training] We can use datamodel predictions as an efficient, closed-form proxy for end-to-end model training. That is, for a test example xx with datamodel gθg_{\theta}, and an arbitrary subset S′S^{\prime} of the training set SS, we can leverage the approximation

We first illustrate the utility of datamodels as a proxy for model training by using them to answer the question: how brittle are model predictions to removing training data? While all useful learning algorithms are data-dependent, cases where model behavior is sensitive to just a few data points are often of particular interest or concern [BGM21a, DKM+06a]. To quantify such sensitivity, we define the data support \textscSupport(x)\textsc{Support}(x) of a target example xx as

Intuitively, examples with a small data support are the examples for which removing a small subset of the training data significantly changes model behavior, i.e., they are “brittle” examples by our criterion of interest. By computing \textscSupport(x)\textsc{Support}(x) for every image in the test set, we can thus get an idea of how brittle model predictions are to removing training data.

One way to compute \textscSupport(x)\textsc{Support}(x) for a given target example xx would be to train several models on every possible subset of the training set SS, then report the largest subset for which the example was misclassified on average—the complement of this set would be exactly \textscSupport(x)\textsc{Support}(x). However, exhaustively computing data support in this manner is simply intractable.

Using datamodels as a proxy for end-to-end model training provides an (efficient) alternative approach. Specifically, rather than training models on every possible subset of the training set, we can use datamodel-predicted outputs gθ(S′)g_{\theta}(S^{\prime}) to perform a guided search, and only train on subsets for which predicted margin on the target example is small. This strategy (described in detail in Algorithm 1 and in Appendix F.3) allows us to compute estimates of the data support while training only a handful of models per target example.

We apply our algorithm to estimate \textscSupport(x)\textsc{Support}(x) for 300 random target examples in the CIFAR-10 test set. For over 90% of these 300 examples, we are able to certify that our estimated data support is strictly larger than the true data support \textscSupport(x)\textsc{Support}(x) (i.e., that we are not over-estimating brittleness) by training several models after excluding the estimated data support and checking that the target example is indeed misclassified on average.

We plot the distribution of estimated data support sizes in Figure 7. Around half of the CIFAR-10 test images have a datamodel-estimated data support comprising 250 images or less, meaning that removing a specific 0.4% of the CIFAR-10 training set induces misclassification. Similarly, 20% of the images had an estimated data support of less than 40 training images (which corresponds to 0.08% of the training set).

To contextualize these findings, we compare our estimates of data support to a few natural baselines. We provide the exact comparison setup in Appendix F.3.1: in summary, each baseline technique can be cast as a swap-in alternative to datamodels for guiding the data support search described above.

It turns out that every baseline we tested provides much looser estimates of data support (Figure 7). For example, even the best-performing baseline predicts that one would need to remove over 600 training images per test image to force misclassification on 20% of the test setMoreover, the data support estimates derived from the baselines are only “certifiable” in the above-described sense (see the beginning of the “Results” paragraph) for 60% of the 300 test examples we study (as opposed to 90% for datamodel-derived estimates).. In contrast, our datamodel-guided estimates indicate that removing 40 train examples is sufficient for misclassifying 20% of test examples.

Note that the brittleness we consider in this section (i.e., brittleness to removing training examples) is substantively different than brittleness to mislabeling examples (as in label-flipping attacks [KL17a, XXE12a, RWR+20a]). In particular, brittleness to removal indicates that there exists a small set of training images whose presence is necessary for correct classification of the target example (thus motivating the term “data support”). Meanwhile, label-flipping attacks can succeed even when the target example has a large data support, as (consistently) mislabeling a set of training examples provide a much stronger signal than simply removing them. Nevertheless, we can easily adapt the above experiment to test brittleness to mislabeling—we do so in Appendix F.4. As one might expect, test predictions are even more brittle to data mislabeling than removal—for 50% of the CIFAR-10 test set, mislabeling 35 target-specific training examples suffices to flip the corresponding prediction (see Figure 24 for a CDF).

As we have already seen, a simple application of datamodels as a proxy for model training (on arbitrary subsets of the training set) enabled us to identify brittle predictions. We now demonstrate another, more intricate application of datamodels as a proxy for end-to-end training: predicting data counterfactuals.

For a fixed target example xx, and a specific subset of the training set R(x)⊂SR(x)\subset S, a data counterfactual is the causal effect of removing the set of examples R(x)R(x) on model outputs for xx. In terms of our notation, this effect is precisely

Such data counterfactuals can be helpful tools for finding brittle predictions (as in the previous subsection), estimating group influence (as done by [KAT+19a] for linear models), and more broadly for understanding how training examples combine (through the lens of the model class) to produce test-time predictions.

Just as in the last section, we again use datamodels beyond the supervised learning regime in which they were developed. In particular, we predict the outcome of a data counterfactual as

where again gθg_{\theta} is the datamodel for a given target example of interest. Since gθg_{\theta} is a linear function in our case, the above predicted data counterfactual actually simplifies to

Our goal now is to demonstrate that datamodels are useful predictors of data counterfactuals across a variety of removed sets R(x)R(x). To accomplish this, we use a large set of target examples. Specifically, for each such target example, we consider different subset sizes kk; for each such kk, we use a variety of heuristics to select a set R(x)R(x) comprising kk “examples of interest.” These heuristics are:

setting R(x)R(x) to be the nearest kk training examples to the target example xx in terms of influence score [KL17a], TracIn score [PLS+20a], or distance in pre-trained representation space [BCV13a]Note that these methods are precisely the ones used as baselines in the previous section.;

setting R(x)R(x) to be the maximizer of the datamodel-predicted counterfactual, i.e.,

(Note that since our datamodels are linear, this simplifies to excluding the training examples corresponding to the top kk coordinates of the datamodel parameter θ\theta.)

setting R(x)R(x) to be the training images corresponding to the bottom (i.e., most negative) kk coordinates of the datamodel weight θ\theta.

We consider six values of kk (the size of the removed subset) ranging from 1010 to 12801280 examples (i.e., 0.02%−2.6%0.02\%-2.6\% of the training set). Thus, the outcome of our procedure is, for each target example, both true and datamodel-predicted data counterfactuals for 30 different training subsets R(x)R(x) (six values of kk and five different heuristics).

In Figure 8, we plot datamodel-predicted data counterfactuals against true data counterfactuals, aggregating across all target examples xx, values of kk, and selection heuristics for R(x)R(x). We find a strong correlation between these two quantities. In particular, across all factors of variation, predicted and true data counterfactuals have Spearman correlation ρ=0.98\rho=0.98 and ρ=0.94\rho=0.94 for CIFAR-10 and FMoW respectively. In fact, the two quantities are correlated roughly linearly: we obtain (Pearson) correlations of r=0.96r=0.96 (CIFAR-10) and r=0.90r=0.90 (FMoW) between counterfactuals and their estimates on aggregate. Correlations are even more pronounced when restricting to any single class of removed sets (i.e., any single hue in Figure 8).

We have seen that datamodels accurately predict the outcome of many natural data counterfactuals, despite only being constructed to predict outcomes for random subsets of a fixed size (α⋅d\alpha\cdot d for α∈(0,1)\alpha\in(0,1) and dd the training set size). Of course, due to both estimation error (i.e., we might not have trained enough models to identify optimal linear datamodels) and misspecification error (i.e., the optimal datamodel might not be linear), we don’t expect a perfect correspondence between datamodel-predicted outputs gθ(1S′)g_{\theta}(\bm{1}_{S^{\prime}}) and true outputs f(x;S)f(x;S) for all 2d2^{d} possible subsets of the training set. Indeed, this is part of the reason why we estimated datamodels for several values of α\alpha, only one of which is shown in Figure 8. As shown in Appendix F.9, datamodels estimated for other values of α\alpha still display strong correlation between true and predicted model outputs, but behave qualitatively differently than the ones shown above (i.e., each value of α\alpha is better or worse at predicting the outcomes of certain types of counterfactuals).

2 Using datamodels to find similar training examples

We now turn to another application of datamodels: identifying training examples that are similar to a given test example. One can use this primitive to identify issues in datasets such as duplicated training examples [LIN+21a] or train-test leakage [BD20a] (test examples that have near-duplicates in the training set).

Recall that in our instantiation of the framework, datamodels predict model output (for a fixed target example) as a linear function of the presence of each training example in the training set. That is, we predict the output of training on a subset S′S^{\prime} of the training set SS as

A benefit of parameterizing datamodels as simple linear functions is that we can use the magnitude of the coordinates of θ\theta to ascertain feature importance [GE03a]. In particular, since in our case each feature coordinate (i.e., each coordinate of 1S′\bm{1}_{S^{\prime}}) actually represents the presence of a particular training example, we can interpret the highest-magnitude coordinates of θ\theta as the indices of the training examples whose presence (or absence) is most predictive of model behavior (again, on the fixed target example in context).

We now show that these high-magnitude training examples (a) they visually resemble the target image, yielding a method for finding similar training examples to a given target; and (b) as a result, datamodels can automatically detect train-test leakage. {appmode}[Train-test similarity] For a test example xx with a linear datamodel gθg_{\theta}, we can interpret the training examples corresponding to the highest-magnitude coordinates of θ\theta as the “nearest neighbors” of xx.

Motivated by the feature importance perspective described above, we visualize (in Figures 9 and 33) a random set of target examples from the CIFAR-10 test set together with the CIFAR-10 training images that correspond to the highest-magnitude datamodel coordinates for each test image.

We find that for a given target example, the highest-magnitude datamodel coordinates—both positive and negative—consistently correspond to visually similar training examples.

Furthermore, the exact training images that are surfaced by looking at high-magnitude weights differ depending on the subsampling parameter α\alpha that we use while constructing the datamodels. (Recall from Section 2 that α\alpha controls the size of the random subsets used to collect the datamodel training set—a datamodel estimated with parameter α\alpha is constructed to predict outcomes of training on random training subsets of size α⋅d\alpha\cdot d, where dd is the training set size.) In Figure 10 (and 34), we consider a pair of target examples from the CIFAR-10 test set, and, for each target example, compare the top training images from two different datamodels: one estimated using α=10%\alpha=10\%, and the other using α=50%\alpha=50\%. We find that in some cases (e.g., Figure 10 left), the α=10%\alpha=10\% datamodel identifies training images that are highly similar to the target example but do not correspond to the highest-magnitude coordinates for the α=50%\alpha=50\% datamodel (in other cases, the reverse is true). Our hypothesis here—which we expand upon in Section 5—is that datamodels estimated with lower α\alpha (i.e., based on smaller random training subsets) find train-test relationships driven by larger groups of examples (and vice-versa).

Another method for finding similar training images is influence functions, which aim to estimate the effect of removing a single training image on the loss (or correctness) for a given test image. A standard technique from robust statistics [HRR+11a] (applied to deep networks by [KL17a]) uses first-order approximation to estimate influence of each training example. We find (cf. Appendix Figure 35), that the high-influence and low-influence examples yielded by this approximation (and similar methods) often fail to find similar training examples for a given test example (also see [BPF21a, HYH+21a]).

Another approach based on empirical influence approximation was used by [FZ20a], who (successfully) use their estimates to identify similar train-test pairs in image datasets as we do above. We discuss empirical influence approximation and its connection with datamodeling in Section 6.1.

We now leverage datamodels’ ability to surface training examples similar to a given target in order to identify same-scene train-test leakage: cases where test examples are near-duplicates of, or clearly come from the same scene as, training examples. Below, we use datamodels to uncover evidence of train-test leakage on both CIFAR and FMoW, and show that datamodels outperform a natural baseline for this task.

To find train-test leakage in CIFAR-10, we collect ten candidate training examples for each image in the CIFAR-10 test set—the ones corresponding to the ten largest coordinates of the test example’s datamodel parameter. We then show crowd annotators (using Amazon Mechanical Turk) tasks that consist of a random CIFAR-10 test example accompanied by its candidate training examples. We ask the annotators to label any of the candidate training images that constitute instances of same-scene leakage (as defined above). We show each task (i.e., each test example) to multiple annotators, and compute the “annotation score” for each of the test example’s candidate training examples as the fraction of annotators who marked it as an instance of leakage. Finally, we compute the “leakage score” for each test example as the highest annotation score (over all of its candidate train images). We use the leakage score as a proxy for whether or not the given image constitutes train-test leakage.

In Figure 11, we plot the distribution of leakage scores over the CIFAR-10 test set, along with random train-test pairs stratified by their annotation score. As the annotation score increases, pairs (qualitatively) appear more likely to correspond to leakage (see Appendix H for more pairs). Furthermore, roughly 10% of test set images were labeled as train-test leakage by over half of the annotators that reviewed them.

To identify train-test leakage on FMoW, we begin with the same candidate-finding process that we used for CIFAR-10. However, FMoW differs from CIFAR in that the examples (satellite images labeled by category, e.g., “port” or “arena”) are annotated with geographic coordinates. These coordinates allow us to avoid crowdsourcing—instead, we compute the geodesic distance between the test image and each of the candidates, and use a simple threshold dd (in miles) to decide whether a given test example constitutes train-test leakage.

Furthermore, we can calculate a “ground-truth” number of train-test leakage instances by counting the test examples whose geodesic nearest-neighbor in the training set is within the specified threshold dd. It turns out that despite having already been de-duplicated, about 20% and 80% of FMoW test images are within 0.25 and 2.6 miles of a training image, respectively—see Appendix Figure 40. Comparing this ground truth to the number of instances of leakage found within the candidate examples yields a qualitative measure of the efficacy of our method (i.e., the quality of candidates we generate).

In Figure 12, we plot this measure of efficacy (# instances found / # ground truth) as a function of the threshold dd, and also visualize examples images from the FMoW test set together with their corresponding datamodel-identified training set candidates. To put our quantitative results into context, we compare the efficacy of candidates derived from top datamodel coordinates (i.e., the ones we use here and for CIFAR-10) to that of candidates derived from nearest neighbors in the representation space of a pretrained neural network [BCV13a, ZIE+18a] (examining such nearest neighbors is a standard way of finding train-test leakage, e.g., used by [BD20a] to study CIFAR-10 and CIFAR-100). Datamodels consistently outperform this baseline.

3 Using datamodels as a feature embedding

Sections 4.1 and 4.2 illustrate the utility of datamodels on a per-example level, i.e., for predicting the outcome of training on arbitrary training subsets and evaluating on a specific target example, or for finding similar training images (again, to a specific target). We’ll conclude this section by demonstrating that datamodels can also help uncover global structure in datasets of interest.

In this section we demonstrate, through two applications, the potential for such datamodel embeddings to discover dataset structure in this way. In Section 4.3.2, we use datamodel embeddings to partition datasets into disjoint clusters, and in Section 4.3.1 we use principal component analysis to get more fine-grained insights into dataset structure. To emphasize our shift in perspective (i.e., from θ\theta being just a parameter of a datamodel gθg_{\theta}, to θ\theta being an embedding for the target example xx), we introduce an embedding function φ(x)↦θ\varphi(x)\mapsto\theta which maps a particular target example to the weights of its corresponding datamodel.

We begin with a simple application of datamodel embeddings, and show that they enable high-quality clustering. Specifically, given two examples x1x_{1} and x2x_{2}, datamodel embeddings induce a natural similarity measure between them:

Finally, we can view this similarity matrix as an adjacency matrix for a (dense) graph connecting all the examples {x1,…xk}\{x_{1},\ldots x_{k}\}: the edge between two examples will be d(xi,xj)d(x_{i},x_{j}), which is in turn the kernelized inner product between their two datamodel weights. We expect similar examples to have high-weight edges between them, and unrelated examples to have (nearly) zero-weight edges between them.

Such a graph unlocks a myriad of graph-theoretic tools for exploring datasets through the lens of datamodels (e.g., cliques in this graph should be examples for which model behavior is driven by the same subset of training examples). However, a complete exploration of these tools is beyond the scope of our work: instead, we focus on just one such tool: spectral clustering.

At a high level, spectral clustering is an algorithm that takes as input any similarity graph GG as well as the number of clusters CC, and outputs a partitioning of the vertices of GG into CC disjoint subsets, in a way that (roughly) minimizes the total weight of inter-cluster edges. We run an off-the-shelf spectral clustering algorithm on the graph induced by the similarity matrix AA above for the images in the CIFAR-10 test set. The result (Figure 13 and Appendix I) demonstrates a simple unsupervised method for uncovering subpopulations in datasets.

We observed above that datamodel embeddings encode enough information about their corresponding examples to cluster them into (at least qualitatively) coherent groups. We now attempt to gain even further insight into the structure of these datamodel embeddings, in the hopes of shedding light on the structure of the underlying dataset itself.

Datamodel embeddings are both high-dimensional and sparse, making analyzing them directly (e.g., by looking at the variation of each coordinate) a daunting task. Instead, we leverage a canonical tool for finding structure in high-dimensional data: principal component analysis (PCA).

each of the kk coordinates of the transformed embeddings is a (fixed) linear combination of the coordinates of the initial datamodel embeddings, i.e., φ~(x)=M⋅φ(x)\widetilde{\varphi}(x)=\bm{M}\cdot\varphi(x) for a fixed k×dk\times d matrix M\bm{M};

Note that in (a), the ii-th coordinate of a transformed embedding is always the same linear combination of the corresponding original embedding (and thus, each coordinate of the transformed embedding has a concrete interpretation as a weighted combination of datamodel coefficients). The exact coefficients of this combination (i.e., the rows of the matrix M\bm{M} above) are called the first kk principal components of the dataset.

Our point of start in analyzing these transformed embeddings is to examine each transformed coordinate separately. In particular, in Figure 14 we visualize, for a few sample coordinate indices i∈[k]i\in[k], the target examples whose transformed embeddings have particularly high or low values of the given coordinate (equivalently, these are the target examples whose datamodel embeddings have the highest or lowest projections onto the ii-th principal component). We find that:

the examples whose transformed embeddings have a large ii-th coordinate all (visually) share a common feature: e.g., the first-row images in Figure 14 share similar pose and color composition;

this (visual) feature is consistent across both train and test set examplesRecall that we computed the PCA transformation to preserve the information in only the training set datamodel embeddings. Thus, this result suggests that the transformed embeddings computed by PCA are not “overfit” to the specific examples that we used to compute it.; and

for a given coordinate, the most positive images and most negative images (i.e., the left and right side of each row of Figure 14, respectively) either (a) have a differing label but share the same common feature or (b) have the same label but differ along the relevant feature.

In Appendix J, we verify that not only are the groups of images found by PCA visually coherent, they are in fact rooted in how the model class makes predictions. In particular, we show that one can find, for any coordinate i∈[k]i\in[k] of the transformed embedding, the training examples that are most important to that coordinate. Furthermore, retraining without these examples significantly decreases (increases) accuracy on the target examples with the most positive (negative) coordinate ii, suggesting that the identified principal components actually reflect model class behavior.

In the context of deep neural networks, the word “embedding” typically refers to features extracted from the penultimate layer of a fixed pre-trained model (see [BCV13a] for an overview). These “deep representations” can serve as an effective proxy for visual similarity [BD20a, ZIE+18a], and also enable a suite of applications such as clustering [GGT+17a] and feature visualization [OMS17a, EIS+19a, ARS+15a, BBC+07a].

Here, we briefly discuss a few advantages of datamodel-based embeddings over their standard penultimate layer-based counterparts.

Axis-alignment: datamodel embeddings are axis-aligned—each embedding component directly corresponds to index into the training set, as opposed to a more abstract or qualitative concept. As a corollary, aggregating or comparing different datamodel embeddings for a given dataset is straightforward, and does not require any alignment tools or additional heuristics. This is not the case for network-based representations, for which the right way to combine representations—even for two models of the same architecture—is still disagreed upon [KNL+19a, BNB21a]. In particular, we can straightforwardly compare datamodel embeddings across different target examples, model architectures, training paradigms, or even datamodel estimation techniques—as long as the set of training examples being stays the same, any resulting datamodel has a uniform interpretation.

Richer representation: the space of datamodel embeddings seems significantly richer than that of standard representation space. In particular, Appendix Figure 43 shows that for standard representation space, ten linear directions suffice to capture 90% of the variation in training set representations. The “effective dimension” of datamodel representations is much higher, with the top 500 principal components explaining only 50% of the variation in training set datamodel embeddings. This difference manifests qualitatively when we redo our PCA study on standard representations (Appendix Figure 46): principal components beyond the 10th lack both the perceptual quality and train-test consistency exhibited by those of datamodel embeddings (e.g., for datamodels even the 76th principal component, shown in Figure 14, exhibits these qualities).

Ingrained causality: datamodel embeddings inherently encode information about how the model class generalizes. Indeed, in Section 4.3.2 we verified via counterfactuals that insights extracted from the principal components of Θ\Theta actually reflect underlying model class behavior.

Discussion: The role of the subsampling fraction α𝛼\alpha

We have used datamodels estimated using several choices of the subsampling fraction α\alpha, and saw that the value of α\alpha corresponding to the most useful datamodels can vary by setting. In particular, the visualizations in Figure 10 suggest that datamodels estimated with lower α\alpha (i.e., based on smaller random training subsets) find train-test relationships driven by larger groups of examples (and vice-versa). Here, we explore this intuition further using thought experiment, toy example, and numerical simulation. Our goal is to intuit how different choices of α\alpha can lead to substantively different datamodels.

First, consider the task of estimating a datamodel for a prototypical image xx—for example, a plane on a blue sky background. As α→1\alpha\to 1, the sets SiS_{i} sampled from DS\mathcal{D}_{S} are relatively large—if these sets have enough other images of planes on blue skies, we will observe little to no variation in fA(x;Si)f_{\mathcal{A}}({x};{S_{i}}), since any predictor trained on SiS_{i} will perform very well on xx. As a result, a datamodel for xx estimated with α→1\alpha\to 1 may assign very little weight to any particular image, even if in reality their total effect is actually significant.

Decreasing α\alpha, then, offers a solution to this problem. In particular, we allow the datamodel to observe cases where entire groups of training examples are not present, and re-distribute the corresponding effect back to the constituents of the group (i.e., assigning them all a share of the weight).

Now, consider a highly atypical yet correctly classified example, whose correctness relies on just the presence of just a few images from the training set. In this setting, datamodels estimated with a small value of α\alpha may be unable to isolate these training points, since they will constantly distribute variation in fA(x;Si)f_{\mathcal{A}}({x};{S_{i}}) among a large group of non-present images. Meanwhile, using a large value of α\alpha allows the estimated datamodel to place weight on the correct training images (since xx will be classified correctly until some of the important training images are not present in SiS_{i}).

In line with this intuition, decreasing α\alpha in Figure 15 (i.e., moving from right to left) leads to datamodels that assign weight to increasingly large neighborhoods of points around the target input. This example and the above reasoning lead us to hypothesize that larger (respectively, smaller) α\alpha are better-suited to cases where model predictions are driven by smaller (respectively, larger) groups of training examples. In Appendix B, we perform a more quantitative analysis of the role of α\alpha, this time by studying an underdetermined linear regression model on data that is organized into overlapping subpopulations. Our findings in this setting (see Figure 16) mirror our intuition thus far—in particular, smaller values of α\alpha result in datamodels that were more predictive on larger subpopulations in the training set, whereas higher values of α\alpha tended to work better smaller subpopulations.

Related work

Datamodels build on a rich and growing body of literature in machine learning, statistics, and interpretability. In this section, we illustrate some of the connections to these fields, highlight a few of the most closely related works to ours.

We start by discussing the particularly important connection between datamodels and another well-studied concept that has recently been applied to the machine learning setting: influence estimators. In particular, a recent line of work aims to compute the empirical influence [HRR+11a] of training points xix_{i} on predictions f(xj)f(x_{j}), i.e.,

where randomness is taken over the training algorithm. Evaluating these influence functions naively requires training C⋅dC\cdot d models where dd is again the size of the train set and CC is the number of samples necessary for an accurate empirical estimate of the probabilities above. To circumvent this prohibitive sample complexity, a recent line of work has proposed approximation schemes for Infl[xi→xj]\text{Infl}[x_{i}\to x_{j}]. We discuss these approximations (and their connection to our work) more generally in Section 6.2, but here we focus on a specific approximation used by [FZ20a] (and in a similar form, by [GZ19a] and [JDW+19a])In fact, (15) is ubiquitous—e.g., in causal inference, it is called the average treatment effect of training on xix_{i} on the correctness of xjx_{j}.:

This estimator improves sample efficiency by reusing the same set of models to compute influences between different input pairs. More precisely, [FZ20a] show that the size of the random subsets trades off sample efficiency (model reuse is maximized when the subsets are exactly half the size of the training set, since this maximizes the number of samples available to estimate each term in (15)) and accuracy with respect to the true empirical influence (which is maximized as the subsets SiS_{i} get larger). Despite its different goal, formulation, and estimation procedure, it turns out that we can cast the difference-of-probabilities estimator (15) above as a rescaled datamodel (in the infinite-sample limit). In particular, in Appendix K.1 we show:

We illustrate this result quantitatively in Appendix K and perform an in-depth study of influence estimators as datamodels. As one might expect given their different goal, influence estimates significantly underperform explicit datamodels in terms of predicting model outputs with respect to every metric we studied (Table 5, Figure 49). We then attempt to explain this performance gap and reconcile it with Lemma 6.1 in terms of the estimation algorithm (OLS vs. LASSO), scale (number of models trained), and output function (0/1 loss vs. margins).

In addition to forging a connection between datamodels and influence estimates, this result also provides an alternate perspective on the parameter α\alpha. Specifically, in light of our discussion in Section 5, it suggests that α\alpha may control the kinds of correlations that are surfaced by empirical influence estimates.

2 Other connections

Above, we contrasted datamodels with empirical influence functions, which measure the counterfactual effect of removing individual training points on a given model output. Specifically, in that section and the corresponding Appendix K, we discussed the subsampled influence estimator of [FZ20a], who use influences to study the memorization behavior of standard vision models. We now provide a brief overview of a variety of other methods for influence estimation developed in prior works.

First-order influence functions are a canonical tool in robust statistics that allows one to approximate the impact of removing a data point on a given parameter without re-estimating the parameter itself [HRR+11a]. [KL17a] apply influence functions to both a variety of classical machine learning models and to penultimate-layer embeddings from neural network architectures, to trace model’s predictions back to individual training examples. In classical settings (namely, for a logistic regression model), [KAT+19a] find that influence functions are also useful for estimating the impact of groups of examples. On the other hand, [BPF21a] finds that approximate influence functions scale poorly to deep neural network architectures; and [FZ20a] argue that understanding the dynamics of the penultimate layer is insufficient for understanding deep models’ decision mechanisms. Other methods for influence approximation (or more generally, instance-level attribution) include gradient-based methods [PLS+20a] and metrics based on representation similarity [CGF+19a, YKY+18a]—see [HYH+21a] for a more detailed overview. Finally, another related line of work [GZ19a, JDW+19a, WZJ+21a] uses Shapley values [Sha51a] to assign a value to datapoints based on their contribution to some aggregate metric (e.g., test accuracy).

As discussed in Section 6.1, datamodels serve a different purpose to influence functions—the former constructs an explicit statistical model, whereas the latter measures the counterfactual value of each training point. Nevertheless, we find that wherever efficient influence approximations and datamodels are quantitatively comparable (e.g., see Section 4.1 or Appendix K) datamodels predict model behavior better.

Datamodels are essentially surrogate models for the function mapping training data to predictions. Surrogate models from pixel-space to predictions are popular tools in machine learning interpretability [RSG16a, LL17a, SHS+19a]. For example, LIME [RSG16a] constructs a local linear model mapping test images to model predictions. Such surrogate models try to understand, for a fixed model, how the features of a given test example change the prediction. In contrast, datamodels hold the test example fixed and instead study how the images present in the training set change the prediction.

In addition to the advantages of our data-based view stated in Section 1, datamodels have two further advantages over pixel-level surrogate models: (a) a clear notion of missingness (i.e., it is easy to remove a training example but usually hard to “remove” pixels [SLL20a, JSW+22a]); and (b) globality of predictions—pixel-level surrogate models are typically accurate within a small neighborhood of a given input in pixel space, whereas datamodels model entire distribution over subsets of the training set, and remain useful both on- and off-distribution.

In other contexts, surrogate models are also used to evaluate data points for active learning and coreset selection [LC94a, CYM+20a]. [CYM+20a] find that shallow neural networks trained with fewer epochs can be a good proxy for a larger model when evaluating data for these applications.

Recall (from Section 1) that datamodels are, in part, inspired by the fact that re-training deep neural networks using the same data and model class leads to models with similar accuracies but vastly different individual predictions. This phenomenon has been observed more broadly. For example, [SYW+21a] make this point explicitly in the context of BERT [DCL+19a] pre-trained language models. Similarly, [NB20a] make note of this non-determinism for networks trained on the same training distribution (but not the same data), while [JNB+21a] find that the same is true for networks trained on the same exact data. [DHM+20a] find that on out-of-distribution data even overall accuracy is highly random. More closely to the spirit to our work, [ZGK+21a] find that non-determinism of individual predictions poses a challenge for comparing different model architectures. (They also propose a set of statistical techniques for overcoming this challenge.) More traditionally, the non-determinism is leveraged by Bayesian [Nea96a] and ensemble methods [LPB17a], which use a distribution over model weights to improve aspects of inference such as calibration of uncertainty.

Recent work (see [Fel19a, Cha18a, ZBH+16a, BN20a] and references therein) brings to light the interplay between learning and memorization, particularly in the context of deep neural networks. While memorization and generalization may seem to be at odds, the picture is more subtle. Indeed, [Cha18a] builds a network of small lookup tables on small vision datasets to show that purely memorization-based systems can still generalize-well. [Fel19a] suggests that memorization of atypical examples may be necessary to generalize well due to a long tail of subpopulations that arises in standard datasets. [FZ20a] find some empirical support for this hypothesis by identifying memorized images on CIFAR-100 and ImageNet and showing that removing them hurts overall generalization. Relatedly, [BBF+21a] proves that for certain natural distributions, memorization of a large fraction of data, even data irrelevant to the task at hand, is necessary for close to optimal generalization. For state of the art models, recent works (e.g., [CLK+19a, CTW+21a]) show that one can indeed extract sensitive training data, indicating models’ tendency to memorize.

Conversely, it has been observed that differentially private (DP) machine learning models—whose aim is precisely to avoid memorizing the training data—tend to exhibit poorer generalization than their memorizing counterparts [ACG+16a]. Moreover, the impact on generalization from DP is disparate across subgroups [BPS19a]. A similar effect has been noted in the context of neural network pruning [HCD+19a]. Datamodeling may be a useful tool for studying these phenomena and, more broadly, the mechanisms mapping data to predictions for modern learning algorithms.

A long line of work in statistics focuses on testing the robustness of statistical conclusions to the omission of datapoints. [BGM21a] study the robustness of econometric analyses to removing a (small) fraction of data. Their method uses a Taylor-approximation based metric to estimate the most influential subset of examples on some target quantity, similar in spirit to our use of datamodels to estimate data support for a target example (as in Figure 7). Datamodels may be a useful tool for extending such robustness analyses to the context of state-of-the-art machine learning models.

Future work

Our instantiation of the datamodeling framework yields both good predictors of model behavior and a variety of direct applications. However, this instantiation is fairly basic and thus leaves significant room for improvement along several axes. More broadly, datamodeling provides a lens under which we can study a variety of questions not addressed in this work. In this section, we identify (a subset of) these questions and provide connections to existing lines of work on them across machine learning and statistics.

Parameter estimation in the presence of such correlated outputs is an active area of research in statistics (see [DDP19a, LLZ19a] and references therein). Applying the corresponding techniques (or modifications thereof) to datamodels may help calibrate predictions and improve sample-efficiency.

Confidence intervals for datamodels. In this work we have focused on attaining point estimates for datamodel parameters via simple linear regression. A natural extension to these results would be to obtain confidence intervals around the datamodel weights. These could, for example, (a) provide interval estimates for model outputs rather than simple point estimates; and (b) decide if a training input is indeed a “significant” predictor for a given test input.

Post-selection inference. Relatedly, the high input-dimensionality of our estimation problem and the sparse nature of the solutions suggests that a two-stage procedure might improve sample efficiency. In such procedures, one first selects (often automatically, e.g., via LASSO) a subset of the coefficients deemed to be “significant” for a given test example, then re-fits a linear model for only these coefficients. This two-stage approach is particularly attractive in settings where the number of subset-output pairs (Si,fA(x;Si))(S_{i},f_{\mathcal{A}}({x};{S_{i}})) is less than the size of the training set ∣S∣|S| being subsampled.

Unfortunately, using the data itself to perform model selection in this manner—a paradigm known as post-selection inference—violates the assumptions of classical statistical inference (in particular, that the model class is chosen independently of the data) and can result in significantly miscalibrated confidence intervals. Applying valid two-stage estimation to datamodeling would be an area for further improvement upon the protocol presented in our work.

Improving subset sampling. Recall (cf. Section 2) that our framework uses a distribution over subsets DS\mathcal{D}_{S} to generate the “datamodel training set.” In this paper, we fixed DS\mathcal{D}_{S} to be random α\alpha-subsets of the training set, and used a nearest-neighbors example (see Figure 15) to provide intuition around the role of α\alpha. While this design choice did yield useful datamodels, it is unclear whether this class of distributions is optimal. In particular, a long line of literature in causal inference focuses on intervention design [ES07a]; drawing upon this line of work may lead to a better choice of subsampling distribution. Furthermore, one might even go beyond a fixed distribution DS\mathcal{D}_{S} and instead choose subsets SiS_{i} adaptively (i.e., based on the datamodels estimated with the previously sampled subsets) in order to reduce sample complexity.

2 Studying generalization

Datamodels also present an opportunity to study generalization more broadly:

Understanding linearity. The key simplifying assumption behind our instantiation of the datamodeling framework is that we can approximate the final output of training a model on a subset of the trainset as a linear function of the presence of each training point. While this assumption certainly leads to a simple estimation procedure, we have very little justification for why such a linear model should be able to capture the complexities of end-to-end model training on data subsets. However, we find that datamodels can accurately predict ground-truth model outputs (cf. Sections 2). In fact, we find a tight linear correlation between datamodel predictions and model outputs even on out-of-distribution (i.e., not in the support of DS\mathcal{D}_{S}) counterfactual datasets. Understanding why a simple linearity assumption leads to effective datamodels for deep neural networks is an interesting open question. Tackling this question may necessitate a better understanding of the training dynamics and implicit biases behind overparameterized training [BMR21a, SRK+20a].

Using sparsity to study generalization. A recent line of work in machine learning studies the interplay between learning, overparameterization, and memorization [Fel19a, Cha18a, ZBH+16a, BN20a, ZBH+20a]. Datamodeling may be a helpful tool in this pursuit, as it connects predictions of machine learning models directly to the data used to train them. For example, the data support introduced in Section 4.1.1 provides a quantitative measure of “how memorized” a given test input is.

Theoretical characterization of the role of α\alpha. In line with our intuitions in Section 5, we have observed both qualitatively (e.g., Figure 10) and quantitatively (e.g., Appendix B) that estimating datamodels using different values of α\alpha identifies correlations at varying granularities. However, despite empirical results around the clear role of α\alpha—Appendix B even isolates its effect on datamodels for simple underdetermined linear regression—we lack a crisp theoretical understanding of how α\alpha affects our estimated datamodels. A better theoretical understanding of the role of α\alpha, even for simple models trained on structured distributions, can provide us with more rigorous intuition for the phenomena observed here, and can in turn guide the development of better choices of sampling distribution for datamodeling.

3 Applying datamodels

Finally, each of the presented perspectives in Section 4 can be taken further to enable even better data and model understanding. For example:

Interpreting predictions. For a given test example, the training images corresponding to the largest-magnitude datamodel weights both (a) share features in common with the test example; and (b) seem to be causally linked to the test example (in the sense that removing the training images flips the test prediction). This immediately suggests the potential utility of datamodels as a tool for interpreting test-time predictions in a counterfactual-centric manner. Establishing them as such requires further evaluation through, for example, human-in-the-loop studies.

Building data exploration tools. In a similar vein, another opportunity for future work is in building user-friendly data exploration tools that leverage datamodel embeddings. In this paper we present the simplest such example in the form of PCA, but leave the vast field of data bias and feature discovery methods (cf. [CAS+19a] and [LSI+21a] for a survey) unexplored.

Conclusion

We present datamodeling, a framework for viewing the output of model training as a simple function of the presence of each training data point. We show that a simple linear instantiation of datamodeling enables us to predict model outputs accurately, and facilitates a variety of applications.

Acknowledgements

We thank Chiyuan Zhang and Vitaly Feldman for providing a set of 5,000 models with which we began our investigation. We also thank Hadi Salman for valuable discussions.

Work supported in part by the NSF grants CCF-1553428 and CNS-1815221, and Open Philanthropy. This material is based upon work supported by the Defense Advanced Research Projects Agency (DARPA) under Contract No. HR001120C0015.

References

References

figuresection tablesection algorithmsection

Appendix B Understanding the Role of α𝛼\alpha through Simulation

At a high level, our intuition for the subsampling fractionSee Section 2 for definition. α\alpha is that datamodels estimated with higher α\alpha tend to detect more local effects (i.e., those driven by smaller groups of examples, such as near-duplicates or small subpopulations), while those estimated with lower α\alpha detect more global effects (i.e., those driven by larger groups of images, such as large subpopulations or subclass biases). To solidify and corroborate this intuition about α\alpha, we analyze a basic simulated setting.

The feature coordinates are distributed as Bernoulli variables of varying frequency:

Each feature k∈[d]k\in[d] naturally defines a subpopulation SkS_{k}, the group of training examples with feature kk active, i.e., Sk≔{xi∈S:xik=1}S_{k}\coloneqq\{x_{i}\in S:x_{ik}=1\}. Features with lower (resp. higher) frequency pkp_{k} are intended to capture more local (resp. more global) effects.

The observed labels are generated according to a linear model y≔Xw+N(0,ϵ),y\coloneqq X\bm{w}+\mathcal{N}(0,\epsilon), where w\bm{w} is the true parameter vector and ϵ>0\epsilon>0 is a constant. We generate samples with d=150,n=125d=150,n=125 and use linear regressionAs the system is underdetermined, we use the pseudoinverse of XX to find the solution with the smallest norm. to estimate w\bm{w}.

Now, to use datamodels to analyze the above “training process” of fitting a linear regression model, we will model the output function fA(;)f_{\mathcal{A}}({};{}) given by the prediction of the linear model at point xjx_{j} when ww is estimated with samples S⊂SS\subset S, e.g.

We generate m=1,000,000m=1,000,000 subsampled training subsetsLarge sample size make sampling error negligible. along with their evaluations, and use ordinary least squares (OLS) to fit the datamodels. (Note that the use of OLS here is separate from the use of linear regression above as the original model class.)

The actual effect of removing the subpopulation SkS_{k} on xjx_{j}, i.e., fA(xj;S)−fA(xj;S∖Sk)f_{\mathcal{A}}({x_{j}};{S})-f_{\mathcal{A}}({x_{j}};{S\setminus S_{k}}),

The datamodel-predicted effect of removing SkS_{k}, i.e., ∑xi∈SΘij⋅1{xi∈Sk}\sum_{x_{i}\in S}\bm{\Theta}_{ij}\cdot\bm{1}\{{x_{i}\in S_{k}}\}.

To quantify the predictiveness of the datamodel at frequency pp, we compute the Pearson correlation between the above two quantities over all features kk with frequency pp and all test examples; see Algorithm 4 for a pseudocode. We repeat this evaluation varying pp and the datamodel (varying α\alpha). According to our intuition, for features kk with lower (resp. higher) frequency pkp_{k}, this correlation should be maximized at higher (resp. lower) values of α\alpha, where the datamodels capture more local (resp. global) effects. Figure 16 accurately reflects this intuition: more local (i.e., less frequent) features are best detected at higher α\alpha.

Appendix C Selecting Output Function to Model

In this section, we outline a heuristic method for selecting the output function fA(x;S)f_{\mathcal{A}}({x};{S}) to model. The heuristic is neither sufficient nor necessary for least-squares regression to work, but may provide some signal as to which output may yield better datamodels.

The first problem we would like to avoid is “output saturation,” i.e., being unable to learn a good datamodel due to insufficient variation in the output. This effect is most pronounced when we measure model correctness: indeed, over 30% of the CIFAR-10 test set is either always correct or always incorrect over all models trained, making datamodel estimation impossible. However, this issue is not unique to correctness. We propose a very simple test inspired by the idealized ordinary least squares model to measure how normally distributed a given type of model output is.

In the idealized ordinary least squares model, the observed outputs fA(x;S)f_{\mathcal{A}}({x};{S}) would follow a normal distribution with fixed mean (θ⋆)⊤1S(\theta^{\star})^{\top}\bm{1}_{S} and unknown variance, where θ⋆\theta^{\star} is the true parameter vector. Although we cannot guarantee this condition, we can measure the “normality” of the outputs (again, for a single fixed subset), with the intuition that the more normal the observed outputs are, the better a least-squares regression will work. Hence, compare different output functions by estimating the noise distribution of datamodels given each choice of output function. We leverage our ability—in contrast to typical settings for regression— to sample multiple response variables fA(x;S)f_{\mathcal{A}}({x};{S}) for a fixed SS (by retraining several models on the same data and recording the output on a fixed test example).

In Figure 17, we show the results of normality test for residuals arising from different choices of fA(⋅;S)f_{\mathcal{A}}({\cdot};{S}): correctness function, confidence on the correct class, cross-entropy loss, and finally correct-class marginCorrect-class margin is the difference between the correct-class logit and the highest incorrect-class logit; it is unbounded by definition, and its sign indicates the correctness of the classification.. Correct-class margins is the only choice of fA(⋅;S)f_{\mathcal{A}}({\cdot};{S}) where the pp-values are distributed nearly uniformly, which is consistent with the outputs being normally distributed. Hence, we choose to use the correct-class margins as the dependent variable for fitting our datamodels.

Appendix D Experimental Setup

We use the standard CIFAR-10 dataset [Kri09a].

FMoW [CFW+18a] is a land use classification dataset based on satellite imagery. WILDS [KSM+20a] uses a subset of FMoW and repurposes it as a benchmark for out-of-distribution (OOD) generalization; we use same the variant (presized to 224x224, single RGB image per example rather than a time sequence). We perform our analysis only on the in-distribution train/test splits (e.g. overlapping years) as our focus is not on OOD settings. Also, we limit our data to the year 2012. (These restrictions are only for convenience, and our framework can easily extend and scale to more general settings.)

Properties of both datasets are summarized in Table 2.

D.2 Models and hyperparameters

We use a ResNet-9 variant from Kakao Brainhttps://github.com/wbaek/torchskeleton/blob/master/bin/dawnbench/cifar10.py optimized for fast training. The hyperparameters (Table 3) were chosen using a grid search. We use the standard batch SGD. For data augmentation, we use random 4px random crop with reflection padding, random horizontal flip, and 8×88\times 8 CutOut [DT17a].

For counterfactual experiments with ResNet-18 (Figure 27), we use the standard variant [HZR+16a].

We use the standard ResNet-18 architecture [HZR+16a]. The hyperparameters (Table 3) were chosen using a grid search, including over different optimizers (SGD, Adam) and learning rate schedules (step decay, cyclic, reduce on plateau). As in [KSM+20a], we do not use any data augmentation. Unlike prior work, we do not initialize from a pre-trained ImageNet model; while this results in lower accuracy, this allows us to focus on the role of the FMoW dataset in isolation.

In Table 4, we show for each dataset the accuracies of the chosen model class (with its specific hyperparameters), across different values of α\alpha.

D.3 Training infrastructure

We train our models on a cluster of machines, each with 9 NVIDIA A100 GPUs and 96 CPU cores. We also use half-precision to increase training speed.

We use FFCV [LIE+22a], which removes the data loading bottleneck for smaller models and allows us achieve a throughput of over 5,000 CIFAR-10 models a day per GPU.

Our datamodel estimation uses (the characteristic vectors) of training subsets and model outputs (margins) on train and test sets. Hence, we do not need to store any model checkpoints, as it suffices to store the training subset and the model outputs after evaluating at the end of training. In particular, training subsets and model outputs can be stored as m×nm\times n or m×dm\times d matrices, with one row for each model instance and one column for each train or test example. All subsequent computations only require the above matrices.

Appendix E Regression

Note that solving large linear systems efficiently is an area of active research ([MT20a]), and as a result we anticipate that datamodel estimation could be significantly improved by applying techniques from numerical optimization. In this paper, however, we take a rather simple approach based on the SAGA algorithm of [GGS19a]. Our starting point is the GPU-enabled implementation of [WSM21a]—while this implementation terminated (unlike the CPU-based off-the-shelf solutions), the regressions are still prohibitively slow (i.e., on the order of several GPU-hours per single datamodel estimation). To address this, we make the following changes:

The first performance bottleneck turns out to be in dataloading. More specifically, SAGA is a minibatch-based algorithm: at each iteration, we have to read BB masks (50,000-dimensional binary vectors) and BB outputs (scalars) and move them onto the GPU for processing. If the masks are read from disk, I/O speed becomes a major bottleneck—on the other hand, if we pre-load the entire set of masks into memory, then we are not able to run multiple regressions on the same machine, since each regression will use essentially the entire RAM disk. To resolve this issue, we use the FFCV library [LIE+22a] for dataloading—FFCV is based on memory mapping, and thus allows for multiple processes to read from the same memory (combining the benefits of the two aforementioned approaches). FFCV also supports batch pre-loading and parallelization of the data processing pipeline out-of-the-box—adapting the SAGA solver to use FFCV cut the runtime significantly.

Next, we leverage the fact that the SAGA algorithm is trivially parallelizable across different instances (sharing the same input matrix), allowing us to estimate multiple datamodels at the same time. In particular, we estimate datamodels for the entire test set in one pass, effectively cutting the runtime of the algorithm by the test set size (e.g., 10,000 for CIFAR-10).

In order to parallelize across test examples, we need to significantly reduce the GPU memory footprint of the SAGA solver. We accomplish this through a combination of simple code optimization (e.g., using in-place operations rather than copies) as well as writing a few custom CUDA kernels to speed up and reduce the memory consumption of algorithms such as soft thresholding or gradient updating.

For each dataset considered, we chose a maximum λ\lambda: 0.010.01 for CIFAR-10 test, 0.10.1 for CIFAR-10 trainset, and 0.050.05 for FMoW datamodels. Next, we chose k=100k=100 logarithmically spaced intermediate values between (λ/100,λ)(\lambda/100,\lambda) as the regularization path. We ran one regression per intermediate λ\lambda, using m−50,000m-50,000 samples (where mm is as in the table in Figure 1 (right)) to estimate the parameters of the model and the remaining 50,00050,000 samples as a validation set. For each image in the test set, we select the λ\lambda corresponding to the best-performing predictor (on the heldout set) along the regularization path. We then re-run the regression once more using these optimal λ\lambda values and the full set of mm samples.

E.2 Omitted results

Appendix F Counterfactual Prediction with Datamodels

For all of our counterfactual experiments, we use a random sample of the respective test datasets. We select at random 300 test images for CIFAR-10 (class-balanced; 30 per class) and 100 test images for FMoW. For the CIFAR-10 baselines, we consider counterfactuals for a 100 image subset of the 300.

For CIFAR-10, we remove top k={10,20,40,80,160,320,640,1280}k=\{10,20,40,80,160,320,640,1280\} images and bottom k={20,40,80,160,320}k=\{20,40,80,160,320\} where applicable. For FMoW, we remove top and bottom k={10,k=\{10, 20,20, 40,40, 80,80, 160,320,640}160,320,640\}.

Each counterfactual (i.e., training models on a given training set S′S^{\prime}) is evaluated over TT trials to reduce the variance that arises purely from non-determinism in model training. We use T=20T=20 for CIFAR-10 and FMoW, and T=10T=10 for CIFAR-10 baselines. In Section F.6, we show that using sufficiently high TT is important for reducing noise.

F.2 Baselines

We describe the baseline methods used to generate data support estimates and counterfactuals. Each of the methods gives a way to select training examples that are most similar or influential to a target example. As in prior work [HYH+21a, PJW+21a], we consider a representative set of baselines spanning both methods based on representation similarity and gradient-based methods, such as influence functions.

In order to more fairly compare with datamodels—so that we can disentangle the variance reduction from using many models and the additional signal captured by datamodels—we also averaged up to 1000 models’ representation distancesWe simply average the ranks from each model, but there are potentially better ways to aggregate them., but this had no discernible difference on the size of the counterfactual effects.

We apply the influence function approximation introduced in [KL17a]. In particular, we use the following first-order approximation for the influence of zz on the loss LL evaluated at ztestz_{\text{test}}:

where θ^\widehat{\theta} is the empirical risk minimizer on the training set and HH is the Hessian of the loss. The influence here is just the dot product of gradients, weighted by the Hessian. We approximate these influence values by using the methods in [KL17a] and as implemented (independently) in pytorch-influence-functions.https://github.com/nimarb/pytorch_influence_functions As in [KL17a], we take a pretrained representation (of a ResNet-9 model, same as that modeled by our datamodels), and compute approximate influence functions with respect to only the parameters in the last linear layer.

[PLS+20a] define an alternative notion of influence: the influence of a training example zz on a test example z′z^{\prime} is the total change in loss on z′z^{\prime} contributed by updates from mini-batches containing zz—intuitively, this measures whether gradient updates from zz are helpful to learning example z′z^{\prime}. They approximate this in practice with TracInCP, which considers checkpoints θt1,...,θtk\theta_{t_{1}},...,\theta_{t_{k}} across training, and sums the dot product of the gradients at zz and z′z^{\prime} at each checkpoint:

One can view TracInCP as a variant of the gradient dot product, but averaged over models at different epochs) and weighted by the learning rate ηi\eta_{i}.

We also consider a random baseline of removing examples from the same class.

F.3 Data support estimation

We use datamodels together with counterfactual evaluations in a guided search to efficiently estimate upper bounds on the size of data supports. For a given target example xx with corresponding datamodel gθg_{\theta}, we want to find candidate training subsets of small size kk whose removal most reduces the classification margin on xx:

Because gθg_{\theta} is a linear model in our case, the solution to the above minimization problem is simply the set corresponding to the largest kk coordinates of the datamodel parameter θ\theta:

(Given that we are using datamodels as surrogates after all, one might wonder if the above counterfactual evaluations are actually necessary—one could instead consider estimating the optimal kk directly from θ\theta. We revisit a heuristic estimation procedure based on this idea at the end of this subsection.)

We assume that the expected margin h(k)≔fA(x;S∖Gk)h(k)\coloneqq f_{\mathcal{A}}({x};{S\setminus G_{k}}) after removing kk examples decreases monotonically in kk; this is expected from the linearity of our datamodels and is further supported empirically (see Figure 22). Then, our goal is to estimate the unique zeroMore precisely, the upper ceiling as data support is defined as an integer quantity. k^\widehat{k} of the above function h(k)h(k) based on (noisy) samples of h(k)h(k) at our chosen values of kk. Note that by definition, k^\widehat{k} is an upper bound on \textscSupport(x)\textsc{Support}(x). Now, because of our monotonicity assumption, we can cast estimating k^\widehat{k} as instance of an isotonic regression problem [RWD88a]); this effectively performs piecewise linear interpolation, while ensuring that monotonicity constraint is not violated. We use sklearn’s IsotonicRegression to fit an estimate h(k)h(k), and use this to estimate k^\widehat{k}.

Due to stochasticity in evaluating counterfactuals, the estimate k^\widehat{k} is noisy. Thus, it is possible that k^\widehat{k} is not a valid upper bound on \textscSupport(x)\textsc{Support}(x), e.g. removing top k^\widehat{k} examples do not misclassify xx. In fact, removing Gk^G_{\widehat{k}} and re-training shows that only 67% of the images are actually misclassified. To establish an upper bound on \textscSupport(x)\textsc{Support}(x) that has sufficient coverage, we evaluate the counterfactuals after removing an additional 20% of highest datamodel weights, e.g. removing top k^×1.2\widehat{k}\times 1.2 examples for each test example. When an additional 20% of training examples are removed, 92% of test examples are misclassified. Hence, we use k^×1.2\widehat{k}\times 1.2 for our final estimates of \textscSupport(x)\textsc{Support}(x).

As baselines, we use the same guided search algorithm described above, but instead of using datamodel-predicted values to guide the search, we select the candidate subset using each of the baselines methods described in Section F.2. In particular, we choose the candidate subset RkR_{k} for a given kk as follows:

Influence estimates (influence functions and TracIn): top kk training examples with highest (most positive) estimated influence on the target example xx:

Random: first kk examples from a random orderingThe random ordering is fixed across different choices of kk, but not across different targets. of training examples from the same class as xx.

While we constructively estimate the data supports by training models on counterfactuals and using the above estimation procedure, we can also consider a simpler and cheaper heuristic to estimate \textscSupport(x)\textsc{Support}(x) assuming the fidelity of the linear datamodels: compute the smallest kk s.t. the sum of the kk highest datamodel weights for xx exceeds the average margin of xx. In Figure 24, we compare the predicted data supports based on this heuristic to the estimated ones from earlier, and find that they are highly correlated. In practice, this can be a more efficient alternative to quantify brittleness without additional model training (beyond the initial ones to estimate the datamodels).

F.4 Brittleness to mislabeling

To study the brittleness of model predictions to mislabeling training images, we take the same 300 random CIFAR-10 test examples and analyze them as follows: First, we find for each example the incorrect class with the highest average logit (across ∼\sim10,000 models trained on the full training set). Then, we construct counterfactual datasets similarly as in Section F.3 where we take the top k={2,4,...,256}k=\{2,4,...,256\} training examples with the highest datamodel weights, but this time mislabel them with the incorrect class identified earlier. After training T=20T=20 models on each counterfactual, for each target example we estimate the number of mislabeled examples at which the expected margin becomes zero, using the same estimation procedure described in Section F.3. The resulting mislabeling brittleness estimates are shown in Figure 24.

F.5 Comparing raw effect sizes

Instead of comparing the data support estimates (which are derived quantities), here we directly compare the average counterfactual effect (i.e. delta margins) of groups selected using different methods. Figure 25 shows again that datamodels identify much larger effects. Among baselines, we see that TracIn performs best, followed by representation distance. We also see that the representation baseline does not gain any additional signal from simple averaging over models.

F.6 Effect of training stochasticity

As described in Section F.1, we re-train up to T=20T=20 models for each counterfactual to reduce noise that arises solely from stochasticity of model training. These additional samples significantly reduces unexplained variance: Figure 27 shows the reduction in variance (“thickness” in the yy-direction) and the resulting increase in correlation as the number of re-training runs is increased from T=1T=1 to T=20T=20.

F.7 Transfer to different architecture

While the main premise of datamodeling is understanding how data is used by a given fixed learning algorithm, it is natural to ask how well datamodels can predict across different learning algorithms. We expect some degradation in predictiveness, as datamodels are fit to a particular learning algorithm; at the same time, we also expect some transfer of predictive power as modern deep neural networks are known to make similar predictions and errors [MMS+19a, TSC+19a].

Here, we study one of the factors in a learning algorithm, the choice of architecture. We take the same counterfactuals and evaluate them on ResNet-18 models, using the same training hyperparameters. As expected, the original datamodels continue to predict accurate counterfactuals for the new model class but with some degradation (Figure 27).

F.8 Stress testing

Section 4.1 showed that datamodels excel at predicting counterfactuals across a variety of removal mechanisms. In an effort to find cases where datamodel predictions are not predictive of data counterfactuals, we evaluate the following additional counterfactuals:

Larger groups of examples (up to 20% of the dataset): we remove k=k= 2560, 5120, 10240 top weights using different datamodels α=0.1,0.2,0.5,0.75\alpha=0.1,0.2,0.5,0.75. The changes in margin have more unexplained variance when larger number of images are removed; nonetheless, the overall correlation remains high ((Figure 28)).

Groups of training examples whose predicted effects are zero: we remove k=k= 20, 40, 80, 160, 320, 640, 1280, 2560 examples with zero weight (α=0.5\alpha=0.5), chosen randomly among all such examples. All of tested counterfactuals had negligible impact on the actual margin, consistent with the prediction of datamodels (Figure 29(a)).

Groups of examples whose predicted effect is negative according to baselines: we test TracIn and influence functions. (We do not consider the representation distance baseline here is there is no obvious way of extracting this information from it.) Correlation degrades but remains high (Figure 29(b)). Note that the relative scale of the effects is much smaller compared to counterfactuals generated using datamodels (Figure 29(a)).

In general, although there is some reduction in datamodels’ predictiveness, we nevertheless find that datamodels continue to be accurate predictors of data counterfactuals.

All of the counterfactuals studied so far are relative a fixed control (the entire training set). Here, we consider counterfactuals relative to a random control S0∼DSS_{0}\sim\mathcal{D}_{S} at α=0.5\alpha=0.5 (i.e. ∣S0∣=α∣S∣|S_{0}|=\alpha|S|). The motivation for considering the shifted control is two folds: first, the counterfactuals generated relative to such S′S^{\prime} are closer in distribution to the original distribution to which datamodels were fit to, so it is natural to study datamodels in this regime; second, this tests whether the counterfactual predictability is robust to the exact choice of the trainset. Latter is desirable, as ultimately we would like to understand how models behave on training sets similar in distribution to SS, not the exact train set.

F.9 Additional plots for different α𝛼\alpha values

Appendix G Nearest Neighbors

In this section we show additional examples of held-out images and their corresponding train image, datamodel weight pairs.

In Figure 33 we show more randomly selected test images along with their positive and negative weight training examples. In Figure 34 we show more examples of test images and their corresponding top train images as we vary α\alpha. In Figure 35 we compare most similar images to given test images identified using various baselines (see Section F.2 for their description).

G.2 FMoW

In Figure 36 we show randomly selected target images along with their top-weight train images, using datamodels of different α\alpha. In Figure 37 we show more examples of test images and their corresponding top train images as we vary α\alpha.

Appendix H Train-Test Leakage

The general setup is as described Section 4.2.2. In the Amazon Mechanical Turk interface (Figure 38), for each test image we displayed the top 5 and bottom 5 train examples by datamodel weight; the vast majority of potential leakage found corresponded to the top 5 examples. Nine different workers filled out each task. We paid 12 cents per task completed and used these qualifications: locale in US/CA/GB and percentage of hits approved >95%>95\%.

Figure 39 shows more examples of (train, test) pairs stratified by annotation score. While there is no ground truth due to lack of metadata, we see that the crowdsourced annotation combined with high quality candidates (as identified by datamodels) can effectively surface leaked examples.

[BD20a] present CIFAIR, a version of CIFAR with fewer duplicates. The authors define duplicates slightly differently than our definition of same scene train-test leakage (cf. Section 3.2 of their work and our interface shown in Figure 38). They identify train-test leakage by using a deep neural network to measure representation space distances between images across training partitions and manually inspecting the lowest distances.

H.2 FMoW

Appendix Figure 40 shows the CDF of each test image’s minimum distance to a train image.

Appendix I Spectral Clustering

We use sklearn’s cluster.SpectralClustering. Internally, this computes similarity scores using the radial basis function (RBF) kernel on the datamodel embeddings. Then, it runs spectral clustering on the graph defined by the similarity matrix AA: it computes a Laplacian LL, represents each node using the first kk eigenvectors of LL, and runs kk-means clustering on the resulting feature representations. We use k=100k=100.

I.2 Omitted results

Figure 41 compares top clusters for the horse class across different α\alpha. Figure 42 shows additional clusters for eight other classes, apart from the ones shown in Figure 13.

Appendix J PCA on Datamodel Embeddings

For the PCA experiments, we use datamodels for the training and test sets estimated with α=0.5\alpha=0.5 unless mentioned otherwise.

In Figure 43, we compare the effective dimensionality of datamodel embeddings with that of a deep representation pretrained on CIFAR-10.

To see whether PCA directions reflect model behavior, we look at how “removing” different principal components affect model predictions. More precisely, we remove training examples corresponding to:

Top kk most positive coordinates of the principal component vector

Top kk most negative coordinates of the principal component vector

Then, for each principal component direction considered, we measure their impact on three groups of held-out samples:

The top 100 examples by most positive projection on the principal component

The bottom 100 examples by most negative projection on the principal component

For each of these groups, we measure the mean change in margin after removing different principal component directions. Our results (Figure 44) show that:

Removing the most positive coordinates of the PC decreases margin on the test set examples with the most positive projections on the PC and increases margin on the examples with the most negative projections on the PC.

Removing the most negative coordinates of the PC has the opposite effect, increasing margin on the positive projection examples and decreasing margin on the negative projection examples.

Increasing the size of each removed set increases the effect magnitude.

Removing PC’s have negligible impact on the aggregate test set, indicating that the impact of different PC’s are roughly “orthogonal,” as one would expect based on the orthogonality of the PCs.

Lastly, Figure 45 shows that datamodels can accurately predict the counterfactual effect of the above removed groups, similarly as in Figure 8.

In Figure 46 we show the top principal components computed using a representation embedding.

In Figure 47 we show additional PCs from a datamodel PCA.

J.2 FMoW

Appendix K Connection between Influence Estimation and Datamodels

For convenience, we introduce the m×nm\times n binary mask matrix A\bm{A} such that Aij\bm{A}_{ij} is an indicator for whether the jj-th training image was included in SiS_{i}. Note that A\bm{A} is a random matrix with fixed row sum of n/2n/2. Next, we define the output vector y∈{0,1}m\bm{y}\in\{0,1\}^{m} that indicates whether a model trained on SiS_{i} was correct on xx. Finally, we introduce the count matrix C=diag(1⊤A)\bm{C}=\text{diag}(\bm{1}^{\top}A), i.e., a diagonal matrix whose entries are the columns sums of A\bm{A}, e.g. the number of times each example appears across mm different masks.

We begin with wOLS\bm{w}_{OLS}. Consider the n×nn\times n matrix Σ=1mZ⊤Z=1m(2⋅A−1m×n)⊤(2⋅A−1m×n)\bm{\Sigma}=\frac{1}{m}\bm{Z}^{\top}\bm{Z}=\frac{1}{m}(2\cdot\bm{A}-\bm{1}_{m\times n})^{\top}(2\cdot\bm{A}-\bm{1}_{m\times n}). The diagonal entries of this matrix are Σii=1\bm{\Sigma}_{ii}=1 (due to A\bm{A} having constant row sum), while the off-diagonal is

By construction, the row sums of Z=2⋅A−1m×n\bm{Z}=2\cdot\bm{A}-\bm{1}_{m\times n} are , and so 1n×n⋅Z⊤=0\bm{1}_{n\times n}\cdot\bm{Z}^{\top}=0. Thus,

We now shift our attention to the empirical influence estimator winfl\bm{w}_{infl}. Using our notation, we can rewrite the (vectorized) empirical influence estimator (15) as:

Now, as m→∞m\to\infty for fixed nn, the random variable mC−1m\bm{C}^{-1} converges to 2⋅I2\cdot\bm{I} with probability 11. Thus,

and the empirical influence estimator winfl→2mZ⊤y\bm{w}_{infl}\to\frac{2}{m}\bm{Z}^{\top}\bm{y}, which completes the proof. ∎

K.2 Evaluating influence estimates as datamodels

Lemma 6.1 suggests that we can re-cast empirical influence estimates as (rescaled) datamodels fit with least-squares loss. Under this view, (i.e., ignoring the difference in conceptual goal), we can differentiate between explicit datamodels and those arising from empirical influences along three axes:

Estimation algorithm: Most importantly, datamodels explicitly minimize the squared error between true and predicted model outputs. Furthermore, datamodels as instantiated here use (a) a sparsity prior and (b) a bias term which may help generalization.

Scale: Driven by their intended applications (where one typically only needs to estimate the highest-influence training points for a given test point), empirical influence estimates are typically computed with relatively few samples (i.e., m<dm<d, in our setting) [FZ20a]. In contrast, we find that for datamodel loss to plateau, one needs to estimate parameters using a much larger set of models.

Output type: Finally, datamodels do not restrict to prediction of a binary correctness variable—in this paper, for example, for deep classification models we find that correct-class margin was best both heuristically and in practice.

In this section, we thus ask: how well do the rescaled datamodels that arise from empirical influence estimates predict model outputs? We address this question in the context of the three axes of variation described above. In order to make results comparable across different outputs types (e.g., correctness vs. correct-class margin), we measure correlation (in the sense of [Spe04a]) between the predicted and true model outputs, in addition to MSE where appropriate. To ensure a conservative comparison, we also measure performance as a predictor of correctness. In particular, we treat w⊤1Si\bm{w}^{\top}\bm{1}_{S_{i}} as a continuous predictor of the binary variable \bm{1}\{\text{model trained onS_{i}iscorrectonis correct onx}\}, and compute the AUC of this predictor (intuitively, this should favor empirical influence estimates since they are computed using correctnesses directly).

In Table 5 we show the difference between empirical influence estimates (first row) and our final datamodel estimates (last row), while disentangling the effect of the three axes above using the rows in between. As expected, there is a vast difference in terms of correlation between the original empirical influence estimates and explicit datamodels.

We further illustrate this point in Figure 49, where we show how the correlation, MSE, and AUC vary with mm for both empirical influence estimates and datamodels, as well as an intermediate estimator that uses the estimation procedure of empirical influence estimates but replaces correctness with margin.

K.3 Testing Lemma 6.1 empirically

In this section, we visualize the performance of empirical influence estimates ([FZ20a]) as datamodels. In Figures 50(a) and 50(b) we plot the distributions of winfl⊤1Si∣yi\bm{w}_{infl}^{\top}\bm{1}_{S_{i}}|y_{i} for different CIFAR-10 test examples; Figure 50(a) shows these “conditional prediction distributions” for subsets SiS_{i} that were used to estimate the empirical influence, while Figure 50(b) shows the corresponding distributions on held-out (unseen) subsets SiS_{i}. The figures suggest that (i) indeed, empirical influences are somewhat predictive of the correctness yiy_{i}, (ii) their predictiveness increases as number of samples m→∞m\to\infty but is still rather low, and (c) a significant amount of the prediction error is generalization error, as the train predictions in Figure 50(a) are significantly better-separated than the heldout predictions in Figure 50(b).

K.4 View of empirical influences as a Taylor approximation

Section 6.1 shows that we can interpret empirical influences as (rescaled) estimates of the weights of a linear datamodel. Here, we give an alternative intuition for why this is the case, even though the definition of empirical influence does not explicitly assume linearity anywhere: we show that the influences define a first-order Taylor approximation of the multilinear extension ff of our target function FF of interest, where the influences (approximately) correspond to first-order derivatives of ff.

f(x)f(x) also has an intuitive interpretation: it is the expected value of F(S)F(S) when SS is chosen by including each xix_{i} in the input with probability xix_{i}.

Next, we take the derivative of ff w.r.t. to the input xix_{i}:

Note that because ff is multilinear, the derivative w.r.t. to xix_{i} is constant in xix_{i}, but not w.r.t. to other xjx_{j}. Now, observe that the above expression evaluated at xj=α,  ∀xjx_{j}=\alpha,\;\forall x_{j} corresponds approximatelyThere are two sources of approximation here. First, the α\alpha-subsampling used in our datamodel definition is defined globally (e.g. α\alpha fraction of entire train set), which is different from the i.i.d. Bern(α)Bern(\alpha) sampling that is considered here. Second, we only observe noisy versions of F(S)F(S). to α\alpha-subsampled influence θi\theta_{i}, of ii on FF: the first term corresponds (using our earlier interpretation) to the expectation of F(S)F(S) conditional on SS including ii, and the second to that conditional on SS excluding ii.

Finally, the first-order Taylor approximation of ff around an xx is given as:

where θi\theta_{i} are the empirical influences.

The above perspective provides an alternative way to think the role of the sampling fraction α\alpha. The weights θi\theta_{i} depend on the regime we are interested in; if we use α\alpha-subsampled influences, then we are effectively taking a local linear approximation of ff in the regime around x⃗=α⋅1⃗\vec{x}=\alpha\cdot\vec{1}.

Though we include the exposition above for completeness, this is a classical derivation that has appeared in similar form in prior works [Owe72a]. Another connection is that Shapley value is equivalent to the integral of ff along the “main diagonal” of the hypercube; it is effectively empirical influences averaged uniformly over the choice of α\alpha.