aiengineering.guideaiengineering.guide

ML coding problem · 6 · Interview prep

Implement softmax (numerically stable)

easyNeural network componentsneural-network-componentsnumerical-stabilityattention

Turn a vector of scores into a probability distribution — the function at the end of every classifier and inside every attention head. The catch is doing it without overflowing.

Problem

The problem

Implement softmax(xs). It takes a list of numbers and returns a new list of the same length where every element is in (0, 1) and the whole list sums to 1 — the standard softmax.

softmax(x)_i = exp(x_i) / Σ_j exp(x_j)

Your implementation must be numerically stable: softmax([1000, 1000, 1000]) should return [0.333…, 0.333…, 0.333…], not a list of nans from an exp overflow. Pure Python — math is available, nothing else.

Concept

softmax maps a vector of real-valued scores to a probability distribution: each output is in (0, 1) and they sum to 1. The naive formula exp(x_i) / Σ exp(x_j) overflows the moment any x_i is large (exp(1000) is inf). The fix is the shift trick: subtract max(x) from every element first. softmax is shift-invariant, so the result is identical, but now the largest exponent is exp(0) = 1 and nothing overflows. This is why every real implementation subtracts the max — it is correctness, not just an optimisation.

Hints

  1. softmax(x)_i = exp(x_i) / Σ_j exp(x_j). Two passes: exponentiate, then divide by the total.
  2. Subtract max(xs) from every element before calling exp(). It changes nothing mathematically (softmax is shift-invariant) but stops exp() from overflowing on large inputs.

Solution

Try it yourself first — the reveal takes two clicks.

solution.pyPython · Pyodide
No browser Python? Run it locally instead.

The runner is a Pyodide worker (~10 MB, WebAssembly) loaded on first run. If your network or browser blocks it, copy your code and the tests below into solution.py andtest_solution.py and run pytest.

# test_solution.py
def test_sums_to_one():
    out = softmax([1.0, 2.0, 3.0])
    assert abs(sum(out) - 1.0) < 1e-9

def test_monotonic():
    out = softmax([1.0, 2.0, 3.0])
    assert out[0] < out[1] < out[2]

def test_uniform_inputs_give_uniform_output():
    out = softmax([5.0, 5.0, 5.0, 5.0])
    for p in out:
        assert abs(p - 0.25) < 1e-9

def test_numerically_stable_on_large_inputs():
    # A naive exp(x)/sum(exp(x)) overflows here; a stable one does not.
    out = softmax([1000.0, 1000.0, 1000.0])
    assert abs(sum(out) - 1.0) < 1e-9
    for p in out:
        assert abs(p - (1.0 / 3.0)) < 1e-9

Enter to go · Esc to close