Strategies for Handling Imbalanced Datasets in Deep Learning: A Practical Guide from d2l-zh
The d2l-zh textbook recommends three core strategies for handling imbalanced datasets in deep learning: reweighting the loss function via importance sampling, using focal loss to down-weight easy examples, and resampling the training distribution to balance class frequencies.
Imbalanced data—where minority classes appear far less frequently than majority ones—can bias deep learning models toward dominant classes. The open-source textbook d2l-zh (Dive into Deep Learning, Chinese edition) provides a comprehensive treatment of strategies for handling imbalanced datasets in deep learning, covering both theoretical foundations in chapter_multilayer-perceptrons/environment.md and practical implementations in chapter_computer-vision/ssd.md.
Weighted Loss and Importance Weighting
The most direct method to counteract imbalance is to modify the loss function so that errors on minority classes incur higher penalties.
Mathematical Foundation
As derived in chapter_multilayer-perceptrons/environment.md (lines 14‑22), the standard empirical risk minimization can be extended to weighted empirical risk minimization by introducing importance weights βᵢ. For each sample, the weight is computed as the ratio of the target density to the source density:
βᵢ = p(xᵢ) / q(xᵢ)
The loss becomes:
L = Σᵢ βᵢ · ℓ(f(xᵢ), yᵢ)
This formulation allows the model to pay more attention to under-represented classes by assigning them higher β values.
PyTorch Implementation
import torch
import torch.nn as nn
# Example: minority class 0 (10%), majority class 1 (90%)
class_counts = torch.tensor([0.1, 0.9])
# Inverse frequency weighting
weights = 1.0 / class_counts
criterion = nn.CrossEntropyLoss(weight=weights)
# Training loop
logits = model(inputs)
loss = criterion(logits, targets)
loss.backward()
Focal Loss for Hard Example Mining
When the imbalance is extreme (e.g., 1:1000 ratios), standard weighted cross-entropy may still be dominated by easy background examples. Focal loss addresses this by down-weighting easy examples and focusing training on hard, misclassified samples.
The Focal Loss Formula
As implemented in chapter_computer-vision/ssd.md (lines 23‑27), the focal loss modifies the standard cross-entropy by adding a modulating factor (1 - pₜ)^γ:
FL(pₜ) = -α(1 - pₜ)^γ log(pₜ)
Where:
- pₜ is the model's estimated probability for the ground-truth class
- γ (gamma) is a focusing parameter (typically 2.0)
- α is an optional weighting factor for class balance
TensorFlow/Keras Implementation
import tensorflow as tf
def focal_loss(gamma=2.0, alpha=0.25):
def loss(y_true, y_pred):
y_true = tf.cast(y_true, tf.float32)
epsilon = tf.keras.backend.epsilon()
y_pred = tf.clip_by_value(y_pred, epsilon, 1.0 - epsilon)
# Cross-entropy
ce = -y_true * tf.math.log(y_pred)
# Focal weighting
weight = alpha * tf.pow(1.0 - y_pred, gamma)
return tf.reduce_mean(weight * ce)
return loss
# Usage
model.compile(optimizer='adam',
loss=focal_loss(gamma=2.0, alpha=0.25),
metrics=['accuracy'])
Resampling and Importance Sampling
Rather than modifying the loss function, you can alter the sampling distribution to balance class frequencies before they reach the model.
WeightedRandomSampler in PyTorch
The WeightedRandomSampler implements importance sampling by drawing examples with probability proportional to their importance weights:
from torch.utils.data import WeightedRandomSampler, DataLoader
# Calculate weights: inverse of class frequency
class_sample_count = torch.bincount(targets)
weight_per_class = 1.0 / class_sample_count.float()
samples_weight = weight_per_class[targets]
# Create sampler with replacement=True for oversampling minority classes
sampler = WeightedRandomSampler(
weights=samples_weight,
num_samples=len(samples_weight),
replacement=True
)
loader = DataLoader(dataset, batch_size=64, sampler=sampler)
Importance-Weighted Empirical Risk
For covariate shift correction (where training and test distributions differ), chapter_multilayer-perceptrons/environment.md derives the importance weight as the density ratio:
import numpy as np
# p(x): target distribution (test set)
# q(x): source distribution (training set)
# Estimated via kernel density estimation or histogram matching
beta = p_x / q_x # importance weights
# Apply to empirical risk
per_sample_loss = loss_fn(predictions, y_true)
weighted_risk = np.mean(beta * per_sample_loss)
Data Augmentation and Preprocessing
While not exclusive to imbalance, data augmentation effectively increases minority class samples. As discussed in chapter_preliminaries/pandas.md (lines 44‑48) regarding data preprocessing, cleaning and augmentation pipelines should be applied strategically to under-represented classes.
Common techniques include:
- Image: Random rotations, flips, color jittering, and mixup/cutmix
- Text: Synonym replacement, back-translation, and random insertion/deletion
- Tabular: SMOTE (Synthetic Minority Over-sampling Technique) and Gaussian noise injection
Evaluation Metrics and Threshold Tuning
When classes are imbalanced, accuracy becomes misleading. The d2l-zh text emphasizes using precision, recall, F1-score, and ROC-AUC to evaluate model performance on minority classes.
Threshold tuning involves adjusting the decision boundary after training to optimize the F1-score or to achieve a specific precision-recall trade-off, rather than using the default 0.5 threshold for binary classification.
Summary
- Weighted loss functions (importance weighting) modify the learning objective to penalize minority class errors more heavily, as formalized in
chapter_multilayer-perceptrons/environment.md. - Focal loss focuses training on hard examples by down-weighting easy majority class predictions, with a concrete implementation shown in
chapter_computer-vision/ssd.md. - Resampling strategies (oversampling minorities via
WeightedRandomSampleror undersampling majorities) alter the training distribution before optimization begins. - Importance sampling corrects for distribution shift using density ratios βᵢ = p(x)/q(x), derived in the environment chapter for covariate shift scenarios.
- Data augmentation and threshold tuning provide additional practical levers for improving minority class recall without architectural changes.
Frequently Asked Questions
What is the most effective strategy for handling imbalanced datasets in deep learning?
There is no universal best strategy; the optimal approach depends on the severity of imbalance and data availability. According to the d2l-zh source code, weighted loss functions and focal loss are generally preferred for moderate imbalance because they preserve all data while adjusting the learning signal. For extreme imbalance (e.g., 1:1000 ratios), combining resampling with focal loss often yields the best results, as demonstrated in the SSD object detection implementation.
How does focal loss differ from weighted cross-entropy?
While weighted cross-entropy applies a static class-specific weight to all samples of a given class, focal loss dynamically scales the loss based on how easily the model classifies each example. As implemented in chapter_computer-vision/ssd.md, focal loss adds a modulating factor (1 - pₜ)^γ that automatically down-weights easy examples (where pₜ is high) and focuses training on hard, misclassified samples, making it particularly effective for dense prediction tasks with extreme background-to-foreground imbalance.
When should I use resampling versus weighted loss?
Use resampling (oversampling minorities or undersampling majorities) when you want to alter the empirical distribution seen by the optimizer, which can help when the absolute number of minority samples is sufficient but their relative frequency is too low. Use weighted loss when you want to preserve the original data distribution but adjust the optimization objective; this is mathematically equivalent to resampling in expectation but more memory-efficient and stable for online learning. The d2l-zh text in chapter_multilayer-perceptrons/environment.md treats both as instances of importance weighting.
Where does d2l-zh discuss these strategies in the codebase?
The theoretical foundation for importance weighting appears in chapter_multilayer-perceptrons/environment.md (lines 14‑22), which derives the βᵢ weights for covariate shift correction. The practical application of focal loss is implemented in chapter_computer-vision/ssd.md (lines 23‑27) for handling class imbalance in object detection. The introductory motivation for these techniques is found in chapter_introduction/index.md (line 138), which identifies imbalanced datasets as a common problem in machine learning. Data preprocessing context appears in chapter_preliminaries/pandas.md (lines 44‑48).
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →