Gradient Descent
Overview
$$ % Colors
% Coordinate vectors and matrices
% Common sets
% Abstract vector symbols
% Norms / absolute value
% Optional: dot product spacing (looks nicer in slides)
% Operators $$
Main idea
- Last time we wrote linear models from balance laws.
- Today we shift from solving equations to minimizing error.
- We start with fitting a line to data.
- Then we use that example to introduce gradient descent.
- This is one natural doorway toward machine learning.
Plan
- Write a prediction model
- Measure the error with a loss function
- Find the best parameters
- Compare exact least squares and gradient descent
- Connect the idea to machine learning
Companion file
- The MATLAB script for today is
scripts/gradient_descent_demo.m - It includes:
- noisy line-fitting data
- the exact least-squares solution
- a gradient descent loop
- loss and contour plots
Fitting
A simple prediction model
- Suppose we have data points \((x_i, y_i)\).
- We try to fit a line \[ \hat y = b + mx. \]
- The parameters are:
- \(b\): intercept
- \(m\): slope
- Different choices of \((b,m)\) give different prediction errors.
Matrix form
- Put the data into a matrix \[ X = \begin{bmatrix} 1 & x_1 \\ 1 & x_2 \\ \vdots & \vdots \\ 1 & x_n \end{bmatrix}, \qquad w = \begin{bmatrix} b \\ m \end{bmatrix}, \qquad y = \begin{bmatrix} y_1 \\ \vdots \\ y_n \end{bmatrix}. \]
- Then our predictions are simply \[ Xw. \]
Loss function
- We want predictions close to the data.
- A standard choice is the mean squared error \[ J(w) = \frac{1}{2n}\|Xw-y\|^2. \]
- Small \(J(w)\) means the line fits the data well.
- So the fitting problem becomes: \[ \text{find } w \text{ that minimizes } J(w). \]
Exact least-squares answer
For this linear model, MATLAB can solve the least-squares problem directly:
w_star = X \ y;That gives the best-fitting line.
This is the same least-squares story you have already seen in linear algebra.
But it is not the only way to think about the problem.
Gradient
Why use an iterative method?
- For a small line-fitting problem, the exact solve is easy.
- But in larger models:
- the data set may be huge
- the model may have many parameters
- the model may no longer be linear
- Then an iterative method becomes very attractive.
- The basic idea is: repeatedly move in a direction that makes the loss smaller.
Gradient of the loss
- For \[ J(w)=\frac{1}{2n}\|Xw-y\|^2, \] the gradient is \[ \nabla J(w)=\frac{1}{n}X^T(Xw-y). \]
- The gradient tells us the direction of steepest increase.
- So to make the loss go down, we move in the opposite direction.
Gradient descent update
\[ w_{k+1} = w_k - \alpha \nabla J(w_k) \]
- \(w_k\): current parameter guess
- \(\nabla J(w_k)\): slope of the loss at that point
- \(\alpha\): learning rate or step size
- Repeat this update many times and the parameters improve
MATLAB code for the gradient
J = @(w) (1/(2*n)) * norm(X*w - y)^2;
gradJ = @(w) (1/n) * X' * (X*w - y);J(w)measures how well the current line fits.gradJ(w)tells us how to change the intercept and slope.
MATLAB code for gradient descent
alpha = 0.4;
num_steps = 25;
w = [-1; -2];
for k = 1:num_steps
w = w - alpha * gradJ(w);
end- Start from a rough guess.
- Each step moves downhill.
- If
alphais too small, progress is slow. - If
alphais too large, the method can oscillate or diverge.
Visuals
What we can plot
- The data and the fitted line
- The loss value versus iteration
- The path of gradient descent in parameter space
- These pictures make the method much easier to understand
Typical contour picture
- For line fitting, the loss depends on just two parameters:
- intercept
- slope
- So we can draw contour lines of \(J(b,m)\).
- Gradient descent walks downhill across those contours toward the minimum.
Exact solve vs gradient descent
| Method | Good when | Main idea |
|---|---|---|
X \ y |
small or moderate linear least-squares problem | direct exact least-squares solve |
| Gradient descent | large models or iterative optimization | improve the parameters step by step |
- For linear regression, both approaches can work.
- Gradient descent becomes more important as the models get bigger.
Toward ML
Why this points toward machine learning
- Machine learning often follows this pattern:
- choose a model with parameters
- define a loss
- use data to reduce the loss
- Linear regression is one of the simplest examples.
- Neural networks use the same broad idea, just with many more parameters.
Fast answers
- Why move opposite the gradient?
- Because the gradient points uphill, so the negative gradient points downhill.
- Why not always use the exact least-squares solve?
- Exact formulas are great for small linear problems, but gradient methods scale better to large or nonlinear models.
- What does the learning rate do?
- It controls the step size.
- What if the learning rate is too big?
- The method can overshoot and fail to settle down.
- Is this already machine learning?
- Yes. Linear regression trained from data is one basic machine-learning model.
Summary
- Least squares turns fitting into a minimization problem.
- The gradient tells us how the loss changes.
- Gradient descent improves the parameters step by step.
- This is a core computational idea behind modern machine learning.
- Linear algebra is still the language underneath the whole story.