ML coding problem · 6 · Interview prep
Implement softmax (numerically stable)
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
- softmax(x)_i = exp(x_i) / Σ_j exp(x_j). Two passes: exponentiate, then divide by the total.
- 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.
import math
def softmax(xs):
m = max(xs)
exps = [math.exp(x - m) for x in xs]
total = sum(exps)
return [e / total for e in exps]
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