Matrix multiplication is associative: (AB)C and A(BC) give the same matrix. It is not equally expensive both ways. For a chain of matrices, the order you choose can change the arithmetic by orders of magnitude while the answer stays identical. Matrix chain multiplication is the problem of finding the cheapest order, and its standard solution is the textbook example of interval dynamic programming.

This article builds the algorithm from first principles: the cost of one product, why trying every order is hopeless, the recurrence, a bottom-up implementation with reconstruction, and a fully worked six-matrix example whose numbers come from running the code shown. It then covers where the problem appears in everyday ML software. NumPy's multi_dot solves it, einsum path optimisation generalises it, and reverse-mode automatic differentiation is a matrix chain ordered the cheap way. It closes with the cost-model caveats and bugs that matter when you use it for real.

What one multiplication costs

Multiplying a p x q matrix by a q x r matrix produces a p x r result. Each of the p*r output entries is a dot product of length q, so the naive algorithm performs p*q*r scalar multiplications. That count is the standard cost model for this problem. It ignores additions, which scale the same way, and memory traffic, which we return to later.

Take three matrices: A1 is 10 x 30, A2 is 30 x 5, A3 is 5 x 60.

  • (A1 A2) A3: first 10*30*5 = 1,500 to get a 10 x 5 matrix, then 10*5*60 = 3,000. Total 4,500.
  • A1 (A2 A3): first 30*5*60 = 9,000 to get a 30 x 60 matrix, then 10*30*60 = 18,000. Total 27,000.

Same answer, six times the work. The expensive order creates a large intermediate, 30 x 60, while the cheap one collapses to a small 10 x 5 matrix early. That intuition, shrink early, is right but not sufficient. In longer chains the best order is not obvious, and greedy rules such as "always multiply the pair with the smallest cost first" fail on some inputs.

A chain of n matrices is described by n+1 numbers. Matrix Ai has dimensions p[i-1] x p[i], so adjacent matrices share a dimension. The three-matrix example is p = [10, 30, 5, 60].

Why trying every order fails

How many ways can you parenthesise n matrices? Choose the last multiplication, which splits the chain into a left part of k matrices and a right part of n-k. Each part can be parenthesised independently. This gives the recurrence for the Catalan numbers, and the count for n matrices is the Catalan number C(n-1):

Matrices nParenthesisations C(n-1)
32
642
104,862
201,767,263,190
301,002,242,216,651,368

The count grows roughly like 4^n, so enumeration is out beyond very short chains. The same split argument that produced the count also produces the fix. The best way to multiply a sub-chain does not depend on how the rest of the chain is handled, so each sub-chain needs to be solved once.

The recurrence

Let m[i][j] be the minimum cost to compute the product Ai...Aj. A single matrix costs nothing, so m[i][i] = 0. For i < j, the final multiplication splits the chain at some k with i ≤ k < j. It multiplies the result of Ai..Ak, a p[i-1] x p[k] matrix, by the result of Ak+1..Aj, a p[k] x p[j] matrix:

m[i][i] = 0
m[i][j] = min over i <= k < j of  m[i][k] + m[k+1][j] + p[i-1] * p[k] * p[j]

Why it is correct (optimal substructure). Suppose the best order for Ai..Aj makes its last split at k. The left part must then be computed in its own cheapest order. If a cheaper way existed, substituting it would lower the total, contradicting optimality. The same holds for the right part. So the optimum is built from optimal sub-solutions, and trying every k covers every possible last split.

Why it needs DP (overlapping subproblems). Plain recursion re-solves the same sub-chains many times. A2..A4 is needed both when A1..A4 splits after A1 and when A2..A5 splits after A4. There are only n(n+1)/2 distinct sub-chains, so storing each answer once makes the problem polynomial.

Fill order. m[i][j] depends only on shorter sub-chains, so fill the table by increasing chain length: all pairs, then all triples, and so on up to the full chain. This is the defining shape of interval DP. In the table it means filling diagonal by diagonal, moving up and to the right.

Bottom-up implementation

The bottom-up implementation below stores the cost table m and a split table s that records the winning k for each sub-chain. Indices are 0-based in the code, so matrix i has dimensions p[i] x p[i+1].

import math

def matrix_chain_order(p: list[int]):
    n = len(p) - 1                       # number of matrices
    m = [[0] * n for _ in range(n)]      # m[i][j]: min cost for A_i..A_j (0-based)
    s = [[0] * n for _ in range(n)]      # s[i][j]: best split k
    for length in range(2, n + 1):       # sub-chain length
        for i in range(n - length + 1):
            j = i + length - 1
            m[i][j] = math.inf
            for k in range(i, j):        # last multiply: (A_i..A_k)(A_k+1..A_j)
                cost = m[i][k] + m[k + 1][j] + p[i] * p[k + 1] * p[j + 1]
                if cost < m[i][j]:
                    m[i][j], s[i][j] = cost, k
    return m, s

