Fundamentals — vision and real data¶
Apply the core workflow to images: trace shapes through a CNN, check the files before training, then decide whether validation evidence supports the model.
Code key: “Runnable toy” includes imports and inputs. Other snippets are excerpts: reuse torch, nn, and the model, loader, tokenizer or helper named in the section. Projects contain the complete runnable scripts.
What this guide connects¶
image files → validated samples → augmented training batches
↓
convolutional feature maps
↓
[batch, classes] logits
↓
validation curves + error inspection
↓
reproducible checkpoint and context
Read Fundamentals — core workflow first if Dataset, logits, loss, or the training update order are unfamiliar.
Study target: after each section, name the input, output, and reason for the operation. A CNN is the same training loop as before; the important new questions are spatial shape, file quality, and generalization.
Convolution and feature maps¶
A dense layer sees one long list of pixels. A convolution keeps the spatial layout and applies the same small kernel at many positions. Reusing weights makes it efficient and allows one learned pattern to be recognized in different locations.

The pictured filter is hand-designed to make its response visible. A trained CNN learns its kernel values through backpropagation.
Channels¶
Convolutions turn input channels into learned feature channels while preserving each sample.
Code and details
Each of the 32 learned filters spans all three RGB input channels. The result is 32 feature maps—not 32 colours. Early filters may respond to edges, colour transitions, or texture; deeper layers combine those maps into patterns useful for the task.
Try one layer with a small batch:
from torch import nn
images = torch.zeros(4, 3, 32, 32) # four RGB images
conv = nn.Conv2d(3, 8, 3, padding=1) # learn eight 3 × 3 filters
maps = conv(images)
print(maps.shape) # [4, 8, 32, 32]
The batch count stays 4; the filter count sets output channels to 8. Padding keeps 32 × 32 here. Check yourself: what changes if padding=0? The spatial output becomes 30 × 30.
Kernel size, stride, and padding¶
- Kernel size controls the local window inspected at once.
- Stride controls how far the kernel moves between positions.
- Padding adds a border so edge pixels participate and output size can be controlled.
- Dilation spaces kernel elements apart to expand the receptive field.
For one spatial dimension:
A 3 × 3 convolution with padding 1, stride 1, and dilation 1 preserves height and width.
CNN architecture and shapes¶

The values in the diagram are an illustrative shape trace. Always inspect the real input and real model.
[batch, 3, 32, 32]
→ convolution: channels grow
→ pooling or stride: width and height shrink
→ classifier: [batch, classes]
A reusable block¶
A convolution learns local features; an activation adds nonlinearity; pooling reduces spatial size.
Code and details
class CNNBlock(nn.Module):
def __init__(self, in_channels, out_channels):
super().__init__()
self.block = nn.Sequential(
nn.Conv2d(in_channels, out_channels, 3, padding=1),
nn.BatchNorm2d(out_channels),
nn.ReLU(),
nn.MaxPool2d(2),
)
def forward(self, x):
return self.block(x)
Convolution learns local features, ReLU adds nonlinearity, normalization stabilizes intermediate distributions, and pooling reduces spatial resolution.
Trace this block on [4, 3, 32, 32]: convolution gives [4, 8, 32, 32]; batch normalization and ReLU keep that shape; MaxPool2d(2) gives [4, 8, 16, 16]. Pooling keeps the strongest value in each 2 × 2 window. It reduces detail and compute, so excessive pooling can discard information. Batch normalization tracks statistics during training and uses stored statistics in evaluation; that is one reason model.eval() matters.
Safer classifier boundaries¶
Hard-coded flatten sizes break when input resolution or feature blocks change. Prefer adaptive pooling when spatial position no longer needs to be preserved:
self.features = nn.Sequential(
CNNBlock(3, 32),
CNNBlock(32, 64),
nn.AdaptiveAvgPool2d((1, 1)),
)
self.classifier = nn.Linear(64, classes)
def forward(self, x):
x = self.features(x)
x = torch.flatten(x, 1)
return self.classifier(x)
When a fixed feature size is intentional, inspect it with a dummy input rather than calculating it mentally:
Trace the first unexpected shape¶
def show_shape(name):
def hook(module, inputs, output):
print(name, tuple(output.shape))
return hook
handle = model.features.register_forward_hook(show_shape("features"))
# run one batch, then remove the temporary hook
handle.remove()
Most shape failures are easier to understand at the first incorrect boundary than at the final linear-layer error.
Inspect the model and its activations¶
Inspect registered modules and intermediate output shapes to locate the first unexpected transformation.
Code and details
| Method | What you see |
|---|---|
print(model) |
layer hierarchy and declared dimensions |
named_children() |
direct child modules |
named_modules() |
nested modules recursively |
named_parameters() |
registered weight/bias names and tensors |
buffers() |
non-parameter state, such as BatchNorm running statistics |
state_dict() |
persisted parameters and buffers |
for name, parameter in model.named_parameters():
print(name, tuple(parameter.shape), parameter.requires_grad)
total = sum(parameter.numel() for parameter in model.parameters())
A linear layer with F inputs and K outputs has F*K + K parameters with bias. An ordinary convolution has out_channels * in_channels * kernel_height * kernel_width + out_channels with bias and groups=1. Pooling and ReLU have no trainable parameters. Parameter count is not the same as compute cost.
Inspect finite values, mean, standard deviation, minimum, maximum and zero fraction after a suspicious block. Mostly zero ReLU outputs or collapsing variance can guide an investigation; they do not prove a bug alone. Hooks are temporary diagnostics: remove their handles and avoid storing tensors with live computation graphs. Start by overfitting one tiny batch before changing the architecture.
Next: compare a dense model with a CNN in EMNIST, then inspect a modular regularized model in Nature CNN.
Reliable image data¶
fit parametersrandom label-preserving transformschoose settings and checkpointdeterministic inputs; no gradient updatesassess selected model oncekeep outside tuningIllustration: shapes and operations, not measured model performance.
A larger architecture cannot repair inconsistent labels, leakage, corrupt files, or transforms applied to the wrong split.
Establish the data contract before training¶
- Define stable class names and a class-to-index mapping.
- Verify file existence, extension, readability, colour mode, and expected input shape.
- Record rejected files with their path and reason.
- Split with a fixed seed before applying random augmentation.
- Inspect class counts in each split.
- Keep validation preprocessing deterministic.
- Split related entities together so near-duplicates cannot leak across sets.
from PIL import Image
def open_rgb(path):
with Image.open(path) as image:
return image.convert("RGB")
Handle failures deliberately¶
Do not recursively request another sample from __getitem__ without a strict limit; a cluster of corrupt files can create infinite recursion. Prefer validating the index before training. When runtime filtering is required, return a controlled sentinel and use a custom collate_fn.
def collate_valid(batch):
valid = [sample for sample in batch if sample is not None]
if not valid:
raise RuntimeError("Batch contains no valid samples")
return torch.utils.data.default_collate(valid)
Make splits reproducible¶
generator = torch.Generator().manual_seed(42)
train_set, validation_set = random_split(
dataset,
[train_size, validation_size],
generator=generator,
)
Record the split seed and per-class distribution. For multiple images of the same person, object, location, or capture burst, split by entity rather than individual image. random_split by itself does not prevent that kind of leakage.
Separate transforms for shared split indices¶
Subset and random_split retain references to their underlying dataset. Assigning a new transform through a shared parent can accidentally augment validation too. With identical folder contents and ordering, use separate dataset objects:
Code and details
from torchvision.datasets import ImageFolder
from torch.utils.data import Subset
train_base = ImageFolder("images", transform=train_transform)
val_base = ImageFolder("images", transform=validation_transform)
assert train_base.samples == val_base.samples
assert train_base.class_to_idx == val_base.class_to_idx
order = torch.randperm(len(train_base), generator=torch.Generator().manual_seed(17))
cut = int(0.8 * len(order))
train_set = Subset(train_base, order[:cut].tolist())
val_set = Subset(val_base, order[cut:].tolist())
Here the transform objects were defined earlier. Use group-aware indices instead of the random permutation when samples are related. A custom per-subset wrapper is another solution; it must receive raw samples so it does not transform an already transformed tensor twice.
Check yourself: if two crops of the same photo land in different splits, validation can reward recognition of that photo rather than a pattern that works on new photos. Group them before splitting.
Monitor data, not only the model¶
- Samples and rejected files per class.
- Pixel range after transforms.
- Label dtype and minimum/maximum values.
- Batch loading time and accelerator idle time.
- Class mapping stored with the run.
- Example transformed images from both training and validation.
Code and details
Decode images in a context manager and convert to RGB explicitly. Image.verify() checks file integrity but requires reopening before decoding/transforming; a successful check does not validate the label. For one-based label files, translate IDs consistently. For filenames starting at 1 and sample indices starting at 0, keep the offset explicit. Do not repeatedly load an entire table inside __getitem__ when only one row is needed.
The robust image-pipeline project downloads public data, builds a folder dataset, adds one known corrupt file, records the validation decision, and trains only on valid examples.
Generalization and regularization¶
Generalization asks whether patterns learned from training examples remain useful on unseen data from the intended environment.
Read learning curves¶
Compare training and validation curves to see whether learning generalizes or begins to overfit.
Code and details
| Observation | Likely issue | First checks |
|---|---|---|
| Training and validation both poor | underfitting or broken pipeline | labels, loss, learning rate, capacity |
| Training improves; validation worsens | overfitting | split quality, augmentation, regularization |
| Loss becomes NaN | numerical instability | invalid inputs, learning rate, gradients |
| High total accuracy; one class fails | imbalance or shortcut learning | per-class recall and confusion matrix |
Four gaps to inspect¶
- Training–validation gap: has the model fitted training-specific detail?
- Validation–deployment gap: does validation represent real inputs?
- Aggregate–class gap: does one headline metric hide weak classes?
- Accuracy–cost gap: is an improvement worth the memory, latency, and complexity?
A model can improve on a standard held-out set and still fail on an external source. Report both domains instead of treating one as the universal truth.
Regularization tools¶
Data augmentation creates realistic label-preserving variation. An upside-down vehicle or severely distorted character may not preserve the label.
Code and details
Weight decay discourages unnecessarily large weights:
Dropout randomly removes activations during training:
It is active in model.train() and disabled in model.eval().
Early stopping ends training after sufficient non-improvement. Saving and restoring the best validation checkpoint is a separate action; stopping alone does not restore those weights.
Change one hypothesis at a time. Use the same split, seed, metrics, and evaluation procedure so an apparent gain does not come from changing the test.
If training loss falls while validation loss rises, first inspect the split and the mistakes. Then try one change, such as realistic augmentation, weight decay, or a smaller model. Dropout is a training-time source of noise; it is not a repair for incorrect labels or leakage.
Saving and restoring¶
Save model parameters¶
Restore them into the same architecture:
model = Classifier(classes=len(class_names))
state = torch.load("model.pth", map_location=device, weights_only=True)
model.load_state_dict(state)
model.to(device)
model.eval()
Save a training checkpoint¶
A resumable checkpoint includes optimizer state and run context as well as model weights.
Code and details
torch.save({
"epoch": epoch,
"model_state": model.state_dict(),
"optimizer_state": optimizer.state_dict(),
"validation_loss": validation_loss,
"class_names": class_names,
"input_shape": input_shape,
"normalization": normalization,
}, "checkpoint.pth")
A usable model is more than its tensors. Retain architecture code, preprocessing, class order, input shape, selected metric, split identity, seed, and package versions. Verify the restored model on a known input.
To resume training consistently, retain scheduler state, mixed-precision scaler state when used, and the next epoch as well. For a best model held in memory, use copy.deepcopy(model.state_dict()) or save immediately: a dictionary of live tensor references can follow later updates. Random state and sampler state matter when exact continuation is required.
Never load an untrusted pickle-based checkpoint. Prefer weight-only loading when the saved format and installed PyTorch version support it.
Vision-workflow checklist¶
Before trusting an image model, confirm:
- Files, labels, classes, and splits were validated before training.
- Training and validation transforms differ only where intended.
- Every CNN boundary has a known shape.
- The classifier produces
[batch, classes]logits. - Metrics include error structure, not only one average.
- Training and validation curves are interpreted together.
- The selected checkpoint can be restored with its preprocessing and class order.
Use the project gallery to move from these patterns to complete, runnable examples, or open the reference when debugging a specific run.
For the next layer of the same workflow, study noise and augmentation and pretrained vision models. Both rely on the shape, split and evaluation checks above.