Gradient Starvation: A Learning Proclivity in Neural Networks

Mohammad Pezeshki, Sékou-Oumar Kaba, Yoshua Bengio, Aaron Courville, Doina Precup, Guillaume Lajoie

Introduction

In 1904, a horse named Hans attracted worldwide attention due to the belief that it was capable of doing arithmetic calculations (Pfungst,, 1911). Its trainer would ask Hans a question, and Hans would reply by tapping on the ground with its hoof. However, it was later revealed that the horse was only noticing subtle but distinctive signals in its trainer’s unconscious behavior, unbeknown to him, and not actually performing arithmetic. An analogous phenomenon has been noticed when training neural networks (e.g. Ribeiro et al.,, 2016; Zhao et al.,, 2017; Jo and Bengio,, 2017; Heinze-Deml and Meinshausen,, 2017; Belinkov and Bisk,, 2017; Baker et al.,, 2018; Gururangan et al.,, 2018; Jacobsen et al.,, 2018; Zech et al.,, 2018; Niven and Kao,, 2019; Ilyas et al.,, 2019; Brendel and Bethge,, 2019; Lapuschkin et al.,, 2019; Oakden-Rayner et al.,, 2020). In many cases, state-of-the-art neural networks appear to focus on low-level superficial correlations, rather than more abstract and robustly informative features of interest (Beery et al.,, 2018; Rosenfeld et al.,, 2018; Hendrycks and Dietterich,, 2019; McCoy et al.,, 2019; Geirhos et al.,, 2020).

The rationale behind this phenomenon is well known by practitioners: given strongly-correlated and fast-to-learn features in training data, gradient descent is biased towards learning them first. However, the precise conditions leading to such learning dynamics, and how one might intervene to control this feature imbalance are not entirely understood. Recent work aims at identifying the reasons behind this phenomenon (Valle-Pérez et al.,, 2018; Nakkiran et al.,, 2019; Cao et al.,, 2019; Nar and Sastry,, 2019; Jacobsen et al.,, 2018; Niven and Kao,, 2019; Wang et al.,, 2019; Shah et al.,, 2020; Rahaman et al.,, 2019; Xu et al., 2019b, ; Hermann and Lampinen,, 2020; Parascandolo et al.,, 2020; Ahuja et al., 2020b, ), while complementary work quantifies resulting shortcomings, including poor generalization to out-of-distribution (OOD) test data, reliance upon spurious correlations, and lack of robustness (Geirhos et al.,, 2020; McCoy et al.,, 2019; Oakden-Rayner et al.,, 2020; Hendrycks and Gimpel,, 2016; Lee et al.,, 2018; Liang et al.,, 2017; Arjovsky et al.,, 2019). However most established work focuses on squared-error loss and its particularities, where results do not readily generalize to other objective forms. This is especially problematic since for several classification applications, cross-entropy is the loss function of choice, yielding very distinct learning dynamics. In this paper, we argue that Gradient Starvation, first coined in Combes et al., (2018), is a leading cause for this feature imbalance in neural networks trained with cross-entropy, and propose a simple approach to mitigate it.

We provide a theoretical framework to study the learning dynamics of linearized neural networks trained with cross-entropy loss in a dual space.

Using perturbation analysis, we formalize Gradient Starvation (GS) in view of the coupling between the dynamics of orthogonal directions in the feature space (Thm. 2).

We leverage our theory to introduce Spectral Decoupling (SD) (Eq. 17) and prove this simple regularizer helps to decouple learning dynamics, mitigating GS.

We support our findings with extensive empirical results on a variety of classification and adversarial attack tasks. All code and experiment details available at GitHub repository.

In the rest of the paper, we first present a simple example to outline the consequences of GS. We then present our theoretical results before outlining a number of numerical experiments. We close with a review of related work followed by a discussion.

Gradient Starvation: A simple example

Consider a 2-D classification task with a training set consisting of two classes, as shown in Figure 1. A two-layer ReLU network with 500 hidden units is trained with cross-entropy loss for two different arrangements of the training points. The difference between the two arrangements is that, in one setting, the data is not linearly separable, but a slight shift makes it linearly separable in the other setting. This small shift allows the network to achieve a negligible loss by only learning to discriminate along the horizontal axis, ignoring the other. This contrasts with the other case, where both features contribute to the learned classification boundary, which arguably matches the data structure better. We observe that training longer or using different regularizers, including weight decay (Krogh and Hertz,, 1992), dropout (Srivastava et al.,, 2014), batch normalization (Ioffe and Szegedy,, 2015), as well as changing the optimization algorithm to Adam (Kingma and Ba,, 2014) or changing the network architecture or the coordinate system, do not encourage the network to learn a curved decision boundary. (See App. B for more details.)

We argue that this occurs because cross-entropy loss leads to gradients “starved” of information from vertical features. Simply put, when one feature is learned faster than the others, the gradient contribution of examples containing that feature is diminished (i.e., they are correctly processed based on that feature alone). This results in a lack of sufficient gradient signal, and hence prevents any remaining features from being learned. This simple mechanism has potential consequences, which we outline below.

In the example above, even in the right plot, the training loss is nearly zero, and the network is very confident in its predictions. However, the decision boundary is located very close to the data points. This could lead to adversarial vulnerability as well as lack of robustness when generalizing to out-of-distribution data.

GS could also result in neural networks that are invariant to task-relevant changes in the input. In the example above, it is possible to obtain a data point with low probability under the data distribution, but that would still be classified with high confidence.

One might argue that according to Occam’s razor, a simpler decision boundary should generalize better. In fact, if both training and test sets share the same dominant feature (in this example, the feature along the horizontal axis), GS naturally prevents the learning of less dominant features that could otherwise result in overfitting. Therefore, depending on our assumptions on the training and test distributions, GS could also act as an implicit regularizer. We provide further discussion on the implicit regularization aspect of GS in Section 5.

Theoretical Results

In this section, we study the learning dynamics of neural networks trained with cross-entropy loss. Particularly, we seek to decompose the learning dynamics along orthogonal directions in the feature space of neural networks, to provide a formal definition of GS, and to derive a simple regularization method to mitigate it. For analytical tractability, we make three key assumptions: (1) we study deep networks in the Neural Tangent Kernel (NTK) regime, (2) we treat a binary classification task, (3) we decompose the interaction between two features. In Section 4, we demonstrate our results hold beyond these simplifying assumptions, for a wide range of practical settings. All derivation details can be found in SM C.

In the wide-width regime, the NTRF changes very little during training Lee et al., (2019), and the output of the neural network can be approximated by a first order Taylor expansion around the initialization parameters \mbox{\boldmath\theta}_{0}. Setting \mbox{\boldmath\Phi}_{0}\equiv\mbox{\boldmath\Phi}\left(\mathbf{X},\mbox{\boldmath\theta}_{0}\right) and then, without loss of generality, centering parameters and the output coordinates to their value at the initialization (θ0\theta_{0} and y^0\hat{\mathbf{y}}_{0}), we get

Dominant directions in the feature space as well as the parameter space are given by principal components of the NTRF matrix \mbox{\boldmath\Phi}_{0}, which are the same as those of the NTK Gram matrix (Yang and Salman,, 2019). We therefore introduce the following definition.

Consider the singular value decomposition (SVD) of the matrix \mathbf{Y}\mbox{\boldmath\Phi}_{0}=\mathbf{U}\mathbf{S}\mathbf{V}^{T}, where Y=diag(y)\mathbf{Y}=\text{diag}\left(\mathbf{y}\right). The jjth feature is given by (VT)j.(\mathbf{V}^{T})_{j.}. The strength of jjth feature is represented by sj=(S)jjs_{j}=(\mathbf{S})_{jj}. Also, (U).j(\mathbf{U})_{.j} contains the weights of this feature in all examples. A neural network’s response to a feature jj is given by zjz_{j} where,

In Eq. 4, the response to feature jj is the sum of the responses to every example in (Yy^\mathbf{Y}\hat{\mathbf{y}}) multiplied by the weight of the feature in that example (UT\mathbf{U}^{T}). For example, if all elements of (U).j(\mathbf{U})_{.j} are positive, it indicates a perfect correlation between this feature and class labels. We are now equipped to formally define GS.

Recall the the model prescribed by Eq. 3. Let zj∗z_{j}^{*} denote the model’s response to feature jj at training optimum θ∗\theta^{*}Training optimum refers to the solution to \nabla_{\mbox{\boldmath\theta}}\mathcal{L}(\mbox{\boldmath\theta})=0.. Feature ii starves the gradient for feature jj if dzj∗/d(si2)<0{dz^{*}_{j}}/{d(s_{i}^{2})}<0.

This definition of GS implies that an increase in the strength of feature ii has a detrimental effect on the learning of feature jj. We now derive conditions for which learning dynamics of system 3 suffer from GS.

2 Training Dynamics

We consider the widely used ridge-regularized cross-entropy loss function,

where 1\mathbf{1} is a vector of size nn with all its elements equal to 1. This vector form simply represents a summation over all the elements of the vector it is multiplied to. λ∈[0,∞)\lambda\in[0,\infty) denotes the weight decay coefficient.

Direct minimization of this loss function using the gradient descent obeys coupled dynamics and is difficult to treat directly (Combes et al.,, 2018). To overcome this problem, we call on a variational approach that leverages the Legendre transformation of the loss function. This allows tractable dynamics that can directly incorporate rates of learning in different feature directions. Following (Jaakkola and Haussler,, 1999), we note the following inequality,

where H(\mbox{\boldmath\alpha})=-\left[\mbox{\boldmath\alpha}\log\mbox{\boldmath\alpha}+(1-\mbox{\boldmath\alpha})\log\left(1-\mbox{\boldmath\alpha}\right)\right] is Shannon’s binary entropy function, \mbox{\boldmath\alpha}\in(0,1)^{n} is a variational parameter defined for each training example, and ⊙\odot denotes the element-wise vector product. Crucially, the equality holds when the maximum of r.h.s. w.r.t α\alpha is achieved at \mbox{\boldmath\alpha}^{*}=\frac{\partial\mathcal{L}}{\partial(\mathbf{Y}\hat{\mathbf{y}})^{T}}, which leads to the following optimization problem,

where the order of min and max can be swapped (see Lemma 3 of Jaakkola and Haussler, (1999)). Since the neural network’s output is approximated by a linear function of θ\theta, the minimization can be performed analytically with an critical value \mbox{\boldmath\theta}^{*T}=\frac{1}{\lambda}\mbox{\boldmath\alpha}\mathbf{Y}\mbox{\boldmath\Phi}_{0}, given by a weighted sum of the training examples. This results in the following maximization problem on the dual variable, i.e., min⁡θL(θ)\operatornamewithlimits{min}_{\theta}\mathcal{L}\left(\theta\right) is equivalent to,

By applying continuous-time gradient ascent on this optimization problem, we derive an autonomous differential equation for the evolution of α\alpha, which can be written in terms of features (see Definition 1),

where η\eta is the learning rate (see SM C.1 for more details). For this dynamical system, we see that the logarithm term acts as barriers that keep αi∈(0,1)\alpha_{i}\in(0,1). The other term depends on the matrix US2UT\mathbf{U}\mathbf{S}^{2}\mathbf{U}^{T}, which is positive definite, and thus pushes the system towards the origin and therefore drives learning.

When λ≪sk2\lambda\ll s_{k}^{2}, where kk is an index over the singular values, the linear term dominates Eq. 9, and the fixed point is drawn closer towards the origin. Approximating dynamics with a first order Taylor expansion around the origin of the second term in Eq. 9, we get

with stability given by the following theorem with proof in SM C.

Any fixed points of the system in Eq. 10 is attractive in the domain αi∈(0,1)\alpha_{i}\in(0,1).

At the fixed point \mbox{\boldmath\alpha}^{*}, corresponding to the optimum of Eq. 8, the feature response of the neural network is given by,

See App. A for further discussions on the distinction between "feature space" and "parameter space". Below, we study how the strength of one feature could impact the response of the network to another feature which leads to GS.

3 Gradient Starvation Regime

In general, we do not expect to find an analytical solution for the dynamics of the coupled non-linear dynamical system of Eq. 10. However, there are at least two cases where a decoupled form for the dynamics allows to find an exact solution. We first introduce these cases and then study their perturbation to outline general lessons.

If the matrix of singular values S2\mathbf{S}^{2} is proportional to the identity: This is the case where all the features have the same strength s2s^{2}. The fixed points are then given by,

where W\mathcal{W} is the Lambert W function.

If the matrix U\mathbf{U} is a permutation matrix: This is the case in which each feature is associated with a single example only. The fixed points are then given by,

To study a minimal case of starvation, we consider a variation of case 2 with the following assumption which implies that each feature is not associated with a single example anymore.

Assume U\mathbf{U} is a perturbed identity matrix (a special case of a permutation matrix) in which the off-diagonal elements are proportional to a small parameter δ>0\delta>0. Then, the fixed point of the dynamical system in Eq. 10 can be approximated by,

where A=λ−1U(S2+λI)UT\mathbf{A}=\lambda^{-1}\mathbf{U}(\mathbf{S^{2}}+\lambda\mathbf{I})\mathbf{U}^{T} and \mbox{\boldmath\alpha}_{0}^{*} is the fixed point of the uncoupled system with δ=0\delta=0.

For sake of ease of derivations, we consider the two dimensional case where,

which is equivalent to a UU matrix with two blocks of features with no intra-block coupling and δ\delta amount of inter-block coupling.

Consider a neural network in the linear regime, trained under cross-entropy loss for a binary classification task. With definition 1, assuming coupling between features 1 and 2 as in Eq. 15 and s12>s22s_{1}^{2}>s_{2}^{2}, we have,

While Thm. 2 outlines conditions for GS in two dimensional feature space, we note that the same rationale naturally extends to higher dimensions, where GS is defined pairwise over feature directions. For a classification task, Thm. 2 indicates that gradient starvation occurs when the data admits different feature strengths, and coupled learning dynamics. GS is thus naturally expected with cross-entropy loss. Its detrimental effects however (as outlined in Sect. 2) arise in settings with large discrepancies between feature strengths, along with network connectivity that couples these features’ directions. This phenomenon readily extends to multi-class settings, and we validate this case with experiments in Sect. 4. Next, we introduce a simple regularizer that encourages feature decoupling, thus mitigating GS by insulating strong features from weaker ones.

4 Spectral Decoupling

By tracing back the equations of the previous section, one may realize that the term UTS2UU^{T}S^{2}U in Eq. 9 is not diagonal in the general case, and consequently introduces coupling between αi\alpha_{i}’s and hence, between the features ziz_{i}’s. We would like to discourage solutions that couple features in this way. To that end, we introduce a simple regularizer: Spectral Decoupling (SD). SD replaces the general L2 weight decay term in Eq. 5 with an L2 penalty exclusively on the network’s logits, yielding

Repeating the same analysis steps taken above, but with SD instead of general L2 penalty, the critical value for \mbox{\boldmath\theta}^{*} becomes \mbox{\boldmath\theta}^{*}=\frac{1}{\lambda}\mbox{\boldmath\alpha}\mbox{\boldmathY}\mbox{\boldmath\Phi}_{0}V\mathbf{S}^{-2}V^{T}. This new expression for \mbox{\boldmath\theta}^{*} results in the following modification of Eq. 9,

where as earlier, log⁡\log and division are taken element-wise on the coordinates of α\alpha.

Note that in contrast to Eq. 9 the matrix multiplication involving UU and SS in Eq. 18 cancels out, leaving αi\alpha_{i} independent of other αj≠i\alpha_{j\neq i}’s. We point out this is true for any initial coupling, without simplifying assumptions. Thus, a simple penalty on output weights promotes decoupled dynamics across the dual parameter αi\alpha_{i}’s, which track learning dynamics of feature responses (see Eq. 7). Together with Thm. 2, Eq. 18 suggests SD should mitigate GS and promote balanced learning dynamics across features. We now verify this in numerical experiments. For further intuition, we provide a simple experiment, summarized in Fig. 6, where directly visualizes the primal vs. the dual dynamics as well as the effect of the proposed spectral decoupling method.

Experiments

The experiments presented here are designed to outline the presence of GS and its consequences, as well as the efficacy of our proposed regularization method to alleviate them. Consequently, we highlight that achieving state-of-the-art results is not the objective. For more details including the scheme for hyper-parameter tuning, see App. B.

Recall the simple 2-D classification task between red and blue data points in Fig. 1. Fig. 1 (c) demonstrates the learned decision boundary when SD is used. SD leads to learning a curved decision boundary with a larger margin in the input space. See App. B for additional details and experiments.

2 CIFAR classification and adversarial robustness

To study the classification margin in deeper networks, we conduct a classification experiment on CIFAR-10, CIFAR-100, and CIFAR-2 (cats vs dogs of CIFAR-10) (Krizhevsky et al.,, 2009) using a convolutional network with ReLU non-linearity. Unlike linear models, the margin to a non-linear decision boundary cannot be computed analytically. Therefore, following the approach in Nar et al., (2019), we use "the norm of input-disturbance required to cross the decision boundary" as a proxy for the margin. The disturbance on the input is computed by projected gradient descent (PGD) (Rauber et al.,, 2017), a well-known adversarial attack.

Table 1 includes the results for IID (original test set) and OOD (perturbed test set by ϵPGD=0.05\epsilon_{\text{PGD}}=0.05). Fig. 2 shows the percentange of mis-classifications as the norm of disturbance is increased for the Cifar-2 dataset. This plot can be interpreted as the cumulative distribution function (CDF) of the margin and hence a lower curve reads as a more robust network with a larger margin. This experiment suggests that when trained with vanilla cross-entropy, even slight disturbances in the input deteriorates the network’s classification accuracy. That is while spectral decoupling (SD) improves the margin considerably. Importantly, this improvement in robustness does not seem to compromise the noise-free test performance. It should also be highlighted that SD does not explicitly aim at maximizing the margin and the observed improvement is in fact a by-product of decoupled learning of latent features. See Section 5 for a discussion on why cross-entropy results in a poor margin while being considered a max-margin classifier in the literature (Soudry et al.,, 2018).

3 Colored MNIST with color bias

We conduct experiments on the Colored MNIST Dataset, proposed in Arjovsky et al., (2019). The task is to predict binary labels y=−1y=-1 for digits 0 to 4 and y=+1y=+1 for digits 5 to 9. A color channel (red, green) is artificially added to each example to deliberately impose a spurious correlation between the color and the label. The task has three environments:

Training env. 1: Color is correlated with the labels with 0.9 probability.

Training env. 2: Color is correlated with the labels with 0.8 probability.

Testing env.: Color is correlated with the labels with 0.1 probability (0.9 reversely correlated).

Because of the opposite correlation between the color and the label in the test set, only learning to classify based on color would be disastrous at testing. For this reason, Empirical Risk Minimization (ERM) performs very poorly on the test set (23.7 % accuracy) as shown in Tab. 2.

Invariant Risk Minimization (IRM) (Arjovsky et al.,, 2019) on the other hand, performs well on the test set with (67.1 % accuracy). However, IRM requires access to multiple (two in this case) separate training environments with varying amount of spurious correlations. IRM uses the variance between environments as a signal for learning to be “invariant” to spurious correlations. Risk Extrapolation (REx) (Krueger et al.,, 2020) is a related training method that encourages learning invariant representations. Similar to IRM, it requires access to multiple training environments in order to quantify the concept of “invariance”.

SD achieves an accuracy of 68.4 %. Its performance is remarkable because unlike IRM and REx, SD does not require access to multiple environments and yet performs well when trained on a single environment (in this case the aggregation of both of the training environments).

A natural question that arises is “How does SD learn to ignore the color feature without having access to multiple environments?” The short answer is that it does not! In fact, we argue that SD learns the color feature but it also learns other predictive features, i.e., the digit shape features. At test time, the predictions resulting from the shape features prevail over the color feature. To validate this hypothesis, we study a trained model with each of these methods (ERM, IRM, SD) on four variants of the test environment: 1) grayscale-digits: No color channel is provided and the network should rely on shape features only. 2) colored-digits: Both color and digit are provided however the color is negatively correlated (opposite of the training set) with the label. 3) grayscale-blank: All images are grayscale and blank and hence do not provide any information. 4) colored-blank: Digit features are removed and only the color feature is kept, also with reverse label compared to training. Fig. 3 summarizes the results. For more discussions see SM B.

As a final remark, we should highlight that, by design, this task assumes access to the test environment for hyperparameter tuning for all the reported methods. This is not a valid assumption in general, and hence the results should be only interpreted as a probe that shows that SD could provide an important level of control over what features are learned.

4 CelebA with gender bias

The CelebA dataset (Liu et al.,, 2015) contains 162k celebrity faces with binary attributes associated with each image. Following the setup of (Sagawa et al.,, 2019), the task is to classify images with respect to their hair color into two classes of blond or dark hair. However, the Gender ∈\in {Male, Female} is spuriously correlated with the HairColor ∈\in {Blond, Dark} in the training data. The rarest group which is blond males represents only 0.85 % of the training data (1387 out of 162k examples). We train a ResNet-50 model (He et al.,, 2016) on this task. Tab. 5 summarizes the results and compares the performance of several methods. A model with vanilla cross-entropy (ERM) appears to generalize well on average but fails to generalize to the rarest group (blond males) which can be considered as “weakly" out-of-distribution (OOD). Our proposed SD improves the performance more than twofold. It should be highlighted that for this task, we use a variant of SD in which, λ2∣∣y^−γ∣∣22\frac{\lambda}{2}||\hat{y}-\gamma||^{2}_{2} is added to the original cross-entropy loss. The hyper-parameters λ\lambda and γ\gamma are tuned separately for each class (a total of four hyper-parameters). This variant of SD does provably decouple the dynamics too but appears to perform better than the original SD in Eq. 17 in this task.

Other proposed methods presented in Tab. 5 also show significant improvements on the performance of the worst group accuracy. The recently proposed “Learning from failure” (LfF) (Nam et al.,, 2020) achieves comparable results to SD, but it requires simultaneous training of two networks. Group DRO (Sagawa et al.,, 2019) is another successful method for this task. However, unlike SD, Group DRO requires explicit information about the spuriously correlated attributes. In most practical tasks, information about the spurious correlations is not provided and, dependence on the spurious correlation goes unrecognized.Recall that it took 3 years for the psychologist, Oskar Pfungst, to realize that Clever Hans was not capable of doing any arithmetic.

Related Work and Discussion

Here, we discuss the related work. Due to space constraints, further discussions are in App. A.

Several works including Saxe et al., (2013, 2019); Advani and Saxe, (2017); Lampinen and Ganguli, (2018) investigate the dynamics of deep linear networks trained with squared-error loss. Different decompositions of the learning process for neural networks have been used: Rahaman et al., (2019); Xu et al., 2019a ; Ronen et al., (2019); Xu et al., 2019b study the learning in the Fourier domain and show that low-frequency functions are learned earlier than high-frequency ones. Saxe et al., (2013); Advani et al., (2020); Gidel et al., (2019) provide closed-form equations for the dynamics of linear networks in terms of the principal components of the input covariance matrix. More recently, with the introduction of neural tangent kernel (NTK) (Jacot et al.,, 2018; Lee et al.,, 2019), a new line of research is to study the convergence properties of gradient descent (e.g. Allen-Zhu et al., 2019b, ; Mei and Montanari,, 2019; Chizat and Bach,, 2018; Du et al., 2018b, ; Allen-Zhu et al., 2019a, ; Huang and Yau,, 2019; Goldt et al.,, 2019; Zou et al.,, 2020; Arora et al., 2019b, ; Vempala and Wilmes,, 2019). Among them, Arora et al., 2019c ; Yang and Salman, (2019); Bietti and Mairal, (2019); Cao et al., (2019) decompose the learning process along the principal components of the NTK. The message in these works is that the training process can be decomposed into independent learning dynamics along the orthogonal directions.

Most of the studies mentioned above focus on the particular squared-error loss. For a linearized network, the squared-error loss results in linear learning dynamics, which often admit an analytical solution. However, the de-facto loss function for many of the practical applications of neural networks is the cross-entropy. Using the cross-entropy as the loss function leads to significantly more complicated and non-linear dynamics, even for a linear neural network. In this work, our focus was the cross-entropy loss.

In the context of robustness in neural networks, state-of-the-art neural networks appear to naturally focus on low-level superficial correlations rather than more abstract and robustly informative features of interest (e.g. Geirhos et al., (2020)). As we argue in this work, Gradient Starvation is likely an important factor contributing to this phenomenon and can result in adversarial vulnerability. There is a rich research literature on adversarial attacks and neural networks’ vulnerability (Szegedy et al.,, 2013; Goodfellow et al.,, 2014; Ilyas et al.,, 2019; Madry et al.,, 2017; Akhtar and Mian,, 2018; Ilyas et al.,, 2018). Interestingly, Nar and Sastry, (2019), Nar et al., (2019) and Jacobsen et al., (2018) draw a similar conclusion and argue that “an insufficiency of the cross-entropy loss” causes excessive invariances to predictive features. Perhaps Shah et al., (2020) is the closest to our work in which authors study the simplicity bias (SB) in stochastic gradient descent. They demonstrate that neural networks exhibit extreme bias that could lead to adversarial vulnerability.

Despite being highly-overparameterized, modern neural networks seem to generalize very well (Zhang et al.,, 2016). Modern neural networks generalize surprisingly well in numerous machine tasks. This is despite the fact that neural networks typically contain orders of magnitude more parameters than the number of examples in a training set and have sufficient capacity to fit a totally randomized dataset perfectly (Zhang et al.,, 2016). The widespread explanation is that the gradient descent has a form of implicit bias towards learning simpler functions that generalize better according to Occam’s razor. Our exposition of GS reinforces this explanation. In essence, when training and test data points are drawn from the same distribution, the top salient features are predictive in both sets. We conjecture that in such a scenario, by not learning the less salient features, GS naturally protects the network from overfitting.

The same phenomenon is referred to as implicit bias, implicit regularization, simplicity bias and spectral bias in several works (Rahaman et al.,, 2019; Neyshabur et al.,, 2014; Gunasekar et al.,, 2017; Neyshabur et al.,, 2017; Nakkiran et al.,, 2019; Ji and Telgarsky,, 2019; Soudry et al.,, 2018; Arora et al., 2019a, ; Arpit et al.,, 2017; Gunasekar et al.,, 2018; Poggio et al.,, 2017; Ma et al.,, 2018).

As an active line of research, numerous studies have provided different explanations for this phenomenon. For example, Nakkiran et al., (2019) justifies the implicit bias of neural networks by showing that stochastic gradient descent learns simpler functions first. Baratin et al., (2020); Oymak et al., (2019) suggests that a form of implicit regularization is induced by an alignment between NTK’s principal components and only a few task-relevant directions. Several other works such as Brutzkus et al., (2017); Gunasekar et al., (2018); Soudry et al., (2018); Chizat and Bach, (2018) recognize the convergence of gradient descent to maximum-margin solution as the essential factor for the generalizability of neural networks. It should be stressed that these work refer to the margin in the hidden space and not in the input space as pointed out in Jolicoeur-Martineau and Mitliagkas, (2019). Indeed, as observed in our experiments, the maximum-margin classifier in the hidden space can be achieved at the expense of a small margin in the input space.

The no free lunch theorem (Shalev-Shwartz and Ben-David,, 2014; Wolpert,, 1996) states that “learning is impossible without making assumptions about training and test distributions”. Perhaps, the most commonly used assumption of machine learning is the i.i.d. assumption (Vapnik and Vapnik,, 1998), which assumes that training and test data are identically distributed. However, in general, this assumption might not hold, and in many practical applications, there are predictive features in the training set that do not generalize to the test set. A natural question that arises is how to favor generalizable features over spurious features? The most common approaches include data augmentation, controlling the inductive biases, using regularizations, and more recently training using multiple environments.

Here, we would like to elaborate on an interesting thought experiment of Parascandolo et al., (2020): Suppose a neural network is provided with a chess book containing examples of chess games with the best movements indicated by a red arrow. The network can take two approaches: 1) learn how to play chess, or 2) learn just the red arrows. Either of these solutions results in zero training loss on the games in the book while only the former is generalizable to new games. With no external knowledge, the network typically learns the simpler solution.

Recent work aims to leverage the invariance principle across several environments to improve robust learning. This is akin to present several chess books to a network, each with markings indicating the best moves for different sets of games. In several studies (Arjovsky et al.,, 2019; Krueger et al.,, 2020; Parascandolo et al.,, 2020; Ahuja et al., 2020a, ), methods are developed to aggregate information from multiple training environments in a way that favors the generalizable / domain-agnostic / invariant solution. We argue that even with having access to only one training environment, there is useful information in the training set that fails to be discovered due to Gradient Starvation. The information on how to actually play chess is already available in any of the chess books. Still, as soon as the network learns the red arrows, the network has no incentive for further learning. Therefore, learning the red arrows is not an issue per se, but not learning to play chess is.

Here, we would like to remind the reader that GS can have both adverse and beneficial consequences. If the learned features are sufficient to generalize to the test data, gradient starvation can be viewed as an implicit regularizer. Otherwise, Gradient Starvation could have an unfavorable effect, which we observe empirically when some predictive features fail to be learned. A better understanding and control of Gradient Starvation and its impact on generalization offers promising avenues to address this issue with minimal assumptions. Indeed, our Spectral Decoupling method requires an assumption about feature imbalance but not to pinpoint them exactly, relying on modulated learning dynamics to achieve balance.

