An Optimization and Generalization Analysis for Max-Pooling Networks
Alon Brutzkus, Amir Globerson
Introduction
Convolutional neural networks (CNNs) have achieved remarkable performance in various computer vision tasks (Krizhevsky et al. 2012; Xu et al. 2015; Taigman et al. 2014). Such networks typically combine convolution and max-pooling layers, and can thus be used for detecting complex patterns in the input. In practice, CNNs typically have more parameters than needed to achieve zero train error (i.e., are overparameterized). Despite the potential problem of non-convexity in optimization and overfitting because of overparameterization, training these models with gradient based methods leads to solutions with low test error. Furthermore, overparameterized CNNs significantly outperform fully connected networks (FCNs) on classifying image data (Malach & Shalev-Shwartz 2020a). Thus, a key question immediately arises:
Why do overparameterized CNNs generalize well on image data and outperform FCNs?
To the best of our knowledge, this question remains largely unanswered. We note that the question contains two significant challenges: the first is to show that minimization of the non-convex training loss leads to high training accuracy (where non-convexity is a result of both max-pooling and ReLU activations), and the other is that over-fitting is avoided despite over-parameterization. The latter challenge is known as the question of inductive bias of gradient descent (Zhang et al. 2017), and understanding it is a key goal of deep learning theory.
In this work, we provide the first results which address the above question. We theoretically analyze learning a simplified pattern recognition task with overparameterized CNNs and overparameterized FCNs. We consider a CNN with a convolution layer, max pooling and fully connected layer and compare it to a one-hidden layer non-linear FCN. Figure 1 shows an example of our setup. We summarize our contributions as follows:
Expressive Power of CNNs with max-pooling: We prove a novel VC dimension lower bound in our setting which is exponential in , the filter dimension of the CNN. This result implies that there exists ERM algorithms which have sample complexity which is exponential in in our setting.
Optimization and Generalization for learning CNNs with max-pooling: We analyze learning overparamaterized CNNs with a layerwise gradient descent optimizer. We show that the algorithm converges to zero training loss and the learning has a sample complexity of . This is despite the above VC result, which shows that general ERM optimizers can potentially overfit. In our proof, we analyze the dynamics of training the first layer. We show that it induces a representation in the last layer which is separable with large margin and thus implies a good generalization guarantee.
Generalization of FCNs: We apply recent results of Brutzkus et al. 2018 which show a generalization bound for overparameterized FC networks that is independent of the network size. We prove that in our setting, their bound can be at best for , and can thus be much larger than the sample complexity we derive for the CNN.
Empirical Evaluation: We empirically validate our theoretical results. We show that CNNs generalize well and significantly outperform FCNs in our setting as predicted by our theory. We empirically confirm that this holds also for several extensions of our setup.
Our results make a significant headway on the challenging problem of understanding why overparameterized CNNs can generalize better than overparameterized FCNs on image classification tasks. In particular, to the best of our knowledge, we provide the first optimization and generalization results for overparameterized CNNs with max pooling.
Related Work
Two recent works have provided theoretical support that that CNNs outperform FCNs. Li et al. 2020 consider a simplified image classification task and prove a sample complexity gap between FCNs and single channel CNNs. Malach & Shalev-Shwartz 2020a prove that for simplified pattern detection tasks, there is a computational separation between overparameterized CNNs and FCNs. Their generalization bound for overparameterized CNNs depends on the number of channels of the CNN. Therefore, both works do not show that over-parameterized CNNs are reslient to over-fitting, which is the main focus of our work.
Several recent works have studied the generalization properties of overparameterized CNNs. Some of these propose generalization bounds that depend on the number of channels (Long & Sedghi 2020; Jiang et al. 2019). Others provide guarantees for CNNs with constraints on the weights (Zhou & Feng 2018; Li et al. 2018). Convergence of gradient descent to KKT points of the max-margin problem is shown in (Lyu & Li 2020) and (Nacson et al. 2019) for homogeneous models. However, their results do not provide generalization guarantees in our setting. Gunasekar et al. 2018 study the inductive bias of linear CNNs.
Yu et al. 2019 study a pattern classification problem similar to ours. However, their analysis their sample complexity guarantee depends on the network size, and thus does not explain why large CNNs do not overfit. Other works have studied learning under certain ground truth distributions. For example, Brutzkus & Globerson 2019 study a simple extension of the XOR problem, showing that overparameterized CNNs generalize better than smaller CNNs. Single-channel CNNs are analyzed in (Du et al. 2018b; Du et al. 2018a; Brutzkus & Globerson 2017; Du et al. 2018c). CNNs were analyzed via the NTK approximation (Li et al. 2019; Arora et al. 2019c). Our analysis does not assume the NTK approximation. For example, we require a mild overparameterization in our results which does not depend on the number of samples, in contrast to NTK analyses. Furthermore, our results hold for sufficiently small initialization, which is not the regime of NTK analysis.
Other works study the inductive bias of gradient descent on fully connected linear or non-linear networks (Ji & Telgarsky 2019a; Arora et al. 2019a; Wei et al. 2019; Brutzkus et al. 2018; Dziugaite & Roy 2017; Allen-Zhu et al. 2019; Chizat & Bach 2020). Fully connected networks were also analyzed via the NTK approximation (Du et al. 2019; Du et al. 2018d; Arora et al. 2019b; Fiat et al. 2019). Kushilevitz & Roth 1996; Shvaytser 1990 study the learnability of visual patterns distribution. However, our focus is on learnability using a specific algorithm and architecture: gradient descent trained on overparameterized CNNs.
Preliminaries
We consider a learning problem that captures a key property of visual classification. Many visual classes are characterized by the existence of certain patterns. For example an 8 will typically contain an x like pattern somewhere in the image. Here we consider an abstraction of this behavior where images consist of a set of patterns. Furthermore, each class is characterized by a pattern that appear exclusively in it. We define this formally below.
Next, we define how labeled points are generated. In our setting we consider three types of patterns: positive, negative and spurious. We will refer to the pattern as positive, the pattern as negative and the patterns as spurious. We let .
Given , a vector is sampled as follows. Randomly sample an index for placing the positive pattern, and set . Then, for each such that , randomly choose and set .
Given , do the same as , using instead of .
Fig. 1(a) shows an example of the above distribution .
CNN Architecture:
CNN Training Algorithm:
For the analysis of learning CNNs, we will consider a layerwise optimization algorithm which performs gradient updates layer-by-layer, starting from the first layer. Layerwise optimization algorithms are used in practice and have been shown to achieve performance that is comparable to end-to-end methods, e.g., on ImageNet (Belilovsky et al. 2019). Furthermore, the assumption on layerwise optimization has been used previously for theoretically analyzing neural networks (Malach & Shalev-Shwartz 2020b).
The layerwise optimization algorithm for learning CNNs is given in Figure 2. The reason we optimize over two losses is technical: we need a fresh IID sample () in the second layer optimization for the generalization analysis (see Section 5).
We define to be the th row of . For , and , define , i.e., corresponds to the pattern in that maximally activates . If , define . Otherwise, define . Notice that the following equality holds:
We note that it is necessary to make assumptions regarding the data distribution because the general case is intractable for optimization (because it includes neural net learning as a special case). We believe that our data generating distribution does reflect core aspects of pattern detection problems. Furthermore, the analysis of overparameterized max pooling networks has not been performed for any task, and analysis of simplified tasks has been shown to be fruitful for understanding CNNs (Li et al. 2020; Malach & Shalev-Shwartz 2020a). Additionally, non-overlapping filters are used in practice, and multiple theoretical works have analyzed CNNs with non-overlapping filters due to their tractability (Sharir & Shashua 2018). Finally, we note that in Section 7 we show that our analysis is in line with the performance of CNNs and FCNs in more complex tasks.
VC Dimension Bound
Thus far we described a data generating distribution and a neural architecture. We now ask how expressive is this neural architecture. Because of the pooling layer, it may seem that the network has limited capacity, even for an unbounded number of channels. However, as we show next the capacity in terms of VC dimension is in fact exponential in in this case. This in turn means that the network can separate datasets of size up to exponential in , and can thus potentially overfit badly. As we show in later sections, overfitting is avoided when learning using gradient descent.
We begin by recalling the definition of the VC dimension.
Let be a hypothesis class of functions from to . For any non-negative in integer , we define:
If , we say that shatters the set . The VC dimension of , denoted by, , is the size of the largest shattered set, or equivalently, the largest such that .
Assume that and , then .
We will construct a set of size that can be shattered. We note that the inclusion will hold for any , . For a given let be its th entry. For any such , define a point such that for any , . Furthermore, arbitrarily choose or and define .
Now, assume that each point has label . We will show that there is a network such that for all . For each , define and , where is the unique solution of the following linear system with equations. For each the system has the following equation:
where for any , is defined such that for all . There is a unique solution because the corresponding matrix of the linear system is the difference between an all 1’s matrix and the identity matrix. By the Sherman-Morrison formula (Sherman & Morrison 1950), this matrix is invertible, where in the formula the outer product rank-1 matrix is the all 1’s matrix and the invertible matrix is minus the identity matrix.
Set to be the matrix with rows followed by rows . Let be the a vector of dimension such that .
Then, for with parameters and any :
by the definition of , the orthogonality of the patterns , and Eq. 5. We have shown that any labeling can be achieved, and hence the set is shattered, completing the proof. ∎
The main limitation of the VC analysis is that it does not take into account the specific implementation of the ERM algorithm (Zhou & Feng 2018). In the next section, we will show a more fine-grained analysis which is specific to the layerwise optimization algorithm, and can thus benefit from the specific inductive bias of this algorithm. As a result, we will obtain a significantly better generalization guarantee.
Generalization Analysis of Gradient Descent
In this section we analyze the optimization and generalization performance of the layer-wise gradient descent algorithm for training overparameterized CNNs (Eq. 1). We will show that it converges to zero training loss and its sample complexity is . This is in contrast to the result of the previous section which shows a VC dimension lower bound which is exponential in , and therefore there are other ERM algorithms that can result in arbitrarily bad test error.
The first part of the theorem is an optimization result stating that the will converge to zero loss. We note that this is despite the non-convexity of the loss . The second part of the theorem states that the learned classifier will have a test error of order . Thus, the sample complexity is linear in . This is contrast to the VC dimension bound which is exponential in .
Before proving the theorem, we make several remarks on the result. First, for simplicity we present asymptotic results for . We can provide convergence rates that depend linearly on by changing the second layer optimization hyper-parameters (initialization and step size) and use recent results of Ji & Telgarsky 2019c. See Section A for details. Second, note that is a mild overparameterization condition, compared to other results which require to depend on the number of samples (Du et al. 2018d; Ji & Telgarsky 2019b).
We will prove the theorem in three parts. We defer the proofs of technical lemmas to the supplementary. We first outline the main ideas of the proof. In the first part we will prove a property of the initialization of the first layer. We show that at initialization there are sufficiently many “lucky” filters in the following sense. Either the pattern in that maximally activates them is and , or the maximum activating pattern is and . In essence, these filters are “good” detectors because they detect the discriminative patterns, with the right sign of .
In the second part we analyze the dynamics of the filters in the first layer. We will show that the “lucky” filters continue to detect the discriminative patterns and their projection on either or becomes larger in each iteration. In contrast, we upper bound the norm of the filters that are ”non-lucky”. Thus, after training the first layer, creates a new representation of the data in the second layer with the following properties: there are sufficiently many discriminative features with sufficiently large absolute values, and the remaining features have a bounded absolute value.
In the third part, we analyze the optimization of the second layer on the new representation. Using the properties of the representation, proved in the second part, we show that this representation induces a distribution on the samples which is linearly separable. Furthermore, it can be classified with margin 1 by a linear classifier of low norm. Then, we apply a result of Soudry et al. 2018, which implies that training the second layer, which is equivalent to logistic regression on the new representation, converges to a low norm solution with zero training loss. Finally, we apply a norm-based generalization bound (Shalev-Shwartz & Ben-David 2014) to obtain the sample complexity guarantee.
Part 1: Properties of the Initialization:
Define the sets , and the following sets:
The sets and correspond to the set of “lucky” filters. We prove a lower and upper bound on the size of these sets.
The following lemma shows the dynamics of the “lucky” neurons that detect the positive patterns.
For all and all the following holds:
.
For all , it holds that .
Furthermore, for all , .
The lemma shows that the projection of the filter on grows significantly, while the projection on other remains small. Finally, it shows that for any positive point in , the pattern which maximally activates the filter is . Thus, the filter is correctly detecting the positive pattern. The proof is technical and shows that the properties above hold by induction on . It is given in Section C.
By the symmetry of our setting we get by Lemma 5.3 a similar result for the “lucky” neurons that detect negative patterns.
With probability at least , for all and all the following holds:
.
For all , it holds that .
Furthermore, for all , .
Finally, we provide a simple bound on the output of all neurons (including the ”non-lucky” ones).
For all , , and sampled from , it holds that .
We conclude the proof of the theorem by analyzing the optimization of the second layer. Here we sketch the analysis and defer the details to Section E.
Using the results of the first layer dynamics, we show that is linearly separable and can be separated with margin 1 by a classifier with . Then, we use recent results on logistic regression (Soudry et al. 2018), to show that by optimizing the second layer, will converge to a low norm solution with zero training loss. Finally, we apply norm-based generalization bounds (Shalev-Shwartz & Ben-David 2014). Since for all , , we obtain a sample complexity guarantee for of order . ∎
Comparison with FCNs
In the previous section we showed that overparameterized CNNs have good sample complexity for learning the pattern distributions in Section 3. How do overparameterized fully connected networks compare with CNNs in our setting? To address this question, we apply recent results of Brutzkus et al. 2018. They provide generalization guarantees for one-hidden layer overparameterized fully connected networks on linearly separable data. We will show that their bound for FC networks can be for any . In contrast, Theorem 5.1 shows a generalization bound for CNNs which is linear in . We note that to fully demonstrate a gap between the methods we also need a lower bound on the FCN for the distribution , and we leave this for future work. Nonetheless, we show empirically, that these generalization bounds predict the performance gap between CNNs and FCNs in our setting.
Assume that is linearly separable with margin 1 by a classifier , i.e., for all , . In Brutzkus et al. 2018 they consider the following fully connected network:
They show that SGD converges to a zero training error solution with sample complexity of , where is the maximum norm of the data, . In our setting it holds that (because each point consists of patterns, each of norm ). Importantly, this bound is independent of the network size .
We note that the bound also holds for the hard-margin linear SVM (Shalev-Shwartz & Ben-David 2014). Therefore, our following conclusions hold for this algorithm as well. In the next section we show experiments that compare CNNs, FCNs and SVMs in our setting and corroborate our findings.
The generalization bound of holds for any which separates with margin 1. Thus, the best bound can be achieved with that has the lowest norm and separates the data with margin 1. Next we show that the lowest norm is at least .
Then .
Assume by contradiction that . Then, there exists such that . Define a positive point such that and for . Similarly, define a negative point such that and for . Then it holds that:
By subtracting Eq. 11 from Eq. 10 we get:
but since , we have by Eq. 12 , which is a contradiction. ∎
Proposition 6.1 implies that the best possible bound of Brutzkus et al. 2018 for FC networks, or margin bound for linear SVM is in our setting. Thus for , the bounds for FC networks and linear SVM are . In contrast, Theorem 5.1 shows a generalization guarantee for CNNs of for any . This gap suggests that CNNs should significantly outperform FCNs and linear SVM in our setting. Next, we provide empirical evidence for this.
Experiments
In this section we provide empirical evaluation of learning with our pooling architecture and compare it to several other models. As baselines we consider:
ConvPool: Our convolution and max-pooling model in Eq. 1. We verified that layer-wise training performs very similarly to standard training, and thus we report results on standard training with Adam (Kingma & Ba 2014) in what follows.
MLP: A standard fully connected neural network with one hidden layer. The network receives the complete as input (with all patterns). We use a number of hidden neurons that results in the same number of parameters as ConvPool.
SVM: A hard-margin linear SVM with as input. This will return zero training errors only when the data is linearly separable. This is the case for our distribution , but no longer the case when we add noise to the patterns (see below).
All experiments used a test set of size , and were repeated times with mean and std reported on figures.
Next, we consider the effect of the number of patterns on performance. As shown in Proposition 6.1, the norm of the max-margin linear classifier is lower bounded by . Thus, increasing is expected to result in worse performance for MLP and SVM by the results in the previous section. In Figure 4, we vary the number of patterns, and indeed observe that performances of MLP and SVM deteriorate while that of ConvPool is only mildly affected (we used the same parameters as above and noise level ).
Finally, we evaluate on the MNIST data set. We create a detection problem as in Fig. 5 where the discriminative patterns are the digits three and five and the spurious patterns are all other digits. Each input image contains four patterns (i.e., four digits). We used a relatively small number of patterns to make the problem not linearly separable for moderate sample sizes. We trained a 3 layer convolutional network as in Eq. 1 with 500 channels. Results in Fig. 3(c) again show excellent performance of the pooling model compared to the baselines.
Discussion
In this paper we presented the first analysis of a convolutional max-pooling architecture in terms of optimization and generalization under over-parameterization. Our analysis is for a natural setting of a detection problem where certain patterns “identify” the class and the others are irrelevant. Our analysis predicts a significant performance gap between CNNs and FCNs, which we observe in experiments.
While our analysis is the first step towards understanding pattern detection architectures, many open problems remain. The first is extending the pattern structure from orthogonal patterns to more general distributions. For example, we can consider the discriminative pattern to be a combination of patterns across the image (e.g., the class of the image is positive only if certain multiple patterns appear in the image). Second, it would be interesting to extend the convolution so that there are overlaps between filters (although this is known to generate local optima even for simpler settings (Brutzkus & Globerson 2017)). Finally, a challenging extension is to a multi-layer architecture with repeated application of pooling.
Acknowledgements
This research is supported by the European Research Council (ERC) under the European Unions Horizon 2020 research and innovation programme (grant ERC HOLI 819080). AB is supported by the Google Doctoral Fellowship in Machine Learning.
References
Appendix A Convergence Rates for Theorem 5.1
In Ji & Telgarsky 2019c, Theorem 4.2, they show the following for logistic regression initialized at zero and a certain learning rate schedule. The margin of the learned classifier is where is the max-margin after iterations. hides a dependency on . They show this for normalized points with norm 1. In our case (see the proof of Theorem 5.1), the max margin after normalizing the points to have norm 1, is . Thus, under their assumptions, after iterations we converge to a solution whose margin is a -multiplicative approximation of the max margin. Therefore, we obtain for this solution, up to a constant, the same generalization guarantees as the max margin classifier (which we provide in the theorem).
Appendix B Proof of Lemma 5.2
where in the last inequality we used the assumption on . Since and for , we get that with probability at least , and . By the symmetry of our problem and definitions of the sets , , , , we similarly get that with probability at least , . Applying the union bound concludes the proof.
Appendix C Proof of Lemma 5.3
We first prove the following two auxiliary lemmas.
For all and all , .
First we notice that for all , . This follows since for all and all , (recall that for ).
Therefore, for all and , . ∎
For all and .
By Lemma C.1 we have for all :
where the last inequality follows by the assumption on . ∎
Lemma 5.3 follows by the following lemma.
With probability at least , for all and all the following holds:
.
For all , it holds that .
We will prove the claim for . We prove the two claims by induction on . In the proof by induction we also show a third claim that: for all , .
For the proof, we condition on the event:
This holds with probability at least by applying Hoeffding’s inequality and a union bound (over positive and negative samples).
For , we have by definition for all , . The second claim holds by the definition of the initialization. The third claim follows by the definition of .
Assume the three claims above hold for . We will prove them for .
Proof of Claim 1. By the gradient update in the first layer, the following holds for :
By Lemma C.2 we have for all . Therefore, for all :
By the induction hypothesis, we have for and all that . Therefore we have:
For all , we have for that depends on . Therefore:
By the facts above we complete the proof of the first claim:
where the last inequality follows from the induction hypothesis.
Proof of Claim 2. Since for all , we have for all , :
By the facts (1) for all and it holds that and (2) for all , we have:
where the right inequality follows by the induction hypothesis.
Proof of Claim 3. Since we conclude by Eq. C and Eq. 23 that for all , . ∎
Appendix D Proof of Lemma 5.5
By Lemma C.1, for all and , . Therefore, for all and sampled from , .
Appendix E Proof of Part 3 of Theorem 5.1
Our goal is to show that is linearly separable and can be separated with a classifier of relatively low norm. Then, we will use recent results on logistic regression, which show that GD converges to low norm solutions. Therefore, by optimizing the second layer, will converge to a low norm solution. Finally, we will apply norm-based generalization bounds to obtain a generalization guarantee for .
where the inequality follows by Lemma 5.2, Lemma 5.3 and Corollary 5.4. By symmetry, we have for all .
Therefore, by this theorem we are guaranteed that:
Specifically, gradient descent converges to zero training loss, i.e., .
By optimality of and Lemma 5.2 we have . Furthermore, by Lemma 5.5. Therefore, we have . Thus, by a standard margin generalization bound (e.g. Theorem 26.13 in Shalev-Shwartz & Ben-David 2014 or Bartlett & Mendelson 2002) we have with probability at least :
where hides an additive term which depends on .