Sulba
000 / 100

Probability

A language model’s output is a probability for every word that could come next, and its training loss is built from those probabilities. Learn what a probability is, how probabilities combine, expected value, why models work with logs of probabilities, and the cross-entropy loss.

Lesson 6 stepsExercise Python, 6 functionsQuiz 4 questions

Step 1 Probability as a fraction

A probability says how likely something is: 0 is never, 1 is certain.When every outcome is equally likely: favourable outcomes ÷ all outcomes.11/621/631/641/651/661/6P(six) = 1 ÷ 6 ≈ 0.1667P(even) = 3 ÷ 6 = 0.5Roll it many times and the shareof sixes settles near 1/6.share of sixes1/660 rolls0.15600 rolls0.166,000 rolls0.1632
01/06

A probability is a number from 0 to 1 that says how likely something is. 0 means it never happens, 1 means it always happens, and 0.5 means it happens half the time.

When every outcome is equally likely, the probability of an event is the number of outcomes in it divided by the number of outcomes in all. A fair die has 6 faces, each equally likely:

P(A)=number of outcomes in Anumber of outcomes in all
SymbolSayMeans
AAAn event: a set of outcomes, such as “an even number”.
P(A)P of AThe probability that A happens, from 0 (never) to 1 (always).
EventOutcomes in itProbability
A six61 ÷ 6 = 0.1667
An even number2, 4, 63 ÷ 6 = 0.5
More than 45, 62 ÷ 6 = 0.3333
Any number1, 2, 3, 4, 5, 66 ÷ 6 = 1
A sevennone0 ÷ 6 = 0

A probability also predicts what you see over many tries. Here is one simulated run of 6,000 rolls of a fair die, with the share of sixes after the first 60, the first 600 and all 6,000:

RollsSixesShare of sixesDistance from 1/6
6099 ÷ 60 = 0.150.0167
6009696 ÷ 600 = 0.160.0067
6,000979979 ÷ 6,000 = 0.16320.0035

The share wanders at first and settles near 1/6 as the rolls add up. In the same way, a model’s accuracy measured on 60 examples can be far from its true accuracy, and on 6,000 it is much closer. The Statistics lesson works out how close.

Try it yourself

Roll one die or two, as many times as you like, and compare the share of each result with its exact probability. Then set the probability a model gave the right answer and see its loss.

123456

share in your rolls exact probability

rolls: 600

sixes: 96

share = 96 ÷ 600 = 0.16

exact = 1 ÷ 6 = 0.1667

average = 2,041 ÷ 600 = 3.4017

expected value = 3.5

Each face has probability 1/6. With few rolls the shares wander; with more, they settle near it.

Try 10 rolls a few times, then 10,000. The shares and the average land closer to the exact values as the rolls add up.

loss = −ln p = −ln 0.6 = 0.5108

Fairly sure and right: a small loss.

Programming exercise

Work out probabilities, expected values, log-probabilities and the cross-entropy loss in plain Python. Save probability.py and test_probability.py in the same folder, fill in each function in probability.py, and run the tests:

python test_probability.py

"""Probability: programming exercise. Work out probabilities, expected values, log-probabilities and thecross-entropy loss in plain Python, then run the tests from this folder:     python test_probability.py math.log(x) is the natural log, ln x.""" import math  # noqa: F401  (you will need math.log)  def probability(favourable, total):    """Return the probability of an event with `favourable` outcomes out of `total` equally likely ones."""    raise NotImplementedError  def share(rolls, value):    """Return the share of the list `rolls` that equal `value`: how many there are, divided by how many rolls."""    raise NotImplementedError  def both(p_a, p_b):    """Return the probability that two independent events both happen."""    raise NotImplementedError  def expected_value(outcomes, probs):    """Return the expected value: each outcome times its probability, added up."""    raise NotImplementedError  def log_prob(probs):    """Return the log of the product of `probs`, without multiplying them: the sum of their logs."""    raise NotImplementedError  def cross_entropy(probs):    """Return the cross-entropy loss: the average of -ln p over the probabilities given to the right answers."""    raise NotImplementedError 

Stuck? math.log(x) is ln x in Python. The Solution tab has one way to write each function.

In practice: probabilities in language models

Language models are trained on cross-entropy. At every position in the training text, the model gives a probability to every token it knows (a token is a word or a piece of a word), and the loss is −ln of the probability it gave the token that actually came next, averaged over all positions. In PyTorch this is torch.nn.functional.cross_entropy, which takes the model’s raw scores and works out the probabilities and their logs itself.

Work in logs. Multiplying hundreds of probabilities rounds to 0, which breaks both training and the scoring of whole sentences. Libraries add log-probabilities instead, using functions such as log_softmax and logsumexp that never form the tiny products at all.

Perplexity. Language-model results are often given as perplexity, which is e raised to the power of the cross-entropy loss. A perplexity of 20 means the model is, on average, as unsure as if it were choosing evenly among 20 tokens.

Test your knowledge

  1. 01A bag has 3 red and 5 blue marbles. What is the probability of drawing a red one?Show answer

    3 ÷ 8 = 0.375. There are 3 favourable outcomes out of 8 equally likely ones.

  2. 02You flip a fair coin 3 times. What is the probability of 3 heads?Show answer

    The flips are independent, so multiply: 0.5 × 0.5 × 0.5 = 0.125, or 1 in 8.

  3. 03A model gives the right answer probability 0.9 on one example and 0.1 on another. What is its cross-entropy loss over the two?Show answer

    −ln 0.9 = 0.1054 and −ln 0.1 = 2.3026. The average is (0.1054 + 2.3026) ÷ 2 = 1.204.

  4. 04Why do models add log-probabilities instead of multiplying probabilities?Show answer

    Multiplying many probabilities below 1 gives numbers too small for the computer to store, so they round to 0. Logs turn the product into a sum of ordinary-sized numbers, and since ln(a × b) = ln a + ln b, nothing is lost.