Weighted sampling means picking items with probability that depends on a weight: training examples by loss, log lines by severity, users by activity, tokens by model probability. It sounds like one problem. It is at least three, and the most common bug is solving the wrong one. Do you want draws with replacement, so a heavy item can appear many times? A set of k distinct items, where heavy items are more likely to be in it? Or a sample from which you will estimate totals, which needs known inclusion probabilities?
This article maps those three cases to the algorithms that solve them, then goes deep on the streaming case: Efraimidis and Spirakis's A-Res and A-ExpJ, which sample k items without replacement from a stream of unknown length in one pass. All code is tested; the probabilities quoted were measured by running it and checked against exact enumeration. Merging samples across machines is covered in the reservoir sampling architecture article, and we link there rather than repeat it.
Three different problems
With replacement. Each draw picks item i with probability w_i / W, where W is the total weight, independently of earlier draws. For a fixed array, compute prefix sums once and binary-search a uniform number in [0, W): O(log n) per draw. If you need millions of draws from one fixed distribution, the alias method gives O(1) per draw after O(n) setup. If weights change between draws, as in prioritized experience replay, a sum tree (a binary tree whose nodes hold the sum of their children) gives O(log n) updates and draws.
Without replacement. Pick k distinct items. The natural definition, and the one A-Res implements, is successive sampling: draw one item with probability proportional to weight, remove it, repeat on the rest. Note what this does not promise: item i is not included with probability k w_i / W. We will compute the actual probabilities below.
For estimation. If you will compute sums or averages from the sample, you need each item's inclusion probability, or an estimator built around the sampler. Priority sampling provides one and is covered in its own section.
A-Res: keep the k largest keys
Efraimidis and Spirakis (2006) proved a neat fact. Give each item a key u^(1/w), with u uniform on (0, 1), and keep the k items with the largest keys. The result has exactly the distribution of successive sampling. Because the rule is just keep the top k of independent keys, it works on a stream in one pass with a min-heap of size k, and it does not need to know n or W in advance.
In code, take logarithms. log(u^(1/w)) = log(u) / w is a monotone transform of the same key, so the top k are unchanged, and it avoids raising a number near 1 to a huge power when w is small, which rounds every key to 1.0 in double precision.
import heapq, math, random
def a_res(stream, k, rng=random):
heap = [] # min-heap of (key, item)
for item, w in stream:
if w <= 0:
continue # zero weight: never sampled
key = math.log(rng.random() or 5e-324) / w
if len(heap) < k:
heapq.heappush(heap, (key, item))
elif key > heap[0][0]:
heapq.heapreplace(heap, (key, item))
return [item for _, item in heap]Cost: O(n log k) time, O(k) memory, one random number per item. The or 5e-324 guard avoids log(0) if the generator ever returns exactly zero. If items are equal-weight, this collapses to a correct uniform reservoir sampler.
Worked example: weights 1, 2, 3, 4
Take four items with weights 1, 2, 3, 4 (W = 10) and k = 2. Under successive sampling, the first draw picks item 4 with probability 0.4. The probability that item 4 ends up in the sample adds every way it can get there: first draw (0.4), or second draw after item j was taken first, 0.1 x 4/9 + 0.2 x 4/8 + 0.3 x 4/7. Summing over all 12 ordered pairs gives exact inclusion probabilities:
| Item weight | Proportional k w / W | Exact successive | A-Res, 200,000 runs | A-ExpJ, 200,000 runs | Gumbel-top-k, 200,000 runs |
|---|---|---|---|---|---|
| 1 | 0.200 | 0.2345 | 0.2345 | 0.2336 | 0.2354 |
| 2 | 0.400 | 0.4413 | 0.4403 | 0.4405 | 0.4415 |
| 3 | 0.600 | 0.6083 | 0.6090 | 0.6084 | 0.6078 |
| 4 | 0.800 | 0.7159 | 0.7161 | 0.7174 | 0.7154 |
The three samplers agree with the exact enumeration to within sampling noise (the standard error at 200,000 runs is about 0.001). And the exact column is not proportional to weight: the heaviest item is under-included (0.72, not 0.80) and the lightest is over-included. That is inherent to without replacement, since no item can be included twice, so heavy items lose share. If downstream code assumes inclusion equals k w / W, its estimates are biased; we measure how much below.
A-ExpJ: jump over items that cannot enter
A-Res draws a random number for every item, even though late in a long stream almost no item enters the heap. A-ExpJ (exponential jumps) skips ahead instead. Let T be the smallest key in the heap. The next item to enter is the first whose key beats T, and the total weight that passes before that happens has a known distribution: draw r uniform and skip a weight of log(r) / log(T). The item where the running weight crosses that amount enters; its key is not drawn fresh but drawn from the part of the distribution above T, which is what the uniform(t_w, 1) line does.
def a_expj(stream, k, rng=random):
heap, it = [], iter(stream)
for item, w in it: # fill exactly as A-Res
if w > 0:
heapq.heappush(heap, (math.log(rng.random() or 5e-324) / w, item))
if len(heap) == k:
break
if len(heap) < k:
return [i for _, i in heap]
log_t = heap[0][0] # log of threshold key T
x = math.log(rng.random() or 5e-324) / log_t # weight to skip
for item, w in it:
if w <= 0:
continue
x -= w
if x <= 0: # this item enters the heap
t_w = math.exp(log_t * w) # T ** w
r2 = rng.uniform(t_w, 1.0) # key conditioned to beat T
heapq.heapreplace(heap, (math.log(r2) / w, item))
log_t = heap[0][0]
x = math.log(rng.random() or 5e-324) / log_t
return [i for _, i in heap]The expected number of entries after the heap fills grows like k ln(n/k), and each costs two random numbers. Measured with k = 100 and weights cycling 1 to 7:
| Stream length n | A-Res random numbers | A-ExpJ random numbers | 2 k ln(n/k) + k |
|---|---|---|---|
| 10,000 | 10,000 | 1,035 | 1,021 |
| 100,000 | 100,000 | 1,537 | 1,482 |
| 1,000,000 | 1,000,000 | 1,971 | 1,942 |
Random numbers are cheap in C; the real saving is that A-ExpJ touches each skipped item only to subtract its weight, which is what you want when the generator is a cryptographic one or when keys are expensive to compute.
The same sampler as Gumbel-top-k
The same sampler hides in machine learning under another name. If g is standard Gumbel noise, g = -log(-log(u)), then the index maximizing log(w_i) + g_i is distributed proportionally to w_i: the Gumbel-max trick. Taking the top k instead of the maximum, Gumbel-top-k, gives successive sampling without replacement. It is the same as A-Res: log(w) - log(-log u) is a monotone function of log(u) / w, so the rankings match exactly. Our run confirms the first-pick probabilities of 0.0999, 0.2006, 0.2999, 0.3995 against the expected 0.1, 0.2, 0.3, 0.4, and the inclusion probabilities in the table above.
This matters in practice because it vectorises. With weights as logits in a tensor, add Gumbel noise and call top-k once, which samples k distinct classes, tokens or negatives in one GPU operation. Kool, van Hoof and Welling (2019) used this to sample sequences without replacement from language models (stochastic beam search).
import numpy as np
def gumbel_top_k(logits, k, rng=np.random.default_rng()):
g = -np.log(-np.log(rng.random(logits.shape)))
return np.argpartition(-(logits + g), k - 1, axis=-1)[..., :k]
Estimating totals: priority sampling
Suppose the sample feeds a dashboard: the total bytes of all requests from one customer, estimated from 200 sampled requests. The textbook Horvitz-Thompson estimator divides each sampled value by its inclusion probability. With A-Res you do not know those probabilities in closed form, and plugging in k w / W is wrong. On 5,000 items with heavy-tailed Pareto weights, estimating the total of every third item over 2,000 runs, that shortcut was biased by +21 percent.
Priority sampling (Duffield, Lund and Thorup, 2007) fixes this. Give each item a priority q = w / u, keep the k + 1 highest, and let tau be the smallest of them. The k items above tau form the sample, and each gets the estimate max(w, tau). The sum of estimates over any subset is an unbiased estimate of that subset's true total, and the variance is close to optimal. In the same experiment priority sampling gave a mean estimate within 0.14 percent of the truth, with a relative standard deviation of 8 percent per run.
def priority_sample(stream, k, rng=random):
heap = [] # (priority, item, weight), size k + 1
for item, w in stream:
q = w / (1.0 - rng.random()) # u in (0, 1]
if len(heap) < k + 1:
heapq.heappush(heap, (q, item, w))
elif q > heap[0][0]:
heapq.heapreplace(heap, (q, item, w))
if len(heap) <= k:
return {item: w for _, item, w in heap} # everything kept, exact totals
tau = heap[0][0]
return {item: max(w, tau) for q, item, w in heap if q > tau}Rule of thumb: inspection samples (show me some error traces, preferring big ones) use A-Res or Gumbel keys; estimation samples (what were total bytes by customer) use priority sampling or keep exact counters alongside.
Operational guidance
- Seed for reproducibility, hash for consistency. For a training data pipeline, seed the generator per epoch and log the seed. To have every service keep the same trace, derive u from a hash of the item id and a shared seed; the reservoir architecture article shows how.
- Validate weights at the edge. Reject negative, NaN and infinite weights explicitly. Decide whether zero means never sampled, and log the count of dropped items.
- Test against enumeration. Keep the four-item test above: exact inclusion probabilities by permutation, then a chi-square or tolerance check on 100,000 runs. It catches inverted keys, which is the most common bug.
- Mind the heap direction. A-Res keeps the largest keys, so the heap root must be the smallest kept key. The exponential-race formulation, -log(u) / w, keeps the smallest keys instead; mixing conventions silently samples the lightest items.
- Store the threshold. Persist T (or tau) with the sample. It makes the sample resumable and gives priority sampling its estimates.
Failure modes
| Symptom | Cause | Fix |
|---|---|---|
| Light items dominate the sample | Kept smallest u^(1/w) or largest -log(u)/w | Fix the heap direction; test on 1, 2, 3, 4 |
| All keys equal 1.0 for tiny weights | Computed u ** (1/w) directly | Use log(u) / w |
| Totals biased upward or downward | Assumed inclusion = k w / W | Priority sampling or exact counters |
| Crash on log(0) | Generator returned exactly 0 | Guard u, or draw from (0, 1] |
| Sample differs per service | Independent random u per process | Hash-derived u with a shared seed |
| Items silently missing from the sample | w = 0 or NaN silently skipped | Validate and count rejected weights |
Trade-offs
Prefix sums and binary search are the simplest with-replacement sampler; alias tables are faster per draw but must be rebuilt when weights change; sum trees handle changing weights at O(log n). For without replacement on a stream, A-Res is the clearest code, A-ExpJ saves random numbers on long streams for a little more code, and Gumbel-top-k is the fastest choice when weights already sit in a tensor. Priority sampling costs one extra heap slot and gives unbiased totals. Every streaming method here keeps O(k) memory and uses a heap; the heap deep dive covers that structure, and the reservoir sampling article covers the uniform special case and its correctness proof.
What to do next
- Write down which of the three questions your sample answers: draws, a distinct set, or estimates.
- For draws from a fixed distribution, use an alias table; for changing weights, a sum tree.
- For a distinct weighted set from a stream, copy a_res and its four-item enumeration test.
- If the stream is long and per-item randomness is expensive, switch to a_expj and rerun the same test.
- If the sample feeds any sum or average, use priority_sample and check its estimates against exact counters for a week.
- On GPUs, replace the loop with Gumbel-top-k over logits.