How to Implement Early Stopping, Checkpointing, and Model Ensembling in Deep Learning
Early stopping halts training when validation metrics plateau, checkpointing preserves the best model weights, and ensembling aggregates predictions from multiple checkpoints to boost accuracy and robustness.
The DeepLearning‑500‑questions repository provides a comprehensive knowledge base that explains how to implement early stopping, checkpointing, and model ensembling as integrated components of production‑grade training pipelines. These three techniques work together to prevent over‑fitting, ensure training resiliency, and improve final model performance.
Early Stopping to Prevent Over‑Fitting
Early stopping monitors a validation metric—typically loss or accuracy—and terminates training when improvement stalls for a predefined number of epochs. According to the source code analysis of ch02_机器学习基础/第二章_机器学习基础.md, early stopping is listed as a classic remedy for over‑fitting around line 1303‑1305.
The mechanism follows four discrete steps:
- Split data into distinct training and validation sets.
- Track the target metric after every epoch.
- Compare current performance against the best‑so‑far value.
- Count consecutive epochs without improvement; when the count exceeds the patience threshold, restore the best weights and exit.
This approach eliminates unnecessary computation and automatically selects the epoch with optimal generalization.
Model Checkpointing for State Recovery
Checkpointing saves model parameters and optimizer state at regular intervals or whenever a metric improves. The repository highlights TensorFlow’s MonitoredTrainingSession in ch18_后端架构选型及应用场景/第十八章_后端架构选型及应用场景.md (lines 188‑189) as a built‑in mechanism that automates checkpoint creation and restoration.
Key implementation aspects include:
- Directory management: A dedicated folder (e.g.,
/tmp/train_logs) stores.ckptfiles. - Callback integration: High‑level APIs use hooks like Keras
ModelCheckpointor PyTorch Lightning equivalents. - Frequency control: Save every epoch, every n steps, or only when the monitored metric improves (best‑only mode).
Checkpointing enables crash recovery and provides the exact weights to reload after early stopping triggers.
Model Ensembling Strategies
Ensembling combines predictions from multiple independently trained models to reduce variance and improve robustness. The repository references "hedge ensemble" as an online heterogeneous ensemble example in ch11_迁移学习/第十一章_迁移学习.md around line 637‑639.
Common architectural patterns include:
- Bagging: Train several models on different bootstrap samples or random seeds, then average predictions (soft‑voting) or take majority votes (hard‑voting).
- Boosting: Sequentially train models where each focuses on correcting the predecessor’s errors (e.g., AdaBoost, Gradient Boosting).
- Stacking: Train a meta‑learner on the concatenated outputs of base models.
The standard workflow integrates all three concepts: use early stopping to find optimal epochs, checkpoint those states, then load multiple checkpoints for ensemble inference.
TensorFlow 2 and Keras Implementation
Below is a complete implementation using Keras callbacks for early stopping and checkpointing, followed by a bagging ensemble.
Training with Callbacks
import tensorflow as tf
from tensorflow.keras import layers, models, callbacks
# Define architecture
model = models.Sequential([
layers.Dense(64, activation='relu', input_shape=(input_dim,)),
layers.Dropout(0.5),
layers.Dense(num_classes, activation='softmax')
])
model.compile(optimizer='adam',
loss='sparse_categorical_crossentropy',
metrics=['accuracy'])
# Early stopping: restore best weights automatically
early_stop = callbacks.EarlyStopping(
monitor='val_loss',
patience=5,
restore_best_weights=True
)
# Checkpoint: save only when validation loss improves
ckpt_path = "checkpoints/epoch-{epoch:02d}-valLoss-{val_loss:.3f}.ckpt"
model_ckpt = callbacks.ModelCheckpoint(
filepath=ckpt_path,
monitor='val_loss',
save_best_only=True,
save_weights_only=False
)
# Execute training
history = model.fit(
train_ds,
epochs=100,
validation_data=val_ds,
callbacks=[early_stop, model_ckpt]
)
Bagging Ensemble in Keras
import numpy as np
def build_and_train(seed):
tf.random.set_seed(seed)
model = models.Sequential([...]) # same architecture
model.compile(optimizer='adam', loss='sparse_categorical_crossentropy')
model.fit(train_ds, validation_data=val_ds,
callbacks=[early_stop, model_ckpt], epochs=100)
return model
# Train multiple instances with different seeds
seeds = [0, 42, 123]
models = [build_and_train(s) for s in seeds]
def ensemble_predict(x):
preds = np.stack([m.predict(x) for m in models], axis=0)
return np.mean(preds, axis=0) # soft-voting
PyTorch Implementation
PyTorch requires manual implementation of early stopping logic, but offers explicit control over checkpoint serialization.
Custom Early Stopping and Checkpointing
import torch
import torch.nn as nn
import torch.optim as optim
class Net(nn.Module):
def __init__(self, input_dim, num_classes):
super().__init__()
self.fc1 = nn.Linear(input_dim, 64)
self.dropout = nn.Dropout(0.5)
self.fc2 = nn.Linear(64, num_classes)
def forward(self, x):
x = torch.relu(self.fc1(x))
x = self.dropout(x)
return torch.softmax(self.fc2(x), dim=1)
class EarlyStopping:
def __init__(self, patience=5, delta=0):
self.patience = patience
self.delta = delta
self.best_loss = None
self.counter = 0
self.best_state = None
def step(self, val_loss, model):
if (self.best_loss is None or
val_loss < self.best_loss - self.delta):
self.best_loss = val_loss
self.counter = 0
# Deep copy state to CPU to avoid GPU memory bloat
self.best_state = {k: v.cpu() for k, v in model.state_dict().items()}
torch.save(self.best_state, "best_checkpoint.pt")
return False
else:
self.counter += 1
return self.counter >= self.patience
# Training loop
model = Net(input_dim, num_classes)
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=1e-3)
early_stopper = EarlyStopping(patience=5)
for epoch in range(100):
model.train()
for xb, yb in train_loader:
optimizer.zero_grad()
loss = criterion(model(xb), yb)
loss.backward()
optimizer.step()
# Validation phase
model.eval()
val_loss = 0.0
with torch.no_grad():
for xb, yb in val_loader:
val_loss += criterion(model(xb), yb).item()
val_loss /= len(val_loader)
if early_stopper.step(val_loss, model):
print(f"Early stopping triggered at epoch {epoch}")
break
Ensembling in PyTorch
def train_one(seed):
torch.manual_seed(seed)
model = Net(input_dim, num_classes)
optimizer = optim.Adam(model.parameters(), lr=1e-3)
# ... training with early stopping as above ...
torch.save(model.state_dict(), f"model_seed{seed}.pt")
return model
seeds = [0, 42, 123]
models = [train_one(s) for s in seeds]
def ensemble_predict(x):
with torch.no_grad():
preds = torch.stack([m(x) for m in models])
return preds.mean(dim=0) # soft-voting
Summary
- Early stopping monitors validation metrics and terminates training when improvement stagnates beyond a patience threshold, reverting to the best weights.
- Checkpointing persists model parameters at optimal states, enabling crash recovery and serving as the bridge between early stopping and final model selection.
- Ensembling aggregates predictions from multiple checkpoints—often generated via different random seeds or data augmentations—to improve generalization and robustness.
- The DeepLearning‑500‑questions repository cites
ch02_机器学习基础/第二章_机器学习基础.mdfor early stopping theory,ch18_后端架构选型及应用场景/第十八章_后端架构选型及应用场景.mdfor checkpoint automation, andch11_迁移学习/第十一章_迁移学习.mdfor ensemble strategies.
Frequently Asked Questions
What is the difference between early stopping and checkpointing?
Early stopping is a training control mechanism that decides when to terminate training, while checkpointing is a persistence mechanism that saves model states. In practice, they operate together: early stopping determines which epoch represents the best model, and checkpointing preserves that specific state to disk so it can be loaded later.
How many models should I include in an ensemble?
Three to five models typically provide the best accuracy‑to‑cost ratio. According to the bagging implementation shown in the code examples, training three models with seeds [0, 42, 123] and averaging their predictions (soft‑voting) usually yields significant error reduction without excessive inference latency.
Does early stopping guarantee the best validation performance?
Early stopping approximates the best performance within the defined patience window. By setting restore_best_weights=True in Keras or manually reloading the best checkpoint in PyTorch, you guarantee that the model retains the weights from the epoch with the lowest validation loss observed during training, not merely the final epoch.
Can I ensemble models trained with different architectures?
Yes. The repository’s reference to "hedge ensemble" in ch11_迁移学习/第十一章_迁移学习.md suggests that heterogeneous ensembles—combining different architectures—can improve robustness. However, ensure all models share the same input and output tensor shapes so predictions can be aggregated via averaging or stacking.
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 →