Concept 02: Autoregressive Next-Token Sampling (Temperature, Top-k, Top-p)

Large Language Models do not generate entire sentences or paragraphs in a single step. Instead, they operate autoregressively: predicting exactly one next token at a time, appending it to the prompt, and feeding the updated sequence back into the model.

At the output of the final Transformer layer, the model produces un-normalized prediction scores called Logits across all 50,000+ vocabulary tokens.

How do we pick the winner?

Open the interactive demo below to adjust Temperature, Top-k, and Top-p sliders, observe the live token probability distribution, and click “Generate Next Token” to sample tokens in real time.


1. Temperature Scaling: Shaping the Distribution

Before applying Softmax, we divide all logits by a positive scalar called Temperature (T):

P(token_i) = exp( zᵢ / T ) / ∑ exp( zⱼ / T )

2. Top-k & Top-p (Nucleus) Filtering

Pure sampling from a 50,000-word vocabulary can occasionally pick bizarre low-probability tokens from the long tail. We filter candidates using two complementary techniques:

Top-k Filtering

Keep only the k highest-probability tokens (e.g. k = 40) and set all other token probabilities to zero.

Top-p (Nucleus) Sampling

Sort all tokens in descending probability order and keep only the smallest group whose cumulative probability reaches threshold p (e.g. p = 0.90):


3. Solving It in Code (Java)

Here is a complete Next-Token Sampler in Java with Temperature and Top-p Nucleus filtering:

import java.util.*;

public class TokenSampler {
    public record Candidate(String token, double prob) {}

    public static String sampleNextToken(double[] logits, String[] vocab, double temperature, double topP) {
        int n = logits.length;

        // 1. Temperature Scaling + Softmax
        double maxLogit = Double.NEGATIVE_INFINITY;
        for (double z : logits) maxLogit = Math.max(maxLogit, z / temperature);

        double expSum = 0.0;
        double[] probs = new double[n];
        for (int i = 0; i < n; i++) {
            probs[i] = Math.exp((logits[i] / temperature) - maxLogit);
            expSum += probs[i];
        }
        for (int i = 0; i < n; i++) probs[i] /= expSum;

        // 2. Sort Candidates Descending
        List<Candidate> candidates = new ArrayList<>();
        for (int i = 0; i < n; i++) candidates.add(new Candidate(vocab[i], probs[i]));
        candidates.sort((a, b) -> Double.compare(b.prob, a.prob));

        // 3. Top-p Nucleus Truncation
        List<Candidate> nucleus = new ArrayList<>();
        double cumulativeProb = 0.0;
        for (Candidate c : candidates) {
            nucleus.add(c);
            cumulativeProb += c.prob;
            if (cumulativeProb >= topP) break;
        }

        // 4. Sample from Nucleus
        double randomVal = Math.random() * cumulativeProb;
        double runningSum = 0.0;
        for (Candidate c : nucleus) {
            runningSum += c.prob;
            if (randomVal <= runningSum) return c.token;
        }

        return nucleus.get(0).token;
    }

    public static void main(String[] args) {
        String[] vocab = {"score", "intake", "align", "wait", "dance"};
        double[] logits = {6.2, 5.1, 4.0, 1.2, -2.5};

        String sampled = sampleNextToken(logits, vocab, 0.7, 0.90);
        System.out.printf("Sampled Next Token: \"%s\"%n", sampled);
    }
}

4. Math! Translation Sidebar

The formal mathematical definition of Nucleus Sampling:

Nucleus S_p = smallest subset of V such that: ∑ P(x) ≥ p

The KV-Cache Optimization:


← Concept 01: RoPE Embeddings
Module 4 Overview
LLM Axon Home →