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:

  1. Split data into distinct training and validation sets.
  2. Track the target metric after every epoch.
  3. Compare current performance against the best‑so‑far value.
  4. 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 .ckpt files.
  • Callback integration: High‑level APIs use hooks like Keras ModelCheckpoint or 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_机器学习基础/第二章_机器学习基础.md for early stopping theory, ch18_后端架构选型及应用场景/第十八章_后端架构选型及应用场景.md for checkpoint automation, and ch11_迁移学习/第十一章_迁移学习.md for 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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →