How to Debug a Deep Neural Network That Is Not Training Properly: A 10-Step Diagnostic Checklist
Systematically debug a deep neural network that is not training properly by verifying data quality, preprocessing, weight initialization, and learning rate, then monitor loss curves and gradient norms to isolate whether the issue stems from data, architecture, or optimization hyperparameters.
Training a deep neural network that is not training properly can stall an entire project, yet most failures stem from ten predictable categories. The scutan90/DeepLearning-500-questions repository provides a battle-tested diagnostic framework in ch12_网络搭建及训练/第十二章_网络搭建及训练.md, mapping specific line numbers to training pitfalls. This guide translates those findings into an actionable debugging workflow with concrete code examples.
1. Verify Data Quality and Preprocessing
Check for Corrupt Samples and Class Balance
According to ch12_网络搭建及训练/第十二章_网络搭建及训练.md (lines 80-84), you must first confirm your dataset contains no obvious corrupt or "dirty" samples and that class distribution is roughly balanced. Bad data can produce NaNs or prevent loss from decreasing.
Validate Preprocessing Pipelines
Incorrect preprocessing, such as forgetting to normalize inputs, leads to exploding gradients. Verify that mean subtraction, variance scaling, and augmentation logic are correctly implemented as emphasized in the section "合适的预处理方法" (lines 85-88).
2. Inspect Weight Initialization and Architecture
Weight Initialization Strategy
Zero initialization causes symmetry that prevents learning entirely. The source code recommends using Xavier or He initialization for ReLU-based networks, documented in "网络的初始化" (lines 89-92). Ensure weights are not all-zero before training begins.
Small-Scale Sanity Checks
Before full training, run a few epochs on a tiny subset (e.g., 100 samples) to confirm the pipeline works. This "小规模数据试练" approach (lines 93-98) quickly catches bugs without wasting compute resources.
3. Optimize Training Dynamics
Learning Rate Calibration
The learning rate is the most sensitive hyper-parameter. If set too large, loss explodes; too small, and no progress occurs. The repository details step decay and cosine annealing strategies in "设置合理Learning Rate" (lines 100-104).
Loss Function Verification
Ensure you are using the correct loss for your task—cross-entropy for classification, MSE/MAE for regression. Mismatched loss functions lead to poor gradients, as noted in the "损失函数" section (lines 106-115).
Gradient Monitoring
Loss should steadily decrease while gradients remain non-zero and finite. Enable gradient clipping to rescue exploding gradients, and log gradient norms to identify vanishing gradient issues, following the loss monitoring discussion (lines 78-84).
4. Validate Execution and Hardware
Graph and Session Integrity (TensorFlow 1.x)
For static graphs, verify all variables are initialized via sess.run(tf.global_variables_initializer()) and that graph nodes are correctly connected. Uninitialized tensors raise runtime errors that stall training, as shown in the TensorFlow code block (lines 44-55).
Hardware Usage Checks
Confirm GPU memory isn't exhausted and that batch sizes fit within available VRAM. Out-of-memory errors can silently abort training loops or cause erratic behavior.
Debugging Tools and Visualization
Leverage TensorBoard for TensorFlow or torch.utils.tensorboard for PyTorch to visualize loss curves, gradient histograms, and weight distributions. Enable PyTorch's anomaly detection with torch.autograd.set_detect_anomaly(True) to catch NaN propagation immediately. These practices align with the repository's emphasis on "完善的文档" (lines 75-78) and comprehensive logging.
Code Examples for Debugging
Below are minimal, framework-specific snippets that demonstrate the core debugging actions from the scutan90/DeepLearning-500-questions analysis.
TensorFlow 1.x Debugging Pattern
import tensorflow as tf
import numpy as np
# 1️⃣ Define a simple model
x = tf.placeholder(tf.float32, [None, 28, 28, 1], name='input')
y = tf.placeholder(tf.int64, [None], name='label')
conv = tf.layers.conv2d(x, 32, 3, activation=tf.nn.relu, name='conv')
flat = tf.layers.flatten(conv, name='flatten')
logits = tf.layers.dense(flat, 10, name='fc')
loss = tf.reduce_mean(tf.nn.sparse_softmax_cross_entropy_with_logits(
logits=logits, labels=y), name='loss')
# 2️⃣ Gradient sanity check
grads = tf.gradients(loss, tf.trainable_variables())
grad_norms = [tf.norm(g) for g in grads if g is not None]
# 3️⃣ Optimizer with learning-rate schedule
global_step = tf.Variable(0, trainable=False)
lr = tf.train.exponential_decay(0.01, global_step, 1000, 0.96, staircase=True)
optimizer = tf.train.AdamOptimizer(lr).minimize(loss, global_step=global_step)
# 4️⃣ TensorBoard summaries
tf.summary.scalar('loss', loss)
for i, g in enumerate(grad_norms):
tf.summary.scalar(f'grad_norm_{i}', g)
merged = tf.summary.merge_all()
with tf.Session() as sess:
sess.run(tf.global_variables_initializer())
writer = tf.summary.FileWriter('./logs', sess.graph)
for epoch in range(5):
# tiny sanity-check on a minibatch
batch_x, batch_y = get_small_batch()
_, step, summary = sess.run([optimizer, global_step, merged],
feed_dict={x: batch_x, y: batch_y})
writer.add_summary(summary, step)
print(f'Epoch {epoch}, step {step}')
Key debugging hooks: gradient norms, learning-rate schedule, TensorBoard logging, and explicit initialization via tf.global_variables_initializer() as referenced in lines 44-55 of the source chapter.
PyTorch Debugging Pattern
import torch, torch.nn as nn, torch.optim as optim
from torch.utils.tensorboard import SummaryWriter
# 1️⃣ Simple CNN
class Net(nn.Module):
def __init__(self):
super().__init__()
self.conv = nn.Conv2d(1, 32, 3)
self.fc = nn.Linear(26*26*32, 10)
def forward(self, x):
x = torch.relu(self.conv(x))
x = x.view(x.size(0), -1)
return self.fc(x)
model = Net()
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=0.01)
# 2️⃣ Enable anomaly detection (gradient NaNs)
torch.autograd.set_detect_anomaly(True)
# 3️⃣ TensorBoard
writer = SummaryWriter(log_dir='./logs')
for epoch in range(5):
inputs, targets = get_small_batch() # tiny sanity-check
outputs = model(inputs)
loss = criterion(outputs, targets)
optimizer.zero_grad()
loss.backward()
# 4️⃣ Log gradient norms
for name, param in model.named_parameters():
if param.grad is not None:
writer.add_scalar(f'grad_norm/{name}', param.grad.norm(), epoch)
optimizer.step()
writer.add_scalar('loss', loss.item(), epoch)
print(f'Epoch {epoch}, loss {loss.item():.4f}')
Key debugging hooks: torch.autograd.set_detect_anomaly, per-parameter gradient norm logging, and small-batch validation before full training.
Key Files in the DeepLearning-500-questions Repository
The following files contain the specific debugging recommendations and line references cited throughout this guide:
| File | Debugging Relevance |
|---|---|
ch12_网络搭建及训练/第十二章_网络搭建及训练.md |
Central checklist covering data quality (lines 80-84), preprocessing (lines 85-88), initialization (lines 89-92), small-scale testing (lines 93-98), learning rate (lines 100-104), and loss functions (lines 106-115). |
README.md |
Overview of the entire question set with navigation links to training-related chapters. |
_sidebar.md |
Navigation file to quickly locate the network training chapter. |
Summary
To effectively debug a deep neural network that is not training properly, follow this prioritized workflow derived from the scutan90/DeepLearning-500-questions analysis:
- Start with data: Verify no corrupt samples exist and that preprocessing (normalization, augmentation) is correctly implemented according to lines 85-88 of the training chapter.
- Validate architecture: Use Xavier/He initialization instead of zeros, and run a small-scale sanity check on 100 samples before full training (lines 89-98).
- Tune optimization: Confirm the learning rate follows a reasonable schedule and that the loss function matches your task type (cross-entropy for classification, MSE for regression) as specified in lines 100-115.
- Monitor continuously: Log loss curves and gradient norms to TensorBoard, enable anomaly detection in PyTorch, and verify all variables are initialized in TensorFlow 1.x sessions.
Frequently Asked Questions
Why does my neural network loss stay flat and not decrease?
A flat loss curve usually indicates a learning rate that is too small, frozen or zero-initialized weights, or a disconnected computation graph. Check that you are using proper Xavier or He initialization instead of zeros, and verify the learning rate is at least 1e-4 for Adam or 1e-2 for SGD. According to the DeepLearning-500-questions repository, running a small-scale sanity check on 100 samples can quickly reveal if the model is capable of overfitting a tiny batch, which isolates whether the issue is model capacity or the data pipeline.
How do I detect vanishing or exploding gradients during training?
Monitor gradient norms after every backward pass. In PyTorch, log param.grad.norm() for each parameter to TensorBoard; in TensorFlow 1.x, compute tf.norm(g) for each gradient tensor. If norms are consistently near zero, you have vanishing gradients—try using ReLU activations or residual connections. If norms exceed 1e3 or become NaN, you have exploding gradients—apply gradient clipping or verify your input normalization. Enable torch.autograd.set_detect_anomaly(True) in PyTorch to catch the exact operation causing NaN propagation immediately.
What is the fastest way to validate my training pipeline before a full run?
Perform a small-scale sanity check by training for 5-10 epochs on a tiny subset of 100-200 samples. According to ch12_网络搭建及训练/第十二章_网络搭建及训练.md (lines 93-98), if the model cannot overfit this small batch to near-zero loss, the bug likely exists in the architecture, initialization, or loss function implementation rather than the data itself. This approach quickly catches logic errors without wasting compute resources on full-scale runs.
Should I use TensorFlow 1.x or PyTorch debugging tools for static vs. dynamic graphs?
For TensorFlow 1.x static graphs, explicitly verify that sess.run(tf.global_variables_initializer()) executes before training, and use tf.summary to log gradient norms and loss values to TensorBoard. Uninitialized tensors raise runtime errors that stall training, as shown in the repository's examples (lines 44-55). For PyTorch dynamic graphs, use torch.autograd.set_detect_anomaly(True) to catch NaN propagation during the backward pass, and log per-parameter gradient norms directly since the graph is rebuilt each iteration. Both approaches require monitoring loss curves and gradient magnitudes, but PyTorch offers more granular per-step debugging while TensorFlow 1.x requires careful session and graph management to avoid uninitialized variable errors.
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 →