← Blog/blog/two-level-softmax-cluster-mass

The fast softmax sampler that weighs a cluster like one item

Suppose a recommender has millions of candidate songs. Exact softmax scores every song, so a common shortcut samples a cluster first and a song second. Bendada and Salha-Galvan show that the shortcut's cluster softmax quietly treats a centroid as though it represented one item. A cluster of ten equally good songs can receive the same mass as a cluster of one.

01

Where the missing mass goes

Softmax assigns each item weight exp(score / temperature). The true mass of a cluster is therefore a sum of exponentials. Standard two-level softmax, abbreviated 2LS, replaces that sum with the exponential of the cluster's mean score. The missing multiplier is the number of items, and the missing shape term comes from variation around the mean.

small-cluster ratiolarge-cluster ratioexact target
Illustrative equal-score clusters. One cluster contains one item; the other grows from one to twelve. Ordinary 2LS keeps giving each cluster 50%, so the small cluster is increasingly oversampled. Ratios are approximate mass divided by exact mass.
02

Two corrections, one compact rule

S-2LS adds cluster size. SD-2LS also adds half the variance of scaled item scores along the current query. The variance term is the Gaussian moment-generating-function correction: if projected scores are normal, their exponential mean is exp(mean + variance / 2).

def cluster_probs(clusters, temperature, method):
    logits = []
    for scores in clusters:
        scaled = [score / temperature for score in scores]
        mean = sum(scaled) / len(scaled)
        variance = sum((x - mean) ** 2 for x in scaled) / len(scaled)
        logit = mean
        if method != "2ls": logit += log(len(scaled))
        if method == "sd-2ls": logit += variance / 2
        logits.append(logit)
    return softmax(logits)

The TypeScript core used by every chart implements this rule directly. The following deterministic toy holds size and mean fixed, then spreads the second cluster's scores symmetrically. It is illustrative, not paper data.

exact softmax massordinary 2LSSD-2LS
Illustrative cluster mass as symmetric score dispersion rises. Ordinary 2LS stays at 50%; over this range, the variance term moves SD-2LS close to the true softmax mass without visiting every item at sampling time.
03

The variance correction cannot see skewness

Mean and variance do not determine an arbitrary distribution's exponential mass. Here, two four-item clusters both have mean zero and variance one. One has scores [-1, -1, 1, 1]; the other has one large negative score and three smaller positive scores. SD-2LS sees identical summaries and assigns 50% to each, while exact softmax does not.

Illustrative equal-mean, equal-variance counterexample at temperature 1. Exact softmax shifts 2.8 percentage points away from 50%, while SD-2LS cannot distinguish the clusters.
ordinary 2LS KLSD-2LS KL
Illustrative KL(approximate || exact) on the skewed counterexample as temperature changes. SD-2LS and ordinary 2LS coincide because both clusters have the same size, mean, and variance; neither can use the missing higher moments.
04

What the large-scale experiments say

The paper evaluates 1.2 million GloVe words, 7.7 million YAMBDA tracks, 19.6 million VK-LSVD videos, and two million-item synthetic corpora. For each dataset and temperature it compares probabilities for 1,000 query embeddings. On the three natural corpora, SD-2LS has the lowest KL at every tested temperature. S-2LS matches ordinary 2LS latency within measurement noise; SD-2LS is slower because it evaluates cluster covariance quadratic forms.

Datasetτ2LS KLS-2LS KLSD-2LS KL
GloVe-1000.100.0615 ± 0.00030.0022 ± 0.00020.0001 ± 0.0000
VK-LSVD0.100.0297 ± 0.00030.0051 ± 0.00010.0002 ± 0.0000
YAMBDA0.100.0300 ± 0.00070.0098 ± 0.00030.0005 ± 0.0000
GloVe-1000.200.0677 ± 0.00010.0001 ± 0.00001.2e−6 ± 9e−8
Selected paper-reported KL(approximate || exact), mean ± 95% confidence interval over 1,000 queries. These are not values from the illustrative charts above.

At temperature 0.1, the reported SD-2LS KL is roughly 59 times smaller than ordinary 2LS on GloVe, 149 times smaller on VK-LSVD, and 60 times smaller on YAMBDA. But the synthetic balanced result at temperature 0.05 is a useful brake on universal claims: top-k is best there, and SD-2LS is slightly worse than ordinary 2LS. Structure determines which approximation wins.

Claim layerWhat it repairsWhat remains open
Size correctionUnequal numbers of items per clusterDifferent score dispersion and shape
Variance correctionSecond moment along the current querySkewness, kurtosis, and other higher moments
Gaussian guaranteeAsymptotic cluster mass under the modelFinite clusters and non-Gaussian score projections
Reported latencySingle-threaded CPU comparisonGPU kernels, batching, memory, and serving load
The corrections are cumulative, but each guarantee has a narrower denominator than 'matches softmax everywhere.'
05

What to probe next

First, stratify residual KL by cluster skewness and kurtosis rather than reporting only an average. Second, compare diagonal, low-rank, and full-rank covariance corrections under a fixed latency budget. Third, test whether errors multiply down deeper hierarchical trees, as the authors conjecture. Finally, measure end-to-end recommendation quality: probability fidelity is a mechanism metric, not proof of better user outcomes.

References

  1. Walid Bendada and Guillaume Salha-Galvan (2026). Two-Level Softmax Sampling Done Right: Correcting Bias from Size Imbalance and Dispersion. NeurIPS 2026 / arXiv:2610.10483
  2. Frederic Morin and Yoshua Bengio (2005). Hierarchical Probabilistic Neural Network Language Model. AISTATS 2005
  3. Jean Pennington, Richard Socher, and Christopher Manning (2014). GloVe: Global Vectors for Word Representation. EMNLP 2014