The Pitfalls of Simplicity Bias in Neural Networks

Harshay Shah, Kaustav Tamuly, Aditi Raghunathan, Prateek Jain, Praneeth Netrapalli

Introduction

Understanding the superior generalization ability of neural networks, despite their high capacity to fit randomly labeled data , has been a subject of intense study. One line of recent work proves that linear neural networks trained with Stochastic Gradient Descent (SGD) on linearly separable data converge to the maximum-margin linear classifier, thereby explaining the superior generalization performance. However, maximum-margin classifiers are inherently robust to perturbations of data at prediction time, and this implication is at odds with concrete evidence that neural networks, in practice, are brittle to adversarial examples and distribution shifts . Hence, the linear setting, while convenient to analyze, is insufficient to capture the non-robustness of neural networks trained on real datasets. Going beyond the linear setting, several works argue that neural networks generalize well because standard training procedures have a bias towards learning simple models. However, the exact notion of “simple" models remains vague and only intuitive. Moreover, the settings studied are insufficient to capture the brittleness of neural networks

Our goal is to formally understand and probe the simplicity bias (SB) of neural networks in a setting that is rich enough to capture known failure modes of neural networks and, at the same time, amenable to theoretical analysis and targeted experiments. Our starting point is the observation that on real-world datasets, there are several distinct ways to discriminate between labels (e.g., by inferring shape, color etc. in image classification) that are (a) predictive of the label to varying extents, and (b) define decision boundaries of varying complexity. For example, in the image classification task of white swans vs. bears, a linear-like “simple" classifier that only looks at color could predict correctly on most instances except white polar bears, while a nonlinear “complex" classifier that infers shape could have almost perfect predictive power. To systematically understand SB, we design modular synthetic and image-based datasets wherein different coordinates (or blocks) define decision boundaries of varying complexity. We refer to each coordinate / block as a feature and define a precise notion of feature simplicity based on the simplicity of the corresponding decision boundary.

Figure 1 illustrates a stylized version of the proposed synthetic dataset with two features, ϕ1\phi_{1} and ϕ2\phi_{2}, that can perfectly predict the label with 100% accuracy, but differ in simplicity.

The simplicity of a feature is precisely determined by the minimum number of linear pieces in the decision boundary that achieves optimal classification accuracy using that feature. For example, in Figure 1, the simple feature ϕ1\phi_{1} requires a linear decision boundary to perfectly predict the label, whereas complex feature ϕ2\phi_{2} requires four linear pieces. Along similar lines, we also introduce a collection of image-based datasets in which each image concatenates MNIST images (simple feature) and CIFAR-10 images (complex feature). The proposed datasets, which incorporate features of varying predictive power and simplicity, allow us to systematically investigate and measure SB in SGD-trained neural networks.

The ideal decision boundary that achieves high accuracy and robustness relies on all features to obtain a large margin (minimum distance from any point to decision boundary). For example, the orange decision boundary in Figure 1 that learns ϕ1\phi_{1} and ϕ2\phi_{2} attains 100% accuracy and exhibits more robustness than the linear boundary because of larger margin. Given the expressive power of large neural networks, one might expect that a network trained on the dataset in Figure 1 would result in the larger-margin orange piecewise linear boundary. However, in practice, we find quite the opposite—trained neural networks have a linear boundary. Surprisingly, neural networks exclusively use feature ϕ1\phi_{1} and remain completely invariant to ϕ2\phi_{2}. More generally, we observe that SB is extreme: neural networks simply ignore several complex predictive features in the presence of few simple predictive features. We first theoretically show that one-hidden-layer neural networks trained on the piecewise linear dataset exhibit SB. Then, through controlled experiments, we validate the extreme nature of SB across model architectures and optimizers.

Theoretical analysis and controlled experiments reveal three major pitfalls of SB in the context of proposed synthetic and image-based datasets, which we conjecture to hold more widely across datasets and domains:

(i) Lack of robustness: Neural networks exclusively latch on to the simplest feature (e.g., background) at the expense of very small margin and completely ignore complex predictive features (e.g., semantics of the object), even when all features have equal predictive power. This results in susceptibility to small adversarial perturbations (due to small margin) and spurious correlations (with simple features). Furthermore, in Section 4, we provide a concrete connection between SB and data-agnostic and model-agnostic universal adversarial perturbations observed in practice.

(ii) Lack of reliable confidence estimates: Ideally, a network should have high confidence only if all predictive features agree in their prediction. However due to extreme SB, the network has high confidence even if several complex predictive features contradict the simple feature, mirroring the widely reported inaccurate and substantially higher confidence estimates reported in practice .

(iii) Suboptimal generalization: Surprisingly, neural networks exclusively rely on the simplest feature even if it less predictive of the label than all complex features in the synthetic datasets. Consequently, contrary to conventional wisdom, extreme SB can hurt robustness as well as generalization.

In contrast, prior works only extol SB by considering settings where all predictive features are simple and hence do not reveal the pitfalls observed in real-world settings. While our results on the pitfalls of SB are established in the context of the proposed datasets, the two design principles underlying these datasets—combining multiple features of varying simplicity & predictive power and capturing multiple failure modes of neural networks in practice—suggest that our conclusions could be justifiable more broadly.

This work makes two key contributions. First, we design datasets that offer a precise stratification of features based on simplicity and predictive power. Second, using the proposed datasets, we provide theoretical and empirical evidence that neural networks exhibit extreme SB, which we postulate as a unifying contributing factor underlying key failure modes of deep learning: poor out-of-distribution performance, adversarial vulnerability and suboptimal generalization. To the best of our knowledge, prior works only focus on the positive aspect of SB: the lack of overfitting in practice. Additionally, we find that standard approaches to improve generalization and robustness—ensembles and adversarial training—do not mitigate simplicity bias and its shortcomings on the proposed datasets. Given the important implications of SB, we hope that the datasets we introduce serve (a) as a useful testbed for devising better training procedures and (b) as a starting point to design more realistic datasets that are amenable to theoretical analysis and controlled experiments.

Organization. We discuss related work in Section 2. Section 3 describes the proposed datasets and metrics. In Section 4, we concretely establish the extreme nature of Simplicity Bias (SB) and its shortcomings through theory and empirics. Section 5 shows that extreme SB can in fact hurt generalization as well. We conclude and discuss the way forward in Section 6.

Related Work

Out-of-Distribution (OOD) performance: Several works demonstrate that NNs tend to learn spurious features & low-level statistical patterns rather than semantic features & high-level abstractions, resulting in poor OOD performance . This phenomenon has been exploited to design backdoor attacks against NNs as well. Recent works that encourage models to learn higher-level features improve OOD performance, but require domain-specific knowledge to penalize reliance on spurious features such as image texture and annotation artifacts in vision & language tasks. Learning robust representations without domain knowledge, however, necessitates formalizing the notion of features and feature reliance; our work takes a step in this direction.

Adversarial robustness: Neural networks exhibit vulnerability to small adversarial perturbations . Standard approaches to mitigate this issue—adversarial training and ensembles —have had limited success on large-scale datasets. Consequently, several works have investigated reasons underlying the existence of adversarial examples: suggests local linearity of trained NNs, indicates insufficient data, suggests inevitability in high dimensions, suggests computational barriers, proposes limitations of neural network architectures, and proposes the presence of non-robust features. Additionally, Jacobsen et al. show that NNs exhibit invariance to large label-relevant perturbations. Prior works have also demonstrated the existence of universal adversarial perturbations (UAPs) that are agnostic to model and data .

Implicit bias of stochastic gradient descent : Brutzkus et al. show that neural networks trained with SGD provably generalize on linearly separable data. Recent works also analyze the limiting direction of gradient descent on logistic regression with linearly separable and non-separable data respectively. Empirical findings provide further evidence to suggest that SGD-trained NNs generalize well because SGD learns models of increasing complexity. Additional recent works investigate the implicit bias of SGD on non-linearly separable data for linear classifiers and infinite-width two-layer NNs , showing convergence to maximum margin classifiers in appropriate spaces. In Section 4, we show that SGD’s implicit bias towards simplicity can result in small-margin and feature-impoverished classifiers instead of large-margin and feature-dense classifiers.

Feature reliance: Two recent works study the relation between inductive biases of training procedures and the set of features that models learn. Hermann et al. use color, shape and texture features in stylized settings to show that standard training procedures can (a) increase reliance on task-relevant features that are partially decodable using untrained networks and (b) suppress reliance on non-discriminative or correlated features. Ortiz et al. show that neural networks learned using standard training (a) develop invariance to non-discriminative features and (b) adversarial training induces a sharp transition in the models’ decision boundaries. In contrast, we develop a precise notion of feature simplicity and subsequently show that SGD-trained models can exhibit invariance to multiple discriminative-but-complex features. We also identify three pitfalls of this phenomenon—poor OOD performance, adversarial vulnerability, suboptimal generalization—and show that adversarial training and standard ensembles do not mitigate the pitfalls of simplicity bias.

Multiple works mentioned above (a) differentially characterize learned features and desired features—statistical regularities vs. high-level concepts , syntactic cues vs. semantic meaning , robust vs. non-robust features —and (b) posit that the mismatch between these features results in non-robustness. In contrast, our work probes why neural networks prefer one set of features over another and unifies the aforementioned feature characterizations through the lens of feature simplicity.

Preliminaries: Setup and Metrics

Next, we introduce two metrics that quantitatively capture the extent to which a model relies on different input coordinates (or features). Let SS denote some subset of coordinates [d][d] and D‾S\overline{\mathcal{D}}^{S} denote the SS-randomized distribution, which is obtained as follows: given DS\mathcal{D}^{S}, the marginal distribution of SS, D‾S\overline{\mathcal{D}}^{S} independently samples ((xS,xSc),y)∼D((x^{S},x^{S^{\mathsf{c}}}),y)\sim\mathcal{D} and x‾S∼DS\overline{x}^{S}\sim\mathcal{D}^{S} and then outputs ((x‾S,xSc),y)((\overline{x}^{S},x^{S^{\mathsf{c}}}),y). In D‾S\overline{\mathcal{D}}^{S}, the coordinates in SS are rendered independent of the label yy. The two metrics are as follows.

Given data distribution D\mathcal{D} and subset of coordinates S⊆[d]S\subseteq[d], the SS-randomized AUC of classifier ff equals the area under the precision-recall curve of distribution D‾S\overline{\mathcal{D}}^{S}.

Our experiments use {S\mathtt{S}, Sc\mathtt{S}^{c} }-randomized metrics—accuracy, AUC, logits—to establish that ff depends exclusively on some features S\mathtt{S} and remains invariant to the rest Sc\mathtt{S}^{c}.

First, if (a) S\mathtt{S}-randomized accuracy and AUC equal 0.50.5 and (b) S\mathtt{S}-randomized logit distribution is a random shuffling of the original distribution (i.e., logits in the original distribution are randomly shuffled across true positives and true negatives), then ff depends exclusively on S\mathtt{S}. Conversely, if (a) Sc\mathtt{S}^{c}-randomized accuracy and AUC are equal to standard accuracy and AUC and (b) Sc\mathtt{S}^{c}-randomized logit distribution is essentially identical to the original distribution, then ff is invariant to Sc\mathtt{S}^{c}; Table 1 summarizes these observations.

1 Datasets

One-dimensional Building Blocks: Our synthetic datasets use three one-dimensional data blocks—linear, noisy linear and kk-slabs—shown in top row of Figure 2. In the linear block, positive and negative examples are uniformly distributed in [0.1,1][0.1,1] and [\scalebox{0.75}[1.0]{-}1,\scalebox{0.75}[1.0]{-}0.1] respectively. In the noisy linear block, given a noise parameter p∈p\in, 1−p1-p fraction of points are distributed like the linear block described above and pp fraction of the examples are uniformly distributed in [\scalebox{0.75}[1.0]{-}0.1,0.1]. In kk-slab blocks, positive and negative examples are distributed in kk well-separated, alternating regions.

Multi-dimensional Synthetic Datasets: We now outline four dd-dimensional datasets wherein each coordinate corresponds to one of three building blocks described above. See Figure 2 for illustration.

LMS-k: Linear and multiple kk-slabs; the first coordinate is a linear block and the remaining d\scalebox{0.75}[1.0]{-}1 coordinates are independent kk-slab blocks; we use LMS-5 & LMS-7 datasets in our analysis.

L̂MS-k: Noisy linear and multiple kk-slab blocks; the first coordinate is a noisy linear block and the remaining d\scalebox{0.75}[1.0]{-}1 coordinates are independent kk-slab blocks. The noise parameter pp is 0.10.1 by default.

MS-(5,7): 5-slab and multiple 7-slab blocks; the first coordinate is a 5-slab block and the remaining d\scalebox{0.75}[1.0]{-}1 coordinates are independent 77-slab blocks, as shown in Figure 2.

MS-5: Multiple 5-slab blocks; all coordinates are independent 5-slab blocks.

We now describe the LSN (linear, 3-slab & noise) dataset, a stylized version of LMS-k that is amenable to theoretical analysis. We note that recent works empirically analyze variants of the LSN dataset. In LSN, conditioned on the label yy, the first and second coordinates of xx are singleton linear and 33-slab blocks: linear and 33-slab blocks have support on \{\scalebox{0.75}[1.0]{-}1,1\} and \{\scalebox{0.75}[1.0]{-}1,0,1\} respectively. The remaining coordinates are standard gaussians and not predictive of the label.

The synthetic datasets comprise features of varying simplicity; in LMS-k, L̂MS-k, and MS-(5,7), the first coordinate is the simplest feature and in MS-5, all features are equally simple. All datasets, even L̂MS-k, can be perfectly classified via piecewise linear classifiers. Though the kk-slab features are special cases of linear periodic functions on which gradient-based methods have been shown to fail for large kk , we note that we use small values of k∈{5,7}k\in\{5,7\} which are quickly learned by SGD in practice. Note that we (a) apply a random rotation matrix to the data and (b) use 5050-dimensional synthetic data (i.e., d=50d=50) by default. Note that all code and datasets are available at the following repository: https://github.com/harshays/simplicitybiaspitfalls.

MNIST-CIFAR Data: The MNIST-CIFAR dataset consists of two classes: images in class \scalebox{0.75}[1.0]{-}1 and class 11 are vertical concatenations of MNIST digit zero & CIFAR-10 automobile and MNIST digit one & CIFAR-10 truck images respectively, as shown in Figure 2. The training and test datasets comprise 50,000 and 10,000 images of size 3×64×323\times 64\times 32. The MNIST-CIFAR dataset mirrors the structure in the synthetic LMS-k dataset—both incorporate simple and complex features. The MNIST and CIFAR blocks correspond to the linear and kk-slab blocks in LMS-k respectively. Also note that MNIST images are zero-padded & replicated across three channels to match CIFAR dimensions before concatenation.

Appendix B provides details about the datasets, models, and optimizers used in our experiments. In Appendix C, we show that our results are robust to the exact choice of MNIST-CIFAR class pairs.

Simplicity Bias (SB) is Extreme and Leads to Non-Robustness

We first establish the extreme nature of SB in neural networks (NNs) on the proposed synthetic datasets using SGD and variants. In particular, we show that for the datasets considered, if all features have full predictive power, NNs rely exclusively on the simplest feature S\mathtt{S} and remain invariant to all complex features Sc\mathtt{S}^{c}. Then, we explain why extreme SB on these datasets results in neural networks that are vulnerable to distribution shifts and data-agnostic & transferable adversarial perturbations.

We consider the LSN dataset (described in Figure 2) that has one linear coordinate and one 3-slab coordinate, both fully predictive of the label on their own; the remaining d\scalebox{0.75}[1.0]{-}2 noise coordinatesI do not have any predictive power. Now, a "large-margin" one-hidden-layer NN with ReLU activation should give equal weight to the linear and 3-slab coordinates. However, we prove that NNs trained with standard mini-batch gradient descent (GD) on the LSN dataset (described in Figure 2) provably learns a classifier that exclusively relies on the “simple" linear coordinate, thus exhibiting simplicity bias at the cost of margin. Further, our claim holds even when the margin in the linear coordinate (minimum distance between linear coordinate of positives and negatives) is significantly smaller than the margin in the slab coordinate. The proof of the following theorem is presented in Appendix F.

Let f(x)=∑j=1kvj⋅ReLU(∑i=1dwi,jxi)f(x)=\sum_{j=1}^{k}v_{j}\cdot\textrm{ReLU}(\sum_{i=1}^{d}w_{i,j}x_{i}) denote a one-hidden-layer neural network with kk hidden units and ReLU activations. Set vj=±\nicefrac1kv_{j}=\pm\nicefrac{{1}}{{\sqrt{k}}} w.p. \nicefrac12\nicefrac{{1}}{{2}} ∀j∈[k]\forall j\in[k]. Let {(xi,yi)}i=1m\{(x^{i},y^{i})\}^{m}_{i=1} denote i.i.d. samples from LSN where m∈[cd2,dα/c]m\in[cd^{2},d^{\alpha}/c] for some α>2\alpha>2. Then, given d>Ω(klog⁡k)d>\Omega(\sqrt{k}\log k) and initial wij∼N(0,1dklog⁡4d)w_{ij}\sim\mathcal{N}(0,\frac{1}{dk\log^{4}d}), after O(1)O(1) iterations, mini-batch gradient descent (over ww) with hinge loss, step size \eta=\Omega{{(\log d)^{\nicefrac{{\scalebox{0.75}[1.0]{-}1}}{{2}}}}}, mini-batch size Θ(m)\Theta(m), satisfies:

Test error is at most \nicefrac1poly(d)\nicefrac{{1}}{{\textrm{poly}(d)}}

The learned weights of hidden units wijw_{ij} satisfy:

with probability greater than 1−1poly(d)1-\frac{1}{\textrm{poly}(d)}. Note that cc is a universal constant.

2 Simplicity Bias (SB) is Extreme in Practice

We now establish the extreme nature of SB on datasets with features of varying simplicity—LMS-5, MS-(5,7), MNIST-CIFAR (described in Section 3)—across multiple model architectures and optimizers. Recall that (a) the simplicity of one-dimensional building blocks is defined as the number of pieces required by a piecewise linear classifier acting only on that block to get optimal accuracy and (b) LMS-5 has one linear block & multiple 55-slabs, MS-(5,7) has one 55-slab and multiple 77-slabs and MNIST-CIFAR concatenates MNIST and CIFAR10 images. We now use S\mathtt{S} to denote the simplest feature in each dataset: linear in LMS-5, 5-slab in MS-(5,7), and MNIST in MNIST-CIFAR.

We first consistently observe that SGD-trained models trained on LMS-5 and MS-(5,7) datasets exhibit extreme SB: they exclusively rely on the simplest feature S\mathtt{S} and remain invariant to all complex features Sc\mathtt{S}^{c}. Using S\mathtt{S}-randomized & Sc\mathtt{S}^{c}-randomized metrics summarized in Table 1, we first establish extreme SB on fully-connected (FCN), convolutional (CNN) & sequential (GRU ) models. We observe that the S\mathtt{S}-randomized AUC is 0.5 across models. That is, unsurprisingly, all models are critically dependent on S\mathtt{S}. Surprisingly, however, Sc\mathtt{S}^{c}-randomized AUC of all models on both datasets equals 1.0. That is, arbitrarily perturbing Sc\mathtt{S}^{c} coordinates has no impact on the class predictions or the ranking of true positives’ logits against true negatives’ logits. One might expect that perturbing Sc\mathtt{S}^{c} would at least bring the logits of positives and negatives closer to each other. Figure 3(a) answers this in negative—the logit distributions over true positives of (100,1)-FCNs (i.e., with width 100 & depth 1) remain unchanged even after randomizing all complex features Sc\mathtt{S}^{c}. Conversely, randomizing the simplest feature S\mathtt{S} randomly shuffles the original logits across true positives as well as true negatives. The two-dimensional projections of FCN decision boundaries in Figure 3(c) visually confirm that FCNs exclusively depend on the simpler coordinate S\mathtt{S} and are invariant to all complex features Sc\mathtt{S}^{c}.

Note that sample size and model architecture do not present any obstacles in learning complex features Sc\mathtt{S}^{c} to achieve 100% accuracy. In fact, if S\mathtt{S} is removed from the dataset, SGD-trained models with the same sample size indeed rely on Sc\mathtt{S}^{c} to attain 100% accuracy. Increasing the number of complex features does not mitigate extreme SB either. Figure 3(b) shows that even when there are 249 complex features and only one simple feature, (2000,1)-FCNs exclusively rely on the simplest feature S\mathtt{S}; randomizing Sc\mathtt{S}^{c} keeps AUC score of 1.01.0 intact but simply randomizing S\mathtt{S} drops the AUC score to 0.50.5. (2000,1)-FCNs exhibit extreme SB despite their expressive power to learn large-margin classifiers that rely on all simple and complex features.

Similarly, on the MNIST-CIFAR dataset, MobileNetV2 , GoogLeNet , ResNet50 and DenseNet121 exhibit extreme SB. All models exclusively latch on to the simpler MNIST block to acheive 100% accuracy and remain invariant to the CIFAR block, even though the CIFAR block alone is almost fully predictive of its label—GoogLeNet attains 95.4% accuracy on the corresponding CIFAR binary classification task. Figure 3(a) shows that randomizing the simpler MNIST block randomly shuffles the logit distribution of true positives whereas randomizing the CIFAR block has no effect—the CIFAR-randomized and original logit distribution over true positives essentially overlap.

3 Extreme Simplicity Bias (SB) leads to Non-Robustness

Now, we discuss how our findings about extreme SB in Section 4.2 can help reconcile poor OOD performance and adversarial vulnerability with superior generalization on the same data distribution.

Poor OOD performance: Given that neural networks tend to heavily rely on spurious features , state-of-the-art accuracies on large and diverse validation sets provide a false sense of security; even benign distributional changes to the data (e.g., domain shifts) during prediction time can drastically degrade or even nullify model performance. This phenomenon, though counter-intuitive, can be easily explained through the lens of extreme SB. Specifically, we hypothesize that spurious features are simple. This hypothesis, when combined with extreme SB, explains the outsized impact of spurious features. For example, Figure 3(b) shows that simply perturbing the simplest (and potentially spurious in practice) feature S\mathtt{S} drops the AUC of trained neural networks to 0.50.5, thereby nullifying model performance. Randomizing all complex features Sc\mathtt{S}^{c}—5-slabs in LMS-5, 7-slabs in MS-(5,7), CIFAR block in MNIST-CIFAR—has negligible effect on the trained neural networks—Sc\mathtt{S}^{c}-randomized and original logits essentially overlap—even though Sc\mathtt{S}^{c} and S\mathtt{S} have equal predictive power This further implies that approaches that aim to detect distribution shifts based on model outputs such as logits or softmax probabilities may themselves fail due to extreme SB.

To summarize, through theoretical analysis and extensive experiments on synthetic and image-based datasets, we (a) establish that SB is extreme in nature across model architectures and datasets and (b) show that extreme SB can result in poor OOD performance and adversarial vulnerability, even when all simple and complex features have equal predictive power.

Extreme Simplicity Bias (SB) can hurt Generalization

In this section, we show that, contrary to conventional wisdom, extreme SB can potentially result in suboptimal generalization of SGD-trained models on the same data distribution as well. This is because exclusive reliance on the simplest feature S\mathtt{S} can persist even when every complex feature in Sc\mathtt{S}^{c} has significantly greater predictive power than S\mathtt{S}.

We verify this phenomenon on L̂MS-7 data defined in Section 3. Recall that L̂MS-7 has one noisy linear coordinate S\mathtt{S} with 95% predictive power (i.e., 10% noise in linear coordinate) and multiple 7-slab coordinates Sc\mathtt{S}^{c}, each with 100% predictive power. Note that our training sample size is large enough for FCNs of depth {1,2}\{1,2\} and width {100,200,300}\{100,200,300\} trained on Sc\mathtt{S}^{c} only (i.e., after removing S\mathtt{S} from data) to attain 100% test accuracy. However, when trained on L̂MS-7 (i.e., including S\mathtt{S}), SGD-trained FCNs exhibit extreme SB and only rely on S\mathtt{S}, the noisy linear coordinate. In Table 2, we report accuracies of SGD-trained FCNs that are selected based on validation accuracy after performing a grid search over four SGD hyperparameters: learning rate, batch size, momentum, and weight decay. The train, test and randomized accuracies in Table 2 collectively show that FCNs exclusively rely on the noisy linear feature and consequently attain 5% generalization error.

To summarize, the mere presence of a simple-but-noisy feature in L̂MS-7 data can significantly degrade the performance of SGD-trained FCNs due to extreme SB. Note that our results show that even an extensive grid search over SGD hyperparameters does not improve the performance of SGD-trained FCNs on L̂MS-7 data but does not necessarily imply that mitigating SB via SGD and its variants is impossible. We provide additional information about the experiment setup in Appendix D.

Conclusion and Discussion

We investigated Simplicity Bias (SB) in SGD-trained neural networks (NNs) using synthetic and image-based datasets that (a) incorporate a precise notion of feature simplicity, (b) are amenable to theoretical analysis and (c) capture the non-robustness of NNs observed in practice. We first showed that one-hidden-layer ReLU NNs provably exhibit SB on the LSN dataset. Then, we analyzed the proposed datasets to empirically demonstrate that SB can be extreme, and can help explain poor OOD performance and adversarial vulnerability of NNs. We also showed that, contrary to conventional wisdom, extreme SB can potentially hurt generalization.

Can we mitigate SB? It is natural to wonder if any modifications to the standard training procedure can help in mitigating extreme SB and its adverse consequences. In Appendix E, we show that well-studied approaches for improving generalization and adversarial robustness—ensemble methods and adversarial training—do not mitigate SB, at least on the proposed datasets. Specifically, “vanilla" ensembles of independent models trained on the proposed datasets mitigate SB to some extent by aggregating predictions, but continue to exclusively rely on the simplest feature. That is, the resulting ensemble remains invariant to all complex features (e.g., 55-Slabs in LMS-5). Our results suggest that in practice, the improvement in generalization due to vanilla ensembles stem from combining multiple simple-but-noisy features (such as color, texture) and not by learning diverse and complex features (such as shape). Similarly, adversarially training FCNs on the proposed datasets increases margin (and hence adversarial robustness) to some extent by combining multiple simple features but does not achieve the maximum possible adversarial robustness; the resulting adversarially trained models remain invariant to all complex features.

Our results collectively motivate the need for novel algorithmic approaches that avoid the pitfalls of extreme SB. Furthermore, the proposed datasets capture the key aspects of training neural networks on real world data, while being amenable to theoretical analysis and controlled experiments, and can serve as an effective testbed to understand deep learning phenomena. evaluating new algorithmic approaches aimed at avoiding the pitfalls of SB.

Broader Impact

Our work is foundational in nature and seeks to improve our understanding of neural networks. We do not foresee any significant societal consequences in the short term. However, in the long term, we believe that a concrete understanding of deep learning phenomena is essential to develop reliable deep learning systems for practical applications that have societal impact.

References

Appendix A Additional Related Work

In this section, we provide a more thorough discussion of relevant work related to margin-based generalization bounds, adversarial attacks and robustness, and out-of-distribution (OOD) examples.

Margin-based generalization bounds: Building up on the classical work of , recent works try to obtain tighter generalization bounds for neural networks in terms of normalized margin . Here, margin is defined as the difference in the probability of the true label and the largest probability of the incorrect labels. While these bounds seem to capture generalization of neural networks at a coarse level, it has been argued that these approaches may be incapable of fully explaining the generalization ability of neural networks. Furthermore, it is unclear if the notion of model complexity used in these works, based on Lipschitz constant, captures generalization ability accurately. In any case, our results suggest that due to extreme simplicity bias (SB), even if a formulation captures both margin and model complexity accurately, current optimization techniques may not be able to find the optimal solution in terms of generalization and robustness-, as they are strongly biased towards small-margin classifiers that exclusively rely on the simplest features.

Detecting OOD Examples: Neural networks trained using standard training procedures tend to rely on low-level features and spurious correlations and hence exhibit brittleness to benign distributional changes to the data. Recent works thus aim to detect OOD examples using generative models , statistical tests , and model confidence scores . Our experiments in Section 4 that validate extreme SB in practice also show that detectors that directly or indirectly rely on model scores to detect OOD examples may not work well as SGD-trained neural networks can exhibit complete invariance to predictive-but-complex features.

Appendix B Experiment Details

In this section, we provide additional details on the datasets, models, optimization methods and training hyperparameters used in our experiments.

One-dimensional Building Blocks: We first describe the data generation process underlying each building block: linear, noisy linear, and kk-slab. Then, we introduce a noisy version of the 55-slab block, which we later use in Appendix D.

Linear(γ,B)(\gamma,B): The linear block is parameterized by the effective margin γ\gamma and width BB. The distribution first samples a label y\in\{\scalebox{0.75}[1.0]{-}1,1\} uniformly at random, and then given yy, xx is sampled as follows: x=y(Bγ+(B−Bγ)⋅U(0,1),x=y(B\gamma+(B-B\gamma)\cdot\text{U}(0,1), where U(0,1)\text{U}(0,1) is the uniform distribution on $$.

NoisyLinear(γ,B,p)(\gamma,B,p): The noisy linear block is parameterized by effective margin γ\gamma, width BB, and noise parameter pp. Linear classifiers can attain the optimal classification accuracy of 1−\nicefracp21-\nicefrac{{p}}{{2}}. Given label y\in\{\scalebox{0.75}[1.0]{-}1,1\} sampled uniformly at random, xx is sampled as follows:

Slab(γ,B,k)(\gamma,B,k): The kk-slab block is parameterized by effective margin γ\gamma, width BB, and number of slabs kk. We use k∈{3,5,7}k\in\{3,5,7\} in our paper. The width of each slab, wk=\nicefrac2B(1−(k−1)γ)kw_{k}=\nicefrac{{2B(1-(k-1)\gamma)}}{{k}}, in the kk-slab block is chosen such that the farthest points are at \scalebox{0.75}[1.0]{-}B and BB. For example, given label y\in\{\scalebox{0.75}[1.0]{-}1,1\} and random sign z\in\{\scalebox{0.75}[1.0]{-}1,1\} sampled unif. at random, we can sample xx from a 33-slab block as follows:

For kk-slab blocks with k∈{5,7}k\in\{5,7\}, the probability of sampling from the two slabs (one on each side) that are farthest away from the origin are \nicefrac14\nicefrac{{1}}{{4}} and \nicefrac18\nicefrac{{1}}{{8}} respectively to ensure that the variance of instances in positive and negative classes, x+x_{+} and x−x_{-}, are equal.

NoisySlab(γ,B,k,p)(\gamma,B,k,p): Analogous to the noisy linear block, the noisy variant of the kk-slab block is additionally parameterized by a noise parameter pp. In this setting, a (k\scalebox{0.75}[1.0]{-}1)-piecewise linear classifier can attain the optimal classification accuracy of 1−\nicefracp21-\nicefrac{{p}}{{2}}. For example, For example, given label y\in\{\scalebox{0.75}[1.0]{-}1,1\} and random sign z\in\{\scalebox{0.75}[1.0]{-}1,1\} sampled uniformly at random, we can sample xx from a pp-noisy 33-slab block as follows:

Datasets: We now outline the default hyperparameters for generating the synthetic datasets used in the paper, provide additional details on the LSN dataset, and introduce two additional synthetic datasets as well as multiple versions of the MNIST-CIFAR dataset (i.e., with different class pairs).

Synthetic Dataset Hyperparameters: Recall that we use four dd-dimensional synthetic datasets—LMS-k, L̂MS-k, MS-(5,7), and MS-5—wherein each coordinate corresponds to one of the building blocks described above. Unless mentioned otherwise, for all four datasets, we set the effective margin parameter γ=0.1\gamma=0.1, width parameter B=1B=1, and noise parameter p=0.1p=0.1 in all blocks/coordinates. Also recall that each dataset comprises at most one “simple" feature S\mathtt{S} and multiple independent complex features Sc\mathtt{S}^{c}. In our experiments, all datasets have sample sizes that are large enough for all models considered in the paper to learn complex features Sc\mathtt{S}^{c} and attain optimal test accuracy, even in the absence of S\mathtt{S}; we use sample sizes of 5000050000 for LMS-5 and MS-5 and 4000040000 for L̂MS-7.

LSN Dataset: Recall that the LSN dataset (described in Section 3) is a stylized version of the LMS-k that is amenable to theoretical analysis. In LSN, conditioned on the label yy, the first and second coordinates of xx are singleton linear and 33-slab blocks: linear and 33-slab blocks have support on \{\scalebox{0.75}[1.0]{-}1,1\} and \{\scalebox{0.75}[1.0]{-}1,0,1\} respectively. The remaining coordinates are standard gaussians and not predictive of the label. Each data point (xi,yi)∈ℜd×{−1,1}(x_{i},y_{i})\in\Re^{d}\times\{-1,1\} can be sampled as follows:

Additional Datasets: We now introduce M̂S-(5,7), the noisy version of MS-(5,7), and three MNIST-CIFAR datasets, each with different MNIST and CIFAR10 classes.

M̂S-(5,7): Noisy 5-slab and multiple noiseless 7-slab blocks; the first coordinate is a noisy 5-slab block and the remaining d\scalebox{0.75}[1.0]{-}1 coordinates are independent 77-slab blocks. Note that this dataset comprises a noisy-but-simpler 55-slab block and multiple noiseless 77-slab blocks; a 66-piecewise linear classifier can attain 100% accuracy by learn any 77-slab block.

MNIST-CIFAR datasets: Recall that images in the MNIST-CIFAR datasets are concatenations of MNIST and CIFAR10 images. We introduce additional variants of the MNIST-CIFAR using different class pairs to show that our results in the paper are robust to the exact choice of pairs:

Models: Here, we briefly describe the models (and its abbreviations) used in the paper. We use fully-connected (FCNs), convolutional (CNNs), and sequential neural networks (GRUs ) on synthetic datasets. Abbreviations (w,d)(w,d)-FCN denotes FCN with width ww and depth dd, (f,k,d)(f,k,d)-CNN denotes dd-layer CNNs with ff filters of size k×kk\times k in each layer with and (h,l,d)(h,l,d)-GRU denotes dd-layer dd-layer GRU with input dimensionality ll and hidden state dimensionality hh. On MNIST-CIFAR, we train MobileNetV2 , GoogLeNet , ResNet50 and DenseNet121 .

Training Procedures: Unless mentioned otherwise, we use the following hyperparameters for standard training and adversarial training on synthetic and MNIST-CIFAR data:

Appendix C Additional Results on the Extreme Nature of Simplicity Bias (SB)

Recall that Section 4 of the paper establishes the extreme nature of SB: If all features have full predictive power, NNs rely exclusively on the simplest feature S\mathtt{S} and remain invariant to all complex features Sc\mathtt{S}^{c}—in Section 4 of the paper. Now, we further validate the extreme nature of SB across model architectures, datasets, optimizers, activation functions and regularization. We also analyze the effect of input dimensionality, number of complex features, choice and scaling of random initialization and non-random initialization.

In this section, we supplement our results in Section 4 of the paper by showing that extreme simplicity bias (SB) persists across several model architectures and on synthetic as well as image-based datasets. In Table 4, we present {S\mathtt{S},Sc\mathtt{S}^{c} }-Randomized AUCs for FCNs, CNNs and GRUs with depth {1,2} trained on LMS and MS-(5,7) datasets and state-of-the-art CNNs trained on MNIST-CIFAR:A. While the Sc\mathtt{S}^{c}-randomized AUC equals 1.001.00 (perfect classification), we see that the S\mathtt{S}-randomized AUCs are approximately 0.50.5 for all models. This is because all models essentially only rely on the simplest feature S\mathtt{S} and remain invariant to all complex features Sc\mathtt{S}^{c}, even though all features have equal predictive power.

C.2 Effect of MNIST-CIFAR Class Pairs

In this section, we supplement our results on MNIST-CIFAR (in Section 4) in order to show that extreme SB observed in MobileNetV2 , GoogLeNet , ResNet50 and DenseNet121 does not depend on the exact choice of MNIST and CIFAR10 class pairs used to construct the MNIST-CIFAR datasets. To do so, we evaluate the MNIST-randomized and CIFAR10-randomized metrics of the aforementioned models on three datasets–MNIST-CIFAR:A, MNIST-CIFAR:B, MNIST-CIFAR:C—described in Appendix B.

Table 5 presents the standard, MNIST-randomized and CIFAR10-randomized AUC values of MobileNetV2, GoogLeNet, ResNet50 and DenseNet121 on three MNIST-CIFAR datasets. We observe that randomizing over the simpler MNIST block is sufficient to fully degrade the predictive power of all models; for instance, randomizing the MNIST block drops the AUC values of ResNet50 from 1.01.0 to 0.50.5 (i.e., equivalent to random classifier). However, randomizing the CIFAR10 block has no effect—standard AUC and CIFAR10-randomized AUCs equal 1.01.0. In contrast, an ideal classifier that relies on MNIST & CIFAR10 would attain non-trivial AUC even when the MNIST block is randomized.

C.3 Effect of Optimizers and Activation Functions

Now, we study the effect of activation function and optimizer on extreme SB. That is, can the usage of different activation functions and optimizer encourage trained neural networks to rely on complex features Sc\mathtt{S}^{c} in addition to the simplest feature S\mathtt{S}?

Table 6 presents the S\mathtt{S}-randomized AUCs of (100,2)(100,2)-FCNs with multiple activation functions—ReLU, Leaky ReLU , PReLU , and Tanh—trained on LMS-7 and MS-(5,7) datasets using multiple commonly-used optimizers: SGD, Adam ,and RMSProp . We observe that for all combinations of activations and optimizers, trained FCNs still only rely on simplest feature S\mathtt{S}; S\mathtt{S}-randomized and Sc\mathtt{S}^{c}-randomized AUCs are approximately 0.500.50 and 1.01.0 respectively for all optimizers and activation functions. Therefore, in addition to SGD, commonly used first-order optimization methods such as Adam and RMSProp cannot jointly learn large-margin classifiers that rely on learn slab-structured features in the presence of a noisy linear structure. To summarize, the experiment in Section C.2 shows that simply altering the choice of optimizer and activation function does not have any effect on extreme SB. Similar to the experiments in Section 4 of the paper, all models exclusively rely on simplest feature S\mathtt{S} and remain invariant to complex features Sc\mathtt{S}^{c}.

C.5 Effect of Input Dimension and Number of Complex Features

In this section, we evaluate the performance of FCNs trained on LMS-7 data using SGD to show the extreme SB persists in the low-dimensional setting (d<10d<10) and also with varying number of 77-slab features 1≤∣Sc∣≤d1\leq|\mathtt{S}^{c}|\leq d.

As shown in Table 8, decreasing the input dimension dd or the number of complex features ∣Sc∣|\mathtt{S}^{c}| has no effect on extreme SB of FCNs trained on the LMS-7 dataset. Similar to our results in Figure 3, the standard and randomized AUCs collectively show that the SGD-trained (100,1)-FCNs exclusively rely on the linear component and do not rely on the 77-slab coordinates.

C.6 Effect of Random Initialization Scale

Now, we analyze the effect of the choice and scale (i.e., magnitude of the weights of randomly initialized FCNs) of random initialization on simplicity bias using FCNs trained on LMS-7 data.

As shown in Table 9, the choice (Kaiming and Xavier) and the scale of random initialization do not alter the extreme SB phenomenon on the LMS-7 dataset. That is, scaling the randomly initialized models by up to 0.10.1 and 10.010.0 has no effect on simplicity bias—SGD-trained (100,1)-FCNs exclusively rely on the linear component and do not rely on the 77-slab coordinates.

C.7 Effect of Non-random Initialization

In this section, we investigate the effect of non-random initialization on simplicity bias using FCNs trained on MS-7 and LMS-7 data. The goal of this experiment is to determine the extent to which simplicity bias persists when the untrained network (at timestep t=0t=0) attains non-random standard accuracy by relying on one or more “complex" 77-slab features.

To vary the degree of non-random initialization α\alpha, we obtain model Mα\mathtt{M}_{\alpha} by linearly interpolating the weights of a randomly initialized network Mrand\mathtt{M}_{\text{rand}} and a network Mslab\mathtt{M}_{\text{slab}} that exclusively relies on one or more “complex" 77-slab features to attain 100%100\% accuracy on the LMS-7 dataset. That is, Mα≡α⋅Mslab+(1−α)⋅Mrand\mathtt{M}_{\alpha}\equiv\alpha\cdot\mathtt{M}_{\text{slab}}+(1-\alpha)\cdot\mathtt{M}_{\text{rand}}. Note that Mslab\mathtt{M}_{\text{slab}} is trained on MS-7 data to attain 100%100\% standard accuracy by relying on one or more 77-slab coordinates. As shown in Figure 5(a), increasing the interpolation constant α\alpha monotonically increases the standard accuracy of Mα\mathtt{M}_{\alpha} on MS-7 data.

Now, we use the linearly interpolated model Mα\mathtt{M}_{\alpha} (for varying values of α\alpha) as initialization and train Mα\mathtt{M}_{\alpha} on LMS-7 data, which additionally consists of a “simple" linearly-separable coordinate. To maintain the input dimensionality, we obtain LMS-7 data by replacing a 77-slab coordinate by the simpler linear coordinate. In order to maintain the non-random accuracy of Mα\mathtt{M}_{\alpha}, we use coordinate-randomized AUCs to choose and replace a 77-slab coordinate that the model does not depend on.

Appendix D Additional Results on the Effect of Extreme SB on Generalization

Recall that in Section 5 of the paper, we showed that extreme SB can result in suboptimal generalization of SGD-trained models on the same data distribution. In this section, we present additional information about the experiment setup used in Section 5 to show that SB can worsen standard generalization.

We now provide additional information about the experimental setup used in Section 5 of the paper, where we show that extreme simplicity bias can result in suboptimal generalization. We train fully-connected networks (FCNs) of width {100,200,300}\{100,200,300\} and depth {1,2}\{1,2\} using SGD on 5050-dimensional L̂MS-7 dataset of 4000040000 samples, which comprises of a noisy linear coordinate (10%10\% noise) and 4949 77-slab coordinates. For each model architecture, we perform a grid search over SGD hyperparameters—learning rate, batch size, momentum, and weight decay—and report standard and randomized test accuracies of the model (in Table 2) that perform best on a L̂MS-7 validation dataset. We perform a grid search over the following SGD hyperparameters:

Learning rate: {0.001,0.01,0.05,0.1,0.3}\{0.001,0.01,0.05,0.1,0.3\}

Appendix E Can we mitigate Simplicity Bias?

In this section, we investigate whether standard approaches for improving generalization error and adversarial robustness—ensembles and adversarial training—help in mitigating SB.

We now study the extent to which ensembles mitigate SB and its adverse effect on generalization. Specifically, we evaluate the performance of ensembles of fully-connected networks (FCNs) that are trained on two datasets: L̂MS-7 and MS-5. Recall that the L̂MS-7 data comprises one simple-but-noisy linear coordinate and multiple relatively complex 77-slab coordinates that have no noise, whereas MS-5 data comprises multiple noiseless 55-slab coordinates only.

To better highlight the effect of ensembles on generalization, we choose a sample size (for both datasets) such that individual models (a) overfit to training data (i.e., non-zero generalization gap) but (b) still attain non-trivial test accuracy. We now discuss the performance of ensembles of independently trained models on MS-5 and L̂MS-7 datasets:

MS-5 data: Recall that MS-5 data comprises multiple independent 5-slab blocks, one in each coordinate, that have equal simplicity and predictive power. Thus, since all features have equal simplicity, independent SGD-trained (100,2)-FCN end up relying on different 55-slab coordinates due to random initialization, as shown in Figure 6(b). As the training sample size is small, FCNs overfit to the training data and attain approximately 75%75\% test accuracy, as shown in Figure 6. Consequently, as shown in Figure 6, ensembles of these models rely on all 55-slab coordinates learned by the individual models and attain better test accuracy by aggregating model predictions and averaging out overfitting. For example, Figure 6 shows that ensembles of size 55 and 1010 improves generalization by approximately 15%15\% and 20%20\% respectively.

L̂MS-7 data: Recall that L̂MS-5 data comprises one simple-but-noisy linear block (with 50%50\% noise) and multiple independent 77-slab blocks that have no noise. Now, due to extreme SB, every independently trained FCN exclusively latches on (and overfits to) the simpler-but-noisy linear block, as shown in Figure 6(b). As a result, all models collectively lack diversity and essentially learn the same decision boundary because of extreme SB. Therefore, ensembles of these models do not improve generalization because the independent models make misclassifications on the same instances. As shown in Figure 6(a), ensembles of size 33, 55 and 1010 do not improve generalization—the test accuracy remains 75%75\%.

The ensemble performance on MS-5 data indicates that when datasets have multiple equally simple features, ensembles of independently trained models mitigate SB to some extent by aggregating predictions of models that rely on simple features. Conversely, the ensemble performance on L̂MS-7 data suggests that when datasets comprise few features that the more noisy and less predictive than the rest, ensembles may not improve generalization. Our results also suggest that the generalization improvements using ensemble methods in practice may stem from combining multiple simple-but-noisy features (such as color, texture) and not by learning complex features (such as shape).

E.2 Adversarial Training

We now investigate the extent to which adversarial training mitigates SB and its adverse effect on adversarial robustness using two datasets: MNIST-CIFAR and AdvMS-(5,7).

Now, we first introduce AdvMS-(5,7), a variant of the MS-(5,7) dataset, and then investigate if adversarially trained FCNs that improve adversarial robustness by some extent also mitigate extreme simplicity bias (SB).

AdvMS-(5,7) dataset: Recall that dd-dimensional MS-(5,7) data, introduced in Section 3, consists of d−1d-1 77-slab coordinates and a single relatively simpler 55-slab coordinate, all of which have perfect predictive power. Similar to MS-(5,7) data, the dd-dimensional AdvMS-(5,7) data comprises 55-slab and 77-slab coordinates. Specifically, the first \nicefracd2\nicefrac{{d}}{{2}} coordinates correspond to independent 55-slabs, each with effective margin γ5\gamma_{5} and the other \nicefracd2\nicefrac{{d}}{{2}} coordinates correspond to independent 77-slabs with effective margin γ7\gamma_{7}. In contrast to MS-(5,7) data, the AdvMS-(5,7) dataset (a) comprises \nicefracd2\nicefrac{{d}}{{2}} 55-slab coordinates and (b) the 55-slabs and 77-slabs do not necessarily share the same effective margin. In our experiments below, we set d=20d=20, γ5=0.05\gamma_{5}=0.05 and γ7=0.15\gamma_{7}=0.15. That is, we conduct our experiments on 2020-dimensional data in which the the simple features S\mathtt{S} and complex features Sc\mathtt{S}^{c} correspond to the 1010 small-margin 55-slab and 1010 large-margin 77-slab coordinates respectively.

Also note that (a) adversarial perturbations are generated using PGD attacks , (b) (200,2)(200,2)-FCNs and (1000,2)(1000,2)-FCNs are expressive enough to learn the maximum-margin classifier on the 2020-dimensional AdvMS-(5,7) data, (c) FCNs are adversarially trained for 40004000 epochs with initial learning rate 0.10.1 that decays by a multiplicative factor of 0.10.1 after every 10001000 epochs, and (d) the training data comprises 60006000 data points, which is enough for SGD-trained (1000,2)(1000,2)-FCNs to learn 77-slab coordinates and attain 100%100\% generalization.

Adversarial trained FCNs do not learn maximum-margin classifiers. When ϵ≤0.1\epsilon\leq 0.1, (200,2)(200,2)-FCNs learn classifiers that attain 100%100\% standard and ϵ\epsilon-robust accuracies. However, when ϵ≥0.2\epsilon\geq 0.2, due to optimization-related issues, adversarially trained (200,2)(200,2)-FCNs are unable to learn a non-trivial classifier that obtains more than 50%50\% standard and ϵ\epsilon-robust accuracy. Increasing the model width from 200200 to 10001000 improves adversarial robustness to some extent—adversarially trained (1000,2)(1000,2)-FCNs learn classifiers with 100%100\% standard and robust accuracies when ϵ≤0.25\epsilon\leq 0.25. However, when ϵ≥γS=0.3\epsilon\geq\gamma_{\mathtt{S}}=0.3 (dashed purple line), adversarially trained (2000,2)(2000,2)-FCNs are unable to learn ϵ\epsilon-robust classifiers as well. We note that further increasing the model width to 20002000 does not improve robustness. Consequently, adversarial training does not result in maximum-margin classifiers that have optimal adversarial robustness (i.e., classifiers with 100%100\% γdata\gamma_{\text{data}}-robust accuracy) on AdvMS-(5,7) data. These results reconcile two phenomena observed in practice: larger capacity models can improve adversarial robustness , but large-epsilon adversarial training can “fail" and result in trivial classifiers due to optimization-related issues .

Adversarial training does not mitigate extreme SB. The {S\mathtt{S},Sc\mathtt{S}^{c} }-randomized accuracies of adversarially trained FCNs collectively show that adversarial training does not mitigate extreme SB. When the perturbation budget ϵ≤γS\epsilon\leq\gamma_{\mathtt{S}}, adversarially trained (2000,2)(2000,2)-FCNs exhibit robustness by exclusively relying on multiple “simple" 55-slab coordinates. That is, randomizing S\mathtt{S} drops the model accuracy to 50%50\%, but randomizing the more complex 77-slab coordinates has no effect on model accuracy. Conversely, when ϵ≥γS\epsilon\geq\gamma_{\mathtt{S}}, classifiers with 100%100\% ϵ\epsilon-robust accuracy must rely on features in S\mathtt{S} and Sc\mathtt{S}^{c}. However, as shown in Figure 7, when ϵ≥γS\epsilon\geq\gamma_{\mathtt{S}}, adversarial training fails and results in trivial classifiers that attain 50%50\% standard and robust accuracy.

E.2.2 Adversarially training CNNs on MNIST-CIFAR data

Table 11 evaluates the standard, ϵ\epsilon-robust and CIFAR10-randomized accuracies of SGD-trained and adversarially trained MobileNetV2, ResNet50 and DenseNet121 using the MNIST-CIFAR dataset. First, we observe that adversarial training with perturbation norm 0.30.3 significantly improves ϵ\epsilon-robust accuracies over those of SGD-trained models, without degrading the models’ standard test accuracies. However, the CIFAR10-randomized accuracies indicate that the adversarial training does not lead to reliance on the CIFAR10 block—adversarially trained CNNs continue to remain invariant to the CIFAR10 block even though it is almost fully predictive of the label. Consequently, these results suggest that adversarial training improves robustness but does not achieve the best ϵ\epsilon-robust accuracy on the MNIST-CIFAR dataset.

To summarize, our experiments on AdvMS-(5,7) and MNIST-CIFAR datasets show that while adversarial training does improve the ϵ\epsilon-robust accuracy over that of SGD-trained model, adversarially trained models continue to remain susceptible to extreme SB and consequently do not achieve maximum possible adversarial robustness.

Appendix F Proof of Theorem 1

In this section, we first re-introduce the data distribution and theorem. Then, we describe the proof sketch and notation, before moving on to the proof.

Linear-Slab-Noise (LSN) data: The LSN dataset is a stylized version of LMS-k that is amenable to theoretical analysis. In LSN, conditioned on the label yy, the first and second coordinates of xx are singleton linear and 33-slab blocks: linear and 33-slab blocks have support on \{\scalebox{0.75}[1.0]{-}1,1\} and \{\scalebox{0.75}[1.0]{-}1,0,1\} respectively. The remaining coordinates are standard gaussians and not predictive of the label. Each data point (xi,yi)∈ℜd×{−1,1}(x_{i},y_{i})\in\Re^{d}\times\{-1,1\} from LSN can be sampled as follows:

According to Theorem 1 (re-stated), one-hidden-layer ReLU neural networks trained with standard mini-batch gradient descent (GD) on the LSN dataset provably learns a classifier that exclusively relies on the “simple" linear coordinate, thus exhibiting simplicity bias at the cost of margin.

Let f(x)=∑j=1kvj⋅ReLU(∑i=1dwi,jxi)f(x)=\sum_{j=1}^{k}v_{j}\cdot\textrm{ReLU}(\sum_{i=1}^{d}w_{i,j}x_{i}) denote a one-hidden-layer neural network with kk hidden units and ReLU activations. Set vj=±\nicefrac1kv_{j}=\pm\nicefrac{{1}}{{\sqrt{k}}} w.p. \nicefrac12\nicefrac{{1}}{{2}} ∀j∈[k]\forall j\in[k]. Let {(xi,yi)}i=1m\{(x^{i},y^{i})\}^{m}_{i=1} denote i.i.d. samples from LSN where m∈[cd2,dα/c]m\in[cd^{2},d^{\alpha}/c] for some α>2\alpha>2. Then, given d>Ω(klog⁡k)d>\Omega(\sqrt{k}\log k) and initial wij∼N(0,1dklog⁡4d)w_{ij}\sim\mathcal{N}(0,\frac{1}{dk\log^{4}d}), after O(1)O(1) iterations, mini-batch gradient descent (over ww) with hinge loss, step size \eta=\Omega{{(\log d)^{\nicefrac{{\scalebox{0.75}[1.0]{-}1}}{{2}}}}}, mini-batch size Θ(m)\Theta(m), satisfies:

Test error is at most \nicefrac1poly(d)\nicefrac{{1}}{{\textrm{poly}(d)}}

The learned weights of hidden units wijw_{ij} satisfy:

with probability greater than 1−1poly(d)1-\frac{1}{\textrm{poly}(d)}. Note that cc is a universal constant.

Proof Sketch Since the number of iterations t=O(1)t=O(1), we partition the dataset into tt minibatches each of size n:=m/tn:=m/t samples. This means that each iteration uses a fresh batch of nn samples and the tt iterations together form a single pass over the data. The overall outline of the proof is as follows. If the step size is η\eta, then for t≲4ηt\lesssim\frac{4}{\eta} iterations, with probability ≥1−1poly(d)\geq 1-\frac{1}{\textrm{poly}(d)},

Lemma 2 shows that the hinge loss is “active" (i.e., yf(x)<1yf(x)<1) for all data points in a given batch.

Under this condition, we derive closed-form expressions for population gradients in Lemmas 4, 5 and 6.

Lemma 1 uses the above lemmas to establish precise estimates of the linear, slab and noise coordinates for all iterations until tt.

The proof is organized as follows. Appendix F.1 presents the main lemmas that will directly lead to Theorem 1. Appendix F.2 derives closed form expressions for population gradients and Appendix F.3 presents auxiliary lemmas that are useful in the main proofs.

The proof directly follows from Lemma 1 and Lemma 2. In Lemma 1, we show that the weights in the linear coordinate are Ω(d)\Omega(\sqrt{d}) larger than the weights in the slab and noise coordinates. Applying Lemma 1 at t^=⌊4η(1−cnlog⁡d)⌋\hat{t}=\lfloor\frac{4}{\eta}(1-\frac{c_{n}}{\sqrt{\log d}})\rfloor gives the following result:

where (a)(a) is due to c0(1+c^)t≤c0ec^t≤c0e1=O(1)c_{0}(1+\hat{c})^{t}\leq c_{0}e^{\hat{c}t}\leq c_{0}e^{1}={O}(1).

The 0−10-1 error of the function ff at timestep t^\hat{t} is small as well, because we can directly use Lemma 2 to get Pr⁡(yf(x)<0)=\nicefrac2c3d6\Pr(yf(x)<0)=\nicefrac{{2}}{{c^{3}d^{6}}}. Therefore, the 0−10-1 error is at most 2c3d6=O(1d6)\frac{2}{c^{3}d^{6}}={O}(\frac{1}{d^{6}}). ∎

In this section, we use proof by induction to show that for the first t=O(\nicefrac1η)t=O(\nicefrac{{1}}{{\eta}}) steps, (1) the hinge loss is “active" for all data points (Lemma 2) and (2) hidden layer weights in the linear coordinate are Ω(d)\Omega(\sqrt{d}) larger than the hidden layer weights in the slab and noise coordinates (Lemma 1).

Let ∣Sn∣∈[cd2,dα/c]|\mathcal{S}_{n}|\in[cd^{2},d^{\alpha}/c] and initialization wij∼N(0,\nicefrac1dklog⁡2d)w_{ij}\sim\mathcal{N}(0,\nicefrac{{1}}{{dk\log^{2}d}}). Also let c^=\nicefracη4\hat{c}=\nicefrac{{\eta}}{{4}}, c0=2c_{0}=2 and cn=5αc0(1+c^)tc_{n}=5\sqrt{\alpha}c_{0}(1+\hat{c})^{t}. Then, for all t≤4η(1−\nicefraccnlog⁡d)t\leq\frac{4}{\eta}(1-\nicefrac{{c_{n}}}{{\sqrt{\log d}}}), d≥exp⁡((\nicefrac8cnη)2)d\geq\exp((\nicefrac{{8c_{n}}}{{\eta}})^{2}), \nicefracdlog⁡3(d)>\nicefrac24kc0c\sqrt{\nicefrac{{d}}{{\log^{3}(d)}}}>\nicefrac{{24\sqrt{k}}}{{c_{0}c}} and i∈[k]i\in[k], w.p. greater than 1−O(1d2)1-O(\frac{1}{d^{2}}), we have:

First, we prove that equations (2), (3) & (4) hold at initialization (i.e., t=0t=0) with high probability. Using 7 and 1:

Therefore, w1i=(0)ηvi2±c0(1+c^)0dklog⁡dw_{1i}=\frac{(0)\eta v_{i}}{2}\pm\frac{c_{0}(1+\hat{c})^{0}}{\sqrt{dk}\log d} and ∣w2i∣≤c0(1+c^)0dklog⁡d|w_{2i}|\leq\frac{c_{0}(1+\hat{c})^{0}}{\sqrt{dk}\log d} and ∣∣wˉi∣∣≤c0(1+c^)0klog⁡d||\bar{w}_{i}||\leq\frac{c_{0}(1+\hat{c})^{0}}{\sqrt{k}\log d}. Since equations (2), (3) & (4) hold at t=0t=0, we can use Lemma 2 to show that the hinge loss is “active" with high probability:

Now, we assume that the inductive hypothesis—equations (1), (2), (3) and (4)—is true after every timestep τ\tau where τ∈{0,⋯ ,t}\tau\in\{0,\cdots,t\}.

We now prove that the inductive hypothesis is true at timestep t+1t+1, after applying gradient descent using the (t+1)th(t+1)^{th} batch. Since z(1) holds at timestep tt, we can use the closed-form expression of the gradient along the linear coordinate (lemma 4) to prove that equation (2) holds at timestep t+1t+1 as well:

where (a)(a) is via equation (12) in Lemma 3, (b)(b) is because \nicefracdlog⁡3(d)≥\nicefrac20c0e1c\nicefrac{{d}}{{\log^{3}(d)}}\geq\nicefrac{{20}}{{c_{0}e^{1}\sqrt{c}}} and (c)(c) is due to ηvi≤c^\eta v_{i}\leq\hat{c}.

Similarly, since equation (1) holds at timestep tt (via the inductive hypothesis), we can use the closed-form expression of the gradient along the slab coordinate (lemma 5) to show that the weights in the slab (i.e., second) coordinate are small (equation (3)) at timestep t+1t+1 as well:

where (a)(a) is due to equations (11) in Lemma 3, (3) and \nicefracdlog⁡3(d)≥\nicefrac20c0e1c\nicefrac{{d}}{{\log^{3}(d)}}\geq\nicefrac{{20}}{{c_{0}e^{1}\sqrt{c}}}.

Finally, we can use the closed-form expression of the gradient along the noise coordinate (lemma 6) to prove that the norm of the gradient along the noise coordinates (i.e., coordinates 33 to dd) is small (equation (4)) at timestep t+1t+1:

We first show that the the first part of the noise gradient, Gˉ1\bar{\mathcal{G}}_{1}, is at most \nicefracηvi2\nicefrac{{\eta v_{i}}}{{2}}:

where (a)(a) is because \nicefracdlog⁡d≥(\nicefrac24kc0c)2\nicefrac{{d}}{{\log d}}\geq(\nicefrac{{24\sqrt{k}}}{{c_{0}c}})^{2}, (b)(b) is due to equation (4) and (c)(c) is because ηvi≤c^\eta v_{i}\leq\hat{c}. ∎

Since equations (2), (3) & (4) hold at timestep tt (from Lemma 1), we can show that the hinge loss is positive (i.e., yf(x)<1yf(x)<1) for all data points with high probability as well.

Let Sn\mathcal{S}_{n} denote a set of n∈[cd2,dα/c]n\in[cd^{2},d^{\alpha}/c] i.i.d. samples from LSN, where α>2\alpha>2 and c>1c>1. Suppose equations (2), (3) & (4) hold at timestep tt. Also let d≥exp⁡((8cnη)2)d\geq\exp((\frac{8c_{n}}{\eta})^{2}) where cn=5αc0(1+c^)tc_{n}=5\sqrt{\alpha}c_{0}(1+\hat{c})^{t}. Then, w.p. greater than 1−2c3d61-\frac{2}{c^{3}d^{6}}, we have:

We use equations (2), (3) & (4) to obtain simplify the dot product between wi(t)w^{(t)}_{i} & xjx_{j} and the indicator \mathds1{wi(t)⋅xj≥0}\mathds{1}{\left\{w^{(t)}_{i}\cdot x_{j}\geq 0\right\}}. First, we show that the dot product between wi(t)w^{(t)}_{i} and xjx_{j} is in the band tηviyj2±cnklog⁡d\frac{t\eta v_{i}y_{j}}{2}\pm\frac{c_{n}}{\sqrt{k\log d}} with high probability:

where (a)(a) is because wˉi(t)⋅xˉj=∣∣wˉi(t)∣∣N(0,1)\bar{w}^{(t)}_{i}\cdot\bar{x}_{j}=||\bar{w}^{(t)}_{i}||\mathcal{N}(0,1), (b)(b) is via lemma 7 & c>1c>1, and (c)(c) is because (yj+yj+12εj)<2(y_{j}+\frac{y_{j}+1}{2}\varepsilon_{j})<2. Next, when d≥exp⁡((8cnη)2)d\geq\exp((\frac{8c_{n}}{\eta})^{2}), we can simplify \mathds1{wi(t)⋅xj≥0}\mathds{1}{\left\{w^{(t)}_{i}\cdot x_{j}\geq 0\right\}} as follows:

We can now use equations (6) & (8) to show that yjf(t)(xj)y_{j}f^{(t)}(x_{j}) is in the band \nicefractη4±O(\nicefrac1log⁡d)\nicefrac{{t\eta}}{{4}}\pm O(\nicefrac{{1}}{{\sqrt{\log d}}}) with high probability:

where (a)(a) is due to ∣{vi  ∣  vi>0}∣=∣{vi  ∣  vi<0}∣=\nicefrack2|\{v_{i}\;|\;v_{i}>0\}|=|\{v_{i}\;|\;v_{i}<0\}|=\nicefrac{{k}}{{2}} and (b)(b) follows from \nicefraccnlog⁡d≤\nicefracη8\nicefrac{{c_{n}}}{{\sqrt{\log d}}}\leq\nicefrac{{\eta}}{{8}} when d≥exp⁡((\nicefrac8cnη)2)d\geq\exp((\nicefrac{{8c_{n}}}{{\eta}})^{2}) ∎

If equations (2), (3) & (4) hold at timestep tt, d>exp⁡((4c0e1η)2)d>\exp((\frac{4c_{0}e^{1}}{\eta})^{2}) and \nicefracdlog⁡d>k\nicefrac{{d}}{{\log d}}>\sqrt{k}, we have:

Let gz(x)=1xφ(zx)g_{z}(x)=\frac{1}{x}\varphi(\frac{z}{x}) and h(x)=max⁡∣δ∣≤∣w2i(t)∣1xφ(w1i(t)+δx)=max⁡∣δ∣≤∣w2i(t)∣gw1it+δ(x)h(x)=\max_{|\delta|\leq|w^{(t)}_{2i}|}\frac{1}{x}\varphi(\frac{w^{(t)}_{1i}+\delta}{x})=\max_{|\delta|\leq|w^{(t)}_{2i}|}g_{w^{t}_{1i}+\delta}(x). To prove Equation 10, we show that an upper bound on ∣∣w(t)∣∣||w^{(t)}|| is less than a lower bound on arg⁡max⁡xh(x)\arg\max_{x}h(x), which subsequently implies that h(∣∣w(t)∣∣)<max⁡xh(x)h(||w^{(t)}||)<\max_{x}h(x) because hh is an increasing function for all ∣x∣≤arg⁡max⁡xh(x)|x|\leq\arg\max_{x}h(x).

First, we find the maximizer x∗x^{*} of h(x)h(x) as follows:

where (a)(a) follows from lemma 11. Next, we lower bound the maximizer x∗x^{*} of h(x)h(x):

From lemma 11, we know that h(x)h(x) is an increasing function for all ∣x∣<x∗|x|<x^{*}. This implies that h(∣∣wi(t)∣∣)≤h(c0(1+c^)tdklog⁡d)≤h(ηvi4)≤h(x∗)h(||w^{(t)}_{i}||)\leq h(\frac{c_{0}(1+\hat{c})^{t}}{\sqrt{dk}\log d})\leq h(\frac{\eta v_{i}}{4})\leq h(x^{*}). Therefore, when d≥exp⁡((4c0e1η)2)d\geq\exp((\frac{4c_{0}e^{1}}{\eta})^{2}) and \nicefracdlog⁡d≥k\nicefrac{{d}}{{\log d}}\geq\sqrt{k}, we obtain the desired result as follows:

Now, we can prove equations (11), (12) and (13) using equation (10) as follows:

F.2 Closed-form Gradient Expressions

In this section, we provide closed-form expressions for gradients along the linear, slab and noise coordinates: ∇w1iLf(Sn)\nabla_{w_{1i}}\mathcal{L}_{f}(\mathcal{S}_{n}), ∇w2iLf(Sn)\nabla_{w_{2i}}\mathcal{L}_{f}(\mathcal{S}_{n}) and ∇wˉiLf(Sn)\nabla_{\bar{w}_{i}}\mathcal{L}_{f}(\mathcal{S}_{n}). First, we provide a closed-form expression for the gradient along the linear coordinate:

If n>cd2n>cd^{2} and yif(xi)<1 ∀(xi,yi)∈Sny_{i}f(x_{i})<1\,\forall(x_{i},y_{i})\in\mathcal{S}_{n}, then w.p. greater than 1−3n1-\frac{3}{n}:

where (a)(a) is due to yixi1=yi2=1y_{i}x_{i1}=y^{2}_{i}=1 & \mathds1{yjf(xj)≤1}=1\mathds{1}{\left\{y_{j}f(x_{j})\leq 1\right\}}=1 and (b)(b) is due to \mathds1{wˉiTxˉj≥k}=\mathds1{∣∣wˉi∣∣Zj≥k}\mathds{1}{\left\{\bar{w}^{T}_{i}\bar{x}_{j}\geq k\right\}}=\mathds{1}{\left\{||\bar{w}_{i}||Z_{j}\geq k\right\}}. ∎

Similarly, we provide a closed-form expression for the gradient along the slab coordinate:

If n>cd2n>cd^{2} and yif(xi)<1 ∀(xi,yi)∈Sny_{i}f(x_{i})<1\,\forall(x_{i},y_{i})\in\mathcal{S}_{n}, then w.p. greater than 1−3n1-\frac{3}{n}:

where (a)(a) is due to yixi2=\mathds1{yi=1}εiy_{i}x_{i2}=\mathds{1}{\left\{y_{i}=1\right\}}\varepsilon_{i} & \mathds1{yjf(xj)≤1}=1\mathds{1}{\left\{y_{j}f(x_{j})\leq 1\right\}}=1 and (b)(b) is due to \mathds1{wˉiTxˉj≥k}=\mathds1{∣∣wˉi∣∣Zj≥k}\mathds{1}{\left\{\bar{w}^{T}_{i}\bar{x}_{j}\geq k\right\}}=\mathds{1}{\left\{||\bar{w}_{i}||Z_{j}\geq k\right\}}. ∎

Next, we provide a closed-form expression for the gradient along the noise coordinates:

If n>cd2n>cd^{2} and yif(xi)<1 ∀(xi,yi)∈Sny_{i}f(x_{i})<1\,\forall(x_{i},y_{i})\in\mathcal{S}_{n}, then w.p. greater than 1−13n1-\frac{1}{3n}:

where ui⊥u^{\perp}_{i} is some unit vector orthogonal to wˉi\bar{w}_{i}.

Next, we show that the projection of ∇wˉiLf(Sn)\nabla_{\bar{w}_{i}}\mathcal{L}_{f}(\mathcal{S}_{n}) onto S⊥S^{\perp} (i.e., case 2) has small norm w.p. greater than 1−1d1-\frac{1}{d}:

where (a)(a) is because xjS⊥xˉSj⊥x^{S}_{j}\perp\bar{x}^{S^{\perp}_{j}}, (b)(b) is via fact 1 and (c)(c) is due to n≥cd2n\geq cd^{2}. Next, we show that the norm of the gradient in the direction of wˉi\bar{w}_{i} (i.e., case 1) is close to Gˉ\bar{\mathcal{G}} w.h.p.:

where (a)(a) is because xˉjS=wˉiTxj∣∣wˉi∣∣2wˉi\bar{x}^{S}_{j}=\frac{\bar{w}^{T}_{i}x_{j}}{||\bar{w}_{i}||^{2}}\bar{w}_{i} and (b)(b) is because (b)(b) is due to \mathds1{wˉiTxˉj≥k}=\mathds1{∣∣wˉi∣∣Zj≥k}\mathds{1}{\left\{\bar{w}^{T}_{i}\bar{x}_{j}\geq k\right\}}=\mathds{1}{\left\{||\bar{w}_{i}||Z_{j}\geq k\right\}}. Therefore, by combining the results in case 1 and 2, the following holds w.p. greater than 1−13n1-\frac{13}{n}:

F.3 Miscellaneous Lemmas

Let Xi∼N(0,σ2)X_{i}\sim\mathcal{N}(0,\sigma^{2}) and δ∈(0,1)\delta\in(0,1). Then, max⁡i∈[k]∣Xi∣≤σ2log⁡(2kδ)\max_{i\in[k]}|X_{i}|\leq\sigma\sqrt{2\log(\frac{2k}{\delta})} with probability greater than 1−δ1-\delta,

Let φ\varphi denote the probability density function of the standard normal. Also let Z∼N(0,1)Z\sim\mathcal{N}(0,1). Then, for t≥1t\geq 1, we have:

where (a)(a) is because \varphi^{\prime}(x)=\scalebox{0.75}[1.0]{-}x\varphi(x). Using union bound with t=2log⁡(2kδ)≥1 ∀δ∈(0,1)t=\sqrt{2\log(\frac{2k}{\delta})}\geq 1\,\forall\delta\in(0,1) gives the desired result. ∎

Let ϕ\phi and φ\varphi denote the cumulative distribution function and the probability density function of the standard gaussian. Then, for any Z∼N(0,1)Z\sim\mathcal{N}(0,1) and k∈\mathdsRk\in\mathds{R}:

The expectation \mathdsE[\mathds1{Z≥k}Z]\mathds{E}[\mathds{1}{\left\{Z\geq k\right\}}Z] can be simplified as follows:

where (a)(a) is due to \varphi^{\prime}(x)=\scalebox{0.75}[1.0]{-}x\varphi(x). ∎

Let bi∼bernoulli(p)b_{i}\sim\text{bernoulli}(p) and Zi∼N(0,1)Z_{i}\sim\mathcal{N}(0,1). Let Xi=bi\mathds1{Zi≥k}X_{i}=b_{i}\mathds{1}{\left\{Z_{i}\geq k\right\}} and Xˉ=1n∑i=1nXi\bar{X}=\frac{1}{n}\sum^{n}_{i=1}X_{i}. Then:

Let bi∼bern(p)b_{i}\sim\text{bern}(p) and Zi∼N(0,1)Z_{i}\sim\mathcal{N}(0,1). Let Xi=bi\mathds1{Zi≥k}ZiX_{i}=b_{i}\mathds{1}{\left\{Z_{i}\geq k\right\}}Z_{i} and Xˉ=1n∑i=1nXi\bar{X}=\frac{1}{n}\sum^{n}_{i=1}X_{i}. Then:

Therefore, Xˉ=pφ(k)±2nlog⁡n\bar{X}=p\varphi(k)\pm\sqrt{\frac{2}{n}}\log n w.p. at least 1−4n1-\frac{4}{n}. ∎

Let g:\mathdsR\textbackslash{0}→\mathdsRg:\mathds{R}\textbackslash\{0\}\to\mathds{R} be defined as g_{z}(x)=\frac{1}{x}\exp(\scalebox{0.75}[1.0]{-}\frac{z^{2}}{2x^{2}}). Then, (1) ∣z∣|z| and \scalebox{0.75}[1.0]{-}|z| are the global maximizer and minimizer respectively, and (2) gg monotonically increases from \scalebox{0.75}[1.0]{-}|z| to ∣z∣|z|.

Note that g^{\prime}_{z}(x)=\frac{1}{x^{2}}\exp(\scalebox{0.75}[1.0]{-}\frac{z^{2}}{2x^{2}})(\frac{z^{2}}{x^{2}}-1). Therefore, the critical points of gg are ∣z∣|z| and \scalebox{0.75}[1.0]{-}|z|. Let S={t:∣t∣≥∣z∣,t∈\mathdsR/{0}}S=\{t:|t|\geq|z|,t\in\mathds{R}/\{0\}\}. Note that gz′(x)<0g^{\prime}_{z}(x)<0 for all x∈Sx\in S and gz′(x)>0g^{\prime}_{z}(x)>0 for all x∈Scx\in S^{c}. Therefore, (1) and (2) hold. ∎

Let X∼N(0,σ2Id)X\sim\mathcal{N}(0,\sigma^{2}I_{d}) denote a dd-dimensional gaussian vector. Then, from , w.p. greater than 1−δ1-\delta: