Theory of the Frequency Principle for General Deep Neural Networks
Tao Luo, Zheng Ma, Zhi-Qin John Xu, Yaoyu Zhang
Introduction
Deep learning has achieved great success as in many fields (LeCun et al., 2015), e.g., speech recognition (Amodei et al., 2016), object recognition (Eitel et al., 2015), natural language processing (Young et al., 2018) and computer game control (Mnih et al., 2015). It has also been adopted into algorithms to solve scientific computing problems (E et al., 2017; Khoo et al., 2017; He et al., 2018; Fan et al., 2018). In principle, the universal approximation theorem states that a commonly-used Deep Neural Network (DNN) of sufficiently large width can approximate any function to a desired precision (Cybenko, 1989). However, it remains a mystery that how a DNN finds a minimum corresponding to such an approximation through the gradient-based training process. To understand the learning behavior of DNNs for the approximation problem, recent works model the gradient flow of parameters in a two-layer ReLU neural networks by a partial differential equation (PDE) in the mean-field limit (Rotskoff & Vanden-Eijnden, 2018; Mei et al., 2018; Sirignano & Spiliopoulos, 2018). However, it is not clear whether this PDE approach, which describes a neural network of one hidden layer of infinite width, can be extended to general DNNs of multiple hidden layers and limited neuron number.
In this work, we take another approach that uses Fourier analysis to study the learning behavior of DNNs based on the phenomenon of Frequency Principle (F-Principle), i.e., a DNN tends to learn a target function from low to high frequencies during the training (Xu et al., 2018; Rahaman et al., 2018; Xu, 2018a, b; Xu et al., 2019; Zhang et al., 2019). Empirically, the F-Principle can be widely observed in general DNNs for both benchmark and synthetic data (Xu et al., 2018, 2019). Conceptually, it provides a qualitative explanation of the success and failure of DNNs (Xu et al., 2019). Based on the F-Principle, a series of works has been done. For example, it is used as an important phenomenon to pursue fundamentally different learning trajectories of meta-learning (Rabinowitz, 2019). It is also used as a tool to observe the performance of adaptive activation function (Jagtap & Karniadakis, 2019). Based on the F-Principle, a numerical algorithm is developed to accelerate the DNN fitting of high frequency functions by shifting high frequencies to lower ones (Cai et al., 2019). Theoretically, an effective model of linear F-Principle dynamics (Zhang et al., 2019), which accurately predicts the learning results of two-layer ReLU neural networks of large widths, leads to an apriori estimate of the generalization bound. In addition, a theorem is provided for the characterization of the initial training stage of a two-layer network (Xu et al., 2019). The same theoretical analysis in Xu et al. (2019) is also adopted in the analysis of DNNs with ReLU activation function (Rahaman et al., 2018) and a nonlinear collaborative scheme of loss functions for DNN training (Zhen et al., 2018). These subsequent works show the importance of the F-Principle. However, a theory of the F-Principle for general DNNs is still missing.
Following the same direction as in Xu et al. (2019), in this work, we propose a theoretical framework of Fourier analysis for the study of the training behavior of general DNNs in the following three stages: the initial stage, the intermediate stage, and the final stage. At all stages, we rigorously characterize the F-Principle by estimating some proper quantities. At the initial and final stages with the MSE loss (mean-squared error, also known as loss), we show that the change of MSE is dominated by low frequencies. Furthermore, in these two stages with general () loss, we show that the change of the DNN output is dominated by the low-frequency part. A key contribution of this work is on the intermediate stage — with loss, the difference of the MSE over a certain period, in which the MSE is reduced by half, is dominated by the low frequencies. In summary, we verify that the F-Principle is universal in the sense that our results not only work for DNNs of multiple layers with any commonly-used activation function, e.g., ReLU, sigmoid, and tanh, but also work for a general population density of data and for a general class of loss functions. The key insight unraveled by our analysis is that the regularity of DNN converts into the decay rate of a loss function in the frequency domain.
Preliminaries
We start with a brief introduction to DNNs and its training dynamics. Under very mild assumptions, we provide some regularity results which are crucial to the proof of the main theorems summarized in the next section.
Consider a DNN with -hidden layers and general activation functions. We regard the -dimensional input as the -th layer and the one-dimensional output as the -th layer. Let be the number of neurons in the -th layer. In particular, and .
The size of the network is the number of the parameters, i.e.,
To define the hypothesis functions in , we need some nonlinear functions which are known as activation functions:
We remark that for the most applications, the activation functions are chosen to be the same, i.e., , , .
For instance, if a one-hidden layer neural network is used, then and the hypothesis function can be written into the following form:
Thus the size of the network which is consistent with (4).
In the sequel, we will also refer to and as the hypothesis and target functions, respectively.
2 Loss Function and Training Dynamics
In this work, we investigate the training dynamics of parameters in DNNs with two cases of loss functions:
(i) The MSE loss function with population measure , i.e.,
In this case, the training dynamics of follows the gradient flow:
(ii) A general loss function with population measure , i.e.,
In the case of MSE loss function, we have
3 Assumptions
The requirements on , , , and are summarized here.
If an activation function is ReLU, then .
For the training dynamics (13) or (15), we suppose the parameters are bounded.
The bound depends on initial parameter .
In the case of MSE loss function, we will further take the following assumption.
The general loss function considered in this work satisfies the following assumption.
4 Regularity
We begin with the integrability of the hypothesis function. To achieve this, we use the “Japanese bracket” of :
The continuity of is neccesary because the composition of two Lebesgue measurable functions need not be Lebesgue measurable.
Suppose that the Assumptions 1 and 2 hold. Then
(c). Let and . Then and . Combining the inequalities in parts (a) and (b), we have
Suppose that the Assumptions 1, 2, and 4 hold. Then
Main Results
In this section, we first propose several quantitative characterization for the F-Principle. Main results are then summarized with numerical illustrations at the end of this section.
For the MSE loss function, a natural quantity to characterize the F-principle is the ratio of the loss function decrements caused by low frequencies and the total loss function decrements. To achieve this, we devide the MSE loss function into two parts, contributed by low and high frequencies, respectively, i.e.,
For a general loss function, the training dynamics leads to
However, there is an issue in this characterization: may not be monotonically decreasing and the denominator in (38) may be zero. To overcome this, a time averaging is required. Indeed, we investigate the following ratio where integrals are taken for both numerator and denominator in (38):
For the general loss function, we also propose another quantity to characterize the F-Principle:
2 Main Theorems
As we mentioned in the introduction, the training dynamics of a DNN has three stages: initial stage, intermediate stage, and final stage. For each stage, we provide a theorem to characterize the F-Principle.
We start with the F-Principle in the initial stage. Clearly, the constants in the estimates depend on the initial parameter and the time .
[F-Principle in the initial stage] ( loss function) Suppose that Assumptions 1, 2, 3, and 4 hold. We consider the training dynamics (13). Then for any and any satifying (if , we further require that ), there is a constant such that
[F-Principle in the intermediate stage] (general loss function) Suppose that Assumptions 1, 2, 3, and 5 hold. We consider the training dynamics (15). Then for any , there is a constant such that for any satisfying , we have
If non-degenerate global minimizers are achieved in the training dynamics, we can obtain global-in-time result which characterizing the training dynamics in the final stage. Here we give the definition for non-degenerate minimizers:
[F-Principle in the final stage] ( loss function) Suppose that Assumptions 1, 2, 3, and 4 hold. We consider the training dynamics (13). If the solution converges to a non-degenerate global minimizer , then for any , there is a constant such that
(general loss function) Suppose that Assumptions 1, 2, 3, and 5 hold. We consider the training dynamics (15). If the solution converges to a non-degenerate global minimizer , then for any , there is a constant such that
3 Discussion and Illustrations
To help the readers get some intuitions of the above theorems, we present a numerical example using the following target function
The training data are uniformly sampled from with sample size . The discrete Fourier transform of is shown in Fig. 1(a), in which we focus on the peak frequencies marked by black squares. First, we use the MSE as the training loss function.
Intermediate stage in Fig. 1 (c). The ratio of the change of the loss function in a certain period, , increases with for a fixed .
Secondly, we use the training loss as shown in Fig. 2. We obtain similar results.
Proof of Theorems
In this section, we focus on the initial stage of the training dynamics. The first result shows that the change of loss function concentrates on low frequencies.
In general, may depend on . In the next section, we will provide a similar result in some situation where does not depend on .
The dynamics for the loss function contributed by high frequency reads as:
The dynamics for the total loss function is
Note that for all . Therefore
By Assumption 3, and
In the situation of Theorem 1 for loss function, we have that for sufficiently large
For sufficiently large , the dynamics of is dissipative because
Next we prove the case of general loss function.
On the one hand, we estimate the numerator by studying the dynamics for :
Taking square and integrating both sides on leads to the upper bound on the numerator
On the other hand, note the dynamics for the hypothesis function
and the dynamics for the total loss function
where we used the Cauchy–Schwarz inequality in the last step. Combining Eqs. (56) and (59), we obtain
Again, by Assumption 3, and
2 F-Principle: Intermediate Stage (Theorem 2)
In this section, we prove the key theorem for the intermediate stage. This theorem then implies several useful corollaries.
The numerator can be controlled as follows
where in the second-to-last step we used the Cauchy–Schwarz inequality and the Plancherel theorem, and in the last step we used the following
By Assumption 3, and
By the assumption that , we have
If , then we choose such that . We have
If the condition is replaced by for any , the estimates in Theorem 2 and the following corollaries still hold.
Under the same assumptions in Theorem 2, for any , there is a constant such that for any satisfying and for all , we have
Similar to the proof of Theorem 2, we have the upper bound for the numerator
where the last inequality is due to the same reason as Theorem 2. ∎
Under the same assumptions in Theorem 2, if the solution converges to a non-degenerate global minimizer , then for any , the above upper bound can be improved to the following: there is a constant such that for any , we have
We skip the proof since this corollary can be obtained directly from Theorem 3.
3 F-Principle: Final Stage (Theorem 3)
In this section, we prove the F-Principle in final stage of the training dynamics.
By Assumption 3, and
where we used the assumption that the minimizer is non-degenerate with the Hessian . ∎
Now we finish the proof for general loss function.
Since , we have