PyTorch Model Evaluation and Debugging

Training and optimizing deep learning models require systematic evaluation and debugging methods. This section introduces how to analyze Loss curves, identify overfitting and underfitting, use confusion matrices to evaluate classification models, and common debugging techniques.


1. Loss Function Analysis

The loss function is a core metric during training, reflecting the gap between model predictions and true values. Correctly analyzing the Loss curve is key to debugging models.

1.1 Comparison of Training and Validation Loss

By comparing training loss and validation loss, the state of the model can be determined:

  • Training loss decreases, validation loss decreases: Normal, indicating the model is learning
  • Training loss decreases, validation loss stable or increases: Overfitting, regularization needed
  • Training loss does not decrease: Learning rate issue or model architecture issue
  • Training loss oscillates: Learning rate too large or batch size too small

1.2 Loss Curve Visualization

Example

import matplotlib.pyplot as plt
import numpy as np

def plot_training_history(history):
    """
Plot training history
history: a dict containing 'train_loss', 'val_loss', 'train_acc', 'val_acc'
    """

    fig, axes = plt.subplots(1, 2, figsize=(14, 5))

    # Loss curve
    axes[0].plot(history['train_loss'], label='Train Loss', alpha=0.8)
    axes[0].plot(history['val_loss'], label='Val Loss', alpha=0.8)
    axes[0].set_xlabel('Epoch')
    axes[0].set_ylabel('Loss')
    axes[0].set_title('Training and Validation Loss')
    axes[0].legend()
    axes[0].grid(True, alpha=0.3)

    # Accuracy curve
    axes[1].plot(history['train_acc'], label='Train Acc', alpha=0.8)
    axes[1].plot(history['val_acc'], label='Val Acc', alpha=0.8)
    axes[1].set_xlabel('Epoch')
    axes[1].set_ylabel('Accuracy')
    axes[1].set_title('Training and Validation Accuracy')
    axes[1].legend()
    axes[1].grid(True, alpha=0.3)

    plt.tight_layout()
    plt.show()


# Simulate training data
history = {
    'train_loss': np.linspace(2.0, 0.2, 50) + np.random.normal(0, 0.05, 50),
    'val_loss': np.linspace(2.0, 0.3, 50) + np.random.normal(0, 0.08, 50),
    'train_acc': np.linspace(0.3, 0.95, 50),
    'val_acc': np.linspace(0.3, 0.88, 50),
}

plot_training_history(history)

1.3 Identifying Overfitting

When validation loss starts to rise while training loss continues to decline, it is the beginning of overfitting. At this point the model starts "memorizing" training data instead of learning general patterns.

Example

import matplotlib.pyplot as plt
import numpy as np

def detect_overfitting(history, patience=5):
    """
Detect overfitting
patience: number of consecutive epochs with validation loss increasing before stopping training
    """

    val_loss = history['val_loss']
    best_loss = float('inf')
    best_epoch = 0
    early_stop_epoch = None

    for epoch in range(patience, len(val_loss)):
        # Check the most recent patience epochs
        recent_losses = val_loss[epoch - patience:epoch]
        current_loss = val_loss[epoch]

        # If the current loss is higher than the most recent minimum loss
        if min(recent_losses) < current_loss:
            early_stop_epoch = epoch
            break

        if current_loss < best_loss:
            best_loss = current_loss
            best_epoch = epoch

    if early_stop_epoch:
        print(f"Overfitting detected! Best epoch: {best_epoch}, should stop at epoch {early_stop_epoch}")
    else:
        print("No overfitting detected")

    return early_stop_epoch, best_epoch


# Overfitting example
overfitting_history = {
    'train_loss': np.linspace(1.5, 0.1, 30),
    'val_loss': np.concatenate([np.linspace(1.5, 0.3, 15), np.linspace(0.3, 0.8, 15)]),
}

detect_overfitting(overfitting_history)

2. Handling Overfitting and Underfitting

2.1 Overfitting Solutions

Method Description Code Implementation
Increase data volume Collect more training data Data augmentation
Dropout Randomly drop neurons nn.Dropout(0.5)
Weight decay L2 regularization weight_decay=1e-4
Early Stopping Stop training when validation loss stops decreasing patience parameter
Reduce model complexity Reduce number of layers or neurons Reduce hidden_size
Data augmentation Random transformations to increase data diversity Rotation, flipping, cropping

2.2 Underfitting Solutions

Method Description
Increase model complexity Increase number of layers or neurons
Train longer Increase number of training epochs
Reduce regularization Lower Dropout or weight_decay
Adjust learning rate Try a larger learning rate

2.3 Implementation of Early Stopping Mechanism

Example

import torch
import torch.nn as nn
import torch.optim as optim

class EarlyStopping:
    """
Early stopping mechanism: stop training when validation loss does not decrease for consecutive epochs
    """

    def __init__(self, patience=7, min_delta=0, mode='min'):
        """
patience: stop after how many consecutive epochs without improvement
min_delta: the minimum change considered as an improvement
mode: 'min' or 'max', whether the metric should be minimized or maximized
        """

        self.patience = patience
        self.min_delta = min_delta
        self.mode = mode
        self.counter = 0
        self.best_score = None
        self.early_stop = False

    def __call__(self, val_loss):
        score = -val_loss if self.mode == 'min' else val_loss

        if self.best_score is None:
            self.best_score = score
        elif score < self.best_score + self.min_delta:
            self.counter += 1
            if self.counter >= self.patience:
                self.early_stop = True
        else:
            self.best_score = score
            self.counter = 0

        return self.early_stop


# Use early stopping
def train_with_early_stopping(model, train_loader, val_loader, patience=7):
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    model = model.to(device)

    criterion = nn.CrossEntropyLoss()
    optimizer = optim.Adam(model.parameters(), lr=1e-3)
    early_stopping = EarlyStopping(patience=patience, mode='min')

    best_val_loss = float('inf')
    best_model_state = None

    for epoch in range(100):
        # Train
        model.train()
        train_loss = 0
        for inputs, labels in train_loader:
            inputs, labels = inputs.to(device), labels.to(device)
            optimizer.zero_grad()
            outputs = model(inputs)
            loss = criterion(outputs, labels)
            loss.backward()
            optimizer.step()
            train_loss += loss.item()

        # Validate
        model.eval()
        val_loss = 0
        with torch.no_grad():
            for inputs, labels in val_loader:
                inputs, labels = inputs.to(device), labels.to(device)
                outputs = model(inputs)
                loss = criterion(outputs, labels)
                val_loss += loss.item()

        val_loss /= len(val_loader)

        print(f"Epoch {epoch+1}: Train Loss={train_loss:.4f}, Val Loss={val_loss:.4f}")

        # Early stopping check
        if early_stopping(val_loss):
            print(f"Early stopping triggered! Validation loss has not decreased for {patience} consecutive epochs")

        # Save best model
        if val_loss < best_val_loss:
            best_val_loss = val_loss
            best_model_state = model.state_dict().copy()

    # Restore best model
    model.load_state_dict(best_model_state)
    return model

3. Confusion Matrix

The confusion matrix is an important tool for evaluating classification models, showing the relationship between model predictions and true labels.

3.1 Confusion Matrix Calculation

Example

import torch
import numpy as np
from sklearn.metrics import confusion_matrix
import matplotlib.pyplot as plt
import seaborn as sns

def calculate_confusion_matrix(model, dataloader, device, num_classes):
    """
Compute confusion matrix
    """

    model.eval()
    all_preds = []
    all_labels = []

    with torch.no_grad():
        for inputs, labels in dataloader:
            inputs = inputs.to(device)
            outputs = model(inputs)
            _, preds = torch.max(outputs, 1)

            all_preds.extend(preds.cpu().numpy())
            all_labels.extend(labels.numpy())

    cm = confusion_matrix(all_labels, all_preds)
    return cm