def parenthesize(s, i: int, j: int) -> str:
    if i == j:
        return f"A{i + 1}"
    k = s[i][j]
    return "(" + parenthesize(s, i, k) + " " + parenthesize(s, k + 1, j) + ")"

def multiply_chain(mats, s, i, j):
    """Evaluate the product in the optimal order (mats are NumPy arrays)."""
    if i == j:
        return mats[i]
    k = s[i][j]
    return multiply_chain(mats, s, i, k) @ multiply_chain(mats, s, k + 1, j)

There are three nested loops over at most n values each, so time is O(n^3) and space is O(n^2). The DP itself does no matrix arithmetic; it only reasons about shapes. For chains of tens or hundreds of matrices it runs in microseconds to milliseconds, negligible next to the multiplications it saves.

A top-down version with memoisation computes the same table and is sometimes easier to adapt. For example, you might add a constraint or skip sub-chains that can never be optimal. Python's functools.cache on a function of (i, j) gives you that in a few lines. Mind the recursion depth for long chains.

Worked example: six matrices

The classic six-matrix instance uses p = [30, 35, 15, 5, 10, 20, 25], so A1 is 30 x 35, A2 is 35 x 15, A3 is 15 x 5, A4 is 5 x 10, A5 is 10 x 20 and A6 is 20 x 25. Running the code above gives this cost table, using 1-based labels in the headers:

m[i][j]j=1j=2j=3j=4j=5j=6
i=1015,7507,8759,37511,87515,125
i=202,6254,3757,12510,500
i=307502,5005,375
i=401,0003,500
i=505,000
i=60

Read one cell to see the recurrence at work. m[1][3] covers A1 A2 A3. Splitting after A1 costs 0 + m[2][3] + 30*35*5 = 2,625 + 5,250 = 7,875. Splitting after A2 costs m[1][2] + 0 + 30*15*5 = 15,750 + 2,250 = 18,000. The minimum, 7,875, is stored with split k = 1.

The top-right cell gives the answer: 15,125 scalar multiplications, with optimal order ((A1 (A2 A3)) ((A4 A5) A6)). The tree in Figure 1 shows how the total decomposes. Naive left-to-right evaluation costs 40,500 and right-to-left costs 47,500, so the DP saves a factor of about 2.7 to 3.1 here. The win comes from reaching the 30 x 5 and 5 x 25 shapes early: A3 has only 5 columns and A4 only 5 rows, so the split at k = 3 makes the final product cheap.

Optimal order for A1..A6 with p = [30, 35, 15, 5, 10, 20, 25]: total 15,125 scalar multiplicationsA1..A6 = 15,125split k=3: 7,875 + 3,500 + 30*5*25A1..A3 = 7,875split k=1: 0 + 2,625 + 30*35*5A4..A6 = 3,500split k=5: 1,000 + 0 + 5*20*25A130 x 35A2 A3 = 2,62535*15*5A4 A5 = 1,0005*10*20A620 x 25A235 x 15A315 x 5A45 x 10A510 x 20Each node's cost = left subtree + right subtree + rows(left) * shared dim * cols(right).Multiplying left to right instead costs 40,500, and right to left costs 47,500.
Figure 1. The optimal parenthesisation as a tree, with each node's cost broken into its two subtrees plus the final multiply.

Where ML code depends on it

You will rarely write this DP by hand, but you use it constantly. Knowing it explains several performance behaviours in ML code.

NumPy's multi_dot. numpy.linalg.multi_dot([A, B, C, D]) picks the evaluation order automatically. In current NumPy source, three arrays are handled by comparing the two possible costs directly, and longer chains go through a private matrix-chain-order routine that implements this DP. The order matters more than people expect. Writing A @ B @ C @ v in Python evaluates left to right and can form large intermediates that A @ (B @ (C @ v)) never creates.

einsum contraction paths. numpy.einsum(..., optimize=True) and numpy.einsum_path choose the order in which to contract a multi-operand expression. This is the same problem generalised from chains to arbitrary tensor networks. There, finding the optimal order is NP-hard, so libraries offer exhaustive search for small expressions and greedy heuristics for large ones. The chain case is the special case where an exact polynomial algorithm exists.

Backpropagation. The gradient of a scalar loss through layers is a chain product: a row vector v (the loss gradient) times Jacobians J2 and J1. Take two layers with 4096 x 4096 Jacobians, so p = [1, 4096, 4096, 4096]. Evaluating (v J2) J1 costs 2 * 4096 * 4096 = 33,554,432. Forming J2 J1 first costs about 68.7 billion. That factor of roughly 2,048 is why training uses reverse mode, which is exactly the left-to-right order the DP picks, and why frameworks compute vector-Jacobian products instead of materialising Jacobians. The backpropagation guide walks through the same chain from the network's side.

Low-rank adapters. A LoRA update applies BA to an activation x, with B of size d x r and A of size r x d. With d = 4096 and r = 16, computing (BA)x per token costs 285,212,672 multiplications, while B(Ax) costs 131,072, about 2,176 times fewer. Merging BA into the base weights once before deployment is a different trade: one large product amortised over every future token.

When FLOP counts mislead

The p*q*r model counts arithmetic. Real hardware also cares about these things:

  • Efficiency at small sizes. GPUs and BLAS libraries run large, well-shaped GEMMs near peak and thin ones far below it. Two orders with similar FLOPs can differ a lot in wall time. Benchmark when costs are within a factor of two.
  • Intermediate memory. The cheapest order can still create a large temporary. If memory is the constraint, add a term to the cost, or reject splits whose intermediate exceeds a budget. The DP structure is unchanged; only the cost expression grows.
  • Sparsity and structure. Diagonal, sparse or low-rank matrices break the dense cost formula. Use the real cost of each product in the recurrence.
  • Repeated chains. If the same shapes recur every training step, compute the order once and cache it. If one operand is fixed, such as merged weights, precompute that sub-product once.

The interval DP family

Matrix chain multiplication is the template for a family of interval DP problems. Each defines a cost on a contiguous range, splits the range at a chosen point, and fills the table by increasing length:

  • Optimal binary search trees: choose the root of each key range to minimise expected search cost. The optimal BST deep dive uses the same table shape.
  • Polygon triangulation: minimum-weight triangulation of a convex polygon is structurally the same recurrence.
  • Burst balloons and merging stones: pick the last element processed in a range, then solve both sides.

Faster algorithms exist for matrix chains specifically. Hu and Shing published an O(n log n) algorithm in 1982 and 1984 based on polygon partitioning, but it is intricate and rarely implemented, because chains in practice are short. Don't assume the Knuth speed-up for optimal BSTs carries over. It relies on a monotone split property that the matrix chain cost does not guarantee in general.

For the broader toolkit, dynamic programming fundamentals covers memoisation versus tabulation, and matrix exponentiation shows the opposite situation: all factors are the same matrix, so repeated squaring wins.

Common bugs

  • Off-by-one in dimensions. n matrices need n+1 numbers in p, and matrix i (0-based) is p[i] x p[i+1]. Mixing 0-based and 1-based indexing in the cost term p[i]*p[k+1]*p[j+1] is the most common bug. Check against the six-matrix answer of 15,125.
  • Wrong fill order. Iterating i and j directly instead of by length reads cells that haven't been computed yet. The result is wrong but plausible-looking.
  • Overflow. In languages with fixed-width integers, costs for large dimensions overflow 32 bits quickly. Use 64-bit integers or floats.
  • Ties. Several splits can share the minimum cost. That is harmless for cost, but it can change which intermediates exist. Break ties deliberately if memory matters.
  • Forgetting to apply the order. Computing s and then multiplying left to right anyway happens more often than you'd think. Evaluate through the split table, as multiply_chain does.

What to do next

  1. Implement matrix_chain_order and confirm 15,125 and ((A1 (A2 A3)) ((A4 A5) A6)) for p = [30, 35, 15, 5, 10, 20, 25].
  2. Write a brute-force recursive version without memoisation, check it agrees on random chains of up to 8 matrices, and time both.
  3. Benchmark A @ B @ C @ v against numpy.linalg.multi_dot for shapes where the orders differ, and compare the measured ratio with the FLOP ratio.
  4. Add an intermediate-memory budget to the recurrence and see how the chosen order changes.
  5. Use numpy.einsum_path on a four-operand expression and read the contraction order it reports.
  6. Solve optimal BST and polygon triangulation with the same table structure to fix the interval-DP pattern.
Key takeaway: The order of a matrix chain does not change the result but can change the cost by orders of magnitude. The interval DP finds the cheapest order in O(n^3) by solving every sub-chain once, in order of length, and a split table recovers the order. The same idea drives multi_dot, einsum path search and reverse-mode autodiff, so check intermediate shapes before trusting a left-to-right chain.