What Is Gradient Descent and How Does It Optimize Neural Networks?

Gradient descent is an iterative optimization algorithm that minimizes a neural network's loss function by repeatedly updating model parameters in the direction opposite to the gradient of the loss, effectively descending the loss landscape to find optimal weights.

According to the HenryNdubuaku/maths-cs-ai-compendium, gradient descent serves as the fundamental engine that powers training in virtually all modern deep learning systems. This comprehensive educational resource maps the theoretical foundations and practical implementations of gradient-based optimization across its Gradient Machine Learning and Optimisation chapters.

How Gradient Descent Works: The Three-Stage Process

The optimization procedure operates through a cyclical three-stage pipeline that gradually reduces prediction error. As implemented in the compendium's machine learning chapter (chapter 06 - machine learning/02. gradient machine learning.md), each iteration follows a precise mathematical protocol.

Forward Pass and Loss Computation

During the forward pass, the network processes a batch of inputs to compute predictions and evaluates a loss function (such as binary-cross-entropy for classification or mean-squared-error for regression). This loss quantifies the discrepancy between current predictions and target values, measuring how far the current parameters deviate from optimal behavior.

The compendium emphasizes that this stage establishes the objective landscape—the high-dimensional surface that gradient descent will navigate.

Backward Pass: Computing Gradients via Back-Propagation

The backward pass employs back-propagation to compute the gradient of the loss with respect to every weight and bias in the network. Using the chain rule of calculus, this process yields the vector ∇ℒ that points in the direction of increasing loss.

According to the source material, this gradient computation is what enables gradient descent to determine which parameters require adjustment and in what direction.

Parameter Update: The Core Update Rule

With gradients computed, parameters are updated by moving a fraction η (the learning rate) opposite to the gradient direction:

w = w - lr * grad

Mathematically, this follows the update rule:

$$\mathbf{w} \leftarrow \mathbf{w} - \eta ,\nabla!\mathcal{L}(\mathbf{w})$$

Repeating these steps across many mini-batches gradually reduces the loss, guiding the network toward a (typically local) optimum. The compendium notes that this simple update rule scales effectively to billions of parameters when accelerated by modern hardware.

From Batch to Stochastic: Mini-Batch Gradient Descent

While batch gradient descent computes gradients using the entire dataset, this approach is computationally prohibitive for large-scale problems. Instead, the compendium advocates for mini-batch stochastic gradient descent (SGD), which estimates the gradient using small random subsets of data.

This method provides a noisy but fast approximation that still converges when the learning rate is properly tuned. The noise introduced by mini-batch sampling often helps escape shallow local minima, a characteristic detailed in chapter 03 - calculus/05. optimisation.md.

Advanced Optimizers: Momentum and Adam

To improve convergence speed and stability, the compendium describes several modifications to the basic update rule:

  • Momentum accumulates a velocity vector that smooths updates across iterations, allowing the optimizer to maintain consistent direction while dampening oscillations in high-curvature directions.
  • Adam (Adaptive Moment Estimation) maintains separate exponential moving averages of first- and second-order moments (mean and variance of gradients), adapting the learning rate per parameter. This makes Adam the default choice for most deep learning projects, as noted in the Gradient Machine Learning chapter.

Implementing Gradient Descent: Code Examples from the Compendium

The compendium provides practical implementations using JAX that demonstrate these concepts for both linear and logistic regression.

Linear Regression with Vanilla Gradient Descent

This example from chapter 06 - machine learning/02. gradient machine learning.md demonstrates fitting a linear model using the basic gradient descent update rule:

import jax
import jax.numpy as jnp
import matplotlib.pyplot as plt

# Synthetic data: y = 3x + 2 + noise

key = jax.random.PRNGKey(42)
n = 100
X = jax.random.uniform(key, (n, 1), minval=0, maxval=10)
y = 3 * X[:, 0] + 2 + jax.random.normal(key, (n,)) * 1.5

# Add bias column

X_b = jnp.column_stack([X, jnp.ones(n)])

# Gradient descent

w = jnp.zeros(2)          # [weight, bias]

lr = 0.005
losses = []
for step in range(500):
    pred = X_b @ w
    error = pred - y
    loss = jnp.mean(error ** 2)
    losses.append(float(loss))
    grad = (2 / n) * X_b.T @ error
    w = w - lr * grad

print(f"Learned weight={w[0]:.4f}, bias={w[1]:.4f}")

# Plot loss convergence

plt.semilogy(losses)
plt.title("GD Loss Convergence")
plt.xlabel("Step")
plt.ylabel("MSE")
plt.show()

Logistic Regression with SGD

The following implementation applies gradient descent to a non-linear classification problem using the sigmoid activation:

import jax
import jax.numpy as jnp
import matplotlib.pyplot as plt
from sklearn.datasets import make_moons

# Data

X, y = make_moons(n_samples=300, noise=0.2, random_state=42)
X, y = jnp.array(X), jnp.array(y, dtype=jnp.float32)

def sigmoid(z):
    return 1 / (1 + jnp.exp(-z))

# Add bias

X_b = jnp.column_stack([X, jnp.ones(len(X))])
w = jnp.zeros(3)
lr = 0.5
losses = []

for step in range(2000):
    z = X_b @ w
    pred = sigmoid(z)
    loss = -jnp.mean(y * jnp.log(pred + 1e-8) + (1 - y) * jnp.log(1 - pred + 1e-8))
    losses.append(float(loss))
    grad = X_b.T @ (pred - y) / len(y)
    w = w - lr * grad

# Decision boundary

xx, yy = jnp.meshgrid(jnp.linspace(-2, 3, 200), jnp.linspace(-1.5, 2, 200))
grid = jnp.column_stack([xx.ravel(), yy.ravel(), jnp.ones(xx.size)])
zz = sigmoid(grid @ w).reshape(xx.shape)

plt.figure(figsize=(8, 6))
plt.contourf(xx, yy, zz, levels=[0, 0.5, 1], alpha=0.3, colors=['#e74c3c', '#3498db'])
plt.contour(xx, yy, zz, levels=[0.5], colors='#9b59b6')
plt.scatter(X[y==0, 0], X[y==0, 1], c='#e74c3c', label='Class 0')
plt.scatter(X[y==1, 0], X[y==1, 1], c='#3498db', label='Class 1')
plt.title("Logistic Regression Decision Boundary")
plt.legend()
plt.show()

Visualizing Optimizer Trajectories: SGD vs Momentum vs Adam

The compendium includes a comparative visualization showing how different optimizers navigate an elongated loss landscape. This example implements vanilla SGD, Momentum, and Adam update rules:

import jax
import jax.numpy as jnp
import matplotlib.pyplot as plt

def loss_fn(w):
    return 0.5 * w[0]**2 + 10 * w[1]**2   # elongated bowl

grad = jax.grad(loss_fn)

def sgd(w0, lr=0.05, steps=80):
    w = w0.copy()
    path = [w.copy()]
    for _ in range(steps):
        w = w - lr * grad(w)
        path.append(w.copy())
    return jnp.stack(path)

def momentum(w0, lr=0.05, beta=0.9, steps=80):
    w, v = w0.copy(), jnp.zeros(2)
    path = [w.copy()]
    for _ in range(steps):
        g = grad(w)
        v = beta * v + (1 - beta) * g
        w = w - lr * v
        path.append(w.copy())
    return jnp.stack(path)

def adam(w0, lr=0.05, b1=0.9, b2=0.999, eps=1e-8, steps=80):
    w, m, v = w0.copy(), jnp.zeros(2), jnp.zeros(2)
    path = [w.copy()]
    for t in range(1, steps + 1):
        g = grad(w)
        m = b1 * m + (1 - b1) * g
        v = b2 * v + (1 - b2) * (g**2)
        m_hat = m / (1 - b1**t)
        v_hat = v / (1 - b2**t)
        w = w - lr * m_hat / (jnp.sqrt(v_hat) + eps)
        path.append(w.copy())
    return jnp.stack(path)

w0 = jnp.array([8.0, 3.0])
sgd_path = sgd(w0)
mom_path = momentum(w0)
adam_path = adam(w0)

# Plot

fig, ax = plt.subplots(figsize=(8, 6))
w1 = jnp.linspace(-10, 10, 100)
w2 = jnp.linspace(-4, 4, 100)
W1, W2 = jnp.meshgrid(w1, w2)
L = 0.5 * W1**2 + 10 * W2**2
ax.contour(W1, W2, L, levels=20, cmap='Greys', alpha=0.4)
ax.plot(sgd_path[:,0], sgd_path[:,1], 'o-', color='#3498db', label='SGD')
ax.plot(mom_path[:,0], mom_path[:,1], 'o-', color='#27ae60', label='Momentum')
ax.plot(adam_path[:,0], adam_path[:,1], 'o-', color='#e74c3c', label='Adam')
ax.plot(0, 0, 'k*', markersize=12, label='Minimum')
ax.set_xlabel('w₁')
ax.set_ylabel('w₂')
ax.set_title('Optimizer Trajectories')
ax.legend()
plt.show()

Mathematical Foundations in the Compendium

The maths-cs-ai-compendium places gradient descent within a broader mathematical framework. The Optimisation chapter (chapter 03 - calculus/05. optimisation.md) contextualizes gradient descent among first- and second-order methods, explaining convexity conditions and detailing when sophisticated algorithms like Newton's method or quasi-Newton methods become advantageous.

Similarly, the Gradient Machine Learning chapter provides detailed derivations showing how back-propagation extends the chain rule to deep networks, linking the abstract mathematics of gradients to concrete neural network implementations.

Summary

  • Gradient descent minimizes loss by iteratively updating parameters opposite to the gradient direction, following the update rule w ← w - η∇L(w).
  • The algorithm operates through three stages: forward pass (loss computation), backward pass (gradient calculation via back-propagation), and parameter update.
  • Mini-batch SGD computes gradients on small data subsets, balancing computational efficiency with convergence stability.
  • Advanced optimizers like Momentum and Adam modify the basic update rule to accelerate convergence and adapt learning rates per parameter.
  • The HenryNdubuaku/maths-cs-ai-compendium provides comprehensive mathematical derivations and JAX implementations in chapter 06 - machine learning/02. gradient machine learning.md and chapter 03 - calculus/05. optimisation.md.

Frequently Asked Questions

What is the difference between batch gradient descent and stochastic gradient descent?

Batch gradient descent computes the gradient using the entire training dataset, providing stable but computationally expensive updates. Stochastic gradient descent (SGD) approximates the gradient using single examples or small mini-batches, introducing noise that speeds up computation and helps escape local minima. According to the compendium, mini-batch SGD is the standard for training modern neural networks because it balances computational efficiency with statistical accuracy.

Why is the learning rate important in gradient descent?

The learning rate (η) controls the step size during parameter updates. If set too high, the optimizer may overshoot the minimum and diverge. If set too low, convergence becomes prohibitively slow. The compendium emphasizes that tuning this hyperparameter is crucial for stable training, particularly when using vanilla SGD.

How does back-propagation relate to gradient descent?

Back-propagation is the algorithm that computes gradients of the loss with respect to each parameter using the chain rule, while gradient descent uses those gradients to update the parameters. Back-propagation answers "which direction reduces loss?" and gradient descent answers "how far should we step?" Together, they enable efficient training of deep networks with millions of parameters.

When should I use Adam instead of vanilla SGD?

Adam is preferred when training deep networks with sparse gradients or noisy data, as it adapts learning rates per parameter and maintains moving averages of past gradients. According to the compendium, Adam is often the default choice for most deep learning projects due to its robustness to hyperparameter choices and faster initial convergence. However, vanilla SGD with momentum sometimes achieves better final generalization performance with proper tuning.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →