Samplers and weighted sampling
A sampler decides which samples the DataLoader fetches and in what order — and a weighted sampler makes rare classes show up as often as common ones without touching the data.
- 8 min read
- 3 reading levels
- Published
Read these first
On this page 5
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
A sampler is the part of the DataLoader that chooses which sample comes next — and you can replace it to change what the model gets to see often.
Think of a classroom with forty loud students and two quiet ones. A teacher who calls on raised hands hears the loud ones all day. The quiet students are present — they are never called on.
A good teacher keeps a list and deliberately calls the quiet ones more often. Nothing about the class changed; only the calling-on rule did.
The sampler is that rule. The DataLoader's first stage picks index numbers, and by default it picks uniformly — every sample an equal chance. A weighted sampler replaces the rule: rare samples get called on more.
Why it exists
Real datasets are lopsided. In a fraud dataset, honest payments outnumber fraud by hundreds to one. Train on it plainly, and the model discovers a lazy trick: predict "honest" always, be right almost always. The rare class — the one you built the system for — is barely rehearsed.
The imbalanced data lesson covers the disease broadly. The sampler is PyTorch's cleanest treatment at the data-feeding level: no data copied, no files duplicated, no dataset edited. The picking rule alone changes, and batches arrive balanced.
How it works
dataset: 950 honest payments, 50 fraud
default sampler: every sample equal chance
→ a batch of 100 holds ~5 fraud cases
weighted sampler: each sample carries a weight,
fraud weighted 19x heavier
→ a batch of 100 holds ~50 fraud casesHeavier samples are drawn more often; the drawing is with replacement, so a rare sample may fairly appear twice in an epoch while a common one sits out.
A real example you have seen
Fraud alerts on UPI apps, spam filters, disease screening — every one of these was trained against brutal imbalance. Rebalancing what the model rehearses, at data-feeding time, is one of the standard fixes running behind all of them.
Remember this
- The sampler chooses which index comes next; replace it to change exposure.
- Weighted sampling makes rare classes rehearsed as often as common ones.
- The data never moves or copies — only the picking rule changes.
What to learn next
- IterableDataset for streams and huge files — what happens when index-picking itself stops being possible.
- Imbalanced data — the full toolbox this lesson is one drawer of.
- Writing a custom loss function — the loss-weighting alternative, and when to prefer it.
Developer — Code and libraries.
Setup
pip install torchWritten and tested against torch 2.5 on CPU. Seeded; exact counts may shift slightly on other versions, the contrast will not.
Before and after, in one script
import torch
from torch.utils.data import DataLoader, TensorDataset, WeightedRandomSampler
torch.manual_seed(0)
# 1000 transactions: 950 normal (label 0), 50 fraud (label 1).
labels = torch.cat([torch.zeros(950, dtype=torch.long),
torch.ones(50, dtype=torch.long)])
features = torch.rand(1000, 4)
ds = TensorDataset(features, labels)
plain = DataLoader(ds, batch_size=100, shuffle=True,
generator=torch.Generator().manual_seed(0))
x, y = next(iter(plain))
print("plain shuffle, one batch: fraud count =", int(y.sum()))
# Give every sample a weight: rare class gets a big one.
class_count = torch.bincount(labels) # tensor([950, 50])
weight_per_class = 1.0 / class_count.float()
sample_weights = weight_per_class[labels] # one weight per sample
sampler = WeightedRandomSampler(sample_weights, num_samples=len(ds), replacement=True,
generator=torch.Generator().manual_seed(0))
balanced = DataLoader(ds, batch_size=100, sampler=sampler)
x, y = next(iter(balanced))
print("weighted sampler, one batch: fraud count =", int(y.sum()))plain shuffle, one batch: fraud count = 6 weighted sampler, one batch: fraud count = 47
Six became forty-seven. The model now rehearses fraud in every batch instead of meeting it as a curiosity.
The walkthrough
The three-line weight recipe is the part to memorise. bincount counts each class; inverting gives per-class weights (rare = heavy); indexing weight_per_class[labels] broadcasts them into one weight per sample. Weights need not sum to one — only their ratios matter.
num_samples=len(ds) defines an "epoch" as 1000 draws. It could be any number: with replacement, an epoch is a budget, not a tour. Expect each fraud sample to appear about ten times per epoch, and some normal samples not at all.
replacement=True is what makes oversampling possible — 500 fraud appearances per epoch from 50 fraud rows. replacement=False with heavy skew would exhaust the rare class and quietly become near-uniform.
sampler= and shuffle= are mutually exclusive. Passing both raises an error, and rightly: shuffle=True is a sampler (a random permutation), as the internals lesson showed. You are replacing exactly that machine.
The rules that keep it honest
Never on validation or test. Weighted sampling is a training-time rehearsal trick. Evaluate on the true, lopsided distribution — the world your model will face. A validation set sampled to balance reports a fantasy accuracy.
Recalibrate if you need probabilities. A model trained on artificially balanced batches learns inflated fraud probabilities. Rankings survive; calibrated probabilities do not. If a downstream decision reads the probability as a number, tune thresholds — on validation data with the true distribution.
One knob at a time. Weighted sampling, class-weighted losses (pos_weight in the loss — the custom loss lesson territory), and data augmentation of the rare class all attack imbalance. Stacking all three over-corrects into the opposite failure. Pick one, measure, then decide.
Common mistakes
Weights from the wrong labels. The weight vector must align index-for-index with the dataset. Shuffle the dataset after computing weights and the weights point at the wrong rows — an especially silent bug.
Forgetting the sampler exists. Months later, someone evaluates on the training loader and cannot reproduce metrics. Name loaders honestly: train_loader_balanced.
Expecting exact 50/50 per batch. Sampling is probabilistic — 47, not 50. If you need guaranteed per-batch class counts, that is a batch_sampler job, built by yielding hand-constructed index lists.
Using it when the imbalance is mild. Below roughly 80/20, the added variance often costs more than the rebalancing earns. Measure against the plain baseline.
Try it yourself
Iterate the full balanced epoch, count total fraud appearances, and compare with the 500 you expect. Then set replacement=False and watch the count collapse toward 50 — the failure mode described above, seen live.
What to learn next
- IterableDataset for streams and huge files — what happens when index-picking itself stops being possible.
- Imbalanced data — the full toolbox this lesson is one drawer of.
- Writing a custom loss function — the loss-weighting alternative, and when to prefer it.
Researcher — Mathematics and papers.
The Sampler abstraction
Sampler is an iterable over indices; BatchSampler wraps any of them into index lists. The stock family: SequentialSampler, RandomSampler (permutation, or i.i.d. with replacement), WeightedRandomSampler, SubsetRandomSampler, and DistributedSampler — which partitions indices across data-parallel ranks and whose set_epoch(e) call is the notorious requirement for reshuffling across epochs in DDP (omit it and every epoch replays epoch 0's permutation, a measured convergence drag). WeightedRandomSampler draws via torch.multinomial over the unnormalised weight vector: $P(i) = w_i / \sum_j w_j$ per draw, i.i.d. with replacement.
Estimator view: sampling as importance reweighting
Uniform sampling makes the empirical risk $\frac{1}{N}\sum_i \ell_i$ an unbiased estimate of risk under the data distribution. Weighted sampling with probabilities $p_i$ instead estimates risk under the tilted distribution $p$ — the deliberate bias. The importance-sampling identity connects the two: sampling by $p_i$ while weighting losses by $1/(N p_i)$ recovers unbiasedness with altered variance. Hence the standard equivalence, exact in expectation: oversampling by $w$ equals loss-weighting by $w$ at equal exposure; they differ in gradient variance — oversampling yields more, smaller, fresher gradient contributions from rare classes; loss weighting yields fewer, larger ones (with correspondingly heavier per-step variance). Buda et al. (2018), A systematic study of the class imbalance problem in CNNs, finds oversampling generally dominant for deep nets and — notably — not to cause the overfitting folklore predicts.
Refinements: class-balanced weighting by effective number of samples $w_c \propto (1-\beta)/(1-\beta^{n_c})$ (Cui et al., 2019), interpolating between inverse-frequency and uniform as $\beta \to 1$; square-root frequency tempering (standard in multilingual and long-tail pretraining); and decoupling representation from classifier — Kang et al. (2020), Decoupling Representation and Classifier for Long-Tailed Recognition, shows uniform sampling can learn the features best while the classifier alone benefits from rebalancing, upending the assumption that balanced batches help end-to-end.
Calibration consequences
Training under tilted priors shifts the learned posterior: with balanced training and true prior $\pi_1$, the model's implicit prior is $1/2$, and recovering calibrated probabilities requires the prior-correction $\log\frac{\pi_1}{1-\pi_1}$ adjustment to logits (equivalently, threshold movement). Rankings (AUC) are invariant; probabilities are not — the formal statement behind the developer block's recalibration warning (Saerens et al., 2002, on prior adjustment).
Interaction with the loader machinery
A custom Sampler composes with everything downstream — workers, collate, pinning — because it lives at the index-generation stage. Length-bucketed batching (collate lesson) is properly a batch_sampler; curriculum schedules (easy-to-hard ordering, Bengio et al., 2009) are samplers whose weights evolve across epochs; hard-example mining is a sampler fed by the model's current losses — the loop closing between training state and data selection.
Reading
- Buda et al. (2018); Cui et al. (2019); Kang et al. (2020) — the imbalance triad.
- Saerens, Latinne, Decaestecker (2002), Adjusting the Outputs of a Classifier to New a Priori Probabilities.
What to learn next
- IterableDataset for streams and huge files — what happens when index-picking itself stops being possible.
- Imbalanced data — the full toolbox this lesson is one drawer of.
- Writing a custom loss function — the loss-weighting alternative, and when to prefer it.