Modern neural networks are being deployed extensively in numerous machine learning tasks. Our models are used in critical applications such as autonomous driving, medical prediction, and even justice system where human lives are at stake. However, neural networks appear to base their predictions on superficial biases in the dataset. Unfortunately, biases in datasets could be neglected and pose negative impacts on our society. In fact, our Celeb-A experiment is an example of the existence of such a bias in the data. As shown in the paper, the gender-specific bias could lead to a superficial high performance and is indeed very hard to detect. Our analysis, although mostly on the theory side, could pave the path for researchers to build machine learning systems that are robust to biases and helps towards fairness in our predictions.

Conclusion

In this paper, we formalized Gradient Starvation (GS) as a phenomenon that emerges when training with cross-entropy loss in neural networks. By analyzing the dynamical system corresponding to the learning process in a dual space, we showed that GS could slow down the learning of certain features, even if they are present in the training set. We derived spectral decoupling (SD) regularization as a possible remedy to GS.

Acknowledgments and Disclosure of Funding

The authors are grateful to Samsung Electronics Co., Ldt., CIFAR, and IVADO for their funding and Calcul Québec and Compute Canada for providing us with the computing resources. We would further like to acknowledge the significance of discussions and supports from Reyhane Askari Hemmat and Faruk Ahmed. MP would like to thank Aristide Baratin, Kostiantyn Lapchevskyi, Seyed Mohammad Mehdi Ahmadpanah, Milad Aghajohari, Kartik Ahuja, Shagun Sodhani, and Emmanuel Bengio for their invaluable help.

References

Appendix A Further discussions

On Primal (parameter space) vs. Dual (feature space) dynamics: Although the cross-entropy loss is convex, it does not admit an analytical solution, even in a simple logistic regression . Importantly, it also does not have a finite solution when the data is linearly separable (which is the case in high dimensions ). As such, our study is concerned with characterizing the solutions that the training algorithm converges to. A dual optimization approach enables us to describe these solutions in terms of contributions of the training examples . While primal and dual dynamics are not guaranteed to match, the solution they converge to is guaranteed to match , and that is what our theory builds upon.

For further intuition, we provide a simple experiment in app C, directly visualizing the primal vs. the dual dynamics as well as the effect of the proposed spectral decoupling method.

The intuition behind Spectral Decoupling (SD): Consider a training datapoint xx in the middle of the training process. Intuitively, the model has two options for decreasing the loss of this example:

Get more confident on a feature that has been learned already by other examples. or,

SD, a simple L2 penalty on the output of the work, would favor (2) over (1). The reason is that (2) does not make the network over-confident on previously learned examples, while (1) results in over-confident predictions. Hence, SD encourages learning more features by penalizing confidence. Our principal novel contribution is to characterize this process formally and to theoretically and empirically demonstrate its effectiveness.

From another perspective, here we describe how one can arrive at Spectral Decoupling. From Thm. 2, we know that Gradient Starvation happens because of the coupling between features (equivalently alphas). We notice that in Eq. 9, if we get rid of S2S^{2}, then the alphas are decoupled. To get rid of S2S^{2} , one can see that instead of ||\mbox{\boldmath\theta}||^{2} as the regularizer, we should have ||SV^{T}\mbox{\boldmath\theta}||^{2}. Luckily, this is exactly equal to ∣∣y^2∣∣||\hat{y}^{2}||, since \hat{y}=\Phi\mbox{\boldmath\theta}=UV^{T}\mbox{\boldmath\theta}. We would like to highlight that ||SV^{T}\mbox{\boldmath\theta}||^{2} as the regularizer means that different directions are penalized according to their strength. It means that we suppress stronger directions more than others which would allow weaker directions to flourish.

Then why not use Squared-error loss for classification too? The biggest obstacle when using squared-error loss for classification is how to select the target. For example, in a cats vs. dogs classification task, not all cats have the same amount of "catty features". However, recent results favor using squared-error loss for classification and show that models trained with squared-error loss are more robust . We conjecture that the improved robustness can be attributed to a lack of gradient starvation.

On using NTK: Theoretical analysis of neural networks in their general form is challenging and generally intractable. Neural Tangent Kernel (NTK) has been an important milestone that has simplified theoretical analysis significantly and provides some mechanistic explanations that are applicable in practice. Inevitably, it imposes a set of restrictions; mainly, NTK is only accurate in the limit of large width. Therefore, the common practice is to provide the theoretical analysis in simplified settings and validate the results empirically in more general cases (see, e.g. ). In this work, we build on the same established practices: Our theories analytically study an NTK linearized network; and we further validate our findings on several standard neural networks. In fact, in all of our experiments, learning is done in the regular "rich" (non-NTK) regime, and we verify that our proposed method, as identified analytically, mitigates learning limitations.

Future Directions: This work takes a step towards understanding the reliance of neural networks upon spurious correlations and shortcuts in the dataset. We believe identifying this reliance in sensitive applications is among the next steps for future research directions. That would have a pronounced real-world impact as neural networks have started to be used in many critical applications. As a recent example, we would like to point to an article by researchers at Cambridge where they study more than 300 papers on detecting whether a patient has COVID or not given their CT Scans. According to the article, none of the papers were able to generalize from one hospital data to another since the models learn to latch on to hospital-specific features. An essential first step is to uncover such reliance and then to design methods such as our proposed spectral decoupling to mitigate the problem.

Appendix B Experimental Details

Here, we provide a simple experiment to study the difference between the primal and dual form dynamics. We also compare the learning dynamics in cases with and without Spectral Decoupling (SD).

Recall that primal dynamics arise from the following optimization,

while the dual dynamics are the result of another optimization,

Also recall that Spectral Decoupling suggests the following optimization,