def plot_confusion_matrix(cm, class_names):
    """
Plot confusion matrix
    """

    plt.figure(figsize=(10, 8))
    sns.heatmap(cm, annot=True, fmt='d', cmap='Blues',
                xticklabels=class_names, yticklabels=class_names)
    plt.xlabel('Predicted Label')
    plt.ylabel('True Label')
    plt.title('Confusion Matrix')
    plt.tight_layout()
    plt.show()


# Usage example
# Assume there are 10 classes
class_names = [f'Class {i}' for i in range(10)]
cm = np.array([
    [45, 1, 0, 0, 0, 0, 2, 1, 1, 0],
    [0, 42, 2, 0, 0, 0, 0, 3, 2, 1],
    [1, 1, 38, 3, 0, 0, 0, 2, 3, 2],
    [0, 0, 1, 41, 2, 1, 0, 1, 1, 3],
    [0, 0, 0, 1, 44, 0, 1, 0, 2, 2],
    [1, 0, 0, 0, 0, 43, 2, 1, 1, 2],
    [2, 1, 0, 0, 1, 1, 39, 2, 2, 2],
    [0, 2, 1, 1, 0, 0, 1, 40, 3, 2],
    [1, 2, 2, 0, 1, 1, 1, 2, 36, 4],
    [0, 1, 1, 2, 2, 2, 2, 1, 2, 37],
])

plot_confusion_matrix(cm, class_names)

3.2 Classification Evaluation Metrics

Multiple evaluation metrics can be computed from the confusion matrix:

Example

import numpy as np

def classification_metrics(cm):
    """
Compute various evaluation metrics from the confusion matrix
    """

    # Accuracy
    accuracy = np.trace(cm) / np.sum(cm)

    # Compute precision, recall, F1 for each class
    num_classes = cm.shape[0]
    precisions = []
    recalls = []
    f1s = []

    for i in range(num_classes):
        tp = cm[i, i]
        fp = np.sum(cm[:, i]) - tp
        fn = np.sum(cm[i, :]) - tp

        precision = tp / (tp + fp) if (tp + fp) > 0 else 0
        recall = tp / (tp + fn) if (tp + fn) > 0 else 0
        f1 = 2 * precision * recall / (precision + recall) if (precision + recall) > 0 else 0

        precisions.append(precision)
        recalls.append(recall)
        f1s.append(f1)

    # Macro Average
    macro_precision = np.mean(precisions)
    macro_recall = np.mean(recalls)
    macro_f1 = np.mean(f1s)

    # Weighted Average
    class_counts = np.sum(cm, axis=1)
    weights = class_counts / np.sum(class_counts)
    weighted_precision = np.average(precisions, weights=weights)
    weighted_recall = np.average(recalls, weights=weights)
    weighted_f1 = np.average(f1s, weights=weights)

    return {
        'accuracy': accuracy,
        'macro_precision': macro_precision,
        'macro_recall': macro_recall,
        'macro_f1': macro_f1,
        'weighted_precision': weighted_precision,
        'weighted_recall': weighted_recall,
        'weighted_f1': weighted_f1,
    }


# Compute metrics
metrics = classification_metrics(cm)

print("=" * 50)
print("Classification evaluation metrics")
print("=" * 50)
print(f"Accuracy: {metrics['accuracy']:.4f}")
print(f"Macro Precision: {metrics['macro_precision']:.4f}")
print(f"Macro Recall: {metrics['macro_recall']:.4f}")
print(f"Macro F1: {metrics['macro_f1']:.4f}")
print(f"Weighted Precision: {metrics['weighted_precision']:.4f}")
print(f"Weighted Recall: {metrics['weighted_recall']:.4f}")
print(f"Weighted F1: {metrics['weighted_f1']:.4f}")

3.3 Classification Report

Example

from sklearn.metrics import classification_report

# Use sklearn's classification report
y_true = [0, 1, 2, 3, 4, 5, 6, 7, 8, 9] * 10  # Simulate true labels
y_pred = [0, 1, 2, 3, 4, 5, 6, 7, 8, 8] * 10  # Simulate predictions

report = classification_report(y_true, y_pred, digits=4)
print(report)

4. Learning Rate Scheduling

Learning rate is one of the most important hyperparameters for training deep learning models. Appropriate learning rate scheduling can significantly improve training effectiveness.

4.1 Learning Rate Scheduling Strategies

Strategy Description Applicable scenarios
StepLR Fixed step decay When the optimal learning rate is known
MultiStepLR Specified epoch decay Non-uniform decay
CosineAnnealingLR Cosine annealing Smooth convergence
ReduceLROnPlateau Decay when validation loss stops decreasing Automatic tuning
Warmup Increase first, then decrease Stabilize early training

4.2 Learning Rate Scheduling Implementation

Example

import torch
import torch.optim as optim
import matplotlib.pyplot as plt

# Simulate training epochs
epochs = 50

# 1. Step LR: decay by half every 10 epochs
scheduler_step = optim.lr_scheduler.StepLR(
    optim.SGD(torch.nn.Linear(10, 10).parameters(), lr=0.1),
    step_size=10, gamma=0.5
)

# 2. MultiStep LR: decay at specified epochs
scheduler_multistep = optim.lr_scheduler.MultiStepLR(
    optim.SGD(torch.nn.Linear(10, 10).parameters(), lr=0.1),
    milestones=[15, 30, 45], gamma=0.1
)

# 3. Cosine Annealing
scheduler_cosine = optim.lr_scheduler.CosineAnnealingLR(
    optim.SGD(torch.nn.Linear(10, 10).parameters(), lr=0.1),
    T_max=50
)

# 4. ReduceLROnPlateau
scheduler_plateau = optim.lr_scheduler.ReduceLROnPlateau(
    optim.SGD(torch.nn.Linear(10, 10).parameters(), lr=0.1),
    mode='min', factor=0.5, patience=5
)

# Plot learning rate curves
fig, axes = plt.subplots(2, 2, figsize=(14, 10))

# Step LR
lr_history = []
for _ in range(epochs):
    lr_history.append(optimizer.param_groups[0]['lr'])
    scheduler_step.step()
axes[0, 0].plot(lr_history)
axes[0, 0].set_title('Step LR')

# Reinitialize
optimizer = optim.SGD(torch.nn.Linear(10, 10).parameters(), lr=0.1)
scheduler = optim.lr_scheduler.MultiStepLR(optimizer, milestones=[15, 30, 45], gamma=0.1)
lr_history = []
for _ in range(epochs):
    lr_history.append(optimizer.param_groups[0]['lr'])
    scheduler.step()
axes[0, 1].plot(lr_history)
axes[0, 1].set_title('MultiStep LR')

# Cosine
optimizer = optim.SGD(torch.nn.Linear(10, 10).parameters(), lr=0.1)
scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=50)
lr_history = []
for _ in range(epochs):
    lr_history.append(optimizer.param_groups[0]['lr'])
    scheduler.step()
axes[1, 0].plot(lr_history)
axes[1, 0].set_title('Cosine Annealing LR')

# Warmup + Cosine
def warmup_cosine(optimizer, warmup_epochs, total_epochs, min_lr=1e-6):
    def lr_lambda(epoch):
        if epoch < warmup_epochs:
            return epoch / warmup_epochs
        return min_lr + 0.5 * (1 - min_lr) * (1 + np.cos(np.pi * (epoch - warmup_epochs) / (total_epochs - warmup_epochs)))

    return optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)

optimizer = optim.SGD(torch.nn.Linear(10, 10).parameters(), lr=0.1)
scheduler = warmup_cosine(optimizer, warmup_epochs=5, total_epochs=50)
lr_history = []
for _ in range(epochs):
    lr_history.append(optimizer.param_groups[0]['lr'])
    scheduler.step()
axes[1, 1].plot(lr_history)
axes[1, 1].set_title('Warmup + Cosine')

for ax in axes.flat:
    ax.set_xlabel('Epoch')
    ax.set_ylabel('Learning Rate')
    ax.grid(True, alpha=0.3)

plt.tight_layout()
plt.show()

5. Gradient Problem Diagnosis

5.1 Gradient Vanishing and Explosion

< h2 class="example">Example
import torch
import torch.nn as nn

