Introduction
Deep learning is black magic. We throw data into PyTorch and we get a model
that seemingly understands language, speech, and vision. However, one
fundamental technique for deep learning has remained the same. That is
gradient descent. The modern optimisers such as Adam and AdamW are refinements
of gradient descent. So, I believe that gradient descent is still a core
technique worth learning.
I've written a post on the relationship between differentiation and
optimisation (https://yasufumimoriya.blogspot.com/2026/04/differentiation-and-optimisation.html). In this post, there was a single parameter to optimise or to find its
slope.
Gradient descent is essentially finding slopes of multiple parameters at once.
This post focuses on how computation of gradient descent proceeds to update
parameters, taking a function with two parameters as an example.
Refresher: Differentiation and Optimisation
This section is a refresher of the relationship between differentiation and
optimisation.
Differentiation
Recall the parabolic function \( y = x^2 \). Its derivative is a function that
can calculate its slope at any point. \[ \frac{dy}{dx} = 2x \] This means that
when \( x = 1 \), its slope is \( 2 \). When \( x= 2 \), its slope is \( 4 \).
This process to find the derivative of a function is differentiation.
Machine learning: data and loss function
The goal of machine learning is to fit a model into the provided examples, so
it can make future predictions. If a model is \( y = wx \) and given examples
are \( (1, 1), (2, 2), (3, 3) \), we want the model weight \( w \) to be \( 1
\). This way, if the next input \( x = 5 \) arrives, the model can correctly
predict \( y = 5 \).
The figure above illustrates potential errors this model can make with a wrong
weight. When \( w = 2 \) and the model is \( y = 2x \), the model is not
perfect fit to the provided examples. When \( w = -2 \) and the model is \( y
= -2x \), the model is further away from a good fit.
We can quantify the "mistake" the model is making by defining a loss function.
Using Mean Squared Error (MSE), the loss for this model is defined below: \[
L(w) = \frac{1}{N} \sum^N_{i=1}(wx_i-y_i)^2 \] where \(N\) is the number of
training examples. This loss is actually a parabolic function with the loss
value \( L(w) \) on the vertical axis and the weight value \( w \) on the
horizontal axis. Visualisation of this parabolic function is below:
As you can see in the figure, the bottom of the parabola is where the slope is
\( 0 \). This means that the weight at that point is the best fit. This is how
machine learning optimisation works: we take the derivative of a loss function
through differentiation, and find the weight that makes the derivative produce
\( 0 \). In this example, the loss happens to be \( 0 \), but the most important
thing is to find the zero slope and not the zero loss.
Real models have more parameters to fit to complex data. Optimisation for such
models keeps track of slopes of all parameters at once, and that vector of
slopes is called a gradient. That's where the naming of gradient descent is
from. Let's look into the prerequisite of gradient descent: partial
differentiation.
Partial Differentiation
Partial differentiation is the process of taking a derivative with respect
to a single parameter, while keeping the other parameters constant. For
example, the function below has two parameters \(x\) and \(y\): \[ f(x, y)
= 3x^2 + 2y^3 + 2 \] The partial derivatives are calculated as follows. \[
\frac{\partial f}{\partial x} = 6x, \quad \frac{\partial f}{\partial
y} = 6y^2 \] So far so good. A single term can have more than one
variable: \[ f(x, y) = 3x^2y + 2y^3 \] In this case, we still take a
derivative with respect to a single parameter, but the other parameter
survives: \[ \frac{\partial f}{\partial x} = 6xy, \quad \frac{\partial
f}{\partial y} = 3x^2 + 6y^2 \] The second case is especially important
for machine learning optimisation.
Gradient Descent
Let's now define a model with two parameters: \[ y = wx + b \] This is a
simple linear model with a single weight \( w \) and a bias term \( b \).
Training examples are: \( (1, 2), (2, 3), (3, 4) \).
The new model with the bias term is more expressive, and it can fit the new
examples. The original model of \( w = 1 \) does not make any predictions
correct, neither nor \( w = 2 \).
The MSE of the new model can be defined as below: \[ L(w, b) = \frac{1}{N}
\sum^N_{i=1} [(wx_i + b) - y_i]^2 \] Firstly, take the derivative with
respect to \(w\). Let's focus on a single data point and define a variable
\(u\): \[ u = (wx_i+b)-y_i \] We apply the chain rule to the partial
derivative with respect to \(w\): \[ \begin{aligned} \frac{\partial
L}{\partial w} &= \frac{\partial L}{\partial u}\frac{\partial
u}{\partial w} \\ &= 2u \cdot x_i \\ &= 2x_i [(wx_i+b)-y_i]
\end{aligned} \] Summing over \(N\) examples gives the full
partial derivative: \[ \frac{\partial L}{\partial w} = \frac{2}{N}
\sum^N_{i=1} x_i [(wx_i+b)-y_i] \] The partial derivative with respect to
\(b\) is easier since \( \frac{\partial u}{\partial b} = 1 \): \[
\frac{\partial L}{\partial b} = \frac{2}{N} \sum^N_{i=1} [(wx_i+b)-y_i] \]
Parameter update
We have a formula to compute gradient with respect to \(w\) and \(b\).
What is left for us is to adjust \(w\) and \(b\) until both slopes reach
zero, \( \frac{\partial L}{\partial w} = 0 \) and \( \frac{\partial
L}{\partial b} = 0 \). As we saw earlier in the refresher, ideal values of
parameters \(w\) and \(b\) produce the zero slopes.
The parameter update rule is as follows: \[ w \leftarrow w - \eta
\frac{\partial L}{\partial w} \] \[ b \leftarrow b - \eta \frac{\partial
L}{\partial b} \] where \( \eta \) is the learning rate which regulates
how large each step of the parameter is.
Gradient descent by hand
Let's now update the parameter \( w \) by hand. The setup is as follows:
- The initial value of the parameter: \( w = 3 \)
-
The second parameter has the correct value for simplicity: \( b = 1 \)
- The learning rate is: \( \eta=0.05 \)
- Training examples are: \( (1, 2), (2, 3), (3, 4) \)
1st update
The gradient value with respect to \(w\):
\[ \begin{aligned} \frac{\partial L}{\partial w} &= \frac{2}{3}
\{[(3+1)-2] + 2[(3\times2+1)-3] + 3[(3\times3+1)-4]\} \\ &=
\frac{2}{3} (2 + 8 + 18) \\ &= 18.67 \end{aligned} \]
The update with \(w\) is:
\[ 3 - 0.05 \times 18.67 = 2.07 \] We can see that the \(w\) value is
closer to fitting examples.
The gradient value with respect to \(w\):
\[ \begin{aligned} \frac{\partial L}{\partial w} &= \frac{2}{3}
\{[(2.07+1)-2] + 2[(2.07\times2+1)-3] + 3[(2.07\times3+1)-4]\} \\ &=
\frac{2}{3} (1.07+4.28+9.63) \\ &= 9.99 \end{aligned} \]
The update with \(w\) is:
\[ 2.07 - 0.05 \times 9.99 = 1.57 \]
3rd update
The gradient value with respect to \(w\):
\[ \begin{aligned} \frac{\partial L}{\partial w} &= \frac{2}{3}
\{[(1.57+1)-2] + 2[(1.57\times2+1)-3] + 3[(1.57\times3+1)-4]\} \\
&= \frac{2}{3} (0.57+2.28+5.13) \\ &= 5.32 \end{aligned} \]
The update with \(w\) is:
\[ 1.57 - 0.05 \times 5.32 = 1.3 \] I'll stop the update process
here.
Notice that gradient descent didn't need the loss to run. In this
example, we knew that the \(w=1.0\) is the value which fits the
provided examples. In real training, however, we don't know the
correct parameter values in advance, and computing the loss value is
important to monitor whether the parameter update process is going
well.
Visualisation of the process
The figure below visualises the gradient descent process.
The initial value of \( w \) is \(3.0\). The parameter updates change
the value of \(w\) to \(2.07\), then to \(1.57\) and then to \(1.3\).
The slope values are shown in the legend box, and we can see the slope
decreasing at each step.
Update multiple parameters
In the previous example, the value of \(b\) was fixed at \(1\) for
simplicity. In practice, we need to update multiple parameters. The slope for \(w\) depends on the current value of \(b\) and vice versa. I will
not include hand calculation of both \(w\) and \(b\) values. Instead,
the figure below visualises the process of finding the optimal
parameters for the provided examples, starting from \(w=3\)
and \(b=2\).
Similar to the single parameter process, gradient descent updates the two parameters towards the optimal point, which is the bottom of this 3D surface.
Summary
This post summarised gradient descent. The refresher showed that finding the point where the slope is zero corresponds to the optimal value of a parameter. We then defined a model with two parameters \( y = wx + b \) and derived the partial derivatives for both \(w\) and \(b\). The example of hand-calculating the parameter \( w \) with setting \(b=1\) demonstrated its move closer towards the optimal value after each update. The 3D plot visualised the gradient descent process, updating the two parameters simultaneously.
Comments
Post a Comment