We conduct experiments on a simple toy classification with two datapoints for which the matrix U\mathbf{U} of Eq. 15 is defined as, U=(0.8−0.60.60.8)\mathbf{U}=\begin{pmatrix}0.8&-0.6\\ 0.6&0.8\end{pmatrix}. The corresponding singular values S=[s1,s2=2]\mathbf{S}=[s_{1},s_{2}=2] where s1∈{2,3,4,5,6}s_{1}\in\{2,3,4,5,6\}. According to Eq. 13, when S=\mathbf{S}=, the dynamics decouple while in other cases starvation occurs. Fig. 6 shows the corresponding features of z1z_{1} and z2z_{2}. It is evident that by increasing the value of s1s_{1}, the value of z1∗z_{1}^{*} increases while z2∗z_{2}^{*} decreases (starves). Fig. 6 (left) also compares the difference between the primal and the dual dynamics. Note that although their dynamics are different, they both share the same fixed points. Fig. 6 (right) also shows that Spectral Decoupling (SD) indeed decouples the learning dynamics of z1z_{1} and z2z_{2} and hence increasing the corresponding singular value of one does not affect the other.

B.2 Two-Moon Classification: Comparison with other regularization methods

We experiment the Two-moon classification example of the main paper with different regularization techniques. The small margin between the two classes allows the network to achieve a negligible loss by only learning to discriminate along the horizontal axis. However, both axes are relevant for the data distribution, and the only reason why the second dimension is not picked up is the fact that the training data allows the learning to explain the labels with only one feature, overlooking the other. Fig. 7 reveals that common regularization strategies including Weight Decay, Dropout and Batch Normalization do not help achieving a larger margin classifier. Unless states otherwise, all the methods are trained with Full batch Gradient Descent with a learning rate of 1e−21e-2 and a momentum of 0.90.9 for 10k10k iterations.

B.3 CIFAR classification

We use a four-layer convolutional network with ReLU non-linearity following the exact setup of . Sweeping λ\lambda from 0 to its optimal value results in a smooth transition from green to orange. However, larger values of λ\lambda will hurt the IID test (zero perturbation) generalization. The value that we cross-validate on is the average of IID and OOD generalization performance.

B.4 Colored MNIST with color bias

For the Colored MNIST task, we aggregate all the examples from both training environments. Table. 3 reports the hyper-parameters used for each method.

As a final remark, we highlight that, by design, this task assumes access to the test environment for hyperparameter tuning for all the reported methods. This is not a valid assumption in general, and hence the results should be only interpreted as a probe that shows that SD could provide an important level of control over what features are learned.

The hyperparameter search has resulted in applying the SD at 450th step. We observe that 450th step is the step at which the traditional (in-distribution) overfitting occurs. This suggests that one might be able to tune hyperparameters without the need to monitor on the test set.

For all the experiments, we use PyTorch . We also use NNGeometry for computing NTK.

B.5 CelebA with gender bias: The experimental details

Figure 5 depicts the learning curves for this task with and without Spectral Decoupling. For the CelebA experiment, we follow the same setup as in and use their released code. We use Adam optimizer for the Spectral Decoupling experiments with a learning rate of 1e−41e-4 and a batch size of 128128. As mentioned in the main text, for this experiment, we use a different variant of Spectral Decoupling which also provably decouples the learning dynamics,

We applied a hyper-parameter search on λ\lambda and γ\gamma for each of the classes separately. Therefore, a total of four hyper-parameters are found. For class zero, λ0=0.088\lambda_{0}=0.088, γ0=0.44\gamma_{0}=0.44 and for class one, λ1=0.012\lambda_{1}=0.012, γ1=2.5\gamma_{1}=2.5 are found to result in the best worst-group performance.

During the experiments, we found that for the CelebA dataset, classes are imbalanced: 10875 examples for class 0 and 1925 examples for class 1; meaning a ratio of 5.65. That is why we decided to penalize examples of each class separately with different coefficients. We also found that penalizing the outputs’ distance to different values γ0\gamma_{0} and γ1\gamma_{1} helps the generalization. As stated in lines 842-844, the hyperparameter search results in the following values: 2.5 and 0.44.

B.6 Computational Resources

For the experiments and hyper-parameter search an approximate number of 800 GPU-hours has been used. GPUs used for the experiments are NVIDIA-V100 mostly on internal cluster and partly on public cloud clusters.

Appendix C Proofs of the Theories and Lemmas

Following , we derive the Legendre transformation of the Cross-Entropy (CE) loss function. Here, we reiterate this transformation as following,

For a variational parameter α∈\alpha\in, the following linear lower bound holds for the cross-entropy loss function,

in which ω:=yy^\omega:=y\hat{y} and H(α)H(\alpha) is the Shannon’s binary entropy. The equality holds for the critical value of α∗=−∇ωL\alpha^{*}=-\nabla_{\omega}\mathcal{L}, i.e., at the maximum of r.h.s. with respect to α\alpha.

The Legendre transformation converts a function L(ω)\mathcal{L}(\omega) to another function g(α)g(\alpha) of conjugate variables α\alpha, L(ω)→g(α)\mathcal{L}(\omega)\to g(\alpha). The idea is to find the expression of the tangent line to L(ω)\mathcal{L}(\omega) at ω0\omega_{0} which is the first-order Taylor expansion of L(ω)\mathcal{L}(\omega),

where t(ω,ω0)t(\omega,\omega_{0}) is the tangent line. According to the Legendre transformation, the function L(ω)\mathcal{L}(\omega) can be written as a function of the intercepts of tangent lines (where ω=0\omega=0). Varying ω0\omega_{0} along the xx-axis provides us with a general equation, representing the intercept as a function of ω\omega,

The cross-entropy loss function can be rewritten as a soft-plus function,

in which ω:=yy^\omega:=y\hat{y}. Letting α:=−∇ωL=σ(−ω)\alpha:=-\nabla_{\omega}\mathcal{L}=\sigma(-\omega) we have,

which allows us to re-write the expression for the intercepts as a function of α\alpha (denoted by g(α)g(\alpha)),

where H(α)H(\alpha) is the binary entropy function.

Now, since L\mathcal{L} is convex, a tangent line is always a lower bound and therefore at its maximum it touches the original function. Consequently, the original function can be recovered as follows,

Note that the lower bound in Eq. 19 is now a linear function of ω:=yy^\omega:=y\hat{y} but at the expense of an additional maximization over the variational parameter α\alpha. An illustration of the lower bound is depicted in Fig. 8. Also a comparison between the dual formulation of other common loss functions is provided in Table. 4.

Building on Eq. 65-71 of , which derives the Legendre transform of multi-class cross-entropy, one can update Eq. 6 of the main paper to

where H(α)H(\alpha) is the entropy function, C=#C=\#classes, and vectors of αc\alpha^{c} are defined for each class. Then Eq. 8 of the paper is then updated to,

With a change of variable αc:=δy−αc\alpha^{c}:=\delta_{y}-\alpha^{c}, the theory of SD should remain unchanged.

C.2 Eq. 8 Dual Dynamics

In Eq. 7, the order of min and max can be swapped as proved in Lemma 3 of , leading to,

The solution to the inner optimization is,

which its substitution into Eq. 7 results in Eq. 8.

C.3 Eq. 9

Simply taking the derivative of Eq. 8 will result in Eq. 9. When we introduce continuous gradient ascent, we must define a learning rate parameter. This term is conceptually equivalent to the learning rate in SGD, but in this continuous setting, it has no influence on the fixed point.

C.4 Eq. 10 Approximate Dynamics

Approximating dynamics of Eq. 9 with a first order Taylor expansion around the origin of the second term, we obtain

Starting from the exact dynamics at Eq. 9,

we perform a first-order Taylor approximation of the second term at α\alpha = 0\mathbf{0} :

C.5 Thm. 1 Attractive Fixed-Points

Any fixed points of the system in Eq. 10 is attractive in the domain αi∈(0,1)\alpha_{i}\in(0,1).

as the gradient function of the autonomous system Eq. 9.

We find the character of possible fixed points by linearization. We compute the jacobian of the gradient function evaluated at the fixed point.

The fixed point is an attractor if the jacobian is a negative-definite matrix. The first term is negative-definite matrix while the second term is negative semi-definite matrix. Since the sum of a negative matrix and negative-semi definite matrix is negative-definite, this completes the proof. ∎

C.6 Eq. 11 Feature Response at Fixed-Point

At the fixed point \mbox{\boldmath\alpha}^{*}, corresponding to the optimum of Eq. 8, the feature response of the neural network is given by,

The solution to the converged \mbox{\boldmath\theta}^{*} at the Fixed-Point \mbox{\boldmath\alpha}^{*} of Eq. 10 is,

which by substitution into Eq. 4, Eq. 11 is derived.

C.7 Eq. 12 Uncoupled Case 1

If the matrix of singular values S2\mathbf{S}^{2} is proportional to the identity, the fixed points of Eq. 10 are given by,

where W\mathcal{W} is the Lambert W function.

When S2=s2I\mathbf{S}^{2}=s^{2}\mathbf{I}, Eq. 10 becomes

Fixed points of this system are obtained when \dot{\mbox{\boldmath\alpha}}=0 :

With z\mathbf{z} given by Eq. 11, we have

C.8 Eq. 13 Uncoupled Case 2

If the matrix U\mathbf{U} is a permutation matrix, the fixed points of Eq. 10 are given by,

When U\mathbf{U} is a permutation matrix, it can be made an identity matrix with a meaningless reordering of the class labels . Without loss of generality, we therefore consider U=I\mathbf{U}=\mathbf{I}

Fixed points of this system are obtained when \dot{\mbox{\boldmath\alpha}}=0

With z\mathbf{z} given by Eq. 11, we have

C.9 Lemma 1 Perturbation Solution

Starting from the autonomous system Eq. 10 and assumption in 1, we have

Since the off-diagonal terms are of order δ\delta, we treat them as a perturbation. The unperturbed system has a solution \mbox{\boldmath\alpha}_{0} given by case 2

We can linearize the autonomous system Eq. 10 around the unperturbed solution to find,

We then apply the perturbation given by off-diagonal terms of A\mathbf{A} to obtain

where \text{diag}\left({\mbox{\boldmath\alpha}_{0}^{*}}^{-1}\right) is the diagonal matrix obtained from {\mbox{\boldmath\alpha}_{0}^{*}}^{-1} and where the inverse is applied element by element.

Solving for \dot{\mbox{\boldmath\alpha}}=0, we obtain the solution

C.10 Thm. 2 Gradient Starvation Regime

Consider a neural network in the linear regime, trained under cross-entropy loss for a binary classification task. With definition 1, assuming coupling between features 1 and 2 as in Eq. 15 and s12>s22s_{1}^{2}>s_{2}^{2}, we have,

From lemma 1, and with U\mathbf{U} given by Eq. 15, we find that the perturbatory solution for the fixed point is

We have found at Eq. 11 that the corresponding steady-state feature response is given by

In the perturbatory regime δ\delta is taken to be a small parameter. We therefore perform a first-order Taylor series expansion of z∗{\mathbf{z}}^{*} around δ=0\delta=0 to obtain

Taking the derivative of z2∗{z}_{2}^{*} with respect to s1{s}_{1}, we find

Knowing that the exponential of the WW Lambert function is a strictly increasing function and that s12>s22s_{1}^{2}>s_{2}^{2}, we find

C.11 Eq. 18 Spectral Decoupling

SD replaces the general L2 weight decay term in Eq. 5 with an L2 penalty exclusively on the network’s logits, yielding

Optimizing \mathcal{L}\left(\mbox{\boldmath\theta}\right) wrt to θ\theta results in the following optimum,

which by substitution into the loss function, the dynamics of gradient ascent leads to,

where log⁡\log and division are taken element-wise on the coordinates of α\alpha and hence dynamics of each αi\alpha_{i} is independent of other αj≠i\alpha_{j\neq i}.