Maxout: Learning the Activation Function

Instead of picking a fixed activation, let each unit take the maximum of several learned linear functions.

Deep Learning- Fundamentals to Advanced Concepts

Every activation we have seen is a fixed curve chosen in advance. Maxout (Goodfellow and colleagues, 2013) lets the network learn its activation function.

Definition

A maxout unit with kk pieces computes kk different linear functions of its input and outputs the largest:

h(x)=max⁡j=1,…,k(wjTx+bj)h(x) = \max_{j = 1,\dots,k}\bigl(w_j^{T}x + b_j\bigr)

Each piece has its own weights and bias, all learned by backpropagation. The unit is a normal neuron with kk weighted sums instead of one, followed by a max. The gradient flows through whichever piece is currently the largest.

It includes ReLU and much more

The maximum of straight lines is always a convex, piecewise-linear function, and the shape is determined by the pieces:

  • With two pieces xx and 00 (that is, w1=1,b1=0w_1 = 1, b_1 = 0 and w2=0,b2=0w_2 = 0, b_2 = 0) we get max⁡(x,0)\max(x, 0), which is ReLU.
  • With the pieces xx and −x-x we get ∣x∣\lvert x\rvert, the absolute value.
  • With more pieces, it can approximate any convex function as closely as we like.

Here is the same idea in code:

import numpy as np

x = np.linspace(-2, 2, 401)
print("ReLU as maxout:", np.max([1 * x + 0, 0 * x + 0], axis=0)[[0, 100, 200, 300, 400]])
print("|x| as maxout: ", np.max([1 * x, -1 * x], axis=0)[[0, 100, 200, 300, 400]])
for pieces in (3, 5, 9):
    points_ = np.linspace(-2, 2, pieces)
    lines = [(2 * p, -p * p) for p in points_]                      # tangent line of x^2 at p
    approx = np.max([a * x + b for a, b in lines], axis=0)
    print(f"{pieces} pieces: largest error approximating x^2 on [-2, 2]: {np.max(np.abs(approx - x ** 2)):.4f}")

Output (values at x=−2,−1,0,1,2x = -2, -1, 0, 1, 2 for the first two lines):

ReLU as maxout: [0. 0. 0. 1. 2.]
|x| as maxout:  [ 2.  1. -0.  1.  2.]
3 pieces: largest error approximating x^2 on [-2, 2]: 1.0000
5 pieces: largest error approximating x^2 on [-2, 2]: 0.2500
9 pieces: largest error approximating x^2 on [-2, 2]: 0.0625

The last three lines use tangent lines to the curve x2x^2, which is convex. More pieces give a closer fit: halving the spacing between the tangent points divides the largest error by four (1, 0.25, 0.0625).

The curve x squared approximated by the maximum of three straight lines and by the maximum of five straight lines, with the five lines drawn faintly
A maximum of tangent lines hugs a convex curve. Five pieces are already close.

Advantages and costs

Advantages

  • No saturation and no dead region. One of the pieces is always active, and each piece has a non-zero slope in general, so gradients flow. (A ReLU unit, in contrast, is flat for negative inputs.)
  • Flexible. The network learns what shape of activation it needs, separately for each unit.
  • Pairs well with dropout. Maxout was designed to be used with dropout, and the combination gave strong results when it was introduced.

Costs

  • kk times as many parameters for each unit, because each piece has its own weights. A maxout layer with k=2k = 2 pieces has twice the weights of a ReLU layer with the same number of units.
  • The functions it represents are convex in each unit, which is less general than it may sound, though a network of many units can still model non-convex functions.

Where it fits

Maxout is a useful conceptual bridge. It shows ReLU and leaky ReLU as special cases of a more general idea, and it shows one way to learn what the earlier lesson chose by hand. In practice, simple ReLU-type activations or GELU are used more often today because they are cheaper, but the idea of learnable piecewise-linear activations appears in several later designs.

Try it yourself
Activations and initialization →

Switch to the maxout view and add more linear pieces.

EasyMaxout

Which two linear pieces make a maxout unit behave exactly like ReLU?

EasyMaxout

Why does a maxout unit with k pieces have k times as many parameters as a ReLU unit?

MediumMaxout

Why can a maxout unit with enough pieces approximate any convex function?