def analyze_gradients(model):
    """
Analyze the gradients of each layer in the model
    """

    grad_stats = {}

    for name, param in model.named_parameters():
        if param.grad is not None:
            grad = param.grad
            grad_stats[name] = {
                'mean': grad.mean().item(),
                'std': grad.std().item(),
                'max': grad.abs().max().item(),
                'min': grad.abs().min().item(),
                'norm': grad.norm().item(),
            }

    return grad_stats


def detect_gradient_issues(model):
    """
Detect gradient issues
    """

    issues = []
    grad_norms = []

    for param in model.parameters():
        if param.grad is not None:
            grad_norms.append(param.grad.norm().item())

    avg_grad_norm = sum(grad_norms) / len(grad_norms) if grad_norms else 0

    if avg_grad_norm < 1e-7:
        issues.append("Gradient vanishing: gradients are too small, the model may be unable to learn")
    elif avg_grad_norm > 100:
        issues.append("Gradient explosion: gradients are too large, training is unstable")

    return issues, avg_grad_norm


# Example: detecting gradients
model = nn.Sequential(
    nn.Linear(100, 50),
    nn.ReLU(),
    nn.Linear(50, 50),
    nn.ReLU(),
    nn.Linear(50, 10)
)

# Simulate forward and backward propagation
x = torch.randn(32, 100)
y = model(x)
loss = y.sum()
loss.backward()

issues, avg_norm = detect_gradient_issues(model)
print(f"Average gradient norm: {avg_norm:.6f}")
if issues:
    for issue in issues:
        print(f"Warning: {issue}")
else:
    print("Gradient status is normal")

5.2 Gradient Clipping

Example

import torch
import torch.nn as nn
import torch.optim as optim
import torch.nn.utils as utils

# Gradient clipping example
def train_with_gradient_clipping(model, dataloader, max_norm=1.0):
    """
Train with gradient clipping
max_norm: maximum gradient norm; gradients exceeding this value will be clipped
    """

    criterion = nn.CrossEntropyLoss()
    optimizer = optim.Adam(model.parameters(), lr=1e-3)

    model.train()
    for inputs, labels in dataloader:
        optimizer.zero_grad()

        outputs = model(inputs)
        loss = criterion(outputs, labels)

        loss.backward()

        # Gradient clipping: prevent gradient explosion
        utils.clip_grad_norm_(model.parameters(), max_norm=max_norm)

        optimizer.step()

    return model


# Element-wise clipping (more conservative)
def clip_grad_by_value(model, clip_value=1.0):
    """
Clip gradients by value
    """

    for param in model.parameters():
        if param.grad is not None:
            param.grad.data.clamp_(min=-clip_value, max=clip_value)

6. Model Performance Analysis

6.1 Computing Parameter Count and FLOPs

Example

import torch
import torch.nn as nn
from thop import profile

def count_parameters(model):
    """Calculate the number of trainable parameters in the model"""
    total = sum(p.numel() for p in model.parameters())
    trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
    return total, trainable


def count_flops(model, input_size=(1, 3, 224, 224)):
    """Calculate FLOPs (requires the thop library)"""
    input_tensor = torch.randn(input_size)
    flops, params = profile(model, inputs=(input_tensor,), verbose=False)
    return flops, params


# Example: analyze the model
model = nn.Sequential(
    nn.Conv2d(3, 64, kernel_size=3, padding=1),
    nn.ReLU(),
    nn.MaxPool2d(2),
    nn.Conv2d(64, 128, kernel_size=3, padding=1),
    nn.ReLU(),
    nn.MaxPool2d(2),
    nn.Conv2d(128, 256, kernel_size=3, padding=1),
    nn.ReLU(),
    nn.AdaptiveAvgPool2d(1),
    nn.Flatten(),
    nn.Linear(256, 10)
)

total, trainable = count_parameters(model)
print(f"Total parameters: {total:,}")
print(f"Trainable parameters: {trainable:,}")
print(f"Model size: {total * 4 / 1024 / 1024:.2f} MB")

# FLOPs calculation
try:
    flops, _ = count_flops(model)
    print(f"FLOPs: {flops / 1e9:.2f} G")
except ImportError:
    print("Please install the thop library: pip install thop")

6.2 Inference Speed Testing

< h2 class="example">Example
import torch
import time

def measure_inference_speed(model, input_size, device='cuda', num_iterations=100):
    """
Measure inference speed
    """

    model.eval()
    model = model.to(device)

    # Warm-up
    dummy_input = torch.randn(input_size).to(device)
    with torch.no_grad():
        for _ in range(10):
            _ = model(dummy_input)

    # Timing
    start = time.time()
    with torch.no_grad():
        for _ in range(num_iterations):
            _ = model(dummy_input)

    if device == 'cuda':
        torch.cuda.synchronize()

    end = time.time()
    avg_time = (end - start) / num_iterations * 1000  # ms

    return avg_time


# Usage example
model = nn.Sequential(
    nn.Conv2d(3, 64, kernel_size=3, padding=1),
    nn.ReLU(),
    nn.Conv2d(64, 64, kernel_size=3, padding=1),
    nn.ReLU(),
    nn.AdaptiveAvgPool2d(1),
    nn.Flatten(),
    nn.Linear(64, 10)
)

# CPU inference speed
cpu_time = measure_inference_speed(model, (1, 3, 224, 224), device='cpu')
print(f"CPU inference time: {cpu_time:.2f} ms/image")

# GPU inference speed (if CUDA is available)
if torch.cuda.is_available():
    gpu_time = measure_inference_speed(model, (1, 3, 224, 224), device='cuda')
    print(f"GPU inference time: {gpu_time:.2f} ms/image")

6.3 Memory Usage Analysis

Example

import torch

def analyze_memory(model, input_size):
    """
Analyze model memory usage
    """

    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    model = model.to(device)

    # Create input
    x = torch.randn(input_size).to(device)

    # Clear cache
    if torch.cuda.is_available():
        torch.cuda.empty_cache()
        torch.cuda.reset_peak_memory_stats()

    # Forward propagation
    with torch.no_grad():
        output = model(x)

    # Get memory statistics
    if torch.cuda.is_available():
        allocated = torch.cuda.memory_allocated() / 1024**2  # MB
        reserved = torch.cuda.memory_reserved() / 1024**2
        max_allocated = torch.cuda.max_memory_allocated() / 1024**2

        print(f"Currently allocated: {allocated:.2f} MB")
        print(f"Reserved memory: {reserved:.2f} MB")
        print(f"Peak usage: {max_allocated:.2f} MB")
    else:
        print("CUDA support required")

    return output


# Usage example
model = nn.Sequential(
    nn.Conv2d(3, 64, kernel_size=3, padding=1),
    nn.ReLU(),
    nn.Conv2d(64, 128, kernel_size=3, padding=1),
    nn.ReLU(),
    nn.Conv2d(128, 256, kernel_size=3, padding=1),
    nn.ReLU(),
)
analyze_memory(model, (8, 3, 224, 224))

7. Common Problems and Debugging Tips

7.1 Quick Reference for Training Problems

Symptom Possible cause Solution
Loss not decreasing Learning rate too small/large, gradient issues Adjust learning rate, check gradients
Loss oscillation Learning rate too large, small batch size Decrease learning rate, increase batch size
NaN appears Division by zero, log(0), gradient explosion Add epsilon, gradient clipping
Overfitting Complex model, insufficient data Add more data, add regularization
Underfitting Model too simple, insufficient training Increase model size, train longer
Insufficient GPU memory Large batch size, large model Reduce batch size, gradient accumulation

7.2 Collection of Debugging Tips

Example

import torch
import random
import numpy as np

def set_seed(seed=42):
    """Set the random seed to ensure reproducible results"""
    random.seed(seed)
    np.random.seed(seed)
    torch.manual_seed(seed)
    torch.cuda.manual_seed(seed)
    torch.cuda.manual_seed_all(seed)
    torch.backends.cudnn.deterministic = True
    torch.backends.cudnn.benchmark = False


def debug_forward(model, x):
    """Debug forward propagation and check the outputs of each layer"""
    print("=" * 50)
    print("Forward propagation debugging")
    print("=" * 50)

    for name, layer in model.named_children():
        x = layer(x)
        if hasattr(x, 'shape'):
            print(f"{name}: shape={x.shape}, ", end='')
            if hasattr(x, 'mean'):
                print(f"mean={x.mean().item():.4f}, std={x.std().item():.4f}")
            if torch.isnan(x).any():
                print(f" &#x26a0;&#xfe0f; NaN detected!")
            if torch.isinf(x).any():
                print(f" &#x26a0;&#xfe0f; Inf detected!")
    print("=" * 50)


def debug_backward(model):
    """Debug backward propagation and check gradients"""
    print("=" * 50)
    print("Backward propagation debugging")
    print("=" * 50)

    for name, param in model.named_parameters():
        if param.grad is not None:
            grad_norm = param.grad.norm().item()
            if torch.isnan(param.grad).any():
                print(f"{name}: &#x26a0;&#xfe0f; Gradient NaN!")
            elif torch.isinf(param.grad).any():
                print(f"{name}: &#x26a0;&#xfe0f; Gradient Inf!")
            elif grad_norm > 10:
                print(f"{name}: &#x26a0;&#xfe0f; Gradient explosion! norm={grad_norm:.2f}")
            elif grad_norm < 1e-7:
                print(f"{name}: &#x26a0;&#xfe0f; Gradient vanishing! norm={grad_norm:.2e}")
            else:
                print(f"{name}: normal norm={grad_norm:.4f}")
    print("=" * 50)


# Usage example
set_seed(42)

model = nn.Sequential(
    nn.Linear(10, 20),
    nn.ReLU(),
    nn.Linear(20, 20),
    nn.ReLU(),
    nn.Linear(20, 5)
)

x = torch.randn(2, 10)
y = model(x)

debug_forward(model, x)

loss = y.sum()
loss.backward()

debug_backward(model)

8. Using PyTorch Profiler

PyTorch Profiler is the official performance analysis tool that can analyze the time consumption and memory usage of each part of the model.

8.1 Basic Usage of Profiler

Example

import torch
import torch.nn as nn
from torch.profiler import profile, ProfilerActivity, schedule

# Simple model
model = nn.Sequential(
    nn.Conv2d(3, 64, 3, padding=1),
    nn.ReLU(),
    nn.Conv2d(64, 64, 3, padding=1),
    nn.ReLU(),
    nn.AdaptiveAvgPool2d(1),
    nn.Flatten(),
    nn.Linear(64, 10)
).cuda()

optimizer = torch.optim.Adam(model.parameters())
criterion = nn.CrossEntropyLoss()

# Use the profiler
with profile(
    activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA],
    schedule=schedule(wait=1, warmup=1, active=3, repeat=1),
    on_trace_ready=torch.profiler.tensorboard_trace_handler('./logs'),
    record_shapes=True,
    profile_memory=True,
    with_stack=True
) as prof:
    for step in range(5):
        inputs = torch.randn(8, 3, 224, 224).cuda()
        labels = torch.randint(0, 10, (8,)).cuda()

        optimizer.zero_grad()
        outputs = model(inputs)
        loss = criterion(outputs, labels)
        loss.backward()
        optimizer.step()
        prof.step()

# Print results
print(prof.key_averages().table(sort_by="cuda_time_total", row_limit=10))

8.2 Visualizing Profiler Results

Example

# Export Chrome trace file
# The generated ./logs directory can be opened and viewed in the Chrome browser

# Or directly print various metrics
print("=" * 80)
print("CPU Time Top 10:")
print(prof.key_averages().table(sort_by="cpu_time_total", row_limit=10))

print("\n" + "=" * 80)
print("CUDA Time Top 10:")
print(prof.key_averages().table(sort_by="cuda_time_total", row_limit=10))

print("\n" + "=" * 80)
print("Memory Usage Top 10:")
print(prof.key_averages().table(sort_by="self_cuda_memory_usage", row_limit=10))
Other